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

371 lines
14 KiB
Go

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)