mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
Merge pull request #41131 from BerriAI/litellm_prompt_security_created_status
fix(prompt_security): keep polling file sanitization through non-terminal statuses
This commit is contained in:
commit
99245f9323
2 changed files with 103 additions and 8 deletions
|
|
@ -27,6 +27,7 @@ if TYPE_CHECKING:
|
|||
|
||||
|
||||
_SANITIZE_FILE_FAIL_OPEN_TIMEOUT_SECONDS: Final = 30.0
|
||||
_SANITIZE_FILE_QUEUED_STATUSES: Final = frozenset({"created", "in progress"})
|
||||
|
||||
|
||||
class PromptSecurityGuardrailMissingSecrets(Exception):
|
||||
|
|
@ -512,16 +513,18 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
"metadata": result.get("metadata", {}),
|
||||
"violations": result.get("metadata", {}).get("violations", []),
|
||||
}
|
||||
elif status == "in progress":
|
||||
verbose_proxy_logger.debug(
|
||||
"Prompt Security Guardrail: File sanitization in progress (attempt %d/%d)",
|
||||
attempt + 1,
|
||||
self.max_poll_attempts,
|
||||
)
|
||||
continue
|
||||
else:
|
||||
|
||||
if status not in _SANITIZE_FILE_QUEUED_STATUSES:
|
||||
raise HTTPException(status_code=500, detail=f"Unexpected sanitization status: {status}")
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"Prompt Security Guardrail: File sanitization status=%s for jobId=%s (attempt %d/%d)",
|
||||
status,
|
||||
job_id,
|
||||
attempt + 1,
|
||||
self.max_poll_attempts,
|
||||
)
|
||||
|
||||
raise HTTPException(status_code=408, detail="File sanitization timeout")
|
||||
|
||||
def _raise_if_file_blocked(self, sanitization_result: _SanitizeResult, resource_name: str) -> None:
|
||||
|
|
|
|||
|
|
@ -497,6 +497,98 @@ async def test_file_sanitization_modify_can_rewrite_when_blocking_disabled(monke
|
|||
assert base64.b64decode(result["file"]["data"]) == b"name,email\nAlice,[REDACTED]\n"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_file_sanitization_keeps_polling_through_queued_statuses(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
|
||||
upload_response = Response(
|
||||
json={"jobId": "queued-job"},
|
||||
status_code=200,
|
||||
request=Request(method="POST", url="https://test.prompt.security/api/sanitizeFile"),
|
||||
)
|
||||
poll_request = Request(method="GET", url="https://test.prompt.security/api/sanitizeFile")
|
||||
poll_responses = [
|
||||
Response(json={"status": "created"}, status_code=200, request=poll_request),
|
||||
Response(json={"status": "in progress"}, status_code=200, request=poll_request),
|
||||
Response(
|
||||
json={"status": "done", "content": "clean", "metadata": {"action": "allow", "violations": []}},
|
||||
status_code=200,
|
||||
request=poll_request,
|
||||
),
|
||||
]
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=upload_response)):
|
||||
with patch.object(guardrail.async_handler, "get", AsyncMock(side_effect=poll_responses)) as poll_mock:
|
||||
result = await guardrail.sanitize_file_content(b"image-content", "image.png")
|
||||
|
||||
assert poll_mock.await_count == 3
|
||||
assert result["action"] == "allow"
|
||||
assert result["content"] == "clean"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_file_sanitization_never_finishing_job_times_out(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, file_sanitization_fail_open=False
|
||||
)
|
||||
guardrail.poll_interval = 0
|
||||
guardrail.max_poll_attempts = 3
|
||||
upload_response = Response(
|
||||
json={"jobId": "stuck-job"},
|
||||
status_code=200,
|
||||
request=Request(method="POST", url="https://test.prompt.security/api/sanitizeFile"),
|
||||
)
|
||||
poll_response = Response(
|
||||
json={"status": "created"},
|
||||
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)) as poll_mock:
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.sanitize_file_content(b"file-content", "document.pdf")
|
||||
|
||||
assert poll_mock.await_count == 3
|
||||
assert exc_info.value.status_code == 408
|
||||
assert exc_info.value.detail == "File sanitization timeout"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("poll_body", [{"status": "failed"}, {}])
|
||||
async def test_file_sanitization_terminal_failure_does_not_fail_open(monkeypatch: pytest.MonkeyPatch, poll_body):
|
||||
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
|
||||
upload_response = Response(
|
||||
json={"jobId": "failed-job"},
|
||||
status_code=200,
|
||||
request=Request(method="POST", url="https://test.prompt.security/api/sanitizeFile"),
|
||||
)
|
||||
poll_response = Response(
|
||||
json=poll_body,
|
||||
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)) as poll_mock:
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.sanitize_file_content(b"file-content", "document.pdf")
|
||||
|
||||
assert poll_mock.await_count == 1
|
||||
assert exc_info.value.status_code == 500
|
||||
assert exc_info.value.detail == f"Unexpected sanitization status: {poll_body.get('status')}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"timeout",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue