diff --git a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py index 4ef6c305cf5..42ddeac5ebc 100644 --- a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py @@ -5,6 +5,7 @@ import json import traceback from collections import deque from collections.abc import AsyncIterator, Iterator, Mapping, Sequence +from dataclasses import dataclass, replace from typing import ( TYPE_CHECKING, Any, @@ -32,7 +33,12 @@ from litellm.types.llms.anthropic import ( UsageDelta, UsageIteration, ) -from litellm.types.utils import AdapterCompletionStreamWrapper, Delta +from litellm.types.utils import ( + AdapterCompletionStreamWrapper, + ChatCompletionDeltaToolCall, + Delta, + Function, +) if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject @@ -91,6 +97,21 @@ def _delta_payload_field(delta_type: StreamingContentBlockDeltaType) -> str: assert_never(delta_type) +# Bounds the arguments held back while several calls are opened together. +_MAX_PENDING_ARGUMENT_CHARS: Final = 4 * 1024 * 1024 + + +@dataclass(frozen=True) +class _PendingCall: + """A tool call opened together with others whose arguments are still arriving.""" + + template: "ModelResponseStream" + call_id: str | None + name: str + index: int + arguments: str + + class _CombinedChunkSplitter: """ Splits a streaming chunk that carries BOTH response content and a @@ -116,6 +137,7 @@ class _CombinedChunkSplitter: self._sync_iter: Iterator[ModelResponseStream] | None = None self._async_iter: AsyncIterator[ModelResponseStream] | None = None self._buffer: deque[ModelResponseStream] = deque() + self._pending: Mapping[int, _PendingCall] = {} @property def chunks(self) -> "list[ModelResponseStream] | None": @@ -213,6 +235,133 @@ class _CombinedChunkSplitter: piece.choices[0].delta = Delta(role=role, **fields) return pieces + @staticmethod + def _call_name(tool_call: object) -> str | None: + name: Final = _optional_attr(_optional_attr(tool_call, "function"), "name") + return name if isinstance(name, str) and name else None + + @staticmethod + def _call_id(tool_call: object) -> str | None: + call_id: Final = _optional_attr(tool_call, "id") + return call_id if isinstance(call_id, str) else None + + @staticmethod + def _call_arguments(tool_call: object) -> str: + arguments: Final = _optional_attr(_optional_attr(tool_call, "function"), "arguments") + return arguments if isinstance(arguments, str) else "" + + @staticmethod + def _call_index(tool_call: object, position: int) -> int: + index: Final = _optional_attr(tool_call, "index") + return index if isinstance(index, int) else position + + @staticmethod + def _is_complete_json(text: str) -> bool: + try: + json.loads(text) + except ValueError: + return False + return bool(text) + + @staticmethod + def _chunk_with_calls( + chunk: "ModelResponseStream", calls: Sequence[object], *, keep_finish: bool + ) -> "ModelResponseStream": + """A copy of ``chunk`` whose only payload is ``calls``; ``chunk`` is not touched.""" + choice: Final = chunk.choices[0] + delta: Final = Delta(role=_optional_attr(choice.delta, "role"), tool_calls=list(calls)) + new_choice: Final = choice.model_copy( + update={"delta": delta, "finish_reason": choice.finish_reason if keep_finish else None} + ) + return chunk.model_copy(update={"choices": [new_choice]}) + + def _tool_calls_of(self, chunk: "ModelResponseStream") -> tuple[object, ...]: + """The tool calls of a single-choice chunk; empty for any other shape.""" + choices: Final = _optional_attr_sequence(chunk, "choices") + if len(choices) != 1: + return () + return tuple(_optional_attr_sequence(_optional_attr(choices[0], "delta"), "tool_calls")) + + def _route_tool_calls(self, chunk: "ModelResponseStream") -> tuple["ModelResponseStream", ...]: + """Emit the chunk, or one chunk per tool call when it opens several at once. + + Anthropic blocks are strictly sequential, so a call's arguments must be whole + when its block is written. Calls opened together with whole arguments are split + at once. If any argument is still incomplete, the calls are held and their + fragments gathered per index, then written in index order at the next chunk + that is not a fragment of them, or at the end of the stream. + """ + calls: Final = self._tool_calls_of(chunk) + if not self._pending and len(calls) < 2: + return (chunk,) # the common case: nothing held and at most one call + if self._pending and calls and all(self._extends_pending(call) for call in calls): + self._pending = self._with_fragments(calls) + return self._flush_pending() if self._pending_size() > _MAX_PENDING_ARGUMENT_CHARS else () + return (*self._flush_pending(), *self._open_calls(chunk, calls)) + + def _extends_pending(self, tool_call: object) -> bool: + index: Final = _optional_attr(tool_call, "index") + return self._call_name(tool_call) is None and isinstance(index, int) and index in self._pending + + def _with_fragments(self, calls: Sequence[object]) -> "Mapping[int, _PendingCall]": + return { + index: replace( + pending, + arguments=pending.arguments + + "".join(self._call_arguments(c) for c in calls if _optional_attr(c, "index") == index), + ) + for index, pending in self._pending.items() + } + + def _pending_size(self) -> int: + return sum(len(pending.arguments) for pending in self._pending.values()) + + def _flush_pending(self) -> tuple["ModelResponseStream", ...]: + """Write the held calls, one chunk each, in index order, and clear them.""" + held: Final = tuple(self._pending[index] for index in sorted(self._pending)) + self._pending = {} + return tuple( + self._chunk_with_calls( + pending.template, + [ + ChatCompletionDeltaToolCall( + id=pending.call_id, + index=pending.index, + type="function", + function=Function(name=pending.name, arguments=pending.arguments), + ) + ], + keep_finish=False, + ) + for pending in held + ) + + def _open_calls(self, chunk: "ModelResponseStream", calls: Sequence[object]) -> tuple["ModelResponseStream", ...]: + indexes: Final = tuple(dict.fromkeys(self._call_index(call, pos) for pos, call in enumerate(calls))) + groups: Final = tuple( + tuple(call for pos, call in enumerate(calls) if self._call_index(call, pos) == index) for index in indexes + ) + opened: Final = tuple(next((c for c in group if self._call_name(c)), None) for group in groups) + # Fewer than two distinct calls, or a fragment with no call to attach to: leave it alone. + if len(groups) < 2 or any(first is None for first in opened): + return (chunk,) + # An opening chunk over the bound is not held at all. + if sum(len(self._call_arguments(c)) for c in calls) > _MAX_PENDING_ARGUMENT_CHARS: + return (chunk,) + if all(self._is_complete_json("".join(self._call_arguments(c) for c in group)) for group in groups): + return tuple(self._chunk_with_calls(chunk, group, keep_finish=True) for group in groups) + self._pending = { + index: _PendingCall( + template=chunk, + call_id=self._call_id(first), + name=self._call_name(first) or "", + index=index, + arguments="".join(self._call_arguments(c) for c in group), + ) + for index, group, first in zip(indexes, groups, opened) + } + return () + @staticmethod def _normalize_reasoning_fields(fields: "dict[str, Any]") -> "dict[str, Any]": """Collapse signature-less ``thinking_blocks`` into ``reasoning_content``. @@ -260,36 +409,47 @@ class _CombinedChunkSplitter: finish_delta.thinking_blocks = None return [content_chunk, finish_chunk] + def _expand(self, chunk: "ModelResponseStream") -> tuple["ModelResponseStream", ...]: + return tuple( + routed + for combined_chunk in self._split(chunk) + for split_chunk in self._split_by_payload_kind(combined_chunk) + for routed in self._route_tool_calls(split_chunk) + ) + def __iter__(self) -> "Iterator[ModelResponseStream]": return self def __next__(self) -> "ModelResponseStream": - if self._buffer: - return self._buffer.popleft() if self._sync_iter is None: self._sync_iter = iter(self._stream) - chunk: Final = next(self._sync_iter) # propagates StopIteration when exhausted - self._buffer.extend( - split_chunk - for combined_chunk in self._split(chunk) - for split_chunk in self._split_by_payload_kind(combined_chunk) - ) + # A chunk held back as part of several calls yields nothing yet, so keep pulling. + while not self._buffer: + try: + chunk = next(self._sync_iter) + except StopIteration: + self._buffer.extend(self._flush_pending()) + if not self._buffer: + raise + break + self._buffer.extend(self._expand(chunk)) return self._buffer.popleft() def __aiter__(self) -> "AsyncIterator[ModelResponseStream]": return self async def __anext__(self) -> "ModelResponseStream": - if self._buffer: - return self._buffer.popleft() if self._async_iter is None: self._async_iter = self._stream.__aiter__() - chunk: Final = await self._async_iter.__anext__() # propagates StopAsyncIteration - self._buffer.extend( - split_chunk - for combined_chunk in self._split(chunk) - for split_chunk in self._split_by_payload_kind(combined_chunk) - ) + while not self._buffer: + try: + chunk = await self._async_iter.__anext__() + except StopAsyncIteration: + self._buffer.extend(self._flush_pending()) + if not self._buffer: + raise + break + self._buffer.extend(self._expand(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 9a6940ff12b..9aae5fb70ab 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 @@ -273,3 +273,182 @@ 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 _finish_chunk(): + return ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(), finish_reason="tool_calls")]) + + +def _text_chunk(text): + return ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(content=text), finish_reason=None)]) + + +def _tool_blocks(chunks): + """The tool_use blocks the Anthropic wrapper writes for these chunks: (name, id, input JSON).""" + wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="mistral-small") + raw = "".join(b.decode() if isinstance(b, bytes) else b for b in wrapper.anthropic_sse_wrapper()) + events = [] + for block in raw.split("\n\n"): + for line in block.splitlines(): + if line.startswith("data: "): + events.append(json.loads(line[len("data: ") :])) + blocks = {} + for event in events: + if event["type"] == "content_block_start" and event["content_block"]["type"] == "tool_use": + blocks[event["index"]] = [event["content_block"]["name"], event["content_block"]["id"], ""] + elif event["type"] == "content_block_delta" and event["delta"]["type"] == "input_json_delta": + blocks[event["index"]][2] += event["delta"]["partial_json"] + return [(name, call_id, json.loads(args)) for name, call_id, args in blocks.values()] + + +def test_parallel_tool_calls_in_one_chunk_become_separate_blocks(): + """Several whole calls in one chunk must yield one tool_use block each (#44029).""" + chunks = [ + _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"}'), + ] + ), + _finish_chunk(), + ] + + blocks = _tool_blocks(chunks) + + assert [(name, args) for name, _, args in blocks] == [ + ("Read", {"file_path": "a.md"}), + ("Read", {"file_path": "b.md"}), + ("Glob", {"pattern": "*.yml"}), + ] + assert len({call_id for _, call_id, _ in blocks}) == 3 + + +def test_calls_opened_together_with_empty_arguments_get_their_own_fragments(): + """Open both, then interleave the argument fragments by index: no fragment may land on the wrong call.""" + chunks = [ + _tool_chunk([_tool_call(0, "call_a", "Read", ""), _tool_call(1, "call_b", "Glob", "")]), + _tool_chunk([_tool_call(0, None, None, '{"file_path":')]), + _tool_chunk([_tool_call(1, None, None, '{"pattern": "*.yml"}')]), + _tool_chunk([_tool_call(0, None, None, ' "a.md"}')]), + _finish_chunk(), + ] + + blocks = _tool_blocks(chunks) + + assert [(name, call_id, args) for name, call_id, args in blocks] == [ + ("Read", blocks[0][1], {"file_path": "a.md"}), + ("Glob", blocks[1][1], {"pattern": "*.yml"}), + ] + assert blocks[0][1] != blocks[1][1] + + +def test_one_whole_call_and_one_partial_call_in_the_opening_chunk(): + chunks = [ + _tool_chunk( + [_tool_call(0, "call_a", "Read", '{"file_path": "a.md"}'), _tool_call(1, "call_b", "Glob", '{"pat')] + ), + _tool_chunk([_tool_call(1, None, None, 'tern": "*.yml"}')]), + _finish_chunk(), + ] + + assert [(name, args) for name, _, args in _tool_blocks(chunks)] == [ + ("Read", {"file_path": "a.md"}), + ("Glob", {"pattern": "*.yml"}), + ] + + +def test_held_calls_are_written_when_the_stream_ends_without_a_finish_chunk(): + chunks = [ + _tool_chunk([_tool_call(0, "call_a", "Read", '{"file_path"'), _tool_call(1, "call_b", "Glob", "")]), + _tool_chunk([_tool_call(0, None, None, ': "a.md"}')]), + _tool_chunk([_tool_call(1, None, None, '{"pattern": "*"}')]), + ] + + assert [(name, args) for name, _, args in _tool_blocks(chunks)] == [ + ("Read", {"file_path": "a.md"}), + ("Glob", {"pattern": "*"}), + ] + + +def test_text_after_held_calls_comes_after_their_blocks(): + chunks = [ + _tool_chunk([_tool_call(0, "call_a", "Read", ""), _tool_call(1, "call_b", "Glob", "")]), + _tool_chunk([_tool_call(0, None, None, "{}")]), + _tool_chunk([_tool_call(1, None, None, "{}")]), + _text_chunk("done"), + _finish_chunk(), + ] + + assert [(name, args) for name, _, args in _tool_blocks(chunks)] == [("Read", {}), ("Glob", {})] + + +def test_a_single_call_with_fragments_stays_one_block(): + chunks = [ + _tool_chunk([_tool_call(0, "call_a", "Read", "")]), + _tool_chunk([_tool_call(0, None, None, '{"file_path":')]), + _tool_chunk([_tool_call(0, None, None, ' "a.md"}')]), + _finish_chunk(), + ] + + assert [(name, args) for name, _, args in _tool_blocks(chunks)] == [("Read", {"file_path": "a.md"})] + + +def test_two_entries_for_the_same_index_stay_one_call(): + chunks = [ + _tool_chunk([_tool_call(0, "call_a", "Read", '{"file_path":'), _tool_call(0, None, None, ' "a.md"}')]), + _finish_chunk(), + ] + + assert [(name, args) for name, _, args in _tool_blocks(chunks)] == [("Read", {"file_path": "a.md"})] + + +def test_an_opening_chunk_over_the_bound_is_not_held(monkeypatch): + from litellm.llms.anthropic.pass_through.adapters import streaming_iterator + + monkeypatch.setattr(streaming_iterator, "_MAX_PENDING_ARGUMENT_CHARS", 10) + splitter = _CombinedChunkSplitter(iter(())) + oversized = _tool_chunk( + [_tool_call(0, "call_a", "Read", '{"file_path": "' + "a" * 50), _tool_call(1, "call_b", "Glob", "")] + ) + + assert splitter._expand(oversized) == (oversized,) + assert not splitter._pending + + +def test_fragments_that_push_held_calls_over_the_bound_flush_them(monkeypatch): + from litellm.llms.anthropic.pass_through.adapters import streaming_iterator + + monkeypatch.setattr(streaming_iterator, "_MAX_PENDING_ARGUMENT_CHARS", 10) + splitter = _CombinedChunkSplitter(iter(())) + opening = _tool_chunk([_tool_call(0, "call_a", "Read", ""), _tool_call(1, "call_b", "Glob", "")]) + fragment = _tool_chunk([_tool_call(0, None, None, '{"file_path": "' + "a" * 50)]) + + assert splitter._expand(opening) == () + flushed = splitter._expand(fragment) + + assert [c.choices[0].delta.tool_calls[0].function.name for c in flushed] == ["Read", "Glob"] + assert not splitter._pending