mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): judge attachment-only input with the prompt injection LLM check
This commit is contained in:
parent
5a878322fd
commit
8c11fd6a84
2 changed files with 50 additions and 22 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue