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:
mateo-berri 2026-09-26 15:35:03 -07:00
parent f241e60497
commit 5a878322fd
7 changed files with 227 additions and 38 deletions

View file

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

View file

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

View file

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

View file

@ -19,7 +19,7 @@ class AzurePromptShieldGuardrailRequestBody(TypedDict):
class AzurePromptShieldAnalysis(BaseModel):
model_config = ConfigDict(frozen=True)
attackDetected: bool = False
attackDetected: bool
class AzurePromptShieldGuardrailResponse(BaseModel):

View file

@ -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()

View file

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

View file

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