fix(guardrails): apply buffered tool argument rewrites

This commit is contained in:
albertbausili 2026-09-21 18:06:10 +02:00
parent 080b5a486a
commit 5c437f7efb
6 changed files with 143 additions and 78 deletions

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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."
),

View file

@ -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

View file

@ -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):