feat: OAuth 绑定安全增强及限流配置优化

- 新增一次性绑定令牌机制,避免在 URL 中暴露 access_token
- 绑定流程改为新标签页打开,完成后自动刷新状态
- 放宽认证相关接口的 IP 限流配置
- 登录添加 429 限流错误的友好提示
- OAuth 绑定列表添加 Provider 图标显示

Close #106
This commit is contained in:
fawney19
2026-01-19 19:34:31 +08:00
parent b53e3f14d6
commit 427f173b38
10 changed files with 259 additions and 42 deletions

View File

@@ -96,6 +96,11 @@ export const oauthApi = {
return response.data.links || [] return response.data.links || []
}, },
async createBindToken(providerType: string): Promise<string> {
const response = await apiClient.post<{ bind_token: string }>(`/api/user/oauth/${providerType}/bind-token`)
return response.data.bind_token
},
async unbind(providerType: string): Promise<{ message: string }> { async unbind(providerType: string): Promise<{ message: string }> {
const response = await apiClient.delete<{ message: string }>(`/api/user/oauth/${providerType}`) const response = await apiClient.delete<{ message: string }>(`/api/user/oauth/${providerType}`)
return response.data return response.data

View File

@@ -242,17 +242,7 @@ import RegisterDialog from './RegisterDialog.vue'
import { authApi } from '@/api/auth' import { authApi } from '@/api/auth'
import { oauthApi, type OAuthProviderInfo } from '@/api/oauth' import { oauthApi, type OAuthProviderInfo } from '@/api/oauth'
import { getApiUrl } from '@/utils/url' import { getApiUrl } from '@/utils/url'
import { getOAuthIcon } from '@/utils/oauth-icons'
// OAuth provider icons
const OAUTH_ICONS: Record<string, string> = {
linuxdo: `<svg viewBox="0 0 120 120" xmlns="http://www.w3.org/2000/svg"><clipPath id="ld"><circle cx="60" cy="60" r="47"/></clipPath><circle fill="#f0f0f0" cx="60" cy="60" r="50"/><rect fill="#1c1c1e" clip-path="url(#ld)" x="10" y="10" width="100" height="30"/><rect fill="#f0f0f0" clip-path="url(#ld)" x="10" y="40" width="100" height="40"/><rect fill="#ffb003" clip-path="url(#ld)" x="10" y="80" width="100" height="30"/></svg>`,
github: `<svg viewBox="0 0 24 24" fill="currentColor"><path d="M12 0c-6.626 0-12 5.373-12 12 0 5.302 3.438 9.8 8.207 11.387.599.111.793-.261.793-.577v-2.234c-3.338.726-4.033-1.416-4.033-1.416-.546-1.387-1.333-1.756-1.333-1.756-1.089-.745.083-.729.083-.729 1.205.084 1.839 1.237 1.839 1.237 1.07 1.834 2.807 1.304 3.492.997.107-.775.418-1.305.762-1.604-2.665-.305-5.467-1.334-5.467-5.931 0-1.311.469-2.381 1.236-3.221-.124-.303-.535-1.524.117-3.176 0 0 1.008-.322 3.301 1.23.957-.266 1.983-.399 3.003-.404 1.02.005 2.047.138 3.006.404 2.291-1.552 3.297-1.23 3.297-1.23.653 1.653.242 2.874.118 3.176.77.84 1.235 1.911 1.235 3.221 0 4.609-2.807 5.624-5.479 5.921.43.372.823 1.102.823 2.222v3.293c0 .319.192.694.801.576 4.765-1.589 8.199-6.086 8.199-11.386 0-6.627-5.373-12-12-12z"/></svg>`,
google: `<svg viewBox="0 0 24 24"><path fill="#4285F4" d="M22.56 12.25c0-.78-.07-1.53-.2-2.25H12v4.26h5.92c-.26 1.37-1.04 2.53-2.21 3.31v2.77h3.57c2.08-1.92 3.28-4.74 3.28-8.09z"/><path fill="#34A853" d="M12 23c2.97 0 5.46-.98 7.28-2.66l-3.57-2.77c-.98.66-2.23 1.06-3.71 1.06-2.86 0-5.29-1.93-6.16-4.53H2.18v2.84C3.99 20.53 7.7 23 12 23z"/><path fill="#FBBC05" d="M5.84 14.09c-.22-.66-.35-1.36-.35-2.09s.13-1.43.35-2.09V7.07H2.18C1.43 8.55 1 10.22 1 12s.43 3.45 1.18 4.93l2.85-2.22.81-.62z"/><path fill="#EA4335" d="M12 5.38c1.62 0 3.06.56 4.21 1.64l3.15-3.15C17.45 2.09 14.97 1 12 1 7.7 1 3.99 3.47 2.18 7.07l3.66 2.84c.87-2.6 3.3-4.53 6.16-4.53z"/></svg>`,
}
function getOAuthIcon(providerType: string): string {
return OAUTH_ICONS[providerType] || OAUTH_ICONS.github
}
const props = defineProps<{ const props = defineProps<{
modelValue: boolean modelValue: boolean

View File

@@ -50,6 +50,9 @@ export const useAuthStore = defineStore('auth', () => {
error.value = '邮箱或密码错误' error.value = '邮箱或密码错误'
} else if (err.response?.status === 422) { } else if (err.response?.status === 422) {
error.value = '请输入有效的邮箱地址' error.value = '请输入有效的邮箱地址'
} else if (err.response?.status === 429) {
// 限流错误,显示后端返回的具体信息
error.value = err.response?.data?.detail || '请求过于频繁,请稍后重试'
} else if (err.response?.status === 500) { } else if (err.response?.status === 500) {
error.value = '服务器错误,请稍后重试' error.value = '服务器错误,请稍后重试'
} else { } else {

View File

@@ -0,0 +1,13 @@
// OAuth provider icons (inline SVG)
export const OAUTH_ICONS: Record<string, string> = {
linuxdo: `<svg viewBox="0 0 120 120" xmlns="http://www.w3.org/2000/svg"><clipPath id="ld"><circle cx="60" cy="60" r="47"/></clipPath><circle fill="#f0f0f0" cx="60" cy="60" r="50"/><rect fill="#1c1c1e" clip-path="url(#ld)" x="10" y="10" width="100" height="30"/><rect fill="#f0f0f0" clip-path="url(#ld)" x="10" y="40" width="100" height="40"/><rect fill="#ffb003" clip-path="url(#ld)" x="10" y="80" width="100" height="30"/></svg>`,
github: `<svg viewBox="0 0 24 24" fill="currentColor"><path d="M12 0c-6.626 0-12 5.373-12 12 0 5.302 3.438 9.8 8.207 11.387.599.111.793-.261.793-.577v-2.234c-3.338.726-4.033-1.416-4.033-1.416-.546-1.387-1.333-1.756-1.333-1.756-1.089-.745.083-.729.083-.729 1.205.084 1.839 1.237 1.839 1.237 1.07 1.834 2.807 1.304 3.492.997.107-.775.418-1.305.762-1.604-2.665-.305-5.467-1.334-5.467-5.931 0-1.311.469-2.381 1.236-3.221-.124-.303-.535-1.524.117-3.176 0 0 1.008-.322 3.301 1.23.957-.266 1.983-.399 3.003-.404 1.02.005 2.047.138 3.006.404 2.291-1.552 3.297-1.23 3.297-1.23.653 1.653.242 2.874.118 3.176.77.84 1.235 1.911 1.235 3.221 0 4.609-2.807 5.624-5.479 5.921.43.372.823 1.102.823 2.222v3.293c0 .319.192.694.801.576 4.765-1.589 8.199-6.086 8.199-11.386 0-6.627-5.373-12-12-12z"/></svg>`,
google: `<svg viewBox="0 0 24 24"><path fill="#4285F4" d="M22.56 12.25c0-.78-.07-1.53-.2-2.25H12v4.26h5.92c-.26 1.37-1.04 2.53-2.21 3.31v2.77h3.57c2.08-1.92 3.28-4.74 3.28-8.09z"/><path fill="#34A853" d="M12 23c2.97 0 5.46-.98 7.28-2.66l-3.57-2.77c-.98.66-2.23 1.06-3.71 1.06-2.86 0-5.29-1.93-6.16-4.53H2.18v2.84C3.99 20.53 7.7 23 12 23z"/><path fill="#FBBC05" d="M5.84 14.09c-.22-.66-.35-1.36-.35-2.09s.13-1.43.35-2.09V7.07H2.18C1.43 8.55 1 10.22 1 12s.43 3.45 1.18 4.93l2.85-2.22.81-.62z"/><path fill="#EA4335" d="M12 5.38c1.62 0 3.06.56 4.21 1.64l3.15-3.15C17.45 2.09 14.97 1 12 1 7.7 1 3.99 3.47 2.18 7.07l3.66 2.84c.87-2.6 3.3-4.53 6.16-4.53z"/></svg>`,
}
// Default icon when provider type is not found
const DEFAULT_ICON = OAUTH_ICONS.github
export function getOAuthIcon(providerType: string): string {
return OAUTH_ICONS[providerType.toLowerCase()] || DEFAULT_ICON
}

View File

@@ -164,7 +164,7 @@
> >
<div <div
:title="item.tooltip" :title="item.tooltip"
:class="['flex items-center gap-1', item.tooltip ? 'cursor-help' : '']" class="flex items-center gap-1"
> >
<span class="text-[10px] text-muted-foreground/60 w-4">{{ item.label }}</span> <span class="text-[10px] text-muted-foreground/60 w-4">{{ item.label }}</span>
<div class="w-12 h-1.5 bg-border rounded-full overflow-hidden"> <div class="w-12 h-1.5 bg-border rounded-full overflow-hidden">

View File

@@ -173,12 +173,18 @@
:key="link.provider_type" :key="link.provider_type"
class="flex items-center justify-between gap-3 rounded-lg border border-border bg-muted/30 p-4" class="flex items-center justify-between gap-3 rounded-lg border border-border bg-muted/30 p-4"
> >
<div class="min-w-0 flex-1"> <div class="flex items-center gap-3 min-w-0 flex-1">
<div class="text-sm font-medium truncate"> <div
{{ link.display_name }} class="oauth-icon shrink-0"
</div> v-html="getOAuthIcon(link.provider_type)"
<div class="text-xs text-muted-foreground truncate"> />
{{ link.provider_username || link.provider_email || '已绑定' }} <div class="min-w-0">
<div class="text-sm font-medium truncate">
{{ link.display_name }}
</div>
<div class="text-xs text-muted-foreground truncate">
{{ link.provider_username || link.provider_email || '已绑定' }}
</div>
</div> </div>
</div> </div>
<Button <Button
@@ -197,12 +203,18 @@
:key="p.provider_type" :key="p.provider_type"
class="flex items-center justify-between gap-3 rounded-lg border border-dashed border-border p-4 hover:border-primary/50 transition-colors" class="flex items-center justify-between gap-3 rounded-lg border border-dashed border-border p-4 hover:border-primary/50 transition-colors"
> >
<div class="min-w-0 flex-1"> <div class="flex items-center gap-3 min-w-0 flex-1">
<div class="text-sm font-medium truncate"> <div
{{ p.display_name }} class="oauth-icon shrink-0"
</div> v-html="getOAuthIcon(p.provider_type)"
<div class="text-xs text-muted-foreground"> />
未绑定 <div class="min-w-0">
<div class="text-sm font-medium truncate">
{{ p.display_name }}
</div>
<div class="text-xs text-muted-foreground">
未绑定
</div>
</div> </div>
</div> </div>
<Button <Button
@@ -432,6 +444,7 @@ import { useAuthStore } from '@/stores/auth'
import { meApi, type Profile } from '@/api/me' import { meApi, type Profile } from '@/api/me'
import { authApi } from '@/api/auth' import { authApi } from '@/api/auth'
import { oauthApi, type OAuthLinkInfo, type OAuthProviderInfo } from '@/api/oauth' import { oauthApi, type OAuthLinkInfo, type OAuthProviderInfo } from '@/api/oauth'
import { getOAuthIcon } from '@/utils/oauth-icons'
import { useDarkMode, type ThemeMode } from '@/composables/useDarkMode' import { useDarkMode, type ThemeMode } from '@/composables/useDarkMode'
import Card from '@/components/ui/card.vue' import Card from '@/components/ui/card.vue'
import Button from '@/components/ui/button.vue' import Button from '@/components/ui/button.vue'
@@ -601,7 +614,42 @@ async function loadOAuthBindings() {
function handleBind(providerType: string) { function handleBind(providerType: string) {
// 保存返回路径OAuth callback 会读取) // 保存返回路径OAuth callback 会读取)
sessionStorage.setItem('redirectPath', route.fullPath) sessionStorage.setItem('redirectPath', route.fullPath)
window.location.href = getApiUrl(`/api/user/oauth/${providerType}/bind`)
// 先获取一次性绑定令牌,再在新标签页打开(避免在 URL 中暴露 access_token
oauthActionLoading.value = true
oauthApi.createBindToken(providerType)
.then((bindToken) => {
// getApiUrl 可能返回相对路径,需要拼接完整 URL
const basePath = getApiUrl(`/api/user/oauth/${providerType}/bind`)
const bindUrl = basePath.startsWith('http')
? new URL(basePath)
: new URL(basePath, window.location.origin)
bindUrl.searchParams.set('bind_token', bindToken)
// 新标签页打开 OAuth 流程
const newTab = window.open(bindUrl.toString(), '_blank')
// 监听标签页关闭,刷新绑定状态
if (newTab) {
const MAX_WAIT_MS = 10 * 60 * 1000 // 10 分钟超时
const startTime = Date.now()
const checkClosed = setInterval(() => {
if (newTab.closed || Date.now() - startTime > MAX_WAIT_MS) {
clearInterval(checkClosed)
oauthActionLoading.value = false
loadOAuthBindings()
}
}, 500)
} else {
// 被浏览器阻止,回退到当前页面跳转
oauthActionLoading.value = false
window.location.href = bindUrl.toString()
}
})
.catch((err) => {
oauthActionLoading.value = false
showError(getErrorMessage(err, '获取绑定令牌失败'))
})
} }
async function handleUnbind(providerType: string) { async function handleUnbind(providerType: string) {
@@ -776,3 +824,15 @@ function formatDate(dateString?: string): string {
}) })
} }
</script> </script>
<style scoped>
.oauth-icon {
width: 24px;
height: 24px;
}
.oauth-icon :deep(svg) {
width: 100%;
height: 100%;
}
</style>

View File

@@ -127,7 +127,7 @@ async def login(request: Request, db: Session = Depends(get_db)):
- **access_token**: 用于后续 API 调用,有效期 24 小时 - **access_token**: 用于后续 API 调用,有效期 24 小时
- **refresh_token**: 用于刷新 access_token - **refresh_token**: 用于刷新 access_token
速率限制: 5次/分钟/IP 速率限制: 20次/分钟/IP
""" """
adapter = AuthLoginAdapter() adapter = AuthLoginAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode) return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -153,7 +153,7 @@ async def register(request: Request, db: Session = Depends(get_db)):
创建新用户账号。需要系统开放注册功能。 创建新用户账号。需要系统开放注册功能。
如果系统开启了邮箱验证,需先通过 /send-verification-code 和 /verify-email 完成邮箱验证。 如果系统开启了邮箱验证,需先通过 /send-verification-code 和 /verify-email 完成邮箱验证。
速率限制: 3次/分钟/IP 速率限制: 10次/分钟/IP
""" """
adapter = AuthRegisterAdapter() adapter = AuthRegisterAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode) return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -202,7 +202,7 @@ async def send_verification_code(request: Request, db: Session = Depends(get_db)
向指定邮箱发送验证码,用于注册前的邮箱验证。 向指定邮箱发送验证码,用于注册前的邮箱验证。
验证码有效期 5 分钟,同一邮箱 60 秒内只能发送一次。 验证码有效期 5 分钟,同一邮箱 60 秒内只能发送一次。
速率限制: 3次/分钟/IP 速率限制: 5次/分钟/IP
""" """
adapter = AuthSendVerificationCodeAdapter() adapter = AuthSendVerificationCodeAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode) return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -216,7 +216,7 @@ async def verify_email(request: Request, db: Session = Depends(get_db)):
验证邮箱收到的验证码是否正确。 验证邮箱收到的验证码是否正确。
验证成功后,邮箱会被标记为已验证状态,可用于注册。 验证成功后,邮箱会被标记为已验证状态,可用于注册。
速率限制: 10次/分钟/IP 速率限制: 20次/分钟/IP
""" """
adapter = AuthVerifyEmailAdapter() adapter = AuthVerifyEmailAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode) return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)

View File

@@ -2,24 +2,30 @@
from __future__ import annotations from __future__ import annotations
from typing import Any, cast from typing import Any, Optional, cast
from fastapi import APIRouter, Depends, Request, status from fastapi import APIRouter, Depends, HTTPException, Request, status
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from starlette.responses import RedirectResponse from starlette.responses import RedirectResponse
from src.api.base.adapter import ApiMode
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
from src.api.base.context import ApiRequestContext from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import ApiRequestPipeline from src.api.base.pipeline import ApiRequestPipeline
from src.clients.redis_client import get_redis_client
from src.database import get_db from src.database import get_db
from src.models.database import User
from src.services.auth.oauth.service import OAuthService from src.services.auth.oauth.service import OAuthService
from src.services.auth.oauth.state import consume_oauth_bind_token, create_oauth_bind_token
router = APIRouter(prefix="/api/user/oauth", tags=["User - OAuth"]) router = APIRouter(prefix="/api/user/oauth", tags=["User - OAuth"])
pipeline = ApiRequestPipeline() pipeline = ApiRequestPipeline()
@router.get("/bindable-providers") @router.get("/bindable-providers")
async def list_bindable_providers(request: Request, db: Session = Depends(get_db)) -> dict[str, Any]: async def list_bindable_providers(
request: Request, db: Session = Depends(get_db)
) -> dict[str, Any]:
adapter = ListBindableProvidersAdapter() adapter = ListBindableProvidersAdapter()
result = await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode) result = await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
return cast(dict[str, Any], result) return cast(dict[str, Any], result)
@@ -32,11 +38,25 @@ async def list_my_oauth_links(request: Request, db: Session = Depends(get_db)) -
return cast(dict[str, Any], result) return cast(dict[str, Any], result)
@router.post("/{provider_type}/bind-token")
async def create_bind_token(
provider_type: str, request: Request, db: Session = Depends(get_db)
) -> dict[str, Any]:
"""创建一次性 OAuth 绑定令牌,用于浏览器跳转场景的安全认证"""
adapter = CreateBindTokenAdapter(provider_type=provider_type)
result = await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
return cast(dict[str, Any], result)
@router.get("/{provider_type}/bind") @router.get("/{provider_type}/bind")
async def bind_oauth_provider( async def bind_oauth_provider(
provider_type: str, request: Request, db: Session = Depends(get_db) provider_type: str,
request: Request,
db: Session = Depends(get_db),
bind_token: Optional[str] = None,
) -> RedirectResponse: ) -> RedirectResponse:
adapter = BindOAuthProviderAdapter(provider_type=provider_type) """发起 OAuth 绑定流程,支持通过 bind_token 参数进行安全认证"""
adapter = BindOAuthProviderAdapter(provider_type=provider_type, bind_token=bind_token)
result = await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode) result = await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
return cast(RedirectResponse, result) return cast(RedirectResponse, result)
@@ -64,13 +84,81 @@ class ListMyOAuthLinksAdapter(AuthenticatedApiAdapter):
return {"links": links} return {"links": links}
class BindOAuthProviderAdapter(AuthenticatedApiAdapter): class CreateBindTokenAdapter(AuthenticatedApiAdapter):
"""创建一次性 OAuth 绑定令牌,用于浏览器跳转场景"""
def __init__(self, provider_type: str): def __init__(self, provider_type: str):
self.provider_type = provider_type self.provider_type = provider_type
async def handle(self, context: ApiRequestContext) -> RedirectResponse: # type: ignore[override] async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
assert context.user is not None assert context.user is not None
url = await OAuthService.build_bind_authorize_url(context.db, context.user, self.provider_type) assert context.user.id is not None
# 验证 provider 是否存在且可绑定
bindable = await OAuthService.list_bindable_providers(context.db, context.user)
if not any(p["provider_type"] == self.provider_type for p in bindable):
raise HTTPException(status_code=400, detail="无法绑定该 Provider")
redis = await get_redis_client(require_redis=True)
if redis is None:
raise HTTPException(status_code=503, detail="Redis 不可用")
token = await create_oauth_bind_token(
redis, user_id=context.user.id, provider_type=self.provider_type
)
return {"bind_token": token}
class BindOAuthProviderAdapter(AuthenticatedApiAdapter):
"""发起 OAuth 绑定流程,支持两种认证方式:
1. Authorization header (标准方式)
2. bind_token 参数 (浏览器跳转场景)
"""
def __init__(self, provider_type: str, bind_token: Optional[str] = None):
self.provider_type = provider_type
self.bind_token = bind_token
self._user_from_bind_token: Optional[User] = None
@property
def mode(self) -> ApiMode: # type: ignore[override]
# 如果有 bind_token使用 PUBLIC mode 跳过 header 认证
if self.bind_token:
return ApiMode.PUBLIC
return ApiMode.USER
def authorize(self, context: ApiRequestContext) -> None:
# 如果是 bind_token 模式,不在这里检查(会在 handle 中验证)
if self.bind_token:
return
# 标准模式,检查用户
if not context.user:
raise HTTPException(status_code=401, detail="未登录")
async def handle(self, context: ApiRequestContext) -> RedirectResponse: # type: ignore[override]
user: Optional[User] = context.user
# 如果使用 bind_token验证并获取用户
if self.bind_token:
redis = await get_redis_client(require_redis=True)
if redis is None:
raise HTTPException(status_code=503, detail="Redis 不可用")
token_data = await consume_oauth_bind_token(redis, self.bind_token)
if not token_data:
raise HTTPException(status_code=401, detail="无效或过期的绑定令牌")
# 验证 provider_type 匹配
if token_data.provider_type != self.provider_type:
raise HTTPException(status_code=400, detail="绑定令牌与 Provider 不匹配")
# 从数据库获取用户
user = context.db.query(User).filter(User.id == token_data.user_id).first()
if not user or not user.is_active or user.is_deleted:
raise HTTPException(status_code=403, detail="用户不存在或已禁用")
assert user is not None
url = await OAuthService.build_bind_authorize_url(context.db, user, self.provider_type)
return RedirectResponse(url=url, status_code=status.HTTP_302_FOUND) return RedirectResponse(url=url, status_code=status.HTTP_302_FOUND)

