mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-11 17:56:37 +08:00
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:
@@ -1,6 +1,6 @@
|
|||||||
"use client"
|
"use client"
|
||||||
|
|
||||||
import {useEffect, useMemo, useState} from "react"
|
import {useEffect, useMemo, useRef, useState} from "react"
|
||||||
import {useMutation, useQuery} from "@tanstack/react-query"
|
import {useMutation, useQuery} from "@tanstack/react-query"
|
||||||
import {useRouter, useSearchParams} from "next/navigation"
|
import {useRouter, useSearchParams} from "next/navigation"
|
||||||
import {toast} from "sonner"
|
import {toast} from "sonner"
|
||||||
@@ -12,6 +12,7 @@ import {Input} from "@/components/ui/input"
|
|||||||
import {Spinner} from "@/components/ui/spinner"
|
import {Spinner} from "@/components/ui/spinner"
|
||||||
import {Field, FieldGroup, FieldLabel} from "@/components/ui/field"
|
import {Field, FieldGroup, FieldLabel} from "@/components/ui/field"
|
||||||
import {AuthHeading} from "@/components/auth/auth-shell"
|
import {AuthHeading} from "@/components/auth/auth-shell"
|
||||||
|
import {CapWidget} from "@/components/auth/cap-widget"
|
||||||
import services from "@/lib/services"
|
import services from "@/lib/services"
|
||||||
import type {RegisterRequest} from "@/lib/services/auth/types"
|
import type {RegisterRequest} from "@/lib/services/auth/types"
|
||||||
import {safeRedirectTarget} from "@/lib/utils"
|
import {safeRedirectTarget} from "@/lib/utils"
|
||||||
@@ -66,6 +67,33 @@ export function RegisterForm() {
|
|||||||
|
|
||||||
const emailRegisterEnabled = configBool(publicConfigQuery.data?.email_register_verification_enabled, false)
|
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
|
// Redirect to login if registration is closed
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
if (publicConfigQuery.isSuccess && !registrationEnabled) {
|
if (publicConfigQuery.isSuccess && !registrationEnabled) {
|
||||||
@@ -75,7 +103,15 @@ export function RegisterForm() {
|
|||||||
}, [publicConfigQuery.isSuccess, registrationEnabled, router])
|
}, [publicConfigQuery.isSuccess, registrationEnabled, router])
|
||||||
|
|
||||||
const registerMutation = useMutation({
|
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) => {
|
onSuccess: (user) => {
|
||||||
setUser(user)
|
setUser(user)
|
||||||
router.replace(redirectTarget)
|
router.replace(redirectTarget)
|
||||||
@@ -83,20 +119,53 @@ export function RegisterForm() {
|
|||||||
},
|
},
|
||||||
onError: (error: Error) => {
|
onError: (error: Error) => {
|
||||||
setErrorMessage(error.message || "注册失败,请重试")
|
setErrorMessage(error.message || "注册失败,请重试")
|
||||||
|
if (capEnabled) {
|
||||||
|
capTokenRef.current = null
|
||||||
|
setCapReady(false)
|
||||||
|
setCapResetKey((key) => key + 1)
|
||||||
|
}
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
|
|
||||||
const sendRegisterCodeMutation = useMutation({
|
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: () => {
|
onSuccess: () => {
|
||||||
setRegisterCooldown(60)
|
setRegisterCooldown(60)
|
||||||
toast.success("验证码已发送至您的邮箱,请查收")
|
toast.success("验证码已发送至您的邮箱,请查收")
|
||||||
|
if (capEnabled) {
|
||||||
|
setCapScope('register')
|
||||||
|
capTokenRef.current = null
|
||||||
|
setCapReady(false)
|
||||||
|
setCapResetKey((key) => key + 1)
|
||||||
|
}
|
||||||
},
|
},
|
||||||
onError: (error: Error) => {
|
onError: (error: Error) => {
|
||||||
toast.error(error.message || "发送验证码失败,请重试")
|
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 handleSendRegisterCode = () => {
|
||||||
const trimmedEmail = email.trim()
|
const trimmedEmail = email.trim()
|
||||||
if (!trimmedEmail) {
|
if (!trimmedEmail) {
|
||||||
@@ -108,6 +177,14 @@ export function RegisterForm() {
|
|||||||
toast.error("请输入有效的邮箱地址")
|
toast.error("请输入有效的邮箱地址")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if (capEnabled && capScope === 'send_email_code' && !capReady) {
|
||||||
|
toast.error(
|
||||||
|
capAutoSolve
|
||||||
|
? "人机验证尚未完成,请稍候…"
|
||||||
|
: "请先点击「开始验证」完成人机验证",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
sendRegisterCodeMutation.mutate(trimmedEmail)
|
sendRegisterCodeMutation.mutate(trimmedEmail)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -135,6 +212,14 @@ export function RegisterForm() {
|
|||||||
toast.error("验证码不能为空")
|
toast.error("验证码不能为空")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if (capEnabled && capScope === 'register' && !capReady) {
|
||||||
|
toast.error(
|
||||||
|
capAutoSolve
|
||||||
|
? "人机验证尚未完成,请稍候…"
|
||||||
|
: "请先点击「开始验证」完成人机验证",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
registerMutation.mutate({
|
registerMutation.mutate({
|
||||||
username: username.trim(),
|
username: username.trim(),
|
||||||
password,
|
password,
|
||||||
@@ -235,7 +320,7 @@ export function RegisterForm() {
|
|||||||
type="button"
|
type="button"
|
||||||
variant="outline"
|
variant="outline"
|
||||||
onClick={handleSendRegisterCode}
|
onClick={handleSendRegisterCode}
|
||||||
disabled={registerCooldown > 0 || sendRegisterCodeMutation.isPending}
|
disabled={sendCodeDisabled}
|
||||||
className="h-10 w-[120px] text-xs [@media(max-height:700px)]:h-9"
|
className="h-10 w-[120px] text-xs [@media(max-height:700px)]:h-9"
|
||||||
>
|
>
|
||||||
{registerCooldown > 0 ? `${registerCooldown}秒后重发` : "获取验证码"}
|
{registerCooldown > 0 ? `${registerCooldown}秒后重发` : "获取验证码"}
|
||||||
@@ -245,6 +330,16 @@ export function RegisterForm() {
|
|||||||
)}
|
)}
|
||||||
</FieldGroup>
|
</FieldGroup>
|
||||||
|
|
||||||
|
{capEnabled && (
|
||||||
|
<CapWidget
|
||||||
|
key={capResetKey}
|
||||||
|
scope={capScope}
|
||||||
|
onToken={handleCapToken}
|
||||||
|
onError={handleCapError}
|
||||||
|
autoStart={capAutoSolve}
|
||||||
|
/>
|
||||||
|
)}
|
||||||
|
|
||||||
{errorMessage ? (
|
{errorMessage ? (
|
||||||
<div className="rounded-lg border border-destructive/30 bg-destructive/5 px-3 py-2 text-sm text-destructive">
|
<div className="rounded-lg border border-destructive/30 bg-destructive/5 px-3 py-2 text-sm text-destructive">
|
||||||
{errorMessage}
|
{errorMessage}
|
||||||
@@ -256,7 +351,7 @@ export function RegisterForm() {
|
|||||||
className="h-10 w-full [@media(max-height:700px)]:h-9"
|
className="h-10 w-full [@media(max-height:700px)]:h-9"
|
||||||
variant="auth"
|
variant="auth"
|
||||||
onClick={handleRegister}
|
onClick={handleRegister}
|
||||||
disabled={registerMutation.isPending}
|
disabled={registerDisabled}
|
||||||
>
|
>
|
||||||
{registerMutation.isPending ? (
|
{registerMutation.isPending ? (
|
||||||
<>
|
<>
|
||||||
|
|||||||
@@ -109,12 +109,12 @@ export class AuthService extends BaseService {
|
|||||||
return this.post<User>('/user/login', request, headers ? ({ headers } as unknown as InternalAxiosRequestConfig) : undefined);
|
return this.post<User>('/user/login', request, headers ? ({ headers } as unknown as InternalAxiosRequestConfig) : undefined);
|
||||||
}
|
}
|
||||||
|
|
||||||
static async register(request: RegisterRequest): Promise<User> {
|
static async register(request: RegisterRequest, headers?: Record<string, string>): Promise<User> {
|
||||||
return this.post<User>('/user/register', request);
|
return this.post<User>('/user/register', request, headers ? ({ headers } as unknown as InternalAxiosRequestConfig) : undefined);
|
||||||
}
|
}
|
||||||
|
|
||||||
static async sendEmailCode(email: string, scene: 'register' | 'login'): Promise<void> {
|
static async sendEmailCode(email: string, scene: 'register' | 'login', headers?: Record<string, string>): Promise<void> {
|
||||||
return this.post<void>('/user/send-email-code', { email, scene });
|
return this.post<void>('/user/send-email-code', { email, scene }, headers ? ({ headers } as unknown as InternalAxiosRequestConfig) : undefined);
|
||||||
}
|
}
|
||||||
|
|
||||||
static async changePassword(request: ChangePasswordRequest): Promise<void> {
|
static async changePassword(request: ChangePasswordRequest): Promise<void> {
|
||||||
|
|||||||
@@ -45,7 +45,7 @@ type ToggleAuthSourceRequest struct {
|
|||||||
// @Failure 500 {object} util.ResponseAny "内部错误"
|
// @Failure 500 {object} util.ResponseAny "内部错误"
|
||||||
// @Router /api/v1/admin/auth-sources [get]
|
// @Router /api/v1/admin/auth-sources [get]
|
||||||
func ListAuthSources(c *gin.Context) {
|
func ListAuthSources(c *gin.Context) {
|
||||||
sources, err := model.GetAuthSources()
|
sources, err := model.GetAuthSources(c.Request.Context())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||||
return
|
return
|
||||||
@@ -84,7 +84,7 @@ func CreateAuthSource(c *gin.Context) {
|
|||||||
Scopes: req.Scopes,
|
Scopes: req.Scopes,
|
||||||
IconURL: req.IconURL,
|
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()))
|
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -133,11 +133,11 @@ func UpdateAuthSource(c *gin.Context) {
|
|||||||
IconURL: req.IconURL,
|
IconURL: req.IconURL,
|
||||||
}
|
}
|
||||||
keepSecret := source.ClientSecret == ""
|
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()))
|
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
updated, err := model.GetAuthSourceByID(id)
|
updated, err := model.GetAuthSourceByID(c.Request.Context(), id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||||
return
|
return
|
||||||
@@ -173,7 +173,7 @@ func ToggleAuthSource(c *gin.Context) {
|
|||||||
return
|
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()))
|
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -198,7 +198,7 @@ func DeleteAuthSource(c *gin.Context) {
|
|||||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||||
return
|
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()))
|
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -210,7 +210,7 @@ func parseSourceID(c *gin.Context) (uint64, error) {
|
|||||||
if raw == "" {
|
if raw == "" {
|
||||||
return 0, errors.New(admin.InvalidAuthSourceID)
|
return 0, errors.New(admin.InvalidAuthSourceID)
|
||||||
}
|
}
|
||||||
source, err := model.GetAuthSourceByName(raw)
|
source, err := model.GetAuthSourceByName(c.Request.Context(), raw)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
return source.ID, nil
|
return source.ID, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -29,8 +29,8 @@ func LogForAudit(ctx context.Context, user *model.User, c *gin.Context) {
|
|||||||
auditJSON, err := json.Marshal(auditLog)
|
auditJSON, err := json.Marshal(auditLog)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorF(ctx, "[LoginRequiredAudit] marshal failed: %v", err)
|
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 {
|
} else {
|
||||||
logger.InfoF(ctx, "[LoginRequiredAudit] %s", auditJSON)
|
logger.DebugF(ctx, "[LoginRequiredAudit] %s", auditJSON)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+13
-13
@@ -6,17 +6,17 @@ package oauth
|
|||||||
|
|
||||||
// OAuth 认证相关错误消息
|
// OAuth 认证相关错误消息
|
||||||
const (
|
const (
|
||||||
InvalidState = "非法登录请求"
|
errInvalidState = "非法登录请求"
|
||||||
IDTokenVerifyFailed = "ID Token 验证失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
errIDTokenVerifyFailed = "ID Token 验证失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||||
IDTokenVerifyFailedFormat = "%s: %w"
|
errIDTokenVerifyFailedFormat = "%s: %w"
|
||||||
NonceMismatch = "nonce 不匹配,可能存在重放攻击"
|
errNonceMismatch = "nonce 不匹配,可能存在重放攻击"
|
||||||
NoActiveAuthSource = "未配置可用认证源"
|
errNoActiveAuthSource = "未配置可用认证源"
|
||||||
ServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试"
|
errServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试"
|
||||||
AuthSourceRequired = "认证源不能为空"
|
errAuthSourceRequired = "认证源不能为空"
|
||||||
DiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
|
errDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
|
||||||
UsernameGenerateFailed = "无法生成可用用户名"
|
errUsernameGenerateFailed = "无法生成可用用户名"
|
||||||
UsernameFromSourceFailed = "无法从认证源获取用户名"
|
errUsernameFromSourceFailed = "无法从认证源获取用户名"
|
||||||
AuthSourceDisabled = "认证源未启用"
|
errAuthSourceDisabled = "认证源未启用"
|
||||||
InvalidExternalAccountBindingID = "绑定记录 ID 无效"
|
errInvalidExternalAccountBindingID = "绑定记录 ID 无效"
|
||||||
TokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -112,10 +112,9 @@ func LoginRequired() gin.HandlerFunc {
|
|||||||
func DisallowTokenAuth() gin.HandlerFunc {
|
func DisallowTokenAuth() gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
if tokenAuth, _ := util.GetFromContext[bool](c, TokenAuthKey); tokenAuth {
|
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
|
return
|
||||||
}
|
}
|
||||||
c.Next()
|
c.Next()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ package oauth
|
|||||||
import (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/logger"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
"github.com/Rain-kl/Wavelet/internal/util"
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
@@ -89,6 +90,11 @@ func UserInfo(c *gin.Context) {
|
|||||||
// @Router /api/v1/oauth/logout [get]
|
// @Router /api/v1/oauth/logout [get]
|
||||||
func Logout(c *gin.Context) {
|
func Logout(c *gin.Context) {
|
||||||
session := sessions.Default(c)
|
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.Options(util.GetSessionOptions(-1))
|
||||||
session.Clear()
|
session.Clear()
|
||||||
if err := session.Save(); err != nil {
|
if err := session.Save(); err != nil {
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/common"
|
"github.com/Rain-kl/Wavelet/internal/common"
|
||||||
"github.com/Rain-kl/Wavelet/internal/config"
|
"github.com/Rain-kl/Wavelet/internal/config"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"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/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
"github.com/Rain-kl/Wavelet/internal/util"
|
||||||
"github.com/coreos/go-oidc/v3/oidc"
|
"github.com/coreos/go-oidc/v3/oidc"
|
||||||
@@ -98,30 +99,28 @@ func isOIDCLoginEnabled(ctx context.Context) bool {
|
|||||||
return enabled
|
return enabled
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) {
|
||||||
|
|
||||||
func resolveAuthSource(sourceName string) (*model.AuthSource, error) {
|
|
||||||
name := strings.TrimSpace(strings.ToLower(sourceName))
|
name := strings.TrimSpace(strings.ToLower(sourceName))
|
||||||
if name == "" {
|
if name == "" {
|
||||||
sources, err := model.GetActiveAuthSources()
|
sources, err := model.GetActiveAuthSources(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if len(sources) == 0 {
|
if len(sources) == 0 {
|
||||||
return nil, errors.New(NoActiveAuthSource)
|
return nil, errors.New(errNoActiveAuthSource)
|
||||||
}
|
}
|
||||||
return &sources[0], nil
|
return &sources[0], nil
|
||||||
}
|
}
|
||||||
return model.GetAuthSourceByName(name)
|
return model.GetAuthSourceByName(ctx, name)
|
||||||
}
|
}
|
||||||
|
|
||||||
func activeLoginSources() []AuthSourceView {
|
func activeLoginSources(ctx context.Context) []AuthSourceView {
|
||||||
enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyOIDCLoginEnabled)
|
enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
|
||||||
if err == nil && !enabled {
|
if err == nil && !enabled {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
dbSources, err := model.GetActiveAuthSources()
|
dbSources, err := model.GetActiveAuthSources(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -143,18 +142,18 @@ func activeLoginSources() []AuthSourceView {
|
|||||||
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
|
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
|
||||||
var sc model.SystemConfig
|
var sc model.SystemConfig
|
||||||
if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err != nil || strings.TrimSpace(sc.Value) == "" {
|
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
|
return strings.TrimRight(sc.Value, "/") + "/login", nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
|
func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
|
||||||
if source == nil {
|
if source == nil {
|
||||||
return nil, nil, errors.New(AuthSourceRequired)
|
return nil, nil, errors.New(errAuthSourceRequired)
|
||||||
}
|
}
|
||||||
|
|
||||||
if source.OpenIDDiscoveryURL == "" {
|
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)
|
// 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) {
|
func uniqueUsername(ctx context.Context, base string) (string, error) {
|
||||||
candidate := strings.TrimSpace(base)
|
base = strings.TrimSpace(base)
|
||||||
if candidate == "" {
|
if base == "" {
|
||||||
candidate = "user"
|
base = "user"
|
||||||
}
|
}
|
||||||
for i := 0; i < 1000; i++ {
|
|
||||||
var count int64
|
var existingUsernames []string
|
||||||
if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", candidate).Count(&count).Error; err != nil {
|
if err := db.DB(ctx).Model(&model.User{}).
|
||||||
return "", err
|
Where("username = ? OR username LIKE ?", base, base+"-%").
|
||||||
}
|
Pluck("username", &existingUsernames).Error; err != nil {
|
||||||
if count == 0 {
|
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
|
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) {
|
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)
|
idToken, verifyErr := verifier.Verify(ctx, rawIDToken)
|
||||||
if verifyErr != nil {
|
if verifyErr != nil {
|
||||||
return fmt.Errorf(IDTokenVerifyFailedFormat, IDTokenVerifyFailed, verifyErr)
|
return fmt.Errorf(errIDTokenVerifyFailedFormat, errIDTokenVerifyFailed, verifyErr)
|
||||||
}
|
}
|
||||||
if nonce != "" && idToken.Nonce != nonce {
|
if nonce != "" && idToken.Nonce != nonce {
|
||||||
return errors.New(NonceMismatch)
|
return errors.New(errNonceMismatch)
|
||||||
}
|
}
|
||||||
if claimsErr := idToken.Claims(userInfo); claimsErr != nil {
|
if claimsErr := idToken.Claims(userInfo); claimsErr != nil {
|
||||||
return claimsErr
|
return claimsErr
|
||||||
@@ -316,7 +332,7 @@ func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error {
|
|||||||
userInfo.Username = userInfo.Sub
|
userInfo.Username = userInfo.Sub
|
||||||
}
|
}
|
||||||
if userInfo.Username == "" {
|
if userInfo.Username == "" {
|
||||||
return errors.New(UsernameFromSourceFailed)
|
return errors.New(errUsernameFromSourceFailed)
|
||||||
}
|
}
|
||||||
if userInfo.Name == "" {
|
if userInfo.Name == "" {
|
||||||
userInfo.Name = userInfo.Username
|
userInfo.Name = userInfo.Username
|
||||||
@@ -344,7 +360,7 @@ func buildCallbackResult(user *model.User, status string) OAuthCallbackResult {
|
|||||||
// @Success 200 {object} util.ResponseAny{data=[]oauth.AuthSourceView} "登录源列表"
|
// @Success 200 {object} util.ResponseAny{data=[]oauth.AuthSourceView} "登录源列表"
|
||||||
// @Router /api/v1/oauth/sources [get]
|
// @Router /api/v1/oauth/sources [get]
|
||||||
func GetLoginSources(c *gin.Context) {
|
func GetLoginSources(c *gin.Context) {
|
||||||
c.JSON(http.StatusOK, util.OK(activeLoginSources()))
|
c.JSON(http.StatusOK, util.OK(activeLoginSources(c.Request.Context())))
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetLoginURL 获取登录授权地址
|
// GetLoginURL 获取登录授权地址
|
||||||
@@ -355,27 +371,26 @@ func GetLoginSources(c *gin.Context) {
|
|||||||
// @Param source query string false "认证源名称,为空使用第一个启用的认证源"
|
// @Param source query string false "认证源名称,为空使用第一个启用的认证源"
|
||||||
// @Success 200 {object} util.ResponseAny{data=oauth.OAuthAuthorizeResponse} "授权 URL"
|
// @Success 200 {object} util.ResponseAny{data=oauth.OAuthAuthorizeResponse} "授权 URL"
|
||||||
// @Failure 400 {object} util.ResponseAny "认证源不存在或未配置"
|
// @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]
|
// @Router /api/v1/oauth/login [get]
|
||||||
func GetLoginURL(c *gin.Context) {
|
func GetLoginURL(c *gin.Context) {
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
if !isOIDCLoginEnabled(ctx) {
|
if !isOIDCLoginEnabled(ctx) {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled))
|
c.JSON(http.StatusBadRequest, util.Err(errAuthSourceDisabled))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
source, err := resolveAuthSource(c.Query("source"))
|
source, err := resolveAuthSource(ctx, c.Query("source"))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if !source.IsActive {
|
if !source.IsActive {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled))
|
c.JSON(http.StatusBadRequest, util.Err(errAuthSourceDisabled))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
session := sessions.Default(c)
|
session := sessions.Default(c)
|
||||||
token, isNew := ensureSessionToken(session)
|
token, isNew := ensureSessionToken(session)
|
||||||
if isNew {
|
if isNew {
|
||||||
@@ -404,7 +419,6 @@ func GetLoginURL(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
|
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
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) {
|
func Authorize(c *gin.Context) {
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
if !isOIDCLoginEnabled(ctx) {
|
if !isOIDCLoginEnabled(ctx) {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled))
|
c.JSON(http.StatusBadRequest, util.Err(errAuthSourceDisabled))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
source, err := resolveAuthSource(c.Param("source"))
|
source, err := resolveAuthSource(ctx, c.Param("source"))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if !source.IsActive {
|
if !source.IsActive {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled))
|
c.JSON(http.StatusBadRequest, util.Err(errAuthSourceDisabled))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose")))
|
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))
|
stateKey := db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State))
|
||||||
payloadRaw, err := db.Redis.Get(ctx, stateKey).Result()
|
payloadRaw, err := db.Redis.Get(ctx, stateKey).Result()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(InvalidState))
|
c.JSON(http.StatusBadRequest, util.Err(errInvalidState))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
_ = db.Redis.Del(ctx, stateKey)
|
_ = db.Redis.Del(ctx, stateKey)
|
||||||
@@ -560,25 +574,22 @@ func Callback(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
if !isOIDCLoginEnabled(ctx) {
|
if !isOIDCLoginEnabled(ctx) {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled))
|
c.JSON(http.StatusBadRequest, util.Err(errAuthSourceDisabled))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
source, err := resolveAuthSource(payload.SourceName)
|
source, err := resolveAuthSource(ctx, payload.SourceName)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if !source.IsActive {
|
if !source.IsActive {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled))
|
c.JSON(http.StatusBadRequest, util.Err(errAuthSourceDisabled))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
redirectURL, err := getFrontendLoginRedirectURL(ctx)
|
redirectURL, err := getFrontendLoginRedirectURL(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
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()))
|
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||||
return
|
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")))
|
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
|
return model.User{}, false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
|
username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
|
||||||
if uniqueErr != nil {
|
if uniqueErr != nil {
|
||||||
c.JSON(http.StatusInternalServerError, util.Err(uniqueErr.Error()))
|
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()))
|
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||||
return model.User{}, false
|
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
|
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]
|
// @Router /api/v1/oauth/external-accounts [get]
|
||||||
func ListExternalAccounts(c *gin.Context) {
|
func ListExternalAccounts(c *gin.Context) {
|
||||||
userID := GetUserIDFromContext(c)
|
userID := GetUserIDFromContext(c)
|
||||||
accounts, err := model.ListExternalAccountsByUserID(userID)
|
accounts, err := model.ListExternalAccountsByUserID(c.Request.Context(), userID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||||
return
|
return
|
||||||
@@ -743,10 +758,10 @@ func DeleteExternalAccount(c *gin.Context) {
|
|||||||
rawID := strings.TrimSpace(c.Param("id"))
|
rawID := strings.TrimSpace(c.Param("id"))
|
||||||
id, err := strconv.ParseUint(rawID, 10, 64)
|
id, err := strconv.ParseUint(rawID, 10, 64)
|
||||||
if err != nil || id == 0 {
|
if err != nil || id == 0 {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(InvalidExternalAccountBindingID))
|
c.JSON(http.StatusBadRequest, util.Err(errInvalidExternalAccountBindingID))
|
||||||
return
|
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()))
|
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -45,7 +45,7 @@ func isEmailRegisterVerificationEnabled(ctx context.Context) bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func isSMTPConfigured(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
|
var scHost model.SystemConfig
|
||||||
if err := scHost.GetByKey(ctx, model.ConfigKeySMTPHost); err == nil {
|
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 {
|
if err := scUser.GetByKey(ctx, model.ConfigKeySMTPUsername); err == nil {
|
||||||
username = scUser.Value
|
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) {
|
func generateVerificationCode() (string, error) {
|
||||||
@@ -241,8 +245,6 @@ func validateRegisterEmailVerification(ctx context.Context, req *registerRequest
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
type updateProfileRequest struct {
|
type updateProfileRequest struct {
|
||||||
Nickname string `json:"nickname"`
|
Nickname string `json:"nickname"`
|
||||||
Email string `json:"email"`
|
Email string `json:"email"`
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/config"
|
"github.com/Rain-kl/Wavelet/internal/config"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
"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/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
"github.com/Rain-kl/Wavelet/internal/util"
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
@@ -124,10 +125,12 @@ func Login(c *gin.Context) {
|
|||||||
var user model.User
|
var user model.User
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
if err := db.DB(ctx).Where("username = ? OR email = ?", req.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 {
|
||||||
|
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))
|
c.JSON(http.StatusOK, util.Err(errUsernameOrPasswordWrong))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if !user.IsActive {
|
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))
|
c.JSON(http.StatusOK, util.Err(common.BannedAccount))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -136,6 +139,7 @@ func Login(c *gin.Context) {
|
|||||||
isPlaintext := !user.IsPasswordEncrypted()
|
isPlaintext := !user.IsPasswordEncrypted()
|
||||||
|
|
||||||
if !user.CheckPassword(req.Password) {
|
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))
|
c.JSON(http.StatusOK, util.Err(errUsernameOrPasswordWrong))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -165,6 +169,8 @@ func Login(c *gin.Context) {
|
|||||||
return
|
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)))
|
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]
|
// @Router /api/v1/user/logout [get]
|
||||||
func Logout(c *gin.Context) {
|
func Logout(c *gin.Context) {
|
||||||
session := sessions.Default(c)
|
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.Options(util.GetSessionOptions(-1))
|
||||||
session.Clear()
|
session.Clear()
|
||||||
if err := session.Save(); err != nil {
|
if err := session.Save(); err != nil {
|
||||||
|
|||||||
@@ -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 {
|
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPUsername).Update("value", "smtpuser").Error; err != nil {
|
||||||
t.Fatalf("set SMTP username failed: %v", err)
|
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
|
// Invalidate the system config cache in Redis
|
||||||
if err := db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err(); err != nil {
|
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 {
|
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||||
t.Fatalf("decode response failed: %v", err)
|
t.Fatalf("decode response failed: %v", err)
|
||||||
}
|
}
|
||||||
if resp.ErrorMsg != oauth.TokenAuthNotAllowed {
|
if resp.ErrorMsg != oauth.ErrTokenAuthNotAllowed {
|
||||||
t.Errorf("expected error message %q, got %q", oauth.TokenAuthNotAllowed, resp.ErrorMsg)
|
t.Errorf("expected error message %q, got %q", oauth.ErrTokenAuthNotAllowed, resp.ErrorMsg)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 4. Test that accessing using a Session succeeds
|
// 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)
|
t.Errorf("expected access token to be revoked (401), got %d", wTokenAfter.Code)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -118,9 +118,9 @@ func (source *AuthSource) Sanitize() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// GetAuthSources 获取所有认证源(已脱敏)
|
// GetAuthSources 获取所有认证源(已脱敏)
|
||||||
func GetAuthSources() ([]AuthSource, error) {
|
func GetAuthSources(ctx context.Context) ([]AuthSource, error) {
|
||||||
var sources []AuthSource
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
for i := range sources {
|
for i := range sources {
|
||||||
@@ -130,9 +130,9 @@ func GetAuthSources() ([]AuthSource, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// GetActiveAuthSources 获取所有已启用的认证源(已脱敏)
|
// GetActiveAuthSources 获取所有已启用的认证源(已脱敏)
|
||||||
func GetActiveAuthSources() ([]AuthSource, error) {
|
func GetActiveAuthSources(ctx context.Context) ([]AuthSource, error) {
|
||||||
var sources []AuthSource
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
for i := range sources {
|
for i := range sources {
|
||||||
@@ -142,12 +142,12 @@ func GetActiveAuthSources() ([]AuthSource, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// GetAuthSourceByID 根据 ID 获取认证源
|
// GetAuthSourceByID 根据 ID 获取认证源
|
||||||
func GetAuthSourceByID(id uint64) (*AuthSource, error) {
|
func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
|
||||||
if id == 0 {
|
if id == 0 {
|
||||||
return nil, errors.New(errAuthSourceIDRequired)
|
return nil, errors.New(errAuthSourceIDRequired)
|
||||||
}
|
}
|
||||||
var source AuthSource
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
source.ClientSecretConfigured = source.ClientSecret != ""
|
source.ClientSecretConfigured = source.ClientSecret != ""
|
||||||
@@ -155,13 +155,13 @@ func GetAuthSourceByID(id uint64) (*AuthSource, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// GetAuthSourceByName 根据名称获取认证源(名称比较不区分大小写)
|
// GetAuthSourceByName 根据名称获取认证源(名称比较不区分大小写)
|
||||||
func GetAuthSourceByName(name string) (*AuthSource, error) {
|
func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error) {
|
||||||
name = strings.TrimSpace(name)
|
name = strings.TrimSpace(name)
|
||||||
if name == "" {
|
if name == "" {
|
||||||
return nil, errors.New(errAuthSourceNameRequired)
|
return nil, errors.New(errAuthSourceNameRequired)
|
||||||
}
|
}
|
||||||
var source AuthSource
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
source.ClientSecretConfigured = source.ClientSecret != ""
|
source.ClientSecretConfigured = source.ClientSecret != ""
|
||||||
@@ -169,20 +169,20 @@ func GetAuthSourceByName(name string) (*AuthSource, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// CreateAuthSource 创建认证源
|
// CreateAuthSource 创建认证源
|
||||||
func CreateAuthSource(source *AuthSource) error {
|
func CreateAuthSource(ctx context.Context, source *AuthSource) error {
|
||||||
if err := source.Validate(); err != nil {
|
if err := source.Validate(); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return db.DB(context.Background()).Create(source).Error
|
return db.DB(ctx).Create(source).Error
|
||||||
}
|
}
|
||||||
|
|
||||||
// UpdateAuthSource 更新认证源,keepSecret 为 true 时保留原密钥
|
// UpdateAuthSource 更新认证源,keepSecret 为 true 时保留原密钥
|
||||||
func UpdateAuthSource(source *AuthSource, keepSecret bool) error {
|
func UpdateAuthSource(ctx context.Context, source *AuthSource, keepSecret bool) error {
|
||||||
if source.ID == 0 {
|
if source.ID == 0 {
|
||||||
return errors.New(errAuthSourceIDRequired)
|
return errors.New(errAuthSourceIDRequired)
|
||||||
}
|
}
|
||||||
var current AuthSource
|
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
|
return err
|
||||||
}
|
}
|
||||||
if keepSecret {
|
if keepSecret {
|
||||||
@@ -191,7 +191,7 @@ func UpdateAuthSource(source *AuthSource, keepSecret bool) error {
|
|||||||
if err := source.Validate(); err != nil {
|
if err := source.Validate(); err != nil {
|
||||||
return err
|
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,
|
"name": source.Name,
|
||||||
"type": source.Type,
|
"type": source.Type,
|
||||||
"display_name": source.DisplayName,
|
"display_name": source.DisplayName,
|
||||||
@@ -205,8 +205,8 @@ func UpdateAuthSource(source *AuthSource, keepSecret bool) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// ToggleAuthSource 切换认证源启用状态
|
// ToggleAuthSource 切换认证源启用状态
|
||||||
func ToggleAuthSource(id uint64, isActive bool) error {
|
func ToggleAuthSource(ctx context.Context, id uint64, isActive bool) error {
|
||||||
source, err := GetAuthSourceByID(id)
|
source, err := GetAuthSourceByID(ctx, id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -214,15 +214,15 @@ func ToggleAuthSource(id uint64, isActive bool) error {
|
|||||||
if err := source.Validate(); err != nil {
|
if err := source.Validate(); err != nil {
|
||||||
return err
|
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 删除认证源及其关联的外部帐号绑定
|
// DeleteAuthSource 删除认证源及其关联的外部帐号绑定
|
||||||
func DeleteAuthSource(id uint64) error {
|
func DeleteAuthSource(ctx context.Context, id uint64) error {
|
||||||
if id == 0 {
|
if id == 0 {
|
||||||
return errors.New(errAuthSourceIDRequired)
|
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 {
|
if err := tx.Where("auth_source_id = ?", id).Delete(&ExternalAccount{}).Error; err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -268,12 +268,12 @@ func BindExternalAccount(ctx context.Context, account *ExternalAccount) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// ListExternalAccountsByUserID 获取指定用户的所有外部帐号绑定视图
|
// ListExternalAccountsByUserID 获取指定用户的所有外部帐号绑定视图
|
||||||
func ListExternalAccountsByUserID(userID uint64) ([]ExternalAccountView, error) {
|
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccountView, error) {
|
||||||
if userID == 0 {
|
if userID == 0 {
|
||||||
return nil, errors.New(errUserIDRequired)
|
return nil, errors.New(errUserIDRequired)
|
||||||
}
|
}
|
||||||
var accounts []ExternalAccount
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
views := make([]ExternalAccountView, 0, len(accounts))
|
views := make([]ExternalAccountView, 0, len(accounts))
|
||||||
@@ -284,7 +284,7 @@ func ListExternalAccountsByUserID(userID uint64) ([]ExternalAccountView, error)
|
|||||||
sourceType = "oidc"
|
sourceType = "oidc"
|
||||||
label = "历史认证源"
|
label = "历史认证源"
|
||||||
} else {
|
} else {
|
||||||
source, err := GetAuthSourceByID(account.AuthSourceID)
|
source, err := GetAuthSourceByID(ctx, account.AuthSourceID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -310,9 +310,9 @@ func ListExternalAccountsByUserID(userID uint64) ([]ExternalAccountView, error)
|
|||||||
}
|
}
|
||||||
|
|
||||||
// DeleteExternalAccountForUser 删除指定用户的外部帐号绑定
|
// DeleteExternalAccountForUser 删除指定用户的外部帐号绑定
|
||||||
func DeleteExternalAccountForUser(id uint64, userID uint64) error {
|
func DeleteExternalAccountForUser(ctx context.Context, id uint64, userID uint64) error {
|
||||||
if id == 0 || userID == 0 {
|
if id == 0 || userID == 0 {
|
||||||
return errors.New(errExternalAccountBindingIDRequired)
|
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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ func GetSessionOptions(maxAge int) sessions.Options {
|
|||||||
MaxAge: maxAge,
|
MaxAge: maxAge,
|
||||||
HttpOnly: config.Config.App.SessionHTTPOnly,
|
HttpOnly: config.Config.App.SessionHTTPOnly,
|
||||||
Secure: config.Config.App.SessionSecure,
|
Secure: config.Config.App.SessionSecure,
|
||||||
|
SameSite: http.SameSiteLaxMode,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user