312 lines
13 KiB
Go
312 lines
13 KiB
Go
package postgres
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"fmt"
|
|
|
|
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/administration"
|
|
)
|
|
|
|
const accountColumns = `id, phone, display_name, role, organization_id, status, password_hash, password_salt, failed_login_count, locked_until, session_version, last_login_at, legacy_subject, created_at, updated_at`
|
|
const organizationColumns = `id, name, status, archive_owner_id, created_at, updated_at`
|
|
const ListAdministrationAccountsSQL = `SELECT ` + accountColumns + ` FROM public.platform_users WHERE ($1::text = '' OR organization_id = $1::text) AND ($2::text = '' OR role = $2::text) AND ($3::boolean OR status = 'active') ORDER BY created_at DESC`
|
|
const GetAdministrationAccountSQL = `SELECT ` + accountColumns + ` FROM public.platform_users WHERE id = $1::text`
|
|
const CreateAdministrationAccountSQL = `INSERT INTO public.platform_users (` + accountColumns + `) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15) RETURNING ` + accountColumns
|
|
const ApplyAdministrationAccountUpdateSQL = `UPDATE public.platform_users SET
|
|
display_name = CASE WHEN $2::boolean THEN $3::text ELSE display_name END,
|
|
role = CASE WHEN $4::boolean THEN $5::text ELSE role END,
|
|
organization_id = CASE WHEN $6::boolean THEN $7::text ELSE organization_id END,
|
|
status = CASE WHEN $8::boolean THEN $9::text ELSE status END,
|
|
password_hash = CASE WHEN $10::boolean THEN $11::text ELSE password_hash END,
|
|
password_salt = CASE WHEN $10::boolean THEN $12::text ELSE password_salt END,
|
|
failed_login_count = CASE WHEN $13::boolean THEN 0 ELSE failed_login_count END,
|
|
locked_until = CASE WHEN $13::boolean THEN NULL ELSE locked_until END,
|
|
session_version = session_version + CASE WHEN $14::boolean THEN 1 ELSE 0 END,
|
|
updated_at = $15::timestamptz
|
|
WHERE id = $1::text
|
|
AND ($16::text = 'super_admin' OR (role = 'user' AND organization_id = $17::text))
|
|
RETURNING ` + accountColumns
|
|
const ListAdministrationOrganizationsSQL = `SELECT ` + organizationColumns + ` FROM public.platform_organizations WHERE ($1::boolean OR status = 'active') ORDER BY created_at ASC`
|
|
const GetAdministrationOrganizationSQL = `SELECT ` + organizationColumns + ` FROM public.platform_organizations WHERE id = $1::text`
|
|
const CreateAdministrationOrganizationSQL = `INSERT INTO public.platform_organizations (` + organizationColumns + `) VALUES ($1,$2,$3,$4,$5,$6) RETURNING ` + organizationColumns
|
|
const UpdateAdministrationOrganizationSQL = `UPDATE public.platform_organizations SET name=$2, status=$3, archive_owner_id=$4, created_at=$5, updated_at=$6 WHERE id=$1::text RETURNING ` + organizationColumns
|
|
const CountAdministrationOrganizationMembersSQL = `SELECT count(*) FROM public.platform_users WHERE organization_id = $1::text`
|
|
const DeleteAdministrationOrganizationSQL = `DELETE FROM public.platform_organizations WHERE id = $1::text RETURNING id`
|
|
const ArchiveAssetsSQL = `UPDATE public.assets SET owner_id = $2::text WHERE owner_id = $1::text`
|
|
const ArchiveGenerationJobsSQL = `UPDATE public.generation_jobs SET owner_id = $2::text WHERE owner_id = $1::text`
|
|
const ArchiveProjectsSQL = `UPDATE public.projects SET owner_id = $2::text WHERE owner_id = $1::text`
|
|
const ArchiveImageTemplatesSQL = `UPDATE public.image_templates SET owner_id = $2::text WHERE owner_id = $1::text`
|
|
const ArchiveUsageEventsSQL = `UPDATE public.usage_events SET owner_id = $2::text WHERE owner_id = $1::text`
|
|
const DeleteAdministrationAccountSQL = `DELETE FROM public.platform_users WHERE id = $1::text`
|
|
|
|
func (db *Database) administrationQuerier() (Querier, error) {
|
|
if db.config.Backend != BackendPostgres || db.querier == nil {
|
|
return nil, fmt.Errorf("PostgreSQL is unavailable when ZHINIAN_DATA_BACKEND=%s", db.config.Backend)
|
|
}
|
|
return db.querier, nil
|
|
}
|
|
func (db *Database) ListAccounts(ctx context.Context, f administration.AccountFilters) ([]administration.Account, error) {
|
|
q, e := db.administrationQuerier()
|
|
if e != nil {
|
|
return nil, e
|
|
}
|
|
rows, e := q.Query(ctx, ListAdministrationAccountsSQL, f.OrganizationID, f.Role, f.IncludeDisabled)
|
|
if e != nil {
|
|
return nil, fmt.Errorf("list administration accounts: %w", e)
|
|
}
|
|
defer rows.Close()
|
|
var out []administration.Account
|
|
for rows.Next() {
|
|
a, e := scanAdministrationAccount(rows)
|
|
if e != nil {
|
|
return nil, e
|
|
}
|
|
out = append(out, a)
|
|
}
|
|
return out, rows.Err()
|
|
}
|
|
func (db *Database) GetAccount(ctx context.Context, id string) (administration.Account, bool, error) {
|
|
q, e := db.administrationQuerier()
|
|
if e != nil {
|
|
return administration.Account{}, false, e
|
|
}
|
|
rows, e := q.Query(ctx, GetAdministrationAccountSQL, id)
|
|
if e != nil {
|
|
return administration.Account{}, false, fmt.Errorf("get administration account: %w", e)
|
|
}
|
|
defer rows.Close()
|
|
if !rows.Next() {
|
|
return administration.Account{}, false, rows.Err()
|
|
}
|
|
a, e := scanAdministrationAccount(rows)
|
|
return a, e == nil, e
|
|
}
|
|
func (db *Database) CreateAccount(ctx context.Context, a administration.Account) (administration.Account, error) {
|
|
return db.writeAccount(ctx, CreateAdministrationAccountSQL, a)
|
|
}
|
|
func (db *Database) UpdateAccount(context.Context, administration.Account) (administration.Account, error) {
|
|
return administration.Account{}, fmt.Errorf("full-row account updates are unsafe; use ApplyAccountUpdate")
|
|
}
|
|
func (db *Database) ApplyAccountUpdate(ctx context.Context, id string, update administration.AccountUpdate) (administration.Account, error) {
|
|
q, err := db.administrationQuerier()
|
|
if err != nil {
|
|
return administration.Account{}, err
|
|
}
|
|
displayName, hasDisplayName := optionalUpdateValue(update.DisplayName)
|
|
role, hasRole := optionalUpdateValue(update.Role)
|
|
organizationID, hasOrganizationID := optionalUpdateValue(update.OrganizationID)
|
|
status, hasStatus := optionalUpdateValue(update.Status)
|
|
passwordHash, passwordSalt, hasPassword := "", "", update.PasswordHash != nil
|
|
if hasPassword {
|
|
passwordHash, passwordSalt = update.PasswordHash.Hash, update.PasswordHash.Salt
|
|
}
|
|
rows, err := q.Query(ctx, ApplyAdministrationAccountUpdateSQL,
|
|
id,
|
|
hasDisplayName, displayName,
|
|
hasRole, role,
|
|
hasOrganizationID, optionalDatabaseText(organizationID),
|
|
hasStatus, status,
|
|
hasPassword, passwordHash, passwordSalt,
|
|
update.ClearLoginLock, update.IncrementSessionVersion,
|
|
update.UpdatedAt,
|
|
update.Actor.Role, update.Actor.OrganizationID,
|
|
)
|
|
if err != nil {
|
|
return administration.Account{}, administrationWriteError(err)
|
|
}
|
|
defer rows.Close()
|
|
if !rows.Next() {
|
|
if err := rows.Err(); err != nil {
|
|
return administration.Account{}, err
|
|
}
|
|
return administration.Account{}, &administration.Error{Kind: administration.ErrorNotFound, Message: "账号不存在或无权操作。"}
|
|
}
|
|
return scanAdministrationAccount(rows)
|
|
}
|
|
|
|
func optionalUpdateValue[T any](value *T) (T, bool) {
|
|
if value == nil {
|
|
var zero T
|
|
return zero, false
|
|
}
|
|
return *value, true
|
|
}
|
|
func (db *Database) writeAccount(ctx context.Context, statement string, a administration.Account) (administration.Account, error) {
|
|
q, e := db.administrationQuerier()
|
|
if e != nil {
|
|
return administration.Account{}, e
|
|
}
|
|
rows, e := q.Query(ctx, statement, accountArgs(a)...)
|
|
if e != nil {
|
|
return administration.Account{}, administrationWriteError(e)
|
|
}
|
|
defer rows.Close()
|
|
if !rows.Next() {
|
|
return administration.Account{}, &administration.Error{Kind: administration.ErrorNotFound, Message: "账号不存在。"}
|
|
}
|
|
return scanAdministrationAccount(rows)
|
|
}
|
|
func accountArgs(a administration.Account) []any {
|
|
return []any{a.ID, a.Phone, a.DisplayName, a.Role, optionalDatabaseText(a.OrganizationID), a.Status, a.PasswordHash, a.PasswordSalt, a.FailedLoginCount, a.LockedUntil, a.SessionVersion, a.LastLoginAt, optionalDatabaseText(a.LegacySubject), a.CreatedAt, a.UpdatedAt}
|
|
}
|
|
func scanAdministrationAccount(rows Rows) (administration.Account, error) {
|
|
var a administration.Account
|
|
var org, legacy sql.NullString
|
|
var locked, last sql.NullTime
|
|
if e := rows.Scan(&a.ID, &a.Phone, &a.DisplayName, &a.Role, &org, &a.Status, &a.PasswordHash, &a.PasswordSalt, &a.FailedLoginCount, &locked, &a.SessionVersion, &last, &legacy, &a.CreatedAt, &a.UpdatedAt); e != nil {
|
|
return a, fmt.Errorf("scan administration account: %w", e)
|
|
}
|
|
if org.Valid {
|
|
a.OrganizationID = org.String
|
|
}
|
|
if legacy.Valid {
|
|
a.LegacySubject = legacy.String
|
|
}
|
|
if locked.Valid {
|
|
a.LockedUntil = &locked.Time
|
|
}
|
|
if last.Valid {
|
|
a.LastLoginAt = &last.Time
|
|
}
|
|
return a, nil
|
|
}
|
|
|
|
func (db *Database) ListOrganizations(ctx context.Context, include bool) ([]administration.Organization, error) {
|
|
q, e := db.administrationQuerier()
|
|
if e != nil {
|
|
return nil, e
|
|
}
|
|
rows, e := q.Query(ctx, ListAdministrationOrganizationsSQL, include)
|
|
if e != nil {
|
|
return nil, fmt.Errorf("list administration organizations: %w", e)
|
|
}
|
|
defer rows.Close()
|
|
var out []administration.Organization
|
|
for rows.Next() {
|
|
o, e := scanAdministrationOrganization(rows)
|
|
if e != nil {
|
|
return nil, e
|
|
}
|
|
out = append(out, o)
|
|
}
|
|
return out, rows.Err()
|
|
}
|
|
func (db *Database) GetOrganization(ctx context.Context, id string) (administration.Organization, bool, error) {
|
|
q, e := db.administrationQuerier()
|
|
if e != nil {
|
|
return administration.Organization{}, false, e
|
|
}
|
|
rows, e := q.Query(ctx, GetAdministrationOrganizationSQL, id)
|
|
if e != nil {
|
|
return administration.Organization{}, false, e
|
|
}
|
|
defer rows.Close()
|
|
if !rows.Next() {
|
|
return administration.Organization{}, false, rows.Err()
|
|
}
|
|
o, e := scanAdministrationOrganization(rows)
|
|
return o, e == nil, e
|
|
}
|
|
func (db *Database) CreateOrganization(ctx context.Context, o administration.Organization) (administration.Organization, error) {
|
|
return db.writeOrganization(ctx, CreateAdministrationOrganizationSQL, o)
|
|
}
|
|
func (db *Database) UpdateOrganization(ctx context.Context, o administration.Organization) (administration.Organization, error) {
|
|
return db.writeOrganization(ctx, UpdateAdministrationOrganizationSQL, o)
|
|
}
|
|
func (db *Database) writeOrganization(ctx context.Context, statement string, o administration.Organization) (administration.Organization, error) {
|
|
q, e := db.administrationQuerier()
|
|
if e != nil {
|
|
return administration.Organization{}, e
|
|
}
|
|
rows, e := q.Query(ctx, statement, o.ID, o.Name, o.Status, o.ArchiveOwnerID, o.CreatedAt, o.UpdatedAt)
|
|
if e != nil {
|
|
return administration.Organization{}, administrationWriteError(e)
|
|
}
|
|
defer rows.Close()
|
|
if !rows.Next() {
|
|
return administration.Organization{}, &administration.Error{Kind: administration.ErrorNotFound, Message: "组织不存在。"}
|
|
}
|
|
return scanAdministrationOrganization(rows)
|
|
}
|
|
func scanAdministrationOrganization(rows Rows) (administration.Organization, error) {
|
|
var o administration.Organization
|
|
if e := rows.Scan(&o.ID, &o.Name, &o.Status, &o.ArchiveOwnerID, &o.CreatedAt, &o.UpdatedAt); e != nil {
|
|
return o, fmt.Errorf("scan administration organization: %w", e)
|
|
}
|
|
return o, nil
|
|
}
|
|
func (db *Database) CountOrganizationMembers(ctx context.Context, id string) (int, error) {
|
|
q, e := db.administrationQuerier()
|
|
if e != nil {
|
|
return 0, e
|
|
}
|
|
rows, e := q.Query(ctx, CountAdministrationOrganizationMembersSQL, id)
|
|
if e != nil {
|
|
return 0, e
|
|
}
|
|
defer rows.Close()
|
|
if !rows.Next() {
|
|
return 0, fmt.Errorf("member count returned no row")
|
|
}
|
|
var count int
|
|
e = rows.Scan(&count)
|
|
return count, e
|
|
}
|
|
func (db *Database) DeleteOrganization(ctx context.Context, id string) error {
|
|
q, e := db.administrationQuerier()
|
|
if e != nil {
|
|
return e
|
|
}
|
|
rows, e := q.Query(ctx, DeleteAdministrationOrganizationSQL, id)
|
|
if e != nil {
|
|
return administrationWriteError(e)
|
|
}
|
|
defer rows.Close()
|
|
if !rows.Next() {
|
|
if e := rows.Err(); e != nil {
|
|
return e
|
|
}
|
|
return &administration.Error{Kind: administration.ErrorNotFound, Message: "组织不存在。"}
|
|
}
|
|
var deletedID string
|
|
if e := rows.Scan(&deletedID); e != nil {
|
|
return fmt.Errorf("scan deleted organization: %w", e)
|
|
}
|
|
return rows.Err()
|
|
}
|
|
func (db *Database) DeleteAccount(ctx context.Context, id, archive string) error {
|
|
if db.config.Backend != BackendPostgres || db.transactions == nil {
|
|
return fmt.Errorf("PostgreSQL is unavailable when ZHINIAN_DATA_BACKEND=%s", db.config.Backend)
|
|
}
|
|
tx, e := db.transactions.Begin(ctx)
|
|
if e != nil {
|
|
return fmt.Errorf("begin delete account transaction: %w", e)
|
|
}
|
|
done := false
|
|
defer func() {
|
|
if !done {
|
|
_ = tx.Rollback(ctx)
|
|
}
|
|
}()
|
|
for _, statement := range []string{ArchiveAssetsSQL, ArchiveGenerationJobsSQL, ArchiveProjectsSQL, ArchiveImageTemplatesSQL, ArchiveUsageEventsSQL} {
|
|
if e = tx.Exec(ctx, statement, id, archive); e != nil {
|
|
return fmt.Errorf("archive account ownership: %w", e)
|
|
}
|
|
}
|
|
if e = tx.Exec(ctx, DeleteAdministrationAccountSQL, id); e != nil {
|
|
return fmt.Errorf("delete administration account: %w", e)
|
|
}
|
|
if e = tx.Commit(ctx); e != nil {
|
|
return fmt.Errorf("commit delete account transaction: %w", e)
|
|
}
|
|
done = true
|
|
return nil
|
|
}
|
|
func administrationWriteError(err error) error {
|
|
if sqlState(err) == "23505" {
|
|
return &administration.Error{Kind: administration.ErrorConflict, Message: "唯一字段已存在。", Err: err}
|
|
}
|
|
return err
|
|
}
|