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

221 lines
7.7 KiB
Go

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)