diff --git a/litellm/main.py b/litellm/main.py index c5af2db75a4..cafa1e4718f 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -8605,6 +8605,12 @@ def _stream_builder_response_cost(response: ModelResponse, logging_obj: Optional return _stream_builder_model_map_cost(response) +def _joined_streamed_citations(streamed_citations: "tuple[object, ...]") -> "list[object]": + if all(isinstance(citation, list) for citation in streamed_citations): + return list(streamed_citations) # mutable-ok: JSON list field + return [list(streamed_citations)] # mutable-ok: JSON list field + + def _stream_builder_model_map_cost(response: ModelResponse) -> float | None: model_name: Final = getattr(response, "model", None) usage: Final = getattr(response, "usage", None) @@ -8861,7 +8867,9 @@ def stream_chunk_builder( fields["citation"] for fields in provider_field_dicts if fields.get("citation") is not None ) citation_fields: Final = ( - {"citations": [list(streamed_citations)]} if streamed_citations else {} # mutable-ok: JSON dict field + {"citations": _joined_streamed_citations(streamed_citations)} # mutable-ok: JSON dict field + if streamed_citations + else {} # mutable-ok: JSON dict field ) combined_provider_fields: Final = { # mutable-ok: Message.provider_specific_fields is a plain dict field key: value diff --git a/tests/test_litellm/test_stream_chunk_builder_citations.py b/tests/test_litellm/test_stream_chunk_builder_citations.py index 55fce727ab2..87774f28d4e 100644 --- a/tests/test_litellm/test_stream_chunk_builder_citations.py +++ b/tests/test_litellm/test_stream_chunk_builder_citations.py @@ -83,3 +83,22 @@ def test_stream_chunk_builder_without_citation_deltas_sets_no_citations_key(): assert fields is not None assert "citations" not in fields assert fields["web_search_results"] == [{"url": "https://example.com"}] + + +def test_stream_chunk_builder_keeps_block_list_citation_deltas_unnested(): + block_one: Final = [dict(_CITATION_ONE), dict(_CITATION_TWO)] + block_two: Final = [dict(_CITATION_ONE)] + chunks: Final = [ + _chunk(Delta(content="Green sky.", role="assistant")), + _chunk(Delta(content="", provider_specific_fields={"citation": block_one})), + _chunk(Delta(content="", provider_specific_fields={"citation": block_two})), + _chunk(Delta(content=""), finish_reason="stop"), + ] + + response: Final = stream_chunk_builder(chunks=chunks) + + assert response is not None + fields: Final = response.choices[0].message.provider_specific_fields + assert fields is not None + assert fields["citations"] == [block_one, block_two] + assert "citation" not in fields