mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-30 22:26:38 +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"
|
||||
|
||||
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 ? (
|
||||
<>
|
||||
|
||||
@@ -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> {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user