fix(guardrails): judge attachment text and the prompt in one LLM call

This commit is contained in:
mateo-berri 2026-09-26 18:01:18 -07:00
parent 8c11fd6a84
commit 888d7ccd83
2 changed files with 74 additions and 48 deletions

View file

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

View file

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