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

@@ -127,7 +127,7 @@ async def login(request: Request, db: Session = Depends(get_db)):
- **access_token**: 用于后续 API 调用,有效期 24 小时
- **refresh_token**: 用于刷新 access_token
速率限制: 5次/分钟/IP
速率限制: 20次/分钟/IP
"""
adapter = AuthLoginAdapter()
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 完成邮箱验证。
速率限制: 3次/分钟/IP
速率限制: 10次/分钟/IP
"""
adapter = AuthRegisterAdapter()
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 秒内只能发送一次。
速率限制: 3次/分钟/IP
速率限制: 5次/分钟/IP
"""
adapter = AuthSendVerificationCodeAdapter()
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()
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 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 starlette.responses import RedirectResponse
from src.api.base.adapter import ApiMode
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import ApiRequestPipeline
from src.clients.redis_client import get_redis_client
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.state import consume_oauth_bind_token, create_oauth_bind_token
router = APIRouter(prefix="/api/user/oauth", tags=["User - OAuth"])
pipeline = ApiRequestPipeline()
@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()
result = await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
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)
@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")
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:
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)
return cast(RedirectResponse, result)
@@ -64,13 +84,81 @@ class ListMyOAuthLinksAdapter(AuthenticatedApiAdapter):
return {"links": links}
class BindOAuthProviderAdapter(AuthenticatedApiAdapter):
class CreateBindTokenAdapter(AuthenticatedApiAdapter):
"""创建一次性 OAuth 绑定令牌,用于浏览器跳转场景"""
def __init__(self, provider_type: str):
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
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)

View File

@@ -8,10 +8,14 @@ from typing import Any, Awaitable, Optional, cast
from redis.asyncio import Redis
OAUTH_STATE_TTL_SECONDS = 600
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"""
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 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": 100, # 默认限制
"login": 5, # 登录接口
"register": 3, # 注册接口
"login": 20, # 登录接口
"register": 10, # 注册接口
"api": 60, # API 接口
"public": 60, # 公共接口
"verification_send": 3, # 发送验证码接口
"verification_verify": 10, # 验证验证码接口
"verification_send": 5, # 发送验证码接口
"verification_verify": 20, # 验证验证码接口
}
@staticmethod