From 80782d478ddbdd2e145b0cb72eb3c653efbe0ed0 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 8 Sep 2026 15:17:36 -0700 Subject: [PATCH] fix(policy_engine): keep legacy stream steps in sync and reject tool-call rewrites Each streaming step drops the response an earlier step's translation stored under request_data["response"], so a later legacy hook sees the stream as the steps before it left it instead of the first step's snapshot. A legacy replacement whose tool calls differ from the scanned chunks is now undeliverable like a text mismatch, so the original stream is released with a warning instead of delivering the text while dropping the tool-call change --- .../proxy/policy_engine/pipeline_executor.py | 17 ++-- .../policy_engine/test_pipeline_executor.py | 78 ++++++++++++++++--- 2 files changed, 81 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/policy_engine/pipeline_executor.py b/litellm/proxy/policy_engine/pipeline_executor.py index 089f0683823..b4f6da18d1c 100644 --- a/litellm/proxy/policy_engine/pipeline_executor.py +++ b/litellm/proxy/policy_engine/pipeline_executor.py @@ -150,8 +150,8 @@ class _LegacyHookStreamAdapter(CustomGuardrail): non-streaming hooks, an exception it raises ends the stream through the executor's fail/error classification, and a replacement response is re-scanned by the same translation so its texts reach the client through the translation's ended-stream write-back. A - replacement whose scanned texts do not line up with the originals is undeliverable, so the - executor releases the original chunks.""" + replacement whose scanned texts do not line up with the originals, or whose tool calls + differ from them, is undeliverable, so the executor releases the original chunks.""" def __init__( self, @@ -184,8 +184,12 @@ class _LegacyHookStreamAdapter(CustomGuardrail): return inputs scanned: Final = _text_snapshot(inputs.get("texts")) rescanned: Final = await self._rescan(replacement, logging_obj) - rewritten: Final = None if rescanned is None else rescanned.get("texts") - if scanned is None or rewritten is None or len(rewritten) != len(scanned): + if scanned is None or rescanned is None: + raise UndeliverableStreamRewrite(self.guardrail_name or "unknown") + rewritten: Final = rescanned.get("texts") + if rewritten is None or len(rewritten) != len(scanned): + raise UndeliverableStreamRewrite(self.guardrail_name or "unknown") + if _tool_call_shapes(rescanned.get("tool_calls")) != _tool_call_shapes(inputs.get("tool_calls")): raise UndeliverableStreamRewrite(self.guardrail_name or "unknown") rewritten_inputs: Final[GenericGuardrailAPIInputs] = {**inputs, "texts": rewritten} return rewritten_inputs @@ -380,7 +384,9 @@ class PipelineExecutor: yet (a tool-call rewrite, a text rewrite on a translation without write-back, or one the translation or adapter refused with ``UndeliverableStreamRewrite``) is discarded: the buffered chunks go back to the originals and the step passes, so the client gets - the stream the merge base sent.""" + the stream the merge base sent. The response an earlier step's translation stored under + ``request_data["response"]`` is dropped first, so this step's hook sees the stream as + the steps before it left it.""" scanner: Final = ( callback if PipelineExecutor.supports_unified_execution(callback) @@ -389,6 +395,7 @@ class PipelineExecutor: observer: Final = _StreamRewriteObserver(scanner) deliver_rewrites: Final = type(endpoint_translation).delivers_ended_stream_text_rewrites originals: Final = copy.deepcopy(streaming_chunks) + hook_input.pop("response", None) # rebind-ok: an earlier step's stored response goes so this step's is stored try: if deliver_rewrites: await endpoint_translation.process_output_streaming_response( 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 a44afc6ef96..860f83a06bb 100644 --- a/tests/test_litellm/proxy/policy_engine/test_pipeline_executor.py +++ b/tests/test_litellm/proxy/policy_engine/test_pipeline_executor.py @@ -675,7 +675,6 @@ async def test_guardrail_not_found_uses_on_fail(monkeypatch): ], ) - result = await PipelineExecutor.execute_steps( steps=pipeline.steps, mode=pipeline.mode, @@ -1298,8 +1297,8 @@ async def test_streaming_step_restores_chunks_when_translation_refuses_the_rewri class _LegacyHookGuardrail(CustomGuardrail): """A guardrail with only the legacy post-call hook: it never defines apply_guardrail.""" - def __init__(self, replacement=None, raises=None): - super().__init__(guardrail_name="masker", event_hook="post_call", default_on=True) + def __init__(self, replacement=None, raises=None, guardrail_name="masker"): + super().__init__(guardrail_name=guardrail_name, event_hook="post_call", default_on=True) self.replacement = replacement self.raises = raises self.calls = [] @@ -1339,7 +1338,7 @@ class _LegacyScanningTranslation: ): request_data.setdefault("response", {"text": responses_so_far[0]["text"]}) outputs = await guardrail_to_apply.apply_guardrail( - inputs={"texts": [responses_so_far[0]["text"]]}, + inputs={"texts": [responses_so_far[0]["text"]], "tool_calls": [dict(responses_so_far[0]["tool_call"])]}, request_data=request_data, input_type="response", logging_obj=litellm_logging_obj, @@ -1350,8 +1349,11 @@ class _LegacyScanningTranslation: async def process_output_response( self, response, guardrail_to_apply, litellm_logging_obj=None, user_api_key_dict=None, request_data=None ): + inputs = {"texts": list(response["texts"])} + if response.get("tool_calls"): + inputs["tool_calls"] = list(response["tool_calls"]) await guardrail_to_apply.apply_guardrail( - inputs={"texts": list(response["texts"])}, + inputs=inputs, request_data={"response": response}, input_type="response", logging_obj=litellm_logging_obj, @@ -1359,10 +1361,26 @@ class _LegacyScanningTranslation: return response +def _legacy_replacement(*texts, tool_calls=None): + return {"texts": list(texts), "tool_calls": [_chunk()["tool_call"]] if tool_calls is None else tool_calls} + + async def _run_legacy_streaming_step(monkeypatch, guardrail, chunks, on_fail="block", on_error="next"): - monkeypatch.setattr(litellm, "callbacks", [guardrail]) + return await _run_legacy_streaming_steps(monkeypatch, [guardrail], chunks, on_fail=on_fail, on_error=on_error) + + +async def _run_legacy_streaming_steps(monkeypatch, guardrails, chunks, on_fail="block", on_error="next"): + monkeypatch.setattr(litellm, "callbacks", list(guardrails)) return await PipelineExecutor.execute_steps( - steps=[PipelineStep(guardrail="masker", on_pass="allow", on_fail=on_fail, on_error=on_error)], + steps=[ + PipelineStep( + guardrail=guardrail.guardrail_name, + on_pass="next" if position + 1 < len(guardrails) else "allow", + on_fail=on_fail, + on_error=on_error, + ) + for position, guardrail in enumerate(guardrails) + ], mode="post_call", data={"model": "m"}, user_api_key_dict=MagicMock(), @@ -1376,7 +1394,7 @@ async def _run_legacy_streaming_step(monkeypatch, guardrail, chunks, on_fail="bl @pytest.mark.asyncio @pytest.mark.parametrize("guardrail_class", [_LegacyHookGuardrail, _NativeHooksGuardrail]) async def test_streaming_step_runs_legacy_hook_and_delivers_its_rewrite(monkeypatch, caplog, guardrail_class): - guardrail = guardrail_class(replacement={"texts": ["[REWRITTEN] hello world"]}) + guardrail = guardrail_class(replacement=_legacy_replacement("[REWRITTEN] hello world")) chunks = [_chunk()] with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): @@ -1434,7 +1452,7 @@ async def test_streaming_step_takes_on_error_when_legacy_hook_crashes(monkeypatc @pytest.mark.asyncio async def test_streaming_step_discards_legacy_rewrite_whose_texts_do_not_line_up(monkeypatch, caplog): - guardrail = _LegacyHookGuardrail(replacement={"texts": ["split", "in two"]}) + guardrail = _LegacyHookGuardrail(replacement=_legacy_replacement("split", "in two")) chunks = [_chunk()] with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): @@ -1442,3 +1460,45 @@ async def test_streaming_step_discards_legacy_rewrite_whose_texts_do_not_line_up _assert_passed_with_discard_warning(result, caplog) assert chunks == [_chunk()] + + +@pytest.mark.asyncio +async def test_streaming_step_discards_legacy_rewrite_that_changes_a_tool_call(monkeypatch, caplog): + masked_tool_call = {"function": {"name": "lookup", "arguments": '{"ssn": "[MASKED]"}'}} + guardrail = _LegacyHookGuardrail( + replacement=_legacy_replacement("[REWRITTEN] hello world", tool_calls=[masked_tool_call]) + ) + chunks = [_chunk()] + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + result = await _run_legacy_streaming_step(monkeypatch, guardrail, chunks) + + _assert_passed_with_discard_warning(result, caplog) + assert chunks == [_chunk()] + + +@pytest.mark.asyncio +async def test_streaming_step_discards_legacy_rewrite_that_drops_the_tool_calls(monkeypatch, caplog): + guardrail = _LegacyHookGuardrail(replacement=_legacy_replacement("[REWRITTEN] hello world", tool_calls=[])) + chunks = [_chunk()] + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + result = await _run_legacy_streaming_step(monkeypatch, guardrail, chunks) + + _assert_passed_with_discard_warning(result, caplog) + assert chunks == [_chunk()] + + +@pytest.mark.asyncio +async def test_later_legacy_step_sees_the_stream_as_the_earlier_step_left_it(monkeypatch): + masker = _LegacyHookGuardrail(replacement=_legacy_replacement("[REWRITTEN] hello world")) + auditor = _LegacyHookGuardrail(replacement=None, guardrail_name="auditor") + chunks = [_chunk()] + + result = await _run_legacy_streaming_steps(monkeypatch, [masker, auditor], chunks, on_fail="next") + + assert result.terminal_action == "allow" + assert [step.outcome for step in result.step_results] == ["pass", "pass"] + assert chunks[0]["text"] == "[REWRITTEN] hello world" + assert [call["response"] for call in masker.calls] == [{"native": True, "text": "hello world"}] + assert [call["response"] for call in auditor.calls] == [{"native": True, "text": "[REWRITTEN] hello world"}]