38 lines
1.3 KiB
Python
38 lines
1.3 KiB
Python
|
|
"""CSRF middleware — Origin header validation for POST/PATCH/DELETE."""
|
||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
from fastapi import Request, status
|
||
|
|
from starlette.middleware.base import BaseHTTPMiddleware
|
||
|
|
from starlette.responses import JSONResponse
|
||
|
|
|
||
|
|
from app.config import get_settings
|
||
|
|
|
||
|
|
|
||
|
|
class CSRFMiddleware(BaseHTTPMiddleware):
|
||
|
|
"""Validate Origin header on all state-changing requests.
|
||
|
|
SameSite=Strict cookie + Origin validation = CSRF protection.
|
||
|
|
"""
|
||
|
|
|
||
|
|
UNSAFE_METHODS = {"POST", "PATCH", "PUT", "DELETE"}
|
||
|
|
|
||
|
|
async def dispatch(self, request: Request, call_next):
|
||
|
|
if request.method in self.UNSAFE_METHODS:
|
||
|
|
origin = request.headers.get("origin")
|
||
|
|
if not origin:
|
||
|
|
return JSONResponse(
|
||
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
||
|
|
content={"detail": "Missing Origin header", "code": "csrf_missing_origin"},
|
||
|
|
)
|
||
|
|
|
||
|
|
settings = get_settings()
|
||
|
|
allowed = settings.cors_origin_list
|
||
|
|
# In production, same-origin is enforced; in dev, explicit origins
|
||
|
|
if origin not in allowed:
|
||
|
|
return JSONResponse(
|
||
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
||
|
|
content={"detail": "Invalid Origin", "code": "csrf_invalid_origin"},
|
||
|
|
)
|
||
|
|
|
||
|
|
return await call_next(request)
|