diff --git a/frontend/components/auth/login-form.tsx b/frontend/components/auth/login-form.tsx index 031ff983..113c0722 100644 --- a/frontend/components/auth/login-form.tsx +++ b/frontend/components/auth/login-form.tsx @@ -10,10 +10,12 @@ import Link from "next/link" import {useAuth} from "@/components/providers/auth-provider" import {Button} from "@/components/ui/button" import {Input} from "@/components/ui/input" +import {Label} from "@/components/ui/label" import {Separator} from "@/components/ui/separator" import {Spinner} from "@/components/ui/spinner" import {Card, CardContent} from "@/components/ui/card" import {CapWidget} from "@/components/auth/cap-widget" +import {OTPForm} from "./otp-form" import services from "@/lib/services" import type {LoginRequest} from "@/lib/services/auth/types" @@ -49,9 +51,9 @@ export function LoginForm() { const [password, setPassword] = useState("") const [code, setCode] = useState("") const [showLoginCodeInput, setShowLoginCodeInput] = useState(false) - const [maskedEmail, setMaskedEmail] = useState("") const [loginCooldown, setLoginCooldown] = useState(0) const [errorMessage, setErrorMessage] = useState("") + const [loginCodeTip, setLoginCodeTip] = useState(null) useEffect(() => { if (loginCooldown > 0) { @@ -105,7 +107,11 @@ export function LoginForm() { const errorMsg = error.message || "" if (errorMsg.startsWith("need_email_code:")) { const emailMasked = errorMsg.substring("need_email_code:".length) - setMaskedEmail(emailMasked) + setLoginCodeTip( + <> + 已向您的安全邮箱 {emailMasked} 发送了登录验证码。 + + ) setShowLoginCodeInput(true) setLoginCooldown(60) toast.success("登录验证码已发送至您的邮箱,请注意查收") @@ -117,6 +123,22 @@ export function LoginForm() { return } + if (errorMsg.startsWith("smtp_invalid:")) { + const tip = errorMsg.substring("smtp_invalid:".length) + setLoginCodeTip( + {tip} + ) + setShowLoginCodeInput(true) + setLoginCooldown(0) + toast.warning(tip) + if (capEnabled) { + capTokenRef.current = null + setCapReady(false) + setCapResetKey((key) => key + 1) + } + return + } + setErrorMessage(errorMsg || "登录失败,请重试") if (capEnabled) { capTokenRef.current = null @@ -130,8 +152,8 @@ export function LoginForm() { setErrorMessage("") const trimmedUsername = username.trim() if (!trimmedUsername || !password) { - toast.error("账号或密码未填写完整", { - description: "请先输入账号 and 密码后再登录", + toast.error("邮箱/用户名或密码未填写完整", { + description: "请先输入邮箱/用户名和密码后再登录", }) return } @@ -203,6 +225,20 @@ export function LoginForm() { ) } + if (showLoginCodeInput) { + return ( + + ) + } + return ( @@ -214,49 +250,31 @@ export function LoginForm() {
-
- setUsername(e.target.value)} - placeholder="用户名" - autoComplete="username" - onKeyDown={(e) => e.key === "Enter" && handlePasswordLogin()} - /> - setPassword(e.target.value)} - type="password" - placeholder="密码" - autoComplete="current-password" - onKeyDown={(e) => e.key === "Enter" && handlePasswordLogin()} - /> +
+
+ + setUsername(e.target.value)} + placeholder="请输入邮箱或用户名" + autoComplete="username" + onKeyDown={(e) => e.key === "Enter" && handlePasswordLogin()} + /> +
+
+ + setPassword(e.target.value)} + type="password" + placeholder="请输入密码" + autoComplete="current-password" + onKeyDown={(e) => e.key === "Enter" && handlePasswordLogin()} + /> +
- {showLoginCodeInput && ( -
-

- 已向您的安全邮箱 {maskedEmail} 发送了登录验证码。 -

-
- setCode(e.target.value)} - placeholder="6 位邮箱验证码" - maxLength={6} - className="flex-1" - onKeyDown={(e) => e.key === "Enter" && handlePasswordLogin()} - /> - -
-
- )}
{/* Cap 人机验证 */} @@ -320,7 +338,7 @@ export function LoginForm() { {registrationEnabled && (
- Don't have an account?{" "} + {"Don't have an account?"}{" "} Sign up diff --git a/frontend/components/auth/otp-form.tsx b/frontend/components/auth/otp-form.tsx new file mode 100644 index 00000000..358722ef --- /dev/null +++ b/frontend/components/auth/otp-form.tsx @@ -0,0 +1,105 @@ +"use client" + +import * as React from "react" +import {RefreshCwIcon} from "lucide-react" + +import {Button} from "@/components/ui/button" +import {Card, CardContent, CardDescription, CardFooter, CardHeader, CardTitle,} from "@/components/ui/card" +import {Field, FieldLabel} from "@/components/ui/field" +import {InputOTP, InputOTPGroup, InputOTPSeparator, InputOTPSlot,} from "@/components/ui/input-otp" +import {cn} from "@/lib/utils" + +interface OTPFormProps { + code: string + setCode: (val: string) => void + loginCodeTip: React.ReactNode + loginCooldown: number + isPending: boolean + onResend: () => void + onSubmit: () => void +} + +export function OTPForm({ + code, + setCode, + loginCodeTip, + loginCooldown, + isPending, + onResend, + onSubmit, +}: OTPFormProps) { + return ( + + + + 验证登录 + + + {loginCodeTip} + + + + +
+ + 验证码 + + +
+
+ + + + + + + + + + + + + +
+
+
+ + +
+ 遇到登录问题?{" "} + + 联系客服 + +
+
+
+ ) +} diff --git a/frontend/components/auth/register-form.tsx b/frontend/components/auth/register-form.tsx index 87d84daa..c7ea424e 100644 --- a/frontend/components/auth/register-form.tsx +++ b/frontend/components/auth/register-form.tsx @@ -10,6 +10,7 @@ import Link from "next/link" import {useAuth} from "@/components/providers/auth-provider" import {Button} from "@/components/ui/button" import {Input} from "@/components/ui/input" +import {Label} from "@/components/ui/label" import {Spinner} from "@/components/ui/spinner" import {Card, CardContent} from "@/components/ui/card" import services from "@/lib/services" @@ -116,17 +117,25 @@ export function RegisterForm() { toast.error("密码长度不能少于 8 位") return } - if (emailRegisterEnabled) { - if (!email.trim() || !code.trim()) { - toast.error("邮箱和验证码不能为空") - return - } + const trimmedEmail = email.trim() + if (!trimmedEmail) { + toast.error("邮箱不能为空") + return + } + const emailRegex = /^[^\s@]+@[^\s@]+\.[^\s@]+$/ + if (!emailRegex.test(trimmedEmail)) { + toast.error("请输入有效的邮箱地址") + return + } + if (emailRegisterEnabled && !code.trim()) { + toast.error("验证码不能为空") + return } registerMutation.mutate({ username: username.trim(), password, nickname: nickname.trim() || undefined, - email: email.trim() || undefined, + email: trimmedEmail, code: code.trim() || undefined, }) } @@ -155,55 +164,78 @@ export function RegisterForm() {
-
- setUsername(e.target.value)} - placeholder="用户名" - autoComplete="username" - onKeyDown={(e) => e.key === "Enter" && handleRegister()} - /> - setNickname(e.target.value)} - placeholder="昵称(可选)" - autoComplete="nickname" - onKeyDown={(e) => e.key === "Enter" && handleRegister()} - /> - setPassword(e.target.value)} - type="password" - placeholder="密码(至少 8 位)" - autoComplete="new-password" - onKeyDown={(e) => e.key === "Enter" && handleRegister()} - /> - setEmail(e.target.value)} - placeholder={emailRegisterEnabled ? "电子邮箱" : "电子邮箱(可选)"} - autoComplete="email" - onKeyDown={(e) => e.key === "Enter" && handleRegister()} - /> +
+
+ + setUsername(e.target.value)} + placeholder="请输入用户名" + autoComplete="username" + onKeyDown={(e) => e.key === "Enter" && handleRegister()} + /> +
+
+ + setNickname(e.target.value)} + placeholder="请输入昵称" + autoComplete="nickname" + onKeyDown={(e) => e.key === "Enter" && handleRegister()} + /> +
+
+ + setPassword(e.target.value)} + type="password" + placeholder="请输入密码(至少 8 位)" + autoComplete="new-password" + onKeyDown={(e) => e.key === "Enter" && handleRegister()} + /> +
+
+ + setEmail(e.target.value)} + placeholder="请输入电子邮箱" + autoComplete="email" + onKeyDown={(e) => e.key === "Enter" && handleRegister()} + /> +
{emailRegisterEnabled && ( -
- setCode(e.target.value)} - placeholder="6 位邮箱验证码" - maxLength={6} - className="flex-1" - onKeyDown={(e) => e.key === "Enter" && handleRegister()} - /> - +
+ +
+ setCode(e.target.value)} + placeholder="请输入 6 位邮箱验证码" + maxLength={6} + className="flex-1" + onKeyDown={(e) => e.key === "Enter" && handleRegister()} + /> + +
)}
diff --git a/internal/apps/oauth/sources.go b/internal/apps/oauth/sources.go index 8371ff8c..d5bc4698 100644 --- a/internal/apps/oauth/sources.go +++ b/internal/apps/oauth/sources.go @@ -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) { diff --git a/internal/apps/user/errs.go b/internal/apps/user/errs.go index a7777a3e..fabd3f71 100644 --- a/internal/apps/user/errs.go +++ b/internal/apps/user/errs.go @@ -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" ) diff --git a/internal/apps/user/logics.go b/internal/apps/user/logics.go index 851161ed..ea0e4bd7 100644 --- a/internal/apps/user/logics.go +++ b/internal/apps/user/logics.go @@ -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 发送邮箱验证码 diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go index 032562de..b55185da 100644 --- a/internal/apps/user/routers.go +++ b/internal/apps/user/routers.go @@ -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 diff --git a/internal/apps/user/routers_test.go b/internal/apps/user/routers_test.go index 914b71db..8784a1e3 100644 --- a/internal/apps/user/routers_test.go +++ b/internal/apps/user/routers_test.go @@ -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) + } +} diff --git a/internal/util/session.go b/internal/util/session.go index 7d3ee7ac..e56ff7c8 100644 --- a/internal/util/session.go +++ b/internal/util/session.go @@ -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 +}