mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(streaming): join fragmented tool call names and ids when rebuilding a stream
stream_chunk_builder kept only the last fragment of a streamed tool call's function name and id, while it already joined the arguments. A name streamed as `get_` then `weather` was rebuilt as `weather`, and an id streamed as `call_` then `9f2c` as `9f2c`, which breaks the next tool-result turn for callers that reuse the rebuilt message and puts the wrong tool name in logs Names and ids are now joined per choice and tool index the same way as arguments, which matches what the OpenAI SDK's stream accumulator does with the same chunks
This commit is contained in:
parent
fe910889f7
commit
765e36087b
2 changed files with 97 additions and 5 deletions
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue