refactor(responses): align streamed identities immutably

This commit is contained in:
Dan Loftus 2026-09-01 17:11:29 -04:00
parent 07dcc007b5
commit bfca46a7db
2 changed files with 37 additions and 28 deletions

View file

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

View file

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