This commit is contained in:
joeym82956 2026-10-04 12:47:47 -07:00 • committed by GitHub
commit 23e80c97a2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 97 additions and 5 deletions

View file

@ -455,20 +455,31 @@ class ChunkProcessor:
continue
if isinstance(tool_call, dict):
key = (choice_index, tool_call.get("index", 0))
if fragment_id := tool_call.get("id"):
yield key, "id", fragment_id
function = tool_call.get("function")
if isinstance(function, dict):
if fragment_arguments := function.get("arguments"):
yield key, "arguments", fragment_arguments
elif function_arguments := getattr(function, "arguments", None):
yield key, "arguments", function_arguments
if fragment_name := function.get("name"):
yield key, "name", fragment_name
else:
if function_arguments := getattr(function, "arguments", None):
yield key, "arguments", function_arguments
if function_name := getattr(function, "name", None):
yield key, "name", function_name
custom = tool_call.get("custom")
if isinstance(custom, dict) and (custom_input := custom.get("input")):
yield key, "custom_input", custom_input
else:
key = (choice_index, getattr(tool_call, "index", 0))
if object_id := getattr(tool_call, "id", None):
yield key, "id", object_id
function = getattr(tool_call, "function", None)
if object_arguments := getattr(function, "arguments", None):
yield key, "arguments", object_arguments
if object_name := getattr(function, "name", None):
yield key, "name", object_name
custom = getattr(tool_call, "custom", None)
if object_custom_input := getattr(custom, "input", None):
yield key, "custom_input", object_custom_input
@ -607,7 +618,7 @@ class ChunkProcessor:
if tool_call_data["id"] and tool_call_data["custom_name"]:
tool_calls_list.append(
ChatCompletionMessageCustomToolCall(
id=tool_call_data["id"],
id=joined_fragments.get((index, "id")) or tool_call_data["id"],
custom=ChatCompletionCustomToolCallPayload(
name=tool_call_data["custom_name"],
input=joined_fragments.get((index, "custom_input"), ""),
@ -620,12 +631,12 @@ class ChunkProcessor:
# Build function - provider_specific_fields should be on tool_call level, not function level
function = Function(
arguments=combined_arguments,
name=tool_call_data["name"],
name=joined_fragments.get((index, "name")) or tool_call_data["name"],
)
# Prepare params for ChatCompletionMessageToolCall
tool_call_params = {
"id": tool_call_data["id"],
"id": joined_fragments.get((index, "id")) or tool_call_data["id"],
"function": function,
"type": tool_call_data["type"] or "function",
}

View file

@ -1309,6 +1309,87 @@ def _choice_tool_call_delta_chunk(choice_index: int, tool_call: dict[str, object
return {"choices": [{"index": choice_index, "delta": {"tool_calls": [tool_call]}}]}
@pytest.mark.parametrize("as_dict", (False, True))
@pytest.mark.parametrize(
"first_name,last_name,expected_name", (("get_", "widget", "get_widget"), ("echo", "echo", "echoecho"))
)
def test_get_combined_tool_content_joins_function_name_deltas(
as_dict: bool, first_name: str, last_name: str, expected_name: str
) -> None:
calls: Final = (
ChatCompletionDeltaToolCall(
index=0, id="call_fragmented", type="function", function=Function(name=first_name, arguments='{"id":')
),
ChatCompletionDeltaToolCall(index=0, function=Function(name=last_name, arguments="17}")),
)
chunks: Final = [_tool_call_delta_chunk(call.model_dump(exclude_none=True) if as_dict else call) for call in calls]
combined: Final = ChunkProcessor(chunks).get_combined_tool_content(chunks)
assert combined == [
ChatCompletionMessageToolCall(
id="call_fragmented", type="function", function=Function(name=expected_name, arguments='{"id":17}')
)
]
@pytest.mark.parametrize("as_dict", (False, True))
@pytest.mark.parametrize(
"first_id,last_id,expected_id", (("call_", "fragmented", "call_fragmented"), ("aa", "aa", "aaaa"))
)
def test_get_combined_tool_content_joins_tool_id_deltas(
as_dict: bool, first_id: str, last_id: str, expected_id: str
) -> None:
calls: Final = (
ChatCompletionDeltaToolCall(
index=0, id=first_id, type="function", function=Function(name="get_widget", arguments='{"id":')
),
ChatCompletionDeltaToolCall(index=0, id=last_id, function=Function(arguments="17}")),
)
chunks: Final = [_tool_call_delta_chunk(call.model_dump(exclude_none=True) if as_dict else call) for call in calls]
combined: Final = ChunkProcessor(chunks).get_combined_tool_content(chunks)
assert combined == [
ChatCompletionMessageToolCall(
id=expected_id, type="function", function=Function(name="get_widget", arguments='{"id":17}')
)
]
@pytest.mark.parametrize("separate_choices", (False, True))
def test_get_combined_tool_content_keeps_interleaved_metadata_with_its_choice_and_tool_index(
separate_choices: bool,
) -> None:
second_choice, second_tool = (1, 0) if separate_choices else (0, 1)
chunks: Final = [
_choice_tool_call_delta_chunk(
0, {"index": 0, "id": "call_", "type": "function", "function": {"name": "get_", "arguments": '{"slot":"'}}
),
_choice_tool_call_delta_chunk(
second_choice,
{
"index": second_tool,
"id": "call_",
"type": "function",
"function": {"name": "get_", "arguments": '{"slot":"'},
},
),
_choice_tool_call_delta_chunk(
second_choice, {"index": second_tool, "id": "beta", "function": {"name": "beta", "arguments": 'beta"}'}}
),
_choice_tool_call_delta_chunk(
0, {"index": 0, "id": "alpha", "function": {"name": "alpha", "arguments": 'alpha"}'}}
),
]
combined: Final = ChunkProcessor(chunks).get_combined_tool_content(chunks)
assert combined == [
ChatCompletionMessageToolCall(
id="call_alpha", function=Function(name="get_alpha", arguments='{"slot":"alpha"}')
),
ChatCompletionMessageToolCall(id="call_beta", function=Function(name="get_beta", arguments='{"slot":"beta"}')),
]
def test_get_combined_tool_content_keeps_each_choices_arguments_apart_when_choices_share_a_tool_index():
processor = ChunkProcessor.__new__(ChunkProcessor)
chunks = [