fix(router): make prompt caching affinity aware of auto-injected cache_control (#37689)

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-08-20 16:10:27 -07:00 • committed by GitHub
parent 8c42d8b97b
commit 33bafd0402
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 303 additions and 2 deletions

View file

@ -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],

View file

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

View file

@ -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."""

View file

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