RICS / app /api /rate_limit.py
StormShadow308's picture
feat: async pipeline, job queue, generation hardening, and docs
732b14f
Raw
History Blame Contribute Delete
4.69 kB
"""Per-tenant sliding-window rate limiter (in-process or Redis-backed)."""
from __future__ import annotations
import logging
import os
import time
from collections import deque
from collections.abc import Awaitable, Callable
from fastapi import HTTPException, Request
from app.config import settings
logger = logging.getLogger(__name__)
_GENERATE_RPM: int = int(
os.environ.get("RATE_LIMIT_GENERATE_RPM", str(settings.rate_limit_generate_rpm))
)
_READ_RPM: int = int(
os.environ.get("RATE_LIMIT_READ_RPM", str(settings.rate_limit_read_rpm))
)
_WINDOW_SECS: float = 60.0
class _SlidingWindowLimiter:
"""In-process per-key sliding-window rate limiter."""
def __init__(self, max_requests: int, window_secs: float) -> None:
self._max = max_requests
self._window = window_secs
self._buckets: dict[str, deque[float]] = {}
def is_allowed(self, key: str) -> bool:
now = time.monotonic()
cutoff = now - self._window
bucket = self._buckets.setdefault(key, deque())
while bucket and bucket[0] < cutoff:
bucket.popleft()
if len(bucket) >= self._max:
return False
bucket.append(now)
return True
def reset(self, key: str) -> None:
self._buckets.pop(key, None)
_generate_mem = _SlidingWindowLimiter(_GENERATE_RPM, _WINDOW_SECS)
_read_mem = _SlidingWindowLimiter(_READ_RPM, _WINDOW_SECS)
_redis_generate = None
_redis_read = None
_use_redis_limits: bool | None = None
async def _ensure_redis_limiters() -> bool:
global _redis_generate, _redis_read, _use_redis_limits
if _use_redis_limits is not None:
return _use_redis_limits
from app.redis_client import redis_configured
if not redis_configured():
_use_redis_limits = False
return False
try:
from app.rate_limit.redis_limiter import RedisSlidingWindowLimiter
from app.redis_client import get_redis
client = await get_redis()
_redis_generate = RedisSlidingWindowLimiter(
client,
key_prefix="rics:rl:generate",
max_requests=_GENERATE_RPM,
window_secs=_WINDOW_SECS,
)
_redis_read = RedisSlidingWindowLimiter(
client,
key_prefix="rics:rl:read",
max_requests=_READ_RPM,
window_secs=_WINDOW_SECS,
)
_use_redis_limits = True
logger.info("Rate limits using Redis (shared across replicas)")
except Exception as exc: # noqa: BLE001
logger.warning("Redis rate limits unavailable, using in-process: %s", exc)
_redis_generate = None
_redis_read = None
_use_redis_limits = False
return _use_redis_limits
def reset_rate_limit_backend_for_tests() -> None:
"""Clear cached backend selection (tests only)."""
global _use_redis_limits, _redis_generate, _redis_read
_use_redis_limits = None
_redis_generate = None
_redis_read = None
_generate_mem.reset("test")
_read_mem.reset("test")
async def _is_allowed_generate(tenant_id: str) -> bool:
if await _ensure_redis_limiters() and _redis_generate is not None:
return await _redis_generate.is_allowed(tenant_id)
return _generate_mem.is_allowed(tenant_id)
async def _is_allowed_read(tenant_id: str) -> bool:
if await _ensure_redis_limiters() and _redis_read is not None:
return await _redis_read.is_allowed(tenant_id)
return _read_mem.is_allowed(tenant_id)
async def check_generate(request: Request) -> None:
tenant_id: str = getattr(request.state, "tenant_id", "anonymous")
if not await _is_allowed_generate(tenant_id):
logger.warning("Rate limit exceeded (generate) for tenant=%s", tenant_id)
raise HTTPException(
status_code=429,
detail=(
f"Rate limit exceeded: at most {_GENERATE_RPM} generation requests "
f"per minute per tenant. Please wait and retry."
),
headers={"Retry-After": "60"},
)
async def check_read(request: Request) -> None:
tenant_id: str = getattr(request.state, "tenant_id", "anonymous")
if not await _is_allowed_read(tenant_id):
logger.warning("Rate limit exceeded (read) for tenant=%s", tenant_id)
raise HTTPException(
status_code=429,
detail=(
f"Rate limit exceeded: at most {_READ_RPM} read requests "
f"per minute per tenant. Please wait and retry."
),
headers={"Retry-After": "60"},
)
def rate_limit_backend_label() -> str:
if _use_redis_limits:
return "redis"
return "memory"