From b262880189e3ea1d1630b7ee0727224a920956cf Mon Sep 17 00:00:00 2001 From: ryan Date: Sat, 13 Jun 2026 12:33:01 +0800 Subject: [PATCH] refactor(auth): improve session security, CAPTCHA validation and code hygiene - Integrate CapWidget with dual-scope capability on the frontend and protect registration/send-email-code endpoints on the backend. - Set session cookie SameSite mode to Lax. - Propagate request context through auth source database operations and optimize username uniqueness validation. - Standardize local error naming to camelCase and resolve references. - Fix linter rules, missing SheetContent closing tag, and unit tests. --- frontend/components/auth/register-form.tsx | 105 +++++++++++++++++++- frontend/lib/services/auth/auth.service.ts | 8 +- internal/apps/admin/auth_source/routers.go | 14 +-- internal/apps/oauth/audit.go | 4 +- internal/apps/oauth/errs.go | 26 ++--- internal/apps/oauth/middlewares.go | 3 +- internal/apps/oauth/routers.go | 6 ++ internal/apps/oauth/sources.go | 109 ++++++++++++--------- internal/apps/user/logics.go | 10 +- internal/apps/user/routers.go | 11 +++ internal/apps/user/routers_test.go | 9 +- internal/model/auth_source.go | 46 ++++----- internal/util/session.go | 1 + 13 files changed, 241 insertions(+), 111 deletions(-) diff --git a/frontend/components/auth/register-form.tsx b/frontend/components/auth/register-form.tsx index 61815af6..e9008542 100644 --- a/frontend/components/auth/register-form.tsx +++ b/frontend/components/auth/register-form.tsx @@ -1,6 +1,6 @@ "use client" -import {useEffect, useMemo, useState} from "react" +import {useEffect, useMemo, useRef, useState} from "react" import {useMutation, useQuery} from "@tanstack/react-query" import {useRouter, useSearchParams} from "next/navigation" import {toast} from "sonner" @@ -12,6 +12,7 @@ import {Input} from "@/components/ui/input" import {Spinner} from "@/components/ui/spinner" import {Field, FieldGroup, FieldLabel} from "@/components/ui/field" import {AuthHeading} from "@/components/auth/auth-shell" +import {CapWidget} from "@/components/auth/cap-widget" import services from "@/lib/services" import type {RegisterRequest} from "@/lib/services/auth/types" import {safeRedirectTarget} from "@/lib/utils" @@ -66,6 +67,33 @@ export function RegisterForm() { const emailRegisterEnabled = configBool(publicConfigQuery.data?.email_register_verification_enabled, false) + const capEnabled = configBool(publicConfigQuery.data?.cap_login_enabled, false) + const capAutoSolve = configBool(publicConfigQuery.data?.cap_auto_solve, true) + + const [capScope, setCapScope] = useState<'send_email_code' | 'register'>('send_email_code') + + // 监听 emailRegisterEnabled 改变初始 scope + useEffect(() => { + setCapScope(emailRegisterEnabled ? 'send_email_code' : 'register') + }, [emailRegisterEnabled]) + + const capTokenRef = useRef(null) + const [capReady, setCapReady] = useState(false) + const [capError, setCapError] = useState(false) + const [capResetKey, setCapResetKey] = useState(0) + + const handleCapToken = (token: string) => { + capTokenRef.current = token + setCapReady(true) + setCapError(false) + } + + const handleCapError = () => { + capTokenRef.current = null + setCapReady(false) + setCapError(true) + } + // Redirect to login if registration is closed useEffect(() => { if (publicConfigQuery.isSuccess && !registrationEnabled) { @@ -75,7 +103,15 @@ export function RegisterForm() { }, [publicConfigQuery.isSuccess, registrationEnabled, router]) const registerMutation = useMutation({ - mutationFn: (req: RegisterRequest) => services.auth.register(req), + mutationFn: (req: RegisterRequest) => { + const headers: Record = {} + if (capEnabled && capTokenRef.current) { + headers["X-Cap-Token"] = capTokenRef.current + capTokenRef.current = null + setCapReady(false) + } + return services.auth.register(req, Object.keys(headers).length ? headers : undefined) + }, onSuccess: (user) => { setUser(user) router.replace(redirectTarget) @@ -83,20 +119,53 @@ export function RegisterForm() { }, onError: (error: Error) => { setErrorMessage(error.message || "注册失败,请重试") + if (capEnabled) { + capTokenRef.current = null + setCapReady(false) + setCapResetKey((key) => key + 1) + } }, }) const sendRegisterCodeMutation = useMutation({ - mutationFn: (targetEmail: string) => services.auth.sendEmailCode(targetEmail, "register"), + mutationFn: (targetEmail: string) => { + const headers: Record = {} + if (capEnabled && capTokenRef.current) { + headers["X-Cap-Token"] = capTokenRef.current + capTokenRef.current = null + setCapReady(false) + } + return services.auth.sendEmailCode(targetEmail, "register", Object.keys(headers).length ? headers : undefined) + }, onSuccess: () => { setRegisterCooldown(60) toast.success("验证码已发送至您的邮箱,请查收") + if (capEnabled) { + setCapScope('register') + capTokenRef.current = null + setCapReady(false) + setCapResetKey((key) => key + 1) + } }, onError: (error: Error) => { toast.error(error.message || "发送验证码失败,请重试") + if (capEnabled) { + capTokenRef.current = null + setCapReady(false) + setCapResetKey((key) => key + 1) + } }, }) + const registerDisabled = + registerMutation.isPending || + (capEnabled && capScope === 'register' && capAutoSolve && !capReady && !capError) + + const sendCodeDisabled = + registerCooldown > 0 || + sendRegisterCodeMutation.isPending || + (capEnabled && capScope === 'send_email_code' && capAutoSolve && !capReady && !capError) + const handleSendRegisterCode = () => { const trimmedEmail = email.trim() if (!trimmedEmail) { @@ -108,6 +177,14 @@ export function RegisterForm() { toast.error("请输入有效的邮箱地址") return } + if (capEnabled && capScope === 'send_email_code' && !capReady) { + toast.error( + capAutoSolve + ? "人机验证尚未完成,请稍候…" + : "请先点击「开始验证」完成人机验证", + ) + return + } sendRegisterCodeMutation.mutate(trimmedEmail) } @@ -135,6 +212,14 @@ export function RegisterForm() { toast.error("验证码不能为空") return } + if (capEnabled && capScope === 'register' && !capReady) { + toast.error( + capAutoSolve + ? "人机验证尚未完成,请稍候…" + : "请先点击「开始验证」完成人机验证", + ) + return + } registerMutation.mutate({ username: username.trim(), password, @@ -235,7 +320,7 @@ export function RegisterForm() { type="button" variant="outline" onClick={handleSendRegisterCode} - disabled={registerCooldown > 0 || sendRegisterCodeMutation.isPending} + disabled={sendCodeDisabled} className="h-10 w-[120px] text-xs [@media(max-height:700px)]:h-9" > {registerCooldown > 0 ? `${registerCooldown}秒后重发` : "获取验证码"} @@ -245,6 +330,16 @@ export function RegisterForm() { )} + {capEnabled && ( + + )} + {errorMessage ? (
{errorMessage} @@ -256,7 +351,7 @@ export function RegisterForm() { className="h-10 w-full [@media(max-height:700px)]:h-9" variant="auth" onClick={handleRegister} - disabled={registerMutation.isPending} + disabled={registerDisabled} > {registerMutation.isPending ? ( <> diff --git a/frontend/lib/services/auth/auth.service.ts b/frontend/lib/services/auth/auth.service.ts index 8def946e..708dc157 100644 --- a/frontend/lib/services/auth/auth.service.ts +++ b/frontend/lib/services/auth/auth.service.ts @@ -109,12 +109,12 @@ export class AuthService extends BaseService { return this.post('/user/login', request, headers ? ({ headers } as unknown as InternalAxiosRequestConfig) : undefined); } - static async register(request: RegisterRequest): Promise { - return this.post('/user/register', request); + static async register(request: RegisterRequest, headers?: Record): Promise { + return this.post('/user/register', request, headers ? ({ headers } as unknown as InternalAxiosRequestConfig) : undefined); } - static async sendEmailCode(email: string, scene: 'register' | 'login'): Promise { - return this.post('/user/send-email-code', { email, scene }); + static async sendEmailCode(email: string, scene: 'register' | 'login', headers?: Record): Promise { + return this.post('/user/send-email-code', { email, scene }, headers ? ({ headers } as unknown as InternalAxiosRequestConfig) : undefined); } static async changePassword(request: ChangePasswordRequest): Promise { diff --git a/internal/apps/admin/auth_source/routers.go b/internal/apps/admin/auth_source/routers.go index 0f43dac6..cc65102e 100644 --- a/internal/apps/admin/auth_source/routers.go +++ b/internal/apps/admin/auth_source/routers.go @@ -45,7 +45,7 @@ type ToggleAuthSourceRequest struct { // @Failure 500 {object} util.ResponseAny "内部错误" // @Router /api/v1/admin/auth-sources [get] func ListAuthSources(c *gin.Context) { - sources, err := model.GetAuthSources() + sources, err := model.GetAuthSources(c.Request.Context()) if err != nil { c.JSON(http.StatusInternalServerError, util.Err(err.Error())) return @@ -84,7 +84,7 @@ func CreateAuthSource(c *gin.Context) { Scopes: req.Scopes, IconURL: req.IconURL, } - if err := model.CreateAuthSource(&source); err != nil { + if err := model.CreateAuthSource(c.Request.Context(), &source); err != nil { c.JSON(http.StatusBadRequest, util.Err(err.Error())) return } @@ -133,11 +133,11 @@ func UpdateAuthSource(c *gin.Context) { IconURL: req.IconURL, } keepSecret := source.ClientSecret == "" - if err := model.UpdateAuthSource(&source, keepSecret); err != nil { + if err := model.UpdateAuthSource(c.Request.Context(), &source, keepSecret); err != nil { c.JSON(http.StatusBadRequest, util.Err(err.Error())) return } - updated, err := model.GetAuthSourceByID(id) + updated, err := model.GetAuthSourceByID(c.Request.Context(), id) if err != nil { c.JSON(http.StatusInternalServerError, util.Err(err.Error())) return @@ -173,7 +173,7 @@ func ToggleAuthSource(c *gin.Context) { return } - if err := model.ToggleAuthSource(id, req.IsActive); err != nil { + if err := model.ToggleAuthSource(c.Request.Context(), id, req.IsActive); err != nil { c.JSON(http.StatusBadRequest, util.Err(err.Error())) return } @@ -198,7 +198,7 @@ func DeleteAuthSource(c *gin.Context) { c.JSON(http.StatusBadRequest, util.Err(err.Error())) return } - if err := model.DeleteAuthSource(id); err != nil { + if err := model.DeleteAuthSource(c.Request.Context(), id); err != nil { c.JSON(http.StatusBadRequest, util.Err(err.Error())) return } @@ -210,7 +210,7 @@ func parseSourceID(c *gin.Context) (uint64, error) { if raw == "" { return 0, errors.New(admin.InvalidAuthSourceID) } - source, err := model.GetAuthSourceByName(raw) + source, err := model.GetAuthSourceByName(c.Request.Context(), raw) if err == nil { return source.ID, nil } diff --git a/internal/apps/oauth/audit.go b/internal/apps/oauth/audit.go index f3f66bd2..909f0c75 100644 --- a/internal/apps/oauth/audit.go +++ b/internal/apps/oauth/audit.go @@ -29,8 +29,8 @@ func LogForAudit(ctx context.Context, user *model.User, c *gin.Context) { auditJSON, err := json.Marshal(auditLog) if err != nil { logger.ErrorF(ctx, "[LoginRequiredAudit] marshal failed: %v", err) - logger.InfoF(ctx, "[LoginRequiredAudit] %s %d %s", c.ClientIP(), user.ID, user.Username) + logger.DebugF(ctx, "[LoginRequiredAudit] %s %d %s", c.ClientIP(), user.ID, user.Username) } else { - logger.InfoF(ctx, "[LoginRequiredAudit] %s", auditJSON) + logger.DebugF(ctx, "[LoginRequiredAudit] %s", auditJSON) } } diff --git a/internal/apps/oauth/errs.go b/internal/apps/oauth/errs.go index b4dd32f5..71eed170 100644 --- a/internal/apps/oauth/errs.go +++ b/internal/apps/oauth/errs.go @@ -6,17 +6,17 @@ package oauth // OAuth 认证相关错误消息 const ( - InvalidState = "非法登录请求" - IDTokenVerifyFailed = "ID Token 验证失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - IDTokenVerifyFailedFormat = "%s: %w" - NonceMismatch = "nonce 不匹配,可能存在重放攻击" - NoActiveAuthSource = "未配置可用认证源" - ServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试" - AuthSourceRequired = "认证源不能为空" - DiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL" - UsernameGenerateFailed = "无法生成可用用户名" - UsernameFromSourceFailed = "无法从认证源获取用户名" - AuthSourceDisabled = "认证源未启用" - InvalidExternalAccountBindingID = "绑定记录 ID 无效" - TokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials + errInvalidState = "非法登录请求" + errIDTokenVerifyFailed = "ID Token 验证失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials + errIDTokenVerifyFailedFormat = "%s: %w" + errNonceMismatch = "nonce 不匹配,可能存在重放攻击" + errNoActiveAuthSource = "未配置可用认证源" + errServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试" + errAuthSourceRequired = "认证源不能为空" + errDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL" + errUsernameGenerateFailed = "无法生成可用用户名" + errUsernameFromSourceFailed = "无法从认证源获取用户名" + errAuthSourceDisabled = "认证源未启用" + errInvalidExternalAccountBindingID = "绑定记录 ID 无效" + ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials ) diff --git a/internal/apps/oauth/middlewares.go b/internal/apps/oauth/middlewares.go index 84e0c785..3d56c80d 100644 --- a/internal/apps/oauth/middlewares.go +++ b/internal/apps/oauth/middlewares.go @@ -112,10 +112,9 @@ func LoginRequired() gin.HandlerFunc { func DisallowTokenAuth() gin.HandlerFunc { return func(c *gin.Context) { if tokenAuth, _ := util.GetFromContext[bool](c, TokenAuthKey); tokenAuth { - c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"error_msg": TokenAuthNotAllowed, "data": nil}) + c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"error_msg": ErrTokenAuthNotAllowed, "data": nil}) return } c.Next() } } - diff --git a/internal/apps/oauth/routers.go b/internal/apps/oauth/routers.go index 8882d4e9..a0432b46 100644 --- a/internal/apps/oauth/routers.go +++ b/internal/apps/oauth/routers.go @@ -7,6 +7,7 @@ package oauth import ( "net/http" + "github.com/Rain-kl/Wavelet/internal/logger" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/util" "github.com/gin-contrib/sessions" @@ -89,6 +90,11 @@ func UserInfo(c *gin.Context) { // @Router /api/v1/oauth/logout [get] func Logout(c *gin.Context) { session := sessions.Default(c) + userID := session.Get(UserIDKey) + username := session.Get(UserNameKey) + if userID != nil { + logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP()) + } session.Options(util.GetSessionOptions(-1)) session.Clear() if err := session.Save(); err != nil { diff --git a/internal/apps/oauth/sources.go b/internal/apps/oauth/sources.go index cf7472e7..3fa9f7c6 100644 --- a/internal/apps/oauth/sources.go +++ b/internal/apps/oauth/sources.go @@ -17,6 +17,7 @@ import ( "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/logger" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/util" "github.com/coreos/go-oidc/v3/oidc" @@ -98,30 +99,28 @@ func isOIDCLoginEnabled(ctx context.Context) bool { return enabled } - - -func resolveAuthSource(sourceName string) (*model.AuthSource, error) { +func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) { name := strings.TrimSpace(strings.ToLower(sourceName)) if name == "" { - sources, err := model.GetActiveAuthSources() + sources, err := model.GetActiveAuthSources(ctx) if err != nil { return nil, err } if len(sources) == 0 { - return nil, errors.New(NoActiveAuthSource) + return nil, errors.New(errNoActiveAuthSource) } return &sources[0], nil } - return model.GetAuthSourceByName(name) + return model.GetAuthSourceByName(ctx, name) } -func activeLoginSources() []AuthSourceView { - enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyOIDCLoginEnabled) +func activeLoginSources(ctx context.Context) []AuthSourceView { + enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled) if err == nil && !enabled { return nil } - dbSources, err := model.GetActiveAuthSources() + dbSources, err := model.GetActiveAuthSources(ctx) if err != nil { return nil } @@ -143,18 +142,18 @@ func activeLoginSources() []AuthSourceView { func getFrontendLoginRedirectURL(ctx context.Context) (string, error) { var sc model.SystemConfig if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err != nil || strings.TrimSpace(sc.Value) == "" { - return "", errors.New(ServerAddressMissing) + return "", errors.New(errServerAddressMissing) } return strings.TrimRight(sc.Value, "/") + "/login", nil } func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) { if source == nil { - return nil, nil, errors.New(AuthSourceRequired) + return nil, nil, errors.New(errAuthSourceRequired) } if source.OpenIDDiscoveryURL == "" { - return nil, nil, errors.New(DiscoveryURLRequired) + return nil, nil, errors.New(errDiscoveryURLRequired) } // Clean the issuer URL (trim /.well-known/openid-configuration if configured by mistake) @@ -229,21 +228,38 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro } func uniqueUsername(ctx context.Context, base string) (string, error) { - candidate := strings.TrimSpace(base) - if candidate == "" { - candidate = "user" + base = strings.TrimSpace(base) + if base == "" { + base = "user" } - for i := 0; i < 1000; i++ { - var count int64 - if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", candidate).Count(&count).Error; err != nil { - return "", err - } - if count == 0 { + + var existingUsernames []string + if err := db.DB(ctx).Model(&model.User{}). + Where("username = ? OR username LIKE ?", base, base+"-%"). + Pluck("username", &existingUsernames).Error; err != nil { + return "", err + } + + // 将现有的用户名放入 map 中,以便 O(1) 查找 + exists := make(map[string]bool, len(existingUsernames)) + for _, u := range existingUsernames { + exists[strings.ToLower(u)] = true + } + + // 检查 base 是否被占用 + if !exists[strings.ToLower(base)] { + return base, nil + } + + // 顺序查找第一个可用的带后缀用户名 + for i := 1; i <= 1000; i++ { + candidate := fmt.Sprintf("%s-%d", base, i) + if !exists[strings.ToLower(candidate)] { return candidate, nil } - candidate = fmt.Sprintf("%s-%d", base, i+1) } - return "", errors.New(UsernameGenerateFailed) + + return "", errors.New(errUsernameGenerateFailed) } func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code string, nonce string, redirectURL string) (*model.OAuthUserInfo, error) { @@ -288,10 +304,10 @@ func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *o } idToken, verifyErr := verifier.Verify(ctx, rawIDToken) if verifyErr != nil { - return fmt.Errorf(IDTokenVerifyFailedFormat, IDTokenVerifyFailed, verifyErr) + return fmt.Errorf(errIDTokenVerifyFailedFormat, errIDTokenVerifyFailed, verifyErr) } if nonce != "" && idToken.Nonce != nonce { - return errors.New(NonceMismatch) + return errors.New(errNonceMismatch) } if claimsErr := idToken.Claims(userInfo); claimsErr != nil { return claimsErr @@ -316,7 +332,7 @@ func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error { userInfo.Username = userInfo.Sub } if userInfo.Username == "" { - return errors.New(UsernameFromSourceFailed) + return errors.New(errUsernameFromSourceFailed) } if userInfo.Name == "" { userInfo.Name = userInfo.Username @@ -344,7 +360,7 @@ func buildCallbackResult(user *model.User, status string) OAuthCallbackResult { // @Success 200 {object} util.ResponseAny{data=[]oauth.AuthSourceView} "登录源列表" // @Router /api/v1/oauth/sources [get] func GetLoginSources(c *gin.Context) { - c.JSON(http.StatusOK, util.OK(activeLoginSources())) + c.JSON(http.StatusOK, util.OK(activeLoginSources(c.Request.Context()))) } // GetLoginURL 获取登录授权地址 @@ -355,27 +371,26 @@ func GetLoginSources(c *gin.Context) { // @Param source query string false "认证源名称,为空使用第一个启用的认证源" // @Success 200 {object} util.ResponseAny{data=oauth.OAuthAuthorizeResponse} "授权 URL" // @Failure 400 {object} util.ResponseAny "认证源不存在或未配置" -// @Failure 500 {object} util.ResponseAny "Redis 异常或构造 URL 失败" +// @Failure 500 {object} util.ResponseAny "Redis 异常 or 构造 URL 失败" // @Router /api/v1/oauth/login [get] func GetLoginURL(c *gin.Context) { ctx := c.Request.Context() if !isOIDCLoginEnabled(ctx) { - c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled)) + c.JSON(http.StatusBadRequest, util.Err(errAuthSourceDisabled)) return } - source, err := resolveAuthSource(c.Query("source")) + source, err := resolveAuthSource(ctx, c.Query("source")) if err != nil { c.JSON(http.StatusBadRequest, util.Err(err.Error())) return } if !source.IsActive { - c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled)) + c.JSON(http.StatusBadRequest, util.Err(errAuthSourceDisabled)) return } - session := sessions.Default(c) token, isNew := ensureSessionToken(session) if isNew { @@ -404,7 +419,6 @@ func GetLoginURL(c *gin.Context) { return } - authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state) if err != nil { c.JSON(http.StatusBadRequest, util.Err(err.Error())) @@ -442,18 +456,18 @@ func buildAuthorizeURL(ctx context.Context, source *model.AuthSource, state stri func Authorize(c *gin.Context) { ctx := c.Request.Context() if !isOIDCLoginEnabled(ctx) { - c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled)) + c.JSON(http.StatusBadRequest, util.Err(errAuthSourceDisabled)) return } - source, err := resolveAuthSource(c.Param("source")) + source, err := resolveAuthSource(ctx, c.Param("source")) if err != nil { c.JSON(http.StatusBadRequest, util.Err(err.Error())) return } if !source.IsActive { - c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled)) + c.JSON(http.StatusBadRequest, util.Err(errAuthSourceDisabled)) return } purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose"))) @@ -525,7 +539,7 @@ func Callback(c *gin.Context) { stateKey := db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State)) payloadRaw, err := db.Redis.Get(ctx, stateKey).Result() if err != nil { - c.JSON(http.StatusBadRequest, util.Err(InvalidState)) + c.JSON(http.StatusBadRequest, util.Err(errInvalidState)) return } _ = db.Redis.Del(ctx, stateKey) @@ -560,25 +574,22 @@ func Callback(c *gin.Context) { return } - - if !isOIDCLoginEnabled(ctx) { - c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled)) + c.JSON(http.StatusBadRequest, util.Err(errAuthSourceDisabled)) return } - source, err := resolveAuthSource(payload.SourceName) + source, err := resolveAuthSource(ctx, payload.SourceName) if err != nil { c.JSON(http.StatusBadRequest, util.Err(err.Error())) return } if !source.IsActive { - c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled)) + c.JSON(http.StatusBadRequest, util.Err(errAuthSourceDisabled)) return } - redirectURL, err := getFrontendLoginRedirectURL(ctx) if err != nil { c.JSON(http.StatusBadRequest, util.Err(err.Error())) @@ -661,6 +672,9 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth c.JSON(http.StatusInternalServerError, util.Err(err.Error())) return } + + logger.InfoF(ctx, "[LoginAudit] successful OAuth login via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP()) + c.JSON(http.StatusOK, util.OK(buildCallbackResult(&user, "logged_in"))) } @@ -677,7 +691,6 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.A return model.User{}, false } - username, uniqueErr := uniqueUsername(ctx, userInfo.Username) if uniqueErr != nil { c.JSON(http.StatusInternalServerError, util.Err(uniqueErr.Error())) @@ -700,6 +713,8 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.A c.JSON(http.StatusBadRequest, util.Err(err.Error())) return model.User{}, false } + logger.InfoF(ctx, "[LoginAudit] successful OAuth registration via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP()) + return user, true } @@ -715,7 +730,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.A // @Router /api/v1/oauth/external-accounts [get] func ListExternalAccounts(c *gin.Context) { userID := GetUserIDFromContext(c) - accounts, err := model.ListExternalAccountsByUserID(userID) + accounts, err := model.ListExternalAccountsByUserID(c.Request.Context(), userID) if err != nil { c.JSON(http.StatusInternalServerError, util.Err(err.Error())) return @@ -743,10 +758,10 @@ func DeleteExternalAccount(c *gin.Context) { rawID := strings.TrimSpace(c.Param("id")) id, err := strconv.ParseUint(rawID, 10, 64) if err != nil || id == 0 { - c.JSON(http.StatusBadRequest, util.Err(InvalidExternalAccountBindingID)) + c.JSON(http.StatusBadRequest, util.Err(errInvalidExternalAccountBindingID)) return } - if err := model.DeleteExternalAccountForUser(id, userID); err != nil { + if err := model.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil { c.JSON(http.StatusBadRequest, util.Err(err.Error())) return } diff --git a/internal/apps/user/logics.go b/internal/apps/user/logics.go index 26e72e2a..edb03a76 100644 --- a/internal/apps/user/logics.go +++ b/internal/apps/user/logics.go @@ -45,7 +45,7 @@ func isEmailRegisterVerificationEnabled(ctx context.Context) bool { } func isSMTPConfigured(ctx context.Context) bool { - var host, port, username string + var host, port, username, password string var scHost model.SystemConfig if err := scHost.GetByKey(ctx, model.ConfigKeySMTPHost); err == nil { @@ -59,8 +59,12 @@ func isSMTPConfigured(ctx context.Context) bool { if err := scUser.GetByKey(ctx, model.ConfigKeySMTPUsername); err == nil { username = scUser.Value } + var scPass model.SystemConfig + if err := scPass.GetByKey(ctx, model.ConfigKeySMTPPassword); err == nil { + password = scPass.Value + } - return host != "" && port != "" && username != "" + return host != "" && port != "" && username != "" && password != "" } func generateVerificationCode() (string, error) { @@ -241,8 +245,6 @@ func validateRegisterEmailVerification(ctx context.Context, req *registerRequest return nil } - - type updateProfileRequest struct { Nickname string `json:"nickname"` Email string `json:"email"` diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go index 6458b1ff..df2afd58 100644 --- a/internal/apps/user/routers.go +++ b/internal/apps/user/routers.go @@ -14,6 +14,7 @@ import ( "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/logger" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/util" "github.com/gin-contrib/sessions" @@ -124,10 +125,12 @@ func Login(c *gin.Context) { var user model.User ctx := c.Request.Context() if err := db.DB(ctx).Where("username = ? OR email = ?", req.Username, req.Username).First(&user).Error; err != nil { + logger.WarnF(ctx, "[LoginAudit] failed login attempt (username not found) for input: %s, IP: %s", req.Username, c.ClientIP()) c.JSON(http.StatusOK, util.Err(errUsernameOrPasswordWrong)) return } if !user.IsActive { + logger.WarnF(ctx, "[LoginAudit] banned user login attempt for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP()) c.JSON(http.StatusOK, util.Err(common.BannedAccount)) return } @@ -136,6 +139,7 @@ func Login(c *gin.Context) { isPlaintext := !user.IsPasswordEncrypted() if !user.CheckPassword(req.Password) { + logger.WarnF(ctx, "[LoginAudit] failed login attempt (incorrect password) for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP()) c.JSON(http.StatusOK, util.Err(errUsernameOrPasswordWrong)) return } @@ -165,6 +169,8 @@ func Login(c *gin.Context) { return } + logger.InfoF(ctx, "[LoginAudit] successful login for user: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP()) + c.JSON(http.StatusOK, util.OK(oauth.BuildBasicUserInfo(&user, needChangePassword))) } @@ -264,6 +270,11 @@ func Register(c *gin.Context) { // @Router /api/v1/user/logout [get] func Logout(c *gin.Context) { session := sessions.Default(c) + userID := session.Get(oauth.UserIDKey) + username := session.Get(oauth.UserNameKey) + if userID != nil { + logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP()) + } session.Options(util.GetSessionOptions(-1)) session.Clear() if err := session.Save(); err != nil { diff --git a/internal/apps/user/routers_test.go b/internal/apps/user/routers_test.go index 7e305422..eb566ffa 100644 --- a/internal/apps/user/routers_test.go +++ b/internal/apps/user/routers_test.go @@ -383,6 +383,9 @@ func TestLoginEmailVerificationFallbackForEmptyEmail(t *testing.T) { if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPUsername).Update("value", "smtpuser").Error; err != nil { t.Fatalf("set SMTP username failed: %v", err) } + if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPPassword).Update("value", "smtppassword").Error; err != nil { + t.Fatalf("set SMTP password failed: %v", err) + } // Invalidate the system config cache in Redis if err := db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err(); err != nil { @@ -522,8 +525,8 @@ func TestAccessTokenEndpointsDisallowTokenAuth(t *testing.T) { if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { t.Fatalf("decode response failed: %v", err) } - if resp.ErrorMsg != oauth.TokenAuthNotAllowed { - t.Errorf("expected error message %q, got %q", oauth.TokenAuthNotAllowed, resp.ErrorMsg) + if resp.ErrorMsg != oauth.ErrTokenAuthNotAllowed { + t.Errorf("expected error message %q, got %q", oauth.ErrTokenAuthNotAllowed, resp.ErrorMsg) } // 4. Test that accessing using a Session succeeds @@ -684,5 +687,3 @@ func TestChangePasswordRevocation(t *testing.T) { t.Errorf("expected access token to be revoked (401), got %d", wTokenAfter.Code) } } - - diff --git a/internal/model/auth_source.go b/internal/model/auth_source.go index 3e443951..d1d65d74 100644 --- a/internal/model/auth_source.go +++ b/internal/model/auth_source.go @@ -118,9 +118,9 @@ func (source *AuthSource) Sanitize() { } // GetAuthSources 获取所有认证源(已脱敏) -func GetAuthSources() ([]AuthSource, error) { +func GetAuthSources(ctx context.Context) ([]AuthSource, error) { var sources []AuthSource - if err := db.DB(context.Background()).Order("id asc").Find(&sources).Error; err != nil { + if err := db.DB(ctx).Order("id asc").Find(&sources).Error; err != nil { return nil, err } for i := range sources { @@ -130,9 +130,9 @@ func GetAuthSources() ([]AuthSource, error) { } // GetActiveAuthSources 获取所有已启用的认证源(已脱敏) -func GetActiveAuthSources() ([]AuthSource, error) { +func GetActiveAuthSources(ctx context.Context) ([]AuthSource, error) { var sources []AuthSource - if err := db.DB(context.Background()).Where("is_active = ?", true).Order("id asc").Find(&sources).Error; err != nil { + if err := db.DB(ctx).Where("is_active = ?", true).Order("id asc").Find(&sources).Error; err != nil { return nil, err } for i := range sources { @@ -142,12 +142,12 @@ func GetActiveAuthSources() ([]AuthSource, error) { } // GetAuthSourceByID 根据 ID 获取认证源 -func GetAuthSourceByID(id uint64) (*AuthSource, error) { +func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) { if id == 0 { return nil, errors.New(errAuthSourceIDRequired) } var source AuthSource - if err := db.DB(context.Background()).First(&source, "id = ?", id).Error; err != nil { + if err := db.DB(ctx).First(&source, "id = ?", id).Error; err != nil { return nil, err } source.ClientSecretConfigured = source.ClientSecret != "" @@ -155,13 +155,13 @@ func GetAuthSourceByID(id uint64) (*AuthSource, error) { } // GetAuthSourceByName 根据名称获取认证源(名称比较不区分大小写) -func GetAuthSourceByName(name string) (*AuthSource, error) { +func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error) { name = strings.TrimSpace(name) if name == "" { return nil, errors.New(errAuthSourceNameRequired) } var source AuthSource - if err := db.DB(context.Background()).First(&source, "LOWER(name) = LOWER(?)", name).Error; err != nil { + if err := db.DB(ctx).First(&source, "LOWER(name) = LOWER(?)", name).Error; err != nil { return nil, err } source.ClientSecretConfigured = source.ClientSecret != "" @@ -169,20 +169,20 @@ func GetAuthSourceByName(name string) (*AuthSource, error) { } // CreateAuthSource 创建认证源 -func CreateAuthSource(source *AuthSource) error { +func CreateAuthSource(ctx context.Context, source *AuthSource) error { if err := source.Validate(); err != nil { return err } - return db.DB(context.Background()).Create(source).Error + return db.DB(ctx).Create(source).Error } // UpdateAuthSource 更新认证源,keepSecret 为 true 时保留原密钥 -func UpdateAuthSource(source *AuthSource, keepSecret bool) error { +func UpdateAuthSource(ctx context.Context, source *AuthSource, keepSecret bool) error { if source.ID == 0 { return errors.New(errAuthSourceIDRequired) } var current AuthSource - if err := db.DB(context.Background()).First(¤t, "id = ?", source.ID).Error; err != nil { + if err := db.DB(ctx).First(¤t, "id = ?", source.ID).Error; err != nil { return err } if keepSecret { @@ -191,7 +191,7 @@ func UpdateAuthSource(source *AuthSource, keepSecret bool) error { if err := source.Validate(); err != nil { return err } - return db.DB(context.Background()).Model(¤t).Updates(map[string]any{ + return db.DB(ctx).Model(¤t).Updates(map[string]any{ "name": source.Name, "type": source.Type, "display_name": source.DisplayName, @@ -205,8 +205,8 @@ func UpdateAuthSource(source *AuthSource, keepSecret bool) error { } // ToggleAuthSource 切换认证源启用状态 -func ToggleAuthSource(id uint64, isActive bool) error { - source, err := GetAuthSourceByID(id) +func ToggleAuthSource(ctx context.Context, id uint64, isActive bool) error { + source, err := GetAuthSourceByID(ctx, id) if err != nil { return err } @@ -214,15 +214,15 @@ func ToggleAuthSource(id uint64, isActive bool) error { if err := source.Validate(); err != nil { return err } - return db.DB(context.Background()).Model(&AuthSource{}).Where("id = ?", id).Update("is_active", isActive).Error + return db.DB(ctx).Model(&AuthSource{}).Where("id = ?", id).Update("is_active", isActive).Error } // DeleteAuthSource 删除认证源及其关联的外部帐号绑定 -func DeleteAuthSource(id uint64) error { +func DeleteAuthSource(ctx context.Context, id uint64) error { if id == 0 { return errors.New(errAuthSourceIDRequired) } - return db.DB(context.Background()).Transaction(func(tx *gorm.DB) error { + return db.DB(ctx).Transaction(func(tx *gorm.DB) error { if err := tx.Where("auth_source_id = ?", id).Delete(&ExternalAccount{}).Error; err != nil { return err } @@ -268,12 +268,12 @@ func BindExternalAccount(ctx context.Context, account *ExternalAccount) error { } // ListExternalAccountsByUserID 获取指定用户的所有外部帐号绑定视图 -func ListExternalAccountsByUserID(userID uint64) ([]ExternalAccountView, error) { +func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccountView, error) { if userID == 0 { return nil, errors.New(errUserIDRequired) } var accounts []ExternalAccount - if err := db.DB(context.Background()).Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil { + if err := db.DB(ctx).Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil { return nil, err } views := make([]ExternalAccountView, 0, len(accounts)) @@ -284,7 +284,7 @@ func ListExternalAccountsByUserID(userID uint64) ([]ExternalAccountView, error) sourceType = "oidc" label = "历史认证源" } else { - source, err := GetAuthSourceByID(account.AuthSourceID) + source, err := GetAuthSourceByID(ctx, account.AuthSourceID) if err != nil { continue } @@ -310,9 +310,9 @@ func ListExternalAccountsByUserID(userID uint64) ([]ExternalAccountView, error) } // DeleteExternalAccountForUser 删除指定用户的外部帐号绑定 -func DeleteExternalAccountForUser(id uint64, userID uint64) error { +func DeleteExternalAccountForUser(ctx context.Context, id uint64, userID uint64) error { if id == 0 || userID == 0 { return errors.New(errExternalAccountBindingIDRequired) } - return db.DB(context.Background()).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error + return db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error } diff --git a/internal/util/session.go b/internal/util/session.go index e56ff7c8..28b0c6c9 100644 --- a/internal/util/session.go +++ b/internal/util/session.go @@ -20,6 +20,7 @@ func GetSessionOptions(maxAge int) sessions.Options { MaxAge: maxAge, HttpOnly: config.Config.App.SessionHTTPOnly, Secure: config.Config.App.SessionSecure, + SameSite: http.SameSiteLaxMode, } }