"""Small in-process limiters for the single-worker API deployment. The reverse proxy remains the first rate-limit layer in production. These limiters provide a second fail-safe around authentication and API credentials. """ from __future__ import annotations import asyncio from dataclasses import dataclass import math import time @dataclass class _Bucket: window_started: float attempts: int = 0 blocked_until: float = 0.0 last_seen: float = 0.0 class FixedWindowLimiter: def __init__( self, *, limit: int, window_seconds: int, block_seconds: int = 0, ) -> None: self.limit = max(1, int(limit)) self.window_seconds = max(1, int(window_seconds)) self.block_seconds = max(0, int(block_seconds)) self._buckets: dict[str, _Bucket] = {} self._lock = asyncio.Lock() async def check(self, key: str) -> int: """Return seconds remaining when blocked, otherwise zero.""" now = time.monotonic() async with self._lock: bucket = self._buckets.get(key) if not bucket: return 0 bucket.last_seen = now if bucket.blocked_until > now: return max(1, math.ceil(bucket.blocked_until - now)) if now - bucket.window_started >= self.window_seconds: self._buckets.pop(key, None) return 0 async def record(self, key: str) -> int: """Count one attempt and return a retry delay if the limit is reached.""" now = time.monotonic() async with self._lock: bucket = self._buckets.get(key) if not bucket or now - bucket.window_started >= self.window_seconds: bucket = _Bucket(window_started=now, last_seen=now) self._buckets[key] = bucket bucket.last_seen = now if bucket.blocked_until > now: return max(1, math.ceil(bucket.blocked_until - now)) bucket.attempts += 1 if bucket.attempts >= self.limit: delay = self.block_seconds or max( 1, math.ceil(self.window_seconds - (now - bucket.window_started)), ) bucket.blocked_until = now + delay return delay self._prune(now) return 0 async def consume(self, key: str) -> int: blocked = await self.check(key) return blocked or await self.record(key) async def reset(self, key: str) -> None: async with self._lock: self._buckets.pop(key, None) def _prune(self, now: float) -> None: if len(self._buckets) < 2_000: return cutoff = now - max(self.window_seconds, self.block_seconds, 60) * 2 stale = [key for key, value in self._buckets.items() if value.last_seen < cutoff] for key in stale: self._buckets.pop(key, None)