Files
NianAIGC/backend/internal/postgres/administration.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
}