mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 7055dc3375 into be4481779e
This commit is contained in:
commit
231f2e076f
2 changed files with 356 additions and 17 deletions
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue