mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
a14af739a4
commit
f8856928b3
2 changed files with 114 additions and 2 deletions
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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]]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue