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:
yucheng-berri 2026-09-14 16:43:37 -07:00 committed by GitHub
commit 99245f9323
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 103 additions and 8 deletions

View file

@ -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:

View file

@ -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",