mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
4ba8517134
commit
8f02361740
3 changed files with 38 additions and 0 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue