fix(router): respect prompt cache affinity ttl

Derive the router affinity cache TTL from the cacheable prompt prefix so one-hour ephemeral cache hints stay routable for their advertised lifetime instead of expiring after the previous five-minute default.

Related: BerriAI/litellm#28427

Tested: python -m pytest tests/router_unit_tests/test_router_prompt_caching.py -q; python -m pytest tests/test_litellm/test_prompt_caching_cache.py -q; python -m ruff check litellm/router_utils/prompt_caching_cache.py tests/router_unit_tests/test_router_prompt_caching.py tests/test_litellm/test_prompt_caching_cache.py; python -m black litellm/router_utils/prompt_caching_cache.py tests/router_unit_tests/test_router_prompt_caching.py tests/test_litellm/test_prompt_caching_cache.py

Co-authored-by: OmX <omx@oh-my-codex.dev>
This commit is contained in:
ByteWise 2026-05-21 17:32:02 +08:00
parent 697a90ea77
commit 14b4e9c0c3
3 changed files with 178 additions and 11 deletions

View file

@ -12,6 +12,12 @@ from litellm.caching.caching import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
PROMPT_CACHING_DEFAULT_TTL_SECONDS = 300
PROMPT_CACHING_CACHE_CONTROL_TTL_SECONDS = {
"5m": 300,
"1h": 3600,
}
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -179,6 +185,53 @@ class PromptCachingCache:
hashed_data = hashlib.sha256(data_to_hash_str.encode()).hexdigest()
return f"deployment:{hashed_data}:prompt_caching"
@staticmethod
def _get_cache_control_ttl_seconds(cache_control: Any) -> Optional[int]:
if (
not isinstance(cache_control, dict)
or cache_control.get("type") != "ephemeral"
):
return None
ttl = cache_control.get("ttl")
if not isinstance(ttl, str):
return None
return PROMPT_CACHING_CACHE_CONTROL_TTL_SECONDS.get(ttl)
@staticmethod
def get_prompt_caching_ttl_seconds(
messages: Optional[List[AllMessageValues]],
) -> int:
if messages is None:
return PROMPT_CACHING_DEFAULT_TTL_SECONDS
cacheable_messages = PromptCachingCache.extract_cacheable_prefix(messages)
ttl_seconds = PROMPT_CACHING_DEFAULT_TTL_SECONDS
for message in cacheable_messages:
message_ttl_seconds = PromptCachingCache._get_cache_control_ttl_seconds(
message.get("cache_control")
)
if message_ttl_seconds is not None:
ttl_seconds = max(ttl_seconds, message_ttl_seconds)
content = message.get("content")
if not isinstance(content, list):
continue
for content_block in content:
if not isinstance(content_block, dict):
continue
block_ttl_seconds = PromptCachingCache._get_cache_control_ttl_seconds(
content_block.get("cache_control")
)
if block_ttl_seconds is not None:
ttl_seconds = max(ttl_seconds, block_ttl_seconds)
return ttl_seconds
def add_model_id(
self,
model_id: str,
@ -193,8 +246,11 @@ class PromptCachingCache:
if cache_key is None:
return None
ttl_seconds = PromptCachingCache.get_prompt_caching_ttl_seconds(messages)
self.cache.set_cache(
cache_key, PromptCachingCacheValue(model_id=model_id), ttl=300
cache_key,
PromptCachingCacheValue(model_id=model_id),
ttl=ttl_seconds,
)
return None
@ -212,10 +268,11 @@ class PromptCachingCache:
if cache_key is None:
return None
ttl_seconds = PromptCachingCache.get_prompt_caching_ttl_seconds(messages)
await self.cache.async_set_cache(
cache_key,
PromptCachingCacheValue(model_id=model_id),
ttl=300, # store for 5 minutes
ttl=ttl_seconds,
)
return None

View file

@ -1,21 +1,12 @@
import sys
import os
import traceback
import asyncio
from dotenv import load_dotenv
from fastapi import Request
from datetime import datetime
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
from litellm import Router
import pytest
import litellm
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
@ -25,6 +16,18 @@ class ExampleModel(BaseModel):
field2: int
class RecordingCache:
def __init__(self):
self.set_calls = []
self.async_set_calls = []
def set_cache(self, key, value, **kwargs):
self.set_calls.append({"key": key, "value": value, "kwargs": kwargs})
async def async_set_cache(self, key, value, **kwargs):
self.async_set_calls.append({"key": key, "value": value, "kwargs": kwargs})
def test_serialize_pydantic_object():
model = ExampleModel(field1="value", field2=42)
serialized = PromptCachingCache.serialize_object(model)
@ -281,3 +284,91 @@ def test_extract_cacheable_prefix_mixed_string_and_list_content():
assert result[0]["role"] == "system"
assert result[1]["content"] == "First cached message"
assert isinstance(result[2]["content"], list)
def test_prompt_caching_ttl_uses_default_for_five_minute_cache_control():
messages = [
{"role": "system", "content": "You are a helpful assistant"},
{
"role": "user",
"content": "Cacheable content",
"cache_control": {"type": "ephemeral", "ttl": "5m"},
},
]
assert PromptCachingCache.get_prompt_caching_ttl_seconds(messages) == 300
def test_prompt_caching_ttl_uses_one_hour_for_message_level_cache_control():
messages = [
{"role": "system", "content": "You are a helpful assistant"},
{
"role": "user",
"content": "Long-lived cacheable content",
"cache_control": {"type": "ephemeral", "ttl": "1h"},
},
]
assert PromptCachingCache.get_prompt_caching_ttl_seconds(messages) == 3600
def test_prompt_caching_ttl_uses_one_hour_for_content_block_cache_control():
messages = [
{
"role": "system",
"content": [
{
"type": "text",
"text": "Short-lived content",
"cache_control": {"type": "ephemeral", "ttl": "5m"},
},
{
"type": "text",
"text": "Long-lived content",
"cache_control": {"type": "ephemeral", "ttl": "1h"},
},
],
},
{"role": "user", "content": "This is outside the cacheable prefix"},
]
assert PromptCachingCache.get_prompt_caching_ttl_seconds(messages) == 3600
def test_add_model_id_writes_prompt_cache_with_effective_ttl():
messages = [
{
"role": "user",
"content": "Long-lived cacheable content",
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
]
recording_cache = RecordingCache()
PromptCachingCache(cache=recording_cache).add_model_id(
model_id="deployment-1",
messages=messages,
tools=None,
)
assert recording_cache.set_calls[0]["kwargs"]["ttl"] == 3600
@pytest.mark.asyncio
async def test_async_add_model_id_writes_prompt_cache_with_effective_ttl():
messages = [
{
"role": "user",
"content": "Long-lived cacheable content",
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
]
recording_cache = RecordingCache()
await PromptCachingCache(cache=recording_cache).async_add_model_id(
model_id="deployment-1",
messages=messages,
tools=None,
)
assert recording_cache.async_set_calls[0]["kwargs"]["ttl"] == 3600

View file

@ -0,0 +1,19 @@
from litellm.router_utils.prompt_caching_cache import PromptCachingCache
def test_prompt_caching_affinity_ttl_respects_one_hour_cache_control():
messages = [
{
"role": "system",
"content": [
{
"type": "text",
"text": "Long-lived prompt cache prefix",
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
],
},
{"role": "user", "content": "This is outside the cacheable prefix"},
]
assert PromptCachingCache.get_prompt_caching_ttl_seconds(messages) == 3600