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:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user