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.
This commit is contained in:
ryan
2026-06-13 12:33:01 +08:00
parent 04280a7b11
commit b262880189
13 changed files with 241 additions and 111 deletions
+100 -5
View File
@@ -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<string | null>(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<string, string> = {}
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<string, string> = {}
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() {
)}
</FieldGroup>
{capEnabled && (
<CapWidget
key={capResetKey}
scope={capScope}
onToken={handleCapToken}
onError={handleCapError}
autoStart={capAutoSolve}
/>
)}
{errorMessage ? (
<div className="rounded-lg border border-destructive/30 bg-destructive/5 px-3 py-2 text-sm text-destructive">
{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 ? (
<>
+4 -4
View File
@@ -109,12 +109,12 @@ export class AuthService extends BaseService {
return this.post<User>('/user/login', request, headers ? ({ headers } as unknown as InternalAxiosRequestConfig) : undefined);
}
static async register(request: RegisterRequest): Promise<User> {
return this.post<User>('/user/register', request);
static async register(request: RegisterRequest, headers?: Record<string, string>): Promise<User> {
return this.post<User>('/user/register', request, headers ? ({ headers } as unknown as InternalAxiosRequestConfig) : undefined);
}
static async sendEmailCode(email: string, scene: 'register' | 'login'): Promise<void> {
return this.post<void>('/user/send-email-code', { email, scene });
static async sendEmailCode(email: string, scene: 'register' | 'login', headers?: Record<string, string>): Promise<void> {
return this.post<void>('/user/send-email-code', { email, scene }, headers ? ({ headers } as unknown as InternalAxiosRequestConfig) : undefined);
}
static async changePassword(request: ChangePasswordRequest): Promise<void> {
+7 -7
View File
@@ -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
}
+2 -2
View File
@@ -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)
}
}
+13 -13
View File
@@ -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
)
+1 -2
View File
@@ -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()
}
}
+6
View File
@@ -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 {
+62 -47
View File
@@ -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
}
+6 -4
View File
@@ -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"`
+11
View File
@@ -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 {
+5 -4
View File
@@ -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)
}
}
+23 -23
View File
@@ -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(&current, "id = ?", source.ID).Error; err != nil {
if err := db.DB(ctx).First(&current, "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(&current).Updates(map[string]any{
return db.DB(ctx).Model(&current).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
}
+1
View File
@@ -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,
}
}