mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(guardrails): configure Prompt Security timeout policy
This commit is contained in:
parent
4cb9599090
commit
3bba1f66a3
4 changed files with 48 additions and 4 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,
|
||||
file_sanitization_fail_open=getattr(litellm_params, "file_sanitization_fail_open", None),
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_prompt_security_callback)
|
||||
|
||||
|
|
|
|||
|
|
@ -92,6 +92,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
system_prompt: str | None = None,
|
||||
check_tool_results: bool | None = None,
|
||||
file_sanitization_timeout: float = _SANITIZE_FILE_FAIL_OPEN_TIMEOUT_SECONDS,
|
||||
file_sanitization_fail_open: bool | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
|
||||
|
|
@ -122,6 +123,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
self.max_poll_attempts = 30 # Maximum number of polling attempts
|
||||
self.poll_interval = 2 # Seconds between polling attempts
|
||||
self.file_sanitization_timeout = file_sanitization_timeout
|
||||
self.file_sanitization_fail_open = file_sanitization_fail_open is not False
|
||||
|
||||
super().__init__(**kwargs)
|
||||
|
||||
|
|
@ -417,6 +419,14 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
timeout=self.file_sanitization_timeout,
|
||||
)
|
||||
except (asyncio.TimeoutError, httpx.TimeoutException, LiteLLMTimeout) as exc:
|
||||
if not self.file_sanitization_fail_open:
|
||||
verbose_proxy_logger.error(
|
||||
"Prompt Security Guardrail: file sanitization for %s timed out with %s; failing closed",
|
||||
filename,
|
||||
type(exc).__name__,
|
||||
)
|
||||
raise HTTPException(status_code=408, detail="File sanitization timeout") from exc
|
||||
|
||||
verbose_proxy_logger.error(
|
||||
"Prompt Security Guardrail: file sanitization for %s timed out with %s; failing open",
|
||||
filename,
|
||||
|
|
|
|||
|
|
@ -12,6 +12,10 @@ class PromptSecurityGuardrailConfigModel(GuardrailConfigModel):
|
|||
default=None,
|
||||
description="The API base for the Prompt Security guardrail. If not provided, the `PROMPT_SECURITY_API_BASE` environment variable is used.",
|
||||
)
|
||||
file_sanitization_fail_open: bool = Field(
|
||||
default=True,
|
||||
description="Whether file sanitization timeouts allow the original file through instead of blocking the request.",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ def test_prompt_security_guard_config(monkeypatch: pytest.MonkeyPatch):
|
|||
"guardrail": "prompt_security",
|
||||
"mode": "during_call",
|
||||
"default_on": True,
|
||||
"file_sanitization_fail_open": False,
|
||||
},
|
||||
}
|
||||
],
|
||||
|
|
@ -41,6 +42,10 @@ def test_prompt_security_guard_config(monkeypatch: pytest.MonkeyPatch):
|
|||
assert registered[0].guardrail_name == "prompt_security"
|
||||
assert registered[0].default_on is True
|
||||
assert registered[0].event_hook == "during_call"
|
||||
assert registered[0].file_sanitization_fail_open is False
|
||||
config_model = registered[0].get_config_model()
|
||||
assert config_model is not None
|
||||
assert config_model().file_sanitization_fail_open is True
|
||||
|
||||
|
||||
def test_prompt_security_guard_config_no_api_key(monkeypatch: pytest.MonkeyPatch):
|
||||
|
|
@ -390,13 +395,28 @@ async def test_file_sanitization(monkeypatch: pytest.MonkeyPatch):
|
|||
),
|
||||
ids=("litellm", "httpx"),
|
||||
)
|
||||
async def test_file_sanitization_fails_open_on_request_timeout(monkeypatch: pytest.MonkeyPatch, timeout: Exception):
|
||||
@pytest.mark.parametrize("fail_open", (True, False), ids=("fail-open", "fail-closed"))
|
||||
async def test_file_sanitization_request_timeout_policy(
|
||||
monkeypatch: pytest.MonkeyPatch, timeout: Exception, fail_open: bool
|
||||
):
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
|
||||
guardrail = PromptSecurityGuardrail(guardrail_name="test-guard", event_hook="pre_call", default_on=True)
|
||||
guardrail = PromptSecurityGuardrail(
|
||||
guardrail_name="test-guard",
|
||||
event_hook="pre_call",
|
||||
default_on=True,
|
||||
file_sanitization_fail_open=fail_open,
|
||||
)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(side_effect=timeout)):
|
||||
if not fail_open:
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.sanitize_file_content(b"file-content", "document.pdf")
|
||||
assert exc_info.value.status_code == 408
|
||||
assert exc_info.value.detail == "File sanitization timeout"
|
||||
return
|
||||
|
||||
result = await guardrail.sanitize_file_content(b"file-content", "document.pdf")
|
||||
|
||||
assert result == {
|
||||
|
|
@ -408,7 +428,8 @@ async def test_file_sanitization_fails_open_on_request_timeout(monkeypatch: pyte
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_file_sanitization_fails_open_on_overall_timeout(monkeypatch: pytest.MonkeyPatch):
|
||||
@pytest.mark.parametrize("fail_open", (True, False), ids=("fail-open", "fail-closed"))
|
||||
async def test_file_sanitization_overall_timeout_policy(monkeypatch: pytest.MonkeyPatch, fail_open: bool):
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
|
||||
|
|
@ -417,13 +438,21 @@ async def test_file_sanitization_fails_open_on_overall_timeout(monkeypatch: pyte
|
|||
event_hook="pre_call",
|
||||
default_on=True,
|
||||
file_sanitization_timeout=0.01,
|
||||
file_sanitization_fail_open=fail_open,
|
||||
)
|
||||
|
||||
async def hanging_post(*args, **kwargs):
|
||||
async def hanging_post(*_args: object, **_kwargs: object) -> None:
|
||||
await asyncio.sleep(60)
|
||||
raise AssertionError("sanitization request should have been cancelled")
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=hanging_post):
|
||||
if not fail_open:
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.sanitize_file_content(b"file-content", "document.pdf")
|
||||
assert exc_info.value.status_code == 408
|
||||
assert exc_info.value.detail == "File sanitization timeout"
|
||||
return
|
||||
|
||||
result = await guardrail.sanitize_file_content(b"file-content", "document.pdf")
|
||||
|
||||
assert result["action"] == "allow"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue