mirror of
https://github.com/tiennm99/litellm.git
synced 2026-07-20 02:18:38 +00:00
Merge pull request #19179 from BerriAI/litellm_fix_streaming_chunk
Fix : test_stream_chunk_builder_litellm_mixed_calls
This commit is contained in:
@@ -147,10 +147,26 @@ class ChunkProcessor:
|
||||
tool_calls = delta.get("tool_calls", [])
|
||||
|
||||
for tool_call in tool_calls:
|
||||
if not tool_call or not hasattr(tool_call, "function"):
|
||||
# Handle both dict and object formats
|
||||
if not tool_call:
|
||||
continue
|
||||
|
||||
# Check if tool_call has function (either as attribute or dict key)
|
||||
has_function = False
|
||||
if isinstance(tool_call, dict):
|
||||
has_function = "function" in tool_call and tool_call["function"] is not None
|
||||
else:
|
||||
has_function = hasattr(tool_call, "function") and tool_call.function is not None
|
||||
|
||||
if not has_function:
|
||||
continue
|
||||
|
||||
index = getattr(tool_call, "index", 0)
|
||||
# Get index (handle both dict and object)
|
||||
if isinstance(tool_call, dict):
|
||||
index = tool_call.get("index", 0)
|
||||
else:
|
||||
index = getattr(tool_call, "index", 0)
|
||||
|
||||
if index not in tool_call_map:
|
||||
tool_call_map[index] = {
|
||||
"id": None,
|
||||
@@ -160,30 +176,56 @@ class ChunkProcessor:
|
||||
"provider_specific_fields": None,
|
||||
}
|
||||
|
||||
if hasattr(tool_call, "id") and tool_call.id:
|
||||
tool_call_map[index]["id"] = tool_call.id
|
||||
if hasattr(tool_call, "type") and tool_call.type:
|
||||
tool_call_map[index]["type"] = tool_call.type
|
||||
if hasattr(tool_call, "function"):
|
||||
if (
|
||||
hasattr(tool_call.function, "name")
|
||||
and tool_call.function.name
|
||||
):
|
||||
tool_call_map[index]["name"] = tool_call.function.name
|
||||
if (
|
||||
hasattr(tool_call.function, "arguments")
|
||||
and tool_call.function.arguments
|
||||
):
|
||||
tool_call_map[index]["arguments"].append(
|
||||
tool_call.function.arguments
|
||||
)
|
||||
# Extract id, type, and function data (handle both dict and object)
|
||||
if isinstance(tool_call, dict):
|
||||
if tool_call.get("id"):
|
||||
tool_call_map[index]["id"] = tool_call["id"]
|
||||
if tool_call.get("type"):
|
||||
tool_call_map[index]["type"] = tool_call["type"]
|
||||
|
||||
function = tool_call.get("function", {})
|
||||
if isinstance(function, dict):
|
||||
if function.get("name"):
|
||||
tool_call_map[index]["name"] = function["name"]
|
||||
if function.get("arguments"):
|
||||
tool_call_map[index]["arguments"].append(function["arguments"])
|
||||
else:
|
||||
# function is an object
|
||||
if hasattr(function, "name") and function.name:
|
||||
tool_call_map[index]["name"] = function.name
|
||||
if hasattr(function, "arguments") and function.arguments:
|
||||
tool_call_map[index]["arguments"].append(function.arguments)
|
||||
else:
|
||||
# tool_call is an object
|
||||
if hasattr(tool_call, "id") and tool_call.id:
|
||||
tool_call_map[index]["id"] = tool_call.id
|
||||
if hasattr(tool_call, "type") and tool_call.type:
|
||||
tool_call_map[index]["type"] = tool_call.type
|
||||
if hasattr(tool_call, "function"):
|
||||
if (
|
||||
hasattr(tool_call.function, "name")
|
||||
and tool_call.function.name
|
||||
):
|
||||
tool_call_map[index]["name"] = tool_call.function.name
|
||||
if (
|
||||
hasattr(tool_call.function, "arguments")
|
||||
and tool_call.function.arguments
|
||||
):
|
||||
tool_call_map[index]["arguments"].append(
|
||||
tool_call.function.arguments
|
||||
)
|
||||
|
||||
# Preserve provider_specific_fields from streaming chunks
|
||||
provider_fields = None
|
||||
if hasattr(tool_call, "provider_specific_fields") and tool_call.provider_specific_fields:
|
||||
provider_fields = tool_call.provider_specific_fields
|
||||
elif hasattr(tool_call, "function") and hasattr(tool_call.function, "provider_specific_fields") and tool_call.function.provider_specific_fields:
|
||||
provider_fields = tool_call.function.provider_specific_fields
|
||||
if isinstance(tool_call, dict):
|
||||
provider_fields = tool_call.get("provider_specific_fields")
|
||||
if not provider_fields and isinstance(tool_call.get("function"), dict):
|
||||
provider_fields = tool_call["function"].get("provider_specific_fields")
|
||||
else:
|
||||
if hasattr(tool_call, "provider_specific_fields") and tool_call.provider_specific_fields:
|
||||
provider_fields = tool_call.provider_specific_fields
|
||||
elif hasattr(tool_call, "function") and hasattr(tool_call.function, "provider_specific_fields") and tool_call.function.provider_specific_fields:
|
||||
provider_fields = tool_call.function.provider_specific_fields
|
||||
|
||||
if provider_fields:
|
||||
# Merge provider_specific_fields if multiple chunks have them
|
||||
|
||||
@@ -422,7 +422,7 @@ def test_streaming_tool_calls_transformation():
|
||||
ChatCompletionDeltaToolCall,
|
||||
Delta,
|
||||
Function,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
)
|
||||
|
||||
@@ -454,7 +454,7 @@ def test_streaming_tool_calls_transformation():
|
||||
delta=mock_delta
|
||||
)
|
||||
|
||||
mock_response = ModelResponse(
|
||||
mock_response = ModelResponseStream(
|
||||
id="test-streaming",
|
||||
choices=[mock_choice],
|
||||
created=1234567890,
|
||||
@@ -493,7 +493,7 @@ def test_streaming_partial_tool_calls_accumulation():
|
||||
ChatCompletionDeltaToolCall,
|
||||
Delta,
|
||||
Function,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
)
|
||||
|
||||
@@ -543,7 +543,7 @@ def test_streaming_partial_tool_calls_accumulation():
|
||||
delta=mock_delta
|
||||
)
|
||||
|
||||
mock_response = ModelResponse(
|
||||
mock_response = ModelResponseStream(
|
||||
id="test-streaming",
|
||||
choices=[mock_choice],
|
||||
created=1234567890,
|
||||
@@ -595,7 +595,7 @@ def test_streaming_multiple_partial_tool_calls():
|
||||
ChatCompletionDeltaToolCall,
|
||||
Delta,
|
||||
Function,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
)
|
||||
|
||||
@@ -642,7 +642,7 @@ def test_streaming_multiple_partial_tool_calls():
|
||||
delta=mock_delta
|
||||
)
|
||||
|
||||
mock_response = ModelResponse(
|
||||
mock_response = ModelResponseStream(
|
||||
id="test-streaming",
|
||||
choices=[mock_choice],
|
||||
created=1234567890,
|
||||
|
||||
Reference in New Issue
Block a user