from dataclasses import dataclass from enum import Enum from functools import cache from limits import RateLimitItem, parse, storage, strategies from assets.connectors import redis_connector @dataclass(frozen=True) class RateLimitResult: ... @dataclass(frozen=True) class Allowed(RateLimitResult): ... @dataclass(frozen=True) class Exceeded(RateLimitResult): message: str @dataclass(frozen=True) class Error(RateLimitResult): message: str class RateLimitedResource(Enum): UPLOAD = "upload" # Rate limit string notation: https://limits.readthedocs.io/en/stable/quickstart.html#rate-limit-string-notation RATE_LIMITS: dict[tuple[RateLimitedResource, str], RateLimitItem] = { (RateLimitedResource.UPLOAD, "oa:179"): parse("120 per 30 seconds"), } class _RateLimiter: def __init__(self) -> None: self._rate_limiter = strategies.FixedWindowRateLimiter( storage.storage_from_string( "redis://", connection_pool=redis_connector.client.connection_pool ) ) def hit( self, resource: RateLimitedResource, key: str ) -> Error | Exceeded | Allowed: rate_limit = RATE_LIMITS.get((resource, key)) if rate_limit is None: return Allowed() try: is_allowed = self._rate_limiter.hit(rate_limit, resource.value, key) if not is_allowed: return Exceeded(f"Rate limit exceeded. Limit: {rate_limit}.") return Allowed() except Exception as e: return Error(message=f"{type(e).__name__}: {e}") @cache def get_rate_limiter() -> _RateLimiter: return _RateLimiter()