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 }