View File

@@ -8,10 +8,14 @@ from typing import Any, Awaitable, Optional, cast
from redis.asyncio import Redis from redis.asyncio import Redis
OAUTH_STATE_TTL_SECONDS = 600 OAUTH_STATE_TTL_SECONDS = 600
OAUTH_STATE_KEY_PREFIX = "oauth_state:" OAUTH_STATE_KEY_PREFIX = "oauth_state:"
# OAuth bind token: 用于安全地在浏览器跳转时传递用户身份
# 短期有效5分钟一次性使用
OAUTH_BIND_TOKEN_TTL_SECONDS = 300
OAUTH_BIND_TOKEN_KEY_PREFIX = "oauth_bind_token:"
CONSUME_STATE_SCRIPT = r""" CONSUME_STATE_SCRIPT = r"""
local value = redis.call("GET", KEYS[1]) local value = redis.call("GET", KEYS[1])
@@ -76,3 +80,57 @@ async def consume_oauth_state(redis: Redis, nonce: str) -> Optional[OAuthStateDa
return None return None
return OAuthStateData.from_dict(parsed) return OAuthStateData.from_dict(parsed)
@dataclass(frozen=True)
class OAuthBindTokenData:
"""OAuth 绑定临时令牌数据,用于浏览器跳转场景的安全认证"""
token: str
user_id: str
provider_type: str
created_at: int
@classmethod
def from_dict(cls, data: dict[str, Any]) -> "OAuthBindTokenData":
return cls(
token=str(data.get("token") or ""),
user_id=str(data.get("user_id") or ""),
provider_type=str(data.get("provider_type") or ""),
created_at=int(data.get("created_at") or 0),
)
def _bind_token_key(token: str) -> str:
return f"{OAUTH_BIND_TOKEN_KEY_PREFIX}{token}"
async def create_oauth_bind_token(redis: Redis, *, user_id: str, provider_type: str) -> str:
"""创建一次性 OAuth 绑定令牌,用于浏览器跳转场景"""
token = secrets.token_urlsafe(32)
data = {
"token": token,
"user_id": user_id,
"provider_type": provider_type,
"created_at": int(time.time()),
}
await redis.setex(_bind_token_key(token), OAUTH_BIND_TOKEN_TTL_SECONDS, json.dumps(data))
return token
async def consume_oauth_bind_token(redis: Redis, token: str) -> Optional[OAuthBindTokenData]:
"""消费验证并删除OAuth 绑定令牌,返回令牌数据或 None"""
if not token:
return None
key = _bind_token_key(token)
raw = await cast(Awaitable[Optional[str]], redis.eval(CONSUME_STATE_SCRIPT, 1, key))
if not raw:
return None
try:
parsed = json.loads(raw)
except json.JSONDecodeError:
return None
return OAuthBindTokenData.from_dict(parsed)

View File

@@ -24,12 +24,12 @@ class IPRateLimiter:
# 默认限制配置(每分钟) # 默认限制配置(每分钟)
DEFAULT_LIMITS = { DEFAULT_LIMITS = {
"default": 100, # 默认限制 "default": 100, # 默认限制
"login": 5, # 登录接口 "login": 20, # 登录接口
"register": 3, # 注册接口 "register": 10, # 注册接口
"api": 60, # API 接口 "api": 60, # API 接口
"public": 60, # 公共接口 "public": 60, # 公共接口
"verification_send": 3, # 发送验证码接口 "verification_send": 5, # 发送验证码接口
"verification_verify": 10, # 验证验证码接口 "verification_verify": 20, # 验证验证码接口
} }
@staticmethod @staticmethod