mirror of
https://github.com/rennf93/roboco.git
synced 2026-08-03 07:23:24 +02:00
[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:
co-authored by
Backend Developer 1
Frontend Developer 1
Renn F
parent
cc4ccb7ea3
commit
98e618c243
@@ -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"
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user