371 lines
14 KiB
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)
|