mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-04 07:06:36 +08:00
邮箱注册要求
This commit is contained in:
@@ -14,6 +14,7 @@ import (
|
||||
"strconv"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common"
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/util"
|
||||
@@ -171,19 +172,32 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro
|
||||
session.Set(UserNameKey, user.Username)
|
||||
|
||||
// 根据系统配置动态设置 Session 过期时间
|
||||
maxAge := 0
|
||||
maxAge := config.Config.App.SessionAge
|
||||
isSessionCookie := false
|
||||
|
||||
ttlHours, err := model.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
|
||||
if err == nil {
|
||||
if ttlHours == -1 {
|
||||
switch {
|
||||
case ttlHours == -1:
|
||||
// 永不过期,设置为 10 年
|
||||
maxAge = 10 * 365 * 24 * 3600
|
||||
} else if ttlHours > 0 {
|
||||
case ttlHours > 0:
|
||||
maxAge = ttlHours * 3600
|
||||
case ttlHours == 0:
|
||||
isSessionCookie = true
|
||||
}
|
||||
}
|
||||
session.Options(util.GetSessionOptions(maxAge))
|
||||
|
||||
return session.Save()
|
||||
if err := session.Save(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if isSessionCookie {
|
||||
util.StripCookieMaxAgeAndExpires(c.Writer.Header(), config.Config.App.SessionCookieName)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func uniqueUsername(ctx context.Context, base string) (string, error) {
|
||||
|
||||
+39
-38
@@ -5,42 +5,43 @@
|
||||
package user
|
||||
|
||||
const (
|
||||
errBindParamsFailed = "参数绑定失败"
|
||||
errInvalidParams = "无效的参数"
|
||||
errPasswordLoginDisabled = "管理员关闭了密码登录" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errUsernameOrPasswordWrong = "用户名或密码错误" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errLoginEmailMissing = "该账号未绑定邮箱,请联系管理员绑定邮箱后再登录"
|
||||
errNeedEmailCodePrefix = "need_email_code:"
|
||||
errEmailCodeInvalidOrExpired = "验证码错误或已过期"
|
||||
errPasswordUpgradeFailed = "升级密码安全算法失败,请重试" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errSaveSessionFailed = "无法保存会话信息,请重试"
|
||||
errRegistrationDisabled = "管理员关闭了注册"
|
||||
errPasswordTooShort = "密码长度不能少于 8 位" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errEmailOrCodeRequired = "邮箱或验证码未填写"
|
||||
errNewPasswordTooShort = "新密码长度不能少于 8 位" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errLoginRequired = "请先登录"
|
||||
errUserNotFound = "未找到该用户"
|
||||
errOldPasswordIncorrect = "原密码不正确" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errPasswordEncryptFailed = "密码加密失败,请重试" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errEmailRequired = "邮箱地址不能为空"
|
||||
errUnsupportedEmailScene = "不支持的验证场景"
|
||||
errEmailAlreadyRegistered = "该邮箱已被注册"
|
||||
errEmailCodeCooldown = "验证码发送频繁,请稍后再试"
|
||||
errEmailFormatInvalid = "邮箱格式不正确"
|
||||
errEmailAlreadyBound = "该邮箱已被其他账号绑定"
|
||||
errRenderEmailTemplateFailed = "渲染验证邮件模板失败:%w"
|
||||
errGenerateEmailCodeFailed = "生成验证码失败,请重试"
|
||||
errDispatchEmailTaskFailed = "投递验证邮件发送任务失败,请重试"
|
||||
errTokenNameRequired = "令牌名称不能为空" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errAccessTokenLimitReached = "已达到访问令牌最大创建数量限制" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errGenerateTokenFailed = "生成令牌失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errInvalidTokenID = "无效的令牌ID" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errTokenNotFoundOrForbidden = "令牌不存在或无权操作" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errAdminTokenRequiresAdmin = "只有管理员才能创建具有管理员权限的令牌" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errTaskPayloadRequired = "任务参数不能为空"
|
||||
errInvalidJSONFormat = "无效的 JSON 格式: %w"
|
||||
errEmailTaskFieldsRequired = "to、subject、body 不能为空"
|
||||
errParseEmailPayloadFailed = "解析邮件发送参数失败: %w"
|
||||
errSMTPConfigIncomplete = "系统 SMTP 邮件服务配置不完整"
|
||||
errSendMailFailed = "发送邮件失败: %w"
|
||||
errBindParamsFailed = "参数绑定失败"
|
||||
errInvalidParams = "无效的参数"
|
||||
errPasswordLoginDisabled = "管理员关闭了密码登录" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errUsernameOrPasswordWrong = "用户名或密码错误" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errLoginEmailMissing = "该账号未绑定邮箱,请联系管理员绑定邮箱后再登录"
|
||||
errNeedEmailCodePrefix = "need_email_code:"
|
||||
errSMTPInvalidUseTempCodePrefix = "smtp_invalid:"
|
||||
errSMTPInvalidUseTempCode = "smtp 配置无效,使用临时码登录"
|
||||
errEmailCodeInvalidOrExpired = "验证码错误或已过期"
|
||||
errSaveSessionFailed = "无法保存会话信息,请重试"
|
||||
errRegistrationDisabled = "管理员关闭了注册"
|
||||
errPasswordTooShort = "密码长度不能少于 8 位" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errEmailOrCodeRequired = "邮箱或验证码未填写"
|
||||
errNewPasswordTooShort = "新密码长度不能少于 8 位" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errLoginRequired = "请先登录"
|
||||
errUserNotFound = "未找到该用户"
|
||||
errOldPasswordIncorrect = "原密码不正确" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errPasswordEncryptFailed = "密码加密失败,请重试" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errEmailRequired = "邮箱地址不能为空"
|
||||
errUnsupportedEmailScene = "不支持的验证场景"
|
||||
errEmailAlreadyRegistered = "该邮箱已被注册"
|
||||
errEmailCodeCooldown = "验证码发送频繁,请稍后再试"
|
||||
errEmailFormatInvalid = "邮箱格式不正确"
|
||||
errEmailAlreadyBound = "该邮箱已被其他账号绑定"
|
||||
errRenderEmailTemplateFailed = "渲染验证邮件模板失败:%w"
|
||||
errGenerateEmailCodeFailed = "生成验证码失败,请重试"
|
||||
errDispatchEmailTaskFailed = "投递验证邮件发送任务失败,请重试"
|
||||
errTokenNameRequired = "令牌名称不能为空" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errAccessTokenLimitReached = "已达到访问令牌最大创建数量限制" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errGenerateTokenFailed = "生成令牌失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errInvalidTokenID = "无效的令牌ID" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errTokenNotFoundOrForbidden = "令牌不存在或无权操作" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errAdminTokenRequiresAdmin = "只有管理员才能创建具有管理员权限的令牌" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errTaskPayloadRequired = "任务参数不能为空"
|
||||
errInvalidJSONFormat = "无效的 JSON 格式: %w"
|
||||
errEmailTaskFieldsRequired = "to、subject、body 不能为空"
|
||||
errParseEmailPayloadFailed = "解析邮件发送参数失败: %w"
|
||||
errSMTPConfigIncomplete = "系统 SMTP 邮件服务配置不完整"
|
||||
errSendMailFailed = "发送邮件失败: %w"
|
||||
)
|
||||
|
||||
@@ -45,17 +45,19 @@ func isEmailRegisterVerificationEnabled(ctx context.Context) bool {
|
||||
}
|
||||
|
||||
func isSMTPConfigured(ctx context.Context) bool {
|
||||
var sc model.SystemConfig
|
||||
var host, port, username string
|
||||
|
||||
if err := sc.GetByKey(ctx, model.ConfigKeySMTPHost); err == nil {
|
||||
host = sc.Value
|
||||
var scHost model.SystemConfig
|
||||
if err := scHost.GetByKey(ctx, model.ConfigKeySMTPHost); err == nil {
|
||||
host = scHost.Value
|
||||
}
|
||||
if err := sc.GetByKey(ctx, model.ConfigKeySMTPPort); err == nil {
|
||||
port = sc.Value
|
||||
var scPort model.SystemConfig
|
||||
if err := scPort.GetByKey(ctx, model.ConfigKeySMTPPort); err == nil {
|
||||
port = scPort.Value
|
||||
}
|
||||
if err := sc.GetByKey(ctx, model.ConfigKeySMTPUsername); err == nil {
|
||||
username = sc.Value
|
||||
var scUser model.SystemConfig
|
||||
if err := scUser.GetByKey(ctx, model.ConfigKeySMTPUsername); err == nil {
|
||||
username = scUser.Value
|
||||
}
|
||||
|
||||
return host != "" && port != "" && username != ""
|
||||
@@ -130,32 +132,44 @@ func verifyEmailCode(ctx context.Context, email, scene, code string) bool {
|
||||
}
|
||||
|
||||
func handleLoginEmailVerification(ctx context.Context, c *gin.Context, req *loginRequest, user *model.User) error {
|
||||
if user.Email == "" {
|
||||
c.JSON(http.StatusOK, util.Err(errLoginEmailMissing))
|
||||
return errors.New("handled")
|
||||
}
|
||||
|
||||
if req.Code == "" {
|
||||
cooldownKey := getEmailCooldownKey("login", user.Email)
|
||||
var temp string
|
||||
err := db.GetJSON(ctx, cooldownKey, &temp)
|
||||
if err != nil {
|
||||
if err := sendEmailVerificationCode(ctx, user.Email, "login", "login_email"); err != nil {
|
||||
c.JSON(http.StatusOK, util.Err(err.Error()))
|
||||
return errors.New("handled")
|
||||
}
|
||||
if req.Code != "" {
|
||||
if !verifyEmailCode(ctx, user.Email, "login", req.Code) {
|
||||
c.JSON(http.StatusOK, util.Err(errEmailCodeInvalidOrExpired))
|
||||
return errors.New("handled")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
maskedEmail := util.MaskEmail(user.Email)
|
||||
c.JSON(http.StatusOK, util.Err(errNeedEmailCodePrefix+maskedEmail))
|
||||
// 如果 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, util.Err(errGenerateEmailCodeFailed))
|
||||
return errors.New("handled")
|
||||
}
|
||||
var msg string
|
||||
if !isSMTPConfigured(ctx) {
|
||||
msg = errSMTPInvalidUseTempCodePrefix + errSMTPInvalidUseTempCode
|
||||
} else {
|
||||
msg = errSMTPInvalidUseTempCodePrefix + "该账号未绑定邮箱,使用临时码登录"
|
||||
}
|
||||
c.JSON(http.StatusOK, util.Err(msg))
|
||||
return errors.New("handled")
|
||||
}
|
||||
|
||||
if !verifyEmailCode(ctx, user.Email, "login", req.Code) {
|
||||
c.JSON(http.StatusOK, util.Err(errEmailCodeInvalidOrExpired))
|
||||
return errors.New("handled")
|
||||
cooldownKey := getEmailCooldownKey("login", user.Email)
|
||||
var temp string
|
||||
err := db.GetJSON(ctx, cooldownKey, &temp)
|
||||
if err != nil {
|
||||
if err := sendEmailVerificationCode(ctx, user.Email, "login", "login_email"); err != nil {
|
||||
c.JSON(http.StatusOK, util.Err(err.Error()))
|
||||
return errors.New("handled")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
|
||||
maskedEmail := util.MaskEmail(user.Email)
|
||||
c.JSON(http.StatusOK, util.Err(errNeedEmailCodePrefix+maskedEmail))
|
||||
return errors.New("handled")
|
||||
}
|
||||
|
||||
// SendEmailCode 发送邮箱验证码
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/Rain-kl/Wavelet/internal/common"
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
@@ -64,14 +65,19 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro
|
||||
session.Set(oauth.UserNameKey, user.Username)
|
||||
|
||||
// 根据系统配置动态设置 Session 过期时间
|
||||
maxAge := 0
|
||||
maxAge := config.Config.App.SessionAge
|
||||
isSessionCookie := false
|
||||
|
||||
ttlHours, err := model.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
|
||||
if err == nil {
|
||||
if ttlHours == -1 {
|
||||
switch {
|
||||
case ttlHours == -1:
|
||||
// 永不过期,设置为 10 年
|
||||
maxAge = 10 * 365 * 24 * 3600
|
||||
} else if ttlHours > 0 {
|
||||
case ttlHours > 0:
|
||||
maxAge = ttlHours * 3600
|
||||
case ttlHours == 0:
|
||||
isSessionCookie = true
|
||||
}
|
||||
}
|
||||
session.Options(util.GetSessionOptions(maxAge))
|
||||
@@ -79,6 +85,11 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro
|
||||
if err := session.Save(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if isSessionCookie {
|
||||
util.StripCookieMaxAgeAndExpires(c.Writer.Header(), config.Config.App.SessionCookieName)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -111,7 +122,7 @@ func Login(c *gin.Context) {
|
||||
|
||||
var user model.User
|
||||
ctx := c.Request.Context()
|
||||
if err := db.DB(ctx).Where("username = ?", req.Username).First(&user).Error; err != nil {
|
||||
if err := db.DB(ctx).Where("username = ? OR email = ?", req.Username, req.Username).First(&user).Error; err != nil {
|
||||
c.JSON(http.StatusOK, util.Err(errUsernameOrPasswordWrong))
|
||||
return
|
||||
}
|
||||
@@ -193,6 +204,10 @@ func Register(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, util.Err(errInvalidParams))
|
||||
return
|
||||
}
|
||||
if req.Email == "" {
|
||||
c.JSON(http.StatusOK, util.Err(errEmailRequired))
|
||||
return
|
||||
}
|
||||
if len(req.Password) < minPasswordLength {
|
||||
c.JSON(http.StatusOK, util.Err(errPasswordTooShort))
|
||||
return
|
||||
|
||||
@@ -5,6 +5,7 @@ package user
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -13,6 +14,7 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/Rain-kl/Wavelet/internal/util"
|
||||
@@ -141,6 +143,7 @@ func TestRegisterCreatesAuthenticatedEncryptedUser(t *testing.T) {
|
||||
Username: "newuser",
|
||||
Password: "newpassword123",
|
||||
Nickname: "New User",
|
||||
Email: "newuser@example.com",
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
|
||||
@@ -240,3 +243,203 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||
t.Errorf("UserInfo() after Login(%q) need_change_password = false, want true", adminUsername)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginEmailVerificationFallbackWhenSMTPUnconfigured(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
const (
|
||||
userID = uint64(222)
|
||||
username = "smtpuser"
|
||||
password = "newpassword123"
|
||||
email = "smtpuser@example.com"
|
||||
)
|
||||
now := time.Now()
|
||||
user := model.User{
|
||||
ID: userID,
|
||||
Username: username,
|
||||
Nickname: "SMTP User",
|
||||
Email: email,
|
||||
IsActive: true,
|
||||
IsAdmin: false,
|
||||
LastLoginAt: now,
|
||||
}
|
||||
if err := user.SetEncryptedPassword(password); 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)
|
||||
}
|
||||
|
||||
// 1. Enable email login verification
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyEmailLoginVerificationEnabled).Update("value", "true").Error; err != nil {
|
||||
t.Fatalf("enable email login verification failed: %v", err)
|
||||
}
|
||||
// 2. Clear SMTP host to simulate unconfigured SMTP
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPHost).Update("value", "").Error; err != nil {
|
||||
t.Fatalf("clear SMTP host failed: %v", err)
|
||||
}
|
||||
|
||||
// 2.5 Invalidate the system config cache in Redis
|
||||
if err := db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err(); err != nil {
|
||||
t.Fatalf("invalidate system config cache failed: %v", err)
|
||||
}
|
||||
|
||||
router := setupUserTestRouter(t)
|
||||
|
||||
// 3. Perform login request without verification code
|
||||
payload := loginRequest{
|
||||
Username: username,
|
||||
Password: password,
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
|
||||
w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("Login(%q) status = %d, want %d. Body: %s", username, w.Code, http.StatusOK, w.Body.String())
|
||||
}
|
||||
|
||||
// Check response error msg
|
||||
var resp struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("unmarshal response failed: %v", err)
|
||||
}
|
||||
expectedError := errSMTPInvalidUseTempCodePrefix + errSMTPInvalidUseTempCode
|
||||
if resp.ErrorMsg != expectedError {
|
||||
t.Errorf("expected error %q, got %q", expectedError, resp.ErrorMsg)
|
||||
}
|
||||
|
||||
// 4. Check that verification code stored in Redis is "888888"
|
||||
ctx := context.Background()
|
||||
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("expected verification code '888888', got %q", storedCode)
|
||||
}
|
||||
|
||||
// 5. Retry login with code "888888"
|
||||
payload.Code = "888888"
|
||||
bodyWithCode, _ := json.Marshal(payload)
|
||||
w = performUserRequest(router, http.MethodPost, "/api/v1/user/login", bodyWithCode, nil)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("Login with code status = %d, want %d. Body: %s", w.Code, http.StatusOK, w.Body.String())
|
||||
}
|
||||
|
||||
var successResp struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &successResp); err != nil {
|
||||
t.Fatalf("unmarshal success response failed: %v", err)
|
||||
}
|
||||
if successResp.ErrorMsg != "" {
|
||||
t.Errorf("expected login success, got error %q", successResp.ErrorMsg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginEmailVerificationFallbackForEmptyEmail(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
const (
|
||||
userID = uint64(223)
|
||||
username = "emptyemailuser"
|
||||
password = "newpassword123"
|
||||
email = ""
|
||||
)
|
||||
now := time.Now()
|
||||
user := model.User{
|
||||
ID: userID,
|
||||
Username: username,
|
||||
Nickname: "Empty Email User",
|
||||
Email: email,
|
||||
IsActive: true,
|
||||
IsAdmin: true,
|
||||
LastLoginAt: now,
|
||||
}
|
||||
if err := user.SetEncryptedPassword(password); 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)
|
||||
}
|
||||
|
||||
// 1. Enable email login verification
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyEmailLoginVerificationEnabled).Update("value", "true").Error; err != nil {
|
||||
t.Fatalf("enable email login verification failed: %v", err)
|
||||
}
|
||||
// 2. Make sure SMTP is configured so we only trigger empty email check
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPHost).Update("value", "smtp.example.com").Error; err != nil {
|
||||
t.Fatalf("set SMTP host failed: %v", err)
|
||||
}
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPPort).Update("value", "587").Error; err != nil {
|
||||
t.Fatalf("set SMTP port failed: %v", err)
|
||||
}
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPUsername).Update("value", "smtpuser").Error; err != nil {
|
||||
t.Fatalf("set SMTP username failed: %v", err)
|
||||
}
|
||||
|
||||
// Invalidate the system config cache in Redis
|
||||
if err := db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err(); err != nil {
|
||||
t.Fatalf("invalidate system config cache failed: %v", err)
|
||||
}
|
||||
|
||||
router := setupUserTestRouter(t)
|
||||
|
||||
// 3. Perform login request without verification code
|
||||
payload := loginRequest{
|
||||
Username: username,
|
||||
Password: password,
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
|
||||
w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("Login(%q) status = %d, want %d. Body: %s", username, w.Code, http.StatusOK, w.Body.String())
|
||||
}
|
||||
|
||||
// Check response error msg
|
||||
var resp struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("unmarshal response failed: %v", err)
|
||||
}
|
||||
expectedError := errSMTPInvalidUseTempCodePrefix + "该账号未绑定邮箱,使用临时码登录"
|
||||
if resp.ErrorMsg != expectedError {
|
||||
t.Errorf("expected error %q, got %q", expectedError, resp.ErrorMsg)
|
||||
}
|
||||
|
||||
// 4. Check that verification code stored in Redis is "888888"
|
||||
ctx := context.Background()
|
||||
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("expected verification code '888888', got %q", storedCode)
|
||||
}
|
||||
|
||||
// 5. Retry login with code "888888"
|
||||
payload.Code = "888888"
|
||||
bodyWithCode, _ := json.Marshal(payload)
|
||||
w = performUserRequest(router, http.MethodPost, "/api/v1/user/login", bodyWithCode, nil)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("Login with code status = %d, want %d. Body: %s", w.Code, http.StatusOK, w.Body.String())
|
||||
}
|
||||
|
||||
var successResp struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &successResp); err != nil {
|
||||
t.Fatalf("unmarshal success response failed: %v", err)
|
||||
}
|
||||
if successResp.ErrorMsg != "" {
|
||||
t.Errorf("expected login success, got error %q", successResp.ErrorMsg)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,6 +5,9 @@
|
||||
package util
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/gin-contrib/sessions"
|
||||
)
|
||||
@@ -19,3 +22,31 @@ func GetSessionOptions(maxAge int) sessions.Options {
|
||||
Secure: config.Config.App.SessionSecure,
|
||||
}
|
||||
}
|
||||
|
||||
// StripCookieMaxAgeAndExpires 从 Set-Cookie 响应头中移除 Max-Age 和 Expires,从而使其成为浏览器会话 Cookie
|
||||
func StripCookieMaxAgeAndExpires(header http.Header, cookieName string) {
|
||||
headers := header["Set-Cookie"]
|
||||
if len(headers) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
newHeaders := make([]string, 0, len(headers))
|
||||
for _, h := range headers {
|
||||
if strings.HasPrefix(h, cookieName+"=") {
|
||||
parts := strings.Split(h, ";")
|
||||
newParts := make([]string, 0, len(parts))
|
||||
for _, p := range parts {
|
||||
trimmed := strings.TrimSpace(p)
|
||||
lower := strings.ToLower(trimmed)
|
||||
if strings.HasPrefix(lower, "max-age=") || strings.HasPrefix(lower, "expires=") {
|
||||
continue
|
||||
}
|
||||
newParts = append(newParts, p)
|
||||
}
|
||||
newHeaders = append(newHeaders, strings.Join(newParts, ";"))
|
||||
} else {
|
||||
newHeaders = append(newHeaders, h)
|
||||
}
|
||||
}
|
||||
header["Set-Cookie"] = newHeaders
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user