236 lines
8.1 KiB
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
|
|
}
|