mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
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:
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:
|
||||
|
||||
+163
-2
@@ -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)
|
||||
|
||||
|
||||
Reference in new issue
Block a user