From 5a878322fd741b8ff2eed24338957d813d0a1513 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 15:35:03 -0700 Subject: [PATCH] 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 --- .../guardrail_translation/handler.py | 35 +++++---- .../guardrail_hooks/azure/prompt_shield.py | 13 +++- .../proxy/hooks/prompt_injection_detection.py | 52 +++++++++---- .../azure/azure_prompt_shield.py | 2 +- .../azure/test_azure_prompt_shield.py | 36 +++++++++ .../hooks/test_prompt_injection_detection.py | 77 +++++++++++++++++++ ...test_openai_responses_guardrail_handler.py | 50 ++++++++++++ 7 files changed, 227 insertions(+), 38 deletions(-) diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 66cebe0175d..9f79f3affa5 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py index 92130eaa8da..012431294a8 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py @@ -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 diff --git a/litellm/proxy/hooks/prompt_injection_detection.py b/litellm/proxy/hooks/prompt_injection_detection.py index bc34ac9c713..4f903c8a4cf 100644 --- a/litellm/proxy/hooks/prompt_injection_detection.py +++ b/litellm/proxy/hooks/prompt_injection_detection.py @@ -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: diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/azure/azure_prompt_shield.py b/litellm/types/proxy/guardrails/guardrail_hooks/azure/azure_prompt_shield.py index 830d75c375a..f88ae063ecf 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/azure/azure_prompt_shield.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/azure/azure_prompt_shield.py @@ -19,7 +19,7 @@ class AzurePromptShieldGuardrailRequestBody(TypedDict): class AzurePromptShieldAnalysis(BaseModel): model_config = ConfigDict(frozen=True) - attackDetected: bool = False + attackDetected: bool class AzurePromptShieldGuardrailResponse(BaseModel): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/azure/test_azure_prompt_shield.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/azure/test_azure_prompt_shield.py index bab11291224..fcf857f577d 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/azure/test_azure_prompt_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/azure/test_azure_prompt_shield.py @@ -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() 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 07c3a17f9ca..9b52a0103a0 100644 --- a/tests/test_litellm/proxy/hooks/test_prompt_injection_detection.py +++ b/tests/test_litellm/proxy/hooks/test_prompt_injection_detection.py @@ -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( diff --git a/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py index a6b930db7a9..54596f081f3 100644 --- a/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py +++ b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py @@ -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"""