470cf75722
Merged upstream/main (418abca, release v0.3.0) into our fork on a
dedicated branch. 21 conflicting files resolved; main worktree untouched.
Resolution policy and key decisions:
- Keep Ai4Sci runtime endpoints, durable dispatch, workspace scopes and
the HITL/DynamicReview approval chain (approval path is product-critical).
- Adopt upstream model registry (llm/registry.py): our 136 model entries
are a strict subset of upstream's 180, so dropping our inline table
loses nothing and gains 44 new models.
- Adopt upstream native EvoChatDeepSeek; drop our obsolete
_patch_deepseek_reasoning_passback monkey patch.
- Keep our six patches.py additions, ported onto upstream's new
_OpenAICompatContent class: stable tool-call ids, tool-history
sanitization, drop_reasoning_metadata, empty-SSE keepalive,
extracted-document-text patch, _has_assistant_tool_protocol.
- Keep our skill-budget middleware path (skills=None) instead of passing
skills through, to avoid double loading.
- Keep sanitized error labels (_safe_error_label) while adopting
upstream's injected MiddlewareEventSink for fallback narration.
- Keep port 3076 and the LANGGRAPH_SERVER_URL override; adopt upstream's
host/probe-host handling and CONFIG_DRIFT_SINCE_LAUNCH.
- Adopt upstream dependency stack: deepagents 0.7.6, langchain-quickjs
0.3.7, langgraph-api 0.14; keep our extra deps (rfc8785, pillow,
firecrawl-anydoc, nest-asyncio).
- Align call sites with upstream APIs: create_tool_selector_middleware
now takes events= instead of track_stream_selection=.
495 lines
16 KiB
Python
495 lines
16 KiB
Python
"""Middleware that implements model fallback on LLM call failures.
|
|
|
|
Uses LangChain's AgentMiddleware to intercept model calls. When the primary
|
|
model raises an exception, the middleware walks the configured fallback chain,
|
|
trying each alternative model in order. Every fallback attempt and its
|
|
outcome is reported to the injected event sink as fallback narration, and the
|
|
frontend sink renders it.
|
|
|
|
Errors that indicate a client-side bug (malformed request / HTTP 400) or a
|
|
context-length breach are not eligible for fallback and are re-raised
|
|
immediately so the correct handler (user or ContextOverflowMapperMiddleware)
|
|
can deal with them.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
import re
|
|
import threading
|
|
from collections.abc import Awaitable, Callable
|
|
from typing import TYPE_CHECKING
|
|
|
|
from langchain.agents.middleware.types import (
|
|
AgentMiddleware,
|
|
ModelRequest,
|
|
ModelResponse,
|
|
)
|
|
|
|
if TYPE_CHECKING:
|
|
from .events import MiddlewareEventSink
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
_fallback_chain_lock = threading.Lock()
|
|
_fallback_chain: list[tuple[str, str]] = []
|
|
"""Ordered list of ``(model_name, provider)`` fallback entries."""
|
|
|
|
_CONTEXT_LIMIT_PATTERNS: list[str] = [
|
|
"context_length_exceeded",
|
|
"context length exceeded",
|
|
"too many tokens",
|
|
"maximum context length",
|
|
"output too large",
|
|
"context_window_exceeded",
|
|
"string_too_long",
|
|
"max_tokens_exceeded",
|
|
]
|
|
"""Substrings that identify a context-length error in provider messages."""
|
|
|
|
_MALFORMED_REQUEST_PATTERNS: list[str] = [
|
|
"invalid_request_error",
|
|
"invalid request",
|
|
"malformed",
|
|
"repetitive tool calls",
|
|
"identical name and arguments",
|
|
]
|
|
"""Substrings that identify a malformed request (client-side bug)."""
|
|
|
|
_AUTH_ERROR_PATTERNS: list[str] = [
|
|
"invalid_api_key",
|
|
"authentication",
|
|
"permission",
|
|
]
|
|
"""Substrings that identify auth/permission errors.
|
|
|
|
These are intentionally *not* treated as non-fallbackable because a different
|
|
provider in the chain may have valid credentials."""
|
|
|
|
_SAFE_ERROR_CODE = re.compile(r"^[A-Z][A-Z0-9_]{1,63}$")
|
|
|
|
|
|
def _safe_error_label(exc: BaseException) -> str:
|
|
"""Describe an error without rendering a provider-controlled response body."""
|
|
label = type(exc).__name__
|
|
status = getattr(exc, "status_code", None)
|
|
if isinstance(status, int) and 100 <= status <= 599:
|
|
label = f"{label} status={status}"
|
|
code = getattr(exc, "code", None)
|
|
if isinstance(code, str) and _SAFE_ERROR_CODE.fullmatch(code):
|
|
label = f"{label} code={code}"
|
|
return label
|
|
|
|
|
|
def get_fallback_chain() -> list[tuple[str, str]]:
|
|
"""Return a snapshot of the current fallback chain.
|
|
|
|
Returns:
|
|
List of ``(model_name, provider)`` tuples in priority order.
|
|
"""
|
|
with _fallback_chain_lock:
|
|
return list(_fallback_chain)
|
|
|
|
|
|
def set_fallback_chain(chain: list[tuple[str, str]]) -> None:
|
|
"""Replace the entire fallback chain.
|
|
|
|
Args:
|
|
chain: New list of ``(model_name, provider)`` tuples.
|
|
"""
|
|
global _fallback_chain
|
|
with _fallback_chain_lock:
|
|
_fallback_chain = list(chain)
|
|
|
|
|
|
def add_fallback(model: str, provider: str) -> bool:
|
|
"""Append a model to the end of the fallback chain.
|
|
|
|
Args:
|
|
model: Short model name (e.g. ``"gpt-5.5"``).
|
|
provider: Provider identifier (e.g. ``"openai"``).
|
|
|
|
Returns:
|
|
``True`` if added, ``False`` if the entry was already present.
|
|
"""
|
|
entry = (model, provider)
|
|
with _fallback_chain_lock:
|
|
if entry in _fallback_chain:
|
|
return False
|
|
_fallback_chain.append(entry)
|
|
return True
|
|
|
|
|
|
def remove_fallback(model: str) -> bool:
|
|
"""Remove all entries matching *model* regardless of provider.
|
|
|
|
Args:
|
|
model: Short model name to remove.
|
|
|
|
Returns:
|
|
``True`` if at least one entry was removed.
|
|
"""
|
|
global _fallback_chain
|
|
with _fallback_chain_lock:
|
|
before = len(_fallback_chain)
|
|
_fallback_chain = [(m, p) for m, p in _fallback_chain if m != model]
|
|
return len(_fallback_chain) < before
|
|
|
|
|
|
def remove_fallback_at(index: int) -> tuple[str, str] | None:
|
|
"""Remove the entry at a 0-based index.
|
|
|
|
Args:
|
|
index: Position in the chain (0-based).
|
|
|
|
Returns:
|
|
The removed ``(model, provider)`` tuple, or ``None`` if out of range.
|
|
"""
|
|
with _fallback_chain_lock:
|
|
if 0 <= index < len(_fallback_chain):
|
|
return _fallback_chain.pop(index)
|
|
return None
|
|
|
|
|
|
def clear_fallbacks() -> None:
|
|
"""Remove every entry from the fallback chain."""
|
|
global _fallback_chain
|
|
with _fallback_chain_lock:
|
|
_fallback_chain = []
|
|
|
|
|
|
def serialize_fallback_chain() -> str:
|
|
"""Serialize the chain to a config-friendly string.
|
|
|
|
Returns:
|
|
Comma-separated ``"model:provider,model:provider"`` string.
|
|
"""
|
|
with _fallback_chain_lock:
|
|
return ",".join(f"{m}:{p}" for m, p in _fallback_chain)
|
|
|
|
|
|
def load_fallback_chain(raw: str) -> None:
|
|
"""Populate the chain from a serialized config string.
|
|
|
|
Args:
|
|
raw: Comma-separated ``"model:provider"`` pairs. Empty or
|
|
whitespace-only segments are silently skipped.
|
|
"""
|
|
global _fallback_chain
|
|
chain: list[tuple[str, str]] = []
|
|
for part in raw.split(","):
|
|
part = part.strip()
|
|
if not part:
|
|
continue
|
|
if ":" in part:
|
|
model, provider = part.rsplit(":", 1)
|
|
chain.append((model.strip(), provider.strip()))
|
|
with _fallback_chain_lock:
|
|
_fallback_chain = chain
|
|
|
|
|
|
def _is_non_fallbackable(exc: Exception) -> str | None:
|
|
"""Determine whether an exception should bypass the fallback chain.
|
|
|
|
Args:
|
|
exc: The exception raised by a model call.
|
|
|
|
Returns:
|
|
A human-readable reason string if the error must *not* trigger
|
|
fallback, or ``None`` if fallback should proceed.
|
|
"""
|
|
from langchain_core.exceptions import ContextOverflowError
|
|
|
|
if getattr(exc, "non_fallbackable", False):
|
|
return f"platform control error: {getattr(exc, 'code', type(exc).__name__)}"
|
|
|
|
if isinstance(exc, ContextOverflowError):
|
|
return "context length exceeded"
|
|
|
|
err_msg = str(exc).lower()
|
|
is_400 = "400" in err_msg or "bad request" in err_msg
|
|
|
|
if is_400 and any(p in err_msg for p in _CONTEXT_LIMIT_PATTERNS):
|
|
return "context length exceeded"
|
|
|
|
if is_400 and any(p in err_msg for p in _MALFORMED_REQUEST_PATTERNS):
|
|
return "malformed request (client-side error)"
|
|
|
|
return None
|
|
|
|
|
|
async def _try_fallbacks(
|
|
request: ModelRequest,
|
|
invoke: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
|
primary_exc: Exception,
|
|
events: MiddlewareEventSink,
|
|
) -> ModelResponse:
|
|
"""Walk the fallback chain, trying each model until one succeeds.
|
|
|
|
Shared implementation for both sync and async middleware entry points.
|
|
The *invoke* callable is an async function that calls the handler with
|
|
a given request — the sync path wraps the synchronous handler in a
|
|
trivial coroutine so both paths converge here.
|
|
|
|
Args:
|
|
request: The original model request.
|
|
invoke: Async callable that invokes the handler on a request.
|
|
primary_exc: The exception raised by the primary model.
|
|
events: Injected event sink for fallback narration.
|
|
|
|
Returns:
|
|
The ``ModelResponse`` from the first successful fallback.
|
|
|
|
Raises:
|
|
Exception: Re-raises the last exception if all fallbacks fail.
|
|
"""
|
|
from ..llm.models import get_chat_model
|
|
|
|
events.emit_fallback_notice(
|
|
f"Primary model failed: {_safe_error_label(primary_exc)}",
|
|
"yellow",
|
|
)
|
|
logger.warning("Primary model failed: %s", _safe_error_label(primary_exc))
|
|
|
|
# Track the request whose model actually raised ``last_exc`` so we
|
|
# can attribute the exception to the failing model, not the
|
|
# original ``request.model``. Without this, a fallback chain
|
|
# ``deepseek → moonshot`` where moonshot exhausts its quota would
|
|
# surface as ``provider: deepseek`` — the model the user never
|
|
# actually saw fail.
|
|
last_exc = primary_exc
|
|
last_failing_request = request
|
|
|
|
for model_name, provider in get_fallback_chain():
|
|
events.emit_fallback_notice(
|
|
f" -> Falling back to {model_name} ({provider}) "
|
|
f"due to: {_safe_error_label(last_exc)}",
|
|
"yellow",
|
|
)
|
|
try:
|
|
fallback_model = get_chat_model(model=model_name, provider=provider)
|
|
fb_request = request.override(model=fallback_model)
|
|
result = await invoke(fb_request)
|
|
events.emit_fallback_notice(
|
|
f" Fallback to {model_name} ({provider}) succeeded",
|
|
"green",
|
|
)
|
|
logger.info("Fallback to %s (%s) succeeded", model_name, provider)
|
|
return result
|
|
except Exception as fb_exc:
|
|
reason = _is_non_fallbackable(fb_exc)
|
|
if reason is not None:
|
|
events.emit_fallback_notice(
|
|
f" {model_name} hit non-fallbackable error ({reason}) "
|
|
f"-- aborting fallback chain",
|
|
"red",
|
|
)
|
|
_raise_normalized(fb_request, fb_exc)
|
|
last_exc = fb_exc
|
|
last_failing_request = fb_request
|
|
events.emit_fallback_notice(
|
|
f" x {model_name} also failed: {_safe_error_label(fb_exc)}",
|
|
"red",
|
|
)
|
|
logger.warning(
|
|
"Fallback %s (provider=%s) failed: %s",
|
|
model_name,
|
|
provider,
|
|
_safe_error_label(fb_exc),
|
|
)
|
|
|
|
events.emit_fallback_notice(
|
|
" All fallbacks exhausted -- re-raising last error", "red"
|
|
)
|
|
_raise_normalized(last_failing_request, last_exc)
|
|
|
|
|
|
def _try_fallbacks_sync(
|
|
request: ModelRequest,
|
|
invoke: Callable[[ModelRequest], ModelResponse],
|
|
primary_exc: Exception,
|
|
events: MiddlewareEventSink,
|
|
) -> ModelResponse:
|
|
"""Synchronous counterpart to :func:`_try_fallbacks`.
|
|
|
|
The synchronous middleware path calls a synchronous model handler. Keeping
|
|
that traversal synchronous avoids manufacturing an event loop solely to
|
|
share the async implementation.
|
|
"""
|
|
from ..llm.models import get_chat_model
|
|
|
|
events.emit_fallback_notice(
|
|
f"Primary model failed: {type(primary_exc).__name__}: {primary_exc}",
|
|
"yellow",
|
|
)
|
|
logger.warning(
|
|
"Primary model failed: %s: %s", type(primary_exc).__name__, primary_exc
|
|
)
|
|
|
|
last_exc = primary_exc
|
|
last_failing_request = request
|
|
|
|
for model_name, provider in get_fallback_chain():
|
|
events.emit_fallback_notice(
|
|
f" -> Falling back to {model_name} ({provider}) due to: "
|
|
f"{type(last_exc).__name__}: {last_exc}",
|
|
"yellow",
|
|
)
|
|
try:
|
|
fallback_model = get_chat_model(model=model_name, provider=provider)
|
|
fb_request = request.override(model=fallback_model)
|
|
result = invoke(fb_request)
|
|
events.emit_fallback_notice(
|
|
f" Fallback to {model_name} ({provider}) succeeded",
|
|
"green",
|
|
)
|
|
logger.info("Fallback to %s (%s) succeeded", model_name, provider)
|
|
return result
|
|
except Exception as fb_exc:
|
|
reason = _is_non_fallbackable(fb_exc)
|
|
if reason is not None:
|
|
events.emit_fallback_notice(
|
|
f" {model_name} hit non-fallbackable error ({reason}) "
|
|
f"-- aborting fallback chain",
|
|
"red",
|
|
)
|
|
_raise_normalized(fb_request, fb_exc)
|
|
last_exc = fb_exc
|
|
last_failing_request = fb_request
|
|
events.emit_fallback_notice(
|
|
f" x {model_name} also failed: {type(fb_exc).__name__}: {fb_exc}",
|
|
"red",
|
|
)
|
|
logger.warning(
|
|
"Fallback %s (provider=%s) failed: %s: %s",
|
|
model_name,
|
|
provider,
|
|
type(fb_exc).__name__,
|
|
fb_exc,
|
|
)
|
|
|
|
events.emit_fallback_notice(
|
|
" All fallbacks exhausted -- re-raising last error", "red"
|
|
)
|
|
_raise_normalized(last_failing_request, last_exc)
|
|
|
|
|
|
def _raise_normalized(request: ModelRequest, exc: Exception) -> None:
|
|
"""Wrap *exc* in a ``ProviderStreamError`` attributed to
|
|
``request.model`` and raise, so the outer chain sees the failure
|
|
tagged with the model that actually raised.
|
|
|
|
Falls back to a plain ``raise`` when the model isn't from a
|
|
recognized provider (``_normalize`` returns None) — nothing useful
|
|
to add.
|
|
"""
|
|
from .error_normalization import _normalize
|
|
|
|
normalized = _normalize(request, exc)
|
|
if normalized is not None:
|
|
raise normalized from exc
|
|
raise exc
|
|
|
|
|
|
def _guard_and_fallback(
|
|
primary_exc: Exception,
|
|
request: ModelRequest,
|
|
invoke: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
|
events: MiddlewareEventSink,
|
|
) -> Awaitable[ModelResponse]:
|
|
"""Check non-fallbackable conditions, then delegate to ``_try_fallbacks``.
|
|
|
|
Args:
|
|
primary_exc: The exception raised by the primary model.
|
|
request: The original model request.
|
|
invoke: Async callable that invokes the handler on a request.
|
|
events: Injected event sink for fallback narration.
|
|
|
|
Returns:
|
|
Coroutine that resolves to the fallback ``ModelResponse``.
|
|
|
|
Raises:
|
|
Exception: Re-raises immediately for non-fallbackable errors.
|
|
"""
|
|
reason = _is_non_fallbackable(primary_exc)
|
|
if reason is not None:
|
|
events.emit_fallback_notice(
|
|
f"Model error ({reason}) -- not eligible for fallback, re-raising",
|
|
"red",
|
|
)
|
|
_raise_normalized(request, primary_exc)
|
|
return _try_fallbacks(request, invoke, primary_exc, events)
|
|
|
|
|
|
def _guard_and_fallback_sync(
|
|
primary_exc: Exception,
|
|
request: ModelRequest,
|
|
invoke: Callable[[ModelRequest], ModelResponse],
|
|
events: MiddlewareEventSink,
|
|
) -> ModelResponse:
|
|
"""Validate and run the native synchronous fallback traversal."""
|
|
reason = _is_non_fallbackable(primary_exc)
|
|
if reason is not None:
|
|
events.emit_fallback_notice(
|
|
f"Model error ({reason}) -- not eligible for fallback, re-raising",
|
|
"red",
|
|
)
|
|
_raise_normalized(request, primary_exc)
|
|
return _try_fallbacks_sync(request, invoke, primary_exc, events)
|
|
|
|
|
|
class ModelFallbackMiddleware(AgentMiddleware):
|
|
"""LangChain AgentMiddleware that retries failed model calls on fallbacks.
|
|
|
|
On each invocation the middleware reads the module-level
|
|
``_fallback_chain`` so that ``/model-fallback add`` takes effect
|
|
immediately without rebuilding the agent.
|
|
|
|
Attributes:
|
|
name: Middleware identifier used by the framework.
|
|
"""
|
|
|
|
name = "model_fallback"
|
|
|
|
def __init__(self, events: MiddlewareEventSink | None = None) -> None:
|
|
super().__init__()
|
|
from .events import NO_OP_SINK
|
|
|
|
self._events = events or NO_OP_SINK
|
|
|
|
def wrap_model_call(
|
|
self,
|
|
request: ModelRequest,
|
|
handler: Callable[[ModelRequest], ModelResponse],
|
|
) -> ModelResponse:
|
|
if not _fallback_chain:
|
|
return handler(request)
|
|
|
|
from .error_normalization import _check_truncated_output
|
|
|
|
def invoke(current_request: ModelRequest) -> ModelResponse:
|
|
return _check_truncated_output(handler(current_request))
|
|
|
|
try:
|
|
return invoke(request)
|
|
except Exception as exc:
|
|
return _guard_and_fallback_sync(exc, request, invoke, self._events)
|
|
|
|
async def awrap_model_call(
|
|
self,
|
|
request: ModelRequest,
|
|
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
|
) -> ModelResponse:
|
|
if not _fallback_chain:
|
|
return await handler(request)
|
|
|
|
from .error_normalization import _check_truncated_output
|
|
|
|
async def invoke(current_request: ModelRequest) -> ModelResponse:
|
|
return _check_truncated_output(await handler(current_request))
|
|
|
|
try:
|
|
return await invoke(request)
|
|
except Exception as exc:
|
|
return await _guard_and_fallback(exc, request, invoke, self._events)
|