254 lines
10 KiB
Go
254 lines
10 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 UpdateAdministrationAccountSQL = `UPDATE public.platform_users SET phone=$2, display_name=$3, role=$4, organization_id=$5, status=$6, password_hash=$7, password_salt=$8, failed_login_count=$9, locked_until=$10, session_version=$11, last_login_at=$12, legacy_subject=$13, created_at=$14, updated_at=$15 WHERE id=$1::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 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(ctx context.Context, a administration.Account) (administration.Account, error) {
|
|
return db.writeAccount(ctx, UpdateAdministrationAccountSQL, a)
|
|
}
|
|
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} {
|
|
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
|
|
}
|