face-intel / api /middleware.py
Marwan
Restructure + add reverse face search (PimEyes-style)
f5eeb1c
Raw
History Blame Contribute Delete
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",
)