From 888d7ccd83f6c69b96dcfa5f622a699700b03db2 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 18:01:18 -0700 Subject: [PATCH] fix(guardrails): judge attachment text and the prompt in one LLM call --- .../proxy/hooks/prompt_injection_detection.py | 67 ++++++++++++------- .../hooks/test_prompt_injection_detection.py | 55 ++++++++------- 2 files changed, 74 insertions(+), 48 deletions(-) diff --git a/litellm/proxy/hooks/prompt_injection_detection.py b/litellm/proxy/hooks/prompt_injection_detection.py index 23290f62b68..09c4b69583c 100644 --- a/litellm/proxy/hooks/prompt_injection_detection.py +++ b/litellm/proxy/hooks/prompt_injection_detection.py @@ -1,6 +1,7 @@ import asyncio from collections.abc import Iterable, Iterator, Mapping, Sequence from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass from difflib import SequenceMatcher from itertools import chain from typing import ClassVar, Final, Literal @@ -109,28 +110,11 @@ def _plain_request_texts(request_data: Mapping[str, object]) -> tuple[str, ...]: return tuple(chain.from_iterable(_strings(request_data.get(field)) for field in PLAIN_TEXT_REQUEST_FIELDS)) -class _PromptInjectionLLMJudge(CustomGuardrail): - def __init__(self, params: LiteLLMPromptInjectionParams, llm_api_name: str, router: Router) -> None: - super().__init__( - guardrail_name=GUARDRAIL_NAME, - supported_event_hooks=[GuardrailEventHooks.during_call], - event_hook=[GuardrailEventHooks.during_call], - default_on=True, - ) - self.params = params - self.llm_api_name = llm_api_name - self.router = router - - async def apply_guardrail( - self, - inputs: GenericGuardrailAPIInputs, - request_data: dict[str, object], - input_type: Literal["request", "response"], - logging_obj: LiteLLMLoggingObj | None = None, - ) -> GenericGuardrailAPIInputs: - if input_type == "request": - await self.reject_injection(inputs.get("texts", ())) - return inputs +@dataclass(frozen=True, slots=True) +class _PromptInjectionLLMJudge: + params: LiteLLMPromptInjectionParams + llm_api_name: str + router: Router async def reject_injection(self, texts: Iterable[str]) -> None: prompt: Final = "\n".join(texts) @@ -157,6 +141,38 @@ class _PromptInjectionLLMJudge(CustomGuardrail): return isinstance(content, str) and fail_call_string in content +class _RequestJudge(CustomGuardrail): + def __init__(self, judge: _PromptInjectionLLMJudge, attachment_texts: tuple[str, ...]) -> None: + super().__init__( + guardrail_name=GUARDRAIL_NAME, + supported_event_hooks=[GuardrailEventHooks.during_call], + event_hook=[GuardrailEventHooks.during_call], + default_on=True, + ) + self.judge = judge + self.attachment_texts = attachment_texts + self.judged = False + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + logging_obj: LiteLLMLoggingObj | None = None, + ) -> GenericGuardrailAPIInputs: + if input_type == "request": + await self.judge_with_attachments(inputs.get("texts", ())) + return inputs + + async def judge_with_attachments(self, texts: Iterable[str]) -> None: + self.judged = True + await self.judge.reject_injection(chain(self.attachment_texts, texts)) + + async def judge_attachments_unless_judged(self) -> None: + if not self.judged: + await self.judge.reject_injection(self.attachment_texts) + + class _OPTIONAL_PromptInjectionDetection(CustomGuardrail): use_native_lifecycle_hooks: ClassVar[bool] = True enforces_request_content: bool = True @@ -288,14 +304,15 @@ 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)) + request_judge: Final = _RequestJudge(judge, _attachment_texts(data)) handler: Final = _translation_handler(call_type) if handler is None: - await judge.reject_injection(_plain_request_texts(data)) + await request_judge.judge_with_attachments(_plain_request_texts(data)) return await handler.process_input_messages( - data=data, guardrail_to_apply=judge, litellm_logging_obj=_logging_obj(data) + data=data, guardrail_to_apply=request_judge, litellm_logging_obj=_logging_obj(data) ) + await request_judge.judge_attachments_unless_judged() async def async_moderation_hook( self, 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 29f4dd3c93f..9fcd32dd3e6 100644 --- a/tests/test_litellm/proxy/hooks/test_prompt_injection_detection.py +++ b/tests/test_litellm/proxy/hooks/test_prompt_injection_detection.py @@ -203,6 +203,17 @@ class _RecordingRouter(Router): return await super().acompletion(model=model, messages=messages, stream=stream) +def _recording_router(verdict: str) -> _RecordingRouter: + return _RecordingRouter( + model_list=[ + { + "model_name": "moderation-model", + "litellm_params": {"model": "openai/gpt-5.6", "api_key": "sk-fake", "mock_response": verdict}, + } + ] + ) + + @pytest.mark.asyncio @pytest.mark.parametrize("call_type", sorted(REQUEST_BY_CALL_TYPE)) async def test_every_unified_call_type_rejects_prompt_injection( @@ -371,24 +382,24 @@ async def test_moderation_hook_skips_llm_check_without_prompt_text(): @pytest.mark.asyncio -async def test_llm_check_scans_text_attachments(monkeypatch: pytest.MonkeyPatch): - router = _RecordingRouter( - model_list=[ - { - "model_name": "moderation-model", - "litellm_params": {"model": "openai/gpt-5.6", "api_key": "sk-fake", "mock_response": "SAFE"}, - } - ] - ) +@pytest.mark.parametrize( + ("call_type", "build"), + [ + ("acompletion", _chat_with_text_and_file), + ("aresponses", _responses_with_text_and_input_file), + ("anthropic_messages", _messages_with_text_and_document), + ], +) +async def test_llm_check_judges_text_attachments_and_prompt_in_one_call( + monkeypatch: pytest.MonkeyPatch, call_type: CallTypesLiteral, build: RequestBuilder +): + router = _recording_router(verdict="SAFE") await _proxy_during_call( - monkeypatch, - _moderation_detector(verdict="SAFE", router=router), - _chat_with_text_and_file("attached text"), - "acompletion", + monkeypatch, _moderation_detector(verdict="SAFE", router=router), build("attached text"), call_type ) - assert router.seen_prompts == ("attached text", SAFE) + assert router.seen_prompts == (f"attached text\n{SAFE}",) @pytest.mark.asyncio @@ -403,11 +414,16 @@ async def test_llm_check_scans_text_attachments(monkeypatch: pytest.MonkeyPatch) async def test_llm_check_judges_attachment_only_input( monkeypatch: pytest.MonkeyPatch, call_type: CallTypesLiteral, build: RequestBuilder ): + router = _recording_router(verdict="UNSAFE") + with pytest.raises(HTTPException) as exc_info: - await _proxy_during_call(monkeypatch, _moderation_detector(verdict="UNSAFE"), build(INJECTION), call_type) + await _proxy_during_call( + monkeypatch, _moderation_detector(verdict="UNSAFE", router=router), build(INJECTION), call_type + ) assert exc_info.value.status_code == 400 assert _error(exc_info.value)["error"] == REJECTION_MESSAGE + assert router.seen_prompts == (INJECTION,) def _moderation(text: str | list[str]) -> dict[str, object]: @@ -468,14 +484,7 @@ async def test_responses_tool_outputs_are_scanned(monkeypatch: pytest.MonkeyPatc @pytest.mark.asyncio async def test_llm_check_judges_responses_tool_outputs(monkeypatch: pytest.MonkeyPatch): - router = _RecordingRouter( - model_list=[ - { - "model_name": "moderation-model", - "litellm_params": {"model": "openai/gpt-5.6", "api_key": "sk-fake", "mock_response": "SAFE"}, - } - ] - ) + router = _recording_router(verdict="SAFE") await _proxy_during_call( monkeypatch,