feat(guardrails): let NeuralTrust transform verdicts reach streaming clients

The native hook never plumbed streaming_transform_mode, so a TrustGuard
transform was silently dropped on streamed tokens while the generic HTTP
path could turn incremental_diff on. Redaction only worked if a customer
kept the adapter the native guardrail is meant to replace.

Expose streaming_transform_mode on the guardrail card and wire it through
the initializer. Under incremental_diff each reply scan withholds the whole
reply via stream_holdback_chars: TrustGuard re-reads the full reply every
round and its redaction spans move as the reply grows, so releasing tokens
early lets a later scan try to rewrite bytes already on the wire, which the
engine can only reject as stream_transform_underflow. Holding them also
means a block lands with nothing streamed.

A scalar transform payload no longer rewrites the request conversation that
the framework attaches to reply scans for context; it rewrites the scanned
reply instead.
This commit is contained in:
albertbausili 2026-09-21 11:15:37 +02:00
parent c955057c65
commit ae61e703d2
8 changed files with 317 additions and 12 deletions

View file

@ -10045,6 +10045,22 @@
"description": "When True (default), after sensitive data is detected and routed, all subsequent requests in the same session will continue routing to the same model.",
"title": "Sticky Session Routing"
},
"streaming_transform_mode": {
"anyOf": [
{
"enum": [
"block_only",
"incremental_diff"
],
"type": "string"
},
{
"type": "null"
}
],
"description": "Whether a guardrail's text rewrite reaches a streaming client. Implemented by guardrail='prompt_security' and 'neuraltrust'; generic_guardrail_api takes the same setting under optional_params. 'block_only' (default) streams the raw model chunks, so a block still ends the stream but rewrites are dropped. 'incremental_diff' withholds those chunks and streams the guardrail's rewritten text instead. OpenAI chat completions streaming only.",
"title": "Streaming Transform Mode"
},
"template_id": {
"anyOf": [
{
@ -10071,7 +10087,7 @@
},
"unreachable_fallback": {
"default": "fail_closed",
"description": "Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', and 'typesafe'. 'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed.",
"description": "Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', 'typesafe', and 'neuraltrust'. 'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed.",
"enum": [
"fail_closed",
"fail_open"
@ -11570,6 +11586,18 @@
"description": "Client secret of the gateway's Entra app registration, used to perform the On-Behalf-Of exchange. Falls back to the AGENT365_CLIENT_SECRET environment variable.",
"title": "Client Secret"
},
"collector_key": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"description": "TrustGuard collector key (tgcol_...). Optional when the API key is bound to a collector. Env: TRUSTGUARD_COLLECTOR_KEY.",
"title": "Collector Key"
},
"confidence_threshold": {
"default": 0.5,
"default_value": 0.5,
@ -12853,6 +12881,22 @@
"description": "When True (default), after sensitive data is detected and routed, all subsequent requests in the same session will continue routing to the same model.",
"title": "Sticky Session Routing"
},
"streaming_transform_mode": {
"anyOf": [
{
"enum": [
"block_only",
"incremental_diff"
],
"type": "string"
},
{
"type": "null"
}
],
"description": "Whether a guardrail's text rewrite reaches a streaming client. Implemented by guardrail='prompt_security' and 'neuraltrust'; generic_guardrail_api takes the same setting under optional_params. 'block_only' (default) streams the raw model chunks, so a block still ends the stream but rewrites are dropped. 'incremental_diff' withholds those chunks and streams the guardrail's rewritten text instead. OpenAI chat completions streaming only.",
"title": "Streaming Transform Mode"
},
"template_id": {
"anyOf": [
{

View file

@ -18,6 +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
default_on: true
```
@ -50,7 +51,16 @@ HTTP 503 entitlements, 401/403, other 4xx/5xx, and unusable TrustGuard verdicts
## Streaming
LiteLLM streaming guardrails default to `block_only`. `block` still fires on streamed calls. `transform` rewrites are not applied to the streamed tokens; use non-streaming requests when DLP redaction must reach the client.
`block` and `ask` end the stream in either mode. `streaming_transform_mode` decides whether a `transform` verdict reaches the client.
| Mode | What the client receives |
| --- | --- |
| `block_only` (default) | the raw model tokens, so the redaction is dropped |
| `incremental_diff` | TrustGuard's rewritten reply |
Under `incremental_diff` the reply is held until the end-of-stream evaluate returns, so the first token arrives with the last. TrustGuard re-reads the whole reply on every scan and its redaction spans move as the reply grows, so releasing tokens early would let a later scan try to rewrite text already on the wire. Holding them also means a blocking verdict ends the stream with nothing sent at all.
`incremental_diff` covers OpenAI chat completions streaming; other surfaces fall back to `block_only`.
## References

View file

@ -19,6 +19,7 @@ def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail) ->
collector_key=litellm_params.collector_key,
unreachable_fallback=litellm_params.unreachable_fallback,
timeout=litellm_params.timeout,
streaming_transform_mode=litellm_params.streaming_transform_mode,
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,

View file

@ -201,6 +201,7 @@ class NeuralTrustGuardrail(CustomGuardrail):
collector_key: str | None = None,
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
timeout: float | None = None,
streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None,
guardrail_name: str | None = None,
event_hook: GuardrailEventHooks | Mode | str | Sequence[str] | None = None,
default_on: bool | None = None,
@ -220,6 +221,10 @@ class NeuralTrustGuardrail(CustomGuardrail):
if resolved_timeout <= 0:
raise ValueError("TrustGuard timeout must be a positive number of seconds.")
self.timeout = resolved_timeout
# Read off the instance by UnifiedLLMGuardrails.async_post_call_streaming_iterator_hook.
self.streaming_transform_mode: Literal["block_only", "incremental_diff"] = (
streaming_transform_mode or "block_only"
)
super().__init__(
guardrail_name=guardrail_name,
supported_event_hooks=self.get_supported_event_hooks(),
@ -235,6 +240,18 @@ class NeuralTrustGuardrail(CustomGuardrail):
request_data: dict, # mutable-ok: CustomGuardrail.apply_guardrail contract
input_type: Literal["request", "response"],
logging_obj: LiteLLMLoggingObj | None = None,
) -> GenericGuardrailAPIInputs:
return self._with_holdback(
await self._evaluate(inputs, request_data, input_type, logging_obj),
input_type,
)
async def _evaluate(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict, # mutable-ok: CustomGuardrail.apply_guardrail contract
input_type: Literal["request", "response"],
logging_obj: LiteLLMLoggingObj | None,
) -> GenericGuardrailAPIInputs:
body: Final = self._evaluate_body(inputs, request_data, input_type, logging_obj)
try:
@ -257,15 +274,32 @@ class NeuralTrustGuardrail(CustomGuardrail):
},
)
if status == STATUS_TRANSFORM:
return self._apply_transform(
inputs,
result.get("transformed_payload"),
sent_count=len(_sent_messages(inputs, input_type)),
)
return self._apply_transform(inputs, result.get("transformed_payload"), input_type=input_type)
if status == STATUS_REPORT:
verbose_proxy_logger.info("TrustGuard report-only findings trace_id=%s", result.get("trace_id"))
return inputs
def _with_holdback(
self,
inputs: GenericGuardrailAPIInputs,
input_type: Literal["request", "response"],
) -> GenericGuardrailAPIInputs:
"""Withhold the whole reply from each streaming round under incremental_diff.
TrustGuard re-reads the full reply every round and its redaction spans move as the reply grows, so a
later round can rewrite text an earlier one already streamed. The engine cannot retract streamed bytes
and answers that with stream_transform_underflow, so nothing is released until the end-of-stream round
forces the holdback to zero and flushes the final redacted reply.
"""
texts: Final = inputs.get("texts")
if input_type == "request" or self.streaming_transform_mode != "incremental_diff" or not texts:
return inputs
held: Final[GenericGuardrailAPIInputs] = { # mutable-ok: GenericGuardrailAPIInputs is a TypedDict
**inputs,
"stream_holdback_chars": [len(text) for text in texts], # mutable-ok: the field is a list
}
return held
def _evaluate_body(
self,
inputs: GenericGuardrailAPIInputs,
@ -370,7 +404,7 @@ class NeuralTrustGuardrail(CustomGuardrail):
inputs: GenericGuardrailAPIInputs,
transformed: object,
*,
sent_count: int,
input_type: Literal["request", "response"],
) -> GenericGuardrailAPIInputs:
if not isinstance(transformed, Mapping):
raise HTTPException(status_code=400, detail=TRANSFORM_MISSING)
@ -378,7 +412,7 @@ class NeuralTrustGuardrail(CustomGuardrail):
raw_messages: Final = transformed.get("messages")
if isinstance(raw_messages, list) and raw_messages:
rewritten_messages: Final = _copy_messages(raw_messages)
if rewritten_messages is None or len(rewritten_messages) != sent_count:
if rewritten_messages is None or len(rewritten_messages) != len(_sent_messages(inputs, input_type)):
raise HTTPException(status_code=400, detail=TRANSFORM_MISSING)
return _inputs_with_messages(inputs, rewritten_messages, replace_tool_calls=True)
@ -386,7 +420,9 @@ class NeuralTrustGuardrail(CustomGuardrail):
if not isinstance(raw_input, str) or not raw_input:
raise HTTPException(status_code=400, detail=TRANSFORM_MISSING)
original_messages: Final = inputs.get("structured_messages")
# structured_messages on a reply scan is the request conversation the framework attached for context,
# so a scalar payload rewrites the scanned reply text instead.
original_messages: Final = inputs.get("structured_messages") if input_type == "request" else None
if isinstance(original_messages, list) and original_messages:
copied: Final = _copy_messages(original_messages)
if copied is None:

View file

@ -1068,6 +1068,17 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
),
)
streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = Field(
default=None,
description=(
"Whether a guardrail's text rewrite reaches a streaming client. "
"Implemented by guardrail='prompt_security' and 'neuraltrust'; generic_guardrail_api takes the "
"same setting under optional_params. 'block_only' (default) streams the raw model chunks, so a "
"block still ends the stream but rewrites are dropped. 'incremental_diff' withholds those chunks "
"and streams the guardrail's rewritten text instead. OpenAI chat completions streaming only."
),
)
extra_headers: list[str] | None = Field(
default=None,
description=(

View file

@ -49,6 +49,17 @@ class NeuralTrustGuardrailConfigModel(GuardrailConfigModel):
),
)
streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = Field(
default=None,
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: "
"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."
),
)
@staticmethod
def ui_friendly_name() -> str:
return "NeuralTrust"

View file

@ -1,4 +1,5 @@
import os
from collections.abc import AsyncIterator, Sequence
from typing import Literal
from unittest.mock import AsyncMock, patch
@ -18,9 +19,18 @@ from litellm.proxy.guardrails.guardrail_hooks.neuraltrust import initialize_guar
from litellm.proxy.guardrails.guardrail_hooks.neuraltrust.neuraltrust import (
NeuralTrustGuardrail,
)
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import UnifiedLLMGuardrails
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.types.guardrails import LitellmParams
from litellm.types.utils import Choices, GenericGuardrailAPIInputs, Message, ModelResponse
from litellm.types.utils import (
Choices,
Delta,
GenericGuardrailAPIInputs,
Message,
ModelResponse,
ModelResponseStream,
StreamingChoices,
)
def _response(payload: object, status_code: int = 200) -> Response:
@ -49,6 +59,7 @@ def _guardrail(
default_on: bool = False,
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
timeout: float | None = None,
streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None,
api_base: str | None = None,
) -> NeuralTrustGuardrail:
return NeuralTrustGuardrail(
@ -59,10 +70,87 @@ def _guardrail(
default_on=default_on,
unreachable_fallback=unreachable_fallback,
timeout=timeout,
streaming_transform_mode=streaming_transform_mode,
api_base=api_base,
)
CARD_NUMBER = "4111 1111 1111 1111"
REPLY_CHUNKS = (
"Here is ",
"the billing ",
"record. ",
"Card ",
"4111 1111 ",
"1111 1111 ",
"is on file.",
)
FULL_REPLY = "".join(REPLY_CHUNKS)
REDACTED_REPLY = FULL_REPLY.replace(CARD_NUMBER, "[REDACTED]")
def _stream_chunk(content: str, finish_reason: str | None = None) -> ModelResponseStream:
return ModelResponseStream(
model="gpt-4o-mini",
choices=[
StreamingChoices(index=0, delta=Delta(content=content, role="assistant"), finish_reason=finish_reason)
],
)
async def _upstream_reply() -> AsyncIterator[ModelResponseStream]:
for chunk in REPLY_CHUNKS:
yield _stream_chunk(chunk)
yield _stream_chunk("", finish_reason="stop")
def _redacting_trustguard(*, on_card: str = "transform") -> AsyncMock:
"""A TrustGuard that only reacts once the whole card number is in the accumulated reply."""
async def _post(*_args: object, **kwargs: object) -> Response:
body = kwargs["json"]
seen = "".join(message["content"] or "" for message in body["payload"]["messages"])
if CARD_NUMBER not in seen:
return _response({"status": "allow"})
if on_card != "transform":
return _response({"status": on_card, "trace_id": "trace-1"})
return _response(
{
"status": "transform",
"transformed_payload": {
"messages": [{"role": "assistant", "content": seen.replace(CARD_NUMBER, "[REDACTED]")}]
},
}
)
return AsyncMock(side_effect=_post)
def _guardrail_stream(guardrail: NeuralTrustGuardrail) -> AsyncIterator[object]:
return UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="tgk_test", request_route="/v1/chat/completions"),
response=_upstream_reply(),
request_data={
"guardrail_to_apply": guardrail,
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "what card is on file?"}],
},
)
async def _drain_into(stream: AsyncIterator[object], sink: list[object]) -> None:
async for item in stream:
sink.append(item)
def _deltas(items: Sequence[object]) -> list[str]:
return [
item.choices[0].delta.content
for item in items
if isinstance(item, ModelResponseStream) and item.choices and item.choices[0].delta.content
]
class TestNeuralTrustGuardrail:
def setup_method(self) -> None:
for key in ("TRUSTGUARD_API_KEY", "TRUSTGUARD_API_BASE", "TRUSTGUARD_COLLECTOR_KEY"):
@ -994,6 +1082,91 @@ class TestNeuralTrustGuardrail:
assert result == inputs
assert mock_post.call_args.kwargs["timeout"] == 12.0
@pytest.mark.asyncio
async def test_incremental_diff_redacts_a_value_split_across_streamed_chunks(self) -> None:
"""The card straddles a sampled scan, so the earlier half must never leave before the redaction lands."""
guardrail = _guardrail(event_hook="post_call", default_on=True, streaming_transform_mode="incremental_diff")
with patch.object(guardrail.async_handler, "post", _redacting_trustguard()):
out = [item async for item in _guardrail_stream(guardrail)]
assert _deltas(out) == [REDACTED_REPLY]
@pytest.mark.asyncio
async def test_block_mid_stream_under_incremental_diff_sends_nothing(self) -> 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", _redacting_trustguard(on_card="block")):
with pytest.raises(HTTPException) as exc_info:
await _drain_into(_guardrail_stream(guardrail), received)
assert exc_info.value.status_code == 400
assert exc_info.value.detail["verdict"] == "block"
assert _deltas(received) == []
@pytest.mark.asyncio
async def test_default_streaming_mode_leaves_the_transform_off_the_wire(self) -> None:
"""Default stays block_only, where the framework streams the raw model chunks and drops rewrites."""
guardrail = _guardrail(event_hook="post_call", default_on=True)
assert guardrail.streaming_transform_mode == "block_only"
with patch.object(guardrail.async_handler, "post", _redacting_trustguard()):
out = [item async for item in _guardrail_stream(guardrail)]
streamed = "".join(_deltas(out))
assert CARD_NUMBER in streamed
assert "[REDACTED]" not in streamed
@pytest.mark.asyncio
@pytest.mark.parametrize(
("mode", "input_type", "expected"),
[
("incremental_diff", "response", [len(REDACTED_REPLY)]),
("incremental_diff", "request", None),
("block_only", "response", None),
],
)
async def test_holdback_covers_streamed_reply_scans_only(
self,
mode: Literal["block_only", "incremental_diff"],
input_type: Literal["request", "response"],
expected: list[int] | None,
) -> None:
guardrail = _guardrail(event_hook="post_call", streaming_transform_mode=mode)
mock_post = AsyncMock(
return_value=_response(
{
"status": "transform",
"transformed_payload": {"messages": [{"role": "assistant", "content": REDACTED_REPLY}]},
}
)
)
with patch.object(guardrail.async_handler, "post", mock_post):
result = await guardrail.apply_guardrail(
inputs={"texts": [FULL_REPLY]},
request_data={},
input_type=input_type,
logging_obj=_logging(),
)
assert result.get("stream_holdback_chars") == expected
@pytest.mark.asyncio
async def test_transform_input_on_a_reply_rewrites_the_reply_not_the_prompt(self) -> None:
"""The framework attaches the request conversation to reply scans; a scalar payload must ignore it."""
guardrail = _guardrail(event_hook="post_call")
mock_post = AsyncMock(
return_value=_response({"status": "transform", "transformed_payload": {"input": "card [REDACTED]"}})
)
with patch.object(guardrail.async_handler, "post", mock_post):
result = await guardrail.apply_guardrail(
inputs={
"texts": ["card 4111 1111 1111 1111"],
"structured_messages": [
{"role": "user", "content": "what card is on file?"},
{"role": "assistant", "content": "card 4111 1111 1111 1111"},
],
},
request_data={},
input_type="response",
logging_obj=_logging(),
)
assert result["texts"] == ["card [REDACTED]"]
def test_get_config_model(self) -> None:
model = NeuralTrustGuardrail.get_config_model()
assert model is not None
@ -1009,10 +1182,12 @@ class TestNeuralTrustGuardrail:
"collector_key",
"unreachable_fallback",
"timeout",
"streaming_transform_mode",
}
assert fields["timeout"]["type"] == "number"
assert fields["timeout"]["default_value"] == 5.0
assert fields["unreachable_fallback"]["options"] == ["fail_closed", "fail_open"]
assert fields["streaming_transform_mode"]["options"] == ["block_only", "incremental_diff"]
def test_timeout_default_stays_local_to_neuraltrust(self) -> None:
assert LitellmParams(guardrail="lakera_v2", mode="pre_call").timeout is None
@ -1035,6 +1210,7 @@ class TestNeuralTrustGuardrail:
collector_key="tgcol_from_params",
unreachable_fallback="fail_open",
timeout=2,
streaming_transform_mode="incremental_diff",
default_on=True,
)
hook = initialize_guardrail(params, {"guardrail_name": "tg-prod"})
@ -1044,6 +1220,7 @@ class TestNeuralTrustGuardrail:
assert hook.collector_key == "tgcol_from_params"
assert hook.unreachable_fallback == "fail_open"
assert hook.timeout == 2.0
assert hook.streaming_transform_mode == "incremental_diff"
assert hook.guardrail_name == "tg-prod"
assert hook.default_on is True
assert hook in litellm.callbacks

View file

@ -24547,6 +24547,11 @@ export interface components {
* @default true
*/
sticky_session_routing: boolean | null;
/**
* Streaming Transform Mode
* @description Whether a guardrail's text rewrite reaches a streaming client. Implemented by guardrail='prompt_security' and 'neuraltrust'; generic_guardrail_api takes the same setting under optional_params. 'block_only' (default) streams the raw model chunks, so a block still ends the stream but rewrites are dropped. 'incremental_diff' withholds those chunks and streams the guardrail's rewritten text instead. OpenAI chat completions streaming only.
*/
streaming_transform_mode?: ("block_only" | "incremental_diff") | null;
/**
* Template Id
* @description The ID of your Model Armor template
@ -24559,7 +24564,7 @@ export interface components {
timeout?: number | null;
/**
* Unreachable Fallback
* @description Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', and 'typesafe'. 'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed.
* @description Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', 'typesafe', and 'neuraltrust'. 'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed.
* @default fail_closed
* @enum {string}
*/
@ -32113,6 +32118,11 @@ export interface components {
* @description Client secret of the gateway's Entra app registration, used to perform the On-Behalf-Of exchange. Falls back to the AGENT365_CLIENT_SECRET environment variable.
*/
client_secret?: string | null;
/**
* Collector Key
* @description TrustGuard collector key (tgcol_...). Optional when the API key is bound to a collector. Env: TRUSTGUARD_COLLECTOR_KEY.
*/
collector_key?: string | null;
/**
* Confidence Threshold
* @description Only block or mask when detection confidence >= this value; below threshold, allow or log_only.
@ -32653,6 +32663,11 @@ export interface components {
* @default true
*/
sticky_session_routing: boolean | null;
/**
* Streaming Transform Mode
* @description Whether a guardrail's text rewrite reaches a streaming client. Implemented by guardrail='prompt_security' and 'neuraltrust'; generic_guardrail_api takes the same setting under optional_params. 'block_only' (default) streams the raw model chunks, so a block still ends the stream but rewrites are dropped. 'incremental_diff' withholds those chunks and streams the guardrail's rewritten text instead. OpenAI chat completions streaming only.
*/
streaming_transform_mode?: ("block_only" | "incremental_diff") | null;
/**
* Template Id
* @description The ID of your Model Armor template