from __future__ import annotations import time from collections import defaultdict from fastapi import Request, Response, status from fastapi.responses import JSONResponse from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint from app.core.config import settings from app.core.logging_config import get_logger logger = get_logger(__name__) class RateLimitMiddleware(BaseHTTPMiddleware): """Token bucket rate limiter per client IP.""" EXCLUDED_PATHS = {"/api/v1/health", "/metrics", "/docs", "/openapi.json", "/redoc"} def __init__(self, app, max_requests: int | None = None, window_seconds: int | None = None) -> None: # noqa: ANN001 super().__init__(app) self.max_requests = max_requests or settings.rate_limit_requests self.window_seconds = window_seconds or settings.rate_limit_window_seconds self._requests: dict[str, list[float]] = defaultdict(list) def _clean_old_requests(self, client_ip: str, now: float) -> None: """Remove requests outside the current window.""" cutoff = now - self.window_seconds self._requests[client_ip] = [ ts for ts in self._requests[client_ip] if ts > cutoff ] async def dispatch(self, request: Request, call_next: RequestResponseEndpoint) -> Response: if request.url.path in self.EXCLUDED_PATHS: return await call_next(request) client_ip = request.client.host if request.client else "unknown" now = time.monotonic() self._clean_old_requests(client_ip, now) if len(self._requests[client_ip]) >= self.max_requests: logger.warning( "rate_limit_exceeded", client_ip=client_ip, path=request.url.path, request_count=len(self._requests[client_ip]), ) return JSONResponse( status_code=status.HTTP_429_TOO_MANY_REQUESTS, content={ "detail": "Rate limit exceeded. Please try again later.", "retry_after_seconds": self.window_seconds, }, headers={ "Retry-After": str(self.window_seconds), "X-RateLimit-Limit": str(self.max_requests), "X-RateLimit-Remaining": "0", "X-RateLimit-Reset": str(int(now + self.window_seconds)), }, ) self._requests[client_ip].append(now) remaining = self.max_requests - len(self._requests[client_ip]) response = await call_next(request) response.headers["X-RateLimit-Limit"] = str(self.max_requests) response.headers["X-RateLimit-Remaining"] = str(remaining) response.headers["X-RateLimit-Reset"] = str(int(now + self.window_seconds)) return response