mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
chore: keep this PR to the shift fix only
Two commits that belong to the malformed-payload PR (#43173) had leaked onto this branch: the widened except clause and its test file. Restore the clause to main's version and drop those tests here, and keep the regression test only in tests/test_litellm/litellm_core_utils/llm_response_utils/ where the unit coverage job runs.
This commit is contained in:
parent
a99918e845
commit
c02f11dd7e
3 changed files with 1 additions and 142 deletions
|
|
@ -413,7 +413,7 @@ def _handle_invalid_parallel_tool_calls(
|
|||
shift += len(replacement) - 1
|
||||
|
||||
return tool_calls
|
||||
except (json.JSONDecodeError, KeyError, TypeError, AttributeError):
|
||||
except json.JSONDecodeError:
|
||||
# if there is a JSONDecodeError, return the original tool_calls
|
||||
return tool_calls
|
||||
|
||||
|
|
|
|||
|
|
@ -2486,85 +2486,3 @@ 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, *calls):
|
||||
return {
|
||||
"id": call_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "multi_tool_use.parallel",
|
||||
"arguments": json.dumps(
|
||||
{
|
||||
"tool_uses": [
|
||||
{"recipient_name": name, "parameters": params}
|
||||
for name, params in calls
|
||||
]
|
||||
}
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
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"),
|
||||
# Two expansions here and one below, so the first splice
|
||||
# changes the list length and the second splice has to land
|
||||
# at the shifted offset.
|
||||
parallel(
|
||||
"m1",
|
||||
("functions.get_time", {"tz": "UTC"}),
|
||||
("functions.get_date", {"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_date",
|
||||
"get_news",
|
||||
"get_quote",
|
||||
"get_forecast",
|
||||
]
|
||||
assert "multi_tool_use.parallel" not in names
|
||||
|
|
|
|||
|
|
@ -1,59 +0,0 @@
|
|||
"""The hallucinated multi_tool_use.parallel expansion must not fail the response."""
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
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 _parallel_instance(arguments: str):
|
||||
return [
|
||||
ChatCompletionMessageToolCall(
|
||||
id="call_1",
|
||||
type="function",
|
||||
function=Function(name="multi_tool_use.parallel", arguments=arguments),
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"arguments",
|
||||
[
|
||||
'{"tool_uses": [{"recipient_name": "functions.get_weather"}]}',
|
||||
'{"tool_uses": "nope"}',
|
||||
'{"tool_uses": [42]}',
|
||||
"{}",
|
||||
],
|
||||
)
|
||||
def test_malformed_tool_uses_returns_original_calls(arguments):
|
||||
"""A hallucinated payload we cannot expand must come back untouched."""
|
||||
tool_calls = _parallel_instance(arguments)
|
||||
|
||||
result = _handle_invalid_parallel_tool_calls(tool_calls)
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0].function.name == "multi_tool_use.parallel"
|
||||
assert result[0].id == "call_1"
|
||||
assert result[0].function.arguments == arguments
|
||||
|
||||
|
||||
def test_valid_parallel_payload_still_expands():
|
||||
"""The guard must not swallow well-formed payloads."""
|
||||
tool_calls = _parallel_instance(
|
||||
json.dumps(
|
||||
{
|
||||
"tool_uses": [
|
||||
{"recipient_name": "functions.get_weather", "parameters": {"city": "NYC"}},
|
||||
{"recipient_name": "functions.get_time", "parameters": {"tz": "EST"}},
|
||||
]
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
result = _handle_invalid_parallel_tool_calls(tool_calls)
|
||||
|
||||
assert [c.function.name for c in result] == ["get_weather", "get_time"]
|
||||
Loading…
Add table
Reference in a new issue