B3: CSRF middleware — add X-CSRF-Token validation via Redis session lookup

This commit is contained in:
2026-07-04 01:23:41 +00:00
parent 2bcf3e5c55
commit 9009135c74
+56 -4
View File
@@ -1,23 +1,32 @@
"""CSRF middleware — Origin header validation for POST/PATCH/DELETE.""" """CSRF middleware — Origin header + CSRF token validation for state-changing requests."""
from __future__ import annotations from __future__ import annotations
import json
import logging
from fastapi import Request, status from fastapi import Request, status
from starlette.middleware.base import BaseHTTPMiddleware from starlette.middleware.base import BaseHTTPMiddleware
from starlette.responses import JSONResponse from starlette.responses import JSONResponse
from app.config import get_settings from app.config import get_settings
logger = logging.getLogger(__name__)
class CSRFMiddleware(BaseHTTPMiddleware): class CSRFMiddleware(BaseHTTPMiddleware):
"""Validate Origin header on all state-changing requests. """Validate Origin header and CSRF token on all state-changing requests.
SameSite=Strict cookie + Origin validation = CSRF protection.
SameSite=Strict cookie + Origin validation + double-submit CSRF token.
The CSRF token is generated at login and stored in the Redis session.
The client must send it via the X-CSRF-Token header on unsafe methods.
""" """
UNSAFE_METHODS = {"POST", "PATCH", "PUT", "DELETE"} UNSAFE_METHODS = {"POST", "PATCH", "PUT", "DELETE"}
async def dispatch(self, request: Request, call_next): async def dispatch(self, request: Request, call_next):
if request.method in self.UNSAFE_METHODS: if request.method in self.UNSAFE_METHODS:
# 1. Origin header check
origin = request.headers.get("origin") origin = request.headers.get("origin")
if not origin: if not origin:
return JSONResponse( return JSONResponse(
@@ -27,11 +36,54 @@ class CSRFMiddleware(BaseHTTPMiddleware):
settings = get_settings() settings = get_settings()
allowed = settings.cors_origin_list allowed = settings.cors_origin_list
# In production, same-origin is enforced; in dev, explicit origins
if origin not in allowed: if origin not in allowed:
return JSONResponse( return JSONResponse(
status_code=status.HTTP_403_FORBIDDEN, status_code=status.HTTP_403_FORBIDDEN,
content={"detail": "Invalid Origin", "code": "csrf_invalid_origin"}, content={"detail": "Invalid Origin", "code": "csrf_invalid_origin"},
) )
# 2. CSRF token validation (double-submit pattern)
# Skip CSRF token check for auth endpoints (login/password-reset)
path = request.url.path
if path.endswith("/auth/login") or path.endswith("/auth/logout") or "/password-reset" in path:
return await call_next(request)
csrf_header = request.headers.get("x-csrf-token")
if not csrf_header:
return JSONResponse(
status_code=status.HTTP_403_FORBIDDEN,
content={"detail": "Missing X-CSRF-Token header", "code": "csrf_missing_token"},
)
# Get session ID from cookie to look up stored CSRF token
session_id = request.cookies.get(settings.session_cookie_name)
if not session_id:
return JSONResponse(
status_code=status.HTTP_403_FORBIDDEN,
content={"detail": "No session for CSRF validation", "code": "csrf_no_session"},
)
# Look up CSRF token from Redis session
import redis.asyncio as aioredis
redis = aioredis.from_url(settings.redis_url, decode_responses=True)
try:
raw = await redis.get(f"session:{session_id}")
if raw is None:
return JSONResponse(
status_code=status.HTTP_403_FORBIDDEN,
content={"detail": "Session expired for CSRF validation", "code": "csrf_session_expired"},
)
session_data = json.loads(raw)
stored_token = session_data.get("csrf_token")
if not stored_token or stored_token != csrf_header:
return JSONResponse(
status_code=status.HTTP_403_FORBIDDEN,
content={"detail": "CSRF token mismatch", "code": "csrf_token_mismatch"},
)
finally:
await redis.close()
return await call_next(request) return await call_next(request)