refactor(user): separate Handler and Logic layer boundaries

Move HTTP handlers out of logics.go and replace gin.Context-coupled
login email verification with context-only processLoginEmailVerification.
Add logics_test.go for pure business logic unit tests.
This commit is contained in:
ryan
2026-06-18 10:31:09 +08:00
parent f826807cf1
commit 50c45db561
3 changed files with 332 additions and 141 deletions
+83 -139
View File
@@ -3,30 +3,48 @@
package user package user
import ("context" import (
"context"
"crypto/rand" "crypto/rand"
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"math/big" "math/big"
"net/http"
"strings" "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/db"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task" "github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/internal/util"
pkgu "github.com/Rain-kl/Wavelet/pkg/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 { const (
Email string `json:"email" binding:"required,email"` // LoginEmailVerificationPassed 验证通过,可继续登录流程。
Scene string `json:"scene" binding:"required"` 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 { func isEmailLoginVerificationEnabled(ctx context.Context) bool {
@@ -136,21 +154,22 @@ func verifyEmailCode(ctx context.Context, email, scene, code string) bool {
return true return true
} }
func handleLoginEmailVerification(ctx context.Context, c *gin.Context, req *loginRequest, user *model.User) error { func processLoginEmailVerification(ctx context.Context, code string, user *model.User) (LoginEmailVerificationResult, error) {
if req.Code != "" { if code != "" {
if !verifyEmailCode(ctx, user.Email, "login", req.Code) { if !verifyEmailCode(ctx, user.Email, "login", code) {
c.JSON(http.StatusOK, response.Err(errEmailCodeInvalidOrExpired)) return LoginEmailVerificationResult{
return errors.New("handled") Status: LoginEmailVerificationRejected,
Message: errEmailCodeInvalidOrExpired,
}, nil
} }
return nil return LoginEmailVerificationResult{Status: LoginEmailVerificationPassed}, nil
} }
// 如果 SMTP 未配置,或者用户没有绑定邮箱(无法发送验证码),则使用临时码 888888 // 如果 SMTP 未配置,或者用户没有绑定邮箱(无法发送验证码),则使用临时码 888888
if !isSMTPConfigured(ctx) || user.Email == "" { if !isSMTPConfigured(ctx) || user.Email == "" {
codeKey := getEmailCodeKey("login", user.Email) codeKey := getEmailCodeKey("login", user.Email)
if err := db.SetJSON(ctx, codeKey, "888888", emailCodeExpiry); err != nil { if err := db.SetJSON(ctx, codeKey, "888888", emailCodeExpiry); err != nil {
c.JSON(http.StatusOK, response.Err(errGenerateEmailCodeFailed)) return LoginEmailVerificationResult{}, errors.New(errGenerateEmailCodeFailed)
return errors.New("handled")
} }
var msg string var msg string
if !isSMTPConfigured(ctx) { if !isSMTPConfigured(ctx) {
@@ -158,173 +177,98 @@ func handleLoginEmailVerification(ctx context.Context, c *gin.Context, req *logi
} else { } else {
msg = errSMTPInvalidUseTempCodePrefix + "该账号未绑定邮箱,使用临时码登录" msg = errSMTPInvalidUseTempCodePrefix + "该账号未绑定邮箱,使用临时码登录"
} }
c.JSON(http.StatusOK, response.Err(msg)) return LoginEmailVerificationResult{
return errors.New("handled") Status: LoginEmailVerificationRejected,
Message: msg,
}, nil
} }
cooldownKey := getEmailCooldownKey("login", user.Email) cooldownKey := getEmailCooldownKey("login", user.Email)
var temp string var temp string
err := db.GetJSON(ctx, cooldownKey, &temp) if err := db.GetJSON(ctx, cooldownKey, &temp); err != nil {
if err != nil {
if err := sendEmailVerificationCode(ctx, user.Email, "login", "login_email"); err != nil { if err := sendEmailVerificationCode(ctx, user.Email, "login", "login_email"); err != nil {
c.JSON(http.StatusOK, response.Err(err.Error())) return LoginEmailVerificationResult{}, err
return errors.New("handled")
} }
} }
maskedEmail := pkgu.MaskEmail(user.Email) maskedEmail := pkgu.MaskEmail(user.Email)
c.JSON(http.StatusOK, response.Err(errNeedEmailCodePrefix+maskedEmail)) return LoginEmailVerificationResult{
return errors.New("handled") Status: LoginEmailVerificationPending,
Message: errNeedEmailCodePrefix + maskedEmail,
}, nil
} }
// SendEmailCode 发送邮箱验证码 func sendRegisterEmailCode(ctx context.Context, email string) error {
// @Summary 发送邮箱验证码 email = strings.TrimSpace(email)
// @Description 向指定邮箱发送验证码(用于注册场景) if email == "" {
// @Tags user return errors.New(errEmailRequired)
// @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()
var count int64 var count int64
if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", req.Email).Count(&count).Error; err != nil { if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", email).Count(&count).Error; err != nil {
c.JSON(http.StatusOK, response.Err(err.Error())) return err
return
} }
if count > 0 { if count > 0 {
c.JSON(http.StatusOK, response.Err(errEmailAlreadyRegistered)) return errors.New(errEmailAlreadyRegistered)
return
} }
cooldownKey := getEmailCooldownKey("register", req.Email) cooldownKey := getEmailCooldownKey("register", email)
var temp string var temp string
err := db.GetJSON(ctx, cooldownKey, &temp) if err := db.GetJSON(ctx, cooldownKey, &temp); err == nil {
if err == nil { return errors.New(errEmailCodeCooldown)
c.JSON(http.StatusOK, response.Err(errEmailCodeCooldown))
return
} }
if err := sendEmailVerificationCode(ctx, req.Email, "register", "register_email"); err != nil { return sendEmailVerificationCode(ctx, email, "register", "register_email")
c.JSON(http.StatusOK, response.Err(err.Error()))
return
}
c.JSON(http.StatusOK, response.OKNil())
} }
func validateRegisterEmailVerification(ctx context.Context, req *registerRequest) error { func validateRegisterEmailVerification(ctx context.Context, email, code string) error {
if !isEmailRegisterVerificationEnabled(ctx) { if !isEmailRegisterVerificationEnabled(ctx) {
return nil return nil
} }
if req.Email == "" || req.Code == "" { if email == "" || code == "" {
return errors.New(errEmailOrCodeRequired) return errors.New(errEmailOrCodeRequired)
} }
if !verifyEmailCode(ctx, req.Email, "register", req.Code) { if !verifyEmailCode(ctx, email, "register", code) {
return errors.New(errEmailCodeInvalidOrExpired) return errors.New(errEmailCodeInvalidOrExpired)
} }
return nil return nil
} }
type updateProfileRequest struct { func updateUserProfile(ctx context.Context, userID uint64, input updateProfileInput) (*model.User, error) {
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()
var dbUser model.User var dbUser model.User
if err := db.DB(ctx).Where("id = ?", userObj.ID).First(&dbUser).Error; err != nil { if err := db.DB(ctx).Where("id = ?", userID).First(&dbUser).Error; err != nil {
c.JSON(http.StatusOK, response.Err(errUserNotFound)) return nil, errors.New(errUserNotFound)
return
} }
req.Email = strings.TrimSpace(req.Email) input.Email = strings.TrimSpace(input.Email)
if req.Email != "" && req.Email != dbUser.Email { if input.Email != "" && input.Email != dbUser.Email {
if !strings.Contains(req.Email, "@") || !strings.Contains(req.Email, ".") { if !strings.Contains(input.Email, "@") || !strings.Contains(input.Email, ".") {
c.JSON(http.StatusOK, response.Err(errEmailFormatInvalid)) return nil, errors.New(errEmailFormatInvalid)
return
} }
var count int64 var count int64
if err := db.DB(ctx).Model(&model.User{}).Where("email = ? AND id != ?", req.Email, dbUser.ID).Count(&count).Error; err != nil { if err := db.DB(ctx).Model(&model.User{}).Where("email = ? AND id != ?", input.Email, dbUser.ID).Count(&count).Error; err != nil {
c.JSON(http.StatusOK, response.Err(err.Error())) return nil, err
return
} }
if count > 0 { if count > 0 {
c.JSON(http.StatusOK, response.Err(errEmailAlreadyBound)) return nil, errors.New(errEmailAlreadyBound)
return
} }
} }
dbUser.Nickname = strings.TrimSpace(req.Nickname) dbUser.Nickname = strings.TrimSpace(input.Nickname)
if dbUser.Nickname == "" { if dbUser.Nickname == "" {
dbUser.Nickname = dbUser.Username dbUser.Nickname = dbUser.Username
} }
dbUser.Email = req.Email dbUser.Email = input.Email
dbUser.AvatarURL = req.AvatarURL dbUser.AvatarURL = input.AvatarURL
dbUser.Bio = req.Bio dbUser.Bio = input.Bio
dbUser.Phone = strings.TrimSpace(req.Phone) dbUser.Phone = strings.TrimSpace(input.Phone)
dbUser.Gender = strings.TrimSpace(req.Gender) dbUser.Gender = strings.TrimSpace(input.Gender)
dbUser.Website = strings.TrimSpace(req.Website) dbUser.Website = strings.TrimSpace(input.Website)
dbUser.Location = strings.TrimSpace(req.Location) dbUser.Location = strings.TrimSpace(input.Location)
if err := db.DB(ctx).Save(&dbUser).Error; err != nil { if err := db.DB(ctx).Save(&dbUser).Error; err != nil {
c.JSON(http.StatusOK, response.Err(err.Error())) return nil, err
return
} }
return &dbUser, nil
session := sessions.Default(c) }
needChange := session.Get("need_change_password") == true
c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(&dbUser, needChange)))
}
+151
View File
@@ -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)
}
}
+98 -2
View File
@@ -37,6 +37,22 @@ type registerRequest struct {
Code string `json:"code"` 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 { func isPasswordLoginEnabled() bool {
enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyPasswordLoginEnabled) enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyPasswordLoginEnabled)
if err != nil { if err != nil {
@@ -146,7 +162,13 @@ func Login(c *gin.Context) {
} }
if isEmailLoginVerificationEnabled(ctx) { 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 return
} }
} }
@@ -223,7 +245,7 @@ func Register(c *gin.Context) {
ctx := c.Request.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())) c.JSON(http.StatusOK, response.Err(err.Error()))
return return
} }
@@ -365,3 +387,77 @@ func ChangePassword(c *gin.Context) {
c.JSON(http.StatusOK, response.OK("密码修改成功")) 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)))
}