feat(sdk): add Bedrock prompt caching middleware (#4108)

Bedrock models now receive `BedrockPromptCachingMiddleware`
automatically when `langchain-aws` is installed.

Adds Python SDK parity with
[deepagentsjs](https://github.com/langchain-ai/deepagentsjs/pull/611) by
wiring Bedrock prompt caching into main and subagent middleware stacks
for AWS Bedrock models.

Docs PR: https://github.com/langchain-ai/docs/pull/4587

Made by [Open
SWE](https://openswe.vercel.app/agents/696a9011-2123-6396-5992-265911c8d2af)

---------

Co-authored-by: Alexander Olsen <13665641+aolsenjazz@users.noreply.github.com>
Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
Co-authored-by: Mason Daugherty <github@mdrxy.com>
Co-authored-by: Mason Daugherty <mason@langchain.dev>
This commit is contained in:
Alexander Olsen
2026-06-25 11:58:38 -04:00
committed by GitHub
parent 7db4316e14
commit a398382c85
4 changed files with 203 additions and 8 deletions
+38 -2
View File
@@ -12,12 +12,24 @@ from deepagents.profiles.provider.provider_profiles import apply_provider_profil
logger = logging.getLogger(__name__)
# LangChain specs and LangSmith params use different provider names for some
# integrations. Canonicalize only known aliases before comparing providers.
_PROVIDER_ALIASES = {
"azure_openai": "azure",
"mistralai": "mistral",
}
"""Known provider aliases between LangChain specs and LangSmith params.
LangChain specs and LangSmith params use different provider names for some
integrations. Canonicalize only known aliases before comparing providers.
"""
_BEDROCK_PROVIDERS = frozenset({"amazon_bedrock", "anthropic_bedrock", "aws", "bedrock", "bedrock_converse"})
"""Normalized provider names that identify AWS Bedrock chat models."""
_BEDROCK_MODEL_CLASSES = frozenset({"ChatAnthropicBedrock", "ChatBedrock", "ChatBedrockConverse", "ChatBedrockNovaSonic"})
"""`langchain-aws` chat model class names that identify AWS Bedrock models."""
_BEDROCK_REGIONAL_PREFIXES = ("apac.", "amer.", "au.", "eu.", "global.", "jp.", "sa.", "us.", "us-gov.")
"""Regional inference profile prefixes stripped from Bedrock model identifiers."""
def resolve_model(model: str | BaseChatModel) -> BaseChatModel:
@@ -106,6 +118,30 @@ def get_model_provider(model: BaseChatModel) -> str | None:
return None
def is_bedrock_model(model: str | BaseChatModel) -> bool:
"""Check whether a model targets AWS Bedrock."""
if isinstance(model, str):
if _is_bedrock_nova_model_id(model):
return True
provider, separator, _ = model.partition(":")
return bool(separator) and _normalize_provider(provider) in _BEDROCK_PROVIDERS
provider = get_model_provider(model)
if provider is not None and _normalize_provider(provider) in _BEDROCK_PROVIDERS:
return True
return type(model).__name__ in _BEDROCK_MODEL_CLASSES
def _is_bedrock_nova_model_id(model: str) -> bool:
"""Check for cache-capable Bedrock Nova model identifiers."""
identifier = model
for prefix in _BEDROCK_REGIONAL_PREFIXES:
if identifier.startswith(prefix):
identifier = identifier.removeprefix(prefix)
break
return identifier.startswith("amazon.nova-")
def model_matches_spec(model: BaseChatModel, spec: str) -> bool:
"""Check whether a model instance already matches a string model spec.
+28 -6
View File
@@ -7,6 +7,7 @@ subagent, and summarization middleware.
import logging
from collections.abc import Callable, Sequence
from importlib import import_module
from typing import Annotated, Any, Required, cast
from langchain.agents import AgentState, create_agent
@@ -185,6 +186,28 @@ def get_default_model() -> ChatAnthropic:
return _build_default_model()
def _create_bedrock_prompt_caching_middleware() -> AgentMiddleware[Any, Any, Any] | None:
"""Create Bedrock prompt caching middleware when `langchain-aws` is installed."""
module_name = "langchain_aws.middleware.prompt_caching"
try:
module = import_module(module_name)
except ImportError as exc:
if exc.name not in {"langchain_aws", "langchain_aws.middleware", module_name}:
raise
logger.debug("Bedrock prompt caching middleware is unavailable.", exc_info=exc)
return None
middleware_cls = module.BedrockPromptCachingMiddleware
return cast("AgentMiddleware[Any, Any, Any]", middleware_cls(unsupported_model_behavior="ignore"))
def _append_prompt_caching_middleware(middleware: list[AgentMiddleware[Any, Any, Any]]) -> None:
"""Append provider-specific prompt caching middleware."""
middleware.append(AnthropicPromptCachingMiddleware(unsupported_model_behavior="ignore"))
bedrock_middleware = _create_bedrock_prompt_caching_middleware()
if bedrock_middleware is not None:
middleware.append(bedrock_middleware)
def _merge_fs_interrupt_on(
fs_interrupt_on: dict[str, InterruptOnConfig],
user_interrupt_on: dict[str, bool | InterruptOnConfig] | None,
@@ -349,6 +372,8 @@ def create_deep_agent( # noqa: C901, PLR0912, PLR0915 # Complex graph assembly
- `_ToolExclusionMiddleware` (if profile has `excluded_tools`)
- [`AnthropicPromptCachingMiddleware`][langchain_anthropic.middleware.AnthropicPromptCachingMiddleware] (unconditional; no-ops for
non-Anthropic models)
- [`BedrockPromptCachingMiddleware`](https://reference.langchain.com/python/langchain-aws/middleware/prompt_caching/BedrockPromptCachingMiddleware)
when `langchain-aws` is installed (no-ops for non-Bedrock models)
- [`MemoryMiddleware`][deepagents.middleware.memory.MemoryMiddleware] (if `memory` is provided)
- [`HumanInTheLoopMiddleware`][langchain.agents.middleware.HumanInTheLoopMiddleware] (if `interrupt_on` is provided)
@@ -635,8 +660,7 @@ def create_deep_agent( # noqa: C901, PLR0912, PLR0915 # Complex graph assembly
if _subagent_profile.excluded_tools:
subagent_middleware.append(_ToolExclusionMiddleware(excluded=_subagent_profile.excluded_tools))
# Prompt caching
subagent_middleware.append(AnthropicPromptCachingMiddleware(unsupported_model_behavior="ignore"))
_append_prompt_caching_middleware(subagent_middleware)
_subagent_matched_classes: set[type[AgentMiddleware[Any, Any, Any]]] = set()
_subagent_matched_names: set[str] = set()
@@ -710,8 +734,7 @@ def create_deep_agent( # noqa: C901, PLR0912, PLR0915 # Complex graph assembly
# Strip excluded tools after all tool-injecting middleware has run
if _profile.excluded_tools:
gp_middleware.append(_ToolExclusionMiddleware(excluded=_profile.excluded_tools))
# Prompt caching is unconditional: "ignore" silently skips non-Anthropic models
gp_middleware.append(AnthropicPromptCachingMiddleware(unsupported_model_behavior="ignore"))
_append_prompt_caching_middleware(gp_middleware)
gp_middleware = _apply_excluded_middleware(
gp_middleware,
@@ -793,8 +816,7 @@ def create_deep_agent( # noqa: C901, PLR0912, PLR0915 # Complex graph assembly
deepagent_middleware.extend(_profile.materialize_extra_middleware())
if _profile.excluded_tools:
deepagent_middleware.append(_ToolExclusionMiddleware(excluded=_profile.excluded_tools))
# Unconditional prompt caching (see general-purpose subagent comment).
deepagent_middleware.append(AnthropicPromptCachingMiddleware(unsupported_model_behavior="ignore"))
_append_prompt_caching_middleware(deepagent_middleware)
if memory is not None:
# MemoryMiddleware applies the cache_control breakpoint only when the
# request model is Anthropic, making it safe to enable unconditionally.
@@ -21,6 +21,7 @@ from deepagents.graph import (
_REQUIRED_MIDDLEWARE_NAMES,
BASE_AGENT_PROMPT,
DeepAgentState,
_create_bedrock_prompt_caching_middleware,
create_deep_agent,
get_default_model,
)
@@ -348,6 +349,88 @@ class TestGeneralPurposeSubagentProfileWiring:
_HARNESS_PROFILES.update(original)
class TestPromptCachingWiring:
"""Tests for provider-specific prompt caching middleware wiring."""
def test_main_and_general_purpose_agents_get_bedrock_prompt_caching(self) -> None:
model = GenericFakeChatModel(messages=iter([AIMessage(content="ok")]))
gp_cache = MagicMock()
main_cache = MagicMock()
fake_agent = MagicMock()
fake_agent.with_config.return_value = "compiled-agent"
with (
patch("deepagents.graph._create_bedrock_prompt_caching_middleware", side_effect=[gp_cache, main_cache]),
patch("deepagents.graph.SubAgentMiddleware", return_value=MagicMock()) as mock_subagents,
patch("deepagents.graph.create_agent", return_value=fake_agent) as mock_create,
):
result = create_deep_agent(model=model)
assert result == "compiled-agent"
subagents = mock_subagents.call_args.kwargs["subagents"]
general_purpose = next(spec for spec in subagents if spec["name"] == "general-purpose")
assert gp_cache in general_purpose["middleware"]
assert main_cache in mock_create.call_args.kwargs["middleware"]
def test_bedrock_explicit_subagent_gets_prompt_caching(self) -> None:
main_model = GenericFakeChatModel(messages=iter([AIMessage(content="main")]))
bedrock_model = GenericFakeChatModel(messages=iter([AIMessage(content="sub")]))
bedrock_model._get_ls_params = MagicMock(return_value={"ls_provider": "amazon_bedrock"})
subagent_cache = MagicMock()
fake_agent = MagicMock()
fake_agent.with_config.return_value = "compiled-agent"
with (
patch("deepagents.graph._create_bedrock_prompt_caching_middleware", return_value=subagent_cache),
patch("deepagents.graph.SubAgentMiddleware", return_value=MagicMock()) as mock_subagents,
patch("deepagents.graph.create_agent", return_value=fake_agent),
):
create_deep_agent(
model=main_model,
subagents=[
{
"name": "bedrock-worker",
"description": "Uses Bedrock.",
"system_prompt": "Help with Bedrock tasks.",
"model": bedrock_model,
}
],
)
subagents = mock_subagents.call_args.kwargs["subagents"]
bedrock_worker = next(spec for spec in subagents if spec["name"] == "bedrock-worker")
assert subagent_cache in bedrock_worker["middleware"]
def test_bedrock_prompt_caching_is_optional_when_middleware_unavailable(self) -> None:
bedrock_model = GenericFakeChatModel(messages=iter([AIMessage(content="ok")]))
bedrock_model._get_ls_params = MagicMock(return_value={"ls_provider": "amazon_bedrock"})
fake_agent = MagicMock()
fake_agent.with_config.return_value = "compiled-agent"
with (
patch(
"deepagents.graph.import_module",
side_effect=ModuleNotFoundError(name="langchain_aws.middleware.prompt_caching"),
),
patch("deepagents.graph.SubAgentMiddleware", return_value=MagicMock()) as mock_subagents,
patch("deepagents.graph.create_agent", return_value=fake_agent) as mock_create,
):
result = create_deep_agent(model=bedrock_model)
assert result == "compiled-agent"
subagents = mock_subagents.call_args.kwargs["subagents"]
general_purpose = next(spec for spec in subagents if spec["name"] == "general-purpose")
assert None not in general_purpose["middleware"]
assert None not in mock_create.call_args.kwargs["middleware"]
def test_bedrock_prompt_caching_preserves_unrelated_import_errors(self) -> None:
with (
patch("deepagents.graph.import_module", side_effect=ImportError(name="missing_transitive")),
pytest.raises(ImportError),
):
_create_bedrock_prompt_caching_middleware()
class TestSystemPromptAssembly:
"""Tests for system prompt assembly: profile base_system_prompt, suffix, and user prompt interaction."""
@@ -13,6 +13,7 @@ from langchain_core.language_models import BaseChatModel
from deepagents._models import (
get_model_identifier,
get_model_provider,
is_bedrock_model,
model_matches_spec,
resolve_model,
)
@@ -44,6 +45,7 @@ from deepagents.profiles.provider.provider_profiles import (
apply_provider_profile,
get_provider_profile,
)
from tests.unit_tests.chat_model import GenericFakeChatModel
_OPENROUTER_AZURE_IGNORE = {"ignore": ["azure"]}
"""Expected default value of `openrouter_provider` injected by the SDK profile."""
@@ -226,6 +228,58 @@ class TestGetModelProvider:
assert get_model_provider(model) is None
class TestIsBedrockModel:
"""Tests for `is_bedrock_model`."""
@pytest.mark.parametrize(
"model",
[
"bedrock:anthropic.claude-3-5-sonnet-20240620-v1:0",
"bedrock_converse:us.anthropic.claude-3-7-sonnet-20250219-v1:0",
"aws:amazon.nova-pro-v1:0",
"anthropic_bedrock:us.anthropic.claude-sonnet-4-6-20251117-v1:0",
"amazon.nova-pro-v1:0",
"us.amazon.nova-pro-v1:0",
],
)
def test_detects_bedrock_provider_strings(self, model: str) -> None:
assert is_bedrock_model(model) is True
@pytest.mark.parametrize(
"model",
[
"amazon.titan-text-express-v1:0",
"anthropic:claude-3-opus",
"openai:gpt-5",
],
)
def test_rejects_non_bedrock_provider_strings(self, model: str) -> None:
assert is_bedrock_model(model) is False
@pytest.mark.parametrize(
"provider",
["amazon_bedrock", "anthropic-bedrock", "bedrock", "bedrock_converse", "aws"],
)
def test_detects_bedrock_model_providers(self, provider: str) -> None:
model = _make_model({})
model._get_ls_params = MagicMock(return_value={"ls_provider": provider})
assert is_bedrock_model(model) is True
@pytest.mark.parametrize(
"model_cls",
["ChatAnthropicBedrock", "ChatBedrock", "ChatBedrockConverse", "ChatBedrockNovaSonic"],
)
def test_detects_bedrock_model_classes_when_provider_unavailable(self, model_cls: str) -> None:
model = type(model_cls, (GenericFakeChatModel,), {})(messages=iter([]))
model._get_ls_params = MagicMock(return_value={})
assert is_bedrock_model(model) is True
def test_rejects_non_bedrock_model_provider(self) -> None:
model = _make_model({})
model._get_ls_params = MagicMock(return_value={"ls_provider": "anthropic"})
assert is_bedrock_model(model) is False
class TestModelMatchesSpec:
"""Tests for `model_matches_spec`."""