mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
refactor(responses): align streamed identities immutably
This commit is contained in:
parent
07dcc007b5
commit
bfca46a7db
2 changed files with 37 additions and 28 deletions
|
|
@ -1,6 +1,7 @@
|
|||
import time
|
||||
import uuid
|
||||
from collections.abc import Sequence
|
||||
from itertools import count
|
||||
from typing import Any, Final, cast
|
||||
|
||||
import litellm
|
||||
|
|
@ -116,8 +117,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
self._tool_args_by_call_id: dict[str, str] = {}
|
||||
self._tool_item_id_by_call_id: dict[str, str] = {} # mutable-ok: filled per call id as tool call events stream
|
||||
self._tool_call_id_by_index: dict[int, str] = {}
|
||||
self._streamed_tool_call_ids_in_order: list[str] = []
|
||||
self._resolved_tool_call_id_by_position: dict[int, str] = {}
|
||||
self._streamed_tool_call_ids_in_order: list[str] = [] # mutable-ok: accumulates call ids across stream chunks
|
||||
self._resolved_tool_call_id_by_position: dict[int, str] = {} # mutable-ok: terminal correlation state
|
||||
self._ambiguous_tool_call_indexes: set[int] = set()
|
||||
self._next_tool_output_index: int = 1 # output_index=0 reserved for the message item
|
||||
self._final_tool_events_queued: bool = False
|
||||
|
|
@ -203,12 +204,9 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
return
|
||||
|
||||
self._tool_args_by_call_id[call_id] = ""
|
||||
streamed_call_ids = getattr(self, "_streamed_tool_call_ids_in_order", None)
|
||||
if streamed_call_ids is None:
|
||||
streamed_call_ids = self._streamed_tool_call_ids_in_order = []
|
||||
streamed_call_ids.append(call_id)
|
||||
self._streamed_tool_call_ids_in_order.append(call_id)
|
||||
|
||||
item_kwargs = build_tool_call_item_kwargs(
|
||||
item_kwargs: Final = build_tool_call_item_kwargs(
|
||||
call_id,
|
||||
tool_name,
|
||||
"",
|
||||
|
|
@ -220,7 +218,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
item_kwargs["namespace"] = tool_namespace
|
||||
|
||||
self._sequence_number += 1
|
||||
event = OutputItemAddedEvent(
|
||||
event: Final = OutputItemAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
|
||||
output_index=output_index,
|
||||
item=BaseLiteLLMOpenAIResponseObject(**item_kwargs),
|
||||
|
|
@ -337,10 +335,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
if not call_id_raw:
|
||||
continue
|
||||
call_id = self._streamed_tool_call_id_for_terminal_call(tc, position) or str(call_id_raw)
|
||||
resolved_call_ids = getattr(self, "_resolved_tool_call_id_by_position", None)
|
||||
if resolved_call_ids is None:
|
||||
resolved_call_ids = self._resolved_tool_call_id_by_position = {}
|
||||
resolved_call_ids[position] = call_id
|
||||
self._resolved_tool_call_id_by_position[position] = call_id
|
||||
output_index = self._get_or_assign_tool_output_index(call_id)
|
||||
|
||||
fn = tc.get("function") if isinstance(tc, dict) else getattr(tc, "function", None)
|
||||
|
|
@ -1215,6 +1210,27 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
chat_completion_delta: Final[ChatCompletionDelta] = choice.delta
|
||||
return chat_completion_delta.content or ""
|
||||
|
||||
def _output_item_with_streamed_tool_identity(self, item: object, tool_position: int) -> object:
|
||||
resolved_call_id: Final = self._resolved_tool_call_id_by_position.get(tool_position)
|
||||
streamed_call_id: Final = (
|
||||
resolved_call_id if resolved_call_id is not None else self._streamed_tool_call_id_at_position(tool_position)
|
||||
)
|
||||
if streamed_call_id is None:
|
||||
return item
|
||||
|
||||
streamed_item_id: Final = self._tool_item_id_by_call_id.get(
|
||||
streamed_call_id,
|
||||
getattr(item, "id", streamed_call_id),
|
||||
)
|
||||
identity_update: Final = { # mutable-ok: Pydantic model_copy requires a mapping update payload
|
||||
"id": streamed_item_id,
|
||||
"call_id": streamed_call_id,
|
||||
}
|
||||
copy_with_identity: Final = getattr(item, "model_copy", None)
|
||||
if not callable(copy_with_identity):
|
||||
return item
|
||||
return copy_with_identity(update=identity_update)
|
||||
|
||||
def _output_with_streamed_item_ids(self, responses_api_response: ResponsesAPIResponse) -> tuple[Any, ...]:
|
||||
"""
|
||||
Reuse the item IDs already emitted by the incremental streaming events in the
|
||||
|
|
@ -1231,21 +1247,13 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
"reasoning",
|
||||
self._cached_reasoning_item_id,
|
||||
)
|
||||
tool_position = 0
|
||||
aligned_items: list[Any] = []
|
||||
for item in reasoning_aligned:
|
||||
if getattr(item, "type", None) in {"function_call", "custom_tool_call"}:
|
||||
resolved_call_ids: Final = getattr(self, "_resolved_tool_call_id_by_position", {})
|
||||
streamed_call_id = resolved_call_ids.get(tool_position)
|
||||
if streamed_call_id is None:
|
||||
streamed_call_id = self._streamed_tool_call_id_at_position(tool_position)
|
||||
tool_position += 1
|
||||
if streamed_call_id is not None:
|
||||
tool_item_ids: Final = getattr(self, "_tool_item_id_by_call_id", {})
|
||||
streamed_item_id = tool_item_ids.get(streamed_call_id, getattr(item, "id", streamed_call_id))
|
||||
item = item.model_copy(update={"id": streamed_item_id, "call_id": streamed_call_id})
|
||||
aligned_items.append(item)
|
||||
return tuple(aligned_items)
|
||||
tool_positions: Final = count()
|
||||
return tuple(
|
||||
self._output_item_with_streamed_tool_identity(item, next(tool_positions))
|
||||
if getattr(item, "type", None) in ("function_call", "custom_tool_call")
|
||||
else item
|
||||
for item in reasoning_aligned
|
||||
)
|
||||
|
||||
def _emit_response_completed_event(self, litellm_model_response: ModelResponse) -> ResponseCompletedEvent | None:
|
||||
if litellm_model_response:
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ import json
|
|||
|
||||
import pytest
|
||||
|
||||
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
TOOL_CALLS_CACHE,
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
|
|
@ -3411,6 +3410,8 @@ class TestEnsureOutputItemContentPartAdded:
|
|||
iterator._tool_args_by_call_id = {}
|
||||
iterator._tool_item_id_by_call_id = {}
|
||||
iterator._tool_call_id_by_index = {}
|
||||
iterator._streamed_tool_call_ids_in_order = []
|
||||
iterator._resolved_tool_call_id_by_position = {}
|
||||
iterator._ambiguous_tool_call_indexes = set()
|
||||
iterator._next_tool_output_index = 1
|
||||
iterator._final_tool_events_queued = False
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue