This commit is contained in:
Ankit Jha 2026-10-05 08:53:45 +08:00 • committed by GitHub
commit 231f2e076f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 356 additions and 17 deletions

View file

@ -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()

View file

@ -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