From e80c20aab744c5b64ee6ba1f394f3377eb456320 Mon Sep 17 00:00:00 2001 From: albertbausili Date: Mon, 21 Sep 2026 14:48:45 +0200 Subject: [PATCH] fix(guardrails): hold streamed tool calls until inspection succeeds --- .../guardrail_hooks/neuraltrust/README.md | 4 +- .../unified_guardrail/unified_guardrail.py | 46 +++++++------------ .../guardrail_hooks/test_neuraltrust.py | 15 +++--- .../test_unified_guardrail.py | 10 +++- 4 files changed, 34 insertions(+), 41 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md index 3ca1aa7900e..620b01bcb55 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md +++ b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md @@ -62,7 +62,9 @@ Under `incremental_diff` the reply is held until the end-of-stream evaluate retu `incremental_diff` covers OpenAI chat completions streaming; other surfaces fall back to `block_only`. -Streamed tool calls are the exception in either mode: LiteLLM forwards the tool-call deltas as they arrive and only sends the assembled call to TrustGuard once the stream ends, so a blocked call can already have reached the client. `incremental_diff` narrows that window, holding back the answer text and the turn's `finish_reason` so the block lands as an error instead of trailing a stream that looks complete. Use non-streaming requests where a tool call must be vetted before the client ever sees it. +Under `incremental_diff`, LiteLLM buffers tool-call deltas until TrustGuard inspects the assembled response. A blocking verdict releases no tool-call arguments or finish signal. Allowed tool calls retain their original deltas and order + +The default `block_only` mode still forwards tool-call deltas before inspection. Use `incremental_diff` or non-streaming requests when tool calls must be checked before delivery. Streamed tool-call rewrites are not supported; use non-streaming requests for transformed tool arguments ## References diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index d68a55f9a88..e5cb8292f3f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -708,40 +708,10 @@ class UnifiedLLMGuardrails(CustomLogger): try: async for item in response: - # v1 transforms only text. A chunk carrying tool_calls is passed - # through raw so function-calling turns are not dropped, but ONLY - # its tool-call fields are forwarded: content is stripped so any - # response text (in the same delta, or in another choice of an n>1 - # chunk) can never bypass the transform. The original chunk is kept - # in responses_so_far so its text is still accumulated + redacted + - # emitted as synthetic deltas, and so the guardrail inspects the - # assembled tool calls at end of stream (see the block inspection - # below), matching block_only. finish_reason rides on the raw - # tool-only chunk, so it is not recorded for the text flush. if self._chunk_has_tool_calls(item): saw_tool_calls = True responses_so_far.append(item) last_chunk = item - # Fix #3 — flush accumulated text BEFORE the tool-call - # passthrough. Without this, a stream of text chunks that - # hasn't yet hit a sampled round can be trailed by a - # tool-call chunk carrying finish_reason="tool_calls"; an - # SSE-compliant client stops reading at that finish_reason - # and drops the end-of-stream text flush that would follow. - if saw_text_content: - async for out in _round(item, is_final=False): - yield out - # Fix #1 — pass finish_reason_per_choice into the - # passthrough so a mixed content+tool_call chunk defers its - # finish_reason to the final text terminator (see the - # _tool_call_passthrough_chunk docstring). - tool_only = self._tool_call_passthrough_chunk( - item, - finish_reason_per_choice=finish_reason_per_choice, - held_choices=_held_choices(held_chars_per_choice), - ) - responses_yielded.append(tool_only) - yield tool_only continue if self._is_trailing_metadata_chunk(item): @@ -759,6 +729,7 @@ class UnifiedLLMGuardrails(CustomLogger): # sampled round here would guardrail the same content twice. if ( not end_of_stream_only + and not saw_tool_calls and not self._chunk_has_finish_reason(item) and chunk_counter % sampling_rate == 0 ): @@ -791,6 +762,21 @@ class UnifiedLLMGuardrails(CustomLogger): ): yield out + if saw_text_content: + async for out in _round(last_chunk, is_final=False): + yield out + for tool_only in ( + self._tool_call_passthrough_chunk( + buffered_item, + finish_reason_per_choice=finish_reason_per_choice, + held_choices=_held_choices(held_chars_per_choice), + ) + for buffered_item in responses_so_far + if self._chunk_has_tool_calls(buffered_item) + ): + responses_yielded.append(tool_only) + yield tool_only + async for out in self._emit_stream_tail( last_chunk=last_chunk, final_round=_round, diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_neuraltrust.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_neuraltrust.py index d34ead6312d..7a2c6fdfc35 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_neuraltrust.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_neuraltrust.py @@ -110,8 +110,8 @@ async def _upstream_reply() -> AsyncIterator[ModelResponseStream]: FORBIDDEN_TOOL = "wire_transfer" -async def _upstream_tool_call() -> AsyncIterator[ModelResponseStream]: - for chunk in REPLY_CHUNKS[:4]: +async def _upstream_tool_call(include_text: bool = True) -> AsyncIterator[ModelResponseStream]: + for chunk in REPLY_CHUNKS[:4] if include_text else (): yield _stream_chunk(chunk) yield ModelResponseStream( model="gpt-4o-mini", @@ -1157,21 +1157,18 @@ class TestNeuralTrustGuardrail: assert _deltas(received) == [] @pytest.mark.asyncio - async def test_tool_call_block_under_incremental_diff_leaves_the_turn_unfinished(self) -> None: - """A streamed tool call is only scanned once the stream ends, so the turn must not look complete. - - Until then the answer text stays withheld and no finish_reason goes out, so a client cannot treat - the turn as done, and the block surfaces as a 400 rather than trailing a finished-looking stream. - """ + @pytest.mark.parametrize("include_text", [True, False]) + async def test_tool_call_block_under_incremental_diff_sends_nothing(self, include_text: bool) -> None: guardrail = _guardrail(event_hook="post_call", default_on=True, streaming_transform_mode="incremental_diff") received: list[object] = [] # mutable-ok: collects what the client saw before the block with patch.object(guardrail.async_handler, "post", _tool_call_blocking_trustguard()): with pytest.raises(HTTPException) as exc_info: - await _drain_into(_guardrail_stream(guardrail, _upstream_tool_call()), received) + await _drain_into(_guardrail_stream(guardrail, _upstream_tool_call(include_text)), received) assert exc_info.value.status_code == 400 assert exc_info.value.detail["verdict"] == "block" assert _deltas(received) == [] assert _finish_reasons(received) == [] + assert received == [] @pytest.mark.asyncio async def test_default_streaming_mode_leaves_the_transform_off_the_wire(self) -> None: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index d1d22d0d7c2..ca507db791d 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -1410,8 +1410,16 @@ class TestStreamingTransform: ], ) + async def upstream(): + yield tool_chunk + + stream: Final = UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/chat/completions"), + response=upstream(), + request_data={"guardrail_to_apply": _ToolCallBlocker(), "model": "gpt-4"}, + ) with pytest.raises(GuardrailRaisedException): - await _drive_stream(UnifiedLLMGuardrails(), _ToolCallBlocker(), [tool_chunk]) + await anext(stream) @pytest.mark.asyncio async def test_mixed_content_and_tool_call_chunk_does_not_leak_text(self):