Files
leocrm/app/core/rate_limit.py
T

216 lines
7.2 KiB
Python
Raw Normal View History

"""Redis-based rate limiting for auth endpoints and general API."""
from __future__ import annotations
import logging
from enum import Enum
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__)
class RateLimitPolicy(Enum):
"""Named rate-limit policies for abuse- and cost-sensitive endpoints.
Each policy maps to a pair of Settings fields:
``rate_limit_<name>_max`` and ``rate_limit_<name>_window``.
"""
AUTH = "auth"
AI = "ai"
UPLOAD = "upload"
WEBHOOK = "webhook"
@property
def _max_field(self) -> str:
return f"rate_limit_{self.value}_max"
@property
def _window_field(self) -> str:
return f"rate_limit_{self.value}_window"
def limits(self) -> tuple[int, int]:
"""Return ``(max_attempts, window_seconds)`` from current settings."""
from app.config import get_settings
s = get_settings()
return getattr(s, self._max_field), getattr(s, self._window_field)
async def check_rate_limit_policy(
redis_key: str,
policy: RateLimitPolicy,
) -> None:
"""Check rate limit using a named :class:`RateLimitPolicy`.
Reads ``max_attempts`` and ``window_seconds`` from application settings
and delegates to :func:`check_rate_limit`.
"""
max_attempts, window_seconds = policy.limits()
await check_rate_limit(redis_key, max_attempts, window_seconds)
async def check_rate_limit(
redis_key: str,
max_attempts: int,
window_seconds: int,
) -> None:
"""Check rate limit using Redis INCR + EXPIRE.
Falls back to in-memory rate limiter when Redis is unavailable.
Raises 429 if limit exceeded.
"""
from app.core.resilience import get_circuit, get_inmemory_limiter
circuit = get_circuit("redis")
if await circuit.can_proceed():
try:
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 {},
)
await circuit.record_success()
return
except HTTPException:
raise
except Exception as exc:
logger.warning("Redis rate limit failed: %s — using in-memory fallback", exc)
await circuit.record_failure()
# In-memory fallback
limiter = get_inmemory_limiter()
allowed, retry_after = await limiter.check(redis_key, max_attempts, window_seconds)
if not allowed:
raise HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail={
"detail": "Rate limit exceeded",
"code": "rate_limited",
"retry_after": retry_after,
},
headers={"Retry-After": str(retry_after)} if retry_after > 0 else {},
)
async def reset_rate_limit(redis_key: str) -> None:
"""Reset a rate limit counter (e.g. on successful login).
Resets both Redis and in-memory counters.
"""
from app.core.resilience import get_circuit, get_inmemory_limiter
# Always reset in-memory
limiter = get_inmemory_limiter()
await limiter.reset(redis_key)
# Try Redis
circuit = get_circuit("redis")
if await circuit.can_proceed():
try:
redis = get_redis()
await redis.delete(redis_key)
await circuit.record_success()
except Exception as exc:
logger.warning("Redis rate limit reset failed: %s", exc)
await circuit.record_failure()
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)
# Use token ID for rate limiting if Bearer token is present, otherwise use IP
auth_header = request.headers.get("Authorization", "")
if auth_header.startswith("Bearer "):
token = auth_header[7:]
# Hash token for privacy in Redis key
import hashlib
token_hash = hashlib.sha256(token.encode()).hexdigest()[:16]
rate_key = f"rate:general:token:{token_hash}"
else:
rate_key = f"rate:general:{ip}"
await check_rate_limit(
rate_key,
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)