33Uses slowapi for IP-based rate limiting with authenticated user bypass support.
44"""
55import logging
6+ import re
7+ import time
8+ from typing import Optional
9+
610from slowapi import Limiter
711from slowapi .util import get_remote_address
812from slowapi .errors import RateLimitExceeded
913from slowapi .middleware import SlowAPIMiddleware
10- from fastapi import Request , Response
14+ from fastapi import Request
1115from fastapi .responses import JSONResponse
16+ from starlette .middleware .base import BaseHTTPMiddleware
17+
1218from .config import get_settings
1319
1420logger = logging .getLogger (__name__ )
1521
1622settings = get_settings ()
1723
24+ RATE_LIMIT_HEADER_NAMES = (
25+ "Retry-After" ,
26+ "X-RateLimit-Limit" ,
27+ "X-RateLimit-Remaining" ,
28+ "X-RateLimit-Reset" ,
29+ )
30+
1831
1932def _get_rate_limit_key (request : Request ) -> str :
20- """
21- Determine rate limit key based on request.
22- If authenticated user bypass is enabled and a valid Bearer token is present,
23- return a user-specific key that receives higher limits.
24- Otherwise, fall back to IP-based limiting.
25- """
2633 if settings .rate_limit_auth_bypass :
2734 auth_header = request .headers .get ("Authorization" , "" )
2835 if auth_header .startswith ("Bearer " ) and len (auth_header ) > 7 :
@@ -34,34 +41,99 @@ def _get_rate_limit_key(request: Request) -> str:
3441 return get_remote_address (request )
3542
3643
44+ def _parse_limit_value (limit_detail : str ) -> int :
45+ match = re .match (r"(\d+)" , limit_detail or "" )
46+ return int (match .group (1 )) if match else 60
47+
48+
49+ def _parse_retry_after_seconds (limit_detail : str ) -> int :
50+ detail = (limit_detail or "" ).lower ()
51+ count_match = re .search (r"(\d+)\s*per" , detail )
52+ window_match = re .search (r"per\s+(\d+)\s*(second|minute|hour|day)" , detail )
53+
54+ count = int (count_match .group (1 )) if count_match else 60
55+ if not window_match :
56+ return max (count , 1 )
57+
58+ amount = int (window_match .group (1 ))
59+ unit = window_match .group (2 )
60+ multipliers = {"second" : 1 , "minute" : 60 , "hour" : 3600 , "day" : 86400 }
61+ window_seconds = amount * multipliers .get (unit , 60 )
62+ return max (window_seconds // max (count , 1 ), 1 )
63+
64+
65+ def build_rate_limit_headers (limit_detail : str , remaining : Optional [int ] = None ) -> dict [str , str ]:
66+ limit_value = _parse_limit_value (limit_detail )
67+ retry_after = _parse_retry_after_seconds (limit_detail )
68+ reset_at = int (time .time ()) + retry_after
69+ resolved_remaining = str (max (remaining , 0 )) if remaining is not None else str (max (limit_value - 1 , 0 ))
70+
71+ return {
72+ "Retry-After" : str (retry_after ),
73+ "X-RateLimit-Limit" : str (limit_value ),
74+ "X-RateLimit-Remaining" : resolved_remaining ,
75+ "X-RateLimit-Reset" : str (reset_at ),
76+ }
77+
78+
3779limiter = Limiter (
3880 key_func = _get_rate_limit_key ,
3981 default_limits = [settings .rate_limit_default ],
4082 storage_uri = settings .redis_url if settings .redis_enabled else "memory://" ,
83+ headers_enabled = True ,
4184)
4285
4386
4487def rate_limit_exceeded_handler (request : Request , exc : RateLimitExceeded ) -> JSONResponse :
45- """Handle rate limit exceeded errors with proper retry information."""
46- retry_after = exc .detail .split ("per" )[- 1 ].strip () if "per" in exc .detail else "60"
47- logger .warning ("Rate limit exceeded for %s on %s" , _get_rate_limit_key (request ), request .url .path )
88+ limit_detail = str (exc .detail )
89+ headers = build_rate_limit_headers (limit_detail , remaining = 0 )
90+ retry_after = int (headers ["Retry-After" ])
91+
92+ logger .warning (
93+ "Rate limit exceeded for %s on %s" ,
94+ _get_rate_limit_key (request ),
95+ request .url .path ,
96+ )
97+
4898 return JSONResponse (
4999 status_code = 429 ,
50100 content = {
51101 "error_code" : "RATE_001" ,
52- "detail" : f"Rate limit exceeded: { exc . detail } " ,
102+ "detail" : f"Rate limit exceeded: { limit_detail } " ,
53103 "retry_after" : retry_after ,
54104 },
55- headers = {
56- "Retry-After" : retry_after ,
57- "X-RateLimit-Limit" : str (exc .detail ),
58- },
105+ headers = headers ,
59106 )
60107
61108
109+ class RateLimitHeaderMiddleware (BaseHTTPMiddleware ):
110+ async def dispatch (self , request : Request , call_next ):
111+ response = await call_next (request )
112+
113+ if response .status_code == 429 :
114+ return response
115+
116+ limit_header = response .headers .get ("X-RateLimit-Limit" )
117+ if limit_header :
118+ for header_name in RATE_LIMIT_HEADER_NAMES :
119+ if header_name in response .headers :
120+ continue
121+ if "X-RateLimit-Remaining" not in response .headers :
122+ response .headers ["X-RateLimit-Remaining" ] = str (max (_parse_limit_value (limit_header ) - 1 , 0 ))
123+ if "X-RateLimit-Reset" not in response .headers :
124+ retry_after = _parse_retry_after_seconds (settings .rate_limit_default )
125+ response .headers ["X-RateLimit-Reset" ] = str (int (time .time ()) + retry_after )
126+
127+ return response
128+
129+
62130def setup_rate_limiting (app ):
63- """Attach rate limiter and exception handler to the FastAPI app."""
64131 app .state .limiter = limiter
132+ app .add_middleware (RateLimitHeaderMiddleware )
65133 app .add_middleware (SlowAPIMiddleware )
66134 app .add_exception_handler (RateLimitExceeded , rate_limit_exceeded_handler )
67- logger .info ("Rate limiting enabled: default=%s, auth=%s" , settings .rate_limit_default , settings .rate_limit_auth )
135+ logger .info (
136+ "Rate limiting enabled: default=%s, auth=%s" ,
137+ settings .rate_limit_default ,
138+ settings .rate_limit_auth ,
139+ )
0 commit comments