mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat: OAuth 绑定安全增强及限流配置优化
- 新增一次性绑定令牌机制,避免在 URL 中暴露 access_token - 绑定流程改为新标签页打开,完成后自动刷新状态 - 放宽认证相关接口的 IP 限流配置 - 登录添加 429 限流错误的友好提示 - OAuth 绑定列表添加 Provider 图标显示 Close #106
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
13
frontend/src/utils/oauth-icons.ts
Normal file
13
frontend/src/utils/oauth-icons.ts
Normal 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
|
||||||
|
}
|
||||||
@@ -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">
|
||||||
|
|||||||
@@ -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>
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user