From 5c437f7efbadfdf5b1b86fd6b0b70c9f8fcc1a6a Mon Sep 17 00:00:00 2001 From: albertbausili Date: Mon, 21 Sep 2026 18:06:10 +0200 Subject: [PATCH] fix(guardrails): apply buffered tool argument rewrites --- .../chat/guardrail_translation/handler.py | 32 +---- .../guardrail_hooks/neuraltrust/README.md | 6 +- .../unified_guardrail/unified_guardrail.py | 37 ++---- .../guardrails/guardrail_hooks/neuraltrust.py | 2 +- .../test_unified_guardrail.py | 111 +++++++++++++++++- .../test_openai_guardrail_handler.py | 33 +++--- 6 files changed, 143 insertions(+), 78 deletions(-) diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index 2b895049743..30d617e3103 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -556,7 +556,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): terminate the stream. Text rewrites are not propagated to the client here (see ``_process_streaming_transform`` for the incremental_diff path) unless ``deliver_ended_stream_rewrites`` opts the ended-stream branch in.""" - has_stream_ended: Final = self._first_choice_has_finished(responses_so_far) + has_stream_ended: Final = deliver_ended_stream_rewrites or self._first_choice_has_finished(responses_so_far) if has_stream_ended: await self._process_ended_stream( @@ -1132,20 +1132,18 @@ class OpenAIChatCompletionsHandler(BaseTranslation): def _function_tool_call_fragments( responses_so_far: Sequence["ModelResponseStream"], ) -> tuple[tuple[ChatCompletionDeltaToolCall, ...], ...]: - """Group the stream's function tool-call fragments by their tool-call index, in - the index order ``stream_chunk_builder`` lists the rebuilt tool calls, keeping - only the indices the builder keeps (an id and a name somewhere in the stream).""" fragments: Final = tuple( - tool_call + (choice.index, tool_call) for response in responses_so_far for choice in response.choices for tool_call in choice.delta.tool_calls or () if isinstance(tool_call, ChatCompletionDeltaToolCall) ) - identified: Final = frozenset(fragment.index for fragment in fragments if fragment.id) - named: Final = frozenset(fragment.index for fragment in fragments if fragment.function.name) + identified: Final = frozenset((choice, fragment.index) for choice, fragment in fragments if fragment.id) + named: Final = frozenset((choice, fragment.index) for choice, fragment in fragments if fragment.function.name) return tuple( - tuple(fragment for fragment in fragments if fragment.index == index) for index in sorted(identified & named) + tuple(fragment for choice, fragment in fragments if (choice, fragment.index) == key) + for key in sorted(identified & named) ) def _write_ended_stream_tool_call_rewrites( @@ -1155,28 +1153,10 @@ class OpenAIChatCompletionsHandler(BaseTranslation): pre_guardrail_tool_calls: tuple[tuple[str | None, str], ...], guardrail_name: str, ) -> None: - """Write ended-stream guardrail tool-call rewrites back across the buffered - chunks: the rewritten name and full arguments land in the tool call's first - fragment and the arguments of its later fragments are blanked, mirroring the - text write-back. A rewrite on a stream carrying more than one distinct choice - index, or whose fragments do not line up with the rebuilt tool calls, is - reported as undeliverable, so the pipeline executor discards it and releases - the original chunks.""" post_guardrail_tool_calls: Final = self._function_tool_call_shapes(guardrailed_response) if post_guardrail_tool_calls == pre_guardrail_tool_calls: return - stream_choice_indices: Final = frozenset( - choice.index for response in responses_so_far for choice in response.choices - ) fragments_by_tool_call: Final = self._function_tool_call_fragments(responses_so_far) - if len(stream_choice_indices) != 1: - from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite - - raise UndeliverableStreamRewrite( - guardrail_name, - f"the stream carries {len(stream_choice_indices)} choices and tool-call rewrites are only written " - "back on single-choice streams", - ) if len(fragments_by_tool_call) != len(post_guardrail_tool_calls): from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite diff --git a/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md index 620b01bcb55..2f3136967de 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md +++ b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md @@ -18,7 +18,7 @@ guardrails: collector_key: os.environ/TRUSTGUARD_COLLECTOR_KEY # tgcol_… ; optional if the API key is bound unreachable_fallback: fail_closed timeout: 5 - streaming_transform_mode: block_only # incremental_diff to stream redacted output + streaming_transform_mode: incremental_diff default_on: true ``` @@ -62,9 +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`. -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 +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. Tool calls retain their IDs and order, with transformed arguments written into the buffered deltas before delivery -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 +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. The example selects `incremental_diff` for inspected text and 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 6be65c2b686..4cc9ff37c73 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -736,28 +736,14 @@ class UnifiedLLMGuardrails(CustomLogger): async for out in _round(item, is_final=False): yield out - # v1 does not transform streamed tool calls, but they must still go - # through the guardrail's block decision. Run the block_only inspection - # over the full assembled response so tool calls cannot bypass it. - # - # Pass a deep copy of responses_so_far — the block path routes through - # ``_process_streaming_block_only`` which mutates ``delta.content`` - # in-place on the chunk objects it receives. For an n>1 chunk carrying - # text on one choice and tool_calls (with finish_reason) on another, - # ``has_stream_ended`` reads ``choices[0]`` alone and can miss the - # terminal signal, letting the block path rewrite the raw accumulator. - # The subsequent final ``_round`` would then re-read the already-mutated - # text, producing double-application for a non-idempotent guardrail or a - # ``stream_transform_underflow`` 400 from mismatched prefixes. A shallow - # list copy wouldn't help — the mutation is on the chunk objects - # themselves — so we deepcopy. if saw_tool_calls: - async for out in self._inspect_full_response_for_block( + inspected_responses: Final = copy.deepcopy(responses_so_far) + async for out in self._inspect_full_response( endpoint_translation=endpoint_translation, guardrail_to_apply=guardrail_to_apply, request_data=request_data, user_api_key_dict=user_api_key_dict, - responses_so_far=copy.deepcopy(responses_so_far), + responses_so_far=inspected_responses, responses_yielded=responses_yielded, ): yield out @@ -771,7 +757,7 @@ class UnifiedLLMGuardrails(CustomLogger): finish_reason_per_choice=finish_reason_per_choice, held_choices=_held_choices(held_chars_per_choice), ) - for buffered_item in responses_so_far + for buffered_item in inspected_responses if self._chunk_has_tool_calls(buffered_item) ) @@ -826,7 +812,7 @@ class UnifiedLLMGuardrails(CustomLogger): responses_yielded.append(trailing) yield trailing - async def _inspect_full_response_for_block( + async def _inspect_full_response( self, *, endpoint_translation: _EndpointTranslation, @@ -836,16 +822,8 @@ class UnifiedLLMGuardrails(CustomLogger): responses_so_far: Sequence[object], responses_yielded: Sequence[object], ) -> AsyncGenerator[object, None]: - """Run the block-only guardrail inspection over the full assembled - response (text + tool calls) so nothing bypasses the block decision. - - The guardrail's returned transforms are discarded here (v1 does not - transform tool calls); only its block decision matters. A block is - surfaced the same way as elsewhere: ModifyResponseException terminates the - stream via the shared block handler; a GenericGuardrailAPI block raises and - propagates, matching block_only. - """ from litellm.integrations.custom_guardrail import ModifyResponseException + from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite try: await endpoint_translation.process_output_streaming_response( @@ -855,7 +833,10 @@ class UnifiedLLMGuardrails(CustomLogger): user_api_key_dict=user_api_key_dict, request_data=request_data, stream_transform_sink=None, + deliver_ended_stream_rewrites=True, ) + except UndeliverableStreamRewrite as exc: + raise HTTPException(status_code=400, detail="Guardrail stream rewrite could not be applied") from exc except ModifyResponseException as e: if e.original_response is None: e.original_response = responses_so_far diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/neuraltrust.py b/litellm/types/proxy/guardrails/guardrail_hooks/neuraltrust.py index ae337367eb5..73f7db0299a 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/neuraltrust.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/neuraltrust.py @@ -54,7 +54,7 @@ class NeuralTrustGuardrailConfigModel(GuardrailConfigModel): description=( "How a `transform` verdict reaches a streaming client. `block_only` (default) streams the raw " "model chunks, so `block` and `ask` still end the stream but the redacted text is dropped. " - "`incremental_diff` withholds the model chunks and streams TrustGuard's rewritten text instead: " + "`incremental_diff` withholds the model chunks and streams TrustGuard's rewritten text and tool arguments: " "the reply arrives once the end-of-stream evaluate returns, and a blocking verdict ends the " "stream with nothing already sent. OpenAI chat completions streaming only." ), 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 c953dea1fb4..59e2b5d7d17 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 @@ -2,7 +2,10 @@ import logging from types import SimpleNamespace -from typing import Final +from typing import Final, Literal +from collections.abc import AsyncIterator + +from pydantic import TypeAdapter import pytest @@ -41,7 +44,10 @@ from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrai ) from litellm.types.guardrails import GuardrailEventHooks from litellm.types.llms.openai import ResponsesAPIResponse -from litellm.types.utils import CallTypes, Delta, ModelResponseStream, StreamingChoices +from litellm.types.utils import ( + CallTypes, ChatCompletionMessageToolCall, Delta, GenericGuardrailAPIInputs, + ModelResponseStream, StreamingChoices, +) class RecordingGuardrail(CustomGuardrail): @@ -947,6 +953,36 @@ def _delta_text(item): return item.choices[0].delta.content or "" +class _ToolRedactingGuardrail(CustomGuardrail): + def __init__(self) -> None: + super().__init__(guardrail_name="tool-redactor", event_hook=GuardrailEventHooks.post_call, default_on=True) + self.streaming_transform_mode = "incremental_diff" + self.streaming_sampling_rate = 1 + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + logging_obj: object | None = None, + ) -> GenericGuardrailAPIInputs: + calls: Final = TypeAdapter(tuple[ChatCompletionMessageToolCall, ...]).validate_python( + inputs.get("tool_calls", ()) + ) + texts: Final = tuple("checked:" + text.replace("SECRET", "MASKED") for text in inputs.get("texts", ())) + return { + **inputs, + "texts": list(texts), + "stream_holdback_chars": [len(text) for text in texts], + "tool_calls": [ + call.model_copy(update={"function": call.function.model_copy(update={ + "arguments": call.function.arguments.replace("SECRET", "MASKED"), + })}) + for call in calls + ], + } + + class TestStreamingTransform: """Streaming text-transformation (incremental_diff) path on the OpenAI chat completions streaming surface.""" @@ -955,6 +991,77 @@ class TestStreamingTransform: def _use_openai_handler_mapping(self, monkeypatch): _patch_translation_mappings(monkeypatch, {CallTypes.acompletion: OpenAIChatCompletionsHandler}) + @pytest.mark.asyncio + @pytest.mark.parametrize("include_text", [False, True]) + @pytest.mark.parametrize("tool_count", [1, 2]) + async def test_buffered_tool_arguments_are_rewritten_before_delivery( + self, include_text: bool, tool_count: int + ) -> None: + chunks: Final = ( + *([_stream_chunk("hello SECRET")] if include_text else []), + ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[ + {"index": index, "id": f"call_{index}", "type": "function", + "function": {"name": "contact", "arguments": '{"contact":"SEC'}} + for index in range(tool_count) + ]))]), + ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[ + {"index": index, "function": {"arguments": 'RET"}'}} for index in range(tool_count) + ]), finish_reason="tool_calls")]), + ModelResponseStream(choices=[], usage={"prompt_tokens": 7, "completion_tokens": 11, "total_tokens": 18}), + ) + out: Final = await _drive_stream(UnifiedLLMGuardrails(), _ToolRedactingGuardrail(), chunks) + calls: Final = tuple( + call for chunk in out for choice in chunk.choices for call in choice.delta.tool_calls or () + ) + for index in range(tool_count): + arguments: Final = "".join(call.function.arguments or "" for call in calls if call.index == index) + assert arguments == '{"contact":"MASKED"}' + assert next(call.id for call in calls if call.index == index and call.id) == f"call_{index}" + assert "".join(_delta_text(chunk) for chunk in out) == ("checked:hello MASKED" if include_text else "") + assert any(choice.finish_reason == "tool_calls" for chunk in out for choice in chunk.choices) + assert out[-1].usage.total_tokens == 18 + assert all("SECRET" not in chunk.model_dump_json() for chunk in out) + + @pytest.mark.asyncio + async def test_tool_rewrites_keep_completion_choices_separate(self) -> None: + async def response() -> AsyncIterator[ModelResponseStream]: + yield ModelResponseStream(choices=[StreamingChoices( + index=index, + delta=Delta(tool_calls=[{"index": 0, "id": f"call_{index}", "type": "function", + "function": {"name": "contact", "arguments": '{"contact":"SECRET"}'}}]), + finish_reason="tool_calls", + ) for index in range(2)]) + + iterator: Final = UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(request_route="/v1/chat/completions"), + response=response(), request_data={"guardrail_to_apply": _ToolRedactingGuardrail()}, + ) + chunks: Final = tuple([chunk async for chunk in iterator]) + calls: Final = tuple( + (choice.index, call.id, call.function.arguments) + for chunk in chunks for choice in chunk.choices for call in choice.delta.tool_calls or () + ) + assert calls == ((0, "call_0", '{"contact":"MASKED"}'), (1, "call_1", '{"contact":"MASKED"}')) + + @pytest.mark.asyncio + async def test_undeliverable_rewrite_is_a_closed_failure(self) -> None: + from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite + + class UndeliverableTranslation(OpenAIChatCompletionsHandler): + async def process_output_streaming_response( + self, *args: object, **kwargs: object + ) -> list[ModelResponseStream]: + raise UndeliverableStreamRewrite("tool-redactor", "unmappable tool fragments") + + iterator: Final = UnifiedLLMGuardrails()._inspect_full_response( + endpoint_translation=UndeliverableTranslation(), guardrail_to_apply=_ToolRedactingGuardrail(), + request_data={}, user_api_key_dict=UserAPIKeyAuth(), responses_so_far=(), responses_yielded=(), + ) + with pytest.raises(unified_module.HTTPException) as error: + await anext(iterator) + assert error.value.status_code == 400 + assert error.value.detail == "Guardrail stream rewrite could not be applied" + @pytest.mark.asyncio async def test_block_only_drops_text_rewrites(self): """Default block_only: the guardrail's uppercasing never reaches the diff --git a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py index 5c85faa5e13..29a6e2ef0eb 100644 --- a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py +++ b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py @@ -7,7 +7,7 @@ with guardrail transformations, including tool calls. import json from collections.abc import Mapping -from typing import Any, Literal, Optional +from typing import Any, Final, Literal, Optional import pytest @@ -1400,24 +1400,21 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput: ] == [("call_1", '{"fruit": "persimmon"}'), ("call_2", '{"fruit": "durian"}')] @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, match="the stream carries 2 choices") as raised: - 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, - ) - - assert raised.value.guardrail_name == "test" - assert raised.value.reason == ( - "the stream carries 2 choices and tool-call rewrites are only written back on single-choice streams" + async def test_deliver_ended_stream_tool_rewrites_keep_choice_indices(self) -> None: + handler: Final = OpenAIChatCompletionsHandler() + chunks: Final = self._two_choice_tool_call_stream_chunks() + 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, ) + arguments: Final = tuple( + "".join(call.function.arguments for chunk in chunks for choice in chunk.choices + if choice.index == index for call in choice.delta.tool_calls or ()) + for index in range(2) + ) + assert arguments == ('{"fruit": "PERSIMMON"}', '{"fruit": "DURIAN"}') @pytest.mark.asyncio async def test_deliver_ended_stream_clean_multi_choice_stream_released_untouched(self):