diff --git a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py index 4ef6c305cf5..1a2080fc1d1 100644 --- a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py @@ -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() diff --git a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_combined_chunk.py b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_combined_chunk.py index 5a7cf652b95..c4dcc56d5e9 100644 --- a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_combined_chunk.py +++ b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_combined_chunk.py @@ -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]]