""" 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", )