diff --git a/src/mnemosyne/gateway.py b/src/mnemosyne/gateway.py index b4bab1e..7aa385d 100644 --- a/src/mnemosyne/gateway.py +++ b/src/mnemosyne/gateway.py @@ -66,6 +66,93 @@ from mnemosyne.tags import ( from mnemosyne.telemetry import Telemetry import re as _re +import threading as _threading + +# --------------------------------------------------------------------------- +# Outbound Rate Limiter — token bucket with circuit breaker. +# Prevents hitting Anthropic's per-minute rate limits (especially on +# Max 5x plans with ~50 RPM ceilings). +# --------------------------------------------------------------------------- + + +class OutboundRateLimiter: + """Token-bucket rate limiter with 429 circuit breaker. + + - **Token bucket**: allows ``max_rpm`` requests per minute with burst + tolerance. Each ``acquire()`` call blocks until a token is available. + - **Retry-After**: when the upstream returns 429 with a ``retry-after`` + header, the limiter pauses all outbound requests for that duration. + - **Circuit breaker**: after ``_CB_THRESHOLD`` consecutive 429s, the + limiter enters open state and pauses for ``_CB_COOLDOWN`` seconds + before allowing any new requests. + """ + + _CB_THRESHOLD = 3 # consecutive 429s to trip the breaker + _CB_COOLDOWN = 30.0 # seconds to wait in open state + + def __init__(self, max_rpm: int = 40): + self._max_rpm = max_rpm + self._interval = 60.0 / max_rpm # seconds between tokens + self._lock = _threading.Lock() + self._next_allowed = 0.0 # monotonic timestamp + self._consecutive_429s = 0 + self._circuit_open_until = 0.0 # monotonic timestamp + + def acquire(self) -> None: + """Block until a request is allowed. Thread-safe.""" + while True: + with self._lock: + now = time.monotonic() + + # Circuit breaker: if open, wait for cooldown + if now < self._circuit_open_until: + wait = self._circuit_open_until - now + else: + # Token bucket: wait until next allowed slot + if now >= self._next_allowed: + self._next_allowed = now + self._interval + return # token acquired + wait = self._next_allowed - now + self._next_allowed += self._interval + + # Sleep outside the lock + time.sleep(wait) + + def record_response(self, status_code: int, retry_after: float | None = None) -> None: + """Record a response status. Call after every outbound request.""" + with self._lock: + if status_code == 429: + self._consecutive_429s += 1 + + # Honor retry-after header + if retry_after and retry_after > 0: + pause_until = time.monotonic() + retry_after + self._next_allowed = max(self._next_allowed, pause_until) + + # Circuit breaker + if self._consecutive_429s >= self._CB_THRESHOLD: + self._circuit_open_until = time.monotonic() + self._CB_COOLDOWN + self._consecutive_429s = 0 # reset counter for next cycle + import sys + + print( + f"[RateLimiter] CIRCUIT OPEN: {self._CB_THRESHOLD} consecutive 429s, " + f"pausing {self._CB_COOLDOWN}s", + file=sys.stderr, + flush=True, + ) + else: + self._consecutive_429s = 0 # reset on any successful response + + @property + def stats(self) -> dict[str, Any]: + with self._lock: + return { + "max_rpm": self._max_rpm, + "consecutive_429s": self._consecutive_429s, + "circuit_open": time.monotonic() < self._circuit_open_until, + } + # --------------------------------------------------------------------------- # SSE Cleanup Filter — strips / tags from @@ -1542,6 +1629,7 @@ def create_app( } ) app.state.clients = clients + app.state.rate_limiter = OutboundRateLimiter(max_rpm=40) # ── Endpoints ──────────────────────────────────────────────────── @@ -1569,6 +1657,7 @@ def create_app( "openai": openai_model_override, }, "sessions": session_summaries, + "rate_limiter": app.state.rate_limiter.stats, } @app.get("/metrics") @@ -2464,10 +2553,17 @@ def create_app( sse_buffer = bytearray() usage: dict[str, Any] = {} try: + rate_limiter: OutboundRateLimiter = app.state.rate_limiter + rate_limiter.acquire() with client.stream( "POST", upstream_path, json=outgoing_body, headers=headers ) as resp: status_code = resp.status_code + retry_after = None + if status_code == 429: + ra = resp.headers.get("retry-after") + retry_after = float(ra) if ra else 1.0 + rate_limiter.record_response(status_code, retry_after) if status_code >= 400: body = resp.read() bytes_out = len(body) @@ -2547,7 +2643,14 @@ def create_app( # Non-streaming try: + rate_limiter: OutboundRateLimiter = app.state.rate_limiter + rate_limiter.acquire() resp = client.post(upstream_path, json=outgoing_body, headers=headers) + retry_after = None + if resp.status_code == 429: + ra = resp.headers.get("retry-after") + retry_after = float(ra) if ra else 1.0 + rate_limiter.record_response(resp.status_code, retry_after) except httpx.HTTPError as e: emit_event( "provider_error",