From 3fb1128a5f6391d339d35a94bd30f8127543909c Mon Sep 17 00:00:00 2001 From: Dan Loftus Date: Tue, 1 Sep 2026 20:32:05 -0400 Subject: [PATCH] fix(responses): correlate completed tools by call id --- .../streaming_iterator.py | 19 +++- .../test_litellm_completion_responses.py | 1 + .../test_streaming_iterator_transformation.py | 99 +++++++++++++++++++ 3 files changed, 116 insertions(+), 3 deletions(-) diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index 193c4f6417c..3e138429f99 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -119,6 +119,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._tool_call_id_by_index: 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._streamed_call_id_by_terminal_id: dict[str, str] = {} # mutable-ok: terminal identity correlation 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 @@ -334,8 +335,10 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): call_id_raw = tc.get("id") if isinstance(tc, dict) else getattr(tc, "id", None) if not call_id_raw: continue - call_id = self._streamed_tool_call_id_for_terminal_call(tc, position) or str(call_id_raw) + terminal_call_id = str(call_id_raw) + call_id = self._streamed_tool_call_id_for_terminal_call(tc, position) or terminal_call_id self._resolved_tool_call_id_by_position[position] = call_id + self._streamed_call_id_by_terminal_id[terminal_call_id] = 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) @@ -1211,9 +1214,19 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): 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) + terminal_call_id: Final = getattr(item, "call_id", None) + resolved_by_call_id: Final = ( + self._streamed_call_id_by_terminal_id.get(terminal_call_id) if isinstance(terminal_call_id, str) else None + ) + resolved_by_position: 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) + resolved_by_call_id + if resolved_by_call_id is not None + else ( + resolved_by_position + if resolved_by_position is not None + else self._streamed_tool_call_id_at_position(tool_position) + ) ) if streamed_call_id is None: return item 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 e5d8a0075ca..e2569c588df 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 @@ -3412,6 +3412,7 @@ class TestEnsureOutputItemContentPartAdded: iterator._tool_call_id_by_index = {} iterator._streamed_tool_call_ids_in_order = [] iterator._resolved_tool_call_id_by_position = {} + iterator._streamed_call_id_by_terminal_id = {} iterator._ambiguous_tool_call_indexes = set() iterator._next_tool_output_index = 1 iterator._final_tool_events_queued = False 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 25a2cd9a6ab..855612a5eda 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 @@ -705,6 +705,105 @@ def test_terminal_only_call_is_not_conflated_with_later_streamed_call(): ] +def test_completed_snapshot_correlates_function_after_server_tool_replacement(): + 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_exec_stream", + "type": "function", + "function": { + "name": "bash_code_execution", + "arguments": '{"command":"printf server"}', + }, + }, + { + "index": 1, + "id": "call_regular_stream", + "type": "function", + "function": { + "name": "lookup_weather", + "arguments": '{"city":"Paris"}', + }, + }, + ] + ) + iterator._pending_tool_events.clear() + 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": "srvtoolu_exec_terminal", + "type": "function", + "function": { + "name": "bash_code_execution", + "arguments": '{"command":"printf server"}', + }, + }, + { + "id": "call_regular_terminal", + "type": "function", + "function": { + "name": "lookup_weather", + "arguments": '{"city":"Paris"}', + }, + }, + ], + "provider_specific_fields": { + "code_interpreter_results": [ + { + "type": "code_interpreter_call", + "id": "srvtoolu_exec_terminal", + "code": "printf server", + "container_id": None, + "status": "completed", + "outputs": [{"type": "logs", "logs": "server"}], + } + ] + }, + }, + } + ], + ) + + iterator._queue_final_tool_call_done_events(terminal_response) + completed = iterator._emit_response_completed_event(terminal_response) + + assert completed is not None + code_calls = [item for item in completed.response.output if item.type == "code_interpreter_call"] + function_calls = [item for item in completed.response.output if item.type == "function_call"] + assert len(code_calls) == 1 + assert len(function_calls) == 1 + code_call = code_calls[0] + assert code_call.id == "srvtoolu_exec_terminal" + assert code_call.code == "printf server" + assert code_call.container_id is None + assert code_call.outputs[0].logs == "server" + function_call = function_calls[0] + assert (function_call.id, function_call.call_id) == ( + "fc_call_regular_stream", + "call_regular_stream", + ) + assert function_call.name == "lookup_weather" + assert function_call.arguments == '{"city":"Paris"}' + + def test_reused_index_with_new_call_id_marks_fallback_ambiguous(): iterator = LiteLLMCompletionStreamingIterator( model="test-model",