mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
fix: correct offset when expanding multiple multi_tool_use.parallel tool calls
_handle_invalid_parallel_tool_calls splices each hallucinated multi_tool_use.parallel call in place with the real calls it encoded. The running offset advanced by len(replacement), but each splice removes one element and inserts len(replacement), so the net shift is len(replacement) - 1. With two or more such calls in one message, the offset over-advanced by one per replacement, leaving later multi_tool_use.parallel wrappers unexpanded and placing their real calls at the wrong position.
This commit is contained in:
parent
79a6b8f7f0
commit
00066a88cc
2 changed files with 89 additions and 1 deletions
|
|
@ -402,7 +402,9 @@ def _handle_invalid_parallel_tool_calls(
|
|||
shift = 0
|
||||
for i, replacement in replacements.items():
|
||||
tool_calls[:] = tool_calls[: i + shift] + replacement + tool_calls[i + shift + 1 :]
|
||||
shift += len(replacement)
|
||||
# Each splice removes one element and inserts len(replacement), so the
|
||||
# net change to subsequent indices is len(replacement) - 1.
|
||||
shift += len(replacement) - 1
|
||||
|
||||
return tool_calls
|
||||
except json.JSONDecodeError:
|
||||
|
|
|
|||
|
|
@ -0,0 +1,86 @@
|
|||
"""
|
||||
Tests for ``_handle_invalid_parallel_tool_calls`` in
|
||||
``litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response``.
|
||||
|
||||
The function replaces each hallucinated ``multi_tool_use.parallel`` tool call
|
||||
with the real tool calls it encoded. These tests focus on the case where a
|
||||
single message carries more than one ``multi_tool_use.parallel`` call, which
|
||||
exercises the running-offset bookkeeping used to splice replacements in place.
|
||||
"""
|
||||
|
||||
import json
|
||||
|
||||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
_handle_invalid_parallel_tool_calls,
|
||||
)
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall, Function
|
||||
|
||||
|
||||
def _multi_tool_use(call_id, tool_uses):
|
||||
return ChatCompletionMessageToolCall(
|
||||
id=call_id,
|
||||
type="function",
|
||||
function=Function(
|
||||
name="multi_tool_use.parallel",
|
||||
arguments=json.dumps({"tool_uses": tool_uses}),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_multiple_multi_tool_use_parallel_all_expanded():
|
||||
"""Two multi_tool_use.parallel calls in one message must both be expanded."""
|
||||
tool_calls = [
|
||||
_multi_tool_use(
|
||||
"call_1",
|
||||
[
|
||||
{"recipient_name": "functions.get_weather", "parameters": {"city": "NYC"}},
|
||||
{"recipient_name": "functions.get_time", "parameters": {"tz": "EST"}},
|
||||
],
|
||||
),
|
||||
_multi_tool_use(
|
||||
"call_2",
|
||||
[
|
||||
{"recipient_name": "functions.get_stock", "parameters": {"ticker": "AAPL"}},
|
||||
],
|
||||
),
|
||||
]
|
||||
|
||||
result = _handle_invalid_parallel_tool_calls(tool_calls)
|
||||
|
||||
# No hallucinated wrapper call should survive.
|
||||
assert all(tc.function.name != "multi_tool_use.parallel" for tc in result)
|
||||
|
||||
names = [tc.function.name for tc in result]
|
||||
assert names == ["get_weather", "get_time", "get_stock"]
|
||||
|
||||
ids = [tc.id for tc in result]
|
||||
assert ids == ["call_1_0", "call_1_1", "call_2_0"]
|
||||
|
||||
|
||||
def test_multi_tool_use_parallel_interleaved_with_real_call():
|
||||
"""A real tool call between two parallel calls keeps its position and value."""
|
||||
tool_calls = [
|
||||
_multi_tool_use(
|
||||
"call_1",
|
||||
[
|
||||
{"recipient_name": "functions.get_weather", "parameters": {"city": "NYC"}},
|
||||
],
|
||||
),
|
||||
ChatCompletionMessageToolCall(
|
||||
id="call_real",
|
||||
type="function",
|
||||
function=Function(name="lookup", arguments='{"q": "x"}'),
|
||||
),
|
||||
_multi_tool_use(
|
||||
"call_3",
|
||||
[
|
||||
{"recipient_name": "functions.get_stock", "parameters": {"ticker": "AAPL"}},
|
||||
],
|
||||
),
|
||||
]
|
||||
|
||||
result = _handle_invalid_parallel_tool_calls(tool_calls)
|
||||
|
||||
assert all(tc.function.name != "multi_tool_use.parallel" for tc in result)
|
||||
assert [tc.function.name for tc in result] == ["get_weather", "lookup", "get_stock"]
|
||||
assert [tc.id for tc in result] == ["call_1_0", "call_real", "call_3_0"]
|
||||
Loading…
Add table
Reference in a new issue