From 33bafd0402bc8a27a4e28acee5f74d36869e72f4 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 20 Aug 2026 16:10:27 -0700 Subject: [PATCH] fix(router): make prompt caching affinity aware of auto-injected cache_control (#37689) Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../anthropic_cache_control_hook.py | 52 ++++- .../prompt_caching_deployment_check.py | 28 ++- .../test_anthropic_cache_control_hook.py | 17 ++ .../test_prompt_caching_deployment_check.py | 208 ++++++++++++++++++ 4 files changed, 303 insertions(+), 2 deletions(-) diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 1258c7593b4..f4f3b00dda0 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -27,11 +27,15 @@ from litellm.types.integrations.anthropic_cache_control_hook import ( CacheControlInjectionPoint, CacheControlMessageInjectionPoint, ) -from litellm.types.llms.anthropic import AnthropicSystemMessageContent +from litellm.types.llms.anthropic import ( + AllAnthropicToolsValues, + AnthropicSystemMessageContent, +) from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionCachedContent, ChatCompletionTextObject, + ChatCompletionToolParam, PromptCacheBreakpoint, PromptCacheOptions, ) @@ -57,6 +61,8 @@ OPENAI_PROMPT_CACHE_BREAKPOINT_BLOCK_TYPES: Final = frozenset( OPENAI_API_HOST: Final = "api.openai.com" OPENAI_API_BASE_ENV_VARS: Final = ("OPENAI_BASE_URL", "OPENAI_API_BASE") +AllToolParamValues = ChatCompletionToolParam | AllAnthropicToolsValues + def supports_openai_prompt_cache_breakpoint(model: str) -> bool: model_map_flag: Final = _model_map_prompt_cache_breakpoint_flag(model) @@ -625,6 +631,50 @@ class AnthropicCacheControlHook(CustomPromptManagement): ] return points + @staticmethod + def messages_with_default_injections( + messages: list[AllMessageValues], + models: Iterable[str], + tools: list[AllToolParamValues] | None = None, + enable_prompt_caching: bool | None = None, + ) -> list[AllMessageValues]: + """Return the messages auto prompt caching will send, default breakpoints included. + + Router cache affinity depends on this. Deployment selection runs before the injection in + `litellm.acompletion`, so it has to reproduce the markers to derive the same cache key the + success event later writes from the sent messages. `models` is every candidate model of the + group: the first that would auto-inject decides, since the default breakpoints (system + prompt and trailing turn) do not depend on which deployment serves the call. Returns the + input list itself when auto-injection would not apply + """ + points: Final = next( + ( + candidate + for candidate in ( + AnthropicCacheControlHook.get_default_injection_points( + messages=messages, + system=None, + model=model, + custom_llm_provider=None, + tools=tools, + enable_prompt_caching=enable_prompt_caching, + ) + for model in models + ) + if candidate + ), + None, + ) + if not points: + return messages + return AnthropicCacheControlHook._apply_message_injections( + points=cast( # cast-ok: the default points are all message-location points + list[CacheControlMessageInjectionPoint], points + ), + messages=copy.deepcopy(messages), + max_blocks=MAX_CACHE_CONTROL_BLOCKS, + ) + @staticmethod def maybe_seed_default_injection_points( non_default_params: dict[str, Any], diff --git a/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py b/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py index e928f4a0c3f..6e8406b2ec7 100644 --- a/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py +++ b/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py @@ -9,6 +9,10 @@ from typing import Final, cast from litellm import verbose_logger from litellm.caching.dual_cache import DualCache from litellm.constants import DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT +from litellm.integrations.anthropic_cache_control_hook import ( + AllToolParamValues, + AnthropicCacheControlHook, +) from litellm.integrations.custom_logger import CustomLogger, Span from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import CallTypes, StandardLoggingPayload @@ -63,8 +67,30 @@ class PromptCachingDeploymentCheck(CustomLogger): cache=self.cache, ) - model_id_dict: Final = await prompt_cache.async_get_model_id( + ## AUTO PROMPT CACHING - the breakpoints this request will carry are injected inside + ## `litellm.acompletion`, after a deployment has been picked, so the affinity key has to + ## be derived from the messages as they will be sent, not as they arrive here. + affinity_messages: Final = AnthropicCacheControlHook.messages_with_default_injections( messages=cast(list[AllMessageValues], messages), + models=( + deployment["litellm_params"]["model"] + for deployment in healthy_deployments + if isinstance(deployment.get("litellm_params"), dict) and deployment["litellm_params"].get("model") + ), + tools=( + cast( # cast-ok: request_kwargs is untyped; the stand-down scan duck-types every tool it reads + list[AllToolParamValues] | None, request_kwargs.get("tools") + ) + if request_kwargs is not None + else None + ), + enable_prompt_caching=( + request_kwargs.get("enable_prompt_caching") is True if request_kwargs is not None else None + ), + ) + + model_id_dict: Final = await prompt_cache.async_get_model_id( + messages=affinity_messages, tools=None, ) if model_id_dict is not None: diff --git a/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py b/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py index e92a368d24b..7bf15f59eb9 100644 --- a/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py +++ b/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py @@ -1738,6 +1738,23 @@ class TestEnableAnthropicPromptCaching: assert result_msgs[-1]["content"][-1]["cache_control"] == {"type": "ephemeral"} assert "cache_control" not in result_msgs[0]["content"][-1] + def test_messages_with_default_injections_leaves_the_caller_list_untouched(self, monkeypatch): + """ + Routing calls this on the live request's own message list to derive the affinity key, before + the request is sent. Marking in place would leak litellm's breakpoints into the caller's + messages, where the real injection pass later reads them back as client-supplied ones. + """ + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + messages = copy.deepcopy(self.MESSAGES) + before = copy.deepcopy(messages) + + injected = AnthropicCacheControlHook.messages_with_default_injections( + messages=messages, models=("claude-sonnet-4-5",) + ) + + assert injected != messages + assert messages == before + class TestPerKeyEnablePromptCaching: """Per-request enable_prompt_caching override (stamped from key metadata) with the global flag off.""" diff --git a/tests/test_litellm/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py b/tests/test_litellm/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py index 6752d76847f..bff6f261020 100644 --- a/tests/test_litellm/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py +++ b/tests/test_litellm/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py @@ -1,3 +1,5 @@ +import asyncio +import copy import os import sys from typing import List, cast @@ -9,6 +11,8 @@ sys.path.insert(0, os.path.abspath("../..")) import litellm from litellm.caching.dual_cache import DualCache from litellm.constants import DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT +from litellm.integrations.anthropic_cache_control_hook import AnthropicCacheControlHook +from litellm.integrations.custom_logger import CustomLogger from litellm.router_utils.pre_call_checks.prompt_caching_deployment_check import ( PromptCachingDeploymentCheck, _get_min_token_count_for_deployments, @@ -187,6 +191,210 @@ async def test_async_filter_deployments_narrows_for_group_whose_model_minimum_is assert filtered == [deployments[1]] +AUTO_CACHING_MODEL = "anthropic/claude-sonnet-4-5" + + +def _auto_caching_messages() -> List[AllMessageValues]: + """A prompt over the model minimum that carries no client cache_control.""" + return cast( + List[AllMessageValues], + [ + {"role": "system", "content": "word " * 3000}, + {"role": "user", "content": "hello"}, + ], + ) + + +def _affinity_messages(messages: List[AllMessageValues]) -> List[AllMessageValues]: + """The messages the check keys deployment affinity on, for a group of `AUTO_CACHING_MODEL`.""" + return AnthropicCacheControlHook.messages_with_default_injections( + messages=messages, + models=(AUTO_CACHING_MODEL,), + ) + + +class _SentMessagesCapture(CustomLogger): + def __init__(self): + self.messages: List[AllMessageValues] | None = None + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + standard_logging_object = kwargs.get("standard_logging_object") + if standard_logging_object is not None: + self.messages = standard_logging_object["messages"] + + +async def _eventually(predicate, timeout: float = 10.0): + """Success callbacks run as tasks, so give the write a bounded window to land.""" + deadline = asyncio.get_running_loop().time() + timeout + while asyncio.get_running_loop().time() < deadline: + result = predicate() + if result: + return result + await asyncio.sleep(0.05) + return predicate() + + +@pytest.mark.asyncio +async def test_affinity_key_matches_the_messages_auto_caching_actually_sends(monkeypatch, local_model_cost_map): + """ + The regression. `enable_anthropic_prompt_caching` injects cache_control inside + `litellm.acompletion`, which runs after routing, so at filter time the messages carried no + marker, `extract_cacheable_prefix` returned [], the key was None, and the check no-opped on + every request. Routing must derive the same key the success event writes from the messages the + request was actually sent with, otherwise auto-injected caching gets no affinity at all. + """ + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + capture = _SentMessagesCapture() + monkeypatch.setattr(litellm, "callbacks", [capture]) + messages = _auto_caching_messages() + + await litellm.acompletion( + model=AUTO_CACHING_MODEL, + messages=copy.deepcopy(messages), + mock_response="ok", + api_key="sk-fake", + ) + sent_messages = await _eventually(lambda: capture.messages) + assert sent_messages is not None + + routing_key = PromptCachingCache.get_prompt_caching_cache_key(_affinity_messages(messages), None) + + assert routing_key is not None + assert routing_key == PromptCachingCache.get_prompt_caching_cache_key(sent_messages, None) + + +@pytest.mark.asyncio +async def test_repeated_auto_cached_prefix_pins_to_one_deployment(monkeypatch, local_model_cost_map): + """ + End to end over the router: identical requests with no client cache_control must stop bouncing + across a multi-deployment group once one deployment has cached the prefix. Bedrock and Anthropic + caches are per account and region, so every bounce paid the cache write premium and never read. + """ + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + router = litellm.Router( + model_list=[ + { + "model_name": MODEL_GROUP_ALIAS, + "litellm_params": {"model": AUTO_CACHING_MODEL, "api_key": "sk-fake"}, + "model_info": {"id": model_id}, + } + for model_id in ("dep-1", "dep-2") + ], + optional_pre_call_checks=["prompt_caching"], + ) + messages = _auto_caching_messages() + + first = await router.acompletion(model=MODEL_GROUP_ALIAS, messages=messages, mock_response="ok") + served_by = first._hidden_params["model_id"] + + affinity_key = PromptCachingCache.get_prompt_caching_cache_key(_affinity_messages(messages), None) + assert await _eventually(lambda: router.cache.get_cache(key=affinity_key)) is not None + + subsequent = [ + (await router.acompletion(model=MODEL_GROUP_ALIAS, messages=messages, mock_response="ok"))._hidden_params[ + "model_id" + ] + for _ in range(4) + ] + + assert subsequent == [served_by] * 4 + + +@pytest.mark.asyncio +async def test_per_request_enable_prompt_caching_reaches_the_affinity_key(monkeypatch, local_model_cost_map): + """ + `enable_prompt_caching` turns auto-injection on for a single request while the global flag stays + off, so routing has to read it too. Ignore it and the key comes off unmarked messages, which is + never what the request goes on to send, and the pin is lost for every per-key enablement. + """ + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", False) + cache = DualCache() + check = PromptCachingDeploymentCheck(cache=cache) + deployments = _deployments(AUTO_CACHING_MODEL, AUTO_CACHING_MODEL) + messages = _auto_caching_messages() + + sent = AnthropicCacheControlHook.messages_with_default_injections( + messages=messages, models=(AUTO_CACHING_MODEL,), enable_prompt_caching=True + ) + assert sent != messages + await PromptCachingCache(cache=cache).async_add_model_id(model_id="dep-2", messages=sent, tools=None) + + filtered = await check.async_filter_deployments( + model=MODEL_GROUP_ALIAS, + healthy_deployments=deployments, + messages=messages, + request_kwargs={"enable_prompt_caching": True}, + ) + + assert filtered == [deployments[1]] + + +@pytest.mark.asyncio +async def test_tool_marked_cache_control_keeps_routing_off_another_requests_prefix(monkeypatch, local_model_cost_map): + """ + Tools carrying the client's own cache_control make auto-injection stand down, so this request + will not carry litellm's breakpoints. Routing must see the tools as well. Ignore them and it + keys off the injected prefix, pinning the request to whichever deployment cached a different, + tool-less request whose prefix it can never actually reuse. + """ + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + cache = DualCache() + check = PromptCachingDeploymentCheck(cache=cache) + deployments = _deployments(AUTO_CACHING_MODEL, AUTO_CACHING_MODEL) + messages = _auto_caching_messages() + cache_marked_tools = [ + { + "type": "function", + "function": {"name": "get_weather", "parameters": {"type": "object", "properties": {}}}, + "cache_control": {"type": "ephemeral"}, + } + ] + + await PromptCachingCache(cache=cache).async_add_model_id( + model_id="dep-2", messages=_affinity_messages(messages), tools=None + ) + + without_tools = await check.async_filter_deployments( + model=MODEL_GROUP_ALIAS, healthy_deployments=deployments, messages=messages + ) + assert without_tools == [deployments[1]] + + with_tools = await check.async_filter_deployments( + model=MODEL_GROUP_ALIAS, + healthy_deployments=deployments, + messages=messages, + request_kwargs={"tools": cache_marked_tools}, + ) + + assert with_tools == deployments + + +def test_client_supplied_cache_control_keeps_its_own_prefix_boundary(monkeypatch, local_model_cost_map): + """ + Auto-injection stands down when the client marks its own breakpoints, so the affinity key must + keep keying off the client's boundary. Injecting on top would push the boundary to the trailing + turn and break affinity for prompts that already worked. + """ + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + messages = cast( + List[AllMessageValues], + [ + { + "role": "system", + "content": [ + {"type": "text", "text": "word " * 3000, "cache_control": {"type": "ephemeral"}}, + ], + }, + {"role": "user", "content": "hello"}, + ], + ) + + for_key = _affinity_messages(messages) + + assert for_key is messages + assert PromptCachingCache.extract_cacheable_prefix(for_key) == messages[:1] + + @pytest.mark.asyncio async def test_wildcard_route_resolves_underlying_model_minimum(local_model_cost_map): from litellm import Router