fix(guardrails): judge attachment-only input with the prompt injection LLM check

This commit is contained in:
mateo-berri 2026-09-26 17:44:21 -07:00
parent 5a878322fd
commit 8c11fd6a84
2 changed files with 50 additions and 22 deletions

View file

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

View file

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