import logging import random import time from collections.abc import Callable from app.dsp.exceptions import RateLimitError, StreamingAPIError from app.dsp.stats import ( InMemoryStatsBackend, RequestStats, StatsBackend, ) logger = logging.getLogger(__name__) DEFAULT_RETRY_AFTER = 5 DEFAULT_MAX_RETRIES = 3 RETRY_AFTER_THRESHOLD = 10 # seconds — escalate to RateLimitError immediately MAX_SLEEP = 30 class RateLimitGuard: def __init__( self, max_retries: int = DEFAULT_MAX_RETRIES, default_retry_after: int = DEFAULT_RETRY_AFTER, stats_backend: StatsBackend | None = None, ) -> None: self._max_retries = max_retries self._default_retry_after = default_retry_after self._stats = stats_backend or InMemoryStatsBackend() @property def stats(self) -> RequestStats: return self._stats.snapshot() def execute[T](self, fn: Callable[[], T]) -> T: self._stats.incr_requests() for attempt in range(self._max_retries + 1): try: return fn() except StreamingAPIError as exc: if exc.status_code != 429: raise self._stats.incr_rate_limited() retry_after = exc.retry_after or self._default_retry_after if retry_after > RETRY_AFTER_THRESHOLD: raise RateLimitError( f"Retry-After {retry_after}s exceeds threshold " f"({RETRY_AFTER_THRESHOLD}s) — aborting batch", status_code=429, retry_after=retry_after, ) from exc if attempt == self._max_retries: raise RateLimitError( f"Rate limit retries exhausted after {attempt + 1} attempts", status_code=429, retry_after=retry_after, ) from exc self._stats.incr_retries() sleep = min(retry_after * (2**attempt), MAX_SLEEP) + random.uniform( 0, 1 ) logger.warning( "Rate limited (429), retrying in %.1fs (attempt %d/%d)", sleep, attempt + 1, self._max_retries, ) time.sleep(sleep) raise AssertionError("unreachable")