From 9e51ec6fddc29ff6d233158cd9e6bf2b0ab67d26 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Wiktor=20Cupia=C5=82?= <98488731+mooshroom4422@users.noreply.github.com> Date: Thu, 7 May 2026 15:35:26 +0200 Subject: [PATCH] feat(cmd): add /model-fallback command (#196) * feat(cmd): add /model-fallback command * fix: apply feedback * fix: lock usage with _fallback_chain * fix: apply feedback * fix: apply feedback * feat: add tests * fix: tests * Update EvoScientist/middleware/model_fallback.py Co-authored-by: dinos --------- Co-authored-by: dinos --- EvoScientist/EvoScientist.py | 10 + EvoScientist/cli/tui_interactive.py | 11 +- .../commands/implementation/__init__.py | 4 +- .../commands/implementation/model_fallback.py | 310 ++++++++++++++ EvoScientist/config/settings.py | 2 + EvoScientist/middleware/__init__.py | 3 + EvoScientist/middleware/model_fallback.py | 378 ++++++++++++++++++ tests/test_model_fallback.py | 327 +++++++++++++++ 8 files changed, 1042 insertions(+), 3 deletions(-) create mode 100644 EvoScientist/commands/implementation/model_fallback.py create mode 100644 EvoScientist/middleware/model_fallback.py create mode 100644 tests/test_model_fallback.py diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index 2151bec..fece8b1 100644 --- a/EvoScientist/EvoScientist.py +++ b/EvoScientist/EvoScientist.py @@ -432,17 +432,22 @@ def _get_default_middleware(): """Build the default middleware list.""" from .middleware import ( ContextOverflowMapperMiddleware, + ModelFallbackMiddleware, ToolErrorHandlerMiddleware, create_context_editing_middleware, create_memory_middleware, create_tool_selector_middleware, + load_fallback_chain, ) cfg = _ensure_config() + if cfg.model_fallbacks: + load_fallback_chain(cfg.model_fallbacks) model = _ensure_chat_model() memory_dir = str(_paths_mod.MEMORIES_DIR) mw = [ create_context_editing_middleware(model), + ModelFallbackMiddleware(), ContextOverflowMapperMiddleware(), ToolErrorHandlerMiddleware(), *create_tool_selector_middleware(model=model), @@ -543,13 +548,17 @@ def create_cli_agent( from .backends import CustomSandboxBackend, MergedSkillsBackend from .middleware import ( ContextOverflowMapperMiddleware, + ModelFallbackMiddleware, ToolErrorHandlerMiddleware, create_context_editing_middleware, create_memory_middleware, create_tool_selector_middleware, + load_fallback_chain, ) cfg = _ensure_config(config) + if cfg.model_fallbacks: + load_fallback_chain(cfg.model_fallbacks) if checkpointer is None: from langgraph.checkpoint.memory import InMemorySaver @@ -600,6 +609,7 @@ def create_cli_agent( model = _ensure_chat_model() mw: list[AgentMiddleware] = [ create_context_editing_middleware(model), + ModelFallbackMiddleware(), ContextOverflowMapperMiddleware(), ToolErrorHandlerMiddleware(), *create_tool_selector_middleware(model=model), diff --git a/EvoScientist/cli/tui_interactive.py b/EvoScientist/cli/tui_interactive.py index 475a8a6..158fe97 100644 --- a/EvoScientist/cli/tui_interactive.py +++ b/EvoScientist/cli/tui_interactive.py @@ -681,6 +681,12 @@ def run_textual_interactive( yield Static("", id="status") def on_mount(self) -> None: + # Register fallback middleware UI callback so messages appear + # as SystemMessage widgets in the chat container. + from ..middleware.model_fallback import set_ui_emit + + set_ui_emit(lambda text, style: self._append_system(text, style)) + self._render_welcome() self._render_status() self.set_interval(1.0, self._render_status) @@ -2533,7 +2539,10 @@ def run_textual_interactive( self._do_exit() def _do_exit(self) -> None: - """Clean up channels and exit.""" + """Clean up channels, unregister callbacks, and exit.""" + from ..middleware.model_fallback import set_ui_emit + + set_ui_emit(None) if self._channel_timer is not None: self._channel_timer.stop() self._channel_timer = None diff --git a/EvoScientist/commands/implementation/__init__.py b/EvoScientist/commands/implementation/__init__.py index 24e3cab..1333598 100644 --- a/EvoScientist/commands/implementation/__init__.py +++ b/EvoScientist/commands/implementation/__init__.py @@ -1,5 +1,5 @@ from __future__ import annotations -from . import channel, general, mcp, model, session, skills +from . import channel, general, mcp, model, model_fallback, session, skills -__all__ = ["channel", "general", "mcp", "model", "session", "skills"] +__all__ = ["channel", "general", "mcp", "model", "model_fallback", "session", "skills"] diff --git a/EvoScientist/commands/implementation/model_fallback.py b/EvoScientist/commands/implementation/model_fallback.py new file mode 100644 index 0000000..718dff5 --- /dev/null +++ b/EvoScientist/commands/implementation/model_fallback.py @@ -0,0 +1,310 @@ +"""Slash command for managing the model fallback chain. + +Provides ``/model-fallback`` (alias ``/fallback``) with subcommands to +add, remove, list, clear, save, and display help for fallback models. +""" + +from __future__ import annotations + +from typing import ClassVar + +from ..base import Argument, Command, CommandContext +from ..manager import manager + + +class ModelFallbackCommand(Command): + """Manage the model fallback chain. + + Subcommands: + add -- Append a model (interactive picker when no args). + remove -- Remove by position (interactive picker in TUI). + list -- Display the current chain. + clear -- Remove all entries. + save -- Persist the chain to the config file. + help -- Show subcommand reference. + """ + + name = "/model-fallback" + alias: ClassVar[list[str]] = ["/fallback"] + description = "Manage fallback models (add/remove/list/clear)" + arguments: ClassVar[list[Argument]] = [ + Argument( + name="action", + type=str, + description="add|remove|list|clear|save|help", + required=False, + ), + ] + + async def execute(self, ctx: CommandContext, args: list[str]) -> None: + from ...llm.models import MODELS + from ...middleware.model_fallback import ( + add_fallback, + clear_fallbacks, + get_fallback_chain, + remove_fallback_at, + serialize_fallback_chain, + ) + + save = "--save" in args + args = [a for a in args if a != "--save"] + + if not args: + await self._show_list(ctx, get_fallback_chain()) + return + + action = args[0].lower() + + if action == "list": + await self._show_list(ctx, get_fallback_chain()) + + elif action == "add": + if len(args) >= 2: + model_name = args[1] + provider = args[2] if len(args) > 2 else None + + if provider is None: + if model_name in MODELS: + _, provider = MODELS[model_name] + else: + ctx.ui.append_system( + f"Unknown model '{model_name}'. Specify provider explicitly: " + f"/model-fallback add {model_name} ", + style="red", + ) + return + else: + picked = await self._pick_model(ctx) + if picked is None: + return + model_name, provider = picked + + if add_fallback(model_name, provider): + ctx.ui.append_system( + f"Added {model_name} ({provider}) to fallback chain", style="green" + ) + else: + ctx.ui.append_system( + f"{model_name} ({provider}) is already in the fallback chain", + style="yellow", + ) + return + + if save: + self._save_to_config(serialize_fallback_chain()) + + elif action == "remove": + chain = get_fallback_chain() + if not chain: + ctx.ui.append_system("Fallback chain is empty", style="yellow") + return + + if len(args) >= 2: + arg = args[1] + try: + idx = int(arg) - 1 + except ValueError: + ctx.ui.append_system( + f"Expected a position number (1-{len(chain)}), got '{arg}'. " + "Use /model-fallback list to see positions.", + style="red", + ) + return + removed = remove_fallback_at(idx) + if removed is None: + ctx.ui.append_system( + f"Invalid position {arg}. " + f"Use a number between 1 and {len(chain)}.", + style="red", + ) + return + model_name, provider = removed + else: + picked = await self._pick_fallback_to_remove(ctx, chain) + if picked is None: + return + model_name, provider = picked + live_chain = get_fallback_chain() + try: + idx = live_chain.index((model_name, provider)) + except ValueError: + ctx.ui.append_system( + f"{model_name} ({provider}) is no longer in the fallback chain", + style="yellow", + ) + return + remove_fallback_at(idx) + + ctx.ui.append_system( + f"Removed {model_name} ({provider}) from fallback chain", + style="green", + ) + + if save: + self._save_to_config(serialize_fallback_chain()) + + elif action == "clear": + clear_fallbacks() + ctx.ui.append_system("Cleared all fallback models", style="green") + + if save: + self._save_to_config("") + + elif action == "save": + self._save_to_config(serialize_fallback_chain()) + ctx.ui.append_system("Fallback chain saved to config", style="green") + + elif action == "help": + self._show_help(ctx) + + else: + self._show_help(ctx) + + async def _pick_model(self, ctx: CommandContext) -> tuple[str, str] | None: + """Open the interactive model picker to select a fallback model. + + Falls back to a usage hint when the UI does not support interactive + widgets (CLI mode without a model argument). + + Args: + ctx: Current command context. + + Returns: + ``(model_name, provider)`` tuple, or ``None`` if cancelled. + """ + if not ctx.ui.supports_interactive: + ctx.ui.append_system( + "Usage: /model-fallback add [provider]", style="yellow" + ) + return None + + from ...EvoScientist import _ensure_config + from ...llm.models import list_models_by_provider + + cfg = _ensure_config() + entries = list_models_by_provider() + + ollama_base_url = getattr(cfg, "ollama_base_url", None) + if ollama_base_url: + from ...llm.ollama_discovery import discover_ollama_models + + detected = await discover_ollama_models(ollama_base_url, timeout=1.5) + for detected_name in detected: + entries.append((detected_name, detected_name, "ollama")) + entries.append(("Custom Ollama model...", "__custom_ollama__", "ollama")) + + result = await ctx.ui.wait_for_model_pick( + entries, + current_model=cfg.model, + current_provider=cfg.provider, + ) + if result is None: + return None + + name, provider = result + if provider == "ollama" and name in ( + "Custom Ollama model...", + "__custom_ollama__", + ): + return None + return name, provider + + async def _pick_fallback_to_remove( + self, ctx: CommandContext, chain: list[tuple[str, str]] + ) -> tuple[str, str] | None: + """Open the model picker populated with the current fallback chain. + + Falls back to a usage hint in CLI mode. + + Args: + ctx: Current command context. + chain: The current fallback chain to choose from. + + Returns: + ``(model_name, provider)`` tuple, or ``None`` if cancelled. + """ + if not ctx.ui.supports_interactive: + ctx.ui.append_system( + "Usage: /model-fallback remove " + "(use /model-fallback list to see positions)", + style="yellow", + ) + return None + + entries = [(m, m, p) for m, p in chain] + result = await ctx.ui.wait_for_model_pick( + entries, current_model=None, current_provider=None + ) + if result is None: + return None + return result + + def _show_help(self, ctx: CommandContext) -> None: + """Render the subcommand reference table. + + Args: + ctx: Current command context. + """ + from rich.text import Text + + text = Text("/model-fallback subcommands:\n", style="bold") + for cmd, desc in ( + ( + "add [model] [provider]", + "Add a fallback model (opens picker if omitted)", + ), + ( + "remove [position]", + "Remove a fallback by position (opens picker in TUI)", + ), + ("list", "Show the current fallback chain"), + ("clear", "Remove all fallback models"), + ("save", "Save current fallback chain to config file"), + ("help", "Show this help message"), + ): + text.append(f" {cmd:<26}", style="cyan") + text.append(f"{desc}\n", style="dim") + text.append( + "\nAdd --save to add/remove/clear to persist the change immediately.\n", + style="dim", + ) + ctx.ui.mount_renderable(text) + + async def _show_list( + self, ctx: CommandContext, chain: list[tuple[str, str]] + ) -> None: + """Display the current fallback chain as a numbered list. + + Args: + ctx: Current command context. + chain: The fallback chain to display. + """ + if not chain: + ctx.ui.append_system("No fallback models configured", style="dim") + ctx.ui.append_system( + "Use /model-fallback add [provider] to add one", + style="dim", + ) + return + + from rich.text import Text + + text = Text("Fallback chain:\n", style="bold") + for idx, (model, provider) in enumerate(chain, 1): + text.append(f" {idx}. ", style="dim") + text.append(model, style="cyan") + text.append(f" ({provider})\n", style="dim") + ctx.ui.mount_renderable(text) + + def _save_to_config(self, value: str) -> None: + """Persist the fallback chain string to the config file. + + Args: + value: Serialized chain (``"model:provider,..."``). + """ + from ...config.settings import set_config_value + + set_config_value("model_fallbacks", value) + + +manager.register(ModelFallbackCommand()) diff --git a/EvoScientist/config/settings.py b/EvoScientist/config/settings.py index c9acb38..daf10f4 100644 --- a/EvoScientist/config/settings.py +++ b/EvoScientist/config/settings.py @@ -86,6 +86,7 @@ class EvoScientistConfig: # LLM Settings provider: str = "anthropic" model: str = "claude-sonnet-4-5" + model_fallbacks: str = "" # "model:provider,model:provider" fallback chain # Async Sub-agent Settings # When True (default), the EvoSci CLI auto-starts a langgraph dev subprocess @@ -436,6 +437,7 @@ _ENV_MAPPINGS = { "default_workdir": "EVOSCIENTIST_WORKSPACE_DIR", "ui_backend": "EVOSCIENTIST_UI_BACKEND", "log_level": "EVOSCIENTIST_LOG_LEVEL", + "model_fallbacks": "EVOSCIENTIST_MODEL_FALLBACKS", "reasoning_effort": "EVOSCIENTIST_REASONING_EFFORT", "channel_debug_tracing": "EVOSCIENTIST_CHANNEL_DEBUG_TRACING", "ccproxy_port": "EVOSCIENTIST_CCPROXY_PORT", diff --git a/EvoScientist/middleware/__init__.py b/EvoScientist/middleware/__init__.py index 7123714..ea20d1e 100644 --- a/EvoScientist/middleware/__init__.py +++ b/EvoScientist/middleware/__init__.py @@ -22,6 +22,7 @@ from .memory import ( ExtractedMemory, create_memory_middleware, ) +from .model_fallback import ModelFallbackMiddleware, load_fallback_chain from .tool_error_handler import ToolErrorHandlerMiddleware from .tool_selector import create_tool_selector_middleware from .utils import disable_thinking @@ -35,6 +36,7 @@ __all__ = [ "EvoMemoryMiddleware", "EvoMemoryState", "ExtractedMemory", + "ModelFallbackMiddleware", "Question", "ToolErrorHandlerMiddleware", "compute_context_editing_trigger", @@ -42,4 +44,5 @@ __all__ = [ "create_memory_middleware", "create_tool_selector_middleware", "disable_thinking", + "load_fallback_chain", ] diff --git a/EvoScientist/middleware/model_fallback.py b/EvoScientist/middleware/model_fallback.py new file mode 100644 index 0000000..0dee6ff --- /dev/null +++ b/EvoScientist/middleware/model_fallback.py @@ -0,0 +1,378 @@ +"""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 surfaced to the user via the registered UI callback. + +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 threading +from collections.abc import Awaitable, Callable + +from langchain.agents.middleware.types import ( + AgentMiddleware, + ModelRequest, + ModelResponse, +) + +logger = logging.getLogger(__name__) + +_ui_emit_fn: Callable[[str, str], None] | None = None +"""UI callback registered by the CLI/TUI entrypoint. ``None`` until set.""" + +_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", +] +"""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.""" + + +def set_ui_emit(fn: Callable[[str, str], None] | None) -> None: + """Register (or clear) the UI callback for fallback status messages. + + Args: + fn: Callable with signature ``fn(text, style)`` where *style* is a + Rich style string (``"yellow"``, ``"red"``, ``"green"``). + Pass ``None`` to unregister. + """ + global _ui_emit_fn + _ui_emit_fn = fn + + +def _emit(text: str, style: str = "yellow") -> None: + """Surface a fallback status message to the user. + + Dispatches to the registered UI callback when available (TUI mode), + otherwise falls back to the shared Rich console on stdout (CLI mode). + + Args: + text: Plain-text message to display. + style: Rich style string applied to the message. + """ + if _ui_emit_fn is not None: + try: + _ui_emit_fn(text, style) + return + except Exception: + pass + + from ..stream.console import console + + console.print(text, style=style) + + +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 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, +) -> 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. + + 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 + + _emit( + f"Primary model failed: {type(primary_exc).__name__}: {primary_exc}", + style="yellow", + ) + logger.warning( + "Primary model failed: %s: %s", type(primary_exc).__name__, primary_exc + ) + + last_exc = primary_exc + for model_name, provider in get_fallback_chain(): + _emit( + f" -> Falling back to {model_name} ({provider}) " + f"due to: {type(last_exc).__name__}: {last_exc}", + style="yellow", + ) + try: + fallback_model = get_chat_model(model=model_name, provider=provider) + fb_request = request.override(model=fallback_model) + result = await invoke(fb_request) + _emit( + f" Fallback to {model_name} ({provider}) succeeded", + style="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: + _emit( + f" {model_name} hit non-fallbackable error ({reason}) " + f"-- aborting fallback chain", + style="red", + ) + raise + last_exc = fb_exc + _emit( + f" x {model_name} also failed: {type(fb_exc).__name__}: {fb_exc}", + style="red", + ) + logger.warning( + "Fallback %s (provider=%s) failed: %s: %s", + model_name, + provider, + type(fb_exc).__name__, + fb_exc, + ) + + _emit(" All fallbacks exhausted -- re-raising last error", style="red") + raise last_exc + + +def _guard_and_fallback( + primary_exc: Exception, + request: ModelRequest, + invoke: Callable[[ModelRequest], Awaitable[ModelResponse]], +) -> 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. + + 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: + _emit( + f"Model error ({reason}) -- not eligible for fallback, re-raising", + style="red", + ) + raise primary_exc + return _try_fallbacks(request, invoke, primary_exc) + + +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 wrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], ModelResponse], + ) -> ModelResponse: + if not _fallback_chain: + return handler(request) + try: + return handler(request) + except Exception as exc: + + async def _sync_invoke(r: ModelRequest) -> ModelResponse: + return handler(r) + + import asyncio + + return asyncio.run(_guard_and_fallback(exc, request, _sync_invoke)) + + async def awrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], Awaitable[ModelResponse]], + ) -> ModelResponse: + if not _fallback_chain: + return await handler(request) + try: + return await handler(request) + except Exception as exc: + return await _guard_and_fallback(exc, request, handler) diff --git a/tests/test_model_fallback.py b/tests/test_model_fallback.py new file mode 100644 index 0000000..19f8f31 --- /dev/null +++ b/tests/test_model_fallback.py @@ -0,0 +1,327 @@ +"""Tests for the model fallback middleware. + +Covers error classification (_is_non_fallbackable) and the end-to-end +fallback chain behaviour via _try_fallbacks / _guard_and_fallback. +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from langchain_core.exceptions import ContextOverflowError +from langchain_core.messages import AIMessage, HumanMessage + +from EvoScientist.middleware.model_fallback import ( + _guard_and_fallback, + _is_non_fallbackable, + _try_fallbacks, + add_fallback, + clear_fallbacks, + set_ui_emit, +) +from tests.conftest import run_async as _run + +# ── Helpers ────────────────────────────────────────────────────── + + +def _fake_request(): + """Build a minimal ModelRequest stub with an .override() method.""" + req = MagicMock() + req.override = MagicMock(side_effect=lambda **kw: req) + req.messages = [HumanMessage(content="hi")] + return req + + +AI_RESPONSE = AIMessage(content="ok") + + +@pytest.fixture(autouse=True) +def _clean_chain(): + """Ensure a clean fallback chain and no UI callback for every test.""" + clear_fallbacks() + set_ui_emit(None) + yield + clear_fallbacks() + set_ui_emit(None) + + +# ═════════════════════════════════════════════════════════════════ +# 1. _is_non_fallbackable — error classification +# ═════════════════════════════════════════════════════════════════ + + +class TestIsNonFallbackable: + """Verify which errors block fallback and which allow it.""" + + # ── Context-length errors: must NOT fallback ──────────────── + + def test_context_overflow_error_instance(self): + exc = ContextOverflowError("too long") + assert _is_non_fallbackable(exc) == "context length exceeded" + + @pytest.mark.parametrize( + "msg", + [ + "Error 400: context_length_exceeded", + "Bad Request: context length exceeded in prompt", + "400 too many tokens for this model", + "Bad Request: maximum context length is 128k", + "Error 400: output too large", + "400 Bad Request: context_window_exceeded", + "400: string_too_long", + "Bad Request: max_tokens_exceeded", + ], + ) + def test_context_limit_400_patterns(self, msg): + assert _is_non_fallbackable(Exception(msg)) == "context length exceeded" + + # ── Malformed request errors: must NOT fallback ───────────── + + @pytest.mark.parametrize( + "msg", + [ + "Error 400: invalid_request_error", + "400 Bad Request: invalid request body", + "400: malformed JSON in request", + ], + ) + def test_malformed_request_400_patterns(self, msg): + assert ( + _is_non_fallbackable(Exception(msg)) + == "malformed request (client-side error)" + ) + + # ── Auth errors: SHOULD fallback (different provider may work) ── + + @pytest.mark.parametrize( + "msg", + [ + "400 Bad Request: invalid_api_key", + "400: authentication failed", + "400 Bad Request: permission denied", + ], + ) + def test_auth_errors_are_fallbackable(self, msg): + assert _is_non_fallbackable(Exception(msg)) is None + + # ── Server / transient errors: SHOULD fallback ────────────── + + @pytest.mark.parametrize( + "msg", + [ + "Error 500: internal server error", + "429 Too Many Requests: rate limit exceeded", + "503 Service Unavailable", + "Connection timed out", + "HTTPSConnectionPool: Read timed out", + "502 Bad Gateway", + "overloaded_error: the server is temporarily overloaded", + ], + ) + def test_server_errors_are_fallbackable(self, msg): + assert _is_non_fallbackable(Exception(msg)) is None + + # ── Edge: 400 without a known pattern → fallbackable ──────── + + def test_400_unknown_pattern_is_fallbackable(self): + assert _is_non_fallbackable(Exception("400: unknown_field 'foo'")) is None + + # ── Edge: pattern present but no 400 → fallbackable ───────── + + def test_context_pattern_without_400_is_fallbackable(self): + exc = Exception("context_length_exceeded (warning only)") + assert _is_non_fallbackable(exc) is None + + def test_malformed_pattern_without_400_is_fallbackable(self): + exc = Exception("invalid_request_error logged for debugging") + assert _is_non_fallbackable(exc) is None + + +# ═════════════════════════════════════════════════════════════════ +# 2. _try_fallbacks — chain walk behaviour +# ═════════════════════════════════════════════════════════════════ + + +class TestTryFallbacks: + """End-to-end tests for the fallback chain traversal.""" + + def test_first_fallback_succeeds(self): + """When the first fallback model works, return its response.""" + add_fallback("fb-model", "fb-provider") + req = _fake_request() + invoke = AsyncMock(return_value=AI_RESPONSE) + + with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: + mock_gcm.return_value = MagicMock() + result = _run(_try_fallbacks(req, invoke, Exception("503 boom"))) + + assert result is AI_RESPONSE + invoke.assert_awaited_once() + mock_gcm.assert_called_once_with(model="fb-model", provider="fb-provider") + + def test_skips_failing_fallback_tries_next(self): + """When the first fallback fails, try the second.""" + add_fallback("fb-bad", "prov-a") + add_fallback("fb-good", "prov-b") + req = _fake_request() + + call_count = 0 + + async def _invoke(r): + nonlocal call_count + call_count += 1 + if call_count == 1: + raise Exception("429 rate limited") + return AI_RESPONSE + + with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: + mock_gcm.return_value = MagicMock() + result = _run(_try_fallbacks(req, _invoke, Exception("503 boom"))) + + assert result is AI_RESPONSE + assert call_count == 2 + + def test_all_fallbacks_exhausted_raises_last(self): + """When every fallback fails, re-raise the last exception.""" + add_fallback("fb-a", "prov-a") + add_fallback("fb-b", "prov-b") + req = _fake_request() + + last_error = Exception("429 from fb-b") + + call_count = 0 + + async def _invoke(r): + nonlocal call_count + call_count += 1 + if call_count == 1: + raise Exception("500 from fb-a") + raise last_error + + with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: + mock_gcm.return_value = MagicMock() + with pytest.raises(Exception, match="429 from fb-b") as exc_info: + _run(_try_fallbacks(req, _invoke, Exception("503 primary"))) + + assert exc_info.value is last_error + + def test_non_fallbackable_in_chain_aborts_immediately(self): + """A non-fallbackable error from a fallback model aborts the chain.""" + add_fallback("fb-a", "prov-a") + add_fallback("fb-b", "prov-b") # should never be reached + req = _fake_request() + + async def _invoke(r): + raise Exception("400: context_length_exceeded") + + with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: + mock_gcm.return_value = MagicMock() + with pytest.raises(Exception, match="context_length_exceeded"): + _run(_try_fallbacks(req, _invoke, Exception("503 primary"))) + + # get_chat_model should only have been called once (for fb-a), + # fb-b should never be reached. + assert mock_gcm.call_count == 1 + + +# ═════════════════════════════════════════════════════════════════ +# 3. _guard_and_fallback — pre-check before chain walk +# ═════════════════════════════════════════════════════════════════ + + +class TestGuardAndFallback: + """Verify that non-fallbackable errors are re-raised before trying the chain.""" + + def test_context_overflow_raises_immediately(self): + add_fallback("fb", "prov") + req = _fake_request() + invoke = AsyncMock() + + with pytest.raises(ContextOverflowError): + _run(_guard_and_fallback(ContextOverflowError("overflow"), req, invoke)) + + invoke.assert_not_awaited() + + def test_malformed_400_raises_immediately(self): + add_fallback("fb", "prov") + req = _fake_request() + invoke = AsyncMock() + + with pytest.raises(Exception, match="invalid_request_error"): + _run( + _guard_and_fallback( + Exception("400: invalid_request_error"), req, invoke + ) + ) + + invoke.assert_not_awaited() + + def test_server_error_proceeds_to_fallback(self): + add_fallback("fb", "prov") + req = _fake_request() + invoke = AsyncMock(return_value=AI_RESPONSE) + + with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: + mock_gcm.return_value = MagicMock() + result = _run(_guard_and_fallback(Exception("503 overloaded"), req, invoke)) + + assert result is AI_RESPONSE + invoke.assert_awaited_once() + + def test_auth_error_proceeds_to_fallback(self): + """Auth errors should try the fallback chain (different provider).""" + add_fallback("fb", "other-prov") + req = _fake_request() + invoke = AsyncMock(return_value=AI_RESPONSE) + + with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: + mock_gcm.return_value = MagicMock() + result = _run( + _guard_and_fallback( + Exception("400 Bad Request: invalid_api_key"), req, invoke + ) + ) + + assert result is AI_RESPONSE + invoke.assert_awaited_once() + + +# ═════════════════════════════════════════════════════════════════ +# 4. UI emit callback +# ═════════════════════════════════════════════════════════════════ + + +class TestUiEmit: + """Verify that fallback events are surfaced via the registered callback.""" + + def test_emit_captures_messages(self): + add_fallback("fb", "prov") + req = _fake_request() + invoke = AsyncMock(return_value=AI_RESPONSE) + + messages: list[tuple[str, str]] = [] + set_ui_emit(lambda text, style: messages.append((text, style))) + + with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: + mock_gcm.return_value = MagicMock() + _run(_try_fallbacks(req, invoke, Exception("503 down"))) + + texts = [t for t, _ in messages] + assert any("Primary model failed" in t for t in texts) + assert any("Falling back to fb (prov)" in t for t in texts) + assert any("succeeded" in t for t in texts) + + def test_emit_shows_non_fallbackable_rejection(self): + add_fallback("fb", "prov") + req = _fake_request() + invoke = AsyncMock() + + messages: list[tuple[str, str]] = [] + set_ui_emit(lambda text, style: messages.append((text, style))) + + with pytest.raises(ContextOverflowError): + _run(_guard_and_fallback(ContextOverflowError("overflow"), req, invoke)) + + texts = [t for t, _ in messages] + assert any("not eligible for fallback" in t for t in texts)