@@ -8,6 +8,7 @@ import (
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
)
|
||||
@@ -119,33 +120,53 @@ func (s *AuthService) PasswordLoginDisabled(ctx context.Context) (bool, error) {
|
||||
}
|
||||
|
||||
// SetPasswordLoginDisabled 保存开关;开启前须至少绑定一个外部身份,防止自锁。
|
||||
// 检查与写入在同一事务内并锁定用户行,防与解绑身份并发绕过「至少一种登录方式」。
|
||||
func (s *AuthService) SetPasswordLoginDisabled(ctx context.Context, username string, disabled bool) error {
|
||||
if s.settings == nil {
|
||||
return errors.New("settings unavailable")
|
||||
}
|
||||
value := ""
|
||||
if disabled {
|
||||
user, err := s.findUser(ctx, username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
n, err := s.identityCount(ctx, user.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n == 0 {
|
||||
return ErrNeedIdentity
|
||||
}
|
||||
value = "1"
|
||||
}
|
||||
if err := s.settings.SetPasswordLoginDisabled(ctx, disabled); err != nil {
|
||||
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
user, err := lockUserForAuthChange(tx, username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if disabled {
|
||||
n, err := identityCountTx(tx, user.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n == 0 {
|
||||
return ErrNeedIdentity
|
||||
}
|
||||
}
|
||||
return saveSettingTx(tx, settingSecPasswordLoginOff, value)
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// 登录策略属敏感变更:递增令牌版本,已签发会话全部失效
|
||||
return s.bumpTokenVersion(ctx, username)
|
||||
}
|
||||
|
||||
func (s *AuthService) identityCount(ctx context.Context, userID uint) (int64, error) {
|
||||
// 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 identityCountTx(tx *gorm.DB, userID uint) (int64, error) {
|
||||
var count int64
|
||||
err := s.db.WithContext(ctx).Model(&model.UserIdentity{}).
|
||||
err := tx.Model(&model.UserIdentity{}).
|
||||
Where("user_id = ?", userID).Count(&count).Error
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("count identities: %w", err)
|
||||
|
||||
Reference in New Issue
Block a user