Download api/middleware.py from Hellbound237/face-intel: direct link, hf CLI and curl.
- Browser
- Download file 3.88 kB
-
https://huggingface.co/Hellbound237/face-intel/resolve/main/api/middleware.py
- Command line
-
hf download hf://Hellbound237/face-intel/api/middleware.py
-
curl -L -o middleware.py https://huggingface.co/Hellbound237/face-intel/resolve/main/api/middleware.py
3.88 kB
| """ | |
| 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", | |
| ) | |