From f9f96c58c00b767a7ca3475526a7fc6c56803490 Mon Sep 17 00:00:00 2001 From: Dan Loftus Date: Tue, 1 Sep 2026 15:16:47 -0400 Subject: [PATCH] fix(responses): preserve streamed function call identity --- .../streaming_iterator.py | 82 ++++- .../test_streaming_iterator_transformation.py | 334 +++++++++++++++++- 2 files changed, 409 insertions(+), 7 deletions(-) diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index b2edf2bf9ed..d389488e71a 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -114,7 +114,10 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._pending_tool_events: list[BaseLiteLLMOpenAIResponseObject] = [] self._tool_output_index_by_call_id: dict[str, int] = {} self._tool_args_by_call_id: dict[str, str] = {} + self._tool_item_id_by_call_id: dict[str, str] = {} 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._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 @@ -153,6 +156,34 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): except (TypeError, ValueError): return None + def _streamed_tool_call_id_at_position(self, position: int) -> str | None: + if position in self._ambiguous_tool_call_indexes: + 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 + # If the stream supplied any indexes, a missing position is a terminal-only + # call. Falling back to arrival order here could conflate parallel calls. + if self._tool_call_id_by_index: + return None + streamed_call_ids: Final = getattr(self, "_streamed_tool_call_ids_in_order", ()) + if position < len(streamed_call_ids): + return streamed_call_ids[position] + return None + + 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.""" + tool_call_index: Final = self._normalize_tool_call_index(tool_call) + if tool_call_index is not None: + if tool_call_index in self._ambiguous_tool_call_indexes: + 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 + if self._tool_call_id_by_index: + return None + return self._streamed_tool_call_id_at_position(position) + def _responses_namespace_tool_call_fields(self, fn_name: str) -> tuple[str, str | None]: mapped: Final = self._namespace_tool_names.get(fn_name) if mapped: @@ -224,9 +255,17 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): if call_id not in self._tool_args_by_call_id: 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._sequence_number += 1 names = self._custom_tool_names item_kwargs = build_tool_call_item_kwargs(call_id, tool_name, "", "in_progress", names) + tool_item_ids = getattr(self, "_tool_item_id_by_call_id", None) + if tool_item_ids is None: + tool_item_ids = self._tool_item_id_by_call_id = {} + tool_item_ids[call_id] = item_kwargs["id"] if tool_namespace: item_kwargs["namespace"] = tool_namespace event = OutputItemAddedEvent( @@ -248,7 +287,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._sequence_number += 1 delta_event: BaseLiteLLMOpenAIResponseObject = FunctionCallArgumentsDeltaEvent( type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, - item_id=call_id, + item_id=self._tool_item_id_by_call_id.get(call_id, call_id), output_index=output_index, delta=delta_chunk, ) @@ -273,11 +312,15 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): if not tool_calls or not isinstance(tool_calls, list): return - for tc in tool_calls: + for position, tc in enumerate(tool_calls): call_id_raw = tc.get("id") if isinstance(tc, dict) else getattr(tc, "id", None) if not call_id_raw: continue - call_id = str(call_id_raw) + 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 output_index = self._get_or_assign_tool_output_index(call_id) fn = tc.get("function") if isinstance(tc, dict) else getattr(tc, "function", None) @@ -300,6 +343,10 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._sequence_number += 1 names = self._custom_tool_names item_kwargs = build_tool_call_item_kwargs(call_id, tool_name, "", "in_progress", names) + tool_item_ids = getattr(self, "_tool_item_id_by_call_id", None) + if tool_item_ids is None: + tool_item_ids = self._tool_item_id_by_call_id = {} + tool_item_ids[call_id] = item_kwargs["id"] if tool_namespace: item_kwargs["namespace"] = tool_namespace event = OutputItemAddedEvent( @@ -325,7 +372,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._sequence_number += 1 delta_event = FunctionCallArgumentsDeltaEvent( type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, - item_id=call_id, + item_id=self._tool_item_id_by_call_id.get(call_id, call_id), output_index=output_index, delta=delta_chunk, ) @@ -335,7 +382,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._sequence_number += 1 done_event = FunctionCallArgumentsDoneEvent( type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE, - item_id=call_id, + item_id=self._tool_item_id_by_call_id.get(call_id, call_id), output_index=output_index, arguments=final_args, ) @@ -345,6 +392,10 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._sequence_number += 1 names = self._custom_tool_names item_kwargs = build_tool_call_item_kwargs(call_id, tool_name, final_args, "completed", names) + tool_item_ids = getattr(self, "_tool_item_id_by_call_id", None) + if tool_item_ids is None: + tool_item_ids = self._tool_item_id_by_call_id = {} + item_kwargs["id"] = tool_item_ids.setdefault(call_id, item_kwargs["id"]) if tool_namespace: item_kwargs["namespace"] = tool_namespace item_done_event = OutputItemDoneEvent( @@ -1161,7 +1212,26 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): "message", self._cached_item_id, ) - return _output_items_with_id(message_aligned, "reasoning", self._cached_reasoning_item_id) + reasoning_aligned: Final = _output_items_with_id( + message_aligned, + "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) def _emit_response_completed_event(self, litellm_model_response: ModelResponse) -> ResponseCompletedEvent | None: if litellm_model_response: 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 01148f627f1..cd54accc616 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 @@ -11,7 +11,7 @@ spend tracking stores, so a follow-up previous_response_id still finds the conve """ import json -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -408,6 +408,338 @@ def test_parallel_tool_calls_without_ids_use_index_mapping(): assert arguments_by_call_id["call_b"] == '{"y":2}' +def test_final_tool_events_and_completed_snapshot_reuse_streamed_call_identity(): + iterator = LiteLLMCompletionStreamingIterator( + model="test-model", + litellm_custom_stream_wrapper=AsyncMock(), + request_input="Test input", + responses_api_request={}, + ) + streamed_ids = ["call_stream_a", "call_stream_b"] + terminal_ids = ["call_terminal_a", "call_terminal_b"] + + iterator._queue_tool_call_delta_events( + [ + { + "index": index, + "id": call_id, + "type": "function", + "function": { + "name": f"tool_{index}", + "arguments": f'{{"value":{index}', + }, + } + for index, call_id in reversed(list(enumerate(streamed_ids))) + ] + ) + # Simulate delivery of all incremental events before the terminal aggregate arrives. + iterator._pending_tool_events.clear() + + iterator.litellm_model_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": [ + { + "id": terminal_id, + "type": "function", + "function": { + "name": f"tool_{index}", + "arguments": f'{{"value":{index}}}', + }, + } + for index, terminal_id in enumerate(terminal_ids) + ], + }, + } + ], + ) + + final_events = [] + for _ in range(20): + event = iterator.common_done_event_logic() + final_events.append(event) + if event.type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED: + break + else: + pytest.fail("response.completed was not emitted") + + final_tool_items = [ + event.item + for event in final_events + if event.type + in { + ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, + } + and getattr(event.item, "type", None) == "function_call" + ] + assert [item.id for item in final_tool_items] == streamed_ids + assert [item.call_id for item in final_tool_items] == streamed_ids + + argument_event_ids = [ + event.item_id + for event in final_events + if event.type + in { + ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, + ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE, + } + ] + assert argument_event_ids + assert set(argument_event_ids) == set(streamed_ids) + + completed = final_events[-1] + completed_calls = [item for item in completed.response.output if item.type == "function_call"] + assert [item.id for item in completed_calls] == streamed_ids + assert [item.call_id for item in completed_calls] == streamed_ids + assert not set(terminal_ids) & { + item_id for item in final_tool_items + completed_calls for item_id in (item.id, item.call_id) + } + + +def test_final_events_preserve_distinct_streamed_item_id_and_call_id(): + from litellm.responses.litellm_completion_transformation.custom_tools import ( + build_tool_call_item_kwargs, + ) + from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, + ) + + iterator = LiteLLMCompletionStreamingIterator( + model="test-model", + litellm_custom_stream_wrapper=AsyncMock(), + request_input="Test input", + responses_api_request={}, + ) + 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": [ + { + "id": "call_terminal", + "type": "function", + "function": {"name": "tool", "arguments": '{"value":1}'}, + } + ], + }, + } + ], + ) + + def distinct_item_id_builder(call_id, *args, **kwargs): + item_kwargs = build_tool_call_item_kwargs(call_id, *args, **kwargs) + if call_id == "call_stream": + item_kwargs["id"] = "fc_stream" + return item_kwargs + + with patch( + "litellm.responses.litellm_completion_transformation.streaming_iterator.build_tool_call_item_kwargs", + side_effect=distinct_item_id_builder, + ): + iterator._queue_tool_call_delta_events( + [ + { + "index": 0, + "id": "call_stream", + "type": "function", + "function": {"name": "tool", "arguments": '{"value":'}, + } + ] + ) + streamed_added = next( + event + for event in iterator._pending_tool_events + if event.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + ) + assert (streamed_added.item.id, streamed_added.item.call_id) == ( + "fc_stream", + "call_stream", + ) + assert { + event.item_id + for event in iterator._pending_tool_events + if event.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA + } == {"fc_stream"} + iterator._pending_tool_events.clear() + iterator._queue_final_tool_call_done_events(terminal_response) + + original_transform = ( + LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response + ) + + def terminal_response_with_distinct_item_id(*args, **kwargs): + response = original_transform(*args, **kwargs) + terminal_call = next(item for item in response.output if item.type == "function_call") + terminal_call.id = "fc_terminal" + terminal_call.call_id = "call_terminal" + return response + + with patch.object( + LiteLLMCompletionResponsesConfig, + "transform_chat_completion_response_to_responses_api_response", + side_effect=terminal_response_with_distinct_item_id, + ): + completed = iterator._emit_response_completed_event(terminal_response) + + assert completed is not None + argument_events = [ + event + for event in iterator._pending_tool_events + if event.type + in { + ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, + ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE, + } + ] + assert argument_events + assert {event.item_id for event in argument_events} == {"fc_stream"} + final_item = next( + event.item + for event in iterator._pending_tool_events + if event.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE + ) + assert (final_item.id, final_item.call_id) == ("fc_stream", "call_stream") + completed_call = next(item for item in completed.response.output if item.type == "function_call") + assert (completed_call.id, completed_call.call_id) == ("fc_stream", "call_stream") + + +def test_terminal_only_tool_calls_keep_terminal_identity(): + iterator = LiteLLMCompletionStreamingIterator( + model="test-model", + litellm_custom_stream_wrapper=AsyncMock(), + request_input="Test input", + responses_api_request={}, + ) + terminal_ids = ["call_terminal_a", "call_terminal_b"] + response = ModelResponse( + id="chatcmpl-terminal-only", + created=123, + model="test-model", + object="chat.completion", + choices=[ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": call_id, + "type": "function", + "function": {"name": f"tool_{index}", "arguments": "{}"}, + } + for index, call_id in enumerate(terminal_ids) + ], + }, + } + ], + ) + + iterator._queue_final_tool_call_done_events(response) + added_items = [ + event.item + for event in iterator._pending_tool_events + if event.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + ] + + assert [item.id for item in added_items] == terminal_ids + assert [item.call_id for item in added_items] == terminal_ids + + +def test_terminal_only_call_is_not_conflated_with_later_streamed_call(): + iterator = LiteLLMCompletionStreamingIterator( + model="test-model", + litellm_custom_stream_wrapper=AsyncMock(), + request_input="Test input", + responses_api_request={}, + ) + iterator._queue_tool_call_delta_events( + [ + { + "index": 1, + "id": "call_streamed", + "type": "function", + "function": {"name": "streamed_tool", "arguments": "{}"}, + } + ] + ) + iterator._pending_tool_events.clear() + response = ModelResponse( + id="chatcmpl-mixed", + created=123, + model="test-model", + object="chat.completion", + choices=[ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_terminal_only", + "type": "function", + "function": {"name": "terminal_tool", "arguments": "{}"}, + }, + { + "id": "call_terminal_drifted", + "type": "function", + "function": {"name": "streamed_tool", "arguments": "{}"}, + }, + ], + }, + } + ], + ) + + iterator._queue_final_tool_call_done_events(response) + added_or_done_items = [ + event.item + for event in iterator._pending_tool_events + if event.type + in { + ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, + } + ] + completed = iterator._emit_response_completed_event(response) + + assert completed is not None + assert {item.id for item in added_or_done_items} == { + "call_terminal_only", + "call_streamed", + } + completed_calls = [item for item in completed.response.output if item.type == "function_call"] + assert [item.id for item in completed_calls] == [ + "call_terminal_only", + "call_streamed", + ] + assert [item.call_id for item in completed_calls] == [ + "call_terminal_only", + "call_streamed", + ] + + def test_reused_index_with_new_call_id_marks_fallback_ambiguous(): iterator = LiteLLMCompletionStreamingIterator( model="test-model",