fix(guardrails): wire streaming_transform_mode through prompt_security initializer

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-08-31 14:02:43 +00:00
parent 4ba8517134
commit 8f02361740
3 changed files with 38 additions and 0 deletions

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),
)
litellm.logging_callback_manager.add_litellm_callback(_prompt_security_callback)

View file

@ -79,6 +79,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,
**kwargs,
):
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
@ -105,6 +106,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

View file

@ -637,3 +637,35 @@ async def test_check_tool_results_enabled(monkeypatch: pytest.MonkeyPatch):
assert "indirect_prompt_injection" in str(excinfo.value.detail)
def test_prompt_security_streaming_transform_mode_from_config(monkeypatch: pytest.MonkeyPatch):
"""streaming_transform_mode in litellm_params must reach the guardrail instance,
otherwise the unified streaming hook stays in block_only and drops redactions."""
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"
def test_prompt_security_streaming_transform_mode_defaults_block_only():
guardrail = PromptSecurityGuardrail(api_key="k", api_base="https://test.prompt.security")
assert guardrail.streaming_transform_mode == "block_only"