fix(llm_response_utils): expand every hallucinated multi_tool_use.parallel call

_handle_invalid_parallel_tool_calls replaces one entry with len(replacement)
entries, so the offsets of everything after it shift by len(replacement) - 1.
Advancing by the full length skipped one entry per expansion: with two
multi_tool_use.parallel calls in one message the second stayed in
message.tool_calls (so agents tried to call a non-existent tool) and the call
after it was overwritten.

Adds a regression test that expands two such calls in one message.
This commit is contained in:
sclfcz 2026-09-25 23:07:18 +08:00
parent 61953318bf
commit 60c2d9f590
2 changed files with 81 additions and 1 deletions

View file

@ -406,7 +406,11 @@ 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)
# One entry is replaced by ``len(replacement)`` entries, so the
# offsets of everything after it move by the difference - not by
# the full length, which would skip one entry per expansion and
# leave the next ``multi_tool_use.parallel`` call in place.
shift += len(replacement) - 1
return tool_calls
except json.JSONDecodeError:

View file

@ -2486,3 +2486,79 @@ class TestConvertToModelResponseObjectCompletion:
},
model_response_object=None,
)
def test_convert_to_model_response_object_expands_every_parallel_tool_call():
"""
Every hallucinated `multi_tool_use.parallel` entry in one message must be
expanded. Replacing one entry with its expansions moves the offsets of the
entries that follow it by `len(expansions) - 1`; advancing by the full
length skipped one entry per expansion, so with two such calls the second
one stayed in place and the call after it was overwritten.
"""
def parallel(call_id, recipient_name, parameters):
return {
"id": call_id,
"type": "function",
"function": {
"name": "multi_tool_use.parallel",
"arguments": json.dumps(
{
"tool_uses": [
{
"recipient_name": recipient_name,
"parameters": parameters,
}
]
}
),
},
}
def plain(call_id, name):
return {
"id": call_id,
"type": "function",
"function": {"name": name, "arguments": "{}"},
}
response_object = {
"id": "chatcmpl-parallel",
"object": "chat.completion",
"created": 1728933352,
"model": "gpt-4o-2024-08-06",
"choices": [
{
"index": 0,
"finish_reason": "tool_calls",
"message": {
"role": "assistant",
"content": None,
"tool_calls": [
plain("0", "get_weather"),
parallel("m1", "functions.get_time", {"tz": "UTC"}),
plain("2", "get_news"),
parallel("m2", "functions.get_quote", {"sym": "AAPL"}),
plain("4", "get_forecast"),
],
},
}
],
}
result = convert_to_model_response_object(
response_object=response_object,
model_response_object=ModelResponse(),
response_type="completion",
)
names = [tc.function.name for tc in result.choices[0].message.tool_calls]
assert names == [
"get_weather",
"get_time",
"get_news",
"get_quote",
"get_forecast",
]
assert "multi_tool_use.parallel" not in names