feat: add Go password session lifecycle
This commit is contained in:
@@ -56,20 +56,36 @@ type Querier interface {
|
||||
Query(context.Context, string, ...any) (Rows, error)
|
||||
}
|
||||
|
||||
type Transaction interface {
|
||||
Querier
|
||||
Exec(context.Context, string, ...any) error
|
||||
Commit(context.Context) error
|
||||
Rollback(context.Context) error
|
||||
}
|
||||
|
||||
type TransactionBeginner interface {
|
||||
Begin(context.Context) (Transaction, error)
|
||||
}
|
||||
|
||||
type Pool interface {
|
||||
Querier
|
||||
Close()
|
||||
}
|
||||
|
||||
type Database struct {
|
||||
config Config
|
||||
querier Querier
|
||||
config Config
|
||||
querier Querier
|
||||
transactions TransactionBeginner
|
||||
}
|
||||
|
||||
type Store = Database
|
||||
|
||||
func NewDatabase(config Config, querier Querier) *Database {
|
||||
return &Database{config: config, querier: querier}
|
||||
db := &Database{config: config, querier: querier}
|
||||
if transactions, ok := querier.(TransactionBeginner); ok {
|
||||
db.transactions = transactions
|
||||
}
|
||||
return db
|
||||
}
|
||||
|
||||
func (db *Database) Readiness(ctx context.Context) error {
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
)
|
||||
|
||||
@@ -68,8 +69,40 @@ func (p *pgxPoolAdapter) Query(ctx context.Context, sql string, args ...any) (Ro
|
||||
return p.pool.Query(ctx, sql, args...)
|
||||
}
|
||||
|
||||
func (p *pgxPoolAdapter) Begin(ctx context.Context) (Transaction, error) {
|
||||
tx, err := p.pool.Begin(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &pgxTransactionAdapter{tx: tx}, nil
|
||||
}
|
||||
|
||||
func (p *pgxPoolAdapter) Close() {
|
||||
p.pool.Close()
|
||||
}
|
||||
|
||||
var _ Pool = (*pgxPoolAdapter)(nil)
|
||||
var _ TransactionBeginner = (*pgxPoolAdapter)(nil)
|
||||
|
||||
type pgxTransactionAdapter struct {
|
||||
tx pgx.Tx
|
||||
}
|
||||
|
||||
func (t *pgxTransactionAdapter) Query(ctx context.Context, sql string, args ...any) (Rows, error) {
|
||||
return t.tx.Query(ctx, sql, args...)
|
||||
}
|
||||
|
||||
func (t *pgxTransactionAdapter) Exec(ctx context.Context, sql string, args ...any) error {
|
||||
_, err := t.tx.Exec(ctx, sql, args...)
|
||||
return err
|
||||
}
|
||||
|
||||
func (t *pgxTransactionAdapter) Commit(ctx context.Context) error {
|
||||
return t.tx.Commit(ctx)
|
||||
}
|
||||
|
||||
func (t *pgxTransactionAdapter) Rollback(ctx context.Context) error {
|
||||
return t.tx.Rollback(ctx)
|
||||
}
|
||||
|
||||
var _ Transaction = (*pgxTransactionAdapter)(nil)
|
||||
|
||||
220
backend/internal/postgres/password_login.go
Normal file
220
backend/internal/postgres/password_login.go
Normal file
@@ -0,0 +1,220 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/subtle"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/identity"
|
||||
"golang.org/x/crypto/scrypt"
|
||||
)
|
||||
|
||||
const SelectPasswordLoginAccountSQL = `SELECT
|
||||
id,
|
||||
phone,
|
||||
display_name,
|
||||
role,
|
||||
organization_id,
|
||||
status,
|
||||
password_hash,
|
||||
password_salt,
|
||||
failed_login_count,
|
||||
locked_until,
|
||||
session_version
|
||||
FROM public.platform_users
|
||||
WHERE phone = $1::text
|
||||
FOR UPDATE`
|
||||
|
||||
const SelectPasswordLoginOrganizationSQL = `SELECT
|
||||
id,
|
||||
name,
|
||||
status
|
||||
FROM public.platform_organizations
|
||||
WHERE id = $1::text`
|
||||
|
||||
const RecordFailedPasswordLoginSQL = `UPDATE public.platform_users
|
||||
SET failed_login_count = $2,
|
||||
locked_until = $3,
|
||||
updated_at = $4
|
||||
WHERE id = $1::text`
|
||||
|
||||
const RecordSuccessfulPasswordLoginSQL = `UPDATE public.platform_users
|
||||
SET failed_login_count = 0,
|
||||
locked_until = NULL,
|
||||
last_login_at = $2,
|
||||
updated_at = $2
|
||||
WHERE id = $1::text`
|
||||
|
||||
const (
|
||||
passwordLoginMaxFailures = 5
|
||||
passwordLoginLockTime = 15 * time.Minute
|
||||
)
|
||||
|
||||
// AttemptPasswordLogin performs each credential attempt under the account
|
||||
// row's PostgreSQL lock. Expected authentication denials are committed so a
|
||||
// failed-password transition cannot be accidentally rolled back by a caller.
|
||||
func (db *Database) AttemptPasswordLogin(ctx context.Context, phone, password string, now time.Time) (identity.LoginAccount, error) {
|
||||
if db.config.Backend != BackendPostgres || db.transactions == nil {
|
||||
return identity.LoginAccount{}, fmt.Errorf("PostgreSQL is unavailable when ZHINIAN_DATA_BACKEND=%s", db.config.Backend)
|
||||
}
|
||||
tx, err := db.transactions.Begin(ctx)
|
||||
if err != nil {
|
||||
return identity.LoginAccount{}, fmt.Errorf("begin password login transaction: %w", err)
|
||||
}
|
||||
finished := false
|
||||
defer func() {
|
||||
if !finished {
|
||||
_ = tx.Rollback(ctx)
|
||||
}
|
||||
}()
|
||||
|
||||
account, passwordHash, passwordSalt, failedCount, lockedUntil, found, err := loadPasswordLoginAccount(ctx, tx, phone)
|
||||
if err != nil {
|
||||
return identity.LoginAccount{}, err
|
||||
}
|
||||
if !found {
|
||||
return commitPasswordLoginDenial(ctx, tx, &finished, identity.LoginFailureInvalidCredentials)
|
||||
}
|
||||
if account.Status != "active" {
|
||||
return commitPasswordLoginDenial(ctx, tx, &finished, identity.LoginFailureAccountDisabled)
|
||||
}
|
||||
result := identity.LoginAccount{Account: account}
|
||||
switch account.Role {
|
||||
case "super_admin":
|
||||
if account.OrganizationID != "" {
|
||||
organization, found, err := loadPasswordLoginOrganization(ctx, tx, account.OrganizationID)
|
||||
if err != nil {
|
||||
return identity.LoginAccount{}, err
|
||||
}
|
||||
if found {
|
||||
result.Organization = &organization
|
||||
}
|
||||
}
|
||||
case "user", "organization_admin":
|
||||
if account.OrganizationID == "" {
|
||||
return commitPasswordLoginDenial(ctx, tx, &finished, identity.LoginFailureOrganizationRequired)
|
||||
}
|
||||
organization, found, err := loadPasswordLoginOrganization(ctx, tx, account.OrganizationID)
|
||||
if err != nil {
|
||||
return identity.LoginAccount{}, err
|
||||
}
|
||||
if !found || organization.ID != account.OrganizationID || organization.Status != "active" {
|
||||
return commitPasswordLoginDenial(ctx, tx, &finished, identity.LoginFailureOrganizationNotActive)
|
||||
}
|
||||
result.Organization = &organization
|
||||
default:
|
||||
return commitPasswordLoginDenial(ctx, tx, &finished, identity.LoginFailureInvalidRole)
|
||||
}
|
||||
if lockedUntil.Valid && lockedUntil.Time.After(now) {
|
||||
return commitPasswordLoginDenial(ctx, tx, &finished, identity.LoginFailureAccountLocked)
|
||||
}
|
||||
|
||||
validPassword, err := verifyNodeScryptPassword(password, passwordHash, passwordSalt)
|
||||
if err != nil {
|
||||
return identity.LoginAccount{}, fmt.Errorf("verify password: %w", err)
|
||||
}
|
||||
if !validPassword {
|
||||
failedCount++
|
||||
var nextLockedUntil any
|
||||
reason := identity.LoginFailureInvalidCredentials
|
||||
if failedCount >= passwordLoginMaxFailures {
|
||||
failedCount = 0
|
||||
nextLockedUntil = now.Add(passwordLoginLockTime)
|
||||
reason = identity.LoginFailureAccountLocked
|
||||
}
|
||||
if err := tx.Exec(ctx, RecordFailedPasswordLoginSQL, account.ID, failedCount, nextLockedUntil, now); err != nil {
|
||||
return identity.LoginAccount{}, fmt.Errorf("record failed password login: %w", err)
|
||||
}
|
||||
return commitPasswordLoginDenial(ctx, tx, &finished, reason)
|
||||
}
|
||||
|
||||
if err := tx.Exec(ctx, RecordSuccessfulPasswordLoginSQL, account.ID, now); err != nil {
|
||||
return identity.LoginAccount{}, fmt.Errorf("record successful password login: %w", err)
|
||||
}
|
||||
if err := tx.Commit(ctx); err != nil {
|
||||
return identity.LoginAccount{}, fmt.Errorf("commit password login transaction: %w", err)
|
||||
}
|
||||
finished = true
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func loadPasswordLoginAccount(ctx context.Context, tx Transaction, phone string) (identity.AccountSnapshot, string, string, int, sql.NullTime, bool, error) {
|
||||
rows, err := tx.Query(ctx, SelectPasswordLoginAccountSQL, phone)
|
||||
if err != nil {
|
||||
return identity.AccountSnapshot{}, "", "", 0, sql.NullTime{}, false, fmt.Errorf("query password login account: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
if !rows.Next() {
|
||||
if err := rows.Err(); err != nil {
|
||||
return identity.AccountSnapshot{}, "", "", 0, sql.NullTime{}, false, fmt.Errorf("read password login account: %w", err)
|
||||
}
|
||||
return identity.AccountSnapshot{}, "", "", 0, sql.NullTime{}, false, nil
|
||||
}
|
||||
|
||||
var account identity.AccountSnapshot
|
||||
var organizationID sql.NullString
|
||||
var passwordHash, passwordSalt string
|
||||
var failedCount int
|
||||
var lockedUntil sql.NullTime
|
||||
if err := rows.Scan(
|
||||
&account.ID, &account.Phone, &account.DisplayName, &account.Role,
|
||||
&organizationID, &account.Status, &passwordHash, &passwordSalt,
|
||||
&failedCount, &lockedUntil, &account.SessionVersion,
|
||||
); err != nil {
|
||||
return identity.AccountSnapshot{}, "", "", 0, sql.NullTime{}, false, fmt.Errorf("scan password login account: %w", err)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return identity.AccountSnapshot{}, "", "", 0, sql.NullTime{}, false, fmt.Errorf("read password login account: %w", err)
|
||||
}
|
||||
if organizationID.Valid {
|
||||
account.OrganizationID = organizationID.String
|
||||
}
|
||||
return account, passwordHash, passwordSalt, failedCount, lockedUntil, true, nil
|
||||
}
|
||||
|
||||
func loadPasswordLoginOrganization(ctx context.Context, tx Transaction, organizationID string) (identity.OrganizationSnapshot, bool, error) {
|
||||
rows, err := tx.Query(ctx, SelectPasswordLoginOrganizationSQL, organizationID)
|
||||
if err != nil {
|
||||
return identity.OrganizationSnapshot{}, false, fmt.Errorf("query password login organization: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
if !rows.Next() {
|
||||
if err := rows.Err(); err != nil {
|
||||
return identity.OrganizationSnapshot{}, false, fmt.Errorf("read password login organization: %w", err)
|
||||
}
|
||||
return identity.OrganizationSnapshot{}, false, nil
|
||||
}
|
||||
var organization identity.OrganizationSnapshot
|
||||
if err := rows.Scan(&organization.ID, &organization.Name, &organization.Status); err != nil {
|
||||
return identity.OrganizationSnapshot{}, false, fmt.Errorf("scan password login organization: %w", err)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return identity.OrganizationSnapshot{}, false, fmt.Errorf("read password login organization: %w", err)
|
||||
}
|
||||
return organization, true, nil
|
||||
}
|
||||
|
||||
func commitPasswordLoginDenial(ctx context.Context, tx Transaction, finished *bool, reason identity.PasswordLoginFailure) (identity.LoginAccount, error) {
|
||||
if err := tx.Commit(ctx); err != nil {
|
||||
return identity.LoginAccount{}, fmt.Errorf("commit password login denial: %w", err)
|
||||
}
|
||||
*finished = true
|
||||
return identity.LoginAccount{}, identity.NewPasswordLoginError(reason)
|
||||
}
|
||||
|
||||
func verifyNodeScryptPassword(password, encodedHash, salt string) (bool, error) {
|
||||
derived, err := scrypt.Key([]byte(password), []byte(salt), 16384, 8, 1, 64)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
expected, decodeErr := hex.DecodeString(encodedHash)
|
||||
if decodeErr != nil || len(expected) != len(derived) {
|
||||
expected = make([]byte, len(derived))
|
||||
}
|
||||
return subtle.ConstantTimeCompare(derived, expected) == 1 && decodeErr == nil, nil
|
||||
}
|
||||
|
||||
var _ identity.CredentialAuthenticator = (*Database)(nil)
|
||||
370
backend/internal/postgres/password_login_test.go
Normal file
370
backend/internal/postgres/password_login_test.go
Normal file
@@ -0,0 +1,370 @@
|
||||
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)
|
||||
Reference in New Issue
Block a user