refactor(providers): simplify retry fallback routing

This commit is contained in:
chengyongru
2026-08-21 16:46:38 +08:00
committed by chengyongru
parent 9a6dc371b3
commit 761e95b659
3 changed files with 319 additions and 617 deletions
+29 -4
View File
@@ -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()
+105 -168
View File
@@ -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,
*,
+185 -445
View File
@@ -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