221 lines
7.7 KiB
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)
|