Spaces:
Runtime error
Runtime error
| """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" | |