"""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"