From 761e95b6593497b610adfa27687de19c49b9789b Mon Sep 17 00:00:00 2001 From: chengyongru Date: Fri, 21 Aug 2026 16:30:07 +0800 Subject: [PATCH] refactor(providers): simplify retry fallback routing --- nanobot/providers/base.py | 33 +- nanobot/providers/fallback_provider.py | 273 +++++------ tests/agent/test_runner_fallback.py | 630 ++++++++----------------- 3 files changed, 319 insertions(+), 617 deletions(-) diff --git a/nanobot/providers/base.py b/nanobot/providers/base.py index f603119b9..f29cd2c82 100644 --- a/nanobot/providers/base.py +++ b/nanobot/providers/base.py @@ -913,10 +913,10 @@ class LLMProvider(ABC): kw["provider_context"] = provider_context if on_stream_recover and getattr(self, "supports_stream_recover_callback", False): kw["on_stream_recover"] = _recover_stream - return await self._run_with_retry( - self._safe_chat_stream, + return await self._run_chat_with_retry( kw, messages, + stream=True, retry_mode=retry_mode, on_retry_wait=on_retry_wait, on_retry_exhausted=on_retry_exhausted or on_retry_wait, @@ -961,15 +961,40 @@ class LLMProvider(ABC): ) if provider_context is not None: kw["provider_context"] = provider_context - return await self._run_with_retry( - self._safe_chat, + return await self._run_chat_with_retry( kw, messages, + stream=False, retry_mode=retry_mode, on_retry_wait=on_retry_wait, on_retry_exhausted=on_retry_exhausted or on_retry_wait, ) + async def _run_chat_with_retry( + self, + kw: dict[str, Any], + original_messages: list[dict[str, Any]], + *, + stream: bool, + retry_mode: str, + on_retry_wait: RetryEventCallback | None, + on_retry_exhausted: RetryEventCallback | None, + should_retry_guard: Callable[[], bool] | None = None, + on_stream_recover: Callable[[], Awaitable[None]] | None = None, + ) -> LLMResponse: + """Run one chat entry point through this provider's retry policy.""" + call = self._safe_chat_stream if stream else self._safe_chat + return await self._run_with_retry( + call, + kw, + original_messages, + retry_mode=retry_mode, + on_retry_wait=on_retry_wait, + on_retry_exhausted=on_retry_exhausted, + should_retry_guard=should_retry_guard, + on_stream_recover=on_stream_recover, + ) + @classmethod def _extract_retry_after(cls, content: str | None) -> float | None: text = (content or "").lower() diff --git a/nanobot/providers/fallback_provider.py b/nanobot/providers/fallback_provider.py index d59217843..7881249aa 100644 --- a/nanobot/providers/fallback_provider.py +++ b/nanobot/providers/fallback_provider.py @@ -92,7 +92,6 @@ _FALLBACK_ERROR_TOKENS = ( FallbackModelObserver = Callable[[str], Awaitable[None]] -_UNSET = object() class FallbackProvider(LLMProvider): @@ -196,48 +195,78 @@ class FallbackProvider(LLMProvider): lambda p, kw: p.chat(**kw), kwargs, has_streamed=None ) - async def chat_with_retry( + async def _run_chat_with_retry( self, - messages: list[dict[str, Any]], - tools: list[dict[str, Any]] | None = None, - model: str | None = None, - max_tokens: object = _UNSET, - temperature: object = _UNSET, - reasoning_effort: object = _UNSET, - tool_choice: str | dict[str, Any] | None = None, - retry_mode: str = "standard", - on_retry_wait: RetryEventCallback | None = None, - provider_context: ProviderCallContext | None = None, - on_retry_exhausted: RetryEventCallback | None = None, + kw: dict[str, Any], + original_messages: list[dict[str, Any]], + *, + stream: bool, + retry_mode: str, + on_retry_wait: RetryEventCallback | None, + on_retry_exhausted: RetryEventCallback | None, + should_retry_guard: Callable[[], bool] | None = None, + on_stream_recover: Callable[[], Awaitable[None]] | None = None, ) -> LLMResponse: - """Exhaust each provider's retries before moving to the next fallback.""" - call_kwargs: dict[str, Any] = { - "messages": messages, - "tools": tools, - "model": model, - "tool_choice": tool_choice, - "retry_mode": retry_mode, - "on_retry_wait": on_retry_wait, - "on_retry_exhausted": on_retry_exhausted, - } - if max_tokens is not _UNSET: - call_kwargs["max_tokens"] = max_tokens - if temperature is not _UNSET: - call_kwargs["temperature"] = temperature - if reasoning_effort is not _UNSET: - call_kwargs["reasoning_effort"] = reasoning_effort - if provider_context is not None: + """Retry each provider before advancing through the fallback chain.""" + call_kwargs = dict(kw) + provider_context = call_kwargs.get("provider_context") + if isinstance(provider_context, ProviderCallContext): call_kwargs["provider_context"] = self._primary_call_context( provider_context, - model, + call_kwargs.get("model"), ) if not self._has_fallbacks: + call_kwargs.update({ + "retry_mode": retry_mode, + "on_retry_wait": on_retry_wait, + "on_retry_exhausted": on_retry_exhausted, + }) + if stream: + return await self._primary.chat_stream_with_retry(**call_kwargs) return await self._primary.chat_with_retry(**call_kwargs) - return await self._route_with_retry_fallback( - lambda p, kw: p.chat_with_retry(**kw), + + has_streamed: list[bool] | None = None + recover_stream = on_stream_recover + if stream: + streamed = [False] + has_streamed = streamed + original_delta = call_kwargs.get("on_content_delta") + + async def _tracking_delta(text: str) -> None: + if text: + streamed[0] = True + if original_delta: + await original_delta(text) + + async def _recover_stream() -> None: + streamed[0] = False + if on_stream_recover: + await on_stream_recover() + + if original_delta is not None: + call_kwargs["on_content_delta"] = _tracking_delta + if on_stream_recover is not None: + call_kwargs["on_stream_recover"] = _recover_stream + recover_stream = _recover_stream + + async def _call_provider( + provider: LLMProvider, + provider_kwargs: dict[str, Any], + ) -> LLMResponse: + if stream: + return await provider.chat_stream_with_retry(**provider_kwargs) + return await provider.chat_with_retry(**provider_kwargs) + + return await self._retry_with_fallback( + _call_provider, call_kwargs, - has_streamed=None, - on_retry_exhausted=on_retry_exhausted or on_retry_wait, + original_messages, + retry_mode=retry_mode, + on_retry_wait=on_retry_wait, + on_retry_exhausted=on_retry_exhausted, + has_streamed=has_streamed, + on_stream_recover=recover_stream, + persistent_retry_guard=should_retry_guard, ) async def chat_with_context( @@ -281,117 +310,62 @@ class FallbackProvider(LLMProvider): on_stream_recover=on_stream_recover, ) - async def chat_stream_with_retry( - self, - messages: list[dict[str, Any]], - tools: list[dict[str, Any]] | None = None, - model: str | None = None, - max_tokens: object = _UNSET, - temperature: object = _UNSET, - reasoning_effort: object = _UNSET, - tool_choice: str | dict[str, Any] | None = None, - on_content_delta: Callable[[str], Awaitable[None]] | None = None, - on_thinking_delta: Callable[[str], Awaitable[None]] | None = None, - on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None, - on_stream_recover: Callable[[], Awaitable[None]] | None = None, - retry_mode: str = "standard", - on_retry_wait: RetryEventCallback | None = None, - provider_context: ProviderCallContext | None = None, - on_retry_exhausted: RetryEventCallback | None = None, - ) -> LLMResponse: - """Exhaust streaming retries on one provider before failing over.""" - call_kwargs: dict[str, Any] = { - "messages": messages, - "tools": tools, - "model": model, - "tool_choice": tool_choice, - "on_content_delta": on_content_delta, - "on_thinking_delta": on_thinking_delta, - "on_tool_call_delta": on_tool_call_delta, - "retry_mode": retry_mode, - "on_retry_wait": on_retry_wait, - "on_retry_exhausted": on_retry_exhausted, - } - if max_tokens is not _UNSET: - call_kwargs["max_tokens"] = max_tokens - if temperature is not _UNSET: - call_kwargs["temperature"] = temperature - if reasoning_effort is not _UNSET: - call_kwargs["reasoning_effort"] = reasoning_effort - if provider_context is not None: - call_kwargs["provider_context"] = self._primary_call_context( - provider_context, - model, - ) - if not self._has_fallbacks: - if on_stream_recover is not None: - call_kwargs["on_stream_recover"] = on_stream_recover - return await self._primary.chat_stream_with_retry(**call_kwargs) - - has_streamed: list[bool] = [False] - has_unrecovered_stream: list[bool] = [False] - original_delta = call_kwargs.get("on_content_delta") - - async def _tracking_delta(text: str) -> None: - if text: - has_streamed[0] = True - has_unrecovered_stream[0] = True - if original_delta: - await original_delta(text) - - async def _recover_stream() -> None: - has_streamed[0] = False - has_unrecovered_stream[0] = False - if on_stream_recover: - await on_stream_recover() - - if original_delta is not None: - call_kwargs["on_content_delta"] = _tracking_delta - if on_stream_recover is not None: - call_kwargs["on_stream_recover"] = _recover_stream - return await self._route_with_retry_fallback( - lambda p, kw: p.chat_stream_with_retry(**kw), - call_kwargs, - has_streamed=has_streamed, - on_stream_recover=_recover_stream if on_stream_recover is not None else None, - persistent_retry_guard=lambda: not has_unrecovered_stream[0], - on_retry_exhausted=on_retry_exhausted or on_retry_wait, - ) - - async def _route_with_retry_fallback( + async def _retry_with_fallback( self, call: Callable[[LLMProvider, dict[str, Any]], Awaitable[LLMResponse]], kwargs: dict[str, Any], + original_messages: list[dict[str, Any]], + *, + retry_mode: str, + on_retry_wait: RetryEventCallback | None, + on_retry_exhausted: RetryEventCallback | None, has_streamed: list[bool] | None, - on_stream_recover: Callable[[], Awaitable[None]] | None = None, - persistent_retry_guard: Callable[[], bool] | None = None, - on_retry_exhausted: RetryEventCallback | None = None, + on_stream_recover: Callable[[], Awaitable[None]] | None, + persistent_retry_guard: Callable[[], bool] | None, ) -> LLMResponse: - """Apply finite retries per provider and persistence to the whole chain.""" - on_retry_wait: RetryEventCallback | None = kwargs.get("on_retry_wait") - if kwargs.get("retry_mode", "standard") != "persistent": - return await self._try_with_retry_fallback( - call, - kwargs, - has_streamed=has_streamed, - on_stream_recover=on_stream_recover, - on_retry_exhausted=on_retry_exhausted, - ) + """Retry each candidate, deferring terminal events until the chain fails.""" async def _call_chain(**chain_kwargs: Any) -> LLMResponse: - chain_kwargs["retry_mode"] = "standard" - return await self._try_with_retry_fallback( - call, + last_exhausted_message: str | None = None + + async def _capture_exhaustion(message: str) -> None: + nonlocal last_exhausted_message + last_exhausted_message = message + + async def _call_candidate( + provider: LLMProvider, + candidate_kwargs: dict[str, Any], + ) -> LLMResponse: + nonlocal last_exhausted_message + last_exhausted_message = None + return await call(provider, { + **candidate_kwargs, + "retry_mode": "standard", + "on_retry_wait": on_retry_wait, + "on_retry_exhausted": _capture_exhaustion, + }) + + response = await self._try_with_fallback( + _call_candidate, chain_kwargs, has_streamed=has_streamed, on_stream_recover=on_stream_recover, - on_retry_exhausted=None, ) + if ( + retry_mode != "persistent" + and response.finish_reason == "error" + and last_exhausted_message + and on_retry_exhausted + ): + await on_retry_exhausted(last_exhausted_message) + return response + if retry_mode != "persistent": + return await _call_chain(**kwargs) return await self._run_with_retry( _call_chain, dict(kwargs), - kwargs["messages"], + original_messages, retry_mode="persistent", on_retry_wait=on_retry_wait, on_retry_exhausted=on_retry_exhausted, @@ -399,43 +373,6 @@ class FallbackProvider(LLMProvider): on_stream_recover=on_stream_recover, ) - async def _try_with_retry_fallback( - self, - call: Callable[[LLMProvider, dict[str, Any]], Awaitable[LLMResponse]], - kwargs: dict[str, Any], - has_streamed: list[bool] | None, - on_stream_recover: Callable[[], Awaitable[None]] | None, - on_retry_exhausted: RetryEventCallback | None, - ) -> LLMResponse: - """Defer a provider's terminal retry event until the chain fails.""" - last_exhausted_message: str | None = None - - async def _capture_exhaustion(message: str) -> None: - nonlocal last_exhausted_message - last_exhausted_message = message - - async def _call_with_deferred_exhaustion( - provider: LLMProvider, - call_kwargs: dict[str, Any], - ) -> LLMResponse: - nonlocal last_exhausted_message - last_exhausted_message = None - candidate_kwargs = { - **call_kwargs, - "on_retry_exhausted": _capture_exhaustion, - } - return await call(provider, candidate_kwargs) - - response = await self._try_with_fallback( - _call_with_deferred_exhaustion, - kwargs, - has_streamed=has_streamed, - on_stream_recover=on_stream_recover, - ) - if response.finish_reason == "error" and last_exhausted_message and on_retry_exhausted: - await on_retry_exhausted(last_exhausted_message) - return response - async def chat_stream_with_context( self, *, diff --git a/tests/agent/test_runner_fallback.py b/tests/agent/test_runner_fallback.py index 667349a41..38eb824b0 100644 --- a/tests/agent/test_runner_fallback.py +++ b/tests/agent/test_runner_fallback.py @@ -45,6 +45,10 @@ def _error_response(content: str = "api error") -> LLMResponse: return _make_response(content, finish_reason="error", error_kind="server_error") +def _retryable_error(content: str = "") -> LLMResponse: + return _make_response(content, finish_reason="error", error_status_code=503) + + def _fallback( model: str, provider: str = "custom", @@ -67,29 +71,40 @@ def _fallback( class _FakeProvider(LLMProvider): """Fake provider for testing.""" - def __init__(self, name: str = "fake", response: LLMResponse | None = None): + def __init__( + self, + name: str = "fake", + response: LLMResponse | None = None, + *, + responses: list[LLMResponse] | None = None, + ): super().__init__() self.name = name self._response = response or _make_response() + self._responses = iter(responses) if responses is not None else None self.chat_calls: list[dict[str, Any]] = [] self.chat_stream_calls: list[dict[str, Any]] = [] self.context_calls: list[ProviderCallContext | None] = [] self.resumable = False self.compact = False + def _next_response(self) -> LLMResponse: + return next(self._responses) if self._responses is not None else self._response + def get_default_model(self) -> str: return f"{self.name}/model" async def chat(self, **kwargs: Any) -> LLMResponse: self.chat_calls.append(dict(kwargs)) - return self._response + return self._next_response() async def chat_stream(self, **kwargs: Any) -> LLMResponse: self.chat_stream_calls.append(dict(kwargs)) + response = self._next_response() on_delta = kwargs.get("on_content_delta") - if on_delta and self._response.content: - await on_delta(self._response.content) - return self._response + if on_delta and response.content: + await on_delta(response.content) + return response async def chat_with_context( self, @@ -745,292 +760,131 @@ class TestFailoverOnTransientError: class TestRetryBeforeFailover: @pytest.mark.asyncio - async def test_retry_entrypoints_accept_positional_messages(self) -> None: - primary = _FakeProvider("primary", _make_response("primary ok")) + @pytest.mark.parametrize("retry_mode", ["standard", "persistent"]) + async def test_primary_recovers_before_fallback(self, retry_mode: str) -> None: + primary = _FakeProvider( + "primary", + responses=[_error_response("rate limited"), _make_response("primary ok")], + ) factory = MagicMock() - fb = FallbackProvider( - primary=primary, - fallback_presets=[_fallback("fallback-a")], - provider_factory=factory, - ) - messages = [{"role": "user", "content": "hi"}] - provider_context = ProviderCallContext() + provider = FallbackProvider(primary, [_fallback("fallback-a")], factory) - chat_result = await fb.chat_with_retry( - messages, - None, - None, - 4096, - 0.7, - None, - None, - "standard", - None, - provider_context, - ) - stream_result = await fb.chat_stream_with_retry(messages) - - assert chat_result.content == "primary ok" - assert stream_result.content == "primary ok" - assert primary.context_calls[0] is not None - factory.assert_not_called() - - @pytest.mark.asyncio - async def test_primary_retries_and_recovers_before_fallback(self) -> None: - primary = _FakeProvider("primary") - fallback = _FakeProvider("fallback", _make_response("fallback ok")) - factory = MagicMock(return_value=fallback) - fb = FallbackProvider( - primary=primary, - fallback_presets=[_fallback("fallback-a")], - provider_factory=factory, - ) - - with ( - patch.object( - primary, - "chat", - new_callable=AsyncMock, - side_effect=[_error_response("rate limited"), _make_response("primary ok")], - ) as primary_chat, - patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock) as sleep, - ): - result = await fb.chat_with_retry( - messages=[{"role": "user", "content": "hi"}], + with patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock): + result = await provider.chat_with_retry( + [{"role": "user", "content": "hi"}], + retry_mode=retry_mode, ) assert result.content == "primary ok" - assert primary_chat.await_count == 2 - sleep.assert_awaited_once_with(1) + assert len(primary.chat_calls) == 2 factory.assert_not_called() @pytest.mark.asyncio - async def test_primary_exhausts_retries_before_fallback(self) -> None: - primary = _FakeProvider("primary") + async def test_primary_exhausts_before_fallback_without_terminal_event(self) -> None: + primary = _FakeProvider( + "primary", + responses=[_retryable_error(f"attempt {attempt}") for attempt in range(4)], + ) fallback = _FakeProvider("fallback", _make_response("fallback ok")) factory = MagicMock(return_value=fallback) - retry_events: list[str] = [] + retry_events = AsyncMock() + provider = FallbackProvider(primary, [_fallback("fallback-a")], factory) - async def _on_retry_event(message: str) -> None: - retry_events.append(message) - - fb = FallbackProvider( - primary=primary, - fallback_presets=[_fallback("fallback-a")], - provider_factory=factory, - ) - - with ( - patch.object( - primary, - "chat", - new_callable=AsyncMock, - side_effect=[ - _make_response( - f"attempt {attempt}", - finish_reason="error", - error_status_code=503, - ) - for attempt in range(4) - ], - ) as primary_chat, - patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock) as sleep, - ): - result = await fb.chat_with_retry( - messages=[{"role": "user", "content": "hi"}], - on_retry_wait=_on_retry_event, + with patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock): + result = await provider.chat_with_retry( + [{"role": "user", "content": "hi"}], + on_retry_wait=retry_events, ) assert result.content == "fallback ok" - assert primary_chat.await_count == 4 - assert [call.args[0] for call in sleep.await_args_list] == [1, 2, 4] - assert not any("giving up" in event for event in retry_events) - factory.assert_called_once_with(_fallback("fallback-a")) - - @pytest.mark.asyncio - async def test_all_models_exhaust_retries_emits_one_terminal_event(self) -> None: - primary = _FakeProvider( - "primary", - _make_response("primary unavailable", finish_reason="error", error_status_code=503), - ) - fallback = _FakeProvider( - "fallback", - _make_response("fallback unavailable", finish_reason="error", error_status_code=503), - ) - retry_events: list[str] = [] - terminal_events: list[str] = [] - - async def _on_retry_event(message: str) -> None: - retry_events.append(message) - - async def _on_retry_exhausted(message: str) -> None: - terminal_events.append(message) - - fb = FallbackProvider( - primary=primary, - fallback_presets=[_fallback("fallback-a")], - provider_factory=MagicMock(return_value=fallback), - ) - - with patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock): - result = await fb.chat_with_retry( - messages=[{"role": "user", "content": "hi"}], - on_retry_wait=_on_retry_event, - on_retry_exhausted=_on_retry_exhausted, - ) - - assert result.finish_reason == "error" assert len(primary.chat_calls) == 4 - assert len(fallback.chat_calls) == 4 - assert not any("giving up" in event for event in retry_events) - assert terminal_events == ["Model request failed after 4 attempts, giving up."] - - @pytest.mark.asyncio - async def test_persistent_mode_retries_each_candidate_before_repeating_chain(self) -> None: - primary = _FakeProvider("primary") - fallback = _FakeProvider("fallback", _make_response("fallback ok")) - factory = MagicMock(return_value=fallback) - fb = FallbackProvider( - primary=primary, - fallback_presets=[_fallback("fallback-a")], - provider_factory=factory, - ) - - with ( - patch.object( - primary, - "chat", - new_callable=AsyncMock, - side_effect=[ - _error_response(f"server_error request-{attempt}") - for attempt in range(4) - ], - ) as primary_chat, - patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock) as sleep, - ): - result = await fb.chat_with_retry( - messages=[{"role": "user", "content": "hi"}], - retry_mode="persistent", - ) - - assert result.content == "fallback ok" - assert primary_chat.await_count == 4 - assert [call.args[0] for call in sleep.await_args_list] == [1, 2, 4] + assert not any("giving up" in call.args[0] for call in retry_events.await_args_list) factory.assert_called_once_with(_fallback("fallback-a")) @pytest.mark.asyncio - async def test_persistent_mode_repeats_the_whole_fallback_chain(self) -> None: - primary = _FakeProvider( - "primary", - _make_response("primary unavailable", finish_reason="error", error_status_code=503), + async def test_all_candidates_exhaust_emit_one_terminal_event(self) -> None: + primary = _FakeProvider("primary", _retryable_error("primary unavailable")) + fallback = _FakeProvider("fallback", _retryable_error("fallback unavailable")) + retry_events = AsyncMock() + terminal_event = AsyncMock() + provider = FallbackProvider( + primary, + [_fallback("fallback-a")], + MagicMock(return_value=fallback), ) - fallback = _FakeProvider( - "fallback", - _make_response("fallback unavailable", finish_reason="error", error_status_code=503), - ) - factory = MagicMock(return_value=fallback) - retry_events: list[str] = [] - - async def _on_retry_event(message: str) -> None: - retry_events.append(message) - - fb = FallbackProvider( - primary=primary, - fallback_presets=[_fallback("fallback-a")], - provider_factory=factory, - ) - fb._PERSISTENT_IDENTICAL_ERROR_LIMIT = 2 with patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock): - result = await fb.chat_with_retry( - messages=[{"role": "user", "content": "hi"}], + result = await provider.chat_with_retry( + [{"role": "user", "content": "hi"}], + on_retry_wait=retry_events, + on_retry_exhausted=terminal_event, + ) + + assert result.finish_reason == "error" + assert len(primary.chat_calls) == len(fallback.chat_calls) == 4 + assert not any("giving up" in call.args[0] for call in retry_events.await_args_list) + terminal_event.assert_awaited_once_with( + "Model request failed after 4 attempts, giving up." + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("factory_fails", [False, True]) + async def test_persistent_mode_repeats_the_whole_chain( + self, + factory_fails: bool, + ) -> None: + primary = _FakeProvider("primary", _retryable_error("primary unavailable")) + fallback = _FakeProvider("fallback", _retryable_error("fallback unavailable")) + factory = ( + MagicMock(side_effect=ValueError("missing fallback credentials")) + if factory_fails + else MagicMock(return_value=fallback) + ) + terminal_event = AsyncMock() + provider = FallbackProvider(primary, [_fallback("fallback-a")], factory) + provider._PERSISTENT_IDENTICAL_ERROR_LIMIT = 2 + + with patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock): + result = await provider.chat_with_retry( + [{"role": "user", "content": "hi"}], retry_mode="persistent", - on_retry_wait=_on_retry_event, + on_retry_exhausted=terminal_event, ) - terminal_events = [ - event - for event in retry_events - if "giving up" in event or "Persistent retry stopped" in event - ] assert result.finish_reason == "error" assert len(primary.chat_calls) == 8 - assert len(fallback.chat_calls) == 8 + assert len(fallback.chat_calls) == (0 if factory_fails else 8) assert factory.call_count == 2 - assert terminal_events == ["Persistent retry stopped after 2 identical errors."] + terminal_event.assert_awaited_once_with( + "Persistent retry stopped after 2 identical errors." + ) @pytest.mark.asyncio - async def test_persistent_mode_repeats_when_fallback_construction_fails(self) -> None: - primary = _FakeProvider( - "primary", - _make_response( - "primary unavailable", - finish_reason="error", - error_status_code=503, - ), - ) - factory = MagicMock(side_effect=ValueError("missing fallback credentials")) - fb = FallbackProvider( - primary=primary, - fallback_presets=[_fallback("fallback-a")], - provider_factory=factory, - ) - fb._PERSISTENT_IDENTICAL_ERROR_LIMIT = 2 - - with patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock): - result = await fb.chat_with_retry( - messages=[{"role": "user", "content": "hi"}], - retry_mode="persistent", - ) - - assert result.content == "primary unavailable" - assert result.error_status_code == 503 - assert len(primary.chat_calls) == 8 - assert factory.call_count == 2 - - @pytest.mark.asyncio - async def test_persistent_mode_waits_for_open_primary_when_fallback_construction_fails( - self, - ) -> None: + async def test_open_primary_circuit_remains_retryable_when_factory_fails(self) -> None: primary = _FakeProvider("primary") factory = MagicMock(side_effect=ValueError("missing fallback credentials")) - retry_events: list[str] = [] - - async def _on_retry_event(message: str) -> None: - retry_events.append(message) - - fb = FallbackProvider( - primary=primary, - fallback_presets=[_fallback("fallback-a")], - provider_factory=factory, - ) - fb._primary_tripped_at = 100.0 + terminal_event = AsyncMock() + provider = FallbackProvider(primary, [_fallback("fallback-a")], factory) + provider._primary_tripped_at = 100.0 + provider._PERSISTENT_IDENTICAL_ERROR_LIMIT = 2 with ( - patch( - "nanobot.providers.fallback_provider.time.monotonic", - return_value=100.0, - ), + patch("nanobot.providers.fallback_provider.time.monotonic", return_value=100.0), patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock), ): - result = await fb.chat_with_retry( - messages=[{"role": "user", "content": "hi"}], + result = await provider.chat_with_retry( + [{"role": "user", "content": "hi"}], retry_mode="persistent", - on_retry_wait=_on_retry_event, + on_retry_exhausted=terminal_event, ) - terminal_events = [ - event for event in retry_events if "Persistent retry stopped" in event - ] - assert result.finish_reason == "error" assert result.error_should_retry is True assert result.error_retry_after_s == 60 assert primary.chat_calls == [] - assert factory.call_count == fb._PERSISTENT_IDENTICAL_ERROR_LIMIT - assert terminal_events == [ - f"Persistent retry stopped after {fb._PERSISTENT_IDENTICAL_ERROR_LIMIT} " - "identical errors." - ] + assert factory.call_count == 2 + terminal_event.assert_awaited_once_with( + "Persistent retry stopped after 2 identical errors." + ) @pytest.mark.asyncio async def test_fallback_retries_before_trying_next_model(self) -> None: @@ -1039,236 +893,99 @@ class TestRetryBeforeFailover: _make_response( "unauthorized", finish_reason="error", - error_status_code=401, error_kind="authentication", error_should_retry=False, ), ) - fallback_a = _FakeProvider("fallback-a") - fallback_b = _FakeProvider("fallback-b", _make_response("fallback b ok")) - factory = MagicMock(side_effect=[fallback_a, fallback_b]) + fallback_a = _FakeProvider( + "fallback-a", + responses=[_error_response("rate limited"), _make_response("fallback a ok")], + ) + factory = MagicMock(side_effect=[fallback_a, _FakeProvider("fallback-b")]) fallback_a_preset = _fallback("fallback-a") - fb = FallbackProvider( - primary=primary, - fallback_presets=[fallback_a_preset, _fallback("fallback-b")], - provider_factory=factory, + provider = FallbackProvider( + primary, + [fallback_a_preset, _fallback("fallback-b")], + factory, ) - with ( - patch.object( - fallback_a, - "chat", - new_callable=AsyncMock, - side_effect=[_error_response("rate limited"), _make_response("fallback a ok")], - ) as fallback_chat, - patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock), - ): - result = await fb.chat_with_retry( - messages=[{"role": "user", "content": "hi"}], - ) + with patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock): + result = await provider.chat_with_retry([{"role": "user", "content": "hi"}]) assert result.content == "fallback a ok" - assert fallback_chat.await_count == 2 + assert len(fallback_a.chat_calls) == 2 factory.assert_called_once_with(fallback_a_preset) @pytest.mark.asyncio - async def test_streaming_primary_retries_before_fallback(self) -> None: - primary = _FakeProvider("primary") + async def test_stream_recovery_keeps_fallback_eligible(self) -> None: + primary = _FakeProvider( + "primary", + responses=[ + _make_response("partial", finish_reason="error", error_kind="timeout"), + *[_retryable_error() for _ in range(3)], + ], + ) fallback = _FakeProvider("fallback", _make_response("fallback ok")) factory = MagicMock(return_value=fallback) - fb = FallbackProvider( - primary=primary, - fallback_presets=[_fallback("fallback-a")], - provider_factory=factory, - ) - responses = iter([ - _make_response("partial", finish_reason="error", error_kind="timeout"), - _make_response("primary ok"), - ]) - streamed: list[str] = [] - recoveries: list[str] = [] + streamed = AsyncMock() + recovered = AsyncMock() + provider = FallbackProvider(primary, [_fallback("fallback-a")], factory) - async def _stream(**kwargs: Any) -> LLMResponse: - response = next(responses) - if callback := kwargs.get("on_content_delta"): - await callback(response.content) - return response - - async def _delta(text: str) -> None: - streamed.append(text) - - async def _recover() -> None: - recoveries.append("recover") - - with ( - patch.object( - primary, - "chat_stream", - new_callable=AsyncMock, - side_effect=_stream, - ) as primary_stream, - patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock), - ): - result = await fb.chat_stream_with_retry( - messages=[{"role": "user", "content": "hi"}], - on_content_delta=_delta, - on_stream_recover=_recover, - ) - - assert result.content == "primary ok" - assert primary_stream.await_count == 2 - assert streamed == ["partial", "primary ok"] - assert recoveries == ["recover"] - factory.assert_not_called() - - @pytest.mark.asyncio - async def test_streaming_primary_exhausts_retries_before_fallback(self) -> None: - primary = _FakeProvider("primary") - fallback = _FakeProvider("fallback", _make_response("fallback ok")) - factory = MagicMock(return_value=fallback) - fb = FallbackProvider( - primary=primary, - fallback_presets=[_fallback("fallback-a")], - provider_factory=factory, - ) - streamed: list[str] = [] - - async def _delta(text: str) -> None: - streamed.append(text) - - with ( - patch.object( - primary, - "chat_stream", - new_callable=AsyncMock, - side_effect=[_error_response(f"server_error {attempt}") for attempt in range(4)], - ) as primary_stream, - patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock), - ): - result = await fb.chat_stream_with_retry( - messages=[{"role": "user", "content": "hi"}], - on_content_delta=_delta, + with patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock): + result = await provider.chat_stream_with_retry( + [{"role": "user", "content": "hi"}], + on_content_delta=streamed, + on_stream_recover=recovered, ) assert result.content == "fallback ok" - assert primary_stream.await_count == 4 - assert streamed == ["fallback ok"] + assert len(primary.chat_stream_calls) == 4 + assert [call.args[0] for call in streamed.await_args_list] == ["partial", "fallback ok"] + recovered.assert_awaited_once_with() factory.assert_called_once_with(_fallback("fallback-a")) @pytest.mark.asyncio - async def test_streaming_without_delta_callback_still_retries_and_falls_back(self) -> None: - primary = _FakeProvider("primary") + async def test_stream_without_delta_callback_retries_before_fallback(self) -> None: + primary = _FakeProvider( + "primary", + responses=[_retryable_error(f"attempt {attempt}") for attempt in range(4)], + ) fallback = _FakeProvider("fallback", _make_response("fallback ok")) factory = MagicMock(return_value=fallback) - fb = FallbackProvider( - primary=primary, - fallback_presets=[_fallback("fallback-a")], - provider_factory=factory, - ) - responses = iter([_error_response(f"server_error {attempt}") for attempt in range(4)]) + provider = FallbackProvider(primary, [_fallback("fallback-a")], factory) - async def _stream(**kwargs: Any) -> LLMResponse: - response = next(responses) - if callback := kwargs.get("on_content_delta"): - await callback("provider-only delta") - return response - - with ( - patch.object( - primary, - "chat_stream", - new_callable=AsyncMock, - side_effect=_stream, - ) as primary_stream, - patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock), - ): - result = await fb.chat_stream_with_retry( - messages=[{"role": "user", "content": "hi"}], + with patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock): + result = await provider.chat_stream_with_retry( + [{"role": "user", "content": "hi"}] ) assert result.content == "fallback ok" - assert primary_stream.await_count == 4 - factory.assert_called_once_with(_fallback("fallback-a")) - - @pytest.mark.asyncio - async def test_streaming_recovery_keeps_fallback_eligible(self) -> None: - primary = _FakeProvider("primary") - fallback = _FakeProvider("fallback", _make_response("fallback ok")) - factory = MagicMock(return_value=fallback) - fb = FallbackProvider( - primary=primary, - fallback_presets=[_fallback("fallback-a")], - provider_factory=factory, - ) - responses = iter([ - _make_response("partial", finish_reason="error", error_kind="timeout"), - *[_error_response(f"server_error {attempt}") for attempt in range(3)], - ]) - streamed: list[str] = [] - recoveries: list[str] = [] - - async def _stream(**kwargs: Any) -> LLMResponse: - response = next(responses) - if response.content == "partial" and (callback := kwargs.get("on_content_delta")): - await callback(response.content) - return response - - async def _delta(text: str) -> None: - streamed.append(text) - - async def _recover() -> None: - recoveries.append("recover") - - with ( - patch.object(primary, "chat_stream", new_callable=AsyncMock, side_effect=_stream), - patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock), - ): - result = await fb.chat_stream_with_retry( - messages=[{"role": "user", "content": "hi"}], - on_content_delta=_delta, - on_stream_recover=_recover, - ) - - assert result.content == "fallback ok" - assert streamed == ["partial", "fallback ok"] - assert recoveries == ["recover"] + assert len(primary.chat_stream_calls) == 4 factory.assert_called_once_with(_fallback("fallback-a")) @pytest.mark.asyncio async def test_unrecovered_stream_keeps_non_timeout_fallback_blocked(self) -> None: - primary = _FakeProvider("primary") - factory = MagicMock() - fb = FallbackProvider( - primary=primary, - fallback_presets=[_fallback("fallback-a")], - provider_factory=factory, + primary = _FakeProvider( + "primary", + responses=[ + _make_response("partial", finish_reason="error", error_kind="timeout"), + _retryable_error(), + _retryable_error(), + _retryable_error("last error"), + ], ) - responses = iter([ - _make_response("partial", finish_reason="error", error_kind="timeout"), - *[_error_response(f"server_error {attempt}") for attempt in range(3)], - ]) - streamed: list[str] = [] + factory = MagicMock() + streamed = AsyncMock() + provider = FallbackProvider(primary, [_fallback("fallback-a")], factory) - async def _stream(**kwargs: Any) -> LLMResponse: - response = next(responses) - if response.content == "partial" and (callback := kwargs.get("on_content_delta")): - await callback(response.content) - return response - - async def _delta(text: str) -> None: - streamed.append(text) - - with ( - patch.object(primary, "chat_stream", new_callable=AsyncMock, side_effect=_stream), - patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock), - ): - result = await fb.chat_stream_with_retry( - messages=[{"role": "user", "content": "hi"}], - on_content_delta=_delta, + with patch("nanobot.providers.base.asyncio.sleep", new_callable=AsyncMock): + result = await provider.chat_stream_with_retry( + [{"role": "user", "content": "hi"}], + on_content_delta=streamed, ) - assert result.content == "server_error 2" - assert streamed == ["partial"] + assert result.content == "last error" + streamed.assert_awaited_once_with("partial") factory.assert_not_called() @@ -1585,6 +1302,29 @@ class TestNoFallbackWhenEmptyList: assert result.finish_reason == "error" factory.assert_not_called() + @pytest.mark.asyncio + async def test_retry_entrypoints_delegate_to_primary(self) -> None: + primary = _FakeProvider("primary") + provider = FallbackProvider(primary, [], MagicMock()) + response = _make_response("primary ok") + + with ( + patch.object( + primary, "chat_with_retry", new_callable=AsyncMock, return_value=response + ) as chat_retry, + patch.object( + primary, + "chat_stream_with_retry", + new_callable=AsyncMock, + return_value=response, + ) as stream_retry, + ): + assert (await provider.chat_with_retry([])) is response + assert (await provider.chat_stream_with_retry([])) is response + + chat_retry.assert_awaited_once() + stream_retry.assert_awaited_once() + class TestChatStreamFailover: @pytest.mark.asyncio