fix(guardrails): key Responses stream tool-call rewrites by call_id

Bridged Responses streams give reasoning and message items output_index 0
and start function calls at 1, so keying rewrites by output_index rewrote
the wrong items. Rewrites now follow each function call's call_id through
the buffered item and argument events, refuse when an event cannot be
resolved to a rewritten call, and the refusal branches on all three
handlers get regression tests
This commit is contained in:
mateo-berri 2026-09-08 13:22:30 -07:00
parent f59354d09b
commit 89e11949c8
4 changed files with 259 additions and 44 deletions

View file

@ -140,6 +140,10 @@ _TERMINAL_ENVELOPE_EVENT_TYPES: Final = frozenset(
)
_FUNCTION_CALL_ARGUMENT_EVENT_TYPES: Final = frozenset(
{"response.function_call_arguments.delta", "response.function_call_arguments.done"}
)
_OUTPUT_ITEM_EVENT_TYPES: Final = frozenset({"response.output_item.added", "response.output_item.done"})
_PATCHABLE_ITEM_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
{"function_call_output": "output", "message": "content"}
)
@ -930,70 +934,110 @@ class OpenAIResponsesHandler(BaseTranslation):
) -> None:
"""Write ended-stream guardrail tool-call rewrites into the completed
envelope's ``function_call`` items and sync the earlier stream events,
keyed by ``output_index``. The guardrail sees the envelope's function
calls in output order, which is how a rewritten call finds its item; a
rewrite whose calls do not line up with the envelope is reported as
undeliverable, so the pipeline executor discards it and releases the
original events."""
keyed by ``call_id``. The guardrail sees the envelope's function calls
in output order, which is how a rewritten call finds its ``call_id``;
the stream events find their call through the ``call_id`` on
``output_item`` events and the ``item_id`` on argument events, since an
event's ``output_index`` need not match the envelope's (the chat bridge
numbers tool calls from 1 while the envelope lists them after the
message). A rewrite whose calls do not line up with the envelope, or
whose events cannot be found, is reported as undeliverable, so the
pipeline executor discards it and releases the original events."""
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
return
function_call_indices: Final = tuple(
output_idx
for output_idx, output_item in enumerate(outputs)
if stream_item_field(output_item, "type") == "function_call"
function_call_items: Final = tuple(
output_item for output_item in outputs if stream_item_field(output_item, "type") == "function_call"
)
if len(function_call_indices) != len(post_guardrail_tool_calls):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
raise UndeliverableStreamRewrite(guardrail_name)
rewrites_by_output_index: Final = MappingProxyType(
call_ids: Final = tuple(
call_id
for output_item in function_call_items
if isinstance(call_id := stream_item_field(output_item, "call_id"), str) and call_id
)
stream_events: Final = responses_so_far[:-1]
call_id_by_item_id: Final = self._function_call_ids_by_item_id(stream_events)
event_call_ids: Final = tuple(
self._function_call_event_call_id(event, call_id_by_item_id) for event in stream_events
)
rewrites_by_call_id: Final = MappingProxyType(
{
output_idx: after
for output_idx, before, after in zip(
function_call_indices, pre_guardrail_tool_calls, post_guardrail_tool_calls
)
call_id: after
for call_id, before, after in zip(call_ids, pre_guardrail_tool_calls, post_guardrail_tool_calls)
if after != before
}
)
for output_idx, rewrite in rewrites_by_output_index.items():
self._write_function_call_item(outputs[output_idx], rewrite.name, rewrite.arguments)
self._sync_stream_events_with_tool_call_rewrites(
stream_events=responses_so_far[:-1],
rewrites_by_output_index=rewrites_by_output_index,
unresolved_argument_event: Final = any(
call_id is None and stream_item_field(event, "type") in _FUNCTION_CALL_ARGUMENT_EVENT_TYPES
for event, call_id in zip(stream_events, event_call_ids)
)
if (
len(call_ids) != len(function_call_items)
or len(frozenset(call_ids)) != len(call_ids)
or len(call_ids) != len(post_guardrail_tool_calls)
or unresolved_argument_event
or not rewrites_by_call_id.keys() <= frozenset(event_call_ids)
):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
def _sync_stream_events_with_tool_call_rewrites(
self,
stream_events: Sequence[object],
rewrites_by_output_index: Mapping[int, _ToolCallShape],
) -> None:
"""Sync pre-completion function-call events with the rewritten completed
response: the first ``function_call_arguments.delta`` for a rewritten call
carries the full rewritten arguments and the rest are blanked, while
``function_call_arguments.done`` and ``output_item.done`` carry the full
rewritten arguments and ``output_item.added`` / ``output_item.done`` the
rewritten name, so every event a client may read agrees with the
rewritten ``response.completed`` payload."""
raise UndeliverableStreamRewrite(guardrail_name)
for output_item, rewrite in (
(output_item, rewrites_by_call_id[call_id])
for output_item, call_id in zip(function_call_items, call_ids)
if call_id in rewrites_by_call_id
):
self._write_function_call_item(output_item, rewrite.name, rewrite.arguments)
delta_replacements: Final = MappingProxyType(
{index: chain((rewrite.arguments,), repeat("")) for index, rewrite in rewrites_by_output_index.items()}
{call_id: chain((rewrite.arguments,), repeat("")) for call_id, rewrite in rewrites_by_call_id.items()}
)
for event in stream_events:
output_index = stream_item_field(event, "output_index")
if not isinstance(output_index, int) or output_index not in rewrites_by_output_index:
for event, call_id in zip(stream_events, event_call_ids):
if call_id not in rewrites_by_call_id:
continue
rewrite = rewrites_by_output_index[output_index]
match stream_item_field(event, "type"):
case "response.function_call_arguments.delta":
self._write_event_field(event, "delta", next(delta_replacements[output_index]))
self._write_event_field(event, "delta", next(delta_replacements[call_id]))
case "response.function_call_arguments.done":
self._write_event_field(event, "arguments", rewrite.arguments)
self._write_event_field(event, "arguments", rewrites_by_call_id[call_id].arguments)
case "response.output_item.added":
self._write_function_call_item(stream_item_field(event, "item"), rewrite.name, None)
self._write_function_call_item(
stream_item_field(event, "item"), rewrites_by_call_id[call_id].name, None
)
case "response.output_item.done":
self._write_function_call_item(stream_item_field(event, "item"), rewrite.name, rewrite.arguments)
self._write_function_call_item(
stream_item_field(event, "item"),
rewrites_by_call_id[call_id].name,
rewrites_by_call_id[call_id].arguments,
)
case _:
pass
@staticmethod
def _function_call_ids_by_item_id(stream_events: Sequence[object]) -> Mapping[str, str]:
items: Final = tuple(
stream_item_field(event, "item")
for event in stream_events
if stream_item_field(event, "type") in _OUTPUT_ITEM_EVENT_TYPES
)
return MappingProxyType(
{
item_id: call_id
for item in items
if stream_item_field(item, "type") == "function_call"
and isinstance(item_id := stream_item_field(item, "id"), str)
and isinstance(call_id := stream_item_field(item, "call_id"), str)
}
)
@staticmethod
def _function_call_event_call_id(event: object, call_id_by_item_id: Mapping[str, str]) -> str | None:
event_type: Final = stream_item_field(event, "type")
if event_type in _FUNCTION_CALL_ARGUMENT_EVENT_TYPES:
item_id: Final = stream_item_field(event, "item_id")
return call_id_by_item_id.get(item_id) if isinstance(item_id, str) else None
if event_type not in _OUTPUT_ITEM_EVENT_TYPES:
return None
item: Final = stream_item_field(event, "item")
call_id: Final = stream_item_field(item, "call_id")
return call_id if stream_item_field(item, "type") == "function_call" and isinstance(call_id, str) else None
@staticmethod
def _write_function_call_item(item: object, name: str | None, arguments: str | None) -> None:
if item is None:

View file

@ -381,6 +381,31 @@ class TestAnthropicMessagesHandlerStreamingOutputProcessing:
assert chunks == original
@pytest.mark.asyncio
async def test_deliver_ended_stream_tool_use_rewrite_with_server_tool_use_block_fails_closed(self):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
handler = AnthropicMessagesHandler()
server_tool_use = [
("content_block_start", {"type": "content_block_start", "index": 0, "content_block": {"type": "server_tool_use", "id": "srvtoolu_1", "name": "web_search", "input": {}}}),
("content_block_delta", {"type": "content_block_delta", "index": 0, "delta": {"type": "input_json_delta", "partial_json": '{"query": "fruit"}'}}),
("content_block_stop", {"type": "content_block_stop", "index": 0}),
]
tool_use = self._ended_tool_use_sse_chunks()
chunks = (
tool_use[:1]
+ [f"event: {name}\ndata: {json.dumps(payload)}\n\n".encode() for name, payload in server_tool_use]
+ [chunk.replace(b'"index": 0', b'"index": 1') for chunk in tool_use[1:]]
)
with pytest.raises(UndeliverableStreamRewrite):
await handler.process_output_streaming_response(
responses_so_far=chunks,
guardrail_to_apply=self._argument_masking_guardrail(),
litellm_logging_obj=MagicMock(),
deliver_ended_stream_rewrites=True,
)
@pytest.mark.asyncio
async def test_ended_stream_rewrite_leaves_chunks_untouched_by_default(self):
handler = AnthropicMessagesHandler()

View file

@ -1252,6 +1252,62 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
deliver_ended_stream_rewrites=True,
)
@staticmethod
def _two_choice_tool_call_stream_chunks() -> list:
from litellm.types.utils import (
ChatCompletionDeltaToolCall,
Delta,
Function,
ModelResponseStream,
StreamingChoices,
)
def chunk(
choice_index: int, tool_call: ChatCompletionDeltaToolCall | None, finish_reason: Optional[str] = None
) -> ModelResponseStream:
return ModelResponseStream(
id="chatcmpl-123",
created=1234567890,
model="gpt-4",
object="chat.completion.chunk",
choices=[
StreamingChoices(
index=choice_index,
delta=Delta(tool_calls=[tool_call] if tool_call else None),
finish_reason=finish_reason,
)
],
)
def fragment(arguments: str, name: Optional[str] = None, call_id: Optional[str] = None):
return ChatCompletionDeltaToolCall(
id=call_id, index=0, type="function", function=Function(name=name, arguments=arguments)
)
return [
chunk(0, fragment("", name="lookup_fruit", call_id="call_1")),
chunk(1, fragment("", name="lookup_fruit", call_id="call_2")),
chunk(0, fragment('{"fruit": "persimmon"}')),
chunk(1, fragment('{"fruit": "durian"}')),
chunk(0, None, finish_reason="tool_calls"),
chunk(1, None, finish_reason="tool_calls"),
]
@pytest.mark.asyncio
async def test_deliver_ended_stream_tool_call_rewrite_on_multi_choice_stream_fails_closed(self):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
handler = OpenAIChatCompletionsHandler()
chunks = self._two_choice_tool_call_stream_chunks()
with pytest.raises(UndeliverableStreamRewrite):
await handler.process_output_streaming_response(
responses_so_far=chunks,
guardrail_to_apply=MockGuardrail(guardrail_name="test"),
litellm_logging_obj=None,
deliver_ended_stream_rewrites=True,
)
@pytest.mark.asyncio
async def test_deliver_ended_stream_clean_multi_choice_stream_released_untouched(self):
handler = OpenAIChatCompletionsHandler()

View file

@ -1314,6 +1314,96 @@ class TestOpenAIResponsesHandlerStreamingOutputProcessing:
assert completed_event.response.output[0].arguments == '{"fruit": "[MASKED]"}'
assert completed_event.response.output[0].name == "lookup_fruit"
@staticmethod
def _bridged_function_call_stream_events() -> List[dict]:
reasoning = {"type": "reasoning", "id": "rs_1", "summary": []}
text = {"type": "output_text", "text": "Looking that up", "annotations": []}
message = {"type": "message", "id": "msg_1", "role": "assistant", "status": "completed", "content": [text]}
def function_call(arguments: str, status: str) -> dict:
return {
"type": "function_call",
"id": "fc_1",
"call_id": "call_1",
"name": "lookup_fruit",
"arguments": arguments,
"status": status,
}
return [
{"type": "response.output_item.added", "output_index": 0, "item": dict(reasoning)},
{"type": "response.output_item.done", "output_index": 0, "item": dict(reasoning)},
{"type": "response.output_item.added", "output_index": 0, "item": {**message, "status": "in_progress", "content": []}},
{"type": "response.output_text.delta", "item_id": "msg_1", "output_index": 0, "content_index": 0, "delta": "Looking that up"},
{"type": "response.output_item.done", "output_index": 0, "item": {**message, "content": [dict(text)]}},
{"type": "response.output_item.added", "output_index": 1, "item": function_call("", "in_progress")},
{"type": "response.function_call_arguments.delta", "item_id": "fc_1", "output_index": 1, "delta": '{"fruit":'},
{"type": "response.function_call_arguments.delta", "item_id": "fc_1", "output_index": 1, "delta": ' "persimmon"}'},
{
"type": "response.function_call_arguments.done",
"item_id": "fc_1",
"output_index": 1,
"arguments": '{"fruit": "persimmon"}',
},
{"type": "response.output_item.done", "output_index": 1, "item": function_call('{"fruit": "persimmon"}', "completed")},
{
"type": "response.completed",
"response": {
"id": "resp_1",
"model": "claude-haiku-4-5",
"output": [
dict(reasoning),
{**message, "content": [dict(text)]},
function_call('{"fruit": "persimmon"}', "completed"),
],
},
},
]
@pytest.mark.asyncio
async def test_deliver_ended_stream_rewrites_keys_bridged_function_call_events_by_call_id(self):
handler = OpenAIResponsesHandler()
events = self._bridged_function_call_stream_events()
await handler.process_output_streaming_response(
responses_so_far=events,
guardrail_to_apply=self._argument_masking_guardrail(),
litellm_logging_obj=None,
deliver_ended_stream_rewrites=True,
)
assert events[6]["delta"] == '{"fruit": "[MASKED]"}'
assert events[7]["delta"] == ""
assert events[8]["arguments"] == '{"fruit": "[MASKED]"}'
assert events[5]["item"]["name"] == "lookup_fruit"
assert events[9]["item"]["arguments"] == '{"fruit": "[MASKED]"}'
assert events[10]["response"]["output"][2]["arguments"] == '{"fruit": "[MASKED]"}'
assert events[3]["delta"] == "Looking that up"
assert events[4]["item"]["content"][0]["text"] == "Looking that up"
assert events[10]["response"]["output"][1]["content"][0]["text"] == "Looking that up"
assert events[1]["item"] == {"type": "reasoning", "id": "rs_1", "summary": []}
@pytest.mark.asyncio
@pytest.mark.parametrize("mismatch", ["orphan_call_id", "duplicate_call_id"])
async def test_deliver_ended_stream_function_call_rewrite_without_matching_events_fails_closed(self, mismatch):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
handler = OpenAIResponsesHandler()
events = self._ended_function_call_stream_events()
envelope_item = events[5]["response"]["output"][0]
if mismatch == "orphan_call_id":
events[5]["response"]["output"] = [{**envelope_item, "call_id": "call_999"}]
else:
events[5]["response"]["output"] = [dict(envelope_item), dict(envelope_item)]
with pytest.raises(UndeliverableStreamRewrite):
await handler.process_output_streaming_response(
responses_so_far=events,
guardrail_to_apply=self._argument_masking_guardrail(),
litellm_logging_obj=None,
deliver_ended_stream_rewrites=True,
)
@pytest.mark.asyncio
async def test_ended_stream_function_call_rewrite_leaves_events_untouched_by_default(self):
handler = OpenAIResponsesHandler()