Files
2026-07-30 12:23:05 +08:00

386 lines
13 KiB
Go

package service
import (
"context"
"errors"
"fmt"
"strings"
"golang.org/x/crypto/bcrypt"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"oci-portal/internal/model"
)
// ErrCredentialConfirm 表示当前密码校验失败;api 层映射 401。
var ErrCredentialConfirm = errors.New("当前密码不正确")
// ErrCredentialInvalid 标记凭据输入非法;api 层映射 400。
var ErrCredentialInvalid = errors.New("凭据输入非法")
// ErrPasswordLoginDisabled 表示密码登录已被禁用;api 层映射 403。
var ErrPasswordLoginDisabled = errors.New("密码登录已禁用,请使用免密方式登录")
// ErrNeedIdentity 表示无任何可用免密登录方式时不可禁用密码登录;api 层映射 409。
var ErrNeedIdentity = errors.New("需先有可用的免密登录方式(通行密钥、钱包或已启用的外部登录),才能禁用密码登录")
// ErrProviderLastLogin 表示密码登录禁用期间不可禁用/清空最后可用的登录方式;api 层映射 409。
var ErrProviderLastLogin = errors.New("密码登录已禁用,该操作将移除最后可用的登录方式;请先允许密码登录")
// ErrLastIdentity 表示密码登录禁用期间不可移除最后一个免密登录方式(防自锁);api 层映射 409。
var ErrLastIdentity = errors.New("密码登录已禁用,不能移除最后一个免密登录方式;请先允许密码登录")
// UpdateCredentialsInput 是修改登录凭据的请求体;两项至少改一项,当前密码必填。
type UpdateCredentialsInput struct {
NewUsername string `json:"newUsername"`
NewPassword string `json:"newPassword"`
CurrentPassword string `json:"currentPassword" binding:"required"`
}
type authenticatedMutation struct {
auth *AuthService
username string
proof TokenProof
}
func (m *authenticatedMutation) lockAndCheck(tx *gorm.DB) error {
user, err := lockUserForAuthChange(tx, m.username)
if err != nil {
return err
}
return m.auth.ensureTokenCurrentTx(tx, user, m.proof)
}
// UpdateCredentials 修改用户名 / 密码:当前密码必验;成功后令牌版本递增
// (全部旧 JWT 立即失效),返回最终用户名供调用方为操作者重签新令牌。
func (s *AuthService) UpdateCredentials(ctx context.Context, username string, in UpdateCredentialsInput, proof TokenProof) (string, error) {
user, err := s.findUser(ctx, username)
if err != nil {
return "", err
}
if bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(in.CurrentPassword)) != nil {
return "", ErrCredentialConfirm
}
newName := strings.TrimSpace(in.NewUsername)
if err := validateCredentialChange(user, newName, in.NewPassword); err != nil {
return "", err
}
updates, finalName, err := s.credentialUpdates(ctx, user, newName, in.NewPassword)
if err != nil {
return "", err
}
updates["token_version"] = gorm.Expr("token_version + 1")
if err := s.applyCredentialUpdates(ctx, username, proof, updates); err != nil {
return "", err
}
return finalName, nil
}
func (s *AuthService) applyCredentialUpdates(
ctx context.Context, username string, proof TokenProof, updates map[string]any,
) error {
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
locked, err := lockUserForAuthChange(tx, username)
if err != nil {
return err
}
if err := s.ensureTokenCurrentTx(tx, locked, proof); err != nil {
return err
}
return updateCredentialsTx(tx, locked.ID, proof.Ver, updates)
})
}
func updateCredentialsTx(tx *gorm.DB, userID uint, version uint, updates map[string]any) error {
res := tx.Model(&model.User{}).
Where("id = ? AND token_version = ?", userID, version).Updates(updates)
if res.Error != nil {
return fmt.Errorf("update credentials: %w", res.Error)
}
if res.RowsAffected == 0 {
return ErrTokenStale
}
return nil
}
// credentialUpdates 组装凭据变更字段并返回最终用户名。
func (s *AuthService) credentialUpdates(ctx context.Context, user *model.User, newName, newPassword string) (map[string]any, string, error) {
updates := map[string]any{}
finalName := user.Username
if newName != "" && newName != user.Username {
taken, err := s.usernameTaken(ctx, newName, user.ID)
if err != nil {
return nil, "", err
}
if taken {
return nil, "", fmt.Errorf("用户名已被占用: %w", ErrCredentialInvalid)
}
updates["username"] = newName
finalName = newName
}
if newPassword != "" {
hash, err := bcrypt.GenerateFromPassword([]byte(newPassword), bcrypt.DefaultCost)
if err != nil {
return nil, "", fmt.Errorf("hash password: %w", err)
}
updates["password_hash"] = string(hash)
}
return updates, finalName, nil
}
// validateCredentialChange 校验改名 / 改密输入;两者均无实际变更时报非法。
func validateCredentialChange(user *model.User, newName, newPassword string) error {
if newName != "" && len(newName) > 64 {
return fmt.Errorf("用户名最长 64 字符: %w", ErrCredentialInvalid)
}
// bcrypt 只取前 72 字节,超长部分静默截断,直接拒绝
if newPassword != "" && (len(newPassword) < 8 || len(newPassword) > 72) {
return fmt.Errorf("新密码长度须在 8-72 之间: %w", ErrCredentialInvalid)
}
if (newName == "" || newName == user.Username) && newPassword == "" {
return fmt.Errorf("没有需要保存的变更: %w", ErrCredentialInvalid)
}
return nil
}
func (s *AuthService) usernameTaken(ctx context.Context, name string, selfID uint) (bool, error) {
var count int64
err := s.db.WithContext(ctx).Model(&model.User{}).
Where("username = ? AND id <> ?", name, selfID).Count(&count).Error
if err != nil {
return false, fmt.Errorf("check username: %w", err)
}
return count > 0, nil
}
// PasswordLoginDisabled 读密码登录禁用开关;settings 未注入(测试)视为未禁用。
func (s *AuthService) PasswordLoginDisabled(ctx context.Context) (bool, error) {
if s.settings == nil {
return false, nil
}
return s.settings.PasswordLoginDisabled(ctx)
}
// SetPasswordLoginDisabled 保存开关;开启前须至少有一种免密登录方式
// (通行密钥或外部身份),防止自锁。检查与写入在同一事务内并锁定用户行,
// 防与解绑身份/删除通行密钥并发绕过「至少一种登录方式」。
func (s *AuthService) SetPasswordLoginDisabled(ctx context.Context, username string, disabled bool, proof TokenProof) error {
if s.settings == nil {
return errors.New("settings unavailable")
}
value := ""
if disabled {
value = "1"
}
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
user, err := lockUserForAuthChange(tx, username)
if err != nil {
return err
}
if err := s.ensureTokenCurrentTx(tx, user, proof); err != nil {
return err
}
if err := s.ensurePasswordlessForToggle(tx, user.ID, disabled); err != nil {
return err
}
if err := saveSettingTx(tx, settingSecPasswordLoginOff, value); err != nil {
return err
}
// 登录策略属敏感变更:版本递增与开关写入同事务提交,不留半程状态
return bumpTokenVersionTx(tx, username)
})
}
func (s *AuthService) ensurePasswordlessForToggle(tx *gorm.DB, userID uint, disabled bool) error {
if !disabled {
return nil
}
origin, err := effectiveOriginTx(tx, s.settings)
if err != nil {
return err
}
ok, err := usablePasswordlessTx(tx, userID, 0, 0, origin)
if err != nil {
return err
}
if !ok {
return ErrNeedIdentity
}
return nil
}
// lockUserForAuthChange 事务内锁定用户行(SQLite 单写天然串行,MySQL/PG 靠行锁),
// 「禁用密码登录」与「解绑身份」都先过这把锁,保证不变量检查与写入不交叉。
func lockUserForAuthChange(tx *gorm.DB, username string) (*model.User, error) {
var user model.User
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
First(&user, "username = ?", username).Error
if err != nil {
return nil, fmt.Errorf("find user %s: %w", username, err)
}
return &user, nil
}
func lockUserByIDForAuthChange(tx *gorm.DB, userID uint) (*model.User, error) {
var user model.User
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&user, userID).Error
if err != nil {
return nil, fmt.Errorf("find user: %w", err)
}
return &user, nil
}
// ensureTokenCurrentTx 行锁下复核请求令牌仍是账号当前状态:版本一致
// (改密 / 撤销全部会递增),且对应会话行未被注销或定点撤销
// (Logout 联动标记 revoked_at,查行即可覆盖两者);存量令牌无行时仅校验版本。
func (s *AuthService) ensureTokenCurrentTx(tx *gorm.DB, user *model.User, proof TokenProof) error {
if user.TokenVersion != proof.Ver {
return ErrTokenStale
}
if proof.Jti == "" {
return nil
}
if s.revokedJti.has(proof.Jti) {
return ErrTokenStale
}
var row model.UserSession
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
Select("id", "revoked_at").Where("token_id = ?", proof.Jti).First(&row).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil
}
return fmt.Errorf("find session: %w", err)
}
if row.RevokedAt != nil {
return ErrTokenStale
}
return nil
}
// lockUsersForAuthChange 锁定全部用户行(单管理员面板即一行):与
// lockUserForAuthChange 竞争同一把行锁,防「改 provider / 面板地址配置」
// 与「禁用密码 / 删除最后因子」并发交错绕过登录方式不变量。
func lockUsersForAuthChange(tx *gorm.DB) error {
var users []model.User
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Find(&users).Error; err != nil {
return fmt.Errorf("lock users: %w", err)
}
return nil
}
// identityProviderCountTx 统计某 provider 的身份总数(单管理员面板,不区分账号)。
func identityProviderCountTx(tx *gorm.DB, provider string) (int64, error) {
var n int64
err := tx.Model(&model.UserIdentity{}).Where("provider = ?", provider).Count(&n).Error
if err != nil {
return 0, fmt.Errorf("count identities: %w", err)
}
return n, nil
}
func identityCountTx(tx *gorm.DB, userID uint) (int64, error) {
var count int64
err := tx.Model(&model.UserIdentity{}).
Where("user_id = ?", userID).Count(&count).Error
if err != nil {
return 0, fmt.Errorf("count identities: %w", err)
}
return count, nil
}
// passkeyCountTx 统计账号的通行密钥数(事务内,防自锁检查用)。
func passkeyCountTx(tx *gorm.DB, userID uint) (int64, error) {
var count int64
err := tx.Model(&model.UserPasskey{}).
Where("user_id = ?", userID).Count(&count).Error
if err != nil {
return 0, fmt.Errorf("count passkeys: %w", err)
}
return count, nil
}
// usablePasswordlessTx 事务内判断(排除给定身份/通行密钥行后)是否仍存在
// 可实际登录的免密方式:通行密钥仅统计注册 origin 与当前面板地址一致的
// (地址变更后旧域名凭据不可登录,不得计入);钱包按地址绑定,恒可用;
// GitHub/OIDC 身份须对应 provider 已配置且未禁用才计入(端到端防自锁)。
// origin 为空时所有方式均不可用:Passkey RP、钱包 SIWE 与 OAuth 回调都依赖面板地址。
func usablePasswordlessTx(tx *gorm.DB, userID uint, excludeIdentity, excludePasskey uint, origin string) (bool, error) {
if origin == "" {
return false, nil
}
pk, err := passkeyCountExcludingTx(tx, userID, excludePasskey, origin)
if err != nil || pk > 0 {
return pk > 0, err
}
if n, err := identityCountByProviderTx(tx, userID, "wallet", excludeIdentity); err != nil || n > 0 {
return n > 0, err
}
for _, p := range []string{"github", "oidc"} {
n, err := identityCountByProviderTx(tx, userID, p, excludeIdentity)
if err != nil {
return false, err
}
if n == 0 {
continue
}
if ok, err := oauthProviderUsableTx(tx, p); err != nil || ok {
return ok, err
}
}
return false, nil
}
// oauthProviderUsableTx 事务内判断 provider 当前可实际登录:
// clientID 与 secret 均已配置(oidc 还需 issuer)且未禁用。
func oauthProviderUsableTx(tx *gorm.DB, provider string) (bool, error) {
need := []string{settingOauthGithubClientID, settingOauthGithubClientSecret}
offKey := settingOauthGithubDisabled
if provider == "oidc" {
need = []string{settingOauthOidcClientID, settingOauthOidcClientSecret, settingOauthOidcIssuer}
offKey = settingOauthOidcDisabled
}
for _, k := range need {
v, err := settingValueTx(tx, k)
if err != nil || v == "" {
return false, err
}
}
off, err := settingValueTx(tx, offKey)
return off != "1", err
}
// identityCountByProviderTx 统计账号某 provider 的身份数,可排除一行(解绑前判定用)。
func identityCountByProviderTx(tx *gorm.DB, userID uint, provider string, excludeID uint) (int64, error) {
q := tx.Model(&model.UserIdentity{}).Where("user_id = ? AND provider = ?", userID, provider)
if excludeID != 0 {
q = q.Where("id <> ?", excludeID)
}
var count int64
if err := q.Count(&count).Error; err != nil {
return 0, fmt.Errorf("count identities: %w", err)
}
return count, nil
}
// passkeyCountExcludingTx 统计「当前地址下可用」的通行密钥数,可排除一行;
// userID 为 0 表示全表(单管理员面板);origin 非空时仅计注册来源一致的凭据。
func passkeyCountExcludingTx(tx *gorm.DB, userID, excludeID uint, origin string) (int64, error) {
q := tx.Model(&model.UserPasskey{})
if userID != 0 {
q = q.Where("user_id = ?", userID)
}
if excludeID != 0 {
q = q.Where("id <> ?", excludeID)
}
if origin != "" {
q = q.Where("origin = ?", origin)
}
var count int64
if err := q.Count(&count).Error; err != nil {
return 0, fmt.Errorf("count passkeys: %w", err)
}
return count, nil
}