From d72c5bbdf4523c6a6e6ef57ced595e046168eb73 Mon Sep 17 00:00:00 2001 From: Vineeth Sai Date: Fri, 7 Aug 2026 11:21:55 -0700 Subject: [PATCH] fix(streaming): keep n>1 choices separate in stream_chunk_builder build_base_response emits a single hardcoded choice and every assembly step reads choices[0], while get_combined_content joins the delta content of every choice of every chunk into one string. A streamed n=2 call therefore came back as one choice holding both completions concatenated, which is text no model produced, plus whichever finish_reason arrived last. The reassembled response is what feeds cost and token accounting, what the caller gets with complete_response=True, and what goes into the response cache, so a cached streamed n=2 call later replays the concatenation. Assemble each choice index on its own by recursing over the chunks filtered to that index, then swap the results in after usage has been calculated. Usage is still computed over the full chunk list, so it is unchanged. A single-choice stream sees no behaviour change: the recursion only runs when the chunks carry more than one distinct index. --- litellm/main.py | 74 +++++++++++++++++++++++++++++++++ tests/test_litellm/test_main.py | 56 +++++++++++++++++++++++++ 2 files changed, 130 insertions(+) diff --git a/litellm/main.py b/litellm/main.py index c70a41c891a..2ddded81a23 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -8439,6 +8439,67 @@ def stream_chunk_builder_text_completion(chunks: list, messages: list | None = N return TextCompletionResponse(**response) +def _streaming_choice_index(choice: Any) -> int | None: + index = choice.get("index") if isinstance(choice, dict) else getattr(choice, "index", None) + return index if isinstance(index, int) else None + + +def _streaming_chunk_choices(chunk: Any) -> list[Any]: + choices = chunk.get("choices") if isinstance(chunk, dict) else getattr(chunk, "choices", None) + return choices or [] + + +def _distinct_streaming_choice_indices(chunks: list[Any]) -> list[int]: + indices: Final = { + _streaming_choice_index(choice) for chunk in chunks for choice in _streaming_chunk_choices(chunk) + } + return sorted(index for index in indices if index is not None) + + +def _replace_chunk_choices(chunk: Any, choices: list[Any]) -> Any: + if not choices: + return chunk + if isinstance(chunk, dict): + return {**chunk, "choices": choices} + return chunk.model_copy(update={"choices": choices}) + + +def _chunks_for_streaming_choice_index(chunks: list[Any], index: int) -> list[Any]: + matched: Final = ( + (chunk, [c for c in _streaming_chunk_choices(chunk) if _streaming_choice_index(c) == index]) + for chunk in chunks + ) + return [ + _replace_chunk_choices(chunk, choices) + for chunk, choices in matched + if choices or not _streaming_chunk_choices(chunk) + ] + + +def _build_streaming_choices_per_index( + chunks: list[Any], + indices: list[int], + messages: list | None, + start_time, + end_time, +) -> list[Choices] | None: + built: Final = [ + stream_chunk_builder( + chunks=_chunks_for_streaming_choice_index(chunks, index), + messages=messages, + start_time=start_time, + end_time=end_time, + ) + for index in indices + ] + if any(not isinstance(response, ModelResponse) or not response.choices for response in built): + return None + return [ + cast(Choices, cast(ModelResponse, response).choices[0]).model_copy(update={"index": index}) + for index, response in zip(indices, built) + ] + + def stream_chunk_builder( chunks: list, messages: list | None = None, @@ -8470,6 +8531,13 @@ def stream_chunk_builder( ): # route to the text completion logic return stream_chunk_builder_text_completion(chunks=chunks, messages=messages) + choice_indices: Final = _distinct_streaming_choice_indices(chunks) + per_choice: Final = ( + _build_streaming_choices_per_index(chunks, choice_indices, messages, start_time, end_time) + if len(choice_indices) > 1 + else None + ) + model: Final = chunks[0]["model"] # Initialize the response dictionary response: Final = processor.build_base_response(chunks) @@ -8521,6 +8589,9 @@ def stream_chunk_builder( ) setattr(response, "usage", usage) + if per_choice is not None: + response.choices = per_choice + # Propagate provider_specific_fields from chunk hidden params when present. for chunk in reversed(chunks): if isinstance(chunk, dict): @@ -8694,6 +8765,9 @@ def stream_chunk_builder( setattr(response, "usage", usage) + if per_choice is not None: + response.choices = per_choice + # Propagate provider_specific_fields from the last chunk (contains provider # metadata like traffic_type set during streaming) for chunk in reversed(chunks): diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index 9e160370048..897bde73bc7 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -1449,6 +1449,62 @@ async def test_async_mock_delay(): assert delay >= 0.01 +def _n_choice_stream_chunk(index, content, finish_reason=None, role=None): + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + delta = {"content": content} + if role is not None: + delta["role"] = role + return ModelResponseStream( + id="chatcmpl-n2", + created=1, + model="gpt-4o", + object="chat.completion.chunk", + choices=[StreamingChoices(index=index, delta=Delta(**delta), finish_reason=finish_reason)], + ) + + +def test_stream_chunk_builder_keeps_n_greater_than_one_choices_separate(): + from litellm import stream_chunk_builder + + chunks = [ + _n_choice_stream_chunk(0, "AAA", role="assistant"), + _n_choice_stream_chunk(1, "BBB", role="assistant"), + _n_choice_stream_chunk(0, " aaa"), + _n_choice_stream_chunk(1, " bbb"), + _n_choice_stream_chunk(0, None, finish_reason="stop"), + _n_choice_stream_chunk(1, None, finish_reason="length"), + ] + + response = stream_chunk_builder(chunks, messages=[{"role": "user", "content": "hi"}]) + + assert response is not None + assert len(response.choices) == 2 + assert [choice.index for choice in response.choices] == [0, 1] + assert response.choices[0].message.content == "AAA aaa" + assert response.choices[1].message.content == "BBB bbb" + assert response.choices[0].finish_reason == "stop" + assert response.choices[1].finish_reason == "length" + + +def test_stream_chunk_builder_single_choice_is_unchanged(): + from litellm import stream_chunk_builder + + chunks = [ + _n_choice_stream_chunk(0, "AAA", role="assistant"), + _n_choice_stream_chunk(0, " aaa"), + _n_choice_stream_chunk(0, None, finish_reason="stop"), + ] + + response = stream_chunk_builder(chunks, messages=[{"role": "user", "content": "hi"}]) + + assert response is not None + assert len(response.choices) == 1 + assert response.choices[0].index == 0 + assert response.choices[0].message.content == "AAA aaa" + assert response.choices[0].finish_reason == "stop" + + def test_stream_chunk_builder_thinking_blocks(): from litellm import stream_chunk_builder from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices