mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(guardrails): scan moderation input and Responses tool outputs, keep Azure Prompt Shield within its request limit
A call type with no guardrail translation handler (/v1/moderations) now has its plain input and prompt strings scanned by both the heuristics and the LLM check. The shared Responses translation handler reads and patches function_call_output and custom_tool_call_output items, so tool outputs reach every guardrail on that route. Azure Prompt Shield sends a userPrompt chunk and its document batch as two requests when together they would pass the 10,000 character limit, and attackDetected is required on every analysis so a verdict-less reply fails closed
This commit is contained in:
parent
f241e60497
commit
5a878322fd
7 changed files with 227 additions and 38 deletions
|
|
@ -241,9 +241,24 @@ _PATCHABLE_ITEM_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
|
|||
{"function_call_output": "output", "custom_tool_call_output": "output", "message": "content"}
|
||||
)
|
||||
|
||||
_TOOL_OUTPUT_ITEM_TYPES: Final = frozenset({"function_call_output", "custom_tool_call_output"})
|
||||
|
||||
_EMPTY_RESPONSES_REQUEST: Final[ResponsesAPIOptionalRequestParams] = {}
|
||||
|
||||
|
||||
def _scanned_text_field(item: Mapping[str, object]) -> str:
|
||||
return "output" if item.get("type") in _TOOL_OUTPUT_ITEM_TYPES else "content"
|
||||
|
||||
|
||||
def _write_scanned_text(item: dict[str, Any], content_idx: int | None, guardrail_response: str) -> None:
|
||||
field: Final = _scanned_text_field(item)
|
||||
content: Final = item.get(field)
|
||||
if isinstance(content, str) and content_idx is None:
|
||||
item[field] = guardrail_response
|
||||
elif isinstance(content, list) and content_idx is not None and isinstance(content[content_idx], dict):
|
||||
content[content_idx]["text"] = guardrail_response
|
||||
|
||||
|
||||
def _item_rewrite_field(item: Mapping[str, object]) -> str | None:
|
||||
item_type: Final = item.get("type")
|
||||
if item_type is None:
|
||||
|
|
@ -634,7 +649,7 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
|
||||
Override this method to customize text/image extraction logic.
|
||||
"""
|
||||
content: Final = message.get("content", None)
|
||||
content: Final = message.get(_scanned_text_field(message))
|
||||
if content is None:
|
||||
return
|
||||
|
||||
|
|
@ -675,22 +690,8 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
|
||||
Override this method to customize how responses are applied.
|
||||
"""
|
||||
for guardrail_response, mapping in zip(responses, task_mappings):
|
||||
msg_idx = cast(int, mapping[0])
|
||||
content_idx_optional = cast(int | None, mapping[1])
|
||||
|
||||
content = messages[msg_idx].get("content", None)
|
||||
if content is None:
|
||||
continue
|
||||
|
||||
if isinstance(content, str) and content_idx_optional is None:
|
||||
# Replace string content with guardrail response
|
||||
messages[msg_idx]["content"] = guardrail_response
|
||||
|
||||
elif isinstance(content, list) and content_idx_optional is not None:
|
||||
# Replace specific text item in list content
|
||||
if isinstance(messages[msg_idx]["content"][content_idx_optional], dict):
|
||||
messages[msg_idx]["content"][content_idx_optional]["text"] = guardrail_response
|
||||
for guardrail_response, (msg_idx, content_idx) in zip(responses, task_mappings):
|
||||
_write_scanned_text(messages[msg_idx], content_idx, guardrail_response)
|
||||
|
||||
async def process_output_response(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -151,16 +151,21 @@ def _document_batches(documents: Sequence[str]) -> tuple[tuple[str, ...], ...]:
|
|||
return reduce(_with_piece, pieces, empty)
|
||||
|
||||
|
||||
def _paired_requests(chunk: str | None, batch: tuple[str, ...] | None) -> tuple[_ShieldRequest, ...]:
|
||||
documents: Final = batch or ()
|
||||
if len(chunk or "") + sum(map(len, documents)) <= AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH:
|
||||
return (_ShieldRequest(user_prompt=chunk, documents=documents),)
|
||||
return (_ShieldRequest(user_prompt=chunk, documents=()), _ShieldRequest(user_prompt=None, documents=documents))
|
||||
|
||||
|
||||
def _shield_requests(user_prompt: str, documents: Sequence[str]) -> tuple[_ShieldRequest, ...]:
|
||||
prompt_chunks: Final = (
|
||||
tuple(AzureGuardrailBase.split_text_by_words(user_prompt, AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH))
|
||||
if user_prompt
|
||||
else ()
|
||||
)
|
||||
return tuple(
|
||||
_ShieldRequest(user_prompt=chunk, documents=batch or ())
|
||||
for chunk, batch in zip_longest(prompt_chunks, _document_batches(documents), fillvalue=None)
|
||||
)
|
||||
pairs: Final = zip_longest(prompt_chunks, _document_batches(documents), fillvalue=None)
|
||||
return tuple(chain.from_iterable(_paired_requests(chunk, batch) for chunk, batch in pairs))
|
||||
|
||||
|
||||
def _add_usage(usage_accumulator: MutableMapping[str, int], texts: Sequence[str]) -> None: # mutable-ok: accumulator
|
||||
|
|
|
|||
|
|
@ -2,10 +2,11 @@ import asyncio
|
|||
from collections.abc import Iterable, Iterator, Mapping, Sequence
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from difflib import SequenceMatcher
|
||||
from itertools import chain
|
||||
from typing import ClassVar, Final, Literal
|
||||
|
||||
from fastapi import HTTPException
|
||||
from pydantic import TypeAdapter
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
|
|
@ -29,6 +30,8 @@ from litellm.types.utils import CallTypes, CallTypesLiteral, GenericGuardrailAPI
|
|||
GUARDRAIL_NAME: Final = "detect_prompt_injection"
|
||||
REJECTION_MESSAGE: Final = "Rejected message. This is a prompt injection attack."
|
||||
SCANNED_REQUEST: Final = TypeAdapter(dict[str, object])
|
||||
REQUEST_ITEMS: Final = TypeAdapter(tuple[object, ...])
|
||||
PLAIN_TEXT_REQUEST_FIELDS: Final = ("input", "prompt")
|
||||
HEURISTICS_EXECUTOR: Final = ThreadPoolExecutor(
|
||||
max_workers=PROMPT_INJECTION_HEURISTICS_MAX_THREADS, thread_name_prefix="prompt-injection-heuristics"
|
||||
)
|
||||
|
|
@ -92,6 +95,20 @@ def _attachment_texts(request_data: Mapping[str, object]) -> tuple[str, ...]:
|
|||
return tuple(attachment.text for attachment in request_attachments(request_data).texts)
|
||||
|
||||
|
||||
def _strings(value: object) -> tuple[str, ...]:
|
||||
if isinstance(value, str):
|
||||
return (value,)
|
||||
try:
|
||||
items: Final = REQUEST_ITEMS.validate_python(value)
|
||||
except ValidationError:
|
||||
return ()
|
||||
return tuple(item for item in items if isinstance(item, str))
|
||||
|
||||
|
||||
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__(
|
||||
|
|
@ -222,16 +239,14 @@ class _OPTIONAL_PromptInjectionDetection(CustomGuardrail):
|
|||
raise _rejection()
|
||||
|
||||
async def _scan_request(self, data: dict[str, object], call_type: str) -> dict[str, object]:
|
||||
handler: Final = _translation_handler(call_type)
|
||||
if handler is None:
|
||||
verbose_proxy_logger.debug(
|
||||
"Prompt injection detection has no translation handler for %s; skipping", call_type
|
||||
)
|
||||
return data
|
||||
attachments: Final = request_attachments(data)
|
||||
if attachments.unscannable and not self._skips_unscannable_attachments():
|
||||
raise _unscannable_rejection(attachments.unscannable)
|
||||
await self._reject_injected_texts(attachment.text for attachment in attachments.texts)
|
||||
handler: Final = _translation_handler(call_type)
|
||||
if handler is None:
|
||||
await self._reject_injected_texts(_plain_request_texts(data))
|
||||
return data
|
||||
return SCANNED_REQUEST.validate_python(
|
||||
await handler.process_input_messages(
|
||||
data=data, guardrail_to_apply=self, litellm_logging_obj=_logging_obj(data)
|
||||
|
|
@ -270,6 +285,19 @@ class _OPTIONAL_PromptInjectionDetection(CustomGuardrail):
|
|||
await self._reject_injected_texts(inputs.get("texts", ()))
|
||||
return inputs
|
||||
|
||||
async def _judge_request(self, judge: _PromptInjectionLLMJudge, data: dict[str, object], call_type: str) -> None:
|
||||
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)
|
||||
)
|
||||
return
|
||||
await judge.apply_guardrail(
|
||||
inputs=GenericGuardrailAPIInputs(texts=list(_plain_request_texts(data))),
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
async def async_moderation_hook(
|
||||
self,
|
||||
data: dict[str, object],
|
||||
|
|
@ -279,16 +307,8 @@ class _OPTIONAL_PromptInjectionDetection(CustomGuardrail):
|
|||
judge: Final = self.llm_judge
|
||||
if judge is None:
|
||||
return
|
||||
handler: Final = _translation_handler(call_type)
|
||||
if handler is None:
|
||||
verbose_proxy_logger.debug(
|
||||
"Prompt injection LLM check has no translation handler for %s; skipping", call_type
|
||||
)
|
||||
return
|
||||
try:
|
||||
await handler.process_input_messages(
|
||||
data=data, guardrail_to_apply=judge, litellm_logging_obj=_logging_obj(data)
|
||||
)
|
||||
await self._judge_request(judge, data, call_type)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ class AzurePromptShieldGuardrailRequestBody(TypedDict):
|
|||
class AzurePromptShieldAnalysis(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
attackDetected: bool = False
|
||||
attackDetected: bool
|
||||
|
||||
|
||||
class AzurePromptShieldGuardrailResponse(BaseModel):
|
||||
|
|
|
|||
|
|
@ -925,6 +925,27 @@ async def test_pre_call_hook_splits_an_oversized_document_by_words():
|
|||
assert set(piece.split()) == {"word"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_hook_keeps_prompt_and_documents_within_the_combined_azure_limit():
|
||||
prompt = "ask " * (AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH * 9 // 40)
|
||||
tool_output = "fact " * (AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH * 3 // 50)
|
||||
assert len(prompt) <= AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH < len(prompt) + len(tool_output)
|
||||
bodies, _ = await _run_pre_call_hook(
|
||||
_shield_guardrail(),
|
||||
[
|
||||
{"role": "user", "content": prompt},
|
||||
_tool_call_message(),
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": tool_output},
|
||||
],
|
||||
)
|
||||
|
||||
for body in bodies:
|
||||
submitted = len(body.get("userPrompt", "")) + sum(len(document) for document in body["documents"])
|
||||
assert submitted <= AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH
|
||||
assert "".join(body.get("userPrompt", "") for body in bodies) == prompt
|
||||
assert [document for body in bodies for document in body["documents"]] == [tool_output]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_billing_counts_document_characters_and_text_records():
|
||||
import math as _math
|
||||
|
|
@ -1000,6 +1021,21 @@ async def test_pre_call_hook_fails_closed_when_a_submitted_document_is_not_analy
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_hook_fails_closed_when_an_analysis_carries_no_verdict():
|
||||
guardrail = _shield_guardrail()
|
||||
response = Mock()
|
||||
response.json.return_value = {"userPromptAnalysis": {}, "documentsAnalysis": []}
|
||||
with patch.object(guardrail.async_handler, "post", return_value=response):
|
||||
with pytest.raises(ValueError, match="attackDetected"):
|
||||
await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="k"),
|
||||
cache=None,
|
||||
data={"messages": [{"role": "user", "content": "Summarize my email."}]},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_sends_request_attachments_as_documents():
|
||||
guardrail = _shield_guardrail()
|
||||
|
|
|
|||
|
|
@ -381,6 +381,83 @@ async def test_llm_check_scans_text_attachments(monkeypatch: pytest.MonkeyPatch)
|
|||
assert router.seen_prompts == (f"{SAFE}\nattached text",)
|
||||
|
||||
|
||||
def _moderation(text: str | list[str]) -> dict[str, object]:
|
||||
return {"model": "omni-moderation-latest", "input": text}
|
||||
|
||||
|
||||
def _responses_with_tool_output(text: str) -> dict[str, object]:
|
||||
return {
|
||||
"model": "test-model",
|
||||
"input": [
|
||||
{"role": "user", "content": [{"type": "input_text", "text": SAFE}]},
|
||||
{"type": "function_call", "call_id": "call_1", "name": "read_email", "arguments": "{}"},
|
||||
{"type": "function_call_output", "call_id": "call_1", "output": text},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("moderation_input", [INJECTION, [SAFE, INJECTION]])
|
||||
async def test_moderation_requests_reject_prompt_injection(
|
||||
monkeypatch: pytest.MonkeyPatch, moderation_input: str | list[str]
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _proxy_pre_call(
|
||||
monkeypatch, _OPTIONAL_PromptInjectionDetection(), _moderation(moderation_input), "moderation"
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert _error(exc_info.value)["error"] == REJECTION_MESSAGE
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_moderation_requests_allow_a_safe_input(monkeypatch: pytest.MonkeyPatch):
|
||||
data = _moderation(SAFE)
|
||||
|
||||
assert await _proxy_pre_call(monkeypatch, _OPTIONAL_PromptInjectionDetection(), data, "moderation") == data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_llm_check_judges_moderation_input(monkeypatch: pytest.MonkeyPatch):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _proxy_during_call(monkeypatch, _moderation_detector(verdict="UNSAFE"), _moderation(SAFE), "moderation")
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert _error(exc_info.value)["error"] == REJECTION_MESSAGE
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_tool_outputs_are_scanned(monkeypatch: pytest.MonkeyPatch):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _proxy_pre_call(
|
||||
monkeypatch, _OPTIONAL_PromptInjectionDetection(), _responses_with_tool_output(INJECTION), "aresponses"
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert _error(exc_info.value)["error"] == REJECTION_MESSAGE
|
||||
|
||||
|
||||
@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"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
await _proxy_during_call(
|
||||
monkeypatch,
|
||||
_moderation_detector(verdict="SAFE", router=router),
|
||||
_responses_with_tool_output("tool output text"),
|
||||
"aresponses",
|
||||
)
|
||||
|
||||
assert router.seen_prompts == (f"{SAFE}\ntool output text",)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_heuristics_check_keeps_event_loop_responsive():
|
||||
detector = _OPTIONAL_PromptInjectionDetection(
|
||||
|
|
|
|||
|
|
@ -249,6 +249,56 @@ class TestOpenAIResponsesHandlerInputProcessing:
|
|||
assert result["input"][1]["content"] == " [GUARDRAILED]"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("call_item", "output_item"),
|
||||
[
|
||||
(
|
||||
{"type": "function_call", "call_id": "call_1", "name": "read_email", "arguments": "{}"},
|
||||
{"type": "function_call_output", "call_id": "call_1", "output": "memo memo"},
|
||||
),
|
||||
(
|
||||
{"type": "custom_tool_call", "call_id": "call_1", "name": "run_script", "input": "ls"},
|
||||
{"type": "custom_tool_call_output", "call_id": "call_1", "output": "memo memo"},
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_process_input_scans_tool_outputs(self, call_item, output_item):
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = MockGuardrail(guardrail_name="test")
|
||||
|
||||
data = {
|
||||
"input": [{"role": "user", "content": "Read my email", "type": "message"}, call_item, output_item],
|
||||
"model": "gpt-4",
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert result["input"][0]["content"] == "Read my email [GUARDRAILED]"
|
||||
assert result["input"][2]["output"] == "memo memo [GUARDRAILED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_input_scans_a_custom_tool_output_content_list(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = MockGuardrail(guardrail_name="test")
|
||||
|
||||
data = {
|
||||
"input": [
|
||||
{"type": "custom_tool_call", "call_id": "call_1", "name": "run_script", "input": "ls"},
|
||||
{
|
||||
"type": "custom_tool_call_output",
|
||||
"call_id": "call_1",
|
||||
"output": [{"type": "input_text", "text": "memo memo"}],
|
||||
},
|
||||
],
|
||||
"model": "gpt-4",
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert result["input"][1]["output"][0]["text"] == "memo memo [GUARDRAILED]"
|
||||
|
||||
|
||||
class TestOpenAIResponsesHandlerOutputProcessing:
|
||||
"""Test output processing functionality"""
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue