From 02f458856d8d0a2b4986a5ad582924731bbff6a2 Mon Sep 17 00:00:00 2001 From: ryan Date: Mon, 8 Jun 2026 20:03:17 +0800 Subject: [PATCH] =?UTF-8?q?=E7=94=A8=E6=88=B7=E4=BC=98=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/apps/oauth/sources.go | 2 +- internal/apps/user/controllers.go | 25 +----------------- internal/model/users.go | 43 +++++++++++++++++++++++++++++-- 3 files changed, 43 insertions(+), 27 deletions(-) diff --git a/internal/apps/oauth/sources.go b/internal/apps/oauth/sources.go index e6e78dee..a38b9010 100644 --- a/internal/apps/oauth/sources.go +++ b/internal/apps/oauth/sources.go @@ -508,7 +508,7 @@ func Callback(c *gin.Context) { return } userInfo.Username = username - if err := user.CreateUser(db.DB(ctx), userInfo); err != nil { + if err := user.CreateUser(ctx, db.DB(ctx), userInfo); err != nil { c.JSON(http.StatusInternalServerError, util.Err(err.Error())) return } diff --git a/internal/apps/user/controllers.go b/internal/apps/user/controllers.go index 766d1cd8..13573436 100644 --- a/internal/apps/user/controllers.go +++ b/internal/apps/user/controllers.go @@ -354,29 +354,6 @@ func Register(c *gin.Context) { _ = db.Redis.Del(ctx, db.PrefixedKey(codeKey)).Err() } - var count int64 - if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", req.Username).Count(&count).Error; err != nil { - c.JSON(http.StatusOK, util.Err(err.Error())) - return - } - if count > 0 { - c.JSON(http.StatusOK, util.Err("用户名已存在")) - return - } - - // 校验邮箱是否已被其他用户使用 - if req.Email != "" { - var emailCount int64 - if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", req.Email).Count(&emailCount).Error; err != nil { - c.JSON(http.StatusOK, util.Err(err.Error())) - return - } - if emailCount > 0 { - c.JSON(http.StatusOK, util.Err("该邮箱已被其他账号绑定")) - return - } - } - user := model.User{ Username: req.Username, Nickname: req.Nickname, @@ -397,7 +374,7 @@ func Register(c *gin.Context) { return } - if err := db.DB(ctx).Create(&user).Error; err != nil { + if err := user.RegisterUser(ctx, db.DB(ctx)); err != nil { c.JSON(http.StatusOK, util.Err(err.Error())) return } diff --git a/internal/model/users.go b/internal/model/users.go index bf03d7ad..816ae18b 100644 --- a/internal/model/users.go +++ b/internal/model/users.go @@ -18,6 +18,7 @@ limitations under the License. package model import ( + "context" "errors" "strconv" "strings" @@ -127,8 +128,13 @@ func (u *User) CheckActive() error { return nil } -// CreateUser 创建新用户 -func (u *User) CreateUser(tx *gorm.DB, oauthInfo *OAuthUserInfo) error { +// CreateUser 创建新用户(用于 OAuth/OIDC 自动注册,含底层权限校验) +func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUserInfo) error { + enabled, err := GetBoolByKey(ctx, ConfigKeyRegistrationEnabled) + if err == nil && !enabled { + return errors.New("注册已关闭") + } + now := time.Now() newUser := User{ ID: oauthInfo.GetID(), @@ -147,3 +153,36 @@ func (u *User) CreateUser(tx *gorm.DB, oauthInfo *OAuthUserInfo) error { *u = newUser return nil } + +// RegisterUser 创建新用户并注册(用于本地密码注册,含全局开关和唯一性多重底层校验) +func (u *User) RegisterUser(ctx context.Context, tx *gorm.DB) error { + enabled, err := GetBoolByKey(ctx, ConfigKeyRegistrationEnabled) + if err == nil && !enabled { + return errors.New("注册已关闭") + } + + // 检查用户名冲突 + var count int64 + if err := tx.Model(&User{}).Where("username = ?", u.Username).Count(&count).Error; err != nil { + return err + } + if count > 0 { + return errors.New("用户名已存在") + } + + // 检查邮箱冲突 + if u.Email != "" { + var emailCount int64 + if err := tx.Model(&User{}).Where("email = ?", u.Email).Count(&emailCount).Error; err != nil { + return err + } + if emailCount > 0 { + return errors.New("该邮箱已被其他账号绑定") + } + } + + if err := tx.Create(u).Error; err != nil { + return err + } + return nil +}