diff --git a/libs/langchain_v1/langchain/agents/middleware/model_fallback.py b/libs/langchain_v1/langchain/agents/middleware/model_fallback.py index 6f5fa29d05..2bd58aa323 100644 --- a/libs/langchain_v1/langchain/agents/middleware/model_fallback.py +++ b/libs/langchain_v1/langchain/agents/middleware/model_fallback.py @@ -1,21 +1,19 @@ """Model fallback middleware for agents. -When a caching middleware such as `AnthropicPromptCachingMiddleware` wraps this -middleware from the outside, it applies Anthropic `cache_control` markers to the -request *before* the fallback loop runs. Those markers are provider-specific and -cause API errors on non-Anthropic fallback models, so this middleware strips them -from fallback attempts — but only when the fallback model itself cannot accept -Anthropic cache markers. When the fallback is another Anthropic model the markers -are valid and preserve prompt caching, so they are left intact. +When an outer caching middleware modifies a request, those changes happen before +the fallback loop runs. Provider-specific cache settings can cause API errors if +the loop reuses them with a different provider. This middleware therefore strips +unsupported Anthropic and Fireworks cache settings from each fallback attempt, +while preserving settings accepted by the selected fallback. -The knowledge of the `cache_control` marker is duplicated here (rather than owned -solely by the Anthropic partner package) because an outer caching middleware -never re-runs during fallback and therefore cannot clean up after itself. +This provider knowledge lives here because an outer caching middleware never +re-runs during fallback and therefore cannot clean up after itself. """ from __future__ import annotations import logging +from collections.abc import Mapping from typing import TYPE_CHECKING, Any from langchain_core.tools import BaseTool @@ -39,6 +37,9 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) +_FIREWORKS_LLM_TYPE = "fireworks-chat" +_FIREWORKS_SESSION_AFFINITY_HEADER = "x-session-affinity" + def _sanitize_content_blocks( content: str | list[str | dict[str, Any]], @@ -109,25 +110,48 @@ def _sanitize_tools( return sanitized_tools if changed else tools -def _sanitize_request_for_fallback(request: ModelRequest[ContextT]) -> ModelRequest[ContextT]: - """Sanitize provider-specific Anthropic cache markers before fallback attempts.""" +def _sanitize_request_for_fallback( + request: ModelRequest[ContextT], + fallback_model: BaseChatModel | None = None, +) -> ModelRequest[ContextT]: + """Sanitize provider-specific cache settings before a fallback attempt.""" overrides: dict[str, Any] = {} + supports_anthropic_cache = _supports_anthropic_cache_control(fallback_model) + + model_settings = request.model_settings + model_settings_changed = False + + if not supports_anthropic_cache: + model_settings, cache_control_changed = _without_cache_control(model_settings) + model_settings_changed = model_settings_changed or cache_control_changed + + if not _supports_fireworks_prompt_cache(fallback_model): + model_settings, fireworks_cache_changed = _without_fireworks_session_affinity( + model_settings + ) + model_settings_changed = model_settings_changed or fireworks_cache_changed + + if not _supports_prompt_cache_key(fallback_model) and "prompt_cache_key" in model_settings: + model_settings = { + key: value for key, value in model_settings.items() if key != "prompt_cache_key" + } + model_settings_changed = True - model_settings, model_settings_changed = _without_cache_control(request.model_settings) if model_settings_changed: overrides["model_settings"] = model_settings - system_message = _sanitize_system_message(request.system_message) - if system_message is not request.system_message: - overrides["system_message"] = system_message + if not supports_anthropic_cache: + system_message = _sanitize_system_message(request.system_message) + if system_message is not request.system_message: + overrides["system_message"] = system_message - messages = _sanitize_messages(request.messages) - if messages is not request.messages: - overrides["messages"] = messages + messages = _sanitize_messages(request.messages) + if messages is not request.messages: + overrides["messages"] = messages - tools = _sanitize_tools(request.tools) - if tools is not request.tools: - overrides["tools"] = tools + tools = _sanitize_tools(request.tools) + if tools is not request.tools: + overrides["tools"] = tools if not overrides: return request @@ -135,7 +159,7 @@ def _sanitize_request_for_fallback(request: ModelRequest[ContextT]) -> ModelRequ # Log only the field names that changed, never request content (may contain # prompt data or PII). logger.debug( - "Stripped Anthropic cache_control markers from %s before fallback attempt", + "Stripped provider-specific cache settings from %s before fallback attempt", sorted(overrides), ) @@ -208,6 +232,30 @@ def _without_cache_control(payload: dict[str, Any]) -> tuple[dict[str, Any], boo ) +def _without_fireworks_session_affinity( + model_settings: dict[str, Any], +) -> tuple[dict[str, Any], bool]: + """Return model settings without the Fireworks-specific affinity header.""" + extra_headers = model_settings.get("extra_headers") + if not isinstance(extra_headers, Mapping): + return model_settings, False + + sanitized_headers = { + key: value + for key, value in extra_headers.items() + if not (isinstance(key, str) and key.lower() == _FIREWORKS_SESSION_AFFINITY_HEADER) + } + if len(sanitized_headers) == len(extra_headers): + return model_settings, False + + sanitized_settings = dict(model_settings) + if sanitized_headers: + sanitized_settings["extra_headers"] = sanitized_headers + else: + sanitized_settings.pop("extra_headers") + return sanitized_settings, True + + def _without_cache_control_from_content_block( block: dict[str, Any], ) -> tuple[dict[str, Any], bool]: @@ -264,7 +312,7 @@ _ANTHROPIC_LLM_TYPES: frozenset[str] = frozenset( ) -def _supports_anthropic_cache_control(model: BaseChatModel) -> bool: +def _supports_anthropic_cache_control(model: BaseChatModel | None) -> bool: """Return whether `model` accepts Anthropic `cache_control` markers. Checked via `_llm_type` so the decision is provider-based rather than @@ -276,6 +324,21 @@ def _supports_anthropic_cache_control(model: BaseChatModel) -> bool: return isinstance(llm_type, str) and llm_type in _ANTHROPIC_LLM_TYPES +def _supports_fireworks_prompt_cache(model: BaseChatModel | None) -> bool: + """Return whether `model` accepts the Fireworks session-affinity header.""" + return getattr(model, "_llm_type", None) == _FIREWORKS_LLM_TYPE + + +def _supports_prompt_cache_key(model: BaseChatModel | None) -> bool: + """Return whether `model` accepts the shared `prompt_cache_key` parameter.""" + llm_type = getattr(model, "_llm_type", None) + return isinstance(llm_type, str) and llm_type in { + _FIREWORKS_LLM_TYPE, + "openai-chat", + "azure-openai-chat", + } + + class ModelFallbackMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, ResponseT]): """Automatic fallback to alternative models on errors. @@ -350,16 +413,12 @@ class ModelFallbackMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, R except Exception as e: last_exception = e - # Try fallback models — sanitize cache markers only when the fallback - # model cannot accept them (i.e. is not an Anthropic-compatible model). + # Try fallback models — sanitize provider-specific cache settings that + # the selected fallback cannot accept. # The request is derived outside the try so a sanitizer or `_llm_type` # bug surfaces directly instead of being masked as a model failure. for fallback_model in self.models: - fallback_request = ( - request - if _supports_anthropic_cache_control(fallback_model) - else _sanitize_request_for_fallback(request) - ) + fallback_request = _sanitize_request_for_fallback(request, fallback_model) try: return handler(fallback_request.override(model=fallback_model)) except GraphBubbleUp: @@ -396,16 +455,12 @@ class ModelFallbackMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, R except Exception as e: last_exception = e - # Try fallback models — sanitize cache markers only when the fallback - # model cannot accept them (i.e. is not an Anthropic-compatible model). + # Try fallback models — sanitize provider-specific cache settings that + # the selected fallback cannot accept. # The request is derived outside the try so a sanitizer or `_llm_type` # bug surfaces directly instead of being masked as a model failure. for fallback_model in self.models: - fallback_request = ( - request - if _supports_anthropic_cache_control(fallback_model) - else _sanitize_request_for_fallback(request) - ) + fallback_request = _sanitize_request_for_fallback(request, fallback_model) try: return await handler(fallback_request.override(model=fallback_model)) except GraphBubbleUp: diff --git a/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_model_fallback.py b/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_model_fallback.py index 5de942978b..320352cf90 100644 --- a/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_model_fallback.py +++ b/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_model_fallback.py @@ -11,6 +11,7 @@ from langchain_core.language_models.fake_chat_models import GenericFakeChatModel from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, SystemMessage from langchain_core.outputs import ChatGeneration, ChatResult from langchain_core.tools import BaseTool, tool +from langchain_openai import AzureChatOpenAI, ChatOpenAI from langgraph.errors import GraphInterrupt from typing_extensions import override @@ -128,6 +129,31 @@ def _make_request_with_cache_markers(primary_model: BaseChatModel) -> ModelReque ) +def _make_request_with_fireworks_cache_settings( + primary_model: BaseChatModel, +) -> ModelRequest: + """Create a request with Fireworks prompt-cache affinity settings.""" + return _make_request().override( + model=primary_model, + model_settings={ + "temperature": 0.3, + "prompt_cache_key": "thread-123", + "extra_headers": { + "X-Request-ID": "request-123", + "X-Session-Affinity": "thread-123", + }, + }, + ) + + +def _assert_fireworks_cache_settings_removed(request: ModelRequest) -> None: + """Assert Fireworks affinity is removed while unrelated settings remain.""" + assert request.model_settings == { + "temperature": 0.3, + "extra_headers": {"X-Request-ID": "request-123"}, + } + + def _assert_request_has_cache_markers(request: ModelRequest) -> None: """Assert request still contains Anthropic-style cache markers.""" assert "cache_control" in request.model_settings @@ -313,6 +339,54 @@ async def test_fallback_sanitizes_cache_markers_async() -> None: _assert_request_has_cache_markers(request) +def test_fallback_removes_fireworks_cache_settings_sync() -> None: + """A non-Fireworks fallback should not receive Fireworks cache settings.""" + primary_model = GenericFakeChatModel(messages=iter([])) + fallback_model = GenericFakeChatModel(messages=iter([])) + middleware = ModelFallbackMiddleware(fallback_model) + request = _make_request_with_fireworks_cache_settings(primary_model) + attempts: list[ModelRequest] = [] + + def mock_handler(req: ModelRequest) -> ModelResponse: + attempts.append(req) + if len(attempts) == 1: + msg = "Primary model failed" + raise ValueError(msg) + + assert req.model is fallback_model + _assert_fireworks_cache_settings_removed(req) + return ModelResponse(result=[AIMessage(content="fallback response")]) + + middleware.wrap_model_call(request, mock_handler) + + assert len(attempts) == 2 + assert request.model_settings["prompt_cache_key"] == "thread-123" + + +async def test_fallback_removes_fireworks_cache_settings_async() -> None: + """Async non-Fireworks fallback should not receive Fireworks cache settings.""" + primary_model = GenericFakeChatModel(messages=iter([])) + fallback_model = GenericFakeChatModel(messages=iter([])) + middleware = ModelFallbackMiddleware(fallback_model) + request = _make_request_with_fireworks_cache_settings(primary_model) + attempts: list[ModelRequest] = [] + + async def mock_handler(req: ModelRequest) -> ModelResponse: + attempts.append(req) + if len(attempts) == 1: + msg = "Primary model failed" + raise ValueError(msg) + + assert req.model is fallback_model + _assert_fireworks_cache_settings_removed(req) + return ModelResponse(result=[AIMessage(content="fallback response")]) + + await middleware.awrap_model_call(request, mock_handler) + + assert len(attempts) == 2 + assert request.model_settings["prompt_cache_key"] == "thread-123" + + def test_sanitize_collapses_emptied_extras_to_none() -> None: """Stripping the only `extras` key (`cache_control`) resets `extras` to None.""" @@ -810,6 +884,14 @@ class _FakeNonStringLlmTypeModel(GenericFakeChatModel): return ["anthropic-chat"] # type: ignore[return-value] +class _FakeFireworksModel(GenericFakeChatModel): + """Fake model that reports the `ChatFireworks` `_llm_type`.""" + + @property + def _llm_type(self) -> str: + return "fireworks-chat" + + _ANTHROPIC_COMPATIBLE_FAKES = [ _FakeAnthropicModel, _FakeBedrockAnthropicModel, @@ -828,6 +910,85 @@ def test_supports_anthropic_cache_control() -> None: assert not _supports_anthropic_cache_control(_FakeNonStringLlmTypeModel(messages=iter([]))) +def test_fallback_preserves_fireworks_cache_settings_for_fireworks() -> None: + """A Fireworks fallback should keep Fireworks prompt-cache affinity settings.""" + primary_model = _FakeFireworksModel(messages=iter([])) + fallback_model = _FakeFireworksModel(messages=iter([])) + middleware = ModelFallbackMiddleware(fallback_model) + request = _make_request_with_fireworks_cache_settings(primary_model) + attempts: list[ModelRequest] = [] + + def mock_handler(req: ModelRequest) -> ModelResponse: + attempts.append(req) + if len(attempts) == 1: + msg = "Primary model failed" + raise ValueError(msg) + + assert req.model is fallback_model + assert req.model_settings == request.model_settings + return ModelResponse(result=[AIMessage(content="fallback response")]) + + middleware.wrap_model_call(request, mock_handler) + + assert len(attempts) == 2 + + +@pytest.mark.parametrize("use_async", [False, True]) +@pytest.mark.parametrize( + ("azure", "extra_headers"), + [ + pytest.param(False, None, id="openai-without-headers"), + pytest.param(False, {"X-Session-Affinity": "thread-123"}, id="openai-affinity-only"), + pytest.param( + True, + {"X-Session-Affinity": "thread-123", "X-Request-ID": "request-123"}, + id="azure-preserves-other-headers", + ), + ], +) +async def test_openai_fallback_preserves_cache_key( + *, use_async: bool, azure: bool, extra_headers: dict[str, str] | None +) -> None: + """OpenAI accepts cache keys independently of Fireworks affinity headers.""" + primary = ChatOpenAI.model_construct(model_name="test-model") + fallback = ( + AzureChatOpenAI.model_construct(model_name="test-model") + if azure + else ChatOpenAI.model_construct(model_name="test-model") + ) + settings: dict[str, Any] = {"prompt_cache_key": "thread-123", "temperature": 0.3} + if extra_headers is not None: + settings["extra_headers"] = extra_headers + request = _make_request().override(model=primary, model_settings=settings) + middleware = ModelFallbackMiddleware(fallback) + attempts: list[ModelRequest] = [] + + def handler(req: ModelRequest) -> ModelResponse: + attempts.append(req) + if req.model is primary: + msg = "primary failed" + raise ValueError(msg) + assert req.model is fallback + expected: dict[str, Any] = {"prompt_cache_key": "thread-123", "temperature": 0.3} + if extra_headers and "X-Request-ID" in extra_headers: + expected["extra_headers"] = {"X-Request-ID": "request-123"} + assert req.model_settings == expected + return ModelResponse(result=[AIMessage(content="ok")]) + + async def ahandler(req: ModelRequest) -> ModelResponse: + return handler(req) + + if use_async: + await middleware.awrap_model_call(request, ahandler) + else: + middleware.wrap_model_call(request, handler) + + assert len(attempts) == 2 + assert request.model_settings == settings + if extra_headers is not None: + assert request.model_settings["extra_headers"]["X-Session-Affinity"] == "thread-123" + + def test_fallback_preserves_cache_markers_for_anthropic_sync() -> None: """Anthropic fallback keeps cache markers; non-Anthropic fallback strips them.""" primary_model = _FakeAnthropicModel(messages=iter([AIMessage(content="primary response")])) @@ -1053,7 +1214,7 @@ def test_fallback_sanitizer_error_is_not_masked_sync( Anthropic fallback succeeds. """ - def _boom(_request: ModelRequest) -> ModelRequest: + def _boom(_request: ModelRequest, _fallback_model: BaseChatModel | None = None) -> ModelRequest: msg = "sanitizer boom" raise RuntimeError(msg) @@ -1086,7 +1247,7 @@ async def test_fallback_sanitizer_error_is_not_masked_async( ) -> None: """Async: a sanitizer bug must surface, not be masked by a later success.""" - def _boom(_request: ModelRequest) -> ModelRequest: + def _boom(_request: ModelRequest, _fallback_model: BaseChatModel | None = None) -> ModelRequest: msg = "sanitizer boom" raise RuntimeError(msg)