diff --git a/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py b/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py index 0954fe1698a..2a02560f6cc 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py +++ b/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py @@ -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") diff --git a/tests/test_litellm/proxy/guardrails/test_prompt_security_guardrails.py b/tests/test_litellm/proxy/guardrails/test_prompt_security_guardrails.py index e650f796f29..2e44b4b91b8 100644 --- a/tests/test_litellm/proxy/guardrails/test_prompt_security_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_prompt_security_guardrails.py @@ -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",