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