fix(responses): preserve streamed function call identity

This commit is contained in:
Dan Loftus 2026-09-01 15:16:47 -04:00
parent 0cf236bebb
commit f9f96c58c0
2 changed files with 409 additions and 7 deletions

View file

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

View file

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