fix(guardrails): scan and rewrite Responses custom_tool_call output items

Post-call guardrails on /v1/responses only treated function_call output
items as tool calls, so a custom_tool_call item (Codex's exec shell tool
on GPT-5.6 models) was never scanned or masked, non-streaming and
streaming alike. Both item types now flow through the shared
tool_call_dict_from_output_item helper, ended-stream delivery syncs the
custom_tool_call_input delta/done events and the item's input field, and
the completed-response scan key fingerprints both kinds of item.

Non-streaming Responses tool-call MASK rewrites were also never written
back to the output item even for function_call; they are now.
This commit is contained in:
mateo-berri 2026-09-09 13:38:51 -07:00
parent ddedb4867b
commit a9cce4f1ed
3 changed files with 322 additions and 72 deletions

View file

@ -227,7 +227,7 @@ class _ChatToolCallDict(ChatCompletionToolCallChunk, total=False):
provider_specific_fields: Mapping[str, object]
def _tool_call_dict_from_output_item(item: Mapping[str, Any], index: int) -> _ChatToolCallDict:
def tool_call_dict_from_output_item(item: Mapping[str, Any], index: int) -> _ChatToolCallDict:
"""Convert a ``function_call`` or ``custom_tool_call`` output item dict to a chat
completions tool_call dict. Custom (grammar/freeform) tool calls carry their raw
string payload in ``input`` rather than ``arguments``; both map to
@ -755,7 +755,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
# Tool calls accumulate into the single trailing tool_calls choice
# like the typed branches above; a choice per call would hide every
# call after choices[0] from chat clients
accumulated_tool_calls.append(_tool_call_dict_from_output_item(raw_item, tool_call_index))
accumulated_tool_calls.append(tool_call_dict_from_output_item(raw_item, tool_call_index))
tool_call_index += 1
elif handle_raw_dict_callback is not None:
choice, index = handle_raw_dict_callback(item=raw_item, index=index)
@ -1409,7 +1409,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
# New output item added
output_item = parsed_chunk.get("item", {})
if output_item.get("type") in ("function_call", "custom_tool_call"):
converted: Final = _tool_call_dict_from_output_item(output_item, parsed_chunk.get("output_index", 0))
converted: Final = tool_call_dict_from_output_item(output_item, parsed_chunk.get("output_index", 0))
provider_specific_fields: Final = converted.get("provider_specific_fields")
function_chunk: Final = ChatCompletionToolCallFunctionChunk(
@ -1484,7 +1484,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
index=0,
delta=Delta(
tool_calls=(
_tool_call_dict_from_output_item(
tool_call_dict_from_output_item(
output_item, parsed_chunk.get("output_index", 0)
),
)

View file

@ -37,7 +37,6 @@ from itertools import accumulate, chain, repeat
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, NamedTuple, Union, cast
from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall
from pydantic import BaseModel, TypeAdapter
from typing_extensions import ReadOnly, TypedDict
@ -45,6 +44,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
OpenAiResponsesToChatCompletionStreamIterator,
tool_call_dict_from_output_item,
)
from litellm.llms.base_llm.guardrail_translation.base_translation import (
BaseTranslation,
@ -84,7 +84,6 @@ from litellm.types.llms.openai import (
)
from litellm.types.responses.main import (
GenericResponseOutputItem,
OutputFunctionToolCall,
OutputText,
)
from litellm.types.utils import GenericGuardrailAPIInputs
@ -140,8 +139,18 @@ _TERMINAL_ENVELOPE_EVENT_TYPES: Final = frozenset(
)
_FUNCTION_CALL_ARGUMENT_EVENT_TYPES: Final = frozenset(
{"response.function_call_arguments.delta", "response.function_call_arguments.done"}
_TOOL_CALL_ITEM_TYPES: Final = frozenset({"function_call", "custom_tool_call"})
_TOOL_CALL_PAYLOAD_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
{"function_call": "arguments", "custom_tool_call": "input"}
)
_TOOL_CALL_PAYLOAD_DELTA_EVENT_TYPES: Final = frozenset(
{"response.function_call_arguments.delta", "response.custom_tool_call_input.delta"}
)
_TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
{"response.function_call_arguments.done": "arguments", "response.custom_tool_call_input.done": "input"}
)
_TOOL_CALL_PAYLOAD_EVENT_TYPES: Final = _TOOL_CALL_PAYLOAD_DELTA_EVENT_TYPES | frozenset(
_TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS
)
_OUTPUT_ITEM_EVENT_TYPES: Final = frozenset({"response.output_item.added", "response.output_item.done"})
_PATCHABLE_ITEM_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
@ -180,8 +189,20 @@ def _rewritten_input_item(item: Mapping[str, object], rewritten: object) -> Mapp
return {**item, field: converted_value} # mutable-ok: request input items must stay JSON-plain dicts
def _is_function_call_item(item: object) -> bool:
return isinstance(item, Mapping) and item.get("type") in ("function_call", "custom_tool_call")
def _is_tool_call_item(item: object) -> bool:
return isinstance(item, Mapping) and item.get("type") in _TOOL_CALL_ITEM_TYPES
def _tool_call_output_item_mapping(item: object) -> Mapping[str, Any] | None:
if stream_item_field(item, "type") not in _TOOL_CALL_ITEM_TYPES:
return None
if isinstance(item, Mapping):
return cast("Mapping[str, Any]", item) # cast-ok: output items are str-keyed JSON objects
return item.model_dump() if isinstance(item, BaseModel) else None
def _is_tool_call_output_item(item: object) -> bool:
return _tool_call_output_item_mapping(item) is not None
def _last_message_role(messages: Sequence[object]) -> str | None:
@ -205,7 +226,7 @@ def _provenance_unit_bounds(
start_indexes: Final = tuple(
index
for index in range(len(raw_input))
if index == 0 or not (_is_function_call_item(raw_input[index]) and trailing_roles[index - 1] == "assistant")
if index == 0 or not (_is_tool_call_item(raw_input[index]) and trailing_roles[index - 1] == "assistant")
)
return tuple(zip(start_indexes, (*start_indexes[1:], len(raw_input))))
@ -603,7 +624,7 @@ class OpenAIResponsesHandler(BaseTranslation):
- response.output is a list of output items
- Each output item can be:
* GenericResponseOutputItem with a content list of OutputText objects
* ResponseFunctionToolCall with tool call data
* ResponseFunctionToolCall or CustomToolCallOutputItem with tool call data
- Each OutputText object has a text field
"""
@ -668,6 +689,7 @@ class OpenAIResponsesHandler(BaseTranslation):
if response_model:
inputs["model"] = response_model
pre_guardrail_tool_calls: Final = _tool_call_shapes(tool_calls_to_check)
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=request_data,
@ -676,6 +698,7 @@ class OpenAIResponsesHandler(BaseTranslation):
)
guardrailed_texts: Final = guardrailed_inputs.get("texts", [])
returned_tool_calls: Final = guardrailed_inputs.get("tool_calls")
# Step 3: Map guardrail responses back to original response structure
await self._apply_guardrail_responses_to_output(
@ -683,6 +706,15 @@ class OpenAIResponsesHandler(BaseTranslation):
responses=guardrailed_texts,
task_mappings=task_mappings,
)
self._write_tool_call_rewrites_to_output(
tool_call_items=tuple(item for item in response_output if _is_tool_call_output_item(item)),
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
post_guardrail_tool_calls=_tool_call_shapes(
returned_tool_calls
if isinstance(returned_tool_calls, list) and len(returned_tool_calls) == len(tool_calls_to_check)
else tool_calls_to_check
),
)
verbose_proxy_logger.debug("OpenAI Responses API: Processed output response: %s", response)
@ -933,11 +965,12 @@ class OpenAIResponsesHandler(BaseTranslation):
guardrail_name: str,
) -> None:
"""Write ended-stream guardrail tool-call rewrites into the completed
envelope's ``function_call`` items and sync the earlier stream 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
envelope's ``function_call`` and ``custom_tool_call`` items and sync the
earlier stream events, keyed by ``call_id``. The guardrail sees the
envelope's tool 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
and custom-input 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
@ -945,18 +978,16 @@ class OpenAIResponsesHandler(BaseTranslation):
pipeline executor discards it and releases the original events."""
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
return
function_call_items: Final = tuple(
output_item for output_item in outputs if stream_item_field(output_item, "type") == "function_call"
)
tool_call_items: Final = tuple(output_item for output_item in outputs if _is_tool_call_output_item(output_item))
call_ids: Final = tuple(
call_id
for output_item in function_call_items
for output_item in tool_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)
call_id_by_item_id: Final = self._tool_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
self._tool_call_event_call_id(event, call_id_by_item_id) for event in stream_events
)
rewrites_by_call_id: Final = MappingProxyType(
{
@ -966,11 +997,11 @@ class OpenAIResponsesHandler(BaseTranslation):
}
)
unresolved_argument_event: Final = any(
call_id is None and stream_item_field(event, "type") in _FUNCTION_CALL_ARGUMENT_EVENT_TYPES
call_id is None and stream_item_field(event, "type") in _TOOL_CALL_PAYLOAD_EVENT_TYPES
for event, call_id in zip(stream_events, event_call_ids)
)
if (
len(call_ids) != len(function_call_items)
len(call_ids) != len(tool_call_items)
or len(frozenset(call_ids)) != len(call_ids)
or len(call_ids) != len(post_guardrail_tool_calls)
or unresolved_argument_event
@ -981,10 +1012,10 @@ class OpenAIResponsesHandler(BaseTranslation):
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)
for output_item, call_id in zip(tool_call_items, call_ids)
if call_id in rewrites_by_call_id
):
self._write_function_call_item(output_item, rewrite.name, rewrite.arguments)
self._write_tool_call_item(output_item, rewrite.name, rewrite.arguments)
delta_replacements: Final = MappingProxyType(
{call_id: chain((rewrite.arguments,), repeat("")) for call_id, rewrite in rewrites_by_call_id.items()}
)
@ -992,16 +1023,18 @@ class OpenAIResponsesHandler(BaseTranslation):
if call_id not in rewrites_by_call_id:
continue
match stream_item_field(event, "type"):
case "response.function_call_arguments.delta":
case str() as event_type if event_type in _TOOL_CALL_PAYLOAD_DELTA_EVENT_TYPES:
self._write_event_field(event, "delta", next(delta_replacements[call_id]))
case "response.function_call_arguments.done":
self._write_event_field(event, "arguments", rewrites_by_call_id[call_id].arguments)
case str() as event_type if event_type in _TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS:
self._write_event_field(
event, _TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS[event_type], rewrites_by_call_id[call_id].arguments
)
case "response.output_item.added":
self._write_function_call_item(
self._write_tool_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(
self._write_tool_call_item(
stream_item_field(event, "item"),
rewrites_by_call_id[call_id].name,
rewrites_by_call_id[call_id].arguments,
@ -1009,8 +1042,23 @@ class OpenAIResponsesHandler(BaseTranslation):
case _:
pass
def _write_tool_call_rewrites_to_output(
self,
tool_call_items: Sequence[object],
pre_guardrail_tool_calls: tuple[_ToolCallShape, ...],
post_guardrail_tool_calls: tuple[_ToolCallShape, ...],
) -> None:
if len(tool_call_items) != len(post_guardrail_tool_calls):
return
for output_item, after in (
(output_item, after)
for output_item, before, after in zip(tool_call_items, pre_guardrail_tool_calls, post_guardrail_tool_calls)
if after != before
):
self._write_tool_call_item(output_item, after.name, after.arguments)
@staticmethod
def _function_call_ids_by_item_id(stream_events: Sequence[object]) -> Mapping[str, str]:
def _tool_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
@ -1020,32 +1068,35 @@ class OpenAIResponsesHandler(BaseTranslation):
{
item_id: call_id
for item in items
if stream_item_field(item, "type") == "function_call"
if stream_item_field(item, "type") in _TOOL_CALL_ITEM_TYPES
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:
def _tool_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:
if event_type in _TOOL_CALL_PAYLOAD_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
return (
call_id if stream_item_field(item, "type") in _TOOL_CALL_ITEM_TYPES and isinstance(call_id, str) else None
)
@staticmethod
def _write_function_call_item(item: object, name: str | None, arguments: str | None) -> None:
def _write_tool_call_item(item: object, name: str | None, payload: str | None) -> None:
if item is None:
return
if name is not None:
OpenAIResponsesHandler._write_event_field(item, "name", name)
if arguments is not None:
OpenAIResponsesHandler._write_event_field(item, "arguments", arguments)
item_type: Final = stream_item_field(item, "type")
if payload is not None and isinstance(item_type, str) and item_type in _TOOL_CALL_PAYLOAD_FIELDS:
OpenAIResponsesHandler._write_event_field(item, _TOOL_CALL_PAYLOAD_FIELDS[item_type], payload)
def _check_streaming_has_ended(self, responses_so_far: Sequence[object]) -> bool:
"""
@ -1073,7 +1124,7 @@ class OpenAIResponsesHandler(BaseTranslation):
def _completed_response_scan_key(response: object) -> StreamingScanKey:
output_items: Final = stream_item_items(response, "output")
message_items: Final = tuple(
item for item in output_items if stream_item_field(item, "type") != "function_call"
item for item in output_items if stream_item_field(item, "type") not in _TOOL_CALL_ITEM_TYPES
)
return StreamingScanKey(
texts=tuple(
@ -1085,7 +1136,7 @@ class OpenAIResponsesHandler(BaseTranslation):
tool_calls=tuple(
stream_item_fingerprint(item)
for item in output_items
if stream_item_field(item, "type") == "function_call"
if stream_item_field(item, "type") in _TOOL_CALL_ITEM_TYPES
),
stream_ended=True,
)
@ -1196,34 +1247,10 @@ class OpenAIResponsesHandler(BaseTranslation):
Override this method to customize text/image/tool extraction logic.
"""
# Check if this is a tool call (OutputFunctionToolCall)
if isinstance(output_item, OutputFunctionToolCall) or (
isinstance(output_item, BaseModel)
and hasattr(output_item, "type")
and getattr(output_item, "type") == "function_call"
):
tool_call_item: Final = _tool_call_output_item_mapping(output_item)
if tool_call_item is not None:
if tool_calls_to_check is not None:
tool_call_dict = (
LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call(
tool_call_item=output_item,
index=output_idx,
)
)
tool_calls_to_check.append(cast(ChatCompletionToolCallChunk, tool_call_dict))
return
elif isinstance(output_item, dict) and output_item.get("type") == "function_call":
# Handle dict representation of tool call
if tool_calls_to_check is not None:
# Convert dict to ResponseFunctionToolCall for processing
try:
tool_call_obj: Final = ResponseFunctionToolCall(**output_item)
tool_call_dict = LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call(
tool_call_item=tool_call_obj,
index=output_idx,
)
tool_calls_to_check.append(cast(ChatCompletionToolCallChunk, tool_call_dict))
except Exception:
pass
tool_calls_to_check.append(tool_call_dict_from_output_item(tool_call_item, output_idx))
return
# Handle both GenericResponseOutputItem and dict

View file

@ -14,7 +14,12 @@ import pytest
from fastapi import HTTPException
from openai.types.responses import ResponseFunctionToolCall
from openai.types.responses import (
ResponseCustomToolCall,
ResponseCustomToolCallInputDeltaEvent,
ResponseCustomToolCallInputDoneEvent,
ResponseFunctionToolCall,
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -27,7 +32,7 @@ from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.responses.main import GenericResponseOutputItem, OutputText
from litellm.types.responses.main import CustomToolCallOutputItem, GenericResponseOutputItem, OutputText
from litellm.types.utils import CallTypes, GenericGuardrailAPIInputs
@ -57,6 +62,37 @@ class MockGuardrail(CustomGuardrail):
return inputs
class PersimmonMaskingGuardrail(CustomGuardrail):
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
tool_calls = [
{
**tool_call,
"function": {
**tool_call["function"],
"arguments": tool_call["function"]["arguments"].replace("persimmon", "[MASKED]"),
},
}
for tool_call in inputs.get("tool_calls", [])
]
return {**inputs, "tool_calls": tool_calls}
CUSTOM_TOOL_CALL_ITEM = {
"type": "custom_tool_call",
"id": "ctc_1",
"call_id": "call_exec_1",
"name": "exec",
"input": "echo persimmon",
"status": "completed",
}
class TestOpenAIResponsesHandlerDiscovery:
"""Test that the handler is properly discovered by the guardrail system"""
@ -628,6 +664,77 @@ class TestOpenAIResponsesHandlerToolCallExtraction:
== '{"location":"Boston, MA","unit":"celsius"}'
)
@pytest.mark.parametrize(
"output_item",
[
dict(CUSTOM_TOOL_CALL_ITEM),
CustomToolCallOutputItem(**CUSTOM_TOOL_CALL_ITEM),
ResponseCustomToolCall(**{key: value for key, value in CUSTOM_TOOL_CALL_ITEM.items() if key != "status"}),
],
ids=["dict", "litellm_typed", "openai_typed"],
)
def test_extract_custom_tool_call_input_as_arguments(self, output_item):
handler = OpenAIResponsesHandler()
texts_to_check: List[str] = []
tool_calls_to_check: List[Any] = []
handler._extract_output_text_and_images(
output_item=output_item,
output_idx=2,
texts_to_check=texts_to_check,
images_to_check=[],
task_mappings=[],
tool_calls_to_check=tool_calls_to_check,
)
assert texts_to_check == []
assert tool_calls_to_check == [
{
"id": "call_exec_1",
"type": "function",
"function": {"name": "exec", "arguments": "echo persimmon"},
"index": 2,
}
]
@pytest.mark.asyncio
@pytest.mark.parametrize("typed", [False, True], ids=["dict", "typed"])
async def test_process_output_response_writes_tool_call_rewrites_back(self, typed):
handler = OpenAIResponsesHandler()
function_call = {
"type": "function_call",
"id": "fc_1",
"call_id": "call_fn_1",
"name": "lookup_fruit",
"arguments": '{"fruit": "persimmon"}',
"status": "completed",
}
message = {
"type": "message",
"id": "msg_1",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": "running persimmon", "annotations": []}],
}
payload = {
"id": "resp_1",
"created_at": 1,
"model": "gpt-5.6",
"object": "response",
"status": "completed",
"output": [message, function_call, dict(CUSTOM_TOOL_CALL_ITEM)],
}
response = ResponsesAPIResponse.model_validate(payload) if typed else payload
result = await handler.process_output_response(response, PersimmonMaskingGuardrail(guardrail_name="mask"))
output = result.output if typed else result["output"]
function_item, custom_item = output[1], output[2]
assert (function_item.arguments if typed else function_item["arguments"]) == '{"fruit": "[MASKED]"}'
assert (custom_item.input if typed else custom_item["input"]) == "echo [MASKED]"
assert (custom_item.name if typed else custom_item["name"]) == "exec"
assert (output[0].content[0].text if typed else output[0]["content"][0]["text"]) == "running persimmon"
@pytest.mark.asyncio
async def test_process_output_response_with_tool_calls(self):
"""Test processing output response containing function tool calls"""
@ -1315,6 +1422,109 @@ class TestOpenAIResponsesHandlerStreamingOutputProcessing:
assert completed_event.response.output[0].arguments == '{"fruit": "[MASKED]"}'
assert completed_event.response.output[0].name == "lookup_fruit"
@staticmethod
def _ended_custom_tool_call_stream_events() -> List[dict]:
def item(input_text: str, status: str) -> dict:
return {**CUSTOM_TOOL_CALL_ITEM, "input": input_text, "status": status}
return [
{"type": "response.output_item.added", "output_index": 0, "item": item("", "in_progress")},
{"type": "response.custom_tool_call_input.delta", "item_id": "ctc_1", "output_index": 0, "delta": "echo "},
{"type": "response.custom_tool_call_input.delta", "item_id": "ctc_1", "output_index": 0, "delta": "persimmon"},
{"type": "response.custom_tool_call_input.done", "item_id": "ctc_1", "output_index": 0, "input": "echo persimmon"},
{"type": "response.output_item.done", "output_index": 0, "item": item("echo persimmon", "completed")},
{
"type": "response.completed",
"response": {
"id": "resp_123",
"created_at": 1,
"model": "gpt-5.6",
"output": [item("echo persimmon", "completed")],
"status": "completed",
},
},
]
@pytest.mark.asyncio
async def test_deliver_ended_stream_rewrites_syncs_custom_tool_call_events(self):
handler = OpenAIResponsesHandler()
events = self._ended_custom_tool_call_stream_events()
result = await handler.process_output_streaming_response(
responses_so_far=events,
guardrail_to_apply=PersimmonMaskingGuardrail(guardrail_name="mask"),
litellm_logging_obj=None,
deliver_ended_stream_rewrites=True,
)
assert result is events
assert events[0]["item"]["input"] == ""
assert events[1]["delta"] == "echo [MASKED]"
assert events[2]["delta"] == ""
assert events[3]["input"] == "echo [MASKED]"
assert events[4]["item"]["input"] == "echo [MASKED]"
assert events[5]["response"]["output"][0]["input"] == "echo [MASKED]"
assert events[5]["response"]["output"][0]["name"] == "exec"
assert "arguments" not in events[5]["response"]["output"][0]
@pytest.mark.asyncio
async def test_deliver_ended_stream_rewrites_syncs_typed_custom_tool_call_events(self):
from litellm.types.llms.openai import (
OutputItemAddedEvent,
OutputItemDoneEvent,
ResponseCompletedEvent,
)
handler = OpenAIResponsesHandler()
typed_events: List[Any] = [
model.model_validate({**event, "sequence_number": sequence_number})
for sequence_number, (model, event) in enumerate(
zip(
(
OutputItemAddedEvent,
ResponseCustomToolCallInputDeltaEvent,
ResponseCustomToolCallInputDeltaEvent,
ResponseCustomToolCallInputDoneEvent,
OutputItemDoneEvent,
ResponseCompletedEvent,
),
self._ended_custom_tool_call_stream_events(),
)
)
]
completed_event = typed_events[5]
assert isinstance(completed_event.response.output[0], CustomToolCallOutputItem)
await handler.process_output_streaming_response(
responses_so_far=typed_events,
guardrail_to_apply=PersimmonMaskingGuardrail(guardrail_name="mask"),
litellm_logging_obj=None,
deliver_ended_stream_rewrites=True,
)
assert typed_events[1].delta == "echo [MASKED]"
assert typed_events[2].delta == ""
assert typed_events[3].input == "echo [MASKED]"
assert typed_events[4].item.input == "echo [MASKED]"
assert completed_event.response.output[0].input == "echo [MASKED]"
assert completed_event.response.output[0].name == "exec"
@pytest.mark.asyncio
async def test_deliver_ended_stream_custom_tool_call_rewrite_without_matching_events_fails_closed(self):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
handler = OpenAIResponsesHandler()
events = self._ended_custom_tool_call_stream_events()
events[5]["response"]["output"] = [{**events[5]["response"]["output"][0], "call_id": "call_999"}]
with pytest.raises(UndeliverableStreamRewrite):
await handler.process_output_streaming_response(
responses_so_far=events,
guardrail_to_apply=PersimmonMaskingGuardrail(guardrail_name="mask"),
litellm_logging_obj=None,
deliver_ended_stream_rewrites=True,
)
@staticmethod
def _bridged_function_call_stream_events() -> List[dict]:
reasoning = {"type": "reasoning", "id": "rs_1", "summary": []}
@ -2747,8 +2957,21 @@ class TestOpenAIResponsesHandlerStreamingScanKey:
assert len(ended_key.tool_calls) == 1 and "get_weather" in ended_key.tool_calls[0]
assert ended_key != open_key
def test_completed_event_with_a_custom_tool_call_changes_the_key(self):
handler = OpenAIResponsesHandler()
message = {"type": "message", "content": [{"type": "output_text", "text": "hi"}]}
ended_key = handler.get_streaming_scan_key(
[self._delta(0, "hi"), self._completed(1, [message, dict(CUSTOM_TOOL_CALL_ITEM)])]
)
rewritten_key = handler.get_streaming_scan_key(
[self._delta(0, "hi"), self._completed(1, [message, {**CUSTOM_TOOL_CALL_ITEM, "input": "echo kumquat"}])]
)
assert ended_key.texts == ("hi",)
assert len(ended_key.tool_calls) == 1 and "echo persimmon" in ended_key.tool_calls[0]
assert rewritten_key != ended_key
def test_completed_event_reads_every_output_text_part(self):
from litellm.types.responses.main import GenericResponseOutputItem, OutputText
from litellm.types.responses.main import CustomToolCallOutputItem, GenericResponseOutputItem, OutputText
item = GenericResponseOutputItem(
type="message",