fix(streaming): word-sliced cache replay for stream=true cache hits (#30216)

* fix(streaming): word-sliced cache replay for stream=true cache hits

* fix(streaming): align mypy and replay happy-path test with word-sliced cache replay

* fix(streaming): short-circuit whitespace-only content in cache replay splitter

* fix(streaming): emit tool_calls/function_call only on first replay slice

* refactor(streaming): drop dead delattr guard in cache replay

A non-None usage on the replay base object always lives in
__pydantic_extra__ (it is attached via setattr earlier in the same
function), so delattr can never raise here; the try/except AttributeError
that silently swallowed a failure was dead defensive code that could only
ever hide a real regression, so it is removed in both the async and sync
generators.

Also switches the new replay annotations from typing.List to the builtin
list to satisfy the strict ruff UP006 gate and drops the unused
PLR0915 noqa directives (the rule is not enabled in this repo's ruff
config, so RUF100 flagged them).

* fix(streaming): drop carried-over metadata from later cache replay slices

The word-sliced cache replay deep-copies the full ModelResponseStream per
slice, so reasoning_content, thinking_blocks, logprobs, enhancements,
annotations and the rest of the per-message metadata rode on every slice, not
just the first. Downstream handlers that accumulate streamed deltas would
collect each one once per slice, e.g. duplicating a cached reasoning trace N
times on a stream=true cache hit.

Later slices are now rebuilt as a content-only delta with choice-level logprobs
and enhancements stripped, so the whole metadata class stays on the first slice.
Adds async (logprobs) and sync (reasoning_content/thinking_blocks/logprobs/
enhancements, plus annotations) regression tests

---------

Co-authored-by: Mateo <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
michelligabriele 2026-06-25 16:13:05 +02:00 • committed by GitHub
parent c712c20d0f
commit 0a8a87afe0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 399 additions and 12 deletions

View file

@ -1,5 +1,6 @@
import asyncio
import json
import re
import time
import traceback
from typing import Dict, Iterable, List, Literal, Optional, Tuple, Union, cast
@ -118,7 +119,44 @@ def convert_tool_call_to_json_mode(
return None, None
async def convert_to_streaming_response_async(response_object: Optional[dict] = None):
# Whitespace-preserving word splitter used by the cache-hit replay generators.
# Each match is any leading whitespace plus a non-whitespace run plus any
# trailing whitespace, so concatenating the matches losslessly reconstructs
# the original string (including content that starts with whitespace).
_REPLAY_CONTENT_SLICE_RE = re.compile(r"\s*\S+\s*", re.UNICODE)
def _split_assembled_content_for_replay(content: Optional[str]) -> list[str]:
"""
Slice an assembled cached completion's ``content`` into word-shaped pieces
for cadence-preserving streaming replay. The split is lossless:
``"".join(_split_assembled_content_for_replay(s)) == s`` for every
non-empty ``s``. Returns ``[]`` for ``None`` / empty / all-whitespace
content.
"""
if not content or content.isspace():
# isspace() guard: on all-whitespace content the regex backtracks
# quadratically before returning no matches.
return []
return _REPLAY_CONTENT_SLICE_RE.findall(content)
def _clear_later_replay_slice_metadata(choice: StreamingChoices) -> None:
# Rebuild the delta as content-only so every accumulate-able field (role,
# tool_calls, reasoning_content, thinking_blocks, audio, images,
# annotations, ...) is dropped on later slices instead of an enumerated
# subset; repeating any of them makes downstream handlers that accumulate
# streamed deltas collect it once per slice, and a field added to Delta
# later can't silently re-introduce the duplication.
choice.delta = Delta(content=choice.delta.content)
choice.logprobs = None # type: ignore[assignment]
if hasattr(choice, "enhancements"):
del choice.enhancements
async def convert_to_streaming_response_async(
response_object: Optional[dict] = None,
):
"""
Asynchronously converts a response object to a streaming response.
@ -215,11 +253,45 @@ async def convert_to_streaming_response_async(response_object: Optional[dict] =
if "model" in response_object:
model_response_object.model = response_object["model"]
yield model_response_object
await asyncio.sleep(0)
# Replay cached content with per-word cadence so stream=true cache hits
# don't arrive as a single SSE frame. Multi-choice (n>1) responses and
# unsplittable content (None/empty/whitespace-free) keep the original
# single-yield behavior.
slices: list[str] = []
if len(model_response_object.choices) == 1:
slices = _split_assembled_content_for_replay(
model_response_object.choices[0].delta.content
)
if len(slices) <= 1:
yield model_response_object
await asyncio.sleep(0)
return
# Detach usage from the base object so we can re-attach it only to the
# final slice chunk. A non-None usage always lives in __pydantic_extra__
# here (set via setattr above), so delattr cannot fail.
original_usage = getattr(model_response_object, "usage", None)
if original_usage is not None:
delattr(model_response_object, "usage")
original_finish_reason = model_response_object.choices[0].finish_reason
last_idx = len(slices) - 1
for i, piece in enumerate(slices):
slice_chunk = model_response_object.model_copy(deep=True)
slice_chunk.choices[0].delta.content = piece
if i > 0:
_clear_later_replay_slice_metadata(slice_chunk.choices[0])
slice_chunk.choices[0].finish_reason = (
original_finish_reason if i == last_idx else None # type: ignore[assignment]
)
if i == last_idx and original_usage is not None:
setattr(slice_chunk, "usage", original_usage)
yield slice_chunk
await asyncio.sleep(0)
def convert_to_streaming_response(response_object: Optional[dict] = None):
def convert_to_streaming_response(
response_object: Optional[dict] = None,
):
# used for yielding Cache hits when stream == True
if response_object is None:
raise Exception("Error in response object format")
@ -278,7 +350,35 @@ def convert_to_streaming_response(response_object: Optional[dict] = None):
if "model" in response_object:
model_response_object.model = response_object["model"]
yield model_response_object
# Replay cached content with per-word cadence on sync cache-hit paths
# (S3Cache, sync completion()). See convert_to_streaming_response_async
# for the full rationale — this mirrors its tail.
slices: list[str] = []
if len(model_response_object.choices) == 1:
slices = _split_assembled_content_for_replay(
model_response_object.choices[0].delta.content
)
if len(slices) <= 1:
yield model_response_object
return
original_usage = getattr(model_response_object, "usage", None)
if original_usage is not None:
delattr(model_response_object, "usage")
original_finish_reason = model_response_object.choices[0].finish_reason
last_idx = len(slices) - 1
for i, piece in enumerate(slices):
slice_chunk = model_response_object.model_copy(deep=True)
slice_chunk.choices[0].delta.content = piece
if i > 0:
_clear_later_replay_slice_metadata(slice_chunk.choices[0])
slice_chunk.choices[0].finish_reason = (
original_finish_reason if i == last_idx else None # type: ignore[assignment]
)
if i == last_idx and original_usage is not None:
setattr(slice_chunk, "usage", original_usage)
yield slice_chunk
from collections import defaultdict

View file

@ -1438,10 +1438,11 @@ class CustomStreamWrapper:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider == "cached_response":
chunk = cast(ModelResponseStream, chunk)
chunk_finish_reason = chunk.choices[0].finish_reason
response_obj = {
"text": chunk.choices[0].delta.content,
"is_finished": True,
"finish_reason": chunk.choices[0].finish_reason,
"is_finished": chunk_finish_reason is not None,
"finish_reason": chunk_finish_reason,
"original_chunk": chunk,
"tool_calls": (
chunk.choices[0].delta.tool_calls

View file

@ -1988,11 +1988,15 @@ class TestConvertToStreamingResponseAsync:
return chunks
chunks = asyncio.run(run())
assert len(chunks) == 1
assert chunks[0].id == "msg_async_1"
assert chunks[0].model == "claude-3"
assert chunks[0].choices[0].delta.content == "Hi there"
assert chunks[0].usage.prompt_tokens == 3
# Cached replay is sliced into word-shaped chunks to preserve
# streaming cadence; joining the slices reconstructs the content.
assert len(chunks) == 2
assert all(c.id == "msg_async_1" for c in chunks)
assert all(c.model == "claude-3" for c in chunks)
assert "".join(c.choices[0].delta.content or "" for c in chunks) == "Hi there"
assert chunks[0].choices[0].finish_reason is None
assert chunks[-1].choices[0].finish_reason == "stop"
assert chunks[-1].usage.prompt_tokens == 3
class TestHandleInvalidParallelToolCalls:

View file

@ -0,0 +1,282 @@
"""
Tests for the cache-hit replay generators in
``litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response``.
These generators are used by ``LLMCachingHandler._convert_cached_stream_response``
to replay a cached non-streaming ``ModelResponse`` as a stream when the
incoming request has ``stream=True``. The fix in this test file ensures the
replay yields multiple word-shaped chunks instead of a single one-shot
content frame, restoring per-token cadence on cache hits.
"""
import pytest
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
_split_assembled_content_for_replay,
convert_to_streaming_response,
convert_to_streaming_response_async,
)
from litellm.types.utils import ModelResponseStream
def _async_payload(content="Hello world! How are you?"):
return {
"id": "chatcmpl-test-async",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4o-mini",
"system_fingerprint": "fp_test",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": content},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 5,
"completion_tokens": 7,
"total_tokens": 12,
},
}
async def _collect_async(payload):
return [chunk async for chunk in convert_to_streaming_response_async(payload)]
# ---------- helper ----------
@pytest.mark.parametrize(
"text",
[
"Hello world!",
"Sure! Here's a list of 25 fruits:\n\n1. Apple\n2. Banana\n3. Orange\n",
" leading whitespace matters",
"Hi",
"你好世界",
],
)
def test_split_is_lossless(text):
assert "".join(_split_assembled_content_for_replay(text)) == text
def test_split_returns_empty_for_none_and_empty():
assert _split_assembled_content_for_replay(None) == []
assert _split_assembled_content_for_replay("") == []
def test_split_returns_empty_for_whitespace_only():
# Must short-circuit before the regex: findall backtracks quadratically
# on all-whitespace input.
assert _split_assembled_content_for_replay(" ") == []
assert _split_assembled_content_for_replay(" \n\t" * 10000) == []
# ---------- async generator ----------
@pytest.mark.asyncio
async def test_async_yields_multiple_content_chunks_with_lossless_join():
text = "Sure! Here's a list of fruits: apple, banana, orange."
chunks = await _collect_async(_async_payload(content=text))
assert len(chunks) > 1
assert all(isinstance(c, ModelResponseStream) for c in chunks)
reassembled = "".join((c.choices[0].delta.content or "") for c in chunks)
assert reassembled == text
@pytest.mark.asyncio
async def test_async_finish_reason_only_on_last_chunk():
chunks = await _collect_async(_async_payload())
finish_reasons = [c.choices[0].finish_reason for c in chunks]
assert finish_reasons[-1] == "stop"
assert all(fr is None for fr in finish_reasons[:-1])
@pytest.mark.asyncio
async def test_async_role_only_on_first_chunk():
chunks = await _collect_async(_async_payload())
assert chunks[0].choices[0].delta.role == "assistant"
for c in chunks[1:]:
assert c.choices[0].delta.role is None
@pytest.mark.asyncio
async def test_async_usage_attached_to_last_chunk_only():
chunks = await _collect_async(_async_payload())
usage_frames = [c for c in chunks if getattr(c, "usage", None) is not None]
assert len(usage_frames) == 1
assert usage_frames[0] is chunks[-1]
assert usage_frames[0].usage.completion_tokens == 7
assert usage_frames[0].usage.prompt_tokens == 5
assert usage_frames[0].usage.total_tokens == 12
assert usage_frames[0].choices[0].finish_reason == "stop"
@pytest.mark.asyncio
async def test_async_empty_content_yields_single_finish_frame():
# tool-calls-only-style response: short-circuit to the single-yield path.
payload = _async_payload(content=None)
payload["usage"] = None
chunks = await _collect_async(payload)
assert len(chunks) == 1
assert chunks[0].choices[0].delta.content is None
assert chunks[0].choices[0].finish_reason == "stop"
assert chunks[0].choices[0].delta.role == "assistant"
assert getattr(chunks[0], "usage", None) is None
@pytest.mark.asyncio
async def test_async_tool_calls_and_function_call_only_on_first_chunk():
# A cached response combining multi-word content with tool_calls must not
# repeat the tool_calls on every slice — downstream handlers accumulate
# tool-call deltas and would collect them N times.
payload = _async_payload(content="I'll look that up for you")
payload["choices"][0]["message"]["tool_calls"] = [
{
"id": "call_1",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
"index": 0,
}
]
chunks = await _collect_async(payload)
assert len(chunks) > 1
first_tool_calls = chunks[0].choices[0].delta.tool_calls
assert first_tool_calls is not None and len(first_tool_calls) == 1
assert first_tool_calls[0].function.name == "lookup"
for c in chunks[1:]:
assert c.choices[0].delta.tool_calls is None
assert c.choices[0].delta.function_call is None
@pytest.mark.asyncio
async def test_async_logprobs_only_on_first_chunk():
payload = _async_payload(content="cached token logprobs")
payload["choices"][0]["logprobs"] = {"content": []}
chunks = await _collect_async(payload)
assert len(chunks) > 1
assert chunks[0].choices[0].logprobs is not None
for c in chunks[1:]:
assert c.choices[0].logprobs is None
@pytest.mark.asyncio
async def test_async_metadata_propagated_to_every_chunk():
chunks = await _collect_async(_async_payload())
for c in chunks:
assert c.id == "chatcmpl-test-async"
assert c.model == "gpt-4o-mini"
assert c.system_fingerprint == "fp_test"
assert c.created == 1700000000
# ---------- sync generator (parity smoke test) ----------
def test_sync_multi_chunk_and_lossless_join():
text = "Sure! Here's a list of fruits: apple, banana, orange."
payload = {
"id": "chatcmpl-test-sync",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4o-mini",
"system_fingerprint": "fp_test",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": text},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 5, "completion_tokens": 7, "total_tokens": 12},
}
chunks = list(convert_to_streaming_response(payload))
assert len(chunks) > 1
reassembled = "".join((c.choices[0].delta.content or "") for c in chunks)
assert reassembled == text
# Sync path must also honor the "usage on last chunk" invariant.
assert chunks[-1].choices[0].finish_reason == "stop"
assert getattr(chunks[-1], "usage", None) is not None
assert chunks[-1].usage.total_tokens == 12
def test_sync_delta_and_choice_metadata_only_on_first_chunk():
thinking_blocks = [
{"type": "thinking", "thinking": "cached thinking", "signature": "sig"}
]
payload = {
"id": "chatcmpl-test-sync",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "cached content slices",
"reasoning_content": "cached reasoning",
"thinking_blocks": thinking_blocks,
},
"finish_reason": "stop",
"logprobs": {"content": []},
"enhancements": {"source": "cache"},
}
],
}
chunks = list(convert_to_streaming_response(payload))
assert len(chunks) > 1
first_choice = chunks[0].choices[0]
assert getattr(first_choice.delta, "reasoning_content", None) == "cached reasoning"
assert getattr(first_choice.delta, "thinking_blocks", None) == thinking_blocks
assert first_choice.logprobs is not None
assert first_choice.enhancements == {"source": "cache"}
for c in chunks[1:]:
choice = c.choices[0]
assert getattr(choice.delta, "reasoning_content", None) is None
assert getattr(choice.delta, "thinking_blocks", None) is None
assert choice.logprobs is None
assert getattr(choice, "enhancements", None) is None
def test_sync_non_enumerated_delta_fields_only_on_first_chunk():
# annotations (and any other Delta field beyond the role/tool_call/reasoning
# set) must also stay on the first slice. Rebuilding later slices as a bare
# content delta drops the whole class, so this holds without enumerating
# every field by hand.
annotations = [
{
"type": "url_citation",
"url_citation": {
"url": "https://example.com",
"title": "Example",
"start_index": 0,
"end_index": 1,
},
}
]
payload = {
"id": "chatcmpl-test-sync",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "cached content slices here",
"annotations": annotations,
},
"finish_reason": "stop",
}
],
}
chunks = list(convert_to_streaming_response(payload))
assert len(chunks) > 1
assert getattr(chunks[0].choices[0].delta, "annotations", None) == annotations
for c in chunks[1:]:
assert getattr(c.choices[0].delta, "annotations", None) is None