fix(providers): preserve retry metadata for raised errors

This commit is contained in:
KDB
2026-09-03 18:31:35 +08:00
committed by Xubin Ren
parent b321c63e1a
commit 480c8dd744
3 changed files with 91 additions and 45 deletions
+57 -2
View File
@@ -943,6 +943,61 @@ class LLMProvider(ABC):
""" """
pass pass
@staticmethod
def _error_response_from_exception(exc: Exception) -> LLMResponse:
"""Convert an unexpected exception while retaining retry metadata."""
error_names = tuple(cls.__name__.lower() for cls in type(exc).__mro__)
error_kind: str | None = None
error_should_retry: bool | None = None
if any("timeout" in name for name in error_names):
error_kind = "timeout"
error_should_retry = True
elif any(
token in name
for name in error_names
for token in ("connect", "connection", "network", "protocol", "transport")
):
error_kind = "connection"
error_should_retry = True
elif any(
"ratelimit" in name or "throttl" in name
for name in error_names
):
error_kind = "rate_limit"
error_should_retry = True
elif any(
"server" in name or "internal" in name
for name in error_names
):
error_kind = "server_error"
error_should_retry = True
elif any(
token in name
for name in error_names
for token in ("auth", "credential", "permissiondenied", "unauthor")
):
error_kind = "authentication"
response = getattr(exc, "response", None)
raw_status = getattr(exc, "status_code", None)
if raw_status is None and response is not None:
raw_status = getattr(response, "status_code", None)
try:
error_status_code = int(raw_status) if raw_status is not None else None
except (TypeError, ValueError):
error_status_code = None
detail = str(exc).strip() or type(exc).__name__
return LLMResponse(
content=f"Error calling LLM: {detail}",
finish_reason="error",
error_status_code=error_status_code,
error_kind=error_kind,
error_type=getattr(exc, "error_type", None),
error_code=getattr(exc, "error_code", None),
error_should_retry=error_should_retry,
)
@classmethod @classmethod
def _is_transient_error(cls, content: str | None) -> bool: def _is_transient_error(cls, content: str | None) -> bool:
err = (content or "").lower() err = (content or "").lower()
@@ -1226,7 +1281,7 @@ class LLMProvider(ABC):
) )
raise raise
except Exception as exc: except Exception as exc:
response = LLMResponse(content=f"Error calling LLM: {exc}", finish_reason="error") response = self._error_response_from_exception(exc)
return self._observe_llm_call( return self._observe_llm_call(
response, response,
kwargs, kwargs,
@@ -1368,7 +1423,7 @@ class LLMProvider(ABC):
) )
raise raise
except Exception as exc: except Exception as exc:
response = LLMResponse(content=f"Error calling LLM: {exc}", finish_reason="error") response = self._error_response_from_exception(exc)
return self._observe_llm_call( return self._observe_llm_call(
_attach_stream_timing(response), _attach_stream_timing(response),
kwargs, kwargs,
+7 -40
View File
@@ -10,7 +10,6 @@ from collections.abc import Awaitable, Callable
from dataclasses import replace from dataclasses import replace
from typing import Any from typing import Any
import httpx
from loguru import logger from loguru import logger
from nanobot.providers.base import ( from nanobot.providers.base import (
@@ -622,45 +621,13 @@ class FallbackProvider(LLMProvider):
except asyncio.CancelledError: except asyncio.CancelledError:
raise raise
except Exception as exc: except Exception as exc:
error_name = type(exc).__name__.lower() response = LLMProvider._error_response_from_exception(exc)
detail = str(exc).strip() or type(exc).__name__ if response.error_kind is None and any(
detail_lower = detail.lower() token in (str(exc).strip() or type(exc).__name__).lower()
error_kind: str | None = None for token in _AUTHENTICATION_ERROR_TOKENS
error_should_retry: bool | None = None ):
if isinstance(exc, (httpx.TimeoutException, asyncio.TimeoutError)): response.error_kind = "authentication"
error_kind = "timeout" return response, exc
error_should_retry = True
elif isinstance(exc, (httpx.NetworkError, httpx.TransportError, ConnectionError)):
error_kind = "connection"
error_should_retry = True
elif any(
token in error_name
for token in ("auth", "credential", "permissiondenied", "unauthor")
) or any(token in detail_lower for token in _AUTHENTICATION_ERROR_TOKENS):
error_kind = "authentication"
elif "ratelimit" in error_name or "throttl" in error_name:
error_kind = "rate_limit"
error_should_retry = True
elif "server" in error_name or "internal" in error_name:
error_kind = "server_error"
error_should_retry = True
response = getattr(exc, "response", None)
raw_status = getattr(exc, "status_code", None)
if raw_status is None and response is not None:
raw_status = getattr(response, "status_code", None)
try:
error_status_code = int(raw_status) if raw_status is not None else None
except (TypeError, ValueError):
error_status_code = None
return LLMResponse(
content=f"Error calling LLM: {detail}",
finish_reason="error",
error_status_code=error_status_code,
error_kind=error_kind,
error_should_retry=error_should_retry,
), exc
async def _notify_fallback_model(self, model: str) -> None: async def _notify_fallback_model(self, model: str) -> None:
if self._fallback_model_observer is None: if self._fallback_model_observer is None:
+27 -3
View File
@@ -353,7 +353,7 @@ class _RaisingProvider(LLMProvider):
"""Provider whose chat/chat_stream raise, like an auth/setup failure.""" """Provider whose chat/chat_stream raise, like an auth/setup failure."""
def __init__(self, name: str = "raiser", exc: BaseException | None = None): def __init__(self, name: str = "raiser", exc: BaseException | None = None):
super().__init__() super().__init__(provider_name=name)
self.name = name self.name = name
self._exc = exc if exc is not None else RuntimeError("GitHub Copilot is not logged in.") self._exc = exc if exc is not None else RuntimeError("GitHub Copilot is not logged in.")
@@ -378,8 +378,8 @@ class _StreamingThenRaisingProvider(_RaisingProvider):
class TestFallbackWhenPrimaryRaises: class TestFallbackWhenPrimaryRaises:
@pytest.mark.parametrize( @pytest.mark.parametrize(
"exc", "exc",
[TimeoutError(), httpx.ReadTimeout(""), httpx.ConnectError("")], [TimeoutError(), httpx.ReadTimeout(""), httpx.ConnectError(""), httpx.ReadError("")],
ids=["asyncio-timeout", "httpx-timeout", "connection"], ids=["asyncio-timeout", "httpx-timeout", "connection", "httpx-network-error"],
) )
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_transient_exception_triggers_fallback(self, exc: Exception) -> None: async def test_transient_exception_triggers_fallback(self, exc: Exception) -> None:
@@ -423,6 +423,30 @@ class TestFallbackWhenPrimaryRaises:
assert result.finish_reason == "stop" assert result.finish_reason == "stop"
factory.assert_called_once_with(_fallback("fallback-a")) factory.assert_called_once_with(_fallback("fallback-a"))
@pytest.mark.asyncio
@pytest.mark.parametrize("stream", [False, True], ids=["chat", "stream"])
async def test_retry_entry_point_preserves_transient_exception_metadata(
self,
stream: bool,
) -> None:
"""A retry wrapper must not hide an empty transient exception from failover."""
primary = _RaisingProvider("primary", TimeoutError())
fallback = _FakeProvider("fallback", _make_response("fallback ok"))
factory = MagicMock(return_value=fallback)
fb = FallbackProvider(
primary=primary,
fallback_presets=[_fallback("fallback-a")],
provider_factory=factory,
)
retry = fb.chat_stream_with_retry if stream else fb.chat_with_retry
result = await retry(messages=[{"role": "user", "content": "hi"}], model="primary-model")
assert result.content == "fallback ok"
assert result.finish_reason == "stop"
factory.assert_called_once_with(_fallback("fallback-a"))
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_authentication_exception_message_is_classified(self) -> None: async def test_authentication_exception_message_is_classified(self) -> None:
primary = _RaisingProvider("primary") primary = _RaisingProvider("primary")