mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-09-02 01:01:52 +03:00
fix(providers): retry before falling back
This commit is contained in:
@@ -7,8 +7,9 @@ import json
|
|||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable, Generator
|
||||||
from contextlib import suppress
|
from contextlib import contextmanager, suppress
|
||||||
|
from contextvars import ContextVar
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
@@ -25,6 +26,26 @@ DEFAULT_STREAM_IDLE_TIMEOUT_S = 90.0
|
|||||||
MAX_STREAM_IDLE_TIMEOUT_S = 3600.0
|
MAX_STREAM_IDLE_TIMEOUT_S = 3600.0
|
||||||
RETRY_AFTER_BUFFER = 1
|
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(
|
def resolve_stream_idle_timeout_s(
|
||||||
*,
|
*,
|
||||||
@@ -910,12 +931,16 @@ class LLMProvider(ABC):
|
|||||||
kw["provider_context"] = provider_context
|
kw["provider_context"] = provider_context
|
||||||
if on_stream_recover and getattr(self, "supports_stream_recover_callback", False):
|
if on_stream_recover and getattr(self, "supports_stream_recover_callback", False):
|
||||||
kw["on_stream_recover"] = _recover_stream
|
kw["on_stream_recover"] = _recover_stream
|
||||||
|
on_retry_exhausted = _RETRY_EXHAUSTED_CALLBACK.get()
|
||||||
return await self._run_with_retry(
|
return await self._run_with_retry(
|
||||||
self._safe_chat_stream,
|
self._safe_chat_stream,
|
||||||
kw,
|
kw,
|
||||||
messages,
|
messages,
|
||||||
retry_mode=retry_mode,
|
retry_mode=retry_mode,
|
||||||
on_retry_wait=on_retry_wait,
|
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,
|
should_retry_guard=lambda: not has_streamed_content,
|
||||||
on_stream_recover=_recover_stream if on_stream_recover else None,
|
on_stream_recover=_recover_stream if on_stream_recover else None,
|
||||||
)
|
)
|
||||||
@@ -956,12 +981,16 @@ class LLMProvider(ABC):
|
|||||||
)
|
)
|
||||||
if provider_context is not None:
|
if provider_context is not None:
|
||||||
kw["provider_context"] = provider_context
|
kw["provider_context"] = provider_context
|
||||||
|
on_retry_exhausted = _RETRY_EXHAUSTED_CALLBACK.get()
|
||||||
return await self._run_with_retry(
|
return await self._run_with_retry(
|
||||||
self._safe_chat,
|
self._safe_chat,
|
||||||
kw,
|
kw,
|
||||||
messages,
|
messages,
|
||||||
retry_mode=retry_mode,
|
retry_mode=retry_mode,
|
||||||
on_retry_wait=on_retry_wait,
|
on_retry_wait=on_retry_wait,
|
||||||
|
on_retry_exhausted=(
|
||||||
|
on_retry_exhausted if on_retry_exhausted is not None else on_retry_wait
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -1067,6 +1096,7 @@ class LLMProvider(ABC):
|
|||||||
*,
|
*,
|
||||||
retry_mode: str,
|
retry_mode: str,
|
||||||
on_retry_wait: Callable[[str], Awaitable[None]] | None,
|
on_retry_wait: Callable[[str], Awaitable[None]] | None,
|
||||||
|
on_retry_exhausted: Callable[[str], Awaitable[None]] | None,
|
||||||
should_retry_guard: Callable[[], bool] | None = None,
|
should_retry_guard: Callable[[], bool] | None = None,
|
||||||
on_stream_recover: Callable[[], Awaitable[None]] | None = None,
|
on_stream_recover: Callable[[], Awaitable[None]] | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
@@ -1154,21 +1184,21 @@ class LLMProvider(ABC):
|
|||||||
identical_error_count,
|
identical_error_count,
|
||||||
(response.content or "")[:120].lower(),
|
(response.content or "")[:120].lower(),
|
||||||
)
|
)
|
||||||
if on_retry_wait:
|
if on_retry_exhausted:
|
||||||
await on_retry_wait(
|
await on_retry_exhausted(
|
||||||
f"Persistent retry stopped after {identical_error_count} identical errors."
|
f"Persistent retry stopped after {identical_error_count} identical errors."
|
||||||
)
|
)
|
||||||
return response
|
return response
|
||||||
|
|
||||||
if not persistent and attempt > len(delays):
|
if not persistent and attempt > len(delays):
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"LLM request failed after {} retries, giving up: {}",
|
"LLM request failed after {} attempts, giving up: {}",
|
||||||
attempt,
|
attempt,
|
||||||
(response.content or "")[:120].lower(),
|
(response.content or "")[:120].lower(),
|
||||||
)
|
)
|
||||||
if on_retry_wait:
|
if on_retry_exhausted:
|
||||||
await on_retry_wait(
|
await on_retry_exhausted(
|
||||||
f"Model request failed after {attempt} retries, giving up."
|
f"Model request failed after {attempt} attempts, giving up."
|
||||||
)
|
)
|
||||||
break
|
break
|
||||||
|
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ from nanobot.providers.base import (
|
|||||||
LLMResponse,
|
LLMResponse,
|
||||||
ProviderCallContext,
|
ProviderCallContext,
|
||||||
ProviderConversationState,
|
ProviderConversationState,
|
||||||
|
retry_exhaustion_callback,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Circuit breaker tuned to match OpenAICompatProvider's Responses API breaker.
|
# Circuit breaker tuned to match OpenAICompatProvider's Responses API breaker.
|
||||||
@@ -91,6 +92,7 @@ _FALLBACK_ERROR_TOKENS = (
|
|||||||
|
|
||||||
|
|
||||||
FallbackModelObserver = Callable[[str], Awaitable[None]]
|
FallbackModelObserver = Callable[[str], Awaitable[None]]
|
||||||
|
_UNSET = object()
|
||||||
|
|
||||||
|
|
||||||
class FallbackProvider(LLMProvider):
|
class FallbackProvider(LLMProvider):
|
||||||
@@ -105,6 +107,7 @@ class FallbackProvider(LLMProvider):
|
|||||||
|
|
||||||
Key design:
|
Key design:
|
||||||
- Failover is request-scoped (the wrapper itself is stateless between turns).
|
- 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,
|
- Skipped when content was already streamed to avoid duplicate output,
|
||||||
except timeout recovery can resume in a new stream segment.
|
except timeout recovery can resume in a new stream segment.
|
||||||
- Recursive failover is prevented by the factory returning plain providers.
|
- 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
|
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(
|
async def chat_with_context(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -234,6 +278,154 @@ class FallbackProvider(LLMProvider):
|
|||||||
on_stream_recover=on_stream_recover,
|
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(
|
async def chat_stream_with_context(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -275,6 +467,7 @@ class FallbackProvider(LLMProvider):
|
|||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
primary_model = kwargs.get("model") or self._primary.get_default_model()
|
primary_model = kwargs.get("model") or self._primary.get_default_model()
|
||||||
primary_was_attempted = False
|
primary_was_attempted = False
|
||||||
|
primary_response: LLMResponse | None = None
|
||||||
primary_error = "unknown error"
|
primary_error = "unknown error"
|
||||||
# A primary error eligible for failover did not return a replacement
|
# A primary error eligible for failover did not return a replacement
|
||||||
# continuation, so the incoming primary state remains reusable.
|
# continuation, so the incoming primary state remains reusable.
|
||||||
@@ -287,6 +480,7 @@ class FallbackProvider(LLMProvider):
|
|||||||
self._primary_failures = 0
|
self._primary_failures = 0
|
||||||
self._primary_tripped_at = None
|
self._primary_tripped_at = None
|
||||||
return response
|
return response
|
||||||
|
primary_response = response
|
||||||
primary_error = (response.content or primary_error)[:120]
|
primary_error = (response.content or primary_error)[:120]
|
||||||
|
|
||||||
if has_streamed is not None and has_streamed[0]:
|
if has_streamed is not None and has_streamed[0]:
|
||||||
@@ -326,7 +520,7 @@ class FallbackProvider(LLMProvider):
|
|||||||
else:
|
else:
|
||||||
logger.debug("Primary model '{}' circuit open; skipping", primary_model)
|
logger.debug("Primary model '{}' circuit open; skipping", primary_model)
|
||||||
|
|
||||||
last_response: LLMResponse | None = None
|
last_response = primary_response
|
||||||
primary_skipped = not primary_was_attempted
|
primary_skipped = not primary_was_attempted
|
||||||
for idx, fallback in enumerate(self._fallback_presets):
|
for idx, fallback in enumerate(self._fallback_presets):
|
||||||
fallback_model = fallback.model
|
fallback_model = fallback.model
|
||||||
@@ -423,11 +617,22 @@ class FallbackProvider(LLMProvider):
|
|||||||
last_response,
|
last_response,
|
||||||
preserve_provider_state_on_error=preserve_primary_state,
|
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(
|
return LLMResponse(
|
||||||
content=f"Primary model '{primary_model}' circuit open and no fallbacks available",
|
content=f"Primary model '{primary_model}' circuit open and no fallbacks available",
|
||||||
finish_reason="error",
|
finish_reason="error",
|
||||||
preserve_provider_state_on_error=preserve_primary_state,
|
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:
|
async def _notify_fallback_model(self, model: str) -> None:
|
||||||
|
|||||||
@@ -3,7 +3,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -743,6 +743,565 @@ class TestFailoverOnTransientError:
|
|||||||
factory.assert_called_once_with(_fallback("fallback-a"))
|
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:
|
class TestFailoverOnArrearageError:
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_non_retryable_quota_tries_configured_fallback(self) -> None:
|
async def test_non_retryable_quota_tries_configured_fallback(self) -> None:
|
||||||
|
|||||||
@@ -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 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
|
@pytest.mark.asyncio
|
||||||
|
|||||||
Reference in New Issue
Block a user