From 2a35dc5217f7fc6697337985702a1401344519dc Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 18:00:03 -0700 Subject: [PATCH] fix(guardrails): keep the undeliverable rewrite reason through copies and name the responses mismatch --- .../guardrail_translation/handler.py | 50 ++++++++++++++----- .../proxy/policy_engine/pipeline_executor.py | 11 ++-- ...test_openai_responses_guardrail_handler.py | 30 +++++++++-- .../policy_engine/test_pipeline_executor.py | 13 +++++ 4 files changed, 83 insertions(+), 21 deletions(-) diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index ddfd2199822..1ef1011591e 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -167,6 +167,34 @@ def _tool_call_rewrite(before: _ToolCallShape, after: _ToolCallShape) -> _ToolCa return _ToolCallShape(name=after.name if after.name != before.name else None, arguments=after.arguments) +def _undeliverable_tool_call_rewrite_reason( + call_ids: Sequence[str], + tool_call_item_count: int, + post_guardrail_tool_call_count: int, + unresolved_argument_event: bool, + rewritten_call_ids: frozenset[str], + event_call_ids: frozenset[str], +) -> str | None: + if len(call_ids) != tool_call_item_count: + return ( + f"{tool_call_item_count - len(call_ids)} of the stream's {tool_call_item_count} tool call items " + "carry no call_id" + ) + if len(frozenset(call_ids)) != len(call_ids): + return "the stream's tool call items repeat a call_id" + if len(call_ids) != post_guardrail_tool_call_count: + return ( + f"the guardrail returned {post_guardrail_tool_call_count} tool calls for the stream's " + f"{len(call_ids)} tool call items" + ) + if unresolved_argument_event: + return "a tool call argument event names an item_id that no output_item event introduced" + missing_call_ids: Final = sorted(rewritten_call_ids - event_call_ids) + if missing_call_ids: + return f"no stream event carries the rewritten call_id {', '.join(missing_call_ids)}" + return None + + class ResponseOutputEnvelope(TypedDict, total=False): """Dict form of a Responses API response, as far as guardrail write-back reads it.""" @@ -1110,20 +1138,18 @@ class OpenAIResponsesHandler(BaseTranslation): 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(tool_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) - ): + undeliverable_reason: Final = _undeliverable_tool_call_rewrite_reason( + call_ids=call_ids, + tool_call_item_count=len(tool_call_items), + post_guardrail_tool_call_count=len(post_guardrail_tool_calls), + unresolved_argument_event=unresolved_argument_event, + rewritten_call_ids=frozenset(rewrites_by_call_id), + event_call_ids=frozenset(call_id for call_id in event_call_ids if call_id is not None), + ) + if undeliverable_reason is not None: from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite - raise UndeliverableStreamRewrite( - guardrail_name, - f"the guardrail returned {len(post_guardrail_tool_calls)} tool calls and the stream's " - f"{len(tool_call_items)} function_call items could not be lined up with them by call_id", - ) + raise UndeliverableStreamRewrite(guardrail_name, undeliverable_reason) for output_item, rewrite in ( (output_item, rewrites_by_call_id[call_id]) for output_item, call_id in zip(tool_call_items, call_ids) diff --git a/litellm/proxy/policy_engine/pipeline_executor.py b/litellm/proxy/policy_engine/pipeline_executor.py index cfb0dc4b617..0b81e7af84d 100644 --- a/litellm/proxy/policy_engine/pipeline_executor.py +++ b/litellm/proxy/policy_engine/pipeline_executor.py @@ -51,13 +51,16 @@ except ImportError: class UndeliverableStreamRewrite(Exception): def __init__(self, guardrail_name: str, reason: str) -> None: - super().__init__( - f"Guardrail '{guardrail_name}' rewrote the streamed response but the rewrite cannot be written " - f"back to the stream: {reason}" - ) + super().__init__(guardrail_name, reason) self.guardrail_name: Final = guardrail_name self.reason: Final = reason + def __str__(self) -> str: + return ( + f"Guardrail '{self.guardrail_name}' rewrote the streamed response but the rewrite cannot be written " + f"back to the stream: {self.reason}" + ) + def _tool_call_shape(tool_call: object) -> tuple[object, object]: plain: Final = tool_call.model_dump() if isinstance(tool_call, BaseModel) else tool_call diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py b/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py index 81adb283dcc..b45cd2ec299 100644 --- a/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py +++ b/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py @@ -1610,13 +1610,14 @@ class TestOpenAIResponsesHandlerStreamingOutputProcessing: 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): + with pytest.raises(UndeliverableStreamRewrite) as undeliverable: 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 undeliverable.value.reason == "no stream event carries the rewritten call_id call_999" @staticmethod def _bridged_function_call_stream_events() -> List[dict]: @@ -1688,8 +1689,21 @@ class TestOpenAIResponsesHandlerStreamingOutputProcessing: 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): + @pytest.mark.parametrize( + ("mismatch", "expected_reason"), + [ + ("orphan_call_id", "no stream event carries the rewritten call_id call_999"), + ("duplicate_call_id", "the stream's tool call items repeat a call_id"), + ("missing_call_id", "1 of the stream's 1 tool call items carry no call_id"), + ( + "unknown_argument_item_id", + "a tool call argument event names an item_id that no output_item event introduced", + ), + ], + ) + async def test_deliver_ended_stream_function_call_rewrite_without_matching_events_fails_closed( + self, mismatch, expected_reason + ): from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite handler = OpenAIResponsesHandler() @@ -1697,16 +1711,22 @@ class TestOpenAIResponsesHandlerStreamingOutputProcessing: envelope_item = events[5]["response"]["output"][0] if mismatch == "orphan_call_id": events[5]["response"]["output"] = [{**envelope_item, "call_id": "call_999"}] - else: + elif mismatch == "duplicate_call_id": events[5]["response"]["output"] = [dict(envelope_item), dict(envelope_item)] + elif mismatch == "missing_call_id": + events[5]["response"]["output"] = [{key: value for key, value in envelope_item.items() if key != "call_id"}] + else: + events[1]["item_id"] = "fc_unknown" - with pytest.raises(UndeliverableStreamRewrite): + with pytest.raises(UndeliverableStreamRewrite) as undeliverable: 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 undeliverable.value.reason == expected_reason + assert str(undeliverable.value).endswith(f"cannot be written back to the stream: {expected_reason}") @pytest.mark.asyncio async def test_ended_stream_function_call_rewrite_leaves_events_untouched_by_default(self): diff --git a/tests/test_litellm/proxy/policy_engine/test_pipeline_executor.py b/tests/test_litellm/proxy/policy_engine/test_pipeline_executor.py index 6af5e9573f2..6aa7eca0f15 100644 --- a/tests/test_litellm/proxy/policy_engine/test_pipeline_executor.py +++ b/tests/test_litellm/proxy/policy_engine/test_pipeline_executor.py @@ -6,6 +6,7 @@ Uses mock guardrails to validate pipeline execution without external services. import copy import logging +import pickle from typing import Literal from unittest.mock import MagicMock @@ -1727,3 +1728,15 @@ async def test_later_legacy_step_sees_the_stream_as_the_earlier_step_left_it(mon assert chunks[0]["text"] == "[REWRITTEN] hello world" assert [call["response"] for call in masker.calls] == [_native("hello world")] assert [call["response"] for call in auditor.calls] == [_native("[REWRITTEN] hello world")] + + +@pytest.mark.parametrize("clone", [copy.deepcopy, lambda exc: pickle.loads(pickle.dumps(exc))], ids=["deepcopy", "pickle"]) +def test_undeliverable_stream_rewrite_keeps_its_reason_through_a_copy(clone): + original = UndeliverableStreamRewrite("masker", "the translation refused it") + + copied = clone(original) + + assert copied.guardrail_name == "masker" + assert copied.reason == "the translation refused it" + assert str(copied) == str(original) + assert str(copied).endswith("cannot be written back to the stream: the translation refused it")