diff --git a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py index 016bb6b1e22..baf71220506 100644 --- a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py +++ b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py @@ -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 diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index e278483d689..03a87fb6a39 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -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 diff --git a/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py b/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py index 9a69f513069..4d67f1e426a 100644 --- a/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py +++ b/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py @@ -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: diff --git a/tests/test_litellm/litellm_core_utils/llm_response_utils/test_convert_to_streaming_response.py b/tests/test_litellm/litellm_core_utils/llm_response_utils/test_convert_to_streaming_response.py new file mode 100644 index 00000000000..2c3dae21058 --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/llm_response_utils/test_convert_to_streaming_response.py @@ -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