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 }