fix(anthropic): emit one tool_use block per call when a chunk carries several

Fixes #44029

Signed-off-by: Ankit Jha <ankit.jha@tradomate.one>
This commit is contained in:
Ankit Jha 2026-10-02 02:33:01 +05:30
parent a14af739a4
commit f8856928b3
No known key found for this signature in database
2 changed files with 114 additions and 2 deletions

View file

@ -213,6 +213,35 @@ class _CombinedChunkSplitter:
piece.choices[0].delta = Delta(role=role, **fields)
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.
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.
"""
choices: Final = _optional_attr_sequence(chunk, "choices")
if len(choices) != 1:
return (chunk,)
delta: Final = _optional_attr(choices[0], "delta")
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,)
by_index: Final[dict[object, list[object]]] = {}
for position, call in enumerate(tool_calls):
call_index = _optional_attr(call, "index")
by_index.setdefault(position if call_index is None else call_index, []).append(call)
if len(by_index) < 2:
return (chunk,)
pieces: Final = tuple(copy.deepcopy(chunk) for _ in by_index)
for piece, calls in zip(pieces, by_index.values()):
piece.choices[0].delta.tool_calls = copy.deepcopy(calls)
return pieces
@staticmethod
def _normalize_reasoning_fields(fields: "dict[str, Any]") -> "dict[str, Any]":
"""Collapse signature-less ``thinking_blocks`` into ``reasoning_content``.
@ -270,9 +299,10 @@ class _CombinedChunkSplitter:
self._sync_iter = iter(self._stream)
chunk: Final = next(self._sync_iter) # propagates StopIteration when exhausted
self._buffer.extend(
split_chunk
tool_chunk
for combined_chunk in self._split(chunk)
for split_chunk in self._split_by_payload_kind(combined_chunk)
for tool_chunk in self._split_parallel_tool_calls(split_chunk)
)
return self._buffer.popleft()
@ -286,9 +316,10 @@ class _CombinedChunkSplitter:
self._async_iter = self._stream.__aiter__()
chunk: Final = await self._async_iter.__anext__() # propagates StopAsyncIteration
self._buffer.extend(
split_chunk
tool_chunk
for combined_chunk in self._split(chunk)
for split_chunk in self._split_by_payload_kind(combined_chunk)
for tool_chunk in self._split_parallel_tool_calls(split_chunk)
)
return self._buffer.popleft()

View file

@ -289,3 +289,84 @@ def test_thinking_then_signature_chunk_does_not_crash_stream():
assert "message_stop" in sse
assert "Done" in sse
def _tool_call(index, call_id=None, name=None, arguments=""):
from litellm.types.utils import ChatCompletionDeltaToolCall, Function
return ChatCompletionDeltaToolCall(
id=call_id,
index=index,
type="function",
function=Function(name=name, arguments=arguments),
)
def _tool_chunk(tool_calls, finish_reason=None):
return ModelResponseStream(
choices=[
StreamingChoices(
index=0,
delta=Delta(role="assistant", tool_calls=tool_calls),
finish_reason=finish_reason,
)
]
)
def _sse_events(raw: str):
events = []
for block in raw.split("\n\n"):
for line in block.splitlines():
if line.startswith("data: "):
events.append(json.loads(line[len("data: ") :]))
return events
def test_parallel_tool_calls_in_one_chunk_become_separate_blocks():
"""Several calls in one chunk must yield one tool_use block each (#44029)."""
chunk = _tool_chunk(
[
_tool_call(0, "call_a", "Read", '{"file_path": "a.md"}'),
_tool_call(1, "call_b", "Read", '{"file_path": "b.md"}'),
_tool_call(2, "call_c", "Glob", '{"pattern": "*.yml"}'),
]
)
wrapper = AnthropicStreamWrapper(completion_stream=iter([chunk]), model="mistral-small")
raw = "".join(b.decode() if isinstance(b, bytes) else b for b in wrapper.anthropic_sse_wrapper())
events = _sse_events(raw)
starts = [e["content_block"] for e in events if e["type"] == "content_block_start"]
tool_starts = [b for b in starts if b["type"] == "tool_use"]
assert [(b["name"]) for b in tool_starts] == ["Read", "Read", "Glob"]
assert len({b["id"] for b in tool_starts}) == 3
inputs = [
e["delta"]["partial_json"]
for e in events
if e["type"] == "content_block_delta" and e["delta"]["type"] == "input_json_delta"
]
assert [json.loads(x) for x in inputs] == [
{"file_path": "a.md"},
{"file_path": "b.md"},
{"pattern": "*.yml"},
]
def test_splitter_keeps_single_call_and_continuation_chunks_whole():
single = _tool_chunk([_tool_call(0, "call_a", "Read", "{}")])
continuation = _tool_chunk([_tool_call(0, None, None, '"x"')])
assert _CombinedChunkSplitter._split_parallel_tool_calls(single) == (single,)
assert _CombinedChunkSplitter._split_parallel_tool_calls(continuation) == (continuation,)
def test_splitter_orders_pieces_by_arrival_and_groups_same_index():
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]]