mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): judge attachment text and the prompt in one LLM call
This commit is contained in:
parent
8c11fd6a84
commit
888d7ccd83
2 changed files with 74 additions and 48 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue