fix(langchain): sanitize cache settings for fallback models (#40886)

`ModelFallbackMiddleware` now removes unsupported cache keys and
Fireworks session-affinity headers before fallback attempts, while
preserving cache settings supported by the selected fallback model.

---

When an agent falls back to another provider, explicitly supplied cache
settings in `ModelRequest.model_settings` can reach a model that does
not accept them and cause the fallback to fail.

`ModelFallbackMiddleware` now removes `x-session-affinity` for
non-Fireworks fallbacks and removes `prompt_cache_key` for providers
outside Fireworks, OpenAI, and Azure OpenAI. It preserves unrelated
settings and headers, retains the existing Anthropic cache-marker
handling, and leaves the original request unchanged. Unit tests cover
synchronous and asynchronous fallback, supported cache-key preservation,
and header cleanup.

Stacked on #38823. This PR contains only the fallback cleanup and its
tests. The Fireworks middleware in the base PR works independently
because its generated affinity is consumed directly by `ChatFireworks`.

Review focus: provider support is determined through `_llm_type`;
`prompt_cache_key` is shared by multiple providers and must not be
treated as Fireworks-only.
This commit is contained in:
Mason Daugherty authored and GitHub committed 2026-09-28 16:09:47 -04:00
1 parent 88b972731a
commit 2ade674ffc
2 files changed
+256 -40

No files matched your search

@@ -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:
@@ -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)