mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
fix(prompt_security): keep polling file sanitization through non-terminal statuses
Prompt Security reports a queued sanitization job as status "created" before it moves to "in progress" and "done". The poller treated anything other than those two known strings as an error and returned HTTP 500 on the first poll, so every image or file request through the guardrail failed while the vendor job was still queued. Only "done" is terminal now. Every other status is logged and polled again until max_poll_attempts or the outer file_sanitization_timeout, after which the existing fail-open or fail-closed (408) policy applies. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
4123b4bc2b
commit
fb60f80c80
2 changed files with 71 additions and 9 deletions
|
|
@ -512,15 +512,14 @@ 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:
|
||||
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")
|
||||
|
||||
|
|
|
|||
|
|
@ -497,6 +497,69 @@ 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(
|
||||
"timeout",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue