mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(caching): defer cache-hit callbacks by replayed result type, not request flags
A converted-stream request whose cache entry is a plain (non-stream) object is replayed as that plain object, so nothing later fires the success callbacks. Decide deferral from the replayed result's type instead of the request kwargs. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
8b86362703
commit
621db91d90
3 changed files with 68 additions and 23 deletions
|
|
@ -114,17 +114,25 @@ def _stream_replay_requested(kwargs: Mapping[str, object]) -> bool:
|
|||
return kwargs.get("stream", False) is True or converted_stream_requested(kwargs)
|
||||
|
||||
|
||||
def _should_defer_streaming_cache_hit_callbacks(*, kwargs: dict[str, object]) -> bool:
|
||||
def _should_defer_streaming_cache_hit_callbacks(*, cached_result: object) -> bool:
|
||||
"""
|
||||
When stream=True, do not run success callbacks at cache-hit time.
|
||||
When the cache hit is replayed as a stream, do not run success callbacks at cache-hit time.
|
||||
|
||||
Cached chat/text completion replay uses CustomStreamWrapper; cached Responses
|
||||
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.
|
||||
spend and callback records. A plain (non-stream) replay logs here, since nothing
|
||||
else will.
|
||||
"""
|
||||
return _stream_replay_requested(kwargs)
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.response_cache import (
|
||||
CachedAnthropicMessagesStreamIterator,
|
||||
)
|
||||
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
|
||||
|
||||
return isinstance(
|
||||
cached_result, (CustomStreamWrapper, BaseResponsesAPIStreamingIterator, CachedAnthropicMessagesStreamIterator)
|
||||
)
|
||||
|
||||
|
||||
def _prompt_tokens_details_as_mapping(details: "PromptTokensDetailsWrapper") -> Mapping[str, object]:
|
||||
|
|
@ -274,7 +282,7 @@ class LLMCachingHandler:
|
|||
custom_llm_provider=kwargs.get("custom_llm_provider", None),
|
||||
args=args,
|
||||
)
|
||||
if not _should_defer_streaming_cache_hit_callbacks(kwargs=kwargs):
|
||||
if not _should_defer_streaming_cache_hit_callbacks(cached_result=cached_result):
|
||||
# LOG SUCCESS
|
||||
self._async_log_cache_hit_on_callbacks(
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -390,7 +398,7 @@ class LLMCachingHandler:
|
|||
is_async=False,
|
||||
)
|
||||
|
||||
if not _should_defer_streaming_cache_hit_callbacks(kwargs=kwargs):
|
||||
if not _should_defer_streaming_cache_hit_callbacks(cached_result=cached_result):
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls(
|
||||
result=cached_result,
|
||||
start_time=start_time,
|
||||
|
|
|
|||
|
|
@ -927,24 +927,14 @@ def test_sync_get_cache_defers_streaming_completion_hit_callbacks():
|
|||
|
||||
|
||||
def test_should_defer_streaming_cache_hit_callbacks_for_any_streaming_request():
|
||||
assert (
|
||||
_should_defer_streaming_cache_hit_callbacks(
|
||||
kwargs={"stream": True},
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
_should_defer_streaming_cache_hit_callbacks(
|
||||
kwargs={"stream": False},
|
||||
)
|
||||
is False
|
||||
)
|
||||
assert (
|
||||
_should_defer_streaming_cache_hit_callbacks(
|
||||
kwargs={},
|
||||
)
|
||||
is False
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
stream_replay = CustomStreamWrapper(
|
||||
completion_stream=iter(()), model="gpt-4o", logging_obj=logging_obj
|
||||
)
|
||||
assert _should_defer_streaming_cache_hit_callbacks(cached_result=stream_replay) is True
|
||||
assert _should_defer_streaming_cache_hit_callbacks(cached_result=ModelResponse()) is False
|
||||
assert _should_defer_streaming_cache_hit_callbacks(cached_result={"id": "msg_1"}) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -693,3 +693,50 @@ async def test_cache_hit_records_the_looked_up_key_as_the_preset_cache_key(monke
|
|||
assert handler.preset_cache_key is not None
|
||||
assert logging_obj.litellm_params["preset_cache_key"] == handler.preset_cache_key
|
||||
assert hit.cached_result._hidden_params["cache_key"] == handler.preset_cache_key
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_converted_stream_cache_hit_replayed_as_plain_object_logs_at_hit_time(monkeypatch):
|
||||
"""A converted-stream Anthropic Messages request that hits a non-stream cache entry gets a plain dict back,
|
||||
so the success callbacks must fire now; nothing else will fire them."""
|
||||
import litellm
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
async def aanthropic_messages(**kwargs):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
|
||||
kwargs = {
|
||||
"model": "claude-sonnet-5",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"max_tokens": 16,
|
||||
"caching": True,
|
||||
"stream": False,
|
||||
"_websearch_interception_converted_stream": True,
|
||||
}
|
||||
cached_message = {
|
||||
"id": "msg_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "hi"}],
|
||||
}
|
||||
await litellm.cache.async_add_cache(cached_message, **kwargs)
|
||||
handler = LLMCachingHandler(original_function=aanthropic_messages, request_kwargs=kwargs, start_time=datetime.now())
|
||||
logging_obj = _build_logging_obj(CallTypes.aanthropic_messages.value, stream=False)
|
||||
logging_obj.async_success_handler = AsyncMock()
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock()
|
||||
|
||||
hit = await handler._async_get_cache(
|
||||
model="claude-sonnet-5",
|
||||
original_function=aanthropic_messages,
|
||||
logging_obj=logging_obj,
|
||||
start_time=datetime.now(),
|
||||
call_type=CallTypes.aanthropic_messages.value,
|
||||
kwargs=kwargs,
|
||||
args=(),
|
||||
)
|
||||
|
||||
assert hit is not None and hit.cached_result == cached_message
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once()
|
||||
assert logging_obj.handle_sync_success_callbacks_for_async_calls.call_args.kwargs["cache_hit"] is True
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue