| """ |
| API middleware — CORS, request-id, rate limiting, request body size limit, |
| global exception handler. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import time |
| import uuid |
| from collections import defaultdict |
|
|
| from fastapi import Request, Response |
| from fastapi.responses import JSONResponse |
| from starlette.middleware.base import BaseHTTPMiddleware |
|
|
|
|
| class RequestContextMiddleware(BaseHTTPMiddleware): |
| """Adds a request-id to every request + measures duration.""" |
|
|
| async def dispatch(self, request: Request, call_next): |
| request_id = request.headers.get("X-Request-ID") or uuid.uuid4().hex[:12] |
| request.state.request_id = request_id |
| t0 = time.perf_counter() |
| response: Response = await call_next(request) |
| elapsed_ms = (time.perf_counter() - t0) * 1000.0 |
| response.headers["X-Request-ID"] = request_id |
| response.headers["X-Response-Time-ms"] = f"{elapsed_ms:.2f}" |
| return response |
|
|
|
|
| class RateLimitMiddleware(BaseHTTPMiddleware): |
| """Simple in-memory per-IP rate limiter. |
| |
| For production, replace with a Redis-backed limiter. |
| """ |
|
|
| def __init__(self, app, requests_per_minute: int = 30): |
| super().__init__(app) |
| self._limit = requests_per_minute |
| self._buckets: dict[str, list[float]] = defaultdict(list) |
|
|
| async def dispatch(self, request: Request, call_next): |
| if request.url.path.startswith("/health"): |
| return await call_next(request) |
| client = request.client.host if request.client else "unknown" |
| now = time.time() |
| window = 60.0 |
| recent = [t for t in self._buckets[client] if now - t < window] |
| if len(recent) >= self._limit: |
| return JSONResponse( |
| status_code=429, |
| content={ |
| "success": False, |
| "error": "rate limit exceeded", |
| "error_type": "RateLimitError", |
| "request_id": getattr(request.state, "request_id", None), |
| }, |
| media_type="application/json", |
| ) |
| recent.append(now) |
| self._buckets[client] = recent |
| return await call_next(request) |
|
|
|
|
| class RequestSizeLimitMiddleware(BaseHTTPMiddleware): |
| """Rejects request bodies larger than max_bytes.""" |
|
|
| def __init__(self, app, max_bytes: int = 25 * 1024 * 1024): |
| super().__init__(app) |
| self._max = max_bytes |
|
|
| async def dispatch(self, request: Request, call_next): |
| cl = request.headers.get("content-length") |
| if cl and cl.isdigit() and int(cl) > self._max: |
| return JSONResponse( |
| status_code=413, |
| content={ |
| "success": False, |
| "error": f"Request body exceeds {self._max} bytes", |
| "error_type": "PayloadTooLarge", |
| "request_id": getattr(request.state, "request_id", None), |
| }, |
| media_type="application/json", |
| ) |
| return await call_next(request) |
|
|
|
|
| class GlobalExceptionMiddleware(BaseHTTPMiddleware): |
| """Catches all uncaught exceptions and returns a structured error response. |
| |
| Prevents stack traces from leaking to clients. |
| """ |
|
|
| async def dispatch(self, request: Request, call_next): |
| try: |
| return await call_next(request) |
| except Exception as e: |
| from loguru import logger |
| logger.exception(f"Unhandled exception on {request.url.path}") |
| return JSONResponse( |
| status_code=500, |
| content={ |
| "success": False, |
| "error": str(e), |
| "error_type": type(e).__name__, |
| "request_id": getattr(request.state, "request_id", None), |
| }, |
| media_type="application/json", |
| ) |
|
|