codekingpro/portable-devtools
114k
1"""Model fallback middleware for agents."""2 3from __future__ import annotations4 5from typing import TYPE_CHECKING6 7from langchain.agents.middleware.types import (8 AgentMiddleware,9 AgentState,10 ContextT,11 ModelRequest,12 ModelResponse,13 ResponseT,14)15from langchain.chat_models import init_chat_model16 17if TYPE_CHECKING:18 from collections.abc import Awaitable, Callable19 20 from langchain_core.language_models.chat_models import BaseChatModel21 from langchain_core.messages import AIMessage22 23 24class ModelFallbackMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, ResponseT]):25 """Automatic fallback to alternative models on errors.26 27 Retries failed model calls with alternative models in sequence until28 success or all models exhausted. Primary model specified in `create_agent`.29 30 Example:31 ```python32 from langchain.agents.middleware import ModelFallbackMiddleware33 from langchain.agents import create_agent34 35 fallback = ModelFallbackMiddleware(36 "openai:gpt-4o-mini", # Try first on error37 "anthropic:claude-sonnet-4-5-20250929", # Then this38 )39 40 agent = create_agent(41 model="openai:gpt-4o", # Primary model42 middleware=[fallback],43 )44 45 # If primary fails: tries gpt-4o-mini, then claude-sonnet-4-5-2025092946 result = await agent.invoke({"messages": [HumanMessage("Hello")]})47 ```48 """49 50 def __init__(51 self,52 first_model: str | BaseChatModel,53 *additional_models: str | BaseChatModel,54 ) -> None:55 """Initialize model fallback middleware.56 57 Args:58 first_model: First fallback model (string name or instance).59 *additional_models: Additional fallbacks in order.60 """61 super().__init__()62 63 # Initialize all fallback models64 all_models = (first_model, *additional_models)65 self.models: list[BaseChatModel] = []66 for model in all_models:67 if isinstance(model, str):68 self.models.append(init_chat_model(model))69 else:70 self.models.append(model)71 72 def wrap_model_call(73 self,74 request: ModelRequest[ContextT],75 handler: Callable[[ModelRequest[ContextT]], ModelResponse[ResponseT]],76 ) -> ModelResponse[ResponseT] | AIMessage:77 """Try fallback models in sequence on errors.78 79 Args:80 request: Initial model request.81 handler: Callback to execute the model.82 83 Returns:84 AIMessage from successful model call.85 86 Raises:87 Exception: If all models fail, re-raises last exception.88 """89 # Try primary model first90 last_exception: Exception91 try:92 return handler(request)93 except Exception as e:94 last_exception = e95 96 # Try fallback models97 for fallback_model in self.models:98 try:99 return handler(request.override(model=fallback_model))100 except Exception as e:101 last_exception = e102 continue103 104 raise last_exception105 106 async def awrap_model_call(107 self,108 request: ModelRequest[ContextT],109 handler: Callable[[ModelRequest[ContextT]], Awaitable[ModelResponse[ResponseT]]],110 ) -> ModelResponse[ResponseT] | AIMessage:111 """Try fallback models in sequence on errors (async version).112 113 Args:114 request: Initial model request.115 handler: Async callback to execute the model.116 117 Returns:118 AIMessage from successful model call.119 120 Raises:121 Exception: If all models fail, re-raises last exception.122 """123 # Try primary model first124 last_exception: Exception125 try:126 return await handler(request)127 except Exception as e:128 last_exception = e129 130 # Try fallback models131 for fallback_model in self.models:132 try:133 return await handler(request.override(model=fallback_model))134 except Exception as e:135 last_exception = e136 continue137 138 raise last_exception139 