package postgres import ( "context" "database/sql" "encoding/hex" "errors" "reflect" "testing" "time" "git.nianxx.cn/wangxuming/NianAIGC/backend/internal/identity" "golang.org/x/crypto/scrypt" ) func TestAttemptPasswordLoginSuccessUsesOneLockedTransaction(t *testing.T) { now := time.Date(2026, 8, 13, 12, 0, 0, 0, time.UTC) hash := nodeCompatibleHash(t, "secret", "001122aabbccddeeff") tx := &loginTransaction{queries: []loginQueryResult{ {rows: loginRows([]any{"user-1", "13800138000", "Name", "user", "org-1", "active", hash, "001122aabbccddeeff", 2, nil, 9})}, {rows: loginRows([]any{"org-1", "Acme", "active"})}, }} db := loginDatabase(tx) got, err := db.AttemptPasswordLogin(context.Background(), "13800138000", "secret", now) if err != nil { t.Fatalf("AttemptPasswordLogin() error = %v", err) } if got.Account.ID != "user-1" || got.Account.SessionVersion != 9 || got.Organization == nil || got.Organization.ID != "org-1" { t.Fatalf("AttemptPasswordLogin() = %#v", got) } wantEvents := []string{"BEGIN", "QUERY user", "QUERY organization", "EXEC success", "COMMIT"} if !reflect.DeepEqual(tx.events, wantEvents) { t.Fatalf("events = %#v, want %#v", tx.events, wantEvents) } if tx.queriesSeen[0].sql != SelectPasswordLoginAccountSQL || !reflect.DeepEqual(tx.queriesSeen[0].args, []any{"13800138000"}) { t.Fatalf("account query = %#v", tx.queriesSeen[0]) } if tx.queriesSeen[1].sql != SelectPasswordLoginOrganizationSQL || !reflect.DeepEqual(tx.queriesSeen[1].args, []any{"org-1"}) { t.Fatalf("organization query = %#v", tx.queriesSeen[1]) } if len(tx.execs) != 1 || tx.execs[0].sql != RecordSuccessfulPasswordLoginSQL || !reflect.DeepEqual(tx.execs[0].args, []any{"user-1", now}) { t.Fatalf("success update = %#v", tx.execs) } } func TestAttemptPasswordLoginPreservesBoundSuperAdministratorOrganizationProfile(t *testing.T) { now := time.Date(2026, 8, 13, 12, 0, 0, 0, time.UTC) hash := nodeCompatibleHash(t, "secret", "salt") for _, status := range []string{"active", "disabled"} { t.Run(status, func(t *testing.T) { tx := &loginTransaction{queries: []loginQueryResult{ {rows: loginRows([]any{"super-1", "13800138000", "Super", "super_admin", "org-1", "active", hash, "salt", 0, nil, 3})}, {rows: loginRows([]any{"org-1", "Bound Organization", status})}, }} got, err := loginDatabase(tx).AttemptPasswordLogin(context.Background(), "13800138000", "secret", now) if err != nil { t.Fatalf("AttemptPasswordLogin() error = %v", err) } if got.Organization == nil || got.Organization.ID != "org-1" || got.Organization.Name != "Bound Organization" || got.Organization.Status != status { t.Fatalf("organization = %#v", got.Organization) } wantEvents := []string{"BEGIN", "QUERY user", "QUERY organization", "EXEC success", "COMMIT"} if !reflect.DeepEqual(tx.events, wantEvents) { t.Fatalf("events = %#v, want %#v", tx.events, wantEvents) } }) } } func TestAttemptPasswordLoginAllowsBoundSuperAdministratorWithMissingOrganization(t *testing.T) { now := time.Date(2026, 8, 13, 12, 0, 0, 0, time.UTC) hash := nodeCompatibleHash(t, "secret", "salt") tx := &loginTransaction{queries: []loginQueryResult{ {rows: loginRows([]any{"super-1", "13800138000", "Super", "super_admin", "deleted-org", "active", hash, "salt", 0, nil, 3})}, {rows: &loginFakeRows{}}, }} got, err := loginDatabase(tx).AttemptPasswordLogin(context.Background(), "13800138000", "secret", now) if err != nil { t.Fatalf("AttemptPasswordLogin() error = %v", err) } if got.Organization != nil || got.Account.OrganizationID != "deleted-org" { t.Fatalf("result = %#v", got) } if len(tx.execs) != 1 || tx.execs[0].sql != RecordSuccessfulPasswordLoginSQL || tx.commits != 1 { t.Fatalf("execs=%#v commits=%d", tx.execs, tx.commits) } } func TestAttemptPasswordLoginCommitsExpectedDenials(t *testing.T) { now := time.Date(2026, 8, 13, 12, 0, 0, 0, time.UTC) hash := nodeCompatibleHash(t, "secret", "salt") tests := []struct { name string row []any password string wantReason identity.PasswordLoginFailure wantExec []any }{ {name: "missing", wantReason: identity.LoginFailureInvalidCredentials}, {name: "disabled", row: []any{"u", "p", "N", "user", "org", "disabled", hash, "salt", 0, nil, 1}, password: "secret", wantReason: identity.LoginFailureAccountDisabled}, {name: "locked", row: []any{"u", "p", "N", "super_admin", nil, "active", hash, "salt", 0, now.Add(time.Minute), 1}, password: "secret", wantReason: identity.LoginFailureAccountLocked}, {name: "bad password", row: []any{"u", "p", "N", "super_admin", nil, "active", hash, "salt", 3, nil, 1}, password: "wrong", wantReason: identity.LoginFailureInvalidCredentials, wantExec: []any{"u", 4, nil, now}}, {name: "fifth failure locks", row: []any{"u", "p", "N", "super_admin", nil, "active", hash, "salt", 4, nil, 1}, password: "wrong", wantReason: identity.LoginFailureAccountLocked, wantExec: []any{"u", 0, now.Add(15 * time.Minute), now}}, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { rows := &loginFakeRows{} if test.row != nil { rows.rows = [][]any{test.row} } tx := &loginTransaction{queries: []loginQueryResult{{rows: rows}}} _, err := loginDatabase(tx).AttemptPasswordLogin(context.Background(), "p", test.password, now) var denial *identity.PasswordLoginError if !errors.As(err, &denial) || denial.Reason != test.wantReason { t.Fatalf("error = %v, want reason %s", err, test.wantReason) } if tx.commits != 1 || tx.rollbacks != 0 { t.Fatalf("commits=%d rollbacks=%d, want 1/0", tx.commits, tx.rollbacks) } if test.wantExec == nil { if len(tx.execs) != 0 { t.Fatalf("unexpected execs %#v", tx.execs) } } else if len(tx.execs) != 1 || tx.execs[0].sql != RecordFailedPasswordLoginSQL || !reflect.DeepEqual(tx.execs[0].args, test.wantExec) { t.Fatalf("failure update = %#v, want args %#v", tx.execs, test.wantExec) } }) } } func TestAttemptPasswordLoginRejectsInactiveOrganizationBeforePasswordState(t *testing.T) { now := time.Date(2026, 8, 13, 12, 0, 0, 0, time.UTC) hash := nodeCompatibleHash(t, "secret", "salt") for _, test := range []struct { name string password string lockedUntil any }{ {name: "wrong password", password: "wrong"}, {name: "locked account", password: "secret", lockedUntil: now.Add(time.Minute)}, } { t.Run(test.name, func(t *testing.T) { tx := &loginTransaction{queries: []loginQueryResult{ {rows: loginRows([]any{"u", "p", "N", "user", "org", "active", hash, "salt", 4, test.lockedUntil, 1})}, {rows: loginRows([]any{"org", "Acme", "disabled"})}, }} _, err := loginDatabase(tx).AttemptPasswordLogin(context.Background(), "p", test.password, now) assertCommittedDenialWithoutUpdate(t, tx, err, identity.LoginFailureOrganizationNotActive) }) } } func TestAttemptPasswordLoginRejectsInvalidRoleAndMissingOrganizationBeforePasswordState(t *testing.T) { now := time.Date(2026, 8, 13, 12, 0, 0, 0, time.UTC) hash := nodeCompatibleHash(t, "secret", "salt") for _, test := range []struct { name string role string orgID any wantReason identity.PasswordLoginFailure }{ {name: "invalid role", role: "root", orgID: "org", wantReason: identity.LoginFailureInvalidRole}, {name: "organization required", role: "user", orgID: nil, wantReason: identity.LoginFailureOrganizationRequired}, } { t.Run(test.name, func(t *testing.T) { tx := &loginTransaction{queries: []loginQueryResult{{rows: loginRows([]any{ "u", "p", "N", test.role, test.orgID, "active", hash, "salt", 4, now.Add(time.Minute), 1, })}}} _, err := loginDatabase(tx).AttemptPasswordLogin(context.Background(), "p", "wrong", now) assertCommittedDenialWithoutUpdate(t, tx, err, test.wantReason) if len(tx.queriesSeen) != 1 { t.Fatalf("queries = %#v, want account query only", tx.queriesSeen) } }) } } func assertCommittedDenialWithoutUpdate(t *testing.T, tx *loginTransaction, err error, wantReason identity.PasswordLoginFailure) { t.Helper() var denial *identity.PasswordLoginError if !errors.As(err, &denial) || denial.Reason != wantReason { t.Fatalf("error = %v, want reason %s", err, wantReason) } if len(tx.execs) != 0 || tx.commits != 1 || tx.rollbacks != 0 { t.Fatalf("execs=%#v commits=%d rollbacks=%d, want no update and commit", tx.execs, tx.commits, tx.rollbacks) } } func TestAttemptPasswordLoginRollsBackInfrastructureFailures(t *testing.T) { now := time.Now() queryErr := errors.New("query failed") tx := &loginTransaction{queries: []loginQueryResult{{err: queryErr}}} _, err := loginDatabase(tx).AttemptPasswordLogin(context.Background(), "p", "secret", now) if !errors.Is(err, queryErr) || tx.commits != 0 || tx.rollbacks != 1 { t.Fatalf("error=%v commits=%d rollbacks=%d", err, tx.commits, tx.rollbacks) } commitErr := errors.New("commit failed") tx = &loginTransaction{queries: []loginQueryResult{{rows: &loginFakeRows{}}}, commitErr: commitErr} _, err = loginDatabase(tx).AttemptPasswordLogin(context.Background(), "p", "secret", now) if !errors.Is(err, commitErr) || tx.commits != 1 { t.Fatalf("commit error=%v commits=%d", err, tx.commits) } } func TestAttemptPasswordLoginFailsClosedWithoutTransactionSupport(t *testing.T) { tests := []struct { name string config Config query Querier }{ {name: "local", config: Config{Backend: BackendLocal}, query: &identityQuerier{}}, {name: "postgres without beginner", config: Config{Backend: BackendPostgres}, query: &identityQuerier{}}, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { _, err := NewDatabase(test.config, test.query).AttemptPasswordLogin(context.Background(), "p", "secret", time.Now()) if err == nil { t.Fatal("AttemptPasswordLogin() error = nil") } }) } } func TestAttemptPasswordLoginMalformedStoredHashIsAnExpectedCredentialDenial(t *testing.T) { tx := &loginTransaction{queries: []loginQueryResult{{rows: loginRows([]any{ "u", "p", "N", "super_admin", nil, "active", "not-hex", "salt", 0, nil, 1, })}}} _, err := loginDatabase(tx).AttemptPasswordLogin(context.Background(), "p", "secret", time.Now()) var denial *identity.PasswordLoginError if !errors.As(err, &denial) || denial.Reason != identity.LoginFailureInvalidCredentials || tx.commits != 1 || len(tx.execs) != 1 { t.Fatalf("error=%v commits=%d execs=%#v", err, tx.commits, tx.execs) } } func TestAttemptPasswordLoginSerializesFiveFailureTransitions(t *testing.T) { now := time.Date(2026, 8, 13, 12, 0, 0, 0, time.UTC) hash := nodeCompatibleHash(t, "secret", "salt") for attempt := 1; attempt <= 5; attempt++ { tx := &loginTransaction{queries: []loginQueryResult{{rows: loginRows([]any{"u", "p", "N", "super_admin", nil, "active", hash, "salt", attempt - 1, nil, 1})}}} _, err := loginDatabase(tx).AttemptPasswordLogin(context.Background(), "p", "wrong", now) var denial *identity.PasswordLoginError if !errors.As(err, &denial) { t.Fatalf("attempt %d error=%v", attempt, err) } wantCount := attempt wantReason := identity.LoginFailureInvalidCredentials if attempt == 5 { wantCount, wantReason = 0, identity.LoginFailureAccountLocked } if denial.Reason != wantReason || tx.execs[0].args[1] != wantCount || tx.commits != 1 { t.Fatalf("attempt %d reason=%s args=%#v commits=%d", attempt, denial.Reason, tx.execs[0].args, tx.commits) } } } func nodeCompatibleHash(t *testing.T, password, salt string) string { t.Helper() derived, err := scrypt.Key([]byte(password), []byte(salt), 16384, 8, 1, 64) if err != nil { t.Fatal(err) } return hex.EncodeToString(derived) } func loginDatabase(tx *loginTransaction) *Database { return NewDatabase(Config{Backend: BackendPostgres}, &loginPool{tx: tx}) } type loginPool struct{ tx *loginTransaction } func (p *loginPool) Query(context.Context, string, ...any) (Rows, error) { return nil, errors.New("query outside transaction") } func (p *loginPool) Begin(context.Context) (Transaction, error) { p.tx.events = append(p.tx.events, "BEGIN") return p.tx, nil } type loginCall struct { sql string args []any } type loginQueryResult struct { rows Rows err error } type loginTransaction struct { queries []loginQueryResult queriesSeen []loginCall execs []loginCall events []string commitErr error commits, rollbacks int } func (t *loginTransaction) Query(_ context.Context, query string, args ...any) (Rows, error) { t.queriesSeen = append(t.queriesSeen, loginCall{query, args}) if query == SelectPasswordLoginAccountSQL { t.events = append(t.events, "QUERY user") } else { t.events = append(t.events, "QUERY organization") } result := t.queries[0] t.queries = t.queries[1:] return result.rows, result.err } func (t *loginTransaction) Exec(_ context.Context, query string, args ...any) error { t.execs = append(t.execs, loginCall{query, args}) if query == RecordSuccessfulPasswordLoginSQL { t.events = append(t.events, "EXEC success") } else { t.events = append(t.events, "EXEC failure") } return nil } func (t *loginTransaction) Commit(context.Context) error { t.commits++ t.events = append(t.events, "COMMIT") return t.commitErr } func (t *loginTransaction) Rollback(context.Context) error { t.rollbacks++ t.events = append(t.events, "ROLLBACK") return nil } func loginRows(row []any) *loginFakeRows { return &loginFakeRows{rows: [][]any{row}} } type loginFakeRows struct { rows [][]any index int err error } func (r *loginFakeRows) Close() {} func (r *loginFakeRows) Err() error { return r.err } func (r *loginFakeRows) Next() bool { return r.index < len(r.rows) } func (r *loginFakeRows) Scan(dest ...any) error { if r.index >= len(r.rows) { return errors.New("scan past end") } row := r.rows[r.index] r.index++ if len(row) != len(dest) { return errors.New("scan arity mismatch") } for i, value := range row { switch target := dest[i].(type) { case *string: *target = value.(string) case *int: *target = value.(int) 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 scan target") } } return nil } var _ identity.CredentialAuthenticator = (*Database)(nil)