mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(anthropic): only split chunks whose tool calls carry complete arguments
Signed-off-by: Ankit Jha <ankit.jha@tradomate.one>
This commit is contained in:
parent
f8856928b3
commit
0462d6654f
2 changed files with 36 additions and 9 deletions
|
|
@ -214,13 +214,22 @@ class _CombinedChunkSplitter:
|
|||
return pieces
|
||||
|
||||
@staticmethod
|
||||
def _split_parallel_tool_calls(chunk: "ModelResponseStream") -> "tuple[ModelResponseStream, ...]":
|
||||
"""Return ``(chunk,)``, or one piece per tool call when a chunk opens several.
|
||||
def _has_complete_arguments(tool_call: object) -> bool:
|
||||
arguments: Final = _optional_attr(_optional_attr(tool_call, "function"), "arguments")
|
||||
if not isinstance(arguments, str) or not arguments:
|
||||
return False
|
||||
try:
|
||||
json.loads(arguments)
|
||||
except ValueError:
|
||||
return False
|
||||
return True
|
||||
|
||||
Some providers send every parallel call in one chunk. The block translators
|
||||
read only the first call's id and name and join all arguments, so the calls
|
||||
would merge into one ``tool_use`` block. Pieces keep call order; entries that
|
||||
share an ``index`` stay together. Argument continuations are left alone.
|
||||
@staticmethod
|
||||
def _split_parallel_tool_calls(chunk: "ModelResponseStream") -> "tuple[ModelResponseStream, ...]":
|
||||
"""Return ``(chunk,)``, or one piece per call when a chunk carries several whole calls.
|
||||
|
||||
Only chunks where every call has its full JSON arguments are split. A later
|
||||
argument fragment could not be routed back to an earlier block.
|
||||
"""
|
||||
choices: Final = _optional_attr_sequence(chunk, "choices")
|
||||
if len(choices) != 1:
|
||||
|
|
@ -229,6 +238,8 @@ class _CombinedChunkSplitter:
|
|||
tool_calls: Final = _optional_attr_sequence(delta, "tool_calls")
|
||||
if sum(1 for call in tool_calls if _optional_attr(_optional_attr(call, "function"), "name")) < 2:
|
||||
return (chunk,)
|
||||
if not all(_CombinedChunkSplitter._has_complete_arguments(call) for call in tool_calls):
|
||||
return (chunk,)
|
||||
|
||||
by_index: Final[dict[object, list[object]]] = {}
|
||||
for position, call in enumerate(tool_calls):
|
||||
|
|
|
|||
|
|
@ -360,13 +360,29 @@ def test_splitter_keeps_single_call_and_continuation_chunks_whole():
|
|||
assert _CombinedChunkSplitter._split_parallel_tool_calls(continuation) == (continuation,)
|
||||
|
||||
|
||||
def test_splitter_orders_pieces_by_arrival_and_groups_same_index():
|
||||
def test_splitter_orders_pieces_by_arrival():
|
||||
chunk = _tool_chunk(
|
||||
[
|
||||
_tool_call(1, "call_b", "Read", "{}"),
|
||||
_tool_call(0, "call_a", "Glob", "{}"),
|
||||
_tool_call(0, None, None, ""),
|
||||
]
|
||||
)
|
||||
pieces = _CombinedChunkSplitter._split_parallel_tool_calls(chunk)
|
||||
assert [[c.index for c in p.choices[0].delta.tool_calls] for p in pieces] == [[1], [0, 0]]
|
||||
assert [[c.index for c in p.choices[0].delta.tool_calls] for p in pieces] == [[1], [0]]
|
||||
|
||||
|
||||
def test_splitter_leaves_calls_that_still_need_arguments_whole():
|
||||
"""A later fragment could not be routed back to an earlier block (Greptile on #44079)."""
|
||||
opening = _tool_chunk([_tool_call(0, "call_a", "Read", ""), _tool_call(1, "call_b", "Read", "")])
|
||||
partial = _tool_chunk([_tool_call(0, "call_a", "Read", '{"file_path":'), _tool_call(1, "call_b", "Read", "{}")])
|
||||
assert _CombinedChunkSplitter._split_parallel_tool_calls(opening) == (opening,)
|
||||
assert _CombinedChunkSplitter._split_parallel_tool_calls(partial) == (partial,)
|
||||
|
||||
|
||||
def test_splitter_leaves_repeated_index_and_multi_choice_chunks_whole():
|
||||
same_index = _tool_chunk([_tool_call(0, "call_a", "Read", "{}"), _tool_call(0, "call_a", "Read", "{}")])
|
||||
assert _CombinedChunkSplitter._split_parallel_tool_calls(same_index) == (same_index,)
|
||||
|
||||
two_choices = _tool_chunk([_tool_call(0, "call_a", "Read", "{}"), _tool_call(1, "call_b", "Read", "{}")])
|
||||
two_choices.choices.append(StreamingChoices(index=1, delta=Delta(content="x")))
|
||||
assert _CombinedChunkSplitter._split_parallel_tool_calls(two_choices) == (two_choices,)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue