mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
f59354d09b
commit
89e11949c8
4 changed files with 259 additions and 44 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue