diff --git a/litellm/proxy/hooks/prompt_injection_detection.py b/litellm/proxy/hooks/prompt_injection_detection.py index 4f903c8a4cf..23290f62b68 100644 --- a/litellm/proxy/hooks/prompt_injection_detection.py +++ b/litellm/proxy/hooks/prompt_injection_detection.py @@ -128,11 +128,14 @@ class _PromptInjectionLLMJudge(CustomGuardrail): input_type: Literal["request", "response"], logging_obj: LiteLLMLoggingObj | None = None, ) -> GenericGuardrailAPIInputs: - if input_type != "request": - return inputs - prompt: Final = "\n".join((*inputs.get("texts", ()), *_attachment_texts(request_data))) + if input_type == "request": + await self.reject_injection(inputs.get("texts", ())) + return inputs + + async def reject_injection(self, texts: Iterable[str]) -> None: + prompt: Final = "\n".join(texts) if not prompt.strip(): - return inputs + return response: Final[ModelResponse] = await self.router.acompletion( model=self.llm_api_name, messages=[ @@ -145,7 +148,6 @@ class _PromptInjectionLLMJudge(CustomGuardrail): ) if self._verdict_is_attack(response): raise _rejection() - return inputs def _verdict_is_attack(self, response: ModelResponse) -> bool: fail_call_string: Final = self.params.llm_api_fail_call_string @@ -286,16 +288,13 @@ class _OPTIONAL_PromptInjectionDetection(CustomGuardrail): return inputs async def _judge_request(self, judge: _PromptInjectionLLMJudge, data: dict[str, object], call_type: str) -> None: + await judge.reject_injection(_attachment_texts(data)) handler: Final = _translation_handler(call_type) - if handler is not None: - await handler.process_input_messages( - data=data, guardrail_to_apply=judge, litellm_logging_obj=_logging_obj(data) - ) + if handler is None: + await judge.reject_injection(_plain_request_texts(data)) return - await judge.apply_guardrail( - inputs=GenericGuardrailAPIInputs(texts=list(_plain_request_texts(data))), - request_data=data, - input_type="request", + await handler.process_input_messages( + data=data, guardrail_to_apply=judge, litellm_logging_obj=_logging_obj(data) ) async def async_moderation_hook( diff --git a/tests/test_litellm/proxy/hooks/test_prompt_injection_detection.py b/tests/test_litellm/proxy/hooks/test_prompt_injection_detection.py index 9b52a0103a0..29f4dd3c93f 100644 --- a/tests/test_litellm/proxy/hooks/test_prompt_injection_detection.py +++ b/tests/test_litellm/proxy/hooks/test_prompt_injection_detection.py @@ -91,18 +91,28 @@ def _chat_with_only_a_file(text: str) -> dict[str, object]: return _chat_with_parts(_file_part(text)) +def _input_file_part(text: str) -> dict[str, object]: + return {"type": "input_file", "filename": "notes.txt", "file_data": _text_data_url(text)} + + +def _document_part(text: str) -> dict[str, object]: + return {"type": "document", "source": {"type": "text", "media_type": "text/plain", "data": text}} + + def _responses_with_text_and_input_file(text: str) -> dict[str, object]: - return _responses_with_parts( - {"type": "input_text", "text": SAFE}, - {"type": "input_file", "filename": "notes.txt", "file_data": _text_data_url(text)}, - ) + return _responses_with_parts({"type": "input_text", "text": SAFE}, _input_file_part(text)) + + +def _responses_with_only_an_input_file(text: str) -> dict[str, object]: + return _responses_with_parts(_input_file_part(text)) def _messages_with_text_and_document(text: str) -> dict[str, object]: - return _messages_with_parts( - {"type": "text", "text": SAFE}, - {"type": "document", "source": {"type": "text", "media_type": "text/plain", "data": text}}, - ) + return _messages_with_parts({"type": "text", "text": SAFE}, _document_part(text)) + + +def _messages_with_only_a_document(text: str) -> dict[str, object]: + return _messages_with_parts(_document_part(text)) PDF_FILE_PART: Final[dict[str, object]] = { @@ -378,7 +388,26 @@ async def test_llm_check_scans_text_attachments(monkeypatch: pytest.MonkeyPatch) "acompletion", ) - assert router.seen_prompts == (f"{SAFE}\nattached text",) + assert router.seen_prompts == ("attached text", SAFE) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("call_type", "build"), + [ + ("acompletion", _chat_with_only_a_file), + ("aresponses", _responses_with_only_an_input_file), + ("anthropic_messages", _messages_with_only_a_document), + ], +) +async def test_llm_check_judges_attachment_only_input( + monkeypatch: pytest.MonkeyPatch, call_type: CallTypesLiteral, build: RequestBuilder +): + with pytest.raises(HTTPException) as exc_info: + await _proxy_during_call(monkeypatch, _moderation_detector(verdict="UNSAFE"), build(INJECTION), call_type) + + assert exc_info.value.status_code == 400 + assert _error(exc_info.value)["error"] == REJECTION_MESSAGE def _moderation(text: str | list[str]) -> dict[str, object]: