Files
roboco/roboco/services/gateway/rate_limit_tracker.py
T
7804e0fafa [f8480831] Batch B: extract route helpers in remaining smaller-offender route files (#760)
* [f8480831] refactor(api): extract route-layer helpers into services/schemas/utils (batch B)

Moves 28 non-@router-decorated helper functions out of 15 route files
(optimal, project, release, dashboard, pitch, x, docs, git, playbooks,
product, provider, research, secretary, system, work_session) into
their paired services module (DB/service-calling helpers), the route's
schemas module as a converter (pure response/request shaping, mirroring
the existing project_to_response/assignment_to_response pattern), or
roboco/utils/converters.py (pure generic helpers). Adds two small
shared role-check helpers to api/deps.py (require_auditor_or_ceo,
require_role_in) for endpoint-specific role gates that had no existing
home. Placement-only: no route paths, schemas, or observable behavior
changed. Fixes the handful of tests that imported the old private
helper names directly.

* [f8480831] docs(map): document Batch B route-helper relocation in api-routes-schemas.md

* [f8480831] docs(map): add Key Symbols rows for require_auditor_or_ceo/require_role_in

---------

Co-authored-by: Backend Developer 2 <be-dev-2@roboco.tech>
Co-authored-by: Backend Documenter <be-doc@roboco.tech>
2026-07-31 20:23:10 +00:00

283 lines
11 KiB
Python

"""Redis-backed rate-limit state tracker for the agent gateway.
State is persisted in Redis as a JSON blob keyed by provider name.
Because it is backed by Redis rather than process memory, state survives
a process restart and a *new* ``RateLimitStateTracker`` instance pointing
at the same Redis URL will read the same values — satisfying the
cross-reconnection persistence requirement.
"""
from __future__ import annotations
import json
from datetime import UTC, datetime, timedelta
from typing import Any
import redis.asyncio as redis
from roboco.config import settings
# Server-side atomic read-modify-write scripts. Redis single-threads a Lua
# ``EVAL``, so the GET → decode → mutate → SET inside one script is indivisible:
# a concurrent op is serialized entirely before or after the script, never
# interleaved between the script's GET and SET. Without this the counter update
# was a non-atomic ``get_state`` → mutate → ``set`` in Python, so a re-park's
# fresh episode blob could be clobbered by the stale increment writing back the
# OLD blob. ``activate`` is itself a Lua merge (not a blind SET): it refreshes
# the episode metadata but carries over the previous ``probe_failures`` count,
# so a re-park can no longer wipe an in-flight increment (resetting the give-up
# / CEO-notify count mid-episode). The counter scripts mutate ONLY
# ``probe_failures`` so every other episode field survives the bump.
_INCREMENT_PROBE_FAILURES = """\
-- roboco:increment_probe_failures
local key = KEYS[1]
local raw = redis.call('GET', key)
if not raw then
redis.call('SET', key, cjson.encode({probe_failures = 1}))
return 1
end
local state = cjson.decode(raw)
local cur = state['probe_failures']
if cur == nil then cur = 0 end
local new_count = cur + 1
state['probe_failures'] = new_count
redis.call('SET', key, cjson.encode(state))
return new_count
"""
_RESET_PROBE_FAILURES = """\
-- roboco:reset_probe_failures
local key = KEYS[1]
local raw = redis.call('GET', key)
if not raw then
redis.call('SET', key, cjson.encode({probe_failures = 0}))
return
end
local state = cjson.decode(raw)
state['probe_failures'] = 0
redis.call('SET', key, cjson.encode(state))
"""
# activate merges: it refreshes the episode metadata (kind / activated_at /
# retry_after / affected_agents) but carries over the previous probe_failures
# count. A blind SET here (the old impl) reset probe_failures to 0, so a
# probe-failure increment that just landed — or was in flight — could be wiped
# by a concurrent re-park, resetting the give-up / CEO-notify count mid-episode.
# Like increment/reset, the read-merge-write runs server-side as one atomic Lua
# EVAL, so it is indivisible w.r.t. the counter scripts.
_ACTIVATE_RATE_LIMIT = """\
-- roboco:activate_rate_limit
local key = KEYS[1]
local fresh = ARGV[1]
local raw = redis.call('GET', key)
if raw then
local prev = cjson.decode(raw)
local old_pf = prev['probe_failures']
if old_pf ~= nil then
local new_state = cjson.decode(fresh)
new_state['probe_failures'] = old_pf
redis.call('SET', key, cjson.encode(new_state))
return old_pf
end
end
redis.call('SET', key, fresh)
return 0
"""
def resume_at(hit_at: str | None, retry_after: float | None) -> str | None:
"""Estimated lift time = hit_at + retry_after, ISO; falls back to hit_at."""
if not hit_at or retry_after is None:
return hit_at
try:
lifted = datetime.fromisoformat(hit_at) + timedelta(seconds=retry_after)
except (ValueError, TypeError):
return hit_at
return lifted.isoformat()
class RateLimitStateTracker:
"""Track rate-limit state for a single AI provider in Redis.
Usage
-----
tracker = RateLimitStateTracker("anthropic")
await tracker.activate(retry_after=60.0, affected_agents=["be-dev-1"])
assert await tracker.is_rate_limited()
A second instance that uses the same Redis URL and provider name
will observe the same state — no in-process singleton required.
"""
_KEY_PREFIX: str = "roboco:rate_limit:"
def __init__(self, provider: str, redis_url: str | None = None) -> None:
"""Construct a tracker for *provider*.
Args:
provider: Logical provider name, e.g. ``"anthropic"`` or
``"ollama_cloud"``. Used as part of the Redis key.
redis_url: Override the Redis URL (defaults to
``settings.redis_url``).
"""
self._provider = provider
self._redis_url = redis_url or settings.redis_url
self._redis: redis.Redis | None = None
# ------------------------------------------------------------------
# Private helpers
# ------------------------------------------------------------------
async def _conn(self) -> redis.Redis:
"""Return a (lazy-connected) redis.asyncio.Redis client."""
if self._redis is None:
self._redis = redis.from_url(self._redis_url)
return self._redis
def _key(self) -> str:
"""Redis key for this provider's state blob."""
return f"{self._KEY_PREFIX}{self._provider}:state"
# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
async def activate(
self,
retry_after: float | None = None,
affected_agents: list[str] | None = None,
kind: str = "rate_limited",
) -> None:
"""Mark the provider as unavailable so new spawns are queued.
Args:
retry_after: Seconds until the provider should accept new
requests, or ``None`` if unknown.
affected_agents: Agent slugs that were active when the limit
was hit (informational; stored in state).
kind: Why the provider is parked — ``"rate_limited"``
(a 429) or ``"overloaded"`` (a persistent 5xx).
Both gate spawns identically; the kind is stored
so the panel / notifications can distinguish them.
"""
r = await self._conn()
state: dict[str, Any] = {
"rate_limited": True,
"kind": kind,
"activated_at": datetime.now(UTC).isoformat(),
"retry_after": retry_after,
"affected_agents": affected_agents or [],
"probe_failures": 0,
}
# Atomic merge (see _ACTIVATE_RATE_LIMIT): the previous probe_failures
# count is carried over so a re-park cannot wipe an in-flight increment.
await r.eval(_ACTIVATE_RATE_LIMIT, 1, self._key(), json.dumps(state))
async def clear(self) -> None:
"""Remove rate-limit state for this provider."""
r = await self._conn()
await r.delete(self._key())
async def is_rate_limited(self) -> bool:
"""Return ``True`` if the provider is currently rate-limited."""
state = await self.get_state()
return bool(state.get("rate_limited", False))
async def get_state(self) -> dict[str, Any]:
"""Return the stored state dict, or ``{}`` if none exists."""
r = await self._conn()
raw = await r.get(self._key())
if raw is None:
return {}
decoded: str = raw.decode() if isinstance(raw, bytes) else str(raw)
result: dict[str, Any] = json.loads(decoded)
return result
async def increment_probe_failures(self) -> int:
"""Increment the probe-failure counter and return the new value.
The probe-failure counter tracks how many successive connectivity
probes have failed since the rate limit was activated. The
orchestrator uses this to decide whether to keep waiting or give
up entirely.
Atomic: the read-modify-write runs server-side as a Lua ``EVAL`` so a
concurrent ``activate()`` re-park cannot interleave between the GET and
SET and clobber the fresh episode blob with a stale one.
"""
r = await self._conn()
new_count = await r.eval(_INCREMENT_PROBE_FAILURES, 1, self._key())
return int(new_count)
async def reset_probe_failures(self) -> None:
"""Reset the probe-failure counter to 0.
Atomic: server-side Lua ``EVAL`` (see ``increment_probe_failures``).
"""
r = await self._conn()
await r.eval(_RESET_PROBE_FAILURES, 1, self._key())
# ------------------------------------------------------------------
# Class-level helpers
# ------------------------------------------------------------------
@classmethod
async def list_rate_limited_providers(
cls,
redis_url: str | None = None,
) -> list[tuple[str, dict[str, Any]]]:
"""Scan Redis for all providers that are currently rate-limited.
Returns a list of ``(provider_name, state_dict)`` tuples — one
entry per provider whose stored state has ``rate_limited == True``.
Returns an empty list when nothing is rate-limited or Redis is
unreachable.
Args:
redis_url: Override the default Redis URL from settings.
"""
url = redis_url or settings.redis_url
pattern = f"{cls._KEY_PREFIX}*:state"
results: list[tuple[str, dict[str, Any]]] = []
# `async with` closes the client on exit (modern redis.asyncio API),
# avoiding a deprecated explicit close in a finally block.
async with redis.from_url(url) as r:
try:
cursor: int = 0
while True:
cursor, keys = await r.scan(cursor, match=pattern, count=100)
for raw_key in keys:
entry = await cls._read_rate_limited_entry(r, raw_key)
if entry is not None:
results.append(entry)
if cursor == 0:
break
except Exception:
pass
return results
@staticmethod
def _decode(value: Any) -> str:
"""Decode a Redis value (bytes or str) to str."""
return value.decode() if isinstance(value, bytes) else str(value)
@classmethod
async def _read_rate_limited_entry(
cls, r: Any, raw_key: Any
) -> tuple[str, dict[str, Any]] | None:
"""``(provider, state)`` for a scan key, iff it holds a rate-limited record.
Returns None for keys that don't match the ``...:{provider}:state`` shape,
have no stored value, or whose state is not currently rate-limited.
"""
key = cls._decode(raw_key)
inner = key[len(cls._KEY_PREFIX) :]
if not inner.endswith(":state"):
return None
provider = inner[: -len(":state")]
raw_val = await r.get(key)
if raw_val is None:
return None
state: dict[str, Any] = json.loads(cls._decode(raw_val))
return (provider, state) if state.get("rate_limited") else None