Files
NianAIGC/backend/internal/postgres/administration_test.go

236 lines
8.1 KiB
Go

package postgres
import (
"context"
"database/sql"
"errors"
"reflect"
"strings"
"testing"
"time"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/administration"
)
func TestApplyAccountUpdateUsesNarrowAtomicSQLAndDatabaseVersionIncrement(t *testing.T) {
now := time.Date(2026, 8, 13, 8, 0, 0, 0, time.UTC)
rows := &administrationRows{rows: [][]any{{
"user-1", "13800138000", "New Name", "user", "org-1", "active",
"new-password-hash", "new-password-salt", 3, now.Add(time.Hour), 9, now.Add(-time.Hour),
nil, now.Add(-24 * time.Hour), now,
}}}
querier := &administrationQuerier{rows: rows}
db := NewDatabase(Config{Backend: BackendPostgres}, querier)
name := "New Name"
got, err := db.ApplyAccountUpdate(context.Background(), "user-1", administration.AccountUpdate{
Actor: administration.Actor{ID: "super", Role: administration.RoleSuperAdmin},
DisplayName: &name,
UpdatedAt: now,
})
if err != nil {
t.Fatal(err)
}
if got.SessionVersion != 9 || got.PasswordHash != "new-password-hash" || got.FailedLoginCount != 3 {
t.Fatalf("updated account=%#v, want database-owned security fields preserved", got)
}
if querier.sql != ApplyAdministrationAccountUpdateSQL {
t.Fatalf("SQL=%q", querier.sql)
}
for _, forbidden := range []string{"phone=$", "created_at=$", "last_login_at=$", "legacy_subject=$"} {
if strings.Contains(querier.sql, forbidden) {
t.Fatalf("atomic PATCH must not overwrite unrelated column: SQL contains %q", forbidden)
}
}
for _, required := range []string{
"password_hash = CASE WHEN $10::boolean THEN $11::text ELSE password_hash END",
"failed_login_count = CASE WHEN $13::boolean THEN 0 ELSE failed_login_count END",
"session_version = session_version + CASE WHEN $14::boolean THEN 1 ELSE 0 END",
} {
if !strings.Contains(querier.sql, required) {
t.Fatalf("atomic PATCH SQL missing %q", required)
}
}
wantArgs := []any{
"user-1",
true, "New Name",
false, administration.Role(""),
false, nil,
false, administration.Status(""),
false, "", "",
false, false,
now,
administration.RoleSuperAdmin, "",
}
if !reflect.DeepEqual(querier.args, wantArgs) {
t.Fatalf("args=%#v want %#v", querier.args, wantArgs)
}
}
func TestApplyAccountUpdateAtomicallyResetsPasswordClearsLockAndIncrementsCurrentVersion(t *testing.T) {
now := time.Date(2026, 8, 13, 8, 0, 0, 0, time.UTC)
rows := &administrationRows{rows: [][]any{{
"user-1", "13800138000", "Name", "user", "org-1", "active",
"replacement-hash", "replacement-salt", 0, nil, 12, nil,
nil, now.Add(-24 * time.Hour), now,
}}}
querier := &administrationQuerier{rows: rows}
db := NewDatabase(Config{Backend: BackendPostgres}, querier)
got, err := db.ApplyAccountUpdate(context.Background(), "user-1", administration.AccountUpdate{
Actor: administration.Actor{Role: administration.RoleOrganizationAdmin, OrganizationID: "org-1"},
PasswordHash: &administration.PasswordHash{Hash: "replacement-hash", Salt: "replacement-salt"},
ClearLoginLock: true,
IncrementSessionVersion: true,
UpdatedAt: now,
})
if err != nil {
t.Fatal(err)
}
if got.SessionVersion != 12 || got.FailedLoginCount != 0 || got.LockedUntil != nil {
t.Fatalf("updated account=%#v", got)
}
if !reflect.DeepEqual(querier.args[9:14], []any{true, "replacement-hash", "replacement-salt", true, true}) {
t.Fatalf("security args=%#v", querier.args[9:14])
}
}
func TestAdministrationListAccountsUsesExplicitColumnsAndParameterizedFilters(t *testing.T) {
rows := &identityRows{}
querier := &identityQuerier{rows: rows}
db := NewDatabase(Config{Backend: BackendPostgres}, querier)
_, err := db.ListAccounts(context.Background(), administration.AccountFilters{OrganizationID: "org-1", Role: administration.RoleUser, IncludeDisabled: true})
if err != nil {
t.Fatal(err)
}
if querier.sql != ListAdministrationAccountsSQL {
t.Fatalf("SQL=%q", querier.sql)
}
if !reflect.DeepEqual(querier.args, []any{"org-1", administration.RoleUser, true}) {
t.Fatalf("args=%#v", querier.args)
}
}
func TestDeleteAccountArchivesAllOwnedDataAndIdentityInOneTransaction(t *testing.T) {
tx := &administrationTransaction{}
db := NewDatabase(Config{Backend: BackendPostgres}, &administrationPool{tx: tx})
if err := db.DeleteAccount(context.Background(), "user-1", "archive:org-1"); err != nil {
t.Fatal(err)
}
wantSQL := []string{ArchiveAssetsSQL, ArchiveGenerationJobsSQL, ArchiveProjectsSQL, ArchiveImageTemplatesSQL, ArchiveUsageEventsSQL, DeleteAdministrationAccountSQL}
if !reflect.DeepEqual(tx.sql, wantSQL) {
t.Fatalf("SQL sequence=%#v want %#v", tx.sql, wantSQL)
}
for _, args := range tx.args {
if !reflect.DeepEqual(args, []any{"user-1", "archive:org-1"}) && !reflect.DeepEqual(args, []any{"user-1"}) {
t.Fatalf("args=%#v", args)
}
}
if tx.commits != 1 || tx.rollbacks != 0 {
t.Fatalf("commits=%d rollbacks=%d", tx.commits, tx.rollbacks)
}
}
func TestDeleteAccountRollsBackArchiveTransactionFailure(t *testing.T) {
tx := &administrationTransaction{failAt: 3, err: errors.New("projects failed")}
db := NewDatabase(Config{Backend: BackendPostgres}, &administrationPool{tx: tx})
err := db.DeleteAccount(context.Background(), "user-1", "archive:org-1")
if !errors.Is(err, tx.err) || tx.commits != 0 || tx.rollbacks != 1 {
t.Fatalf("err=%v commits=%d rollbacks=%d", err, tx.commits, tx.rollbacks)
}
}
func TestAdministrationStoreFailsClosedWithoutPostgres(t *testing.T) {
db := NewDatabase(Config{Backend: BackendLocal}, &identityQuerier{})
if _, err := db.ListOrganizations(context.Background(), true); err == nil {
t.Fatal("ListOrganizations error=nil")
}
if err := db.DeleteAccount(context.Background(), "u", "a"); err == nil {
t.Fatal("DeleteAccount error=nil")
}
}
var _ administration.Store = (*Database)(nil)
type administrationPool struct{ tx *administrationTransaction }
func (p *administrationPool) Query(context.Context, string, ...any) (Rows, error) {
return nil, errors.New("outside transaction")
}
func (p *administrationPool) Begin(context.Context) (Transaction, error) { return p.tx, nil }
type administrationTransaction struct {
sql []string
args [][]any
failAt int
err error
commits, rollbacks int
}
func (t *administrationTransaction) Query(context.Context, string, ...any) (Rows, error) {
return nil, errors.New("unexpected query")
}
func (t *administrationTransaction) Exec(_ context.Context, query string, args ...any) error {
t.sql = append(t.sql, query)
t.args = append(t.args, args)
if t.failAt == len(t.sql) {
return t.err
}
return nil
}
func (t *administrationTransaction) Commit(context.Context) error { t.commits++; return nil }
func (t *administrationTransaction) Rollback(context.Context) error { t.rollbacks++; return nil }
type administrationQuerier struct {
rows *administrationRows
sql string
args []any
}
func (q *administrationQuerier) Query(_ context.Context, query string, args ...any) (Rows, error) {
q.sql, q.args = query, args
return q.rows, nil
}
type administrationRows struct {
rows [][]any
idx int
}
func (r *administrationRows) Close() {}
func (r *administrationRows) Err() error { return nil }
func (r *administrationRows) Next() bool { return r.idx < len(r.rows) }
func (r *administrationRows) Scan(dest ...any) error {
if r.idx >= len(r.rows) || len(dest) != len(r.rows[r.idx]) {
return errors.New("invalid administration row scan")
}
row := r.rows[r.idx]
r.idx++
for i, target := range dest {
value := row[i]
switch target := target.(type) {
case *string:
*target = value.(string)
case *int:
*target = value.(int)
case *administration.Role:
*target = administration.Role(value.(string))
case *administration.Status:
*target = administration.Status(value.(string))
case *time.Time:
*target = value.(time.Time)
case *sql.NullString:
if value != nil {
*target = sql.NullString{String: value.(string), Valid: true}
}
case *sql.NullTime:
if value != nil {
*target = sql.NullTime{Time: value.(time.Time), Valid: true}
}
default:
return errors.New("unsupported administration scan target")
}
}
return nil
}