feat: add Go password session lifecycle

This commit is contained in:
2026-08-13 15:34:46 +08:00
parent 772795e7eb
commit d0207fcebe
19 changed files with 2483 additions and 16 deletions

View File

@@ -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 {

View File

@@ -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)

View 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)

View 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)