[aaac85d2] Rate limit guardrails for Anthropic and Ollama providers (#104)

* [25aa5b24] Implement rate-limit Zustand store, Axios interceptor, WebSocket hook, banner component, and page-load sync (#99) (#101)

* [25aa5b24] feat(rate-limits): add types, Zustand store, Axios 429 interceptor, WS hook, sync hook, and banner component

- panel/src/types/rate-limits.ts: RateLimitEntry, RateLimitHitEvent, RateLimitLiftedEvent, RateLimitApiResponse
- panel/src/store/rate-limit-store.ts: useRateLimitStore with Map state, hitRateLimit/liftRateLimit/syncFromApi
- panel/src/lib/api/rate-limits.ts: GET /api/system/rate-limits with isMockMode guard
- panel/src/lib/api/client.ts: 429 interceptor dispatches to store first, Sonner toast on retry exhaustion
- panel/src/hooks/use-rate-limit-websocket.ts: RATE_LIMIT_HIT/LIFTED events + onReconnect callback
- panel/src/hooks/use-rate-limit-sync.ts: mount sync + no-op with console.warn when endpoint unavailable
- panel/src/components/rate-limit/rate-limit-banner.tsx: amber rows with countdown, no dismiss button
- panel/src/app/(dashboard)/layout.tsx: RateLimitBanner mounted below Header
- store/index.ts, hooks/index.ts: export new store and hooks

* [25aa5b24] fix(rate-limit-banner): use lint-clean countdown pattern (computeSecondsLeft outside render)

* [25aa5b24] fix(client): add real retry loop to 429 interceptor so Sonner toast fires on exhaustion

- Increment error.config._retryCount and return api(error.config) when
  retryCount < RATE_LIMIT_MAX_RETRIES, actually retrying the request.
- Toast fires only when retryCount >= RATE_LIMIT_MAX_RETRIES (3 attempts).
- Fixes AC4: toast was dead code because without return api(error.config)
  every 429 saw retryCount=1, permanently below the threshold of 3.

---------

Co-authored-by: Frontend Developer 1 <fe-dev-1@agents.roboco.dev>

* [4112cd34] feat(rate-limit): add RateLimitError with 5-retry exponential backoff at all LLM call sites (#102) (#103)

- Create roboco/services/exceptions.py with RateLimitError(provider, retry_after),
  HTTP_TOO_MANY_REQUESTS, MAX_RATE_LIMIT_RETRIES constants, and
  parse_retry_after_header() helper
- extraction.py: extract _call_anthropic_with_retry() helper; retry Anthropic
  call 5x on 429 with exponential backoff; re-raise RateLimitError from outer
  except instead of swallowing it
- ollama_embedder.py: 5-retry outer loop (429) wrapping existing 3-retry inner
  loop (ConnectError/Timeout) for all 4 call sites; two concerns kept isolated
- indexes/base.py, mentor.py, validator.py: replace magic 429 literals with
  HTTP_TOO_MANY_REQUESTS; 5-retry loop on 429 for LLM calls
- middleware.py: add rate_limit_exception_handler returning HTTP 429 with
  Retry-After response header
- tests/unit/services/test_rate_limit_retry.py: 28 tests covering exhaustion,
  Retry-After header sleep, partial retries then success, ConnectError isolation

Co-authored-by: Backend Developer 1 <be-dev-1@agents.roboco.dev>

* [18107054] feat(rate-limit): Redis rate-limit state tracker + i_am_blocked rate_limited path (#105) (#106)

- Add RateLimitStateTracker in roboco/services/gateway/rate_limit_tracker.py
  with activate(), clear(), is_rate_limited(), get_state(),
  increment_probe_failures(), reset_probe_failures() backed by redis.asyncio
- Add RATE_LIMIT_HIT = "rate_limit.hit" to EventType StrEnum in events.py
- Add _handle_rate_limited_parking() to Choreographer: intercepts
  i_am_blocked(reason='rate_limited') before block state transition,
  parks all active agents sharing affected provider via mark_waiting_long,
  publishes RATE_LIMIT_HIT event to StreamEventBus, task stays in_progress
- Add get_provider_for_agent() and get_active_agent_slugs_for_provider()
  helper methods to AgentOrchestrator
- Wire orchestrator and stream_bus into ChoreographerDeps via deps.py
- Add test_rate_limit_tracker.py (basic ops, probe failures, cross-reconnection
  persistence, provider isolation) and test_i_am_blocked_rate_limited.py
  (AC3/AC4/AC5 coverage: task stays in_progress, mark_waiting_long call count
  equals active agent count, RATE_LIMIT_HIT event payload structure)

Co-authored-by: Backend Developer 1 <be-dev-1@agents.roboco.dev>

* [5501e4b4] Wire RateLimitStateTracker into live orchestrator paths — 4 CEO-identified integration gaps (#109)

* [8451ca50] feat(gateway): wire RateLimitStateTracker.activate() into i_am_blocked rate-limited path and add provider-rate-limit gate to decide_spawn() (#107)

- Add provider/provider_rate_limited optional fields to TriggerContext (backward-compatible defaults)
- Insert rule 2 in decide_spawn(): QUEUE when trigger.provider_rate_limited is True with reason 'provider X rate-limited'
- Call RateLimitStateTracker(provider).activate() in _handle_rate_limited_parking() after mark_waiting_long loop (wrapped in contextlib.suppress for Redis fault tolerance)
- Extend gateway_pre_spawn_check() with optional provider param; check RateLimitStateTracker.is_rate_limited() when provider is known
- Pass provider=self.get_provider_for_agent(agent_id) from orchestrator call site
- Add TestProviderRateLimitGate (6 tests) to test_trigger_filter.py
- Add TestRateLimitTrackerActivateOnParking (6 tests) to test_i_am_blocked_rate_limited.py
- All 38 unit tests pass; ruff and mypy clean on changed files

Co-authored-by: Backend Developer 1 <be-dev-1@agents.roboco.dev>

* [e9cef0f0] feat(rate-limits): sweeper probe loop, CEO notification, and GET /api/system/rate-limits endpoint (AC4, AC8, AC9) (#108)

- Add RATE_LIMIT_LIFTED event type to EventType enum in models/events.py
- Add RateLimitStateTracker.list_rate_limited_providers() classmethod to scan
  Redis for all currently rate-limited providers (used by the new endpoint)
- Add orchestrator._rate_limit_probe_loop(): background task started/stopped in
  start()/stop(), runs _sweep_rate_limit_probes() every 30s
- Add orchestrator._probe_one_provider(): checks estimated_lift_at gate, calls
  _do_probe(); on success: tracker.clear(), resolve_wait() for all parked agents
  with waiting_for='rate_limit_lifted' matching the provider, publishes
  RATE_LIMIT_LIFTED event; on failure: increments probe_failures counter, sends
  CEO notification at threshold 10 (once per episode via _rate_limit_ceo_notified)
- Add orchestrator._make_tracker(): injectable factory for RateLimitStateTracker
- Add orchestrator._do_probe(): overridable async bool probe (default: True)
- Add orchestrator._notify_rate_limit_ceo(): high-priority notification to CEO
  containing provider name, duration since activation, and paused agent count
- Add roboco/api/routes/system.py with GET /rate-limits endpoint (AC9)
- Register system_router in app.py under /api/system prefix
- Add 17 unit tests in tests/unit/runtime/test_rate_limit_sweep.py covering
  all AC4/AC8/AC9 paths: probe success/failure, CEO threshold, endpoint schema

Co-authored-by: Backend Developer 1 <be-dev-1@agents.roboco.dev>

---------

Co-authored-by: Backend Developer 1 <be-dev-1@agents.roboco.dev>

---------

Co-authored-by: Frontend Developer 1 <fe-dev-1@agents.roboco.dev>
Co-authored-by: Backend Developer 1 <be-dev-1@agents.roboco.dev>
Co-authored-by: Renn F <rennf93@users.noreply.github.com>
This commit is contained in:
Renzo F
2026-06-11 09:21:26 +02:00
committed by GitHub
co-authored by Backend Developer 1 Frontend Developer 1 Renn F
parent cc4ccb7ea3
commit 98e618c243
30 changed files with 3919 additions and 273 deletions
@@ -0,0 +1,508 @@
"""Unit tests for the rate-limited path in Choreographer.i_am_blocked.
Acceptance criteria verified here:
- AC1: i_am_blocked(reason='rate_limited') calls RateLimitStateTracker.activate()
and stores affected agent IDs; all active agents on the rate-limited
provider are subsequently marked waiting-long.
- AC3: POST /v1/i_am_blocked with reason='rate_limited' does NOT transition
the task to 'blocked'; the task remains in its current status
(in_progress) and the calling agent is parked via
mark_waiting_long(waiting_for='rate_limit_lifted').
- AC4: mark_waiting_long is called for every orchestrator-tracked active agent
sharing the affected provider — call count equals active agent count.
- AC5: A RATE_LIMIT_HIT event is published to the StreamEventBus with fields
provider, affectedAgents, retryAfterSeconds, and timestamp.
"""
from __future__ import annotations
from unittest.mock import AsyncMock, MagicMock, patch
from uuid import uuid4
from roboco.models.events import EventType
from roboco.services.gateway.choreographer import Choreographer, ChoreographerDeps
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
_ACTIVE_AGENTS = ["be-dev-1", "be-dev-2", "be-qa"]
_PROVIDER = "anthropic"
def _make_evidence_repo() -> AsyncMock:
repo = AsyncMock()
for method in (
"list_unread_a2a",
"list_unread_mentions",
"list_pending_notifications",
"task_metadata_gaps",
"recent_team_activity",
"blockers_in_lane",
"journal_highlights_for_task",
):
getattr(repo, method).return_value = []
return repo
def _make_task_svc(agent_id: object, task_id: object) -> AsyncMock:
t = MagicMock(
id=task_id,
status="in_progress",
assigned_to=agent_id,
pre_block_state=None,
task_type="code",
team="backend",
# Avoid issues with spec iteration in claim guards
dependency_ids=[],
# acceptance_criteria needed by some paths
acceptance_criteria=[],
quick_context=None,
)
task_svc = AsyncMock()
task_svc.session = MagicMock()
task_svc.session.begin_nested = MagicMock(
return_value=MagicMock(
__aenter__=AsyncMock(return_value=None),
__aexit__=AsyncMock(return_value=False),
)
)
task_svc.get.return_value = t
task_svc.agent_for.return_value = MagicMock(
id=agent_id,
role="developer",
team="backend",
slug="be-dev-1", # calling agent's slug
)
return task_svc
def _make_orchestrator(
active_agents: list[str] | None = None,
provider: str = _PROVIDER,
) -> MagicMock:
"""Build a synchronous/async orchestrator mock."""
agents = active_agents if active_agents is not None else _ACTIVE_AGENTS
orch = MagicMock()
orch.get_provider_for_agent = MagicMock(return_value=provider)
orch.get_active_agent_slugs_for_provider = MagicMock(return_value=agents)
orch.mark_waiting_long = AsyncMock(return_value=None)
return orch
def _make_stream_bus() -> AsyncMock:
bus = AsyncMock()
bus.publish = AsyncMock(return_value="msg-id-1")
return bus
def _make_deps(
agent_id: object,
task_id: object,
orchestrator: MagicMock | None = None,
stream_bus: AsyncMock | None = None,
) -> ChoreographerDeps:
return ChoreographerDeps(
task=_make_task_svc(agent_id, task_id),
work_session=AsyncMock(),
git=AsyncMock(),
a2a=AsyncMock(),
journal=AsyncMock(),
audit=AsyncMock(),
evidence_repo=_make_evidence_repo(),
orchestrator=orchestrator,
stream_bus=stream_bus,
)
# ---------------------------------------------------------------------------
# AC3: Task stays in in_progress, agent parked via mark_waiting_long
# ---------------------------------------------------------------------------
class TestRateLimitedDoesNotBlockTask:
async def test_task_status_remains_in_progress(self) -> None:
"""reason='rate_limited' must NOT transition the task to 'blocked'."""
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator()
deps = _make_deps(agent_id, task_id, orchestrator=orch)
c = Choreographer(deps)
env = await c.i_am_blocked(agent_id, task_id, "rate_limited")
assert env.error is None
assert env.status == "in_progress"
async def test_verb_runner_block_action_not_called(self) -> None:
"""The `block` action (task.escalate) must NOT run on rate_limited path."""
agent_id = uuid4()
task_id = uuid4()
deps = _make_deps(agent_id, task_id)
c = Choreographer(deps)
await c.i_am_blocked(agent_id, task_id, "rate_limited")
# The VerbRunner calls task.escalate for the normal block path.
# In the rate-limited path this must NOT happen.
deps.task.escalate.assert_not_awaited()
async def test_calling_agent_parked_via_mark_waiting_long(self) -> None:
"""mark_waiting_long must be called with waiting_for='rate_limit_lifted'."""
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator(active_agents=["be-dev-1"])
deps = _make_deps(agent_id, task_id, orchestrator=orch)
c = Choreographer(deps)
await c.i_am_blocked(agent_id, task_id, "rate_limited")
# Verify that at least one mark_waiting_long call uses the right reason.
# The implementation calls mark_waiting_long(slug, waiting_for=..., ...)
# so waiting_for is always a keyword argument.
waiting_for_values = [
c.kwargs.get("waiting_for") for c in orch.mark_waiting_long.call_args_list
]
assert "rate_limit_lifted" in waiting_for_values
async def test_case_insensitive_reason_match(self) -> None:
"""reason='Rate_Limited' (any case) should trigger the special path."""
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator(active_agents=["be-dev-1"])
deps = _make_deps(agent_id, task_id, orchestrator=orch)
c = Choreographer(deps)
env = await c.i_am_blocked(agent_id, task_id, "Rate_Limited")
assert env.error is None
assert env.status == "in_progress"
async def test_struggle_journal_still_written(self) -> None:
"""journal.write_struggle must still be written on the rate_limited path."""
agent_id = uuid4()
task_id = uuid4()
deps = _make_deps(agent_id, task_id)
c = Choreographer(deps)
await c.i_am_blocked(agent_id, task_id, "rate_limited")
deps.journal.write_struggle.assert_awaited_once()
# ---------------------------------------------------------------------------
# AC4: mark_waiting_long called for every active agent on affected provider
# ---------------------------------------------------------------------------
class TestMarkWaitingLongCallCount:
async def test_call_count_equals_active_agent_count(self) -> None:
"""mark_waiting_long must be called once per active agent."""
agent_id = uuid4()
task_id = uuid4()
active = ["be-dev-1", "be-dev-2", "be-dev-3"]
orch = _make_orchestrator(active_agents=active)
deps = _make_deps(agent_id, task_id, orchestrator=orch)
c = Choreographer(deps)
await c.i_am_blocked(agent_id, task_id, "rate_limited")
assert orch.mark_waiting_long.call_count == len(active)
async def test_call_count_with_single_active_agent(self) -> None:
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator(active_agents=["be-dev-1"])
deps = _make_deps(agent_id, task_id, orchestrator=orch)
c = Choreographer(deps)
await c.i_am_blocked(agent_id, task_id, "rate_limited")
assert orch.mark_waiting_long.call_count == 1
async def test_no_calls_when_no_active_agents(self) -> None:
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator(active_agents=[])
deps = _make_deps(agent_id, task_id, orchestrator=orch)
c = Choreographer(deps)
await c.i_am_blocked(agent_id, task_id, "rate_limited")
assert orch.mark_waiting_long.call_count == 0
async def test_no_calls_when_orchestrator_is_none(self) -> None:
"""When orchestrator is not wired in, no parking happens but no crash."""
agent_id = uuid4()
task_id = uuid4()
deps = _make_deps(agent_id, task_id, orchestrator=None)
c = Choreographer(deps)
env = await c.i_am_blocked(agent_id, task_id, "rate_limited")
# Should still succeed; no orchestrator = no parking
assert env.error is None
assert env.status == "in_progress"
async def test_mark_waiting_long_receives_waiting_for_arg(self) -> None:
"""Every mark_waiting_long call must carry waiting_for='rate_limit_lifted'."""
agent_id = uuid4()
task_id = uuid4()
active = ["be-dev-1", "be-qa"]
orch = _make_orchestrator(active_agents=active)
deps = _make_deps(agent_id, task_id, orchestrator=orch)
c = Choreographer(deps)
await c.i_am_blocked(agent_id, task_id, "rate_limited")
for c_args in orch.mark_waiting_long.call_args_list:
# mark_waiting_long(slug, waiting_for=..., ...) — waiting_for is a kwarg
assert c_args.kwargs.get("waiting_for") == "rate_limit_lifted"
# ---------------------------------------------------------------------------
# AC5: RATE_LIMIT_HIT event published with correct payload structure
# ---------------------------------------------------------------------------
class TestRateLimitHitEventPublished:
async def test_stream_bus_publish_called_once(self) -> None:
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator()
bus = _make_stream_bus()
deps = _make_deps(agent_id, task_id, orchestrator=orch, stream_bus=bus)
c = Choreographer(deps)
await c.i_am_blocked(agent_id, task_id, "rate_limited")
bus.publish.assert_awaited_once()
async def test_event_type_is_rate_limit_hit(self) -> None:
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator()
bus = _make_stream_bus()
deps = _make_deps(agent_id, task_id, orchestrator=orch, stream_bus=bus)
c = Choreographer(deps)
await c.i_am_blocked(agent_id, task_id, "rate_limited")
event = bus.publish.call_args.args[0]
assert event.type == EventType.RATE_LIMIT_HIT
async def test_event_data_has_provider_field(self) -> None:
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator(provider="anthropic")
bus = _make_stream_bus()
deps = _make_deps(agent_id, task_id, orchestrator=orch, stream_bus=bus)
c = Choreographer(deps)
await c.i_am_blocked(agent_id, task_id, "rate_limited")
event = bus.publish.call_args.args[0]
assert "provider" in event.data
assert event.data["provider"] == "anthropic"
async def test_event_data_has_affected_agents_list(self) -> None:
agent_id = uuid4()
task_id = uuid4()
active = ["be-dev-1", "be-dev-2"]
orch = _make_orchestrator(active_agents=active)
bus = _make_stream_bus()
deps = _make_deps(agent_id, task_id, orchestrator=orch, stream_bus=bus)
c = Choreographer(deps)
await c.i_am_blocked(agent_id, task_id, "rate_limited")
event = bus.publish.call_args.args[0]
assert "affectedAgents" in event.data
assert isinstance(event.data["affectedAgents"], list)
assert event.data["affectedAgents"] == active
async def test_event_data_has_retry_after_seconds_null_by_default(self) -> None:
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator()
bus = _make_stream_bus()
deps = _make_deps(agent_id, task_id, orchestrator=orch, stream_bus=bus)
c = Choreographer(deps)
await c.i_am_blocked(agent_id, task_id, "rate_limited")
event = bus.publish.call_args.args[0]
assert "retryAfterSeconds" in event.data
assert event.data["retryAfterSeconds"] is None
async def test_event_data_retry_after_parsed_from_what_needed(self) -> None:
"""If what_needed is a numeric string, it becomes retryAfterSeconds."""
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator()
bus = _make_stream_bus()
deps = _make_deps(agent_id, task_id, orchestrator=orch, stream_bus=bus)
c = Choreographer(deps)
await c.i_am_blocked(agent_id, task_id, "rate_limited", what_needed="30")
event = bus.publish.call_args.args[0]
assert event.data["retryAfterSeconds"] == float("30")
async def test_event_data_has_timestamp_iso_string(self) -> None:
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator()
bus = _make_stream_bus()
deps = _make_deps(agent_id, task_id, orchestrator=orch, stream_bus=bus)
c = Choreographer(deps)
await c.i_am_blocked(agent_id, task_id, "rate_limited")
event = bus.publish.call_args.args[0]
assert "timestamp" in event.data
# ISO string: must be a non-empty string
ts = event.data["timestamp"]
assert isinstance(ts, str) and len(ts) > 0
async def test_no_publish_when_stream_bus_is_none(self) -> None:
"""When stream_bus is not wired in, no publish is attempted."""
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator()
# stream_bus=None: no bus
deps = _make_deps(agent_id, task_id, orchestrator=orch, stream_bus=None)
c = Choreographer(deps)
env = await c.i_am_blocked(agent_id, task_id, "rate_limited")
# Should still succeed
assert env.error is None
assert env.status == "in_progress"
# ---------------------------------------------------------------------------
# AC1: RateLimitStateTracker.activate() called on rate_limited path
# ---------------------------------------------------------------------------
_TRACKER_PATCH = "roboco.services.gateway.rate_limit_tracker.RateLimitStateTracker"
class TestRateLimitTrackerActivateOnParking:
"""Verify that _handle_rate_limited_parking() calls activate()."""
async def test_activate_called_when_provider_known(self) -> None:
"""activate() must be called once when provider != 'unknown'."""
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator(active_agents=["be-dev-1"], provider=_PROVIDER)
deps = _make_deps(agent_id, task_id, orchestrator=orch)
c = Choreographer(deps)
mock_tracker = AsyncMock()
mock_tracker.activate = AsyncMock(return_value=None)
mock_tracker_cls = MagicMock(return_value=mock_tracker)
with patch(_TRACKER_PATCH, mock_tracker_cls):
await c.i_am_blocked(agent_id, task_id, "rate_limited")
mock_tracker_cls.assert_called_once_with(_PROVIDER)
mock_tracker.activate.assert_awaited_once()
async def test_activate_receives_affected_agents(self) -> None:
"""activate() must be called with the affected_agents list."""
agent_id = uuid4()
task_id = uuid4()
active = ["be-dev-1", "be-dev-2"]
orch = _make_orchestrator(active_agents=active, provider=_PROVIDER)
deps = _make_deps(agent_id, task_id, orchestrator=orch)
c = Choreographer(deps)
mock_tracker = AsyncMock()
mock_tracker.activate = AsyncMock(return_value=None)
mock_tracker_cls = MagicMock(return_value=mock_tracker)
with patch(_TRACKER_PATCH, mock_tracker_cls):
await c.i_am_blocked(agent_id, task_id, "rate_limited")
call_kwargs = mock_tracker.activate.call_args.kwargs
assert call_kwargs.get("affected_agents") == active
async def test_activate_receives_retry_after_from_what_needed(self) -> None:
"""activate() must receive retry_after parsed from what_needed."""
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator(active_agents=["be-dev-1"], provider=_PROVIDER)
deps = _make_deps(agent_id, task_id, orchestrator=orch)
c = Choreographer(deps)
mock_tracker = AsyncMock()
mock_tracker.activate = AsyncMock(return_value=None)
mock_tracker_cls = MagicMock(return_value=mock_tracker)
with patch(_TRACKER_PATCH, mock_tracker_cls):
await c.i_am_blocked(agent_id, task_id, "rate_limited", what_needed="45")
call_kwargs = mock_tracker.activate.call_args.kwargs
assert call_kwargs.get("retry_after") == float("45")
async def test_activate_retry_after_none_when_what_needed_not_numeric(self) -> None:
"""activate() must receive retry_after=None when what_needed is not a number."""
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator(active_agents=["be-dev-1"], provider=_PROVIDER)
deps = _make_deps(agent_id, task_id, orchestrator=orch)
c = Choreographer(deps)
mock_tracker = AsyncMock()
mock_tracker.activate = AsyncMock(return_value=None)
mock_tracker_cls = MagicMock(return_value=mock_tracker)
with patch(_TRACKER_PATCH, mock_tracker_cls):
await c.i_am_blocked(
agent_id, task_id, "rate_limited", what_needed="retry soon"
)
call_kwargs = mock_tracker.activate.call_args.kwargs
assert call_kwargs.get("retry_after") is None
async def test_activate_skipped_when_provider_unknown(self) -> None:
"""activate() must NOT be called when provider resolves to 'unknown'."""
agent_id = uuid4()
task_id = uuid4()
# get_provider_for_agent returns None → provider stays 'unknown'
orch = MagicMock()
orch.get_provider_for_agent = MagicMock(return_value=None)
orch.get_active_agent_slugs_for_provider = MagicMock(return_value=[])
orch.mark_waiting_long = AsyncMock(return_value=None)
deps = _make_deps(agent_id, task_id, orchestrator=orch)
c = Choreographer(deps)
mock_tracker = AsyncMock()
mock_tracker.activate = AsyncMock(return_value=None)
mock_tracker_cls = MagicMock(return_value=mock_tracker)
with patch(_TRACKER_PATCH, mock_tracker_cls):
env = await c.i_am_blocked(agent_id, task_id, "rate_limited")
# No crash, no activate call
assert env.error is None
mock_tracker.activate.assert_not_awaited()
async def test_activate_failure_does_not_crash_path(self) -> None:
"""If activate() raises, _handle_rate_limited_parking must still succeed."""
agent_id = uuid4()
task_id = uuid4()
orch = _make_orchestrator(active_agents=["be-dev-1"], provider=_PROVIDER)
deps = _make_deps(agent_id, task_id, orchestrator=orch)
c = Choreographer(deps)
mock_tracker = AsyncMock()
mock_tracker.activate = AsyncMock(side_effect=RuntimeError("redis down"))
mock_tracker_cls = MagicMock(return_value=mock_tracker)
with patch(_TRACKER_PATCH, mock_tracker_cls):
env = await c.i_am_blocked(agent_id, task_id, "rate_limited")
assert env.error is None
assert env.status == "in_progress"
+103 -1
View File
@@ -34,17 +34,21 @@ def _task(
return t
def _trigger(
def _trigger( # noqa: PLR0913
kind: TriggerKind,
skill: str | None = None,
recent_spawns_for_task: int = 0,
recent_spawns_for_role: int = 0,
provider: str | None = None,
provider_rate_limited: bool = False,
) -> TriggerContext:
return TriggerContext(
kind=kind,
skill=skill,
recent_spawns_for_task=recent_spawns_for_task,
recent_spawns_for_role=recent_spawns_for_role,
provider=provider,
provider_rate_limited=provider_rate_limited,
)
@@ -148,3 +152,101 @@ class TestCooldown:
)
assert decision.outcome == SpawnDecision.QUEUE
assert "rate" in decision.reason.lower()
class TestProviderRateLimitGate:
"""Rule 2: provider rate-limit gate fires before claimant-lock/cooldown."""
def test_queues_when_provider_rate_limited(self) -> None:
"""QUEUE outcome when provider_rate_limited=True."""
t = _task(status="in_progress")
decision = decide_spawn(
task=t,
trigger=_trigger(
TriggerKind.NOTIFICATION,
provider="anthropic",
provider_rate_limited=True,
),
config=_DEFAULT_CONFIG,
)
assert decision.outcome == SpawnDecision.QUEUE
assert "provider anthropic rate-limited" in decision.reason
def test_reason_contains_provider_name(self) -> None:
"""Reason string must contain the provider name."""
t = _task(status="pending")
decision = decide_spawn(
task=t,
trigger=_trigger(
TriggerKind.SCAN,
provider="ollama_cloud",
provider_rate_limited=True,
),
config=_DEFAULT_CONFIG,
)
assert "ollama_cloud" in decision.reason
def test_reason_contains_unknown_when_no_provider_name(self) -> None:
"""When provider is None, reason still contains 'unknown'."""
t = _task(status="in_progress")
decision = decide_spawn(
task=t,
trigger=_trigger(
TriggerKind.NOTIFICATION,
provider=None,
provider_rate_limited=True,
),
config=_DEFAULT_CONFIG,
)
assert decision.outcome == SpawnDecision.QUEUE
assert "unknown" in decision.reason
def test_no_queue_injection_when_not_rate_limited(self) -> None:
"""SPAWN when provider_rate_limited=False and all other gates clear."""
t = _task(status="in_progress")
decision = decide_spawn(
task=t,
trigger=_trigger(
TriggerKind.NOTIFICATION,
provider="anthropic",
provider_rate_limited=False,
),
config=_DEFAULT_CONFIG,
)
assert decision.outcome == SpawnDecision.SPAWN
def test_stale_drop_fires_before_rate_limit_gate(self) -> None:
"""Rule 1 (stale-drop) fires before rule 2 (rate-limit gate)."""
t = _task(status="completed")
decision = decide_spawn(
task=t,
trigger=_trigger(
TriggerKind.NOTIFICATION,
provider="anthropic",
provider_rate_limited=True,
),
config=_DEFAULT_CONFIG,
)
# Rule 1 fires first — outcome must be DROP, not QUEUE
assert decision.outcome == SpawnDecision.DROP
def test_rate_limit_gate_fires_before_claimant_lock(self) -> None:
"""Rule 2 (rate-limit gate) fires before rule 3 (single-claimant invariant)."""
recent = datetime.now(tz=UTC)
t = _task(
status="in_progress",
active_claimant_id=uuid4(),
last_heartbeat_at=recent,
)
decision = decide_spawn(
task=t,
trigger=_trigger(
TriggerKind.NOTIFICATION,
provider="anthropic",
provider_rate_limited=True,
),
config=_DEFAULT_CONFIG,
)
# Both gates would QUEUE but reason must come from rate-limit (rule 2)
assert decision.outcome == SpawnDecision.QUEUE
assert "rate-limited" in decision.reason
+550
View File
@@ -0,0 +1,550 @@
"""Unit tests for the rate-limit sweeper probe loop (AC4, AC8).
Tests cover:
- probe-success path: tracker.clear() + resolve_wait + RATE_LIMIT_LIFTED event
- probe-failure path: increment_probe_failures is called
- CEO notification fires at threshold 10 exactly once per episode
- ``_do_probe`` / ``_make_tracker`` are injectable boundaries for mocking
"""
from __future__ import annotations
import fnmatch
import json
from datetime import UTC, datetime
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
from uuid import uuid4
from httpx import ASGITransport, AsyncClient
from roboco.api.app import create_app
from roboco.models.events import EventType
from roboco.models.runtime import WaitingRecord
from roboco.runtime.orchestrator import AgentOrchestrator
from roboco.services.gateway.rate_limit_tracker import RateLimitStateTracker
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_redis_mock(initial_store: dict[str, Any] | None = None) -> AsyncMock:
"""Fake redis.asyncio.Redis backed by a plain dict."""
store: dict[str, Any] = initial_store if initial_store is not None else {}
async def _get(key: str) -> bytes | None:
val = store.get(key)
if val is None:
return None
return str(val).encode() if not isinstance(val, bytes) else val
async def _set(key: str, value: Any) -> None:
store[key] = value
async def _delete(key: str) -> int:
return 1 if store.pop(key, None) is not None else 0
async def _scan(
_cursor: int, match: str = "*", count: int = 100 # noqa: ARG001
) -> tuple[int, list[bytes]]:
# Simple in-memory scan: return all matching keys in one shot
matches = [k.encode() for k in store if fnmatch.fnmatch(k, match)]
return (0, matches)
async def _aclose() -> None:
pass
mock = AsyncMock()
mock.get = AsyncMock(side_effect=_get)
mock.set = AsyncMock(side_effect=_set)
mock.delete = AsyncMock(side_effect=_delete)
mock.scan = AsyncMock(side_effect=_scan)
mock.aclose = AsyncMock(side_effect=_aclose)
mock._store = store
return mock
def _make_orchestrator() -> AgentOrchestrator:
"""Build a minimal orchestrator via __new__ (no __init__ side-effects)."""
orch = AgentOrchestrator.__new__(AgentOrchestrator)
orch._running = True
orch._waiting_records: dict[str, WaitingRecord] = {}
orch._instances: dict[str, Any] = {}
orch._rate_limit_ceo_notified: set[str] = set()
return orch
def _make_tracker_mock(failure_return: int = 1) -> AsyncMock:
"""Create an async mock RateLimitStateTracker instance."""
mock = AsyncMock()
mock.clear = AsyncMock()
mock.increment_probe_failures = AsyncMock(return_value=failure_return)
mock.reset_probe_failures = AsyncMock()
return mock
def _make_active_state(
_provider: str = "anthropic",
retry_after: float | None = None,
probe_failures: int = 0,
activated_at: datetime | None = None,
) -> dict[str, Any]:
"""Build a tracker state dict."""
at = activated_at or datetime.now(UTC)
return {
"rate_limited": True,
"activated_at": at.isoformat(),
"retry_after": retry_after,
"affected_agents": ["be-dev-1"],
"probe_failures": probe_failures,
}
def _waiting_record(
agent_id: str,
provider: str = "anthropic",
task_id: str | None = None,
) -> WaitingRecord:
return WaitingRecord(
agent_id=agent_id,
task_id=task_id or str(uuid4()),
waiting_for="rate_limit_lifted",
waiting_since=datetime.now(UTC),
context={"provider": provider},
)
# ---------------------------------------------------------------------------
# Tests: probe-success path (AC4)
# ---------------------------------------------------------------------------
class TestProbeSuccessPath:
"""When _do_probe returns True the rate limit should be cleared and
all parked agents resolved."""
async def test_tracker_clear_called_on_success(self) -> None:
"""tracker.clear() is invoked when the probe succeeds."""
orch = _make_orchestrator()
provider = "anthropic"
state = _make_active_state(provider, retry_after=None)
tracker_mock = _make_tracker_mock()
orch._make_tracker = MagicMock(return_value=tracker_mock) # type: ignore[method-assign]
orch.resolve_wait = AsyncMock(return_value=None)
async def fake_do_probe(_p: str) -> bool:
return True
orch._do_probe = fake_do_probe # type: ignore[method-assign]
with patch("roboco.events.get_event_bus") as mock_bus_fn:
bus_mock = AsyncMock()
bus_mock.publish = AsyncMock()
mock_bus_fn.return_value = bus_mock
await orch._probe_one_provider(provider, state)
tracker_mock.clear.assert_awaited_once()
async def test_resolve_wait_called_for_parked_agents(self) -> None:
"""resolve_wait is called for each agent waiting for rate_limit_lifted."""
orch = _make_orchestrator()
provider = "anthropic"
state = _make_active_state(provider, retry_after=None)
agent1 = "be-dev-1"
agent2 = "be-dev-2"
orch._waiting_records = {
agent1: _waiting_record(agent1, provider),
agent2: _waiting_record(agent2, provider),
"be-qa-1": _waiting_record(
"be-qa-1", "other-provider"
), # different provider
}
orch.resolve_wait = AsyncMock(return_value=None)
tracker_mock = _make_tracker_mock()
orch._make_tracker = MagicMock(return_value=tracker_mock) # type: ignore[method-assign]
async def fake_do_probe(_p: str) -> bool:
return True
orch._do_probe = fake_do_probe # type: ignore[method-assign]
with patch("roboco.events.get_event_bus") as mock_bus_fn:
bus_mock = AsyncMock()
bus_mock.publish = AsyncMock()
mock_bus_fn.return_value = bus_mock
await orch._probe_one_provider(provider, state)
# Only the two anthropic-parked agents should be resolved
assert orch.resolve_wait.await_count == 2 # noqa: PLR2004
resolved_ids = {call.args[0] for call in orch.resolve_wait.call_args_list}
assert agent1 in resolved_ids
assert agent2 in resolved_ids
assert "be-qa-1" not in resolved_ids
async def test_rate_limit_lifted_event_published(self) -> None:
"""RATE_LIMIT_LIFTED event is published to the bus on probe success."""
orch = _make_orchestrator()
provider = "anthropic"
state = _make_active_state(provider, retry_after=None)
orch.resolve_wait = AsyncMock(return_value=None)
tracker_mock = _make_tracker_mock()
orch._make_tracker = MagicMock(return_value=tracker_mock) # type: ignore[method-assign]
published_events: list[Any] = []
async def fake_do_probe(_p: str) -> bool:
return True
orch._do_probe = fake_do_probe # type: ignore[method-assign]
with patch("roboco.events.get_event_bus") as mock_bus_fn:
bus_mock = AsyncMock()
bus_mock.publish = AsyncMock(side_effect=published_events.append)
mock_bus_fn.return_value = bus_mock
await orch._probe_one_provider(provider, state)
assert len(published_events) == 1
event = published_events[0]
assert event.type == EventType.RATE_LIMIT_LIFTED
assert event.data["provider"] == provider
async def test_ceo_notified_flag_cleared_on_success(self) -> None:
"""_rate_limit_ceo_notified is cleared when probe succeeds."""
orch = _make_orchestrator()
provider = "anthropic"
orch._rate_limit_ceo_notified.add(provider) # simulates prior episode
state = _make_active_state(provider, retry_after=None)
orch.resolve_wait = AsyncMock(return_value=None)
tracker_mock = _make_tracker_mock()
orch._make_tracker = MagicMock(return_value=tracker_mock) # type: ignore[method-assign]
async def fake_do_probe(_p: str) -> bool:
return True
orch._do_probe = fake_do_probe # type: ignore[method-assign]
with patch("roboco.events.get_event_bus") as mock_bus_fn:
bus_mock = AsyncMock()
bus_mock.publish = AsyncMock()
mock_bus_fn.return_value = bus_mock
await orch._probe_one_provider(provider, state)
assert provider not in orch._rate_limit_ceo_notified
async def test_probe_skipped_before_estimated_lift_at(self) -> None:
"""When retry_after has not elapsed yet the probe is skipped entirely."""
orch = _make_orchestrator()
provider = "anthropic"
# Set activated_at to now; retry_after = 300s → estimated lift in future
state = _make_active_state(
provider,
retry_after=300.0,
activated_at=datetime.now(UTC),
)
probe_called = []
async def fake_do_probe(_p: str) -> bool:
probe_called.append(_p)
return True
orch._do_probe = fake_do_probe # type: ignore[method-assign]
await orch._probe_one_provider(provider, state)
assert probe_called == [] # probe was gated by time
# ---------------------------------------------------------------------------
# Tests: probe-failure path (AC4)
# ---------------------------------------------------------------------------
class TestProbeFailurePath:
"""When _do_probe returns False the failure counter should be incremented."""
async def test_increment_probe_failures_called_on_failure(self) -> None:
"""increment_probe_failures is called when the probe fails."""
orch = _make_orchestrator()
provider = "anthropic"
state = _make_active_state(provider, retry_after=None)
tracker_mock = _make_tracker_mock(failure_return=1)
orch._make_tracker = MagicMock(return_value=tracker_mock) # type: ignore[method-assign]
async def fake_do_probe(_p: str) -> bool:
return False
orch._do_probe = fake_do_probe # type: ignore[method-assign]
orch._notify_rate_limit_ceo = AsyncMock()
await orch._probe_one_provider(provider, state)
tracker_mock.increment_probe_failures.assert_awaited_once()
async def test_clear_not_called_on_failure(self) -> None:
"""tracker.clear() must NOT be called when the probe fails."""
orch = _make_orchestrator()
provider = "anthropic"
state = _make_active_state(provider, retry_after=None)
tracker_mock = _make_tracker_mock(failure_return=1)
orch._make_tracker = MagicMock(return_value=tracker_mock) # type: ignore[method-assign]
async def fake_do_probe(_p: str) -> bool:
return False
orch._do_probe = fake_do_probe # type: ignore[method-assign]
orch._notify_rate_limit_ceo = AsyncMock()
await orch._probe_one_provider(provider, state)
tracker_mock.clear.assert_not_awaited()
# ---------------------------------------------------------------------------
# Tests: CEO notification threshold (AC8)
# ---------------------------------------------------------------------------
class TestCEONotificationThreshold:
"""CEO notification fires at count==10 exactly once per episode."""
async def test_notification_fires_at_exactly_10_failures(self) -> None:
"""_notify_rate_limit_ceo is called when failure count hits 10."""
orch = _make_orchestrator()
provider = "anthropic"
state = _make_active_state(provider, retry_after=None)
# simulate already at 9 failures; next increment returns 10
tracker_mock = _make_tracker_mock(failure_return=10)
orch._make_tracker = MagicMock(return_value=tracker_mock) # type: ignore[method-assign]
orch._notify_rate_limit_ceo = AsyncMock()
async def fake_do_probe(_p: str) -> bool:
return False
orch._do_probe = fake_do_probe # type: ignore[method-assign]
await orch._probe_one_provider(provider, state)
orch._notify_rate_limit_ceo.assert_awaited_once()
async def test_notification_not_fired_before_threshold(self) -> None:
"""No CEO notification below threshold 10."""
orch = _make_orchestrator()
provider = "anthropic"
state = _make_active_state(provider, retry_after=None)
tracker_mock = _make_tracker_mock(failure_return=9)
orch._make_tracker = MagicMock(return_value=tracker_mock) # type: ignore[method-assign]
orch._notify_rate_limit_ceo = AsyncMock()
async def fake_do_probe(_p: str) -> bool:
return False
orch._do_probe = fake_do_probe # type: ignore[method-assign]
await orch._probe_one_provider(provider, state)
orch._notify_rate_limit_ceo.assert_not_awaited()
async def test_notification_sent_only_once_per_episode(self) -> None:
"""Even if failures keep accumulating, the CEO is notified only once."""
orch = _make_orchestrator()
provider = "anthropic"
state = _make_active_state(provider, retry_after=None)
# Mark this episode as already notified
orch._rate_limit_ceo_notified.add(provider)
tracker_mock = _make_tracker_mock(failure_return=15)
orch._make_tracker = MagicMock(return_value=tracker_mock) # type: ignore[method-assign]
orch._notify_rate_limit_ceo = AsyncMock()
async def fake_do_probe(_p: str) -> bool:
return False
orch._do_probe = fake_do_probe # type: ignore[method-assign]
await orch._probe_one_provider(provider, state)
orch._notify_rate_limit_ceo.assert_not_awaited()
async def test_new_episode_allows_new_notification(self) -> None:
"""After a rate-limit clears (success) a new episode starts fresh."""
orch = _make_orchestrator()
provider = "anthropic"
# Episode 1: had a notification
orch._rate_limit_ceo_notified.add(provider)
success_state = _make_active_state(provider, retry_after=None)
orch.resolve_wait = AsyncMock(return_value=None)
tracker_mock = _make_tracker_mock(failure_return=10)
orch._make_tracker = MagicMock(return_value=tracker_mock) # type: ignore[method-assign]
notify_mock = AsyncMock()
orch._notify_rate_limit_ceo = notify_mock
async def fake_do_probe_success(_p: str) -> bool:
return True
orch._do_probe = fake_do_probe_success # type: ignore[method-assign]
with patch("roboco.events.get_event_bus") as mock_bus_fn:
bus_mock = AsyncMock()
bus_mock.publish = AsyncMock()
mock_bus_fn.return_value = bus_mock
# Success clears the episode flag
await orch._probe_one_provider(provider, success_state)
assert provider not in orch._rate_limit_ceo_notified
# Episode 2: simulate a new failure reaching threshold 10
async def fake_do_probe_fail(_p: str) -> bool:
return False
orch._do_probe = fake_do_probe_fail # type: ignore[method-assign]
failure_state = _make_active_state(provider, retry_after=None)
await orch._probe_one_provider(provider, failure_state)
# Notification SHOULD fire for the new episode
notify_mock.assert_awaited_once()
# ---------------------------------------------------------------------------
# Tests: list_rate_limited_providers
# ---------------------------------------------------------------------------
class TestListRateLimitedProviders:
"""list_rate_limited_providers scans Redis for active rate-limit keys."""
async def test_returns_empty_when_no_keys(self) -> None:
redis_mock = _make_redis_mock()
with patch("redis.asyncio.from_url", return_value=redis_mock):
result = await RateLimitStateTracker.list_rate_limited_providers()
assert result == []
async def test_returns_active_provider(self) -> None:
state = {
"rate_limited": True,
"activated_at": datetime.now(UTC).isoformat(),
"retry_after": 60.0,
"affected_agents": ["be-dev-1"],
"probe_failures": 0,
}
store = {"roboco:rate_limit:anthropic:state": json.dumps(state).encode()}
redis_mock = _make_redis_mock(store)
with patch("redis.asyncio.from_url", return_value=redis_mock):
result = await RateLimitStateTracker.list_rate_limited_providers()
assert len(result) == 1
provider, returned_state = result[0]
assert provider == "anthropic"
assert returned_state["rate_limited"] is True
async def test_ignores_cleared_providers(self) -> None:
state = {
"rate_limited": False,
"activated_at": datetime.now(UTC).isoformat(),
"retry_after": 60.0,
"affected_agents": [],
"probe_failures": 2,
}
store = {"roboco:rate_limit:anthropic:state": json.dumps(state).encode()}
redis_mock = _make_redis_mock(store)
with patch("redis.asyncio.from_url", return_value=redis_mock):
result = await RateLimitStateTracker.list_rate_limited_providers()
assert result == []
# ---------------------------------------------------------------------------
# Tests: GET /api/system/rate-limits endpoint schema (AC9)
# ---------------------------------------------------------------------------
class TestRateLimitsEndpoint:
"""GET /api/system/rate-limits returns correct schema."""
async def test_returns_empty_list_when_no_rate_limits(self) -> None:
app = create_app()
with patch(
"roboco.api.routes.system.RateLimitStateTracker"
".list_rate_limited_providers",
new_callable=AsyncMock,
return_value=[],
):
async with AsyncClient(
transport=ASGITransport(app=app), base_url="http://test"
) as client:
resp = await client.get("/api/system/rate-limits")
assert resp.status_code == 200 # noqa: PLR2004
assert resp.json() == []
async def test_returns_provider_state_when_rate_limited(self) -> None:
app = create_app()
state = {
"rate_limited": True,
"activated_at": "2026-06-11T00:00:00+00:00",
"retry_after": 60.0,
"affected_agents": ["be-dev-1"],
"probe_failures": 3,
}
with patch(
"roboco.api.routes.system.RateLimitStateTracker"
".list_rate_limited_providers",
new_callable=AsyncMock,
return_value=[("anthropic", state)],
):
async with AsyncClient(
transport=ASGITransport(app=app), base_url="http://test"
) as client:
resp = await client.get("/api/system/rate-limits")
assert resp.status_code == 200 # noqa: PLR2004
data = resp.json()
assert len(data) == 1
entry = data[0]
assert entry["provider"] == "anthropic"
assert entry["rate_limited"] is True
assert entry["probe_failures"] == 3 # noqa: PLR2004
assert entry["retry_after"] == 60.0 # noqa: PLR2004
async def test_endpoint_not_404(self) -> None:
"""The endpoint must be registered in app.py — no 404."""
app = create_app()
with patch(
"roboco.api.routes.system.RateLimitStateTracker"
".list_rate_limited_providers",
new_callable=AsyncMock,
return_value=[],
):
async with AsyncClient(
transport=ASGITransport(app=app), base_url="http://test"
) as client:
resp = await client.get("/api/system/rate-limits")
assert resp.status_code != 404 # noqa: PLR2004
@@ -0,0 +1,691 @@
"""
Unit tests for rate-limit retry behaviour across all LLM call sites.
Covers acceptance criteria:
- 5-retry exhaustion raises RateLimitError
- Retry-After header drives the sleep duration
- Partial retries then success returns the correct result
- ConnectError / TimeoutException in OllamaEmbedder does NOT trigger the 429
retry path (the two concerns are composed without double-retrying)
"""
from __future__ import annotations
import json
import sys
import types
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
from uuid import uuid4
import anthropic as anthropic_mod
import httpx
import pytest
import pytest_asyncio # noqa: F401 - registers asyncio mode
# ---------------------------------------------------------------------------
# Ensure piragi stubs are present before the optimal_brain modules are imported
# ---------------------------------------------------------------------------
_PIRAGI_STUB_NAMES = (
"piragi",
"piragi.types",
"piragi.stores",
"piragi.stores.postgres",
"piragi.chunking",
"piragi.semantic_chunking",
)
def _stub_piragi() -> None:
mock = MagicMock()
for name in _PIRAGI_STUB_NAMES:
if name not in sys.modules:
mod = types.ModuleType(name)
mod.__dict__.update(
{
"AsyncRagi": mock,
"Citation": mock,
"Document": mock,
"Chunk": mock,
"PostgresStore": mock,
}
)
sys.modules[name] = mod
_stub_piragi()
# ---------------------------------------------------------------------------
# Module imports (after stubs are injected)
# ---------------------------------------------------------------------------
from roboco.models.extraction import ExtractionContext # noqa: E402
from roboco.models.optimal import IndexType # noqa: E402
from roboco.services.exceptions import ( # noqa: E402
MAX_RATE_LIMIT_RETRIES,
RateLimitError,
parse_retry_after_header,
)
from roboco.services.extraction import ExtractionService # noqa: E402
from roboco.services.optimal_brain.indexes.journals import ( # noqa: E402
JournalsIndexPlugin,
)
from roboco.services.optimal_brain.mentor import MentorService # noqa: E402
from roboco.services.optimal_brain.ollama_embedder import ( # noqa: E402
MAX_RETRIES,
OllamaConnectionError,
OllamaEmbedder,
)
from roboco.services.optimal_brain.validator import ValidatorService # noqa: E402
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
_RETRY_AFTER_FLOAT = 30.0
_RETRY_AFTER_FLOAT_2 = 2.5
_RETRY_AFTER_7 = 7.0
_RETRY_AFTER_9 = 9.0
_RETRY_AFTER_12 = 12.0
_RETRY_AFTER_5 = 5.0
_EMBED_DIM = 4 # zero-vector dimension in test responses
_CALLS_2RL_1_SUCCESS = 3 # 2 rate-limit errors then 1 success
_EMBED_PATH = (
"roboco.services.optimal_brain.ollama_embedder.OllamaEmbedder._create_async_client"
)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_response(
status_code: int,
body: Any = None,
retry_after: str | None = None,
) -> httpx.Response:
"""Build a minimal httpx.Response for use in mocks."""
headers: dict[str, str] = {}
if retry_after is not None:
headers["retry-after"] = retry_after
content = json.dumps(body or {}).encode()
return httpx.Response(
status_code=status_code,
headers=headers,
content=content,
)
def _success_embed_response(n: int = 1) -> httpx.Response:
"""Return a valid Ollama /api/embed response with *n* zero-vectors."""
body = {"embeddings": [[0.0] * _EMBED_DIM] * n}
return _make_response(200, body)
def _429_response(retry_after: str | None = None) -> httpx.Response:
return _make_response(429, {"error": "rate limited"}, retry_after=retry_after)
def _make_async_client_mock(post_return: Any = None) -> AsyncMock:
"""Return an async context manager mock exposing `.post`."""
client = AsyncMock()
client.__aenter__ = AsyncMock(return_value=client)
client.__aexit__ = AsyncMock(return_value=False)
if post_return is not None:
client.post = AsyncMock(return_value=post_return)
return client
def _make_extraction_context() -> ExtractionContext:
return ExtractionContext(
content="This is a test message that is long enough.",
agent_id=uuid4(),
channel_id=uuid4(),
session_id=uuid4(),
group_id=uuid4(),
)
def _make_anthropic_rl_exc(
retry_after: str | None = None,
) -> anthropic_mod.RateLimitError:
"""Build a minimal anthropic.RateLimitError."""
fake_resp = MagicMock()
fake_resp.headers = {} if retry_after is None else {"retry-after": retry_after}
return anthropic_mod.RateLimitError(
message="rate limit",
response=fake_resp,
body={},
)
def _make_journal_plugin() -> JournalsIndexPlugin:
"""Create a minimal JournalsIndexPlugin without initialising piragi."""
plugin = JournalsIndexPlugin.__new__(JournalsIndexPlugin)
plugin._config = MagicMock()
plugin._config.llm_base_url = "http://ollama-test:11434/v1"
plugin._config.llm_model = "glm-5:cloud"
plugin._ragi = MagicMock()
plugin._initialized = True
return plugin
def _make_search_outcome(content: str = "context text") -> Any:
mock_outcome = MagicMock()
mock_outcome.success = True
mock_outcome.results = [
MagicMock(content=content, source="src", score=0.9, index_type=None)
]
return mock_outcome
def _make_sources() -> list[Any]:
return [
MagicMock(content="ctx", source="s", score=0.9, index_type=IndexType.JOURNALS)
]
def _make_standards() -> list[Any]:
s = MagicMock()
s.content = "### PY-001: Use Type Hints\nMust add return type annotations."
return [s]
# ===========================================================================
# 1. parse_retry_after_header
# ===========================================================================
class TestParseRetryAfterHeader:
def test_integer_seconds(self) -> None:
resp = _429_response(retry_after="30")
assert parse_retry_after_header(resp) == _RETRY_AFTER_FLOAT
def test_float_seconds(self) -> None:
resp = _429_response(retry_after="2.5")
assert parse_retry_after_header(resp) == _RETRY_AFTER_FLOAT_2
def test_missing_header_returns_none(self) -> None:
resp = _make_response(429)
assert parse_retry_after_header(resp) is None
def test_non_numeric_returns_none(self) -> None:
resp = _429_response(retry_after="Wed, 21 Oct 2015 07:28:00 GMT")
assert parse_retry_after_header(resp) is None
# ===========================================================================
# 2. RateLimitError class
# ===========================================================================
class TestRateLimitError:
def test_fields(self) -> None:
err = RateLimitError(provider="anthropic", retry_after=_RETRY_AFTER_FLOAT)
assert err.provider == "anthropic"
assert err.retry_after == _RETRY_AFTER_FLOAT
def test_message_includes_provider(self) -> None:
err = RateLimitError(provider="ollama")
assert "ollama" in str(err)
def test_message_includes_retry_after_when_set(self) -> None:
err = RateLimitError(provider="ollama", retry_after=15.0)
assert "15" in str(err)
def test_none_retry_after(self) -> None:
err = RateLimitError(provider="anthropic", retry_after=None)
assert err.retry_after is None
# ===========================================================================
# 3. OllamaEmbedder - async path (aembed_query)
# ===========================================================================
class TestOllamaEmbedderAembed:
"""Tests for the async aembed_query method."""
def _make_embedder(self) -> OllamaEmbedder:
return OllamaEmbedder(base_url="http://ollama-test:11434")
async def test_five_consecutive_429s_raise_rate_limit_error(self) -> None:
"""After MAX_RATE_LIMIT_RETRIES attempts all 429 -> RateLimitError."""
embedder = self._make_embedder()
mock_c = _make_async_client_mock(post_return=_429_response())
with (
patch(_EMBED_PATH, return_value=mock_c),
patch("asyncio.sleep", new_callable=AsyncMock),
pytest.raises(RateLimitError) as exc_info,
):
await embedder.aembed_query("hello")
assert exc_info.value.provider == "ollama"
assert mock_c.post.call_count == MAX_RATE_LIMIT_RETRIES
async def test_retry_after_header_respected_as_sleep_duration(self) -> None:
"""Retry-After: 7 -> asyncio.sleep(7.0) on each inter-attempt gap."""
embedder = self._make_embedder()
sleep_calls: list[float] = []
async def _fake_sleep(secs: float) -> None:
sleep_calls.append(secs)
mock_c = _make_async_client_mock(post_return=_429_response(retry_after="7"))
with (
patch(_EMBED_PATH, return_value=mock_c),
patch("asyncio.sleep", side_effect=_fake_sleep),
pytest.raises(RateLimitError),
):
await embedder.aembed_query("hello")
assert all(s == _RETRY_AFTER_7 for s in sleep_calls)
assert len(sleep_calls) == MAX_RATE_LIMIT_RETRIES - 1
async def test_partial_retries_then_success(self) -> None:
"""Two 429s then a 200 -> returns the embedding list."""
embedder = self._make_embedder()
mock_c = AsyncMock()
mock_c.__aenter__ = AsyncMock(return_value=mock_c)
mock_c.__aexit__ = AsyncMock(return_value=False)
mock_c.post = AsyncMock(
side_effect=[
_429_response(),
_429_response(),
_success_embed_response(),
]
)
with (
patch(_EMBED_PATH, return_value=mock_c),
patch("asyncio.sleep", new_callable=AsyncMock),
):
result = await embedder.aembed_query("hello")
assert isinstance(result, list)
assert len(result) == _EMBED_DIM
async def test_connect_error_does_not_trigger_429_retry_path(
self,
) -> None:
"""ConnectError -> OllamaConnectionError after MAX_RETRIES=3, not 5."""
embedder = self._make_embedder()
mock_c = AsyncMock()
mock_c.__aenter__ = AsyncMock(return_value=mock_c)
mock_c.__aexit__ = AsyncMock(return_value=False)
mock_c.post = AsyncMock(side_effect=httpx.ConnectError("refused"))
with (
patch(_EMBED_PATH, return_value=mock_c),
patch("asyncio.sleep", new_callable=AsyncMock),
pytest.raises(OllamaConnectionError),
):
await embedder.aembed_query("hello")
assert mock_c.post.call_count == MAX_RETRIES
async def test_timeout_does_not_trigger_429_retry_path(self) -> None:
"""TimeoutException -> OllamaConnectionError after MAX_RETRIES=3."""
embedder = self._make_embedder()
mock_c = AsyncMock()
mock_c.__aenter__ = AsyncMock(return_value=mock_c)
mock_c.__aexit__ = AsyncMock(return_value=False)
mock_c.post = AsyncMock(side_effect=httpx.TimeoutException("timed out"))
with (
patch(_EMBED_PATH, return_value=mock_c),
patch("asyncio.sleep", new_callable=AsyncMock),
pytest.raises(OllamaConnectionError),
):
await embedder.aembed_query("hello")
assert mock_c.post.call_count == MAX_RETRIES
# ===========================================================================
# 4. OllamaEmbedder - sync path (embed_query)
# ===========================================================================
class TestOllamaEmbedderSync:
"""Tests for the synchronous embed_query method."""
def _make_embedder(self) -> OllamaEmbedder:
return OllamaEmbedder(base_url="http://ollama-test:11434")
def test_five_consecutive_429s_raise_rate_limit_error(self) -> None:
embedder = self._make_embedder()
mock_client = MagicMock()
mock_client.post.return_value = _429_response()
with (
patch.object(embedder, "_get_sync_client", return_value=mock_client),
patch("time.sleep"),
pytest.raises(RateLimitError) as exc_info,
):
embedder.embed_query("hello")
assert exc_info.value.provider == "ollama"
assert mock_client.post.call_count == MAX_RATE_LIMIT_RETRIES
def test_retry_after_header_respected_as_sleep_duration(self) -> None:
embedder = self._make_embedder()
sleep_calls: list[float] = []
mock_client = MagicMock()
mock_client.post.return_value = _429_response(retry_after="9")
with (
patch.object(embedder, "_get_sync_client", return_value=mock_client),
patch("time.sleep", side_effect=sleep_calls.append),
pytest.raises(RateLimitError),
):
embedder.embed_query("hello")
assert all(s == _RETRY_AFTER_9 for s in sleep_calls)
assert len(sleep_calls) == MAX_RATE_LIMIT_RETRIES - 1
def test_partial_retries_then_success(self) -> None:
embedder = self._make_embedder()
mock_client = MagicMock()
mock_client.post.side_effect = [_429_response(), _success_embed_response()]
with (
patch.object(embedder, "_get_sync_client", return_value=mock_client),
patch("time.sleep"),
):
result = embedder.embed_query("hello")
assert isinstance(result, list)
def test_connect_error_does_not_trigger_rate_limit_retry(self) -> None:
embedder = self._make_embedder()
mock_client = MagicMock()
mock_client.post.side_effect = httpx.ConnectError("refused")
with (
patch.object(embedder, "_get_sync_client", return_value=mock_client),
patch("time.sleep"),
pytest.raises(OllamaConnectionError),
):
embedder.embed_query("hello")
assert mock_client.post.call_count == MAX_RETRIES
# ===========================================================================
# 5. OllamaEmbedder - _embed_batch_sync
# ===========================================================================
class TestOllamaEmbedBatchSync:
def _make_embedder(self) -> OllamaEmbedder:
return OllamaEmbedder(base_url="http://ollama-test:11434")
def test_five_429s_raise_rate_limit_error(self) -> None:
embedder = self._make_embedder()
mock_client = MagicMock()
mock_client.post.return_value = _429_response()
with patch("time.sleep"), pytest.raises(RateLimitError):
embedder._embed_batch_sync(mock_client, ["a", "b"], batch_index=0)
assert mock_client.post.call_count == MAX_RATE_LIMIT_RETRIES
def test_connect_error_raises_connection_error_not_rate_limit(self) -> None:
embedder = self._make_embedder()
mock_client = MagicMock()
mock_client.post.side_effect = httpx.ConnectError("refused")
with patch("time.sleep"), pytest.raises(OllamaConnectionError):
embedder._embed_batch_sync(mock_client, ["a"], batch_index=0)
assert mock_client.post.call_count == MAX_RETRIES
# ===========================================================================
# 6. extraction.py - Anthropic rate-limit retry
# ===========================================================================
class TestExtractionAnthropicRetry:
"""Tests for ExtractionService.extract_with_llm Anthropic retry logic.
AsyncAnthropic is imported *inside* extract_with_llm, so we patch at the
anthropic module level. The client is NOT used as a context manager there.
"""
async def test_five_rate_limit_errors_raises_rate_limit_error(self) -> None:
"""RateLimitError raised 5 times -> our RateLimitError propagated."""
svc = ExtractionService()
api_exc = _make_anthropic_rl_exc(retry_after="5")
mock_client = MagicMock()
mock_client.messages.create = AsyncMock(side_effect=api_exc)
with (
patch("anthropic.AsyncAnthropic") as mock_cls,
patch("roboco.config.settings") as mock_settings,
patch("asyncio.sleep", new_callable=AsyncMock),
pytest.raises(RateLimitError) as exc_info,
):
mock_cls.return_value = mock_client
mock_settings.anthropic_api_key = "test-key"
await svc.extract_with_llm(_make_extraction_context())
assert exc_info.value.provider == "anthropic"
assert mock_client.messages.create.call_count == MAX_RATE_LIMIT_RETRIES
async def test_retry_after_header_drives_sleep_duration(self) -> None:
"""Retry-After: 12 -> asyncio.sleep(12) called on each gap."""
svc = ExtractionService()
sleep_calls: list[float] = []
api_exc = _make_anthropic_rl_exc(retry_after="12")
mock_client = MagicMock()
mock_client.messages.create = AsyncMock(side_effect=api_exc)
async def _fake_sleep(secs: float) -> None:
sleep_calls.append(secs)
with (
patch("anthropic.AsyncAnthropic") as mock_cls,
patch("roboco.config.settings") as mock_settings,
patch("asyncio.sleep", side_effect=_fake_sleep),
pytest.raises(RateLimitError),
):
mock_cls.return_value = mock_client
mock_settings.anthropic_api_key = "test-key"
await svc.extract_with_llm(_make_extraction_context())
assert all(s == _RETRY_AFTER_12 for s in sleep_calls)
assert len(sleep_calls) == MAX_RATE_LIMIT_RETRIES - 1
async def test_partial_rate_limit_then_success_returns_result(self) -> None:
"""Two RateLimitErrors then success -> result returned, no raise."""
svc = ExtractionService()
api_exc = _make_anthropic_rl_exc()
text_block = MagicMock()
text_block.text = "[N,]{type,content,confidence}:\nreasoning,Hello world,0.9"
success_response = MagicMock()
success_response.content = [text_block]
mock_client = MagicMock()
mock_client.messages.create = AsyncMock(
side_effect=[api_exc, api_exc, success_response]
)
raised = False
with (
patch("anthropic.AsyncAnthropic") as mock_cls,
patch("roboco.config.settings") as mock_settings,
patch("asyncio.sleep", new_callable=AsyncMock),
):
mock_cls.return_value = mock_client
mock_settings.anthropic_api_key = "test-key"
try:
result = await svc.extract_with_llm(_make_extraction_context())
assert result is not None
except RateLimitError:
raised = True
assert not raised, "RateLimitError raised even though 3rd attempt succeeded"
assert mock_client.messages.create.call_count == _CALLS_2RL_1_SUCCESS
# ===========================================================================
# 7. indexes/base.py - BaseIndexPlugin.ask() LLM 429 retry
# ===========================================================================
class TestIndexAsk429Retry:
"""Tests for the LLM call in BaseIndexPlugin.ask()."""
async def test_ask_raises_rate_limit_after_five_429s(self) -> None:
plugin = _make_journal_plugin()
mock_c = _make_async_client_mock(post_return=_429_response())
with (
patch.object(
plugin, "search", AsyncMock(return_value=_make_search_outcome())
),
patch("httpx.AsyncClient", return_value=mock_c),
patch("asyncio.sleep", new_callable=AsyncMock),
pytest.raises(RateLimitError) as exc_info,
):
await plugin.ask("what is X")
assert exc_info.value.provider == "ollama"
assert mock_c.post.call_count == MAX_RATE_LIMIT_RETRIES
async def test_ask_partial_retry_then_success(self) -> None:
plugin = _make_journal_plugin()
success_body = {"choices": [{"message": {"content": "Here is the answer."}}]}
success_resp = _make_response(200, success_body)
mock_c = AsyncMock()
mock_c.__aenter__ = AsyncMock(return_value=mock_c)
mock_c.__aexit__ = AsyncMock(return_value=False)
mock_c.post = AsyncMock(side_effect=[_429_response(), success_resp])
with (
patch.object(
plugin, "search", AsyncMock(return_value=_make_search_outcome())
),
patch("httpx.AsyncClient", return_value=mock_c),
patch("asyncio.sleep", new_callable=AsyncMock),
):
answer, results = await plugin.ask("what is X")
assert answer == "Here is the answer."
assert results is not None
# ===========================================================================
# 8. mentor.py - MentorService._synthesize_answer() 429 retry
# ===========================================================================
class TestMentorSynthesizeAnswer429:
async def test_raises_rate_limit_after_five_429s(self) -> None:
mentor = MentorService()
mentor._optimal_service = MagicMock()
mock_c = _make_async_client_mock(post_return=_429_response())
with (
patch("httpx.AsyncClient", return_value=mock_c),
patch("asyncio.sleep", new_callable=AsyncMock),
pytest.raises(RateLimitError) as exc_info,
):
await mentor._synthesize_answer(
question="test",
sources=_make_sources(),
conversation_context="",
agent_profile=None,
journal_context=[],
)
assert exc_info.value.provider == "ollama"
assert mock_c.post.call_count == MAX_RATE_LIMIT_RETRIES
async def test_partial_retry_then_success(self) -> None:
mentor = MentorService()
mentor._optimal_service = MagicMock()
success_body = {"choices": [{"message": {"content": "Great answer."}}]}
mock_c = AsyncMock()
mock_c.__aenter__ = AsyncMock(return_value=mock_c)
mock_c.__aexit__ = AsyncMock(return_value=False)
mock_c.post = AsyncMock(
side_effect=[_429_response(), _make_response(200, success_body)]
)
with (
patch("httpx.AsyncClient", return_value=mock_c),
patch("asyncio.sleep", new_callable=AsyncMock),
):
answer = await mentor._synthesize_answer(
question="test",
sources=_make_sources(),
conversation_context="",
agent_profile=None,
journal_context=[],
)
assert answer == "Great answer."
# ===========================================================================
# 9. validator.py - ValidatorService._validate_with_llm() 429 retry
# ===========================================================================
class TestValidatorLLMRetry:
async def test_raises_rate_limit_after_five_429s(self) -> None:
validator = ValidatorService()
validator._optimal_service = MagicMock()
validator._llm_available = True
mock_c = _make_async_client_mock(post_return=_429_response())
with (
patch("httpx.AsyncClient", return_value=mock_c),
patch("asyncio.sleep", new_callable=AsyncMock),
pytest.raises(RateLimitError) as exc_info,
):
await validator._validate_with_llm(
action_type="create_endpoint",
context="def foo(): pass",
standards=_make_standards(),
)
assert exc_info.value.provider == "ollama"
assert mock_c.post.call_count == MAX_RATE_LIMIT_RETRIES
async def test_partial_retry_then_success(self) -> None:
validator = ValidatorService()
validator._optimal_service = MagicMock()
validator._llm_available = True
success_content = '{"violations": [], "summary": "ok"}'
success_body = {"choices": [{"message": {"content": success_content}}]}
mock_c = AsyncMock()
mock_c.__aenter__ = AsyncMock(return_value=mock_c)
mock_c.__aexit__ = AsyncMock(return_value=False)
mock_c.post = AsyncMock(
side_effect=[_429_response(), _make_response(200, success_body)]
)
with (
patch("httpx.AsyncClient", return_value=mock_c),
patch("asyncio.sleep", new_callable=AsyncMock),
):
violations, warnings = await validator._validate_with_llm(
action_type="create_endpoint",
context="def foo(): pass",
standards=_make_standards(),
)
assert violations == []
assert warnings == []
@@ -0,0 +1,244 @@
"""Unit tests for RateLimitStateTracker.
These tests use mock Redis clients (no real Redis server required) to
verify the state-management logic. The cross-reconnection persistence
test constructs *two* RateLimitStateTracker instances that share the
same mock Redis store, proving that state written by one instance is
visible to a fresh instance.
"""
from __future__ import annotations
import json
from typing import Any
from unittest.mock import AsyncMock, MagicMock
import pytest
from roboco.services.gateway.rate_limit_tracker import RateLimitStateTracker
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_redis_mock(initial_store: dict[str, Any] | None = None) -> AsyncMock:
"""Build an async Redis mock backed by a plain dict.
The mock supports ``get``, ``set``, and ``delete`` with the same
semantics as the real redis.asyncio.Redis client.
"""
# Use the dict AS-IS (no copy) so that two mocks sharing the same
# dict object see each other's writes and deletes — this is what the
# cross-reconnection persistence tests rely on.
store: dict[str, Any] = initial_store if initial_store is not None else {}
async def _get(key: str) -> bytes | None:
val = store.get(key)
if val is None:
return None
if isinstance(val, bytes):
return val
return str(val).encode()
async def _set(key: str, value: Any) -> None:
store[key] = value
async def _delete(key: str) -> int:
return 1 if store.pop(key, None) is not None else 0
mock = AsyncMock()
mock.get = AsyncMock(side_effect=_get)
mock.set = AsyncMock(side_effect=_set)
mock.delete = AsyncMock(side_effect=_delete)
# Stash the backing store so tests can inspect raw state
mock._store = store
return mock
def _make_tracker(
provider: str = "anthropic",
redis_mock: AsyncMock | None = None,
) -> RateLimitStateTracker:
"""Build a tracker with an injected mock Redis client."""
tracker = RateLimitStateTracker(provider=provider, redis_url="redis://unused")
if redis_mock is not None:
tracker._redis = redis_mock # type: ignore[assignment]
return tracker
# ---------------------------------------------------------------------------
# Tests: basic operations
# ---------------------------------------------------------------------------
class TestActivateAndRead:
async def test_is_rate_limited_false_by_default(self) -> None:
mock = _make_redis_mock()
tracker = _make_tracker(redis_mock=mock)
assert await tracker.is_rate_limited() is False
async def test_get_state_empty_by_default(self) -> None:
mock = _make_redis_mock()
tracker = _make_tracker(redis_mock=mock)
assert await tracker.get_state() == {}
async def test_activate_sets_rate_limited_true(self) -> None:
mock = _make_redis_mock()
tracker = _make_tracker(redis_mock=mock)
await tracker.activate()
assert await tracker.is_rate_limited() is True
async def test_activate_stores_retry_after(self) -> None:
mock = _make_redis_mock()
tracker = _make_tracker(redis_mock=mock)
await tracker.activate(retry_after=30.0)
state = await tracker.get_state()
assert state["retry_after"] == 30.0
async def test_activate_stores_affected_agents(self) -> None:
mock = _make_redis_mock()
tracker = _make_tracker(redis_mock=mock)
await tracker.activate(affected_agents=["be-dev-1", "be-dev-2"])
state = await tracker.get_state()
assert state["affected_agents"] == ["be-dev-1", "be-dev-2"]
async def test_activate_initialises_probe_failures_zero(self) -> None:
mock = _make_redis_mock()
tracker = _make_tracker(redis_mock=mock)
await tracker.activate()
state = await tracker.get_state()
assert state["probe_failures"] == 0
async def test_clear_removes_state(self) -> None:
mock = _make_redis_mock()
tracker = _make_tracker(redis_mock=mock)
await tracker.activate()
await tracker.clear()
assert await tracker.is_rate_limited() is False
assert await tracker.get_state() == {}
class TestProbeFailures:
async def test_increment_starts_from_zero(self) -> None:
mock = _make_redis_mock()
tracker = _make_tracker(redis_mock=mock)
await tracker.activate()
count = await tracker.increment_probe_failures()
assert count == 1
async def test_increment_accumulates(self) -> None:
mock = _make_redis_mock()
tracker = _make_tracker(redis_mock=mock)
await tracker.activate()
await tracker.increment_probe_failures()
await tracker.increment_probe_failures()
count = await tracker.increment_probe_failures()
assert count == 3
async def test_reset_sets_zero(self) -> None:
mock = _make_redis_mock()
tracker = _make_tracker(redis_mock=mock)
await tracker.activate()
await tracker.increment_probe_failures()
await tracker.increment_probe_failures()
await tracker.reset_probe_failures()
state = await tracker.get_state()
assert state["probe_failures"] == 0
# ---------------------------------------------------------------------------
# Tests: cross-reconnection persistence
# ---------------------------------------------------------------------------
#
# AC2: "State persists across client reconnection: a test writes state via
# activate(), creates a new RateLimitStateTracker instance pointing at the
# same Redis URL, calls is_rate_limited() and get_state() and gets back the
# same values — proving state survives a process restart."
#
# We simulate this by sharing the same backing dict between two mock Redis
# clients — one injected into the first tracker and one injected into the
# second. Both clients read from and write to the same dict, so the second
# tracker "sees" everything the first wrote.
# ---------------------------------------------------------------------------
class TestStatePersistsAcrossReconnection:
async def test_is_rate_limited_survives_reconnection(self) -> None:
shared_store: dict[str, Any] = {}
# First "connection": write rate-limit state
mock_a = _make_redis_mock(initial_store=shared_store)
tracker_a = _make_tracker(provider="anthropic", redis_mock=mock_a)
await tracker_a.activate(retry_after=60.0, affected_agents=["be-dev-1"])
# The mock writes into shared_store directly (our _set stores raw).
# We need to seed the second mock from the same backing store.
# Because mock_a._store IS shared_store (same dict object), we only
# need to give mock_b access to the same dict.
mock_b = _make_redis_mock(initial_store=mock_a._store)
tracker_b = RateLimitStateTracker(
provider="anthropic", redis_url="redis://unused"
)
tracker_b._redis = mock_b # type: ignore[assignment]
assert await tracker_b.is_rate_limited() is True
async def test_get_state_survives_reconnection(self) -> None:
shared_store: dict[str, Any] = {}
mock_a = _make_redis_mock(initial_store=shared_store)
tracker_a = _make_tracker(provider="anthropic", redis_mock=mock_a)
await tracker_a.activate(retry_after=45.0, affected_agents=["be-dev-2"])
mock_b = _make_redis_mock(initial_store=mock_a._store)
tracker_b = RateLimitStateTracker(
provider="anthropic", redis_url="redis://unused"
)
tracker_b._redis = mock_b # type: ignore[assignment]
state = await tracker_b.get_state()
assert state["rate_limited"] is True
assert state["retry_after"] == 45.0
assert state["affected_agents"] == ["be-dev-2"]
async def test_clear_via_first_instance_visible_to_second(self) -> None:
shared_store: dict[str, Any] = {}
mock_a = _make_redis_mock(initial_store=shared_store)
tracker_a = _make_tracker(provider="anthropic", redis_mock=mock_a)
await tracker_a.activate()
# Second instance points at the same store
mock_b = _make_redis_mock(initial_store=mock_a._store)
tracker_b = RateLimitStateTracker(
provider="anthropic", redis_url="redis://unused"
)
tracker_b._redis = mock_b # type: ignore[assignment]
# Write clear via tracker_a
await tracker_a.clear()
# tracker_b observes the cleared state
assert await tracker_b.is_rate_limited() is False
assert await tracker_b.get_state() == {}
# ---------------------------------------------------------------------------
# Tests: different providers are isolated
# ---------------------------------------------------------------------------
class TestProviderIsolation:
async def test_activating_one_provider_does_not_affect_another(self) -> None:
store: dict[str, Any] = {}
mock_a = _make_redis_mock(initial_store=store)
mock_b = _make_redis_mock(initial_store=store)
tracker_anthropic = _make_tracker(provider="anthropic", redis_mock=mock_a)
tracker_ollama = _make_tracker(provider="ollama_cloud", redis_mock=mock_b)
await tracker_anthropic.activate()
assert await tracker_anthropic.is_rate_limited() is True
assert await tracker_ollama.is_rate_limited() is False