mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix: merge gemini stream provider fields
This commit is contained in:
parent
ebd335da67
commit
79e4542703
2 changed files with 149 additions and 9 deletions
|
|
@ -7373,6 +7373,74 @@ def stream_chunk_builder_text_completion(
|
|||
return TextCompletionResponse(**response)
|
||||
|
||||
|
||||
def _merge_server_side_tool_invocations(existing: Optional[Any], incoming: Any) -> Any:
|
||||
if not isinstance(incoming, list):
|
||||
return incoming
|
||||
|
||||
if not isinstance(existing, list):
|
||||
existing = []
|
||||
|
||||
merged_invocations: List[Any] = []
|
||||
invocations_by_id: Dict[str, Dict[str, Any]] = {}
|
||||
|
||||
for invocation in existing:
|
||||
if not isinstance(invocation, dict):
|
||||
merged_invocations.append(invocation)
|
||||
continue
|
||||
|
||||
invocation_copy = dict(invocation)
|
||||
invocation_id = invocation_copy.get("id")
|
||||
if isinstance(invocation_id, str) and invocation_id:
|
||||
invocations_by_id[invocation_id] = invocation_copy
|
||||
merged_invocations.append(invocation_copy)
|
||||
|
||||
for invocation in incoming:
|
||||
if not isinstance(invocation, dict):
|
||||
merged_invocations.append(invocation)
|
||||
continue
|
||||
|
||||
invocation_id = invocation.get("id")
|
||||
if isinstance(invocation_id, str) and invocation_id in invocations_by_id:
|
||||
existing_invocation = invocations_by_id[invocation_id]
|
||||
for key, value in invocation.items():
|
||||
if key not in existing_invocation or existing_invocation[key] is None:
|
||||
existing_invocation[key] = value
|
||||
continue
|
||||
|
||||
invocation_copy = dict(invocation)
|
||||
if isinstance(invocation_id, str) and invocation_id:
|
||||
invocations_by_id[invocation_id] = invocation_copy
|
||||
merged_invocations.append(invocation_copy)
|
||||
|
||||
return merged_invocations
|
||||
|
||||
|
||||
def _merge_streaming_provider_specific_field(
|
||||
combined_provider_fields: Dict[str, Any], key: str, value: Any
|
||||
) -> None:
|
||||
if key == "server_side_tool_invocations":
|
||||
combined_provider_fields[key] = _merge_server_side_tool_invocations(
|
||||
combined_provider_fields.get(key), value
|
||||
)
|
||||
return
|
||||
|
||||
if key == "thought_signatures" and isinstance(value, list):
|
||||
existing = combined_provider_fields.get(key)
|
||||
if isinstance(existing, list):
|
||||
existing.extend(value)
|
||||
else:
|
||||
combined_provider_fields[key] = list(value)
|
||||
return
|
||||
|
||||
if key not in combined_provider_fields:
|
||||
combined_provider_fields[key] = value
|
||||
elif isinstance(value, list) and isinstance(combined_provider_fields[key], list):
|
||||
# For lists like web_search_results, take the last (most complete) one
|
||||
combined_provider_fields[key] = value
|
||||
else:
|
||||
combined_provider_fields[key] = value
|
||||
|
||||
|
||||
def stream_chunk_builder( # noqa: PLR0915
|
||||
chunks: list,
|
||||
messages: Optional[list] = None,
|
||||
|
|
@ -7613,15 +7681,9 @@ def stream_chunk_builder( # noqa: PLR0915
|
|||
fields = chunk["choices"][0]["delta"]["provider_specific_fields"]
|
||||
if isinstance(fields, dict):
|
||||
for key, value in fields.items():
|
||||
if key not in combined_provider_fields:
|
||||
combined_provider_fields[key] = value
|
||||
elif isinstance(value, list) and isinstance(
|
||||
combined_provider_fields[key], list
|
||||
):
|
||||
# For lists like web_search_results, take the last (most complete) one
|
||||
combined_provider_fields[key] = value
|
||||
else:
|
||||
combined_provider_fields[key] = value
|
||||
_merge_streaming_provider_specific_field(
|
||||
combined_provider_fields, key, value
|
||||
)
|
||||
|
||||
if combined_provider_fields:
|
||||
_choice = cast(Choices, response.choices[0])
|
||||
|
|
|
|||
|
|
@ -613,3 +613,81 @@ def test_stream_chunk_builder_dict_snapshot_preserves_hidden_provider_fields():
|
|||
assert (
|
||||
response._hidden_params["provider_specific_fields"]["traffic_type"] == "default"
|
||||
)
|
||||
|
||||
|
||||
def test_stream_chunk_builder_merges_gemini_server_side_tool_invocations():
|
||||
chunk1 = ModelResponseStream(
|
||||
id="chatcmpl-gemini-server-tools",
|
||||
created=1,
|
||||
model="gemini-3-flash-preview",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason=None,
|
||||
index=0,
|
||||
delta=Delta(
|
||||
content="",
|
||||
role="assistant",
|
||||
provider_specific_fields={
|
||||
"server_side_tool_invocations": [
|
||||
{
|
||||
"id": "srvtoolu_1",
|
||||
"type": "toolCall",
|
||||
"name": "googleSearch",
|
||||
"args": {"query": "Brisbane population"},
|
||||
"thought_signature": "call-signature",
|
||||
}
|
||||
],
|
||||
"thought_signatures": ["signature-1"],
|
||||
"web_search_results": [{"title": "older results"}],
|
||||
},
|
||||
),
|
||||
)
|
||||
],
|
||||
)
|
||||
chunk2 = ModelResponseStream(
|
||||
id="chatcmpl-gemini-server-tools",
|
||||
created=2,
|
||||
model="gemini-3-flash-preview",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(
|
||||
content=None,
|
||||
role=None,
|
||||
provider_specific_fields={
|
||||
"server_side_tool_invocations": [
|
||||
{
|
||||
"id": "srvtoolu_1",
|
||||
"type": "toolResponse",
|
||||
"response": {"results": [{"title": "Brisbane"}]},
|
||||
"thought_signature": "response-signature",
|
||||
}
|
||||
],
|
||||
"thought_signatures": ["signature-2"],
|
||||
"web_search_results": [{"title": "newer results"}],
|
||||
},
|
||||
),
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
response = stream_chunk_builder(chunks=[chunk1, chunk2])
|
||||
|
||||
assert response is not None
|
||||
provider_fields = response.choices[0].message.provider_specific_fields
|
||||
assert provider_fields is not None
|
||||
assert provider_fields["thought_signatures"] == ["signature-1", "signature-2"]
|
||||
assert provider_fields["web_search_results"] == [{"title": "newer results"}]
|
||||
assert provider_fields["server_side_tool_invocations"] == [
|
||||
{
|
||||
"id": "srvtoolu_1",
|
||||
"type": "toolCall",
|
||||
"name": "googleSearch",
|
||||
"args": {"query": "Brisbane population"},
|
||||
"thought_signature": "call-signature",
|
||||
"response": {"results": [{"title": "Brisbane"}]},
|
||||
}
|
||||
]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue