mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 765e36087b into 461a58c40a
This commit is contained in:
commit
23e80c97a2
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