fix(caching): replay cache hits for converted streams as streams

A deployment hook (Headroom, code interpreter, web search) can downgrade
kwargs["stream"] to False while the caller still expects to iterate the
result. The cache handler keyed stream replay and callback deferral off
the raw flag, so a cache hit returned a plain object to a caller that
iterates, and the Responses iterator never persisted the converted
stream in the first place. Key both off the conversion marker as well

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-15 01:42:49 +00:00
parent 95ef538789
commit 8b86362703
4 changed files with 126 additions and 12 deletions

View file

@ -35,6 +35,7 @@ from litellm.litellm_core_utils.logging_utils import (
_assemble_complete_response_from_streaming_chunks,
)
from litellm.types.caching import CachedEmbedding
from litellm.types.integrations.custom_logger import converted_stream_requested
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.rerank import RerankResponse
from litellm.types.utils import (
@ -107,6 +108,12 @@ def _is_chat_completion_cached_dict(cached_result: dict) -> bool:
return "choices" in cached_result
def _stream_replay_requested(kwargs: Mapping[str, object]) -> bool:
"""True when the caller must receive a stream, including when a deployment hook downgraded
`kwargs["stream"]` to False for the provider call."""
return kwargs.get("stream", False) is True or converted_stream_requested(kwargs)
def _should_defer_streaming_cache_hit_callbacks(*, kwargs: dict[str, object]) -> bool:
"""
When stream=True, do not run success callbacks at cache-hit time.
@ -117,7 +124,7 @@ def _should_defer_streaming_cache_hit_callbacks(*, kwargs: dict[str, object]) ->
handlers when the stream finishes; firing them here too would double-count
spend and callback records.
"""
return kwargs.get("stream", False) is True
return _stream_replay_requested(kwargs)
def _prompt_tokens_details_as_mapping(details: "PromptTokensDetailsWrapper") -> Mapping[str, object]:
@ -823,7 +830,7 @@ class LLMCachingHandler:
if (call_type == CallTypes.acompletion.value or call_type == CallTypes.completion.value) and isinstance(
cached_result, dict
):
if kwargs.get("stream", False) is True:
if _stream_replay_requested(kwargs):
cached_result = self._convert_cached_stream_response(
cached_result=cached_result,
call_type=call_type,
@ -838,7 +845,7 @@ class LLMCachingHandler:
if (
call_type == CallTypes.atext_completion.value or call_type == CallTypes.text_completion.value
) and isinstance(cached_result, dict):
if kwargs.get("stream", False) is True:
if _stream_replay_requested(kwargs):
cached_result = self._convert_cached_stream_response(
cached_result=cached_result,
call_type=call_type,
@ -893,7 +900,7 @@ class LLMCachingHandler:
elif (call_type == "aresponses" or call_type == "responses") and isinstance(cached_result, dict):
use_chat_completion_cache: Final = _is_chat_completion_cached_dict(cached_result)
if use_chat_completion_cache:
if kwargs.get("stream", False) is True:
if _stream_replay_requested(kwargs):
bridge_call_type: Final = (
CallTypes.acompletion.value if call_type == "aresponses" else CallTypes.completion.value
)
@ -921,7 +928,7 @@ class LLMCachingHandler:
):
response_obj._hidden_params["cache_hit"] = True
if kwargs.get("stream", False) is True:
if _stream_replay_requested(kwargs):
cached_result = CachedResponsesAPIStreamingIterator(
response=response_obj,
logging_obj=logging_obj,

View file

@ -33,6 +33,7 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
from litellm.litellm_core_utils.thread_pool_executor import executor
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.responses.utils import ResponseAPILoggingUtils, ResponsesAPIRequestUtils
from litellm.types.integrations.custom_logger import converted_stream_requested
from litellm.types.llms.openai import (
PART_UNION_TYPES,
ResponseAPIUsage,
@ -626,7 +627,9 @@ class BaseResponsesAPIStreamingIterator:
return
request_kwargs = getattr(caching_handler, "request_kwargs", None)
if not _is_json_object(request_kwargs) or request_kwargs.get("stream") is not True:
if not _is_json_object(request_kwargs):
return
if request_kwargs.get("stream") is not True and not converted_stream_requested(request_kwargs):
return
request_kwargs = request_kwargs.copy()
preset_cache_key = getattr(caching_handler, "preset_cache_key", None)

View file

@ -855,6 +855,11 @@ def _is_converted_stream_result(result: object) -> bool:
return isinstance(result, (CustomStreamWrapper, BaseResponsesAPIStreamingIterator))
def _mark_logging_as_stream(logging_obj: LiteLLMLoggingObject) -> None:
logging_obj.stream = True
logging_obj.model_call_details["stream"] = True
# Runs once per call to check if the user wants to send their data anywhere - PostHog/Sentry/Slack/etc.
def function_setup(
original_function: str,
@ -1898,6 +1903,8 @@ def client(original_function):
_caching_handler_response.cached_result is not None
and _caching_handler_response.final_embedding_cached_response is None
):
if _is_converted_stream_result(_caching_handler_response.cached_result):
_mark_logging_as_stream(logging_obj)
return _caching_handler_response.cached_result
elif _caching_handler_response.embedding_all_elements_cache_hit is True:
@ -1956,8 +1963,7 @@ def client(original_function):
end_time = datetime.datetime.now()
if _is_streaming_request(kwargs=kwargs, call_type=call_type) or _is_converted_stream_result(result):
logging_obj.stream = True
logging_obj.model_call_details["stream"] = True
_mark_logging_as_stream(logging_obj)
if "complete_response" in kwargs and kwargs["complete_response"] is True:
chunks: Final = []
for idx, chunk in enumerate(result):

View file

@ -20,6 +20,8 @@ from jsonschema import validate
import litellm
from litellm._internal_context import is_internal_call
from litellm.caching.caching import Cache
from litellm.caching.caching_handler import _PENDING_CACHE_WRITES
from litellm.constants import DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT
from litellm._logging import (
CorrelationContextFilter,
@ -5324,12 +5326,18 @@ class _SuccessKwargsCapture(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.success_kwargs: list[dict[str, object]] = []
self.stream_event_responses: list[object] = []
async def async_log_success_event(
self, kwargs: dict[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self.success_kwargs.append(kwargs)
async def async_log_stream_event(
self, kwargs: dict[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self.stream_event_responses.append(response_obj)
def _install_converted_stream_callbacks(monkeypatch: pytest.MonkeyPatch) -> _SuccessKwargsCapture:
capture: Final = _SuccessKwargsCapture()
@ -5341,13 +5349,23 @@ def _install_converted_stream_callbacks(monkeypatch: pytest.MonkeyPatch) -> _Suc
return capture
async def _wait_for_success_kwargs(capture: _SuccessKwargsCapture) -> dict[str, object]:
async def _wait_for_success_kwargs(capture: _SuccessKwargsCapture, count: int = 1) -> dict[str, object]:
for _ in range(50):
if capture.success_kwargs:
if len(capture.success_kwargs) >= count and not _PENDING_CACHE_WRITES:
break
await asyncio.sleep(0.05)
(success_kwargs,) = capture.success_kwargs
return success_kwargs
await asyncio.sleep(0.2)
assert len(capture.success_kwargs) == count
return capture.success_kwargs[-1]
def _assert_cache_hit_logged_as_stream(capture: _SuccessKwargsCapture, success_kwargs: dict[str, object]) -> None:
standard_logging_object: Final = success_kwargs["standard_logging_object"]
assert isinstance(standard_logging_object, dict)
assert standard_logging_object["cache_hit"] is True
assert standard_logging_object["stream"] is True
assert success_kwargs["stream"] is True
assert capture.stream_event_responses == []
@pytest.mark.asyncio
@ -5424,6 +5442,86 @@ async def test_wrapper_async_logs_converted_responses_stream_with_standard_loggi
assert success_kwargs["stream"] is True
@pytest.mark.asyncio
async def test_wrapper_async_replays_cached_converted_chat_stream_as_stream(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A cache hit for a converted stream must replay as a stream: the caller still iterates the
result even though the deployment hook set kwargs["stream"] to False."""
capture: Final = _install_converted_stream_callbacks(monkeypatch)
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
request: Final = {
"model": "gpt-5.6",
"messages": [{"role": "user", "content": "replay me from cache"}],
"stream": True,
"mock_response": "converted stream body",
"num_retries": 0,
}
first: Final = await litellm.acompletion(**request)
first_chunks: Final = [chunk async for chunk in first]
assert "".join(chunk.choices[0].delta.content or "" for chunk in first_chunks) == "converted stream body"
await _wait_for_success_kwargs(capture)
replay: Final = await litellm.acompletion(**request)
assert isinstance(replay, CustomStreamWrapper)
replay_chunks: Final = [chunk async for chunk in replay]
assert "".join(chunk.choices[0].delta.content or "" for chunk in replay_chunks) == "converted stream body"
_assert_cache_hit_logged_as_stream(capture, await _wait_for_success_kwargs(capture, count=2))
@pytest.mark.asyncio
@respx.mock
async def test_wrapper_async_replays_cached_converted_responses_stream_as_stream(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Responses surface of the cache-hit replay: the hit must come back as a streaming iterator."""
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
capture: Final = _install_converted_stream_callbacks(monkeypatch)
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
route: Final = respx.post("https://api.openai.com/v1/responses").respond(
json={
"id": "resp_cached_converted",
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-5.6",
"output": [
{
"type": "message",
"id": "msg_cached_converted",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "converted stream body", "annotations": []}],
}
],
"usage": {"input_tokens": 3, "output_tokens": 4, "total_tokens": 7},
}
)
request: Final = {
"model": "openai/gpt-5.6",
"input": "replay me from cache",
"stream": True,
"api_key": "sk-test",
"num_retries": 0,
}
first: Final = await litellm.aresponses(**request)
assert [event async for event in first][-1].type == "response.completed"
await _wait_for_success_kwargs(capture)
replay: Final = await litellm.aresponses(**request)
assert isinstance(replay, BaseResponsesAPIStreamingIterator)
assert [event async for event in replay][-1].type == "response.completed"
assert route.call_count == 1
_assert_cache_hit_logged_as_stream(capture, await _wait_for_success_kwargs(capture, count=2))
def test_function_setup_failure_after_logging_construction_restores_context(monkeypatch):
"""If function_setup() constructs Logging() (which already mutated
trace_id_var/session_id_var in __init__) but then raises before returning,