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)