fix(caching): cache anthropic /v1/messages responses, including streaming

anthropic_messages was missing from the cache's supported call types, so every /v1/messages request went to the provider. Adding it alone is not enough: the cache key is built from the OpenAI-ish param set, which has no system, top_k or stop_sequences, so two requests differing only by system prompt shared an entry and the second got the first one's answer. The Anthropic Messages request shape now feeds the key set as well.

Streaming responses return to the caller before async_set_cache runs, so they are teed on the way out and the SSE events are stored verbatim once the stream reaches message_stop without a provider error. A hit replays those bytes and logs the request as a cache hit with zero cost.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
shivam 2026-07-25 00:09:21 +00:00
parent 79c5c169d8
commit ab997e04eb
10 changed files with 457 additions and 53 deletions

View file

@ -67,20 +67,7 @@ class Cache:
default_in_memory_ttl: Optional[float] = None,
default_in_redis_ttl: Optional[float] = None,
similarity_threshold: Optional[float] = None,
supported_call_types: Optional[List[CachingSupportedCallTypes]] = [
"completion",
"acompletion",
"embedding",
"aembedding",
"atranscription",
"transcription",
"atext_completion",
"text_completion",
"arerank",
"rerank",
"responses",
"aresponses",
],
supported_call_types: list[CachingSupportedCallTypes] | None = list(DEFAULT_CACHING_SUPPORTED_CALL_TYPES),
# s3 Bucket, boto3 configuration
azure_account_url: Optional[str] = None,
azure_blob_container: Optional[str] = None,
@ -930,20 +917,7 @@ def enable_cache(
host: Optional[str] = None,
port: Optional[str] = None,
password: Optional[str] = None,
supported_call_types: Optional[List[CachingSupportedCallTypes]] = [
"completion",
"acompletion",
"embedding",
"aembedding",
"atranscription",
"transcription",
"atext_completion",
"text_completion",
"arerank",
"rerank",
"responses",
"aresponses",
],
supported_call_types: list[CachingSupportedCallTypes] | None = list(DEFAULT_CACHING_SUPPORTED_CALL_TYPES),
**kwargs,
):
"""
@ -990,20 +964,7 @@ def update_cache(
host: Optional[str] = None,
port: Optional[str] = None,
password: Optional[str] = None,
supported_call_types: Optional[List[CachingSupportedCallTypes]] = [
"completion",
"acompletion",
"embedding",
"aembedding",
"atranscription",
"transcription",
"atext_completion",
"text_completion",
"arerank",
"rerank",
"responses",
"aresponses",
],
supported_call_types: list[CachingSupportedCallTypes] | None = list(DEFAULT_CACHING_SUPPORTED_CALL_TYPES),
**kwargs,
):
"""

View file

@ -116,7 +116,8 @@ def _should_defer_streaming_cache_hit_callbacks(*, kwargs: Dict[str, Any]) -> bo
When stream=True, do not run success callbacks at cache-hit time.
Cached chat/text completion replay uses CustomStreamWrapper; cached Responses
replay uses CachedResponsesAPIStreamingIterator. Both invoke logging success
replay uses CachedResponsesAPIStreamingIterator; cached Anthropic Messages
replay uses CachedAnthropicMessagesStreamIterator. All invoke logging success
handlers when the stream finishes; firing them here too would double-count
spend and callback records.
"""
@ -848,6 +849,18 @@ class LLMCachingHandler:
response_type="audio_transcription",
hidden_params=hidden_params,
)
elif (
call_type == CallTypes.anthropic_messages.value or call_type == CallTypes.aanthropic_messages.value
) and isinstance(cached_result, dict):
from litellm.llms.anthropic.experimental_pass_through.messages.response_cache import (
convert_cached_anthropic_messages_result,
)
cached_result = convert_cached_anthropic_messages_result(
cached_result=cached_result,
logging_obj=logging_obj,
kwargs=kwargs,
)
elif (call_type == "aresponses" or call_type == "responses") and isinstance(cached_result, dict):
use_chat_completion_cache = _is_chat_completion_cached_dict(cached_result)
if use_chat_completion_cache:
@ -1044,6 +1057,32 @@ class LLMCachingHandler:
and (kwargs.get("cache", {}).get("no-store", False) is not True)
)
def wrap_streaming_result_for_cache(self, result: Any, call_type: str) -> Any:
"""
Tee a streaming result so it still reaches the cache.
Streaming responses are returned to the caller before ``async_set_cache``
runs. Chat/text completion streams are teed inside ``CustomStreamWrapper``
and Responses API streams inside their own iterator; Anthropic Messages
streams have no such hook, so they are wrapped here.
"""
if call_type not in (
CallTypes.anthropic_messages.value,
CallTypes.aanthropic_messages.value,
):
return result
if litellm.cache is None or not self._should_store_result_in_cache(
original_function=self.original_function, kwargs=self.request_kwargs
):
return result
if not hasattr(result, "__anext__"):
return result
from litellm.llms.anthropic.experimental_pass_through.messages.response_cache import (
AnthropicMessagesStreamCacheWriter,
)
return AnthropicMessagesStreamCacheWriter(stream=result, caching_handler=self)
def _is_call_type_supported_by_cache(
self,
original_function: Callable,

View file

@ -18,6 +18,7 @@ from openai.types.responses.response_create_params import (
)
from litellm._logging import verbose_logger
from litellm.types.llms.anthropic import AnthropicMessagesRequest
from litellm.types.rerank import RerankRequest
@ -40,7 +41,7 @@ class ModelParamHelper:
@staticmethod
def get_exclude_params_for_model_parameters() -> Set[str]:
return set(["messages", "prompt", "input"])
return set(["messages", "prompt", "input", "system"])
@staticmethod
def _get_relevant_args_to_use_for_logging() -> Set[str]:
@ -73,6 +74,7 @@ class ModelParamHelper:
transcription_kwargs = ModelParamHelper._get_litellm_supported_transcription_kwargs()
rerank_kwargs = ModelParamHelper._get_litellm_supported_rerank_kwargs()
responses_api_kwargs = ModelParamHelper._get_litellm_supported_responses_api_kwargs()
anthropic_messages_kwargs = ModelParamHelper._get_litellm_supported_anthropic_messages_kwargs()
exclude_kwargs = ModelParamHelper._get_exclude_kwargs()
combined_kwargs = chat_completion_kwargs.union(
@ -81,6 +83,7 @@ class ModelParamHelper:
transcription_kwargs,
rerank_kwargs,
responses_api_kwargs,
anthropic_messages_kwargs,
)
combined_kwargs = combined_kwargs.difference(exclude_kwargs)
return combined_kwargs
@ -167,12 +170,24 @@ class ModelParamHelper:
streaming_params: Set[str] = set(getattr(ResponseCreateParamsStreaming, "__annotations__", {}).keys())
return non_streaming_params.union(streaming_params)
@staticmethod
def _get_litellm_supported_anthropic_messages_kwargs() -> set[str]:
"""
Get the litellm supported Anthropic /v1/messages kwargs
This follows the Anthropic Messages API spec. `system`, `top_k` and
`stop_sequences` have no OpenAI equivalent, so without them the cache key
for a /v1/messages request ignores them and collides across requests that
differ only by system prompt.
"""
return set(getattr(AnthropicMessagesRequest, "__annotations__", {}).keys())
@staticmethod
def _get_exclude_kwargs() -> Set[str]:
"""
Get the kwargs to exclude from the cache key
"""
return set(["metadata"])
return set(["metadata", "litellm_metadata"])
ModelParamHelper._relevant_logging_args = frozenset(ModelParamHelper._get_relevant_args_to_use_for_logging())

View file

@ -0,0 +1,163 @@
"""
Response caching for Anthropic Messages (`/v1/messages`) requests.
Non-streaming responses are plain dicts and are stored by the generic caching
handler. Streaming responses are returned to the caller before
``LLMCachingHandler.async_set_cache`` runs, so they are teed here instead: the
SSE events are buffered while they are forwarded and persisted verbatim once the
stream completes, and a hit replays exactly what the provider sent.
"""
from collections.abc import AsyncIterator
from typing import TYPE_CHECKING, Any, cast
import litellm
from litellm._logging import verbose_logger
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
BaseAnthropicMessagesStreamingIterator,
_is_message_stop_chunk,
_is_provider_error_chunk,
aclose_if_supported,
)
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
if TYPE_CHECKING:
from litellm.caching.caching_handler import LLMCachingHandler
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
else:
LLMCachingHandler = Any
LiteLLMLoggingObj = Any
CACHED_STREAM_EVENTS_KEY = "litellm_cached_anthropic_sse_events"
def _decode(chunk: bytes | str) -> str:
return chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk
class AnthropicMessagesStreamCacheWriter:
"""
Forwards a `/v1/messages` SSE stream unchanged while buffering it, then
writes the collected events to the response cache on normal completion.
Only a stream that ran to a ``message_stop`` without a provider ``error``
event is written, so partial or failed responses cannot be replayed.
"""
def __init__(
self,
stream: AsyncIterator[bytes | str],
caching_handler: "LLMCachingHandler",
) -> None:
self.stream = stream
self.caching_handler = caching_handler
self.collected_events: list[str] = []
self.saw_message_stop = False
self.saw_provider_error = False
self.persisted = False
self._hidden_params: dict[str, Any] = getattr(stream, "_hidden_params", {}) or {}
def __aiter__(self) -> "AnthropicMessagesStreamCacheWriter":
return self
async def __anext__(self) -> bytes | str:
try:
chunk = await self.stream.__anext__()
except StopAsyncIteration:
await self._persist()
raise
chunk_bytes = chunk.encode("utf-8") if isinstance(chunk, str) else chunk
self.saw_message_stop = self.saw_message_stop or _is_message_stop_chunk(chunk_bytes)
self.saw_provider_error = self.saw_provider_error or _is_provider_error_chunk(chunk_bytes)
self.collected_events.append(_decode(chunk))
return chunk
async def aclose(self) -> None:
await aclose_if_supported(self.stream)
async def _persist(self) -> None:
if self.persisted or litellm.cache is None:
return
if not self.saw_message_stop or self.saw_provider_error:
return
self.persisted = True
request_kwargs = dict(self.caching_handler.request_kwargs)
if not self.caching_handler._should_store_result_in_cache(
original_function=self.caching_handler.original_function,
kwargs=request_kwargs,
):
return
preset_cache_key = self.caching_handler.preset_cache_key
if preset_cache_key is not None:
request_kwargs["cache_key"] = preset_cache_key
try:
await litellm.cache.async_add_cache(
{CACHED_STREAM_EVENTS_KEY: self.collected_events},
dynamic_cache_object=self.caching_handler.dual_cache,
**request_kwargs,
)
except Exception as e:
verbose_logger.exception("Anthropic Messages stream cache write failed: %s", e)
class CachedAnthropicMessagesStreamIterator(BaseAnthropicMessagesStreamingIterator):
"""
Replays cached `/v1/messages` SSE events and logs the request as a cache hit
once the replay finishes, mirroring what the live stream logs at end of stream.
"""
def __init__(
self,
events: list[str],
litellm_logging_obj: LiteLLMLoggingObj,
request_body: dict[str, Any],
) -> None:
super().__init__(litellm_logging_obj=litellm_logging_obj, request_body=request_body)
self.chunks: list[bytes] = [event.encode("utf-8") for event in events]
self.current_index = 0
self._hidden_params: dict[str, Any] = {"cache_hit": True}
litellm_logging_obj.model_call_details["cache_hit"] = True
def __aiter__(self) -> "CachedAnthropicMessagesStreamIterator":
return self
async def __anext__(self) -> bytes:
if self.current_index >= len(self.chunks):
await self._handle_streaming_logging(self.chunks)
raise StopAsyncIteration
chunk = self.chunks[self.current_index]
self.current_index += 1
return chunk
def get_cached_stream_events(cached_result: dict[str, Any]) -> list[str] | None:
events = cached_result.get(CACHED_STREAM_EVENTS_KEY)
if isinstance(events, list):
return [_decode(event) for event in events]
return None
def convert_cached_anthropic_messages_result(
cached_result: dict[str, Any],
logging_obj: LiteLLMLoggingObj,
kwargs: dict[str, Any],
) -> AnthropicMessagesResponse | CachedAnthropicMessagesStreamIterator:
"""
Turn a cached `/v1/messages` entry back into what the caller expects: an
SSE replay iterator for a streamed entry, otherwise the response itself
(``AnthropicMessagesResponse`` is a TypedDict, i.e. a dict at runtime).
"""
events = get_cached_stream_events(cached_result)
if events is not None:
return CachedAnthropicMessagesStreamIterator(
events=events,
litellm_logging_obj=logging_obj,
request_body=kwargs,
)
return cast( # cast-ok: AnthropicMessagesResponse is a TypedDict; validating would drop provider fields we must replay verbatim
AnthropicMessagesResponse, cached_result
)

View file

@ -255,12 +255,16 @@ class AnthropicPassthroughLoggingHandler:
litellm_params=(logging_obj.litellm_params if hasattr(logging_obj, "litellm_params") else None)
)
response_cost = litellm.completion_cost(
completion_response=litellm_model_response,
model=model_for_cost,
custom_llm_provider=custom_llm_provider,
custom_pricing=custom_pricing,
router_model_id=router_model_id,
response_cost = (
0.0
if logging_obj.model_call_details.get("cache_hit") is True
else litellm.completion_cost(
completion_response=litellm_model_response,
model=model_for_cost,
custom_llm_provider=custom_llm_provider,
custom_pricing=custom_pricing,
router_model_id=router_model_id,
)
)
kwargs["response_cost"] = response_cost

View file

@ -161,7 +161,7 @@ class PassThroughStreamingHandler:
result=standard_logging_response_object,
start_time=start_time,
end_time=end_time,
cache_hit=False,
cache_hit=litellm_logging_obj.model_call_details.get("cache_hit") is True,
prefer_async_handlers=True,
**kwargs,
)

View file

@ -30,8 +30,27 @@ CachingSupportedCallTypes = Literal[
"rerank",
"responses",
"aresponses",
"anthropic_messages",
"aanthropic_messages",
]
DEFAULT_CACHING_SUPPORTED_CALL_TYPES: tuple[CachingSupportedCallTypes, ...] = (
"completion",
"acompletion",
"embedding",
"aembedding",
"atranscription",
"transcription",
"atext_completion",
"text_completion",
"arerank",
"rerank",
"responses",
"aresponses",
"anthropic_messages",
"aanthropic_messages",
)
class RedisPipelineIncrementOperation(TypedDict):
"""

