fix(guardrails): stream Prompt Security post_call redactions in incremental_diff mode

Forward streaming_transform_mode from guardrail litellm_params into PromptSecurityGuardrail so incremental_diff is reachable from config; the default stays block_only. In incremental_diff the guardrail now returns stream_holdback_chars alongside the rewritten texts so that a value split across streamed chunks (or across an abbreviation period) is never partially released before the vendor rewrite arrives. Each response text gets its own protect call so modified_text maps back to the right choice when n > 1, and custom_guardrail no longer logs a clean response as mask just because the guardrail attached holdback metadata

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-17 00:47:17 +00:00
parent 38676aa599
commit 7815719de7
6 changed files with 279 additions and 27 deletions

View file

@ -1379,8 +1379,9 @@ class CustomGuardrail(CustomLogger):
raise e
def _inputs_were_modified(self, original_inputs: Mapping[str, object], response: Mapping[str, object]) -> bool:
"""True when any key of either mapping differs between them (mask), False otherwise (allow)."""
return any(original_inputs.get(key) != response.get(key) for key in original_inputs.keys() | response.keys())
"""True when any content key of either mapping differs between them (mask), False otherwise (allow)."""
compared_keys: Final = (original_inputs.keys() | response.keys()) - _STREAM_CONTROL_KEYS
return any(original_inputs.get(key) != response.get(key) for key in compared_keys)
def mask_content_in_string(
self,
@ -1490,6 +1491,7 @@ def _sync_guardrail_info_to_logging_obj(request_data: dict, logging_obj: object)
_PRE_CALL_CONTENT_KEYS: Final = frozenset(
{"messages", "input", "prompt", "system", "instructions", "tools", "functions", "function_call", "tool_choice"}
)
_STREAM_CONTROL_KEYS: Final = frozenset({"stream_holdback_chars"})
def _original_inputs_for(

View file

@ -20,6 +20,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
streaming_transform_mode=getattr(litellm_params, "streaming_transform_mode", None),
file_sanitization_fail_open=getattr(litellm_params, "file_sanitization_fail_open", None),
block_on_file_modify=getattr(litellm_params, "block_on_file_modify", None),
)

View file

@ -38,6 +38,11 @@ class PromptSecurityGuardrailMissingSecrets(Exception):
pass
def _modified_or_original(text: str, verdict: "_ProtectVerdict") -> str:
modified_text: Final = verdict.get("modified_text") if verdict.get("action") == "modify" else None
return text if modified_text is None else modified_text
def _inputs_with_structured_messages(
inputs: GenericGuardrailAPIInputs, rewritten_messages: Sequence[AllMessageValues] | None
) -> GenericGuardrailAPIInputs:
@ -119,6 +124,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
user: str | None = None,
system_prompt: str | None = None,
check_tool_results: bool | None = None,
streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None,
file_sanitization_timeout: float = _SANITIZE_FILE_FAIL_OPEN_TIMEOUT_SECONDS,
file_sanitization_fail_open: bool | None = None,
block_on_file_modify: bool | None = None,
@ -148,6 +154,10 @@ class PromptSecurityGuardrail(CustomGuardrail):
)
raise PromptSecurityGuardrailMissingSecrets(msg)
self.streaming_transform_mode: Literal["block_only", "incremental_diff"] = (
"block_only" if streaming_transform_mode is None else streaming_transform_mode
)
# Configuration for file sanitization
self.max_poll_attempts = 30 # Maximum number of polling attempts
self.poll_interval = 2 # Seconds between polling attempts
@ -342,16 +352,46 @@ class PromptSecurityGuardrail(CustomGuardrail):
texts: list[str],
user_api_key_alias: str | None,
) -> GenericGuardrailAPIInputs:
"""Handle response-side guardrail checks."""
"""Handle response-side guardrail checks, one protect verdict per text.
Prompt Security rewrites a single string, so texts from several choices must be scanned separately
or one ``modified_text`` cannot be mapped back onto the choice it came from. It also returns no span
offsets, so on a stream every text is held back in full until the final verdict: a value the vendor
redacts later may start anywhere in text that looked clean so far, and streamed bytes cannot be recalled.
"""
if not texts:
return inputs
# Combine all texts for response checking
combined_text: Final = "\n".join(texts)
verdicts: Final = await asyncio.gather(
*(self._protect_response_text(text, user_api_key_alias) for text in texts)
)
violations: Final = tuple(
violation
for verdict in verdicts
if verdict.get("action") == "block"
for violation in verdict.get("violations", ())
)
if any(verdict.get("action") == "block" for verdict in verdicts):
raise HTTPException(
status_code=400,
detail="Blocked by Prompt Security, Violations: " + ", ".join(violations),
)
returned_texts: Final = [ # mutable-ok: GenericGuardrailAPIInputs.texts is list[str]
_modified_or_original(text, verdict) for text, verdict in zip(texts, verdicts, strict=True)
]
patched: Final[GenericGuardrailAPIInputs] = {
**inputs,
"texts": returned_texts,
"stream_holdback_chars": [ # mutable-ok: GenericGuardrailAPIInputs.stream_holdback_chars is list[int]
len(text) for text in returned_texts
],
}
return patched
async def _protect_response_text(self, text: str, user_api_key_alias: str | None) -> _ProtectVerdict:
headers: Final = self._build_headers(user_api_key_alias)
payload: Final = {
"response": combined_text,
"response": text,
"user": user_api_key_alias or self.user,
"system_prompt": self.system_prompt,
}
@ -360,7 +400,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
method="POST",
url=f"{self.api_base}/api/protect",
headers=headers,
payload={"response_length": len(combined_text)},
payload={"response_length": len(text)},
)
response: Final = await self.async_handler.post(
@ -377,26 +417,8 @@ class PromptSecurityGuardrail(CustomGuardrail):
payload={"result": res.get("result")},
)
result: Final = res.get("result", {}).get("response", {})
if result is None:
return inputs
action: Final = result.get("action")
violations: Final = result.get("violations", [])
if action == "block":
raise HTTPException(
status_code=400,
detail="Blocked by Prompt Security, Violations: " + ", ".join(violations),
)
elif action == "modify":
modified_text: Final = result.get("modified_text")
if modified_text is not None:
# If we combined multiple texts, return the modified version as single text
# The framework will handle distributing it back
inputs["texts"] = [modified_text]
return inputs
verdict: Final = res.get("result", {}).get("response", {})
return {} if verdict is None else verdict
def _extract_texts_from_messages(self, messages: Sequence[Mapping[str, object]]) -> list[str]:
return [text for message in messages for text in message_slot_texts(message)]

View file

@ -1,3 +1,5 @@
from typing import Literal
from pydantic import Field
from .base import GuardrailConfigModel
@ -20,6 +22,16 @@ class PromptSecurityGuardrailConfigModel(GuardrailConfigModel):
default=True,
description="Whether a file sanitization `modify` verdict blocks the request instead of replacing the file content.",
)
streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = Field(
default=None,
description=(
"How post_call `modify` verdicts reach a streaming client. `block_only` (default) streams the raw upstream "
"chunks and only a `block` verdict ends the stream, so `modified_text` is dropped. `incremental_diff` "
"buffers the whole response and sends the redacted text once the final verdict is in, so the first token "
"arrives with the last, while a `block` verdict still ends the stream early. "
"OpenAI chat completions streaming only."
),
)
@staticmethod
def ui_friendly_name() -> str:

View file

@ -3130,3 +3130,22 @@ class TestPreCallHookResponseIsNotLoggedVerbatim:
)
assert self._logged_response(data) == "mask"
@pytest.mark.asyncio
async def test_apply_guardrail_adding_only_stream_holdback_logs_allow(self):
class HoldbackOnlyGuardrail(CustomGuardrail):
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> GenericGuardrailAPIInputs:
return {**inputs, "stream_holdback_chars": [6]}
data = self._request()
await HoldbackOnlyGuardrail(guardrail_name="g").apply_guardrail(
inputs={"texts": ["SECRET_PROMPT"]}, request_data=data, input_type="response"
)
assert self._logged_response(data) == "allow"

View file

@ -8,12 +8,15 @@ from fastapi.exceptions import HTTPException
from httpx import ReadTimeout, Request, Response
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.prompt_security.prompt_security import (
PromptSecurityGuardrail,
PromptSecurityGuardrailMissingSecrets,
)
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import UnifiedLLMGuardrails
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
def test_prompt_security_guard_config(monkeypatch: pytest.MonkeyPatch):
@ -415,6 +418,199 @@ async def test_apply_guardrail_modify_response(monkeypatch: pytest.MonkeyPatch):
assert result["texts"] == ["Your SSN is [REDACTED]"]
@pytest.mark.asyncio
async def test_apply_guardrail_modify_response_keeps_multi_choice_texts_aligned():
"""With n>1 each choice text gets its own verdict, so a rewrite lands on the choice it came from."""
guardrail = PromptSecurityGuardrail(
guardrail_name="test-guard",
event_hook="post_call",
default_on=True,
api_key="test-key",
api_base="https://test.prompt.security",
)
async def mock_post(*args, **kwargs):
text = kwargs["json"]["response"]
redacted = text.replace("123-45-6789", "[REDACTED]")
mock_response = Response(
json={
"result": {
"response": {
"action": "modify" if redacted != text else "log",
"violations": [],
"modified_text": redacted,
}
}
},
status_code=200,
request=Request(method="POST", url="https://test.prompt.security/api/protect"),
)
mock_response.raise_for_status = lambda: None
return mock_response
with patch.object(guardrail.async_handler, "post", side_effect=mock_post):
result = await guardrail.apply_guardrail(
inputs={"texts": ["all clear", "SSN 123-45-6789 on file"]},
request_data={},
input_type="response",
)
assert result["texts"] == ["all clear", "SSN [REDACTED] on file"]
assert result["stream_holdback_chars"] == [len("all clear"), len("SSN [REDACTED] on file")]
def test_prompt_security_streaming_transform_mode_from_config(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, "guardrail_name_config_map", {})
monkeypatch.setattr(litellm, "callbacks", [])
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
init_guardrails_v2(
all_guardrails=[
{
"guardrail_name": "prompt_security_streaming",
"litellm_params": {
"guardrail": "prompt_security",
"mode": "post_call",
"default_on": True,
"streaming_transform_mode": "incremental_diff",
},
}
],
config_file_path="",
)
registered = [c for c in litellm.callbacks if isinstance(c, PromptSecurityGuardrail)]
assert len(registered) == 1
assert registered[0].streaming_transform_mode == "incremental_diff"
assert PromptSecurityGuardrail(api_key="k", api_base="https://b").streaming_transform_mode == "block_only"
def _stream_chunk(content: str, finish_reason: str | None = None) -> ModelResponseStream:
return ModelResponseStream(
choices=[StreamingChoices(index=0, delta=Delta(content=content, role="assistant"), finish_reason=finish_reason)]
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("chunks", "secret", "redacted_output"),
[
pytest.param(
(
"Sure. I checked the billing record for this account and confirmed the details below. Card 4111 1111 ",
"1111 1111 is on file.",
),
"4111 1111 1111 1111",
"Sure. I checked the billing record for this account and confirmed the details below. "
"Card [REDACTED] is on file.",
id="spaced_value_after_full_sentence",
),
pytest.param(
("Ship to 12 Main St. ", "Springfield 62704 today."),
"12 Main St. Springfield 62704",
"Ship to [REDACTED] today.",
id="value_spanning_abbreviation_period",
),
pytest.param(
(
"Customer record follows.\nName: John Smith\n"
"Address: 12 Main St, Springfield IL 62704, United States\n",
"SSN: 123-45-6789\nThat is all.",
),
"Name: John Smith\nAddress: 12 Main St, Springfield IL 62704, United States\nSSN: 123-45-6789",
"Customer record follows.\n[REDACTED]\nThat is all.",
id="multi_line_record_redacted_as_one_span",
),
],
)
async def test_prompt_security_incremental_diff_redacts_value_split_across_chunks(
chunks: tuple[str, ...],
secret: str,
redacted_output: str,
):
"""A modify verdict reaches the client redacted even when the value straddles a sampled scan."""
guardrail = PromptSecurityGuardrail(
guardrail_name="prompt_security_streaming",
event_hook="post_call",
default_on=True,
api_key="test-key",
api_base="https://test.prompt.security",
streaming_transform_mode="incremental_diff",
)
guardrail.streaming_sampling_rate = 1
async def mock_post(*args, **kwargs):
text = kwargs["json"]["response"]
redacted = text.replace(secret, "[REDACTED]")
mock_response = Response(
json={
"result": {
"response": {
"action": "modify" if redacted != text else "log",
"violations": ["pii"] if redacted != text else [],
"modified_text": redacted,
}
}
},
status_code=200,
request=Request(method="POST", url="https://test.prompt.security/api/protect"),
)
mock_response.raise_for_status = lambda: None
return mock_response
async def _upstream():
for chunk in chunks:
yield _stream_chunk(chunk)
yield _stream_chunk("", finish_reason="stop")
with patch.object(guardrail.async_handler, "post", side_effect=mock_post):
out = [
item
async for item in 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": guardrail, "model": "gpt-4"},
)
]
assert all(isinstance(item, ModelResponseStream) for item in out)
deltas = [item.choices[0].delta.content for item in out if item.choices and item.choices[0].delta.content]
assert deltas == [redacted_output]
assert all(secret[:6] not in delta for delta in deltas)
@pytest.mark.asyncio
async def test_prompt_security_clean_non_streaming_response_logs_allow():
"""A log verdict keeps the text (even if modified_text is present) and is logged as allow."""
guardrail = PromptSecurityGuardrail(
guardrail_name="prompt_security_streaming",
event_hook="post_call",
default_on=True,
api_key="test-key",
api_base="https://test.prompt.security",
streaming_transform_mode="incremental_diff",
)
mock_response = Response(
json={"result": {"response": {"action": "log", "violations": [], "modified_text": "order noted"}}},
status_code=200,
request=Request(method="POST", url="https://test.prompt.security/api/protect"),
)
mock_response.raise_for_status = lambda: None
request_data = {"metadata": {}}
with patch.object(guardrail.async_handler, "post", return_value=mock_response):
result = await guardrail.apply_guardrail(
inputs={"texts": ["order confirmed"]},
request_data=request_data,
input_type="response",
)
assert result["texts"] == ["order confirmed"]
info = request_data["metadata"]["standard_logging_guardrail_information"]
assert [entry["guardrail_response"] for entry in info] == ["allow"]
@pytest.mark.asyncio
async def test_file_sanitization(monkeypatch: pytest.MonkeyPatch):
"""Test file sanitization for images"""