diff --git a/nanobot/providers/base.py b/nanobot/providers/base.py index 0c6a52978..20364b687 100644 --- a/nanobot/providers/base.py +++ b/nanobot/providers/base.py @@ -7,8 +7,9 @@ import json import os import re from abc import ABC, abstractmethod -from collections.abc import Awaitable, Callable -from contextlib import suppress +from collections.abc import Awaitable, Callable, Generator +from contextlib import contextmanager, suppress +from contextvars import ContextVar from copy import deepcopy from dataclasses import dataclass, field from datetime import datetime, timezone @@ -25,6 +26,26 @@ DEFAULT_STREAM_IDLE_TIMEOUT_S = 90.0 MAX_STREAM_IDLE_TIMEOUT_S = 3600.0 RETRY_AFTER_BUFFER = 1 +RetryEventCallback = Callable[[str], Awaitable[None]] +_RETRY_EXHAUSTED_CALLBACK: ContextVar[RetryEventCallback | None] = ContextVar( + "nanobot_retry_exhausted_callback", + default=None, +) + + +@contextmanager +def retry_exhaustion_callback(callback: RetryEventCallback) -> Generator[None, None, None]: + """Redirect terminal retry events within one async call context. + + Provider wrappers use this internal scope to defer a candidate's terminal + notification without changing the public retry-method signatures. + """ + token = _RETRY_EXHAUSTED_CALLBACK.set(callback) + try: + yield + finally: + _RETRY_EXHAUSTED_CALLBACK.reset(token) + def resolve_stream_idle_timeout_s( *, @@ -910,12 +931,16 @@ 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 + on_retry_exhausted = _RETRY_EXHAUSTED_CALLBACK.get() return await self._run_with_retry( self._safe_chat_stream, kw, messages, retry_mode=retry_mode, on_retry_wait=on_retry_wait, + on_retry_exhausted=( + on_retry_exhausted if on_retry_exhausted is not None else on_retry_wait + ), should_retry_guard=lambda: not has_streamed_content, on_stream_recover=_recover_stream if on_stream_recover else None, ) @@ -956,12 +981,16 @@ class LLMProvider(ABC): ) if provider_context is not None: kw["provider_context"] = provider_context + on_retry_exhausted = _RETRY_EXHAUSTED_CALLBACK.get() return await self._run_with_retry( self._safe_chat, kw, messages, retry_mode=retry_mode, on_retry_wait=on_retry_wait, + on_retry_exhausted=( + on_retry_exhausted if on_retry_exhausted is not None else on_retry_wait + ), ) @classmethod @@ -1067,6 +1096,7 @@ class LLMProvider(ABC): *, retry_mode: str, on_retry_wait: Callable[[str], Awaitable[None]] | None, + on_retry_exhausted: Callable[[str], Awaitable[None]] | None, should_retry_guard: Callable[[], bool] | None = None, on_stream_recover: Callable[[], Awaitable[None]] | None = None, ) -> LLMResponse: @@ -1154,21 +1184,21 @@ class LLMProvider(ABC): identical_error_count, (response.content or "")[:120].lower(), ) - if on_retry_wait: - await on_retry_wait( + if on_retry_exhausted: + await on_retry_exhausted( f"Persistent retry stopped after {identical_error_count} identical errors." ) return response if not persistent and attempt > len(delays): logger.warning( - "LLM request failed after {} retries, giving up: {}", + "LLM request failed after {} attempts, giving up: {}", attempt, (response.content or "")[:120].lower(), ) - if on_retry_wait: - await on_retry_wait( - f"Model request failed after {attempt} retries, giving up." + if on_retry_exhausted: + await on_retry_exhausted( + f"Model request failed after {attempt} attempts, giving up." ) break diff --git a/nanobot/providers/fallback_provider.py b/nanobot/providers/fallback_provider.py index 24cf2785a..513cf0d01 100644 --- a/nanobot/providers/fallback_provider.py +++ b/nanobot/providers/fallback_provider.py @@ -17,6 +17,7 @@ from nanobot.providers.base import ( LLMResponse, ProviderCallContext, ProviderConversationState, + retry_exhaustion_callback, ) # Circuit breaker tuned to match OpenAICompatProvider's Responses API breaker. @@ -91,6 +92,7 @@ _FALLBACK_ERROR_TOKENS = ( FallbackModelObserver = Callable[[str], Awaitable[None]] +_UNSET = object() class FallbackProvider(LLMProvider): @@ -105,6 +107,7 @@ class FallbackProvider(LLMProvider): Key design: - Failover is request-scoped (the wrapper itself is stateless between turns). + - Retrying entry points exhaust one provider's retry policy before failover. - Skipped when content was already streamed to avoid duplicate output, except timeout recovery can resume in a new stream segment. - Recursive failover is prevented by the factory returning plain providers. @@ -193,6 +196,47 @@ class FallbackProvider(LLMProvider): lambda p, kw: p.chat(**kw), kwargs, has_streamed=None ) + async def 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: Callable[[str], Awaitable[None]] | None = None, + provider_context: ProviderCallContext | 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, + } + 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: + return await self._primary.chat_with_retry(**call_kwargs) + return await self._route_with_retry_fallback( + lambda p, kw: p.chat_with_retry(**kw), + call_kwargs, + has_streamed=None, + ) + async def chat_with_context( self, *, @@ -234,6 +278,154 @@ 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: Callable[[str], Awaitable[None]] | None = None, + provider_context: ProviderCallContext | 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, + } + 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], + ) + + async def _route_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 = None, + persistent_retry_guard: Callable[[], bool] | None = None, + ) -> LLMResponse: + """Apply finite retries per provider and persistence to the whole chain.""" + on_retry_wait: Callable[[str], Awaitable[None]] | 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_wait, + ) + + async def _call_chain(**chain_kwargs: Any) -> LLMResponse: + chain_kwargs["retry_mode"] = "standard" + return await self._try_with_retry_fallback( + call, + chain_kwargs, + has_streamed=has_streamed, + on_stream_recover=on_stream_recover, + on_retry_exhausted=None, + ) + + return await self._run_with_retry( + _call_chain, + dict(kwargs), + kwargs["messages"], + retry_mode="persistent", + on_retry_wait=on_retry_wait, + on_retry_exhausted=on_retry_wait, + should_retry_guard=persistent_retry_guard, + 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: Callable[[str], Awaitable[None]] | 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 + with retry_exhaustion_callback(_capture_exhaustion): + return await call(provider, call_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, *, @@ -275,6 +467,7 @@ class FallbackProvider(LLMProvider): ) -> LLMResponse: primary_model = kwargs.get("model") or self._primary.get_default_model() primary_was_attempted = False + primary_response: LLMResponse | None = None primary_error = "unknown error" # A primary error eligible for failover did not return a replacement # continuation, so the incoming primary state remains reusable. @@ -287,6 +480,7 @@ class FallbackProvider(LLMProvider): self._primary_failures = 0 self._primary_tripped_at = None return response + primary_response = response primary_error = (response.content or primary_error)[:120] if has_streamed is not None and has_streamed[0]: @@ -326,7 +520,7 @@ class FallbackProvider(LLMProvider): else: logger.debug("Primary model '{}' circuit open; skipping", primary_model) - last_response: LLMResponse | None = None + last_response = primary_response primary_skipped = not primary_was_attempted for idx, fallback in enumerate(self._fallback_presets): fallback_model = fallback.model @@ -423,11 +617,22 @@ class FallbackProvider(LLMProvider): last_response, preserve_provider_state_on_error=preserve_primary_state, ) - # Primary was tripped and we have no fallbacks — synthesize an error. + # Primary was skipped and no fallback returned a response. Keep the result + # transient until the primary circuit is eligible for another probe. + retry_after_s = ( + max( + 0.1, + _PRIMARY_COOLDOWN_S - (time.monotonic() - self._primary_tripped_at), + ) + if self._primary_tripped_at is not None + else None + ) return LLMResponse( content=f"Primary model '{primary_model}' circuit open and no fallbacks available", finish_reason="error", preserve_provider_state_on_error=preserve_primary_state, + error_retry_after_s=retry_after_s, + error_should_retry=True, ) async def _notify_fallback_model(self, model: str) -> None: diff --git a/tests/agent/test_runner_fallback.py b/tests/agent/test_runner_fallback.py index 9afe50cbd..30f8eadc0 100644 --- a/tests/agent/test_runner_fallback.py +++ b/tests/agent/test_runner_fallback.py @@ -3,7 +3,7 @@ from __future__ import annotations from typing import Any -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest from loguru import logger @@ -743,6 +743,565 @@ class TestFailoverOnTransientError: factory.assert_called_once_with(_fallback("fallback-a")) +class TestRetryBeforeFailover: + @pytest.mark.asyncio + async def test_retry_entrypoints_accept_positional_messages(self) -> None: + primary = _FakeProvider("primary", _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() + + 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"}], + ) + + assert result.content == "primary ok" + assert primary_chat.await_count == 2 + sleep.assert_awaited_once_with(1) + factory.assert_not_called() + + @pytest.mark.asyncio + async def test_primary_exhausts_retries_before_fallback(self) -> None: + primary = _FakeProvider("primary") + fallback = _FakeProvider("fallback", _make_response("fallback ok")) + 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, + ) + + 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, + ) + + 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] = [] + + async def _on_retry_event(message: str) -> None: + retry_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, + ) + + terminal_events = [event for event in retry_events if "giving up" in event] + assert result.finish_reason == "error" + assert len(primary.chat_calls) == 4 + assert len(fallback.chat_calls) == 4 + 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] + 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), + ) + 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"}], + retry_mode="persistent", + on_retry_wait=_on_retry_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 factory.call_count == 2 + assert terminal_events == ["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: + 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 + + with ( + 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"}], + retry_mode="persistent", + on_retry_wait=_on_retry_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." + ] + + @pytest.mark.asyncio + async def test_legacy_retry_override_does_not_receive_internal_callback(self) -> None: + class LegacyRetryProvider(_FakeProvider): + def __init__(self) -> None: + super().__init__("legacy", _make_response("legacy ok")) + self.retry_calls = 0 + + async def chat_with_retry( + self, + messages: list[dict[str, Any]], + tools: list[dict[str, Any]] | None = None, + model: str | None = None, + max_tokens: object = None, + temperature: object = None, + reasoning_effort: object = None, + tool_choice: str | dict[str, Any] | None = None, + retry_mode: str = "standard", + on_retry_wait: Any = None, + provider_context: ProviderCallContext | None = None, + ) -> LLMResponse: + self.retry_calls += 1 + return await self.chat(messages=messages) + + primary = LegacyRetryProvider() + fb = FallbackProvider( + primary=primary, + fallback_presets=[_fallback("fallback-a")], + provider_factory=MagicMock(), + ) + + result = await fb.chat_with_retry(messages=[{"role": "user", "content": "hi"}]) + + assert result.content == "legacy ok" + assert primary.retry_calls == 1 + + @pytest.mark.asyncio + async def test_fallback_retries_before_trying_next_model(self) -> None: + primary = _FakeProvider( + "primary", + _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_preset = _fallback("fallback-a") + fb = FallbackProvider( + primary=primary, + fallback_presets=[fallback_a_preset, _fallback("fallback-b")], + provider_factory=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"}], + ) + + assert result.content == "fallback a ok" + assert fallback_chat.await_count == 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") + 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] = [] + + 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, + ) + + assert result.content == "fallback ok" + assert primary_stream.await_count == 4 + assert streamed == ["fallback ok"] + 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") + 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)]) + + 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"}], + ) + + 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"] + 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, + ) + 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] = [] + + 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, + ) + + assert result.content == "server_error 2" + assert streamed == ["partial"] + factory.assert_not_called() + + class TestFailoverOnArrearageError: @pytest.mark.asyncio async def test_non_retryable_quota_tries_configured_fallback(self) -> None: diff --git a/tests/providers/test_provider_retry.py b/tests/providers/test_provider_retry.py index dac7e6e60..dfba6a1f5 100644 --- a/tests/providers/test_provider_retry.py +++ b/tests/providers/test_provider_retry.py @@ -129,7 +129,7 @@ async def test_chat_with_retry_emits_terminal_progress_when_standard_retries_exh ) assert response.content == "503 final server error" - assert progress[-1] == "Model request failed after 4 retries, giving up." + assert progress[-1] == "Model request failed after 4 attempts, giving up." @pytest.mark.asyncio