Files
backend/backend/app/core/csrf.py
T

75 lines
2.9 KiB
Python
Raw Normal View History

"""CSRF 防护中间件 — Double-Submit Cookie 模式"""
import secrets
from fastapi.responses import JSONResponse
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.requests import Request
from app.config import settings
SAFE_METHODS = {"GET", "HEAD", "OPTIONS"}
CSRF_EXEMPT_PREFIXES = (
"/api/v1/auth/", "/api/v1/webhooks/", "/api/v1/ws/",
"/api/v1/captcha/", "/api/v1/public/", "/api/v1/features/",
"/api/v1/domains/", "/api/v1/approvals/", "/api/v1/subscriptions/plans",
"/docs", "/openapi.json", "/redoc", "/health",
)
class CSRFProtectionMiddleware(BaseHTTPMiddleware):
"""双提交 Cookie CSRF 防护。
- 安全方法(GET/HEAD/OPTIONS)+ 豁免前缀:不验证,但自动注入 csrf_token cookie
- Bearer Token 请求:自动豁免(SPA token-based auth 天然防 CSRF
- 其他 state-changing 请求:验证 Cookie 中的 csrf_token 与 X-CSRF-Token 请求头一致
"""
async def dispatch(self, request: Request, call_next):
path = request.url.path.rstrip('/') or '/'
is_safe = request.method in SAFE_METHODS
is_exempt = any(path.startswith(p.rstrip('/')) for p in CSRF_EXEMPT_PREFIXES)
# 安全方法 + 豁免路径:确保 csrf_token cookie 已设置
if is_safe or is_exempt:
response = await call_next(request)
if not request.cookies.get("csrf_token"):
cookie_secure = request.headers.get("x-forwarded-proto", request.url.scheme) == "https"
response.set_cookie(
key="csrf_token",
value=secrets.token_hex(32),
# NOT HttpOnly — JS 需要读取并设置 X-CSRF-Token 请求头
httponly=False,
samesite="lax",
secure=cookie_secure,
path="/",
max_age=86400 * 7, # 7 天
)
return response
# Bearer Token 天然防 CSRF(浏览器不会自动附加)
if request.headers.get("Authorization", "").startswith("Bearer "):
return await call_next(request)
# State-changing 请求:验证双提交 Token
cookie_token = request.cookies.get("csrf_token")
header_token = request.headers.get("X-CSRF-Token")
if not cookie_token or not header_token or not _constant_time_compare(cookie_token, header_token):
return JSONResponse(
status_code=403,
content={"success": False, "error": {"code": "CSRF_FAILED", "message": "CSRF token missing or invalid"}},
)
return await call_next(request)
def _constant_time_compare(a: str, b: str) -> bool:
"""防止时序攻击的字符串比较"""
if len(a) != len(b):
return False
result = 0
for x, y in zip(a, b, strict=True):
result |= ord(x) ^ ord(y)
return result == 0