diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index f2810ac0abe..5542760a738 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -1,6 +1,7 @@ import time import uuid from collections.abc import Sequence +from dataclasses import dataclass from itertools import count from typing import Any, Final, cast @@ -51,6 +52,39 @@ from litellm.types.utils import ( ) +@dataclass(frozen=True, slots=True) +class _StreamedToolCallMetadata: + call_type: str | None + tool_name: str | None + tool_namespace: str | None + ambiguous: bool = False + + def matches(self, incoming: "_StreamedToolCallMetadata") -> bool: + return all( + existing_value is None or incoming_value is None or existing_value == incoming_value + for existing_value, incoming_value in zip( + (self.call_type, self.tool_name, self.tool_namespace), + (incoming.call_type, incoming.tool_name, incoming.tool_namespace), + ) + ) + + def merged_with(self, incoming: "_StreamedToolCallMetadata") -> "_StreamedToolCallMetadata": + return _StreamedToolCallMetadata( + call_type=self.call_type or incoming.call_type, + tool_name=self.tool_name or incoming.tool_name, + tool_namespace=self.tool_namespace or incoming.tool_namespace, + ambiguous=self.ambiguous, + ) + + def marked_ambiguous(self) -> "_StreamedToolCallMetadata": + return _StreamedToolCallMetadata( + call_type=self.call_type, + tool_name=self.tool_name, + tool_namespace=self.tool_namespace, + ambiguous=True, + ) + + def _index_of_output_item_type(items: Sequence[object], item_type: str) -> int | None: return next( (index for index, item in enumerate(items) if getattr(item, "type", None) == item_type), @@ -117,6 +151,9 @@ 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._tool_call_metadata_by_index: dict[ + int, _StreamedToolCallMetadata + ] = {} # mutable-ok: streamed metadata and ambiguity accumulate across chunks 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._streamed_call_id_by_terminal_id: dict[str, str] = {} # mutable-ok: terminal identity correlation @@ -158,6 +195,9 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return None def _streamed_tool_call_id_at_position(self, position: int) -> str | None: + indexed_metadata: Final = self._tool_call_metadata_by_index.get(position) + if indexed_metadata is not None and indexed_metadata.ambiguous: + return None indexed_call_id: Final = self._tool_call_id_by_index.get(position) if indexed_call_id is not None: return indexed_call_id @@ -172,8 +212,17 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): def _streamed_tool_call_id_for_terminal_call(self, tool_call: object, position: int) -> str | None: """Match a terminal aggregate tool call to the identity emitted while streaming.""" + call_id_raw: Final = tool_call.get("id") if isinstance(tool_call, dict) else getattr(tool_call, "id", None) + if call_id_raw: + call_id_match: Final = self._streamed_call_id_by_terminal_id.get(str(call_id_raw)) + if call_id_match is not None: + return call_id_match + tool_call_index: Final = self._normalize_tool_call_index(tool_call) if tool_call_index is not None: + indexed_metadata: Final = self._tool_call_metadata_by_index.get(tool_call_index) + if indexed_metadata is not None and indexed_metadata.ambiguous: + return None indexed_call_id: Final = self._tool_call_id_by_index.get(tool_call_index) if indexed_call_id is not None: return indexed_call_id @@ -181,22 +230,55 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return None return self._streamed_tool_call_id_at_position(position) - def _resolve_streamed_tool_call_id(self, tool_call_index: int | None, call_id_raw: object) -> str | None: + def _resolve_streamed_tool_call_id( + self, + tool_call_index: int | None, + call_id_raw: object, + metadata: _StreamedToolCallMetadata, + ) -> str | None: if tool_call_index is None: - return str(call_id_raw) if call_id_raw else None + if not call_id_raw: + return None + call_id: Final = str(call_id_raw) + self._streamed_call_id_by_terminal_id[call_id] = call_id + return call_id + + incoming_call_id: Final = str(call_id_raw) if call_id_raw else None + known_call_id: Final = ( + self._streamed_call_id_by_terminal_id.get(incoming_call_id) if incoming_call_id is not None else None + ) indexed_call_id: Final = self._tool_call_id_by_index.get(tool_call_index) if indexed_call_id is not None: - if call_id_raw: - self._streamed_call_id_by_terminal_id[str(call_id_raw)] = indexed_call_id - return indexed_call_id + if known_call_id is not None and known_call_id != indexed_call_id: + return known_call_id + indexed_metadata: Final = self._tool_call_metadata_by_index.get( + tool_call_index, + _StreamedToolCallMetadata(None, None, None), + ) + if incoming_call_id is None: + return None if indexed_metadata.ambiguous else indexed_call_id - if not call_id_raw: + if indexed_metadata.matches(metadata): + self._tool_call_metadata_by_index[tool_call_index] = indexed_metadata.merged_with(metadata) + self._streamed_call_id_by_terminal_id[incoming_call_id] = indexed_call_id + return indexed_call_id + + self._tool_call_metadata_by_index[tool_call_index] = indexed_metadata.marked_ambiguous() + if incoming_call_id == indexed_call_id: + return None + self._streamed_call_id_by_terminal_id[incoming_call_id] = incoming_call_id + return incoming_call_id + + if incoming_call_id is None: return None + if known_call_id is not None: + return known_call_id - call_id: Final = str(call_id_raw) - self._tool_call_id_by_index[tool_call_index] = call_id - return call_id + self._tool_call_id_by_index[tool_call_index] = incoming_call_id + self._tool_call_metadata_by_index[tool_call_index] = metadata + self._streamed_call_id_by_terminal_id[incoming_call_id] = incoming_call_id + return incoming_call_id def _responses_namespace_tool_call_fields(self, fn_name: str) -> tuple[str, str | None]: mapped: Final = self._namespace_tool_names.get(fn_name) @@ -267,10 +349,6 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): for tc in tool_calls: tc_index = self._normalize_tool_call_index(tc) call_id_raw = tc.get("id") if isinstance(tc, dict) else getattr(tc, "id", None) - call_id = self._resolve_streamed_tool_call_id(tc_index, call_id_raw) - if call_id is None: - continue - fn = tc.get("function") if isinstance(tc, dict) else getattr(tc, "function", None) fn_name = "" fn_args_delta = "" @@ -281,6 +359,15 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): fn_name = str(getattr(fn, "name", "") or "") fn_args_delta = serialize_tool_call_arguments(getattr(fn, "arguments", "")) tool_name, tool_namespace = self._responses_namespace_tool_call_fields(fn_name) + call_type_raw = tc.get("type") if isinstance(tc, dict) else getattr(tc, "type", None) + metadata = _StreamedToolCallMetadata( + call_type=str(call_type_raw) if call_type_raw else None, + tool_name=tool_name or None, + tool_namespace=tool_namespace, + ) + call_id = self._resolve_streamed_tool_call_id(tc_index, call_id_raw, metadata) + if call_id is None: + continue output_index = self._get_or_assign_tool_output_index(call_id) self._queue_first_streamed_tool_call_event( diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py index e2569c588df..a106c3ffa53 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -3410,6 +3410,7 @@ class TestEnsureOutputItemContentPartAdded: iterator._tool_args_by_call_id = {} iterator._tool_item_id_by_call_id = {} iterator._tool_call_id_by_index = {} + iterator._tool_call_metadata_by_index = {} iterator._streamed_tool_call_ids_in_order = [] iterator._resolved_tool_call_id_by_position = {} iterator._streamed_call_id_by_terminal_id = {} diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py b/tests/test_litellm/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py index 4ed8b478e1c..227d672d02c 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py @@ -804,7 +804,10 @@ def test_completed_snapshot_correlates_function_after_server_tool_replacement(): assert function_call.arguments == '{"city":"Paris"}' -def test_reused_index_with_new_call_id_preserves_first_streamed_identity(): +@pytest.mark.parametrize("include_replacement_metadata", (True, False)) +def test_reused_index_with_new_call_id_preserves_first_streamed_identity( + include_replacement_metadata: bool, +): iterator = LiteLLMCompletionStreamingIterator( model="test-model", litellm_custom_stream_wrapper=AsyncMock(), @@ -822,15 +825,22 @@ def test_reused_index_with_new_call_id_preserves_first_streamed_identity(): } ] ) + replacement_call = ( + { + "index": 0, + "id": "call_b", + "type": "function", + "function": {"name": "tool_a", "arguments": "1"}, + } + if include_replacement_metadata + else { + "index": 0, + "id": "call_b", + "function": {"arguments": "1"}, + } + ) iterator._queue_tool_call_delta_events( - [ - { - "index": 0, - "id": "call_b", - "type": "function", - "function": {"name": "tool_a", "arguments": "1"}, - } - ] + [replacement_call] ) iterator._queue_tool_call_delta_events( [ @@ -901,6 +911,119 @@ def test_reused_index_with_new_call_id_preserves_first_streamed_identity(): assert (completed_call.id, completed_call.call_id) == ("fc_call_a", "call_a") +@pytest.mark.parametrize( + ("replacement_type", "replacement_name"), + (("function", "tool_b"), ("custom", "tool_a")), +) +def test_reused_index_with_changed_tool_metadata_starts_separate_identity( + replacement_type: str, + replacement_name: str, +): + iterator = LiteLLMCompletionStreamingIterator( + model="test-model", + litellm_custom_stream_wrapper=AsyncMock(), + request_input="Test input", + responses_api_request={}, + ) + + iterator._queue_tool_call_delta_events( + [ + { + "index": 0, + "id": "call_a", + "type": "function", + "function": {"name": "tool_a", "arguments": '{"safe":'}, + } + ] + ) + iterator._queue_tool_call_delta_events( + [ + { + "index": 0, + "id": "call_b", + "type": "function", + "function": {"name": "tool_a", "arguments": ""}, + } + ] + ) + iterator._queue_tool_call_delta_events( + [ + { + "index": 0, + "id": "call_b", + "type": replacement_type, + "function": {"name": replacement_name, "arguments": '{"privileged":'}, + } + ] + ) + iterator._queue_tool_call_delta_events( + [{"index": 0, "id": "call_b", "function": {"arguments": "true}"}}] + ) + iterator._queue_tool_call_delta_events( + [{"index": 0, "function": {"arguments": "ignored"}}] + ) + terminal_response = ModelResponse( + id="chatcmpl-terminal", + created=123, + model="test-model", + object="chat.completion", + choices=[ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "index": 0, + "id": "call_b", + "type": "function", + "function": { + "name": replacement_name, + "arguments": '{"privileged":true}', + }, + } + ], + }, + } + ], + ) + iterator._queue_final_tool_call_done_events(terminal_response) + completed = iterator._emit_response_completed_event(terminal_response) + + added_items = [ + event.item + for event in iterator._pending_tool_events + if event.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + ] + delta_events = [ + event + for event in iterator._pending_tool_events + if event.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA + ] + done_item = next( + event.item + for event in iterator._pending_tool_events + if event.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE + ) + + assert completed is not None + assert [(item.id, item.call_id) for item in added_items] == [ + ("fc_call_a", "call_a"), + ("fc_call_b", "call_b"), + ] + assert "".join(event.delta for event in delta_events if event.item_id == "fc_call_a") == '{"safe":' + assert "".join(event.delta for event in delta_events if event.item_id == "fc_call_b") == '{"privileged":true}' + assert (done_item.id, done_item.call_id, done_item.name) == ("fc_call_b", "call_b", replacement_name) + completed_call = next(item for item in completed.response.output if item.type == "function_call") + assert (completed_call.id, completed_call.call_id, completed_call.name) == ( + "fc_call_b", + "call_b", + replacement_name, + ) + + @pytest.mark.asyncio async def test_streaming_events_share_the_chat_completion_response_id(): """