View file

@ -1708,7 +1708,10 @@ def client(original_function):
start_time=start_time,
end_time=end_time,
)
return result
return _llm_caching_handler.wrap_streaming_result_for_cache(
result=result,
call_type=call_type,
)
elif call_type == CallTypes.arealtime.value:
return result
### POST-CALL RULES ###

View file

@ -1,6 +1,8 @@
import logging
import re
import pytest
from litellm.caching.caching import Cache
from litellm.types.caching import LiteLLMCacheType
from litellm.types.utils import Embedding, EmbeddingResponse, Usage
@ -146,3 +148,22 @@ def test_exact_cache_key_still_includes_prompt():
model="gpt-4o-mini", messages=[{"role": "user", "content": "b"}]
)
assert key_a != key_b
@pytest.mark.parametrize(
"anthropic_param",
[
{"system": "answer ALPHA"},
{"top_k": 5},
{"stop_sequences": ["STOP"]},
],
)
def test_exact_cache_key_includes_anthropic_messages_params(anthropic_param):
"""Anthropic /v1/messages params with no OpenAI equivalent must still key the
cache; without them two requests that differ only by system prompt collide."""
cache = Cache(type=LiteLLMCacheType.LOCAL)
messages = [{"role": "user", "content": "which greek letter?"}]
baseline = cache.get_cache_key(model="claude-sonnet-4-5", messages=messages)
assert baseline != cache.get_cache_key(
model="claude-sonnet-4-5", messages=messages, **anthropic_param
)

View file

@ -0,0 +1,179 @@
import asyncio
import os
import sys
from typing import Any, AsyncIterator, Dict, List
import pytest
sys.path.insert(0, os.path.abspath("../../../../.."))
import litellm
from litellm.caching.caching import Cache, LiteLLMCacheType
from litellm.llms.anthropic.experimental_pass_through.messages import handler
STREAM_EVENTS: List[bytes] = [
b'event: message_start\ndata: {"type": "message_start", "message": {"id": "msg_stream_1", "type": "message", '
b'"role": "assistant", "model": "claude-sonnet-4-5", "content": [], "stop_reason": null, '
b'"usage": {"input_tokens": 10, "output_tokens": 0}}}\n\n',
b'event: content_block_start\ndata: {"type": "content_block_start", "index": 0, '
b'"content_block": {"type": "text", "text": ""}}\n\n',
b'event: content_block_delta\ndata: {"type": "content_block_delta", "index": 0, '
b'"delta": {"type": "text_delta", "text": "ALPHA"}}\n\n',
b'event: content_block_stop\ndata: {"type": "content_block_stop", "index": 0}\n\n',
b'event: message_delta\ndata: {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, '
b'"usage": {"output_tokens": 3}}\n\n',
b'event: message_stop\ndata: {"type": "message_stop"}\n\n',
]
def _anthropic_response(message_id: str, text: str) -> Dict[str, Any]:
return {
"id": message_id,
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-5",
"content": [{"type": "text", "text": text}],
"stop_reason": "end_turn",
"usage": {"input_tokens": 10, "output_tokens": 3},
}
class _CountingHandler:
"""Stands in for the provider dispatch so cache hits are observable as skipped calls."""
def __init__(self, results: List[Any]) -> None:
self.results = results
self.calls: List[Dict[str, Any]] = []
def __call__(self, *args: Any, **kwargs: Any) -> Any:
self.calls.append(kwargs)
return self.results[min(len(self.calls) - 1, len(self.results) - 1)]
async def _byte_stream(chunks: List[bytes]) -> AsyncIterator[bytes]:
for chunk in chunks:
yield chunk
async def _collect(stream: AsyncIterator[bytes]) -> List[bytes]:
return [chunk async for chunk in stream]
@pytest.fixture
def local_cache():
previous_cache = litellm.cache
litellm.cache = Cache(type=LiteLLMCacheType.LOCAL)
yield litellm.cache
litellm.cache = previous_cache
@pytest.fixture
def request_kwargs() -> Dict[str, Any]:
return {
"model": "anthropic/claude-sonnet-4-5",
"custom_llm_provider": "anthropic",
"api_key": "fake-key",
"max_tokens": 64,
"messages": [{"role": "user", "content": "which greek letter?"}],
}
@pytest.mark.asyncio
async def test_non_streaming_request_is_served_from_cache(local_cache, request_kwargs, monkeypatch):
fake_handler = _CountingHandler([_anthropic_response("msg_1", "ALPHA"), _anthropic_response("msg_2", "BETA")])
monkeypatch.setattr(handler, "anthropic_messages_handler", fake_handler)
first = await litellm.anthropic_messages(**request_kwargs)
await asyncio.sleep(0)
second = await litellm.anthropic_messages(**request_kwargs)
assert len(fake_handler.calls) == 1
assert first == second
assert second["content"][0]["text"] == "ALPHA"
@pytest.mark.asyncio
async def test_cache_key_separates_different_system_prompts(local_cache, request_kwargs, monkeypatch):
"""`system` has no OpenAI equivalent; if it is dropped from the cache key the
second request is answered with the first system prompt's response."""
fake_handler = _CountingHandler([_anthropic_response("msg_1", "ALPHA"), _anthropic_response("msg_2", "BETA")])
monkeypatch.setattr(handler, "anthropic_messages_handler", fake_handler)
first = await litellm.anthropic_messages(**request_kwargs, system="Always answer ALPHA")
await asyncio.sleep(0)
second = await litellm.anthropic_messages(**request_kwargs, system="Always answer BETA")
assert len(fake_handler.calls) == 2
assert first["content"][0]["text"] == "ALPHA"
assert second["content"][0]["text"] == "BETA"
@pytest.mark.parametrize("anthropic_param", [{"top_k": 5}, {"stop_sequences": ["STOP"]}])
@pytest.mark.asyncio
async def test_cache_key_separates_anthropic_native_params(local_cache, request_kwargs, monkeypatch, anthropic_param):
fake_handler = _CountingHandler([_anthropic_response("msg_1", "ALPHA"), _anthropic_response("msg_2", "BETA")])
monkeypatch.setattr(handler, "anthropic_messages_handler", fake_handler)
await litellm.anthropic_messages(**request_kwargs)
await asyncio.sleep(0)
await litellm.anthropic_messages(**request_kwargs, **anthropic_param)
assert len(fake_handler.calls) == 2
@pytest.mark.asyncio
async def test_streaming_request_is_replayed_from_cache(local_cache, request_kwargs, monkeypatch):
fake_handler = _CountingHandler([_byte_stream(STREAM_EVENTS), _byte_stream([b"event: never_used\n\n"])])
monkeypatch.setattr(handler, "anthropic_messages_handler", fake_handler)
first = await _collect(await litellm.anthropic_messages(**request_kwargs, stream=True))
second_stream = await litellm.anthropic_messages(**request_kwargs, stream=True)
second = await _collect(second_stream)
assert len(fake_handler.calls) == 1
assert first == STREAM_EVENTS
assert second == STREAM_EVENTS
assert second_stream._hidden_params["cache_hit"] is True
@pytest.mark.asyncio
async def test_streaming_cache_is_not_shared_with_non_streaming(local_cache, request_kwargs, monkeypatch):
fake_handler = _CountingHandler([_byte_stream(STREAM_EVENTS), _anthropic_response("msg_2", "ALPHA")])
monkeypatch.setattr(handler, "anthropic_messages_handler", fake_handler)
await _collect(await litellm.anthropic_messages(**request_kwargs, stream=True))
non_streaming = await litellm.anthropic_messages(**request_kwargs)
assert len(fake_handler.calls) == 2
assert non_streaming["content"][0]["text"] == "ALPHA"
@pytest.mark.asyncio
async def test_failed_stream_is_not_cached(local_cache, request_kwargs, monkeypatch):
error_events = STREAM_EVENTS[:3] + [
b'event: error\ndata: {"type": "error", "error": {"type": "overloaded_error", "message": "overloaded"}}\n\n'
]
fake_handler = _CountingHandler([_byte_stream(error_events), _byte_stream(STREAM_EVENTS)])
monkeypatch.setattr(handler, "anthropic_messages_handler", fake_handler)
failed = await _collect(await litellm.anthropic_messages(**request_kwargs, stream=True))
replayed = await _collect(await litellm.anthropic_messages(**request_kwargs, stream=True))
assert failed == error_events
assert len(fake_handler.calls) == 2
assert replayed == STREAM_EVENTS
@pytest.mark.asyncio
async def test_abandoned_stream_is_not_cached(local_cache, request_kwargs, monkeypatch):
fake_handler = _CountingHandler([_byte_stream(STREAM_EVENTS), _byte_stream(STREAM_EVENTS)])
monkeypatch.setattr(handler, "anthropic_messages_handler", fake_handler)
partial_stream = await litellm.anthropic_messages(**request_kwargs, stream=True)
await partial_stream.__anext__()
await partial_stream.aclose()
replayed = await _collect(await litellm.anthropic_messages(**request_kwargs, stream=True))
assert len(fake_handler.calls) == 2
assert replayed == STREAM_EVENTS