| import base64 | |
| import json | |
| import logging | |
| from fastapi import Request | |
| from slowapi import Limiter | |
| from slowapi.util import get_remote_address | |
| logger = logging.getLogger(__name__) | |
| def _rate_limit_key(request: Request) -> str: | |
| """Use user ID from JWT for authenticated requests, fall back to IP.""" | |
| auth = request.headers.get("Authorization", "") | |
| if auth.startswith("Bearer "): | |
| try: | |
| token = auth[7:] | |
| parts = token.split(".") | |
| if len(parts) == 3: | |
| payload = parts[1] | |
| padding = 4 - len(payload) % 4 | |
| if padding != 4: | |
| payload += "=" * padding | |
| decoded = base64.urlsafe_b64decode(payload) | |
| claims = json.loads(decoded) | |
| uid = claims.get("sub") | |
| if uid: | |
| return f"user:{uid}" | |
| except Exception: | |
| pass | |
| return get_remote_address(request) | |
| limiter = Limiter(key_func=_rate_limit_key, default_limits=[]) | |