"""Redis-based rate limiting for auth endpoints and general API.""" from __future__ import annotations import logging from fastapi import HTTPException, Request, status from starlette.middleware.base import BaseHTTPMiddleware from starlette.responses import JSONResponse from app.core.auth import get_redis logger = logging.getLogger(__name__) async def check_rate_limit( redis_key: str, max_attempts: int, window_seconds: int, ) -> None: """Check rate limit using Redis INCR + EXPIRE. Raises 429 if limit exceeded. """ redis = get_redis() current = await redis.incr(redis_key) if current == 1: await redis.expire(redis_key, window_seconds) if current > max_attempts: ttl = await redis.ttl(redis_key) raise HTTPException( status_code=status.HTTP_429_TOO_MANY_REQUESTS, detail={ "detail": "Rate limit exceeded", "code": "rate_limited", "retry_after": ttl, }, headers={"Retry-After": str(ttl)} if ttl > 0 else {}, ) async def reset_rate_limit(redis_key: str) -> None: """Reset a rate limit counter (e.g. on successful login).""" redis = get_redis() await redis.delete(redis_key) def get_client_ip(request: Request) -> str: """Extract client IP from request. Only trusts X-Forwarded-For if the direct client is a trusted proxy (configured via TRUSTED_PROXY_CIDRS env var, comma-separated CIDRs). This prevents IP spoofing to bypass rate limits. """ direct_ip = request.client.host if request.client else "unknown" # Check if the direct client is a trusted proxy from app.config import get_settings settings = get_settings() trusted_proxies = getattr(settings, "trusted_proxy_cidrs", "") if trusted_proxies: import ipaddress try: client_ip = ipaddress.ip_address(direct_ip) for cidr in trusted_proxies.split(","): cidr = cidr.strip() if cidr and client_ip in ipaddress.ip_network(cidr, strict=False): # Trusted proxy — use X-Forwarded-For forwarded = request.headers.get("x-forwarded-for") if forwarded: # Use the leftmost (original client) IP return forwarded.split(",")[0].strip() # Fallback to X-Real-IP real_ip = request.headers.get("x-real-ip") if real_ip: return real_ip.strip() break except (ValueError, TypeError): pass # Not a trusted proxy or no trusted proxies configured — use direct IP return direct_ip class GeneralRateLimitMiddleware(BaseHTTPMiddleware): """Apply general rate limiting to all API routes.""" # Paths to skip rate limiting SKIP_PATHS = {"/api/v1/health", "/api/v1/health/live", "/api/v1/health/ready", "/api/v1/metrics"} async def dispatch(self, request: Request, call_next): from app.config import get_settings settings = get_settings() path = request.url.path # Skip health and metrics endpoints if path in self.SKIP_PATHS or path.startswith("/docs") or path.startswith("/redoc"): return await call_next(request) # Only rate limit API routes if not path.startswith("/api/"): return await call_next(request) try: ip = get_client_ip(request) await check_rate_limit( f"rate:general:{ip}", settings.rate_limit_general_max, settings.rate_limit_general_window, ) except HTTPException as exc: return JSONResponse( status_code=exc.status_code, content=exc.detail, headers=exc.headers, ) return await call_next(request)