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.
This commit is contained in:
Vineeth Sai 2026-08-07 11:21:55 -07:00
parent bd289c151c
commit d72c5bbdf4
2 changed files with 130 additions and 0 deletions

View file

@ -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):

View file

@ -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