mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
8c42d8b97b
commit
33bafd0402
4 changed files with 303 additions and 2 deletions
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue