diff --git a/internal/apps/user/logics.go b/internal/apps/user/logics.go index 02923f41..37fc259e 100644 --- a/internal/apps/user/logics.go +++ b/internal/apps/user/logics.go @@ -3,30 +3,48 @@ package user -import ("context" +import ( + "context" "crypto/rand" "encoding/json" "errors" "fmt" "math/big" - "net/http" "strings" - "github.com/gin-contrib/sessions" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/task" - "github.com/Rain-kl/Wavelet/internal/util" pkgu "github.com/Rain-kl/Wavelet/pkg/util" - "github.com/gin-gonic/gin" +) - "github.com/Rain-kl/Wavelet/internal/common/response") +// LoginEmailVerificationStatus 登录邮箱验证的处理结果。 +type LoginEmailVerificationStatus int -type sendEmailCodeRequest struct { - Email string `json:"email" binding:"required,email"` - Scene string `json:"scene" binding:"required"` +const ( + // LoginEmailVerificationPassed 验证通过,可继续登录流程。 + LoginEmailVerificationPassed LoginEmailVerificationStatus = iota + // LoginEmailVerificationPending 需要用户输入邮箱验证码。 + LoginEmailVerificationPending + // LoginEmailVerificationRejected 验证被拒绝(验证码错误、临时码提示等)。 + LoginEmailVerificationRejected +) + +// LoginEmailVerificationResult 登录邮箱验证的业务结果。 +type LoginEmailVerificationResult struct { + Status LoginEmailVerificationStatus + Message string +} + +type updateProfileInput struct { + Nickname string + Email string + AvatarURL string + Bio string + Phone string + Gender string + Website string + Location string } func isEmailLoginVerificationEnabled(ctx context.Context) bool { @@ -136,21 +154,22 @@ func verifyEmailCode(ctx context.Context, email, scene, code string) bool { return true } -func handleLoginEmailVerification(ctx context.Context, c *gin.Context, req *loginRequest, user *model.User) error { - if req.Code != "" { - if !verifyEmailCode(ctx, user.Email, "login", req.Code) { - c.JSON(http.StatusOK, response.Err(errEmailCodeInvalidOrExpired)) - return errors.New("handled") +func processLoginEmailVerification(ctx context.Context, code string, user *model.User) (LoginEmailVerificationResult, error) { + if code != "" { + if !verifyEmailCode(ctx, user.Email, "login", code) { + return LoginEmailVerificationResult{ + Status: LoginEmailVerificationRejected, + Message: errEmailCodeInvalidOrExpired, + }, nil } - return nil + return LoginEmailVerificationResult{Status: LoginEmailVerificationPassed}, nil } // 如果 SMTP 未配置,或者用户没有绑定邮箱(无法发送验证码),则使用临时码 888888 if !isSMTPConfigured(ctx) || user.Email == "" { codeKey := getEmailCodeKey("login", user.Email) if err := db.SetJSON(ctx, codeKey, "888888", emailCodeExpiry); err != nil { - c.JSON(http.StatusOK, response.Err(errGenerateEmailCodeFailed)) - return errors.New("handled") + return LoginEmailVerificationResult{}, errors.New(errGenerateEmailCodeFailed) } var msg string if !isSMTPConfigured(ctx) { @@ -158,173 +177,98 @@ func handleLoginEmailVerification(ctx context.Context, c *gin.Context, req *logi } else { msg = errSMTPInvalidUseTempCodePrefix + "该账号未绑定邮箱,使用临时码登录" } - c.JSON(http.StatusOK, response.Err(msg)) - return errors.New("handled") + return LoginEmailVerificationResult{ + Status: LoginEmailVerificationRejected, + Message: msg, + }, nil } cooldownKey := getEmailCooldownKey("login", user.Email) var temp string - err := db.GetJSON(ctx, cooldownKey, &temp) - if err != nil { + if err := db.GetJSON(ctx, cooldownKey, &temp); err != nil { if err := sendEmailVerificationCode(ctx, user.Email, "login", "login_email"); err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) - return errors.New("handled") + return LoginEmailVerificationResult{}, err } } maskedEmail := pkgu.MaskEmail(user.Email) - c.JSON(http.StatusOK, response.Err(errNeedEmailCodePrefix+maskedEmail)) - return errors.New("handled") + return LoginEmailVerificationResult{ + Status: LoginEmailVerificationPending, + Message: errNeedEmailCodePrefix + maskedEmail, + }, nil } -// SendEmailCode 发送邮箱验证码 -// @Summary 发送邮箱验证码 -// @Description 向指定邮箱发送验证码(用于注册场景) -// @Tags user -// @Accept json -// @Produce json -// @Param request body user.sendEmailCodeRequest true "发送验证码请求参数" -// @Success 200 {object} response.Any "发送成功" -// @Failure 400 {object} response.Any "参数错误" -// @Router /api/v1/user/send-email-code [post] -func SendEmailCode(c *gin.Context) { - var req sendEmailCodeRequest - if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) - return +func sendRegisterEmailCode(ctx context.Context, email string) error { + email = strings.TrimSpace(email) + if email == "" { + return errors.New(errEmailRequired) } - req.Email = strings.TrimSpace(req.Email) - if req.Email == "" { - c.JSON(http.StatusOK, response.Err(errEmailRequired)) - return - } - - if req.Scene != "register" { - c.JSON(http.StatusOK, response.Err(errUnsupportedEmailScene)) - return - } - - ctx := c.Request.Context() - var count int64 - if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", req.Email).Count(&count).Error; err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) - return + if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", email).Count(&count).Error; err != nil { + return err } if count > 0 { - c.JSON(http.StatusOK, response.Err(errEmailAlreadyRegistered)) - return + return errors.New(errEmailAlreadyRegistered) } - cooldownKey := getEmailCooldownKey("register", req.Email) + cooldownKey := getEmailCooldownKey("register", email) var temp string - err := db.GetJSON(ctx, cooldownKey, &temp) - if err == nil { - c.JSON(http.StatusOK, response.Err(errEmailCodeCooldown)) - return + if err := db.GetJSON(ctx, cooldownKey, &temp); err == nil { + return errors.New(errEmailCodeCooldown) } - if err := sendEmailVerificationCode(ctx, req.Email, "register", "register_email"); err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) - return - } - - c.JSON(http.StatusOK, response.OKNil()) + return sendEmailVerificationCode(ctx, email, "register", "register_email") } -func validateRegisterEmailVerification(ctx context.Context, req *registerRequest) error { +func validateRegisterEmailVerification(ctx context.Context, email, code string) error { if !isEmailRegisterVerificationEnabled(ctx) { return nil } - if req.Email == "" || req.Code == "" { + if email == "" || code == "" { return errors.New(errEmailOrCodeRequired) } - if !verifyEmailCode(ctx, req.Email, "register", req.Code) { + if !verifyEmailCode(ctx, email, "register", code) { return errors.New(errEmailCodeInvalidOrExpired) } return nil } -type updateProfileRequest struct { - Nickname string `json:"nickname"` - Email string `json:"email"` - AvatarURL string `json:"avatar_url"` - Bio string `json:"bio"` - Phone string `json:"phone"` - Gender string `json:"gender"` - Website string `json:"website"` - Location string `json:"location"` -} - -// UpdateProfile 修改当前登录用户的个人资料 -// @Summary 修改当前登录用户的个人资料 -// @Description 修改当前登录用户的昵称、邮箱、头像、简介、电话、性别、个人网站和所在地。 -// @Tags user -// @Accept json -// @Produce json -// @Param request body user.updateProfileRequest true "更新请求参数" -// @Success 200 {object} response.Any{data=oauth.BasicUserInfo} "修改成功,返回更新后的用户信息" -// @Failure 400 {object} response.Any "邮箱已被占用或参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Router /api/v1/user/profile [put] -func UpdateProfile(c *gin.Context) { - var req updateProfileRequest - if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) - return - } - - userObj, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) - if userObj == nil { - c.JSON(http.StatusUnauthorized, response.Err(errLoginRequired)) - return - } - - ctx := c.Request.Context() +func updateUserProfile(ctx context.Context, userID uint64, input updateProfileInput) (*model.User, error) { var dbUser model.User - if err := db.DB(ctx).Where("id = ?", userObj.ID).First(&dbUser).Error; err != nil { - c.JSON(http.StatusOK, response.Err(errUserNotFound)) - return + if err := db.DB(ctx).Where("id = ?", userID).First(&dbUser).Error; err != nil { + return nil, errors.New(errUserNotFound) } - req.Email = strings.TrimSpace(req.Email) - if req.Email != "" && req.Email != dbUser.Email { - if !strings.Contains(req.Email, "@") || !strings.Contains(req.Email, ".") { - c.JSON(http.StatusOK, response.Err(errEmailFormatInvalid)) - return + input.Email = strings.TrimSpace(input.Email) + if input.Email != "" && input.Email != dbUser.Email { + if !strings.Contains(input.Email, "@") || !strings.Contains(input.Email, ".") { + return nil, errors.New(errEmailFormatInvalid) } var count int64 - if err := db.DB(ctx).Model(&model.User{}).Where("email = ? AND id != ?", req.Email, dbUser.ID).Count(&count).Error; err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) - return + if err := db.DB(ctx).Model(&model.User{}).Where("email = ? AND id != ?", input.Email, dbUser.ID).Count(&count).Error; err != nil { + return nil, err } if count > 0 { - c.JSON(http.StatusOK, response.Err(errEmailAlreadyBound)) - return + return nil, errors.New(errEmailAlreadyBound) } } - dbUser.Nickname = strings.TrimSpace(req.Nickname) + dbUser.Nickname = strings.TrimSpace(input.Nickname) if dbUser.Nickname == "" { dbUser.Nickname = dbUser.Username } - dbUser.Email = req.Email - dbUser.AvatarURL = req.AvatarURL - dbUser.Bio = req.Bio - dbUser.Phone = strings.TrimSpace(req.Phone) - dbUser.Gender = strings.TrimSpace(req.Gender) - dbUser.Website = strings.TrimSpace(req.Website) - dbUser.Location = strings.TrimSpace(req.Location) + dbUser.Email = input.Email + dbUser.AvatarURL = input.AvatarURL + dbUser.Bio = input.Bio + dbUser.Phone = strings.TrimSpace(input.Phone) + dbUser.Gender = strings.TrimSpace(input.Gender) + dbUser.Website = strings.TrimSpace(input.Website) + dbUser.Location = strings.TrimSpace(input.Location) if err := db.DB(ctx).Save(&dbUser).Error; err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) - return + return nil, err } - - session := sessions.Default(c) - needChange := session.Get("need_change_password") == true - - c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(&dbUser, needChange))) -} + return &dbUser, nil +} \ No newline at end of file diff --git a/internal/apps/user/logics_test.go b/internal/apps/user/logics_test.go new file mode 100644 index 00000000..127cbc6f --- /dev/null +++ b/internal/apps/user/logics_test.go @@ -0,0 +1,151 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package user + +import ( + "context" + "testing" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/testhelper" +) + +func TestProcessLoginEmailVerificationSMTPFallback(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + + const email = "smtpuser@example.com" + now := time.Now() + user := model.User{ + ID: 222, + Username: "smtpuser", + Nickname: "SMTP User", + Email: email, + IsActive: true, + LastLoginAt: now, + } + if err := user.SetEncryptedPassword("newpassword123"); err != nil { + t.Fatalf("set encrypted password failed: %v", err) + } + if err := dbConn.Create(&user).Error; err != nil { + t.Fatalf("create test user failed: %v", err) + } + + if err := dbConn.Model(&model.SystemConfig{}). + Where("key = ?", model.ConfigKeySMTPHost). + Update("value", "").Error; err != nil { + t.Fatalf("clear SMTP host failed: %v", err) + } + if err := db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err(); err != nil { + t.Fatalf("invalidate system config cache failed: %v", err) + } + + ctx := context.Background() + result, err := processLoginEmailVerification(ctx, "", &user) + if err != nil { + t.Fatalf("processLoginEmailVerification() error = %v, want nil", err) + } + expected := errSMTPInvalidUseTempCodePrefix + errSMTPInvalidUseTempCode + if result.Status != LoginEmailVerificationRejected || result.Message != expected { + t.Fatalf("processLoginEmailVerification() = %+v, want rejected with %q", result, expected) + } + + codeKey := getEmailCodeKey("login", email) + var storedCode string + if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil { + t.Fatalf("get stored verification code failed: %v", err) + } + if storedCode != "888888" { + t.Errorf("stored verification code = %q, want %q", storedCode, "888888") + } + + passed, err := processLoginEmailVerification(ctx, "888888", &user) + if err != nil { + t.Fatalf("processLoginEmailVerification(valid code) error = %v, want nil", err) + } + if passed.Status != LoginEmailVerificationPassed { + t.Fatalf("processLoginEmailVerification(valid code) status = %v, want passed", passed.Status) + } +} + +func TestProcessLoginEmailVerificationEmptyEmailFallback(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + + now := time.Now() + user := model.User{ + ID: 223, + Username: "emptyemailuser", + Nickname: "Empty Email User", + Email: "", + IsActive: true, + LastLoginAt: now, + } + if err := user.SetEncryptedPassword("newpassword123"); err != nil { + t.Fatalf("set encrypted password failed: %v", err) + } + if err := dbConn.Create(&user).Error; err != nil { + t.Fatalf("create test user failed: %v", err) + } + + for _, cfg := range []struct { + key string + value string + }{ + {model.ConfigKeySMTPHost, "smtp.example.com"}, + {model.ConfigKeySMTPPort, "587"}, + {model.ConfigKeySMTPUsername, "smtpuser"}, + {model.ConfigKeySMTPPassword, "smtppassword"}, + } { + if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", cfg.key).Update("value", cfg.value).Error; err != nil { + t.Fatalf("set %s failed: %v", cfg.key, err) + } + } + if err := db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err(); err != nil { + t.Fatalf("invalidate system config cache failed: %v", err) + } + + ctx := context.Background() + result, err := processLoginEmailVerification(ctx, "", &user) + if err != nil { + t.Fatalf("processLoginEmailVerification() error = %v, want nil", err) + } + expected := errSMTPInvalidUseTempCodePrefix + "该账号未绑定邮箱,使用临时码登录" + if result.Status != LoginEmailVerificationRejected || result.Message != expected { + t.Fatalf("processLoginEmailVerification() = %+v, want rejected with %q", result, expected) + } +} + +func TestProcessLoginEmailVerificationInvalidCode(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + + const email = "codeduser@example.com" + now := time.Now() + user := model.User{ + ID: 224, + Username: "codeduser", + Email: email, + IsActive: true, + LastLoginAt: now, + } + if err := dbConn.Create(&user).Error; err != nil { + t.Fatalf("create test user failed: %v", err) + } + + ctx := context.Background() + if err := db.SetJSON(ctx, getEmailCodeKey("login", email), "123456", emailCodeExpiry); err != nil { + t.Fatalf("seed verification code failed: %v", err) + } + + result, err := processLoginEmailVerification(ctx, "000000", &user) + if err != nil { + t.Fatalf("processLoginEmailVerification() error = %v, want nil", err) + } + if result.Status != LoginEmailVerificationRejected || result.Message != errEmailCodeInvalidOrExpired { + t.Fatalf("processLoginEmailVerification() = %+v, want rejected invalid code", result) + } +} \ No newline at end of file diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go index fad35df2..ff595f45 100644 --- a/internal/apps/user/routers.go +++ b/internal/apps/user/routers.go @@ -37,6 +37,22 @@ type registerRequest struct { Code string `json:"code"` } +type sendEmailCodeRequest struct { + Email string `json:"email" binding:"required,email"` + Scene string `json:"scene" binding:"required"` +} + +type updateProfileRequest struct { + Nickname string `json:"nickname"` + Email string `json:"email"` + AvatarURL string `json:"avatar_url"` + Bio string `json:"bio"` + Phone string `json:"phone"` + Gender string `json:"gender"` + Website string `json:"website"` + Location string `json:"location"` +} + func isPasswordLoginEnabled() bool { enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyPasswordLoginEnabled) if err != nil { @@ -146,7 +162,13 @@ func Login(c *gin.Context) { } if isEmailLoginVerificationEnabled(ctx) { - if emailErr := handleLoginEmailVerification(ctx, c, &req, &user); emailErr != nil { + result, err := processLoginEmailVerification(ctx, req.Code, &user) + if err != nil { + c.JSON(http.StatusOK, response.Err(err.Error())) + return + } + if result.Status != LoginEmailVerificationPassed { + c.JSON(http.StatusOK, response.Err(result.Message)) return } } @@ -223,7 +245,7 @@ func Register(c *gin.Context) { ctx := c.Request.Context() // 邮箱注册验证校验 - if err := validateRegisterEmailVerification(ctx, &req); err != nil { + if err := validateRegisterEmailVerification(ctx, req.Email, req.Code); err != nil { c.JSON(http.StatusOK, response.Err(err.Error())) return } @@ -365,3 +387,77 @@ func ChangePassword(c *gin.Context) { c.JSON(http.StatusOK, response.OK("密码修改成功")) } + +// SendEmailCode 发送邮箱验证码 +// @Summary 发送邮箱验证码 +// @Description 向指定邮箱发送验证码(用于注册场景) +// @Tags user +// @Accept json +// @Produce json +// @Param request body user.sendEmailCodeRequest true "发送验证码请求参数" +// @Success 200 {object} response.Any "发送成功" +// @Failure 400 {object} response.Any "参数错误" +// @Router /api/v1/user/send-email-code [post] +func SendEmailCode(c *gin.Context) { + var req sendEmailCodeRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, response.Err(err.Error())) + return + } + + req.Email = strings.TrimSpace(req.Email) + if req.Email == "" { + c.JSON(http.StatusOK, response.Err(errEmailRequired)) + return + } + + if req.Scene != "register" { + c.JSON(http.StatusOK, response.Err(errUnsupportedEmailScene)) + return + } + + ctx := c.Request.Context() + if err := sendRegisterEmailCode(ctx, req.Email); err != nil { + c.JSON(http.StatusOK, response.Err(err.Error())) + return + } + + c.JSON(http.StatusOK, response.OKNil()) +} + +// UpdateProfile 修改当前登录用户的个人资料 +// @Summary 修改当前登录用户的个人资料 +// @Description 修改当前登录用户的昵称、邮箱、头像、简介、电话、性别、个人网站和所在地。 +// @Tags user +// @Accept json +// @Produce json +// @Param request body user.updateProfileRequest true "更新请求参数" +// @Success 200 {object} response.Any{data=oauth.BasicUserInfo} "修改成功,返回更新后的用户信息" +// @Failure 400 {object} response.Any "邮箱已被占用或参数错误" +// @Failure 401 {object} response.Any "未登录" +// @Router /api/v1/user/profile [put] +func UpdateProfile(c *gin.Context) { + var req updateProfileRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, response.Err(err.Error())) + return + } + + userObj, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) + if userObj == nil { + c.JSON(http.StatusUnauthorized, response.Err(errLoginRequired)) + return + } + + ctx := c.Request.Context() + dbUser, err := updateUserProfile(ctx, userObj.ID, updateProfileInput(req)) + if err != nil { + c.JSON(http.StatusOK, response.Err(err.Error())) + return + } + + session := sessions.Default(c) + needChange := session.Get("need_change_password") == true + + c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(dbUser, needChange))) +}