feat(webui): show the actual fallback model (#5017)

This commit is contained in:
chengyongru
2026-07-23 15:57:13 +08:00
committed by GitHub
parent 4188ffc88d
commit 96eb965aae
15 changed files with 352 additions and 8 deletions
+19
View File
@@ -82,6 +82,9 @@ _FALLBACK_ERROR_TOKENS = (
)
FallbackModelObserver = Callable[[str], Awaitable[None]]
class FallbackProvider(LLMProvider):
"""Wrap a primary provider and transparently failover to fallback models.
@@ -108,10 +111,12 @@ class FallbackProvider(LLMProvider):
primary: LLMProvider,
fallback_presets: list[Any],
provider_factory: Callable[[Any], LLMProvider],
fallback_model_observer: FallbackModelObserver | None = None,
):
self._primary = primary
self._fallback_presets = list(fallback_presets)
self._provider_factory = provider_factory
self._fallback_model_observer = fallback_model_observer
self._has_fallbacks = bool(fallback_presets)
self._primary_failures = 0
self._primary_tripped_at: float | None = None
@@ -127,6 +132,10 @@ class FallbackProvider(LLMProvider):
def get_default_model(self) -> str:
return self._primary.get_default_model()
def set_fallback_model_observer(self, observer: FallbackModelObserver | None) -> None:
"""Attach a process-level observer without changing request call signatures."""
self._fallback_model_observer = observer
@property
def supports_progress_deltas(self) -> bool:
return bool(getattr(self._primary, "supports_progress_deltas", False))
@@ -268,6 +277,8 @@ class FallbackProvider(LLMProvider):
)
continue
await self._notify_fallback_model(fallback_model)
original_values = {
name: kwargs.get(name, _MISSING)
for name in ("model", "max_tokens", "temperature", "reasoning_effort")
@@ -315,6 +326,14 @@ class FallbackProvider(LLMProvider):
finish_reason="error",
)
async def _notify_fallback_model(self, model: str) -> None:
if self._fallback_model_observer is None:
return
try:
await self._fallback_model_observer(model)
except Exception:
logger.exception("fallback model observer failed for '{}'", model)
@staticmethod
def _should_fallback(response: LLMResponse) -> bool:
if LLMProvider.is_arrearage_response(response):