fix(providers): retry before falling back

This commit is contained in:
chengyongru
2026-08-21 16:46:38 +08:00
committed by chengyongru
parent 98660c19cc
commit f93d4c3ae4
4 changed files with 806 additions and 12 deletions
+560 -1
View File
@@ -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: