fix(responses): enforce tool metadata identity boundaries

This commit is contained in:
Dan Loftus 2026-09-01 22:12:10 -04:00
parent fa9cb19307
commit cddb23dc5d
3 changed files with 233 additions and 22 deletions

View file

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

View file

@ -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 = {}

View file

@ -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():
"""