fix(router): hash the prompt caching affinity prefix off the event loop

Offload the per-block hashing through offload_token_count on both the pre-call
read and the success-event write, hash raw bytes as base64 instead of raising,
drop the unused serialize_object helper, and bind the chained digest, the
message envelope, and the bytes path in the regression tests
This commit is contained in:
mateo-berri 2026-09-19 20:22:28 -07:00
parent 517fff5bb7
commit 3ffe6272c9
4 changed files with 43 additions and 85 deletions

View file

@ -399,8 +399,6 @@ MINIMUM_PROMPT_CACHE_TOKEN_COUNT: Final = (
if MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE is not None
else DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT
)
# Anthropic checks at most 20 block positions behind a breakpoint for a cached prefix, a run of tool_use
# or tool_result blocks counting as one position, so deployment affinity probes the same window
PROMPT_CACHE_LOOKBACK_POSITIONS: Final = 20
DEFAULT_TRIM_RATIO: Final = float(
os.getenv("DEFAULT_TRIM_RATIO", 0.75)

View file

@ -15,10 +15,8 @@ from typing_extensions import TypedDict
from litellm.caching.caching import DualCache
from litellm.constants import PROMPT_CACHE_LOOKBACK_POSITIONS
from litellm.litellm_core_utils.logging_utils import (
truncate_base64_in_messages,
truncate_base64_in_messages_async,
)
from litellm.litellm_core_utils.logging_utils import truncate_base64_in_messages
from litellm.litellm_core_utils.token_counter import offload_token_count
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
if TYPE_CHECKING:
@ -88,7 +86,9 @@ def _seed(tools: Sequence[ChatCompletionToolParam] | None) -> bytes:
if tools is None:
return hashlib.sha256(b"").digest()
return hashlib.sha256(
_canonical_bytes(_TOOLS_ADAPTER.validate_python(to_jsonable_python(tools, serialize_unknown=True)))
_canonical_bytes(
_TOOLS_ADAPTER.validate_python(to_jsonable_python(tools, serialize_unknown=True, bytes_mode="base64"))
)
).digest()
@ -132,23 +132,6 @@ class PromptCachingCache:
def __init__(self, cache: DualCache):
self.cache = cache
@staticmethod
def serialize_object(obj: Any) -> object:
"""Helper function to serialize Pydantic objects, dictionaries, or fallback to string."""
if hasattr(obj, "dict"):
# If the object is a Pydantic model, use its `dict()` method
return obj.dict()
elif isinstance(obj, dict):
# If the object is a dictionary, serialize it with sorted keys
return json.dumps(obj, sort_keys=True, separators=(",", ":")) # Standardize serialization
elif isinstance(obj, list):
# Serialize lists by ensuring each element is handled properly
return [PromptCachingCache.serialize_object(item) for item in obj]
elif isinstance(obj, (int, float, bool)):
return obj # Keep primitive types as-is
return str(obj)
@staticmethod
def extract_cacheable_prefix(
messages: list[AllMessageValues],
@ -263,6 +246,7 @@ class PromptCachingCache:
to_jsonable_python(
truncate_base64_in_messages(PromptCachingCache.extract_cacheable_prefix(messages)),
serialize_unknown=True,
bytes_mode="base64",
)
),
tools,
@ -275,15 +259,7 @@ class PromptCachingCache:
) -> tuple[PrefixPosition, ...]:
if not messages:
return ()
return _positions_of(
_PREFIX_ADAPTER.validate_python(
to_jsonable_python(
await truncate_base64_in_messages_async(PromptCachingCache.extract_cacheable_prefix(messages)),
serialize_unknown=True,
)
),
tools,
)
return await offload_token_count(PromptCachingCache.prefix_positions)(messages, tools)
@staticmethod
def get_prompt_caching_cache_key(

View file

@ -11,57 +11,9 @@ from unittest.mock import patch, MagicMock, AsyncMock
from create_mock_standard_logging_payload import create_standard_logging_payload
from litellm.types.utils import StandardLoggingPayload
import unittest
from pydantic import BaseModel
from litellm.router_utils.prompt_caching_cache import PromptCachingCache
class ExampleModel(BaseModel):
field1: str
field2: int
def test_serialize_pydantic_object():
model = ExampleModel(field1="value", field2=42)
serialized = PromptCachingCache.serialize_object(model)
assert serialized == {"field1": "value", "field2": 42}
def test_serialize_dict():
obj = {"b": 2, "a": 1}
serialized = PromptCachingCache.serialize_object(obj)
assert serialized == '{"a":1,"b":2}' # JSON string with sorted keys
def test_serialize_nested_dict():
obj = {"z": {"b": 2, "a": 1}, "x": [1, 2, {"c": 3}]}
serialized = PromptCachingCache.serialize_object(obj)
expected = '{"x":[1,2,{"c":3}],"z":{"a":1,"b":2}}' # JSON string with sorted keys
assert serialized == expected
def test_serialize_list():
obj = ["item1", {"a": 1, "b": 2}, 42]
serialized = PromptCachingCache.serialize_object(obj)
expected = ["item1", '{"a":1,"b":2}', 42]
assert serialized == expected
def test_serialize_fallback():
obj = 12345 # Simple non-serializable object
serialized = PromptCachingCache.serialize_object(obj)
assert serialized == 12345
def test_serialize_non_serializable():
class CustomClass:
def __str__(self):
return "custom_object"
obj = CustomClass()
serialized = PromptCachingCache.serialize_object(obj)
assert serialized == "custom_object" # Fallback to string conversion
@pytest.mark.asyncio
async def test_router_prompt_caching_same_cacheable_prefix_routes_to_same_deployment():
"""

View file

@ -708,7 +708,10 @@ async def test_a_run_of_tool_blocks_counts_as_one_lookback_position():
@pytest.mark.asyncio
async def test_an_edited_earlier_block_does_not_inherit_the_pin():
"""Walking back must still bind every block's content, or an edited conversation pins to a stale cache."""
"""
Every key must bind the whole prefix before its block, not the block alone, or a conversation
that repeats a pinned block after an edit walks back onto a cache the provider no longer holds.
"""
prompt_cache = PromptCachingCache(cache=DualCache())
await prompt_cache.async_add_model_id(
model_id="dep-1", messages=_turn({"role": "user", "content": [_marked("original")]}), tools=None
@ -716,12 +719,41 @@ async def test_an_edited_earlier_block_does_not_inherit_the_pin():
edited = _turn(
{"role": "user", "content": [_text("edited")]},
{"role": "assistant", "content": "ok"},
{"role": "user", "content": [_marked("next")]},
{"role": "user", "content": [_marked("original")]},
)
assert await prompt_cache.async_get_model_id(messages=edited, tools=None) is None
@pytest.mark.asyncio
async def test_swapped_roles_do_not_inherit_the_pin():
"""The message envelope is part of what the provider caches, so the same blocks under other roles key apart."""
prompt_cache = PromptCachingCache(cache=DualCache())
pinned = _turn(
{"role": "user", "content": [_text("question")]},
{"role": "assistant", "content": [_marked("answer")]},
)
swapped = _turn(
{"role": "assistant", "content": [_text("question")]},
{"role": "user", "content": [_marked("answer")]},
)
await prompt_cache.async_add_model_id(model_id="dep-1", messages=pinned, tools=None)
assert await prompt_cache.async_get_model_id(messages=pinned, tools=None) == {"model_id": "dep-1"}
assert await prompt_cache.async_get_model_id(messages=swapped, tools=None) is None
@pytest.mark.asyncio
async def test_raw_bytes_in_a_block_hash_instead_of_failing_the_request():
"""A block carrying raw bytes must key like any other block rather than raising out of the router filter."""
prompt_cache = PromptCachingCache(cache=DualCache())
binary_block = {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": b"\xff\xfe"}}
turn = _turn({"role": "user", "content": [binary_block, _marked("describe")]})
await prompt_cache.async_add_model_id(model_id="dep-1", messages=turn, tools=None)
assert await prompt_cache.async_get_model_id(messages=turn, tools=None) == {"model_id": "dep-1"}
class _BrokenBatchReadCache(DualCache):
async def async_batch_get_cache(self, keys, parent_otel_span=None, local_only=False, **kwargs):
return None
@ -786,11 +818,11 @@ async def test_claude_code_style_session_stays_on_one_deployment_across_turns(lo
"litellm_params": {"model": AUTO_CACHING_MODEL, "api_key": "sk-fake"},
"model_info": {"id": model_id},
}
for model_id in ("dep-1", "dep-2", "dep-3")
for model_id in (f"dep-{number}" for number in range(1, 7))
],
optional_pre_call_checks=["prompt_caching"],
)
user_turns = [LONG_PROMPT, *(f"follow-up {number}" for number in range(1, 6))]
user_turns = [LONG_PROMPT, *(f"follow-up {number}" for number in range(1, 9))]
history: List[AllMessageValues] = []
served: List[str] = []
for text in user_turns: