mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge fa15acf355 into b03e913ccf
This commit is contained in:
commit
53172f9766
4 changed files with 288 additions and 31 deletions
|
|
@ -20,6 +20,8 @@ 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),
|
||||
block_on_file_modify=getattr(litellm_params, "block_on_file_modify", None),
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_prompt_security_callback)
|
||||
|
||||
|
|
|
|||
|
|
@ -4,10 +4,12 @@ import os
|
|||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Final, Literal, Optional
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.exceptions import Timeout as LiteLLMTimeout
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
|
|
@ -24,6 +26,9 @@ if TYPE_CHECKING:
|
|||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
||||
|
||||
_SANITIZE_FILE_FAIL_OPEN_TIMEOUT_SECONDS: Final = 30.0
|
||||
|
||||
|
||||
class PromptSecurityGuardrailMissingSecrets(Exception):
|
||||
pass
|
||||
|
||||
|
|
@ -63,6 +68,13 @@ class _SanitizeStatusResponse(TypedDict, total=False):
|
|||
metadata: ReadOnly[_SanitizeMetadata]
|
||||
|
||||
|
||||
class _SanitizeResult(TypedDict):
|
||||
action: ReadOnly[str]
|
||||
content: ReadOnly[str | None]
|
||||
metadata: ReadOnly[_SanitizeMetadata]
|
||||
violations: ReadOnly[Sequence[str]]
|
||||
|
||||
|
||||
class PromptSecurityGuardrail(CustomGuardrail):
|
||||
@classmethod
|
||||
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
|
||||
|
|
@ -79,6 +91,9 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
user: str | None = None,
|
||||
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,
|
||||
block_on_file_modify: bool | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
|
||||
|
|
@ -108,6 +123,9 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
# Configuration for file sanitization
|
||||
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
|
||||
self.block_on_file_modify = block_on_file_modify is not False
|
||||
|
||||
super().__init__(**kwargs)
|
||||
|
||||
|
|
@ -356,13 +374,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
result = await self.sanitize_file_content(
|
||||
file_data, filename, user_api_key_alias=user_api_key_alias
|
||||
)
|
||||
|
||||
if result.get("action") == "block":
|
||||
violations = result.get("violations", [])
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Image blocked by Prompt Security. Violations: {', '.join(violations)}",
|
||||
)
|
||||
self._raise_if_file_blocked(result, "Image")
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
|
|
@ -392,11 +404,44 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
file_data: bytes,
|
||||
filename: str,
|
||||
user_api_key_alias: str | None = None,
|
||||
) -> dict:
|
||||
) -> _SanitizeResult:
|
||||
"""
|
||||
Sanitize file content using Prompt Security API.
|
||||
Returns: dict with keys 'action', 'content', 'metadata'
|
||||
"""
|
||||
try:
|
||||
return await asyncio.wait_for(
|
||||
self._sanitize_file_content(file_data, filename, user_api_key_alias),
|
||||
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,
|
||||
type(exc).__name__,
|
||||
)
|
||||
fail_open_result: Final[_SanitizeResult] = {
|
||||
"action": "allow",
|
||||
"content": None,
|
||||
"metadata": {},
|
||||
"violations": (),
|
||||
}
|
||||
return fail_open_result
|
||||
|
||||
async def _sanitize_file_content(
|
||||
self,
|
||||
file_data: bytes,
|
||||
filename: str,
|
||||
user_api_key_alias: str | None,
|
||||
) -> _SanitizeResult:
|
||||
headers: Final = {"APP-ID": self.api_key}
|
||||
if user_api_key_alias:
|
||||
headers["X-LiteLLM-Key-Alias"] = user_api_key_alias
|
||||
|
|
@ -479,6 +524,17 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
|
||||
raise HTTPException(status_code=408, detail="File sanitization timeout")
|
||||
|
||||
def _raise_if_file_blocked(self, sanitization_result: _SanitizeResult, resource_name: str) -> None:
|
||||
action: Final = sanitization_result.get("action")
|
||||
if action != "block" and not (action == "modify" and self.block_on_file_modify):
|
||||
return
|
||||
|
||||
violations: Final = sanitization_result.get("violations", ())
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"{resource_name} blocked by Prompt Security. Violations: {', '.join(violations)}",
|
||||
)
|
||||
|
||||
async def _process_image_url_item(self, item: dict, user_api_key_alias: str | None) -> dict:
|
||||
"""Process and sanitize image_url items."""
|
||||
image_url_data: Final = item.get("image_url", {})
|
||||
|
|
@ -498,13 +554,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
file_data, filename, user_api_key_alias=user_api_key_alias
|
||||
)
|
||||
action: Final = sanitization_result.get("action")
|
||||
|
||||
if action == "block":
|
||||
violations: Final = sanitization_result.get("violations", [])
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"File blocked by Prompt Security. Violations: {', '.join(violations)}",
|
||||
)
|
||||
self._raise_if_file_blocked(sanitization_result, "File")
|
||||
|
||||
if action == "modify":
|
||||
sanitized_content: Final = sanitization_result.get("content", "")
|
||||
|
|
@ -566,13 +616,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
file_data, filename, user_api_key_alias=user_api_key_alias
|
||||
)
|
||||
action: Final = sanitization_result.get("action")
|
||||
|
||||
if action == "block":
|
||||
violations: Final = sanitization_result.get("violations", [])
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Document blocked by Prompt Security. Violations: {', '.join(violations)}",
|
||||
)
|
||||
self._raise_if_file_blocked(sanitization_result, "Document")
|
||||
|
||||
if action == "modify":
|
||||
sanitized_content: Final = sanitization_result.get("content", "")
|
||||
|
|
|
|||
|
|
@ -12,6 +12,14 @@ 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.",
|
||||
)
|
||||
block_on_file_modify: bool = Field(
|
||||
default=True,
|
||||
description="Whether a file sanitization `modify` verdict blocks the request instead of replacing the file content.",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
|
|
|
|||
|
|
@ -1,16 +1,16 @@
|
|||
from fastapi.exceptions import HTTPException
|
||||
from unittest.mock import patch, AsyncMock
|
||||
from httpx import Response, Request
|
||||
import asyncio
|
||||
import base64
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.guardrails.guardrail_hooks.prompt_security.prompt_security import (
|
||||
PromptSecurityGuardrailMissingSecrets,
|
||||
PromptSecurityGuardrail,
|
||||
)
|
||||
from fastapi.exceptions import HTTPException
|
||||
from httpx import ReadTimeout, Request, Response
|
||||
|
||||
import litellm
|
||||
from litellm.proxy.guardrails.guardrail_hooks.prompt_security.prompt_security import (
|
||||
PromptSecurityGuardrail,
|
||||
PromptSecurityGuardrailMissingSecrets,
|
||||
)
|
||||
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
||||
|
||||
|
||||
|
|
@ -30,6 +30,8 @@ def test_prompt_security_guard_config(monkeypatch: pytest.MonkeyPatch):
|
|||
"guardrail": "prompt_security",
|
||||
"mode": "during_call",
|
||||
"default_on": True,
|
||||
"file_sanitization_fail_open": False,
|
||||
"block_on_file_modify": False,
|
||||
},
|
||||
}
|
||||
],
|
||||
|
|
@ -41,6 +43,12 @@ 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
|
||||
assert registered[0].block_on_file_modify is False
|
||||
config_model = registered[0].get_config_model()
|
||||
assert config_model is not None
|
||||
assert config_model().file_sanitization_fail_open is True
|
||||
assert config_model().block_on_file_modify is True
|
||||
|
||||
|
||||
def test_prompt_security_guard_config_no_api_key(monkeypatch: pytest.MonkeyPatch):
|
||||
|
|
@ -374,6 +382,201 @@ async def test_file_sanitization(monkeypatch: pytest.MonkeyPatch):
|
|||
assert result is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_file_sanitization_modify_blocks_by_default(monkeypatch: pytest.MonkeyPatch):
|
||||
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
|
||||
)
|
||||
csv_data = b"name,email\nAlice,alice@example.com\n"
|
||||
item = {
|
||||
"type": "file",
|
||||
"file": {
|
||||
"data": base64.b64encode(csv_data).decode(),
|
||||
"mime_type": "text/csv",
|
||||
},
|
||||
}
|
||||
upload_response = Response(
|
||||
json={"jobId": "modify-job"},
|
||||
status_code=200,
|
||||
request=Request(method="POST", url="https://test.prompt.security/api/sanitizeFile"),
|
||||
)
|
||||
poll_response = Response(
|
||||
json={
|
||||
"status": "done",
|
||||
"content": "name,email\nAlice,[REDACTED]\n",
|
||||
"metadata": {"action": "modify", "violations": ["Sensitive Data"]},
|
||||
},
|
||||
status_code=200,
|
||||
request=Request(method="GET", url="https://test.prompt.security/api/sanitizeFile"),
|
||||
)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=upload_response)):
|
||||
with patch.object(guardrail.async_handler, "get", AsyncMock(return_value=poll_response)):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail._process_document_item(item, None)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.detail == "Document blocked by Prompt Security. Violations: Sensitive Data"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_standalone_image_sanitization_modify_blocks_by_default(monkeypatch: pytest.MonkeyPatch):
|
||||
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.poll_interval = 0
|
||||
image_url = "data:image/png;base64," + base64.b64encode(b"image-content").decode()
|
||||
upload_response = Response(
|
||||
json={"jobId": "modify-image-job"},
|
||||
status_code=200,
|
||||
request=Request(method="POST", url="https://test.prompt.security/api/sanitizeFile"),
|
||||
)
|
||||
poll_response = Response(
|
||||
json={
|
||||
"status": "done",
|
||||
"content": "Email: [REDACTED]",
|
||||
"metadata": {"action": "modify", "violations": ["Sensitive Data"]},
|
||||
},
|
||||
status_code=200,
|
||||
request=Request(method="GET", url="https://test.prompt.security/api/sanitizeFile"),
|
||||
)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=upload_response)):
|
||||
with patch.object(guardrail.async_handler, "get", AsyncMock(return_value=poll_response)):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail._process_standalone_images([image_url], None)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.detail == "Image blocked by Prompt Security. Violations: Sensitive Data"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_file_sanitization_modify_can_rewrite_when_blocking_disabled(monkeypatch: pytest.MonkeyPatch):
|
||||
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,
|
||||
block_on_file_modify=False,
|
||||
)
|
||||
csv_data = b"name,email\nAlice,alice@example.com\n"
|
||||
item = {
|
||||
"type": "file",
|
||||
"file": {
|
||||
"data": base64.b64encode(csv_data).decode(),
|
||||
"mime_type": "text/csv",
|
||||
},
|
||||
}
|
||||
upload_response = Response(
|
||||
json={"jobId": "modify-job"},
|
||||
status_code=200,
|
||||
request=Request(method="POST", url="https://test.prompt.security/api/sanitizeFile"),
|
||||
)
|
||||
poll_response = Response(
|
||||
json={
|
||||
"status": "done",
|
||||
"content": "name,email\nAlice,[REDACTED]\n",
|
||||
"metadata": {"action": "modify", "violations": ["Sensitive Data"]},
|
||||
},
|
||||
status_code=200,
|
||||
request=Request(method="GET", url="https://test.prompt.security/api/sanitizeFile"),
|
||||
)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=upload_response)):
|
||||
with patch.object(guardrail.async_handler, "get", AsyncMock(return_value=poll_response)):
|
||||
result = await guardrail._process_document_item(item, None)
|
||||
|
||||
assert base64.b64decode(result["file"]["data"]) == b"name,email\nAlice,[REDACTED]\n"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"timeout",
|
||||
(
|
||||
litellm.Timeout(
|
||||
message="Prompt Security upload timed out",
|
||||
model="default-model-name",
|
||||
llm_provider="litellm-httpx-handler",
|
||||
),
|
||||
ReadTimeout(
|
||||
"Prompt Security poll timed out",
|
||||
request=Request(method="GET", url="https://test.prompt.security/api/sanitizeFile"),
|
||||
),
|
||||
),
|
||||
ids=("litellm", "httpx"),
|
||||
)
|
||||
@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,
|
||||
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 == {
|
||||
"action": "allow",
|
||||
"content": None,
|
||||
"metadata": {},
|
||||
"violations": (),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@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")
|
||||
|
||||
guardrail = PromptSecurityGuardrail(
|
||||
guardrail_name="test-guard",
|
||||
event_hook="pre_call",
|
||||
default_on=True,
|
||||
file_sanitization_timeout=0.01,
|
||||
file_sanitization_fail_open=fail_open,
|
||||
)
|
||||
|
||||
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"
|
||||
assert result["content"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_file_sanitization_block(monkeypatch: pytest.MonkeyPatch):
|
||||
"""Test that file sanitization blocks malicious files"""
|
||||
|
|
@ -544,7 +747,7 @@ async def test_role_filtering(monkeypatch: pytest.MonkeyPatch):
|
|||
return mock_response
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=mock_post):
|
||||
result = await guardrail.apply_guardrail(
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue