mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(responses): preserve streamed function call identity
This commit is contained in:
parent
0cf236bebb
commit
f9f96c58c0
2 changed files with 409 additions and 7 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue