fix(guardrails): pick Azure prompt source by call type so a messages stub cannot hide Responses input

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-30 19:35:23 +00:00
parent 6788afd267
commit 29fa9f8ac4
6 changed files with 337 additions and 52 deletions

View file

@ -12,6 +12,7 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import AllMessageValues, ResponseInputParam
from litellm.types.utils import CallTypes, CallTypesLiteral
# Azure Content Safety APIs have a 10,000 character limit per request.
AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH: Final = 10000
@ -23,6 +24,8 @@ AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH: Final = 1000
AZURE_CONTENT_SAFETY_DEFAULT_API_VERSION: Final = "2024-09-01"
JAVELIN_API_VERSION_STORED_BY_OLDER_RELEASES: Final = "v1"
_RESPONSES_API_CALL_TYPES: Final = frozenset({CallTypes.responses, CallTypes.aresponses})
def resolve_content_safety_api_version(configured: str | None) -> str:
if not configured or configured == JAVELIN_API_VERSION_STORED_BY_OLDER_RELEASES:
@ -131,15 +134,15 @@ class AzureGuardrailBase:
return chunks
def get_user_prompt_from_request(self, data: Mapping[str, object]) -> str | None:
def get_user_prompt_from_request(self, data: Mapping[str, object], call_type: CallTypesLiteral) -> str | None:
if call_type in _RESPONSES_API_CALL_TYPES:
responses_input: Final = data.get("input")
if not isinstance(responses_input, (str, list)):
return None
validated_input: Final = cast(ResponseInputParam, responses_input)
return get_last_user_message(ResponsesAPIRequestUtils.responses_input_to_chat_messages(validated_input))
messages: Final = data.get("messages")
if isinstance(messages, list):
return get_last_user_message(cast(list[AllMessageValues], messages))
responses_input: Final = data.get("input")
if not isinstance(responses_input, (str, list)):
if not isinstance(messages, list):
return None
validated_input: Final = cast(ResponseInputParam, responses_input)
chat_messages: Final = ResponsesAPIRequestUtils.responses_input_to_chat_messages(validated_input)
return get_last_user_message(chat_messages)
return get_last_user_message(cast(list[AllMessageValues], messages))

View file

@ -249,7 +249,7 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai
"Azure Prompt Shield: Running pre-call prompt scan, on call_type: %s",
call_type,
)
user_prompt: Final = self.get_user_prompt_from_request(data)
user_prompt: Final = self.get_user_prompt_from_request(data, call_type)
if user_prompt:
verbose_proxy_logger.debug("Azure Prompt Shield: User prompt: %s", user_prompt)

View file

@ -231,7 +231,7 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr
"Azure Text Moderation: Running pre-call prompt scan, on call_type: %s",
call_type,
)
user_prompt: Final = self.get_user_prompt_from_request(data)
user_prompt: Final = self.get_user_prompt_from_request(data, call_type)
if user_prompt:
verbose_proxy_logger.info("Azure Text Moderation: User prompt: %s", user_prompt)

View file

@ -0,0 +1,210 @@
import json
import uuid
from collections.abc import Iterator
from contextlib import ExitStack
from pathlib import Path
from typing import Final
import pytest
import yaml
from integration._support.client import Gateway, eventually, gateway_from_environment, object_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, Wire, wire_server
_ATTACK_MARKER: Final = "synthetic-attack-marker"
_AZURE_TARGET_PREFIX: Final = "/contentsafety/text:shieldPrompt?api-version="
def _azure_shield(request: Request) -> Reply:
assert request.method == "POST"
assert request.target.startswith(_AZURE_TARGET_PREFIX), request.target
user_prompt: Final = object_value(json.loads(request.body))["userPrompt"]
assert isinstance(user_prompt, str)
return Reply(
body=json.dumps(
{
"userPromptAnalysis": {"attackDetected": _ATTACK_MARKER in user_prompt},
"documentsAnalysis": [],
}
).encode()
)
def _provider(request: Request) -> Reply:
assert request.method == "POST"
if request.target == "/v1/messages":
return Reply(
body=json.dumps(
{
"id": "msg_" + uuid.uuid4().hex,
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-5-20250929",
"content": [{"type": "text", "text": "permitted response"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 11, "output_tokens": 4},
}
).encode()
)
if request.target == "/v1/responses":
return Reply(
body=json.dumps(
{
"id": "resp_" + uuid.uuid4().hex,
"object": "response",
"created_at": 1700000000,
"status": "completed",
"model": "gpt-4.1-mini",
"output": [
{
"type": "message",
"id": "msg_" + uuid.uuid4().hex,
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "permitted response", "annotations": []}],
}
],
"usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15},
}
).encode()
)
assert request.target == "/v1/chat/completions", request.target
return Reply(
body=json.dumps(
{
"id": "chatcmpl-" + uuid.uuid4().hex,
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4.1-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "permitted response"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15},
}
).encode()
)
@pytest.fixture(scope="module")
def azure_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[tuple[Gateway, Wire, Wire]]:
directory: Final = tmp_path_factory.mktemp("azure-shield")
with ExitStack() as stack:
gateway: Final = stack.enter_context(gateway_from_environment())
azure: Final = stack.enter_context(wire_server(_azure_shield))
provider: Final = stack.enter_context(wire_server(_provider))
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["guardrails"] = [
{
"guardrail_name": "azure-shield-" + uuid.uuid4().hex,
"litellm_params": {
"guardrail": "azure/prompt_shield",
"mode": "pre_call",
"default_on": True,
"api_base": azure.url,
"api_key": "synthetic-azure-key",
"cost_tier": "paid",
"price_per_1000_text_records": 0.38,
},
}
]
path: Final = directory / "azure-shield.yaml"
path.write_text(yaml.safe_dump(config))
candidate: Final = stack.enter_context(owned_proxy(gateway, directory, {}, config=path))
yield candidate, azure, provider
@pytest.fixture(autouse=True)
def _clear_wires(azure_rig: tuple[Gateway, Wire, Wire]) -> None:
azure_rig[1].drain()
azure_rig[2].drain()
def _scanned_prompts(azure: Wire) -> list[str]:
return [
object_value(json.loads(scan.body))["userPrompt"]
for scan in azure.drain()
if scan.target.startswith(_AZURE_TARGET_PREFIX)
]
def _guardrail_entry(model: str) -> dict:
rows: Final = eventually(
lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
lambda values: len(values) == 1,
seconds=70,
)
saved: Final = object_value(rows[0]["metadata"])
entries: Final = saved["guardrail_information"]
assert isinstance(entries, list) and len(entries) == 1, saved
return object_value(entries[0])
@pytest.mark.parametrize(
("path", "body_shape", "model_provider"),
[
pytest.param("/v1/chat/completions", "chat", "openai", id="chat-completions-messages"),
pytest.param("/v1/messages", "chat", "anthropic", id="anthropic-messages"),
pytest.param("/v1/responses", "responses-string", "openai", id="responses-string-input"),
pytest.param("/v1/responses", "responses-list", "openai", id="responses-list-input"),
pytest.param(
"/v1/responses", "responses-string-with-empty-messages", "openai", id="responses-empty-messages-stub"
),
],
)
def test_azure_prompt_shield_scans_the_user_prompt_on_every_endpoint(
azure_rig: tuple[Gateway, Wire, Wire], path: str, body_shape: str, model_provider: str
) -> None:
candidate, azure, provider = azure_rig
prompt: Final = f"synthetic prompt {body_shape} {uuid.uuid4().hex}"
with candidate.scenario() as scenario:
model: Final = scenario.model(
model=("anthropic/claude-sonnet-4-5-20250929" if model_provider == "anthropic" else "openai/gpt-4.1-mini"),
api_base=provider.url if model_provider == "anthropic" else provider.url + "/v1",
api_key="synthetic-provider-key",
)
body: Final = {
"responses-string": {"input": prompt},
"responses-list": {"input": [{"role": "user", "content": [{"type": "input_text", "text": prompt}]}]},
"responses-string-with-empty-messages": {"messages": [], "input": prompt},
}.get(
body_shape,
{"messages": [{"role": "user", "content": prompt}], "max_tokens": 16},
)
response: Final = candidate.request("POST", path, {"model": model, **body})
assert response.status_code == 200, response.text
assert "permitted response" in response.text
assert _scanned_prompts(azure) == [prompt]
assert len(provider.drain()) == 1
entry: Final = _guardrail_entry(model)
assert entry["guardrail_status"] == "success", entry
assert entry["guardrail_usage"] == {
"requests": 1,
"input_characters": len(prompt),
"text_records": 1,
}, entry
assert entry["guardrail_cost"] == pytest.approx(0.38 / 1000), entry
def test_azure_prompt_shield_blocks_attack_in_responses_input(
azure_rig: tuple[Gateway, Wire, Wire],
) -> None:
candidate, azure, provider = azure_rig
prompt: Final = f"synthetic prompt {_ATTACK_MARKER} {uuid.uuid4().hex}"
with candidate.scenario() as scenario:
model: Final = scenario.model(
model="openai/gpt-4.1-mini",
api_base=provider.url + "/v1",
api_key="synthetic-provider-key",
)
response: Final = candidate.request("POST", "/v1/responses", {"model": model, "input": prompt})
assert response.status_code == 400, response.text
assert "Violated Azure Prompt Shield guardrail policy" in response.text
assert _scanned_prompts(azure) == [prompt]
assert provider.drain() == ()

View file

@ -405,6 +405,49 @@ async def test_responses_input_is_scanned_and_billing_is_logged(responses_input:
assert entry["guardrail_cost_in_spend"] is False
@pytest.mark.asyncio
async def test_empty_messages_stub_does_not_hide_responses_input() -> None:
"""Cursor sends /v1/responses bodies with an empty messages list plus the real
input; a messages-first selector would scan nothing and let the prompt through."""
guardrail: Final = _priced_shield_guardrail(cost_tier="paid", price_per_1000_text_records=0.38)
prompt: Final = "summarize the thread"
data: Final[dict[str, object]] = {"messages": [], "input": prompt}
with patch.object(guardrail.async_handler, "post", return_value=_shield_response(False)) as mock_post:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="k"),
cache=None,
data=data,
call_type="aresponses",
)
mock_post.assert_called_once()
assert mock_post.call_args.kwargs["json"]["userPrompt"] == prompt
entry: Final = _recorded_guardrail_info(data)
assert entry["guardrail_usage"] == {"requests": 1, "input_characters": len(prompt), "text_records": 1}
assert entry["guardrail_cost"] == pytest.approx(0.00038)
@pytest.mark.asyncio
async def test_chat_call_type_scans_messages_not_input() -> None:
guardrail: Final = _shield_guardrail()
data: Final[dict[str, object]] = {
"messages": [{"role": "user", "content": "chat prompt"}],
"input": "unrelated responses input",
}
with patch.object(guardrail.async_handler, "post", return_value=_shield_response(False)) as mock_post:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="k"),
cache=None,
data=data,
call_type="acompletion",
)
mock_post.assert_called_once()
assert mock_post.call_args.kwargs["json"]["userPrompt"] == "chat prompt"
@pytest.mark.asyncio
async def test_responses_input_attack_detected_raises_http_exception() -> None:
guardrail: Final = _priced_shield_guardrail(cost_tier="paid", price_per_1000_text_records=0.38)

View file

@ -20,9 +20,7 @@ async def test_azure_text_moderation_guardrail_pre_call_hook():
api_key="azure_text_moderation_api_key",
api_base="azure_text_moderation_api_base",
)
with patch.object(
azure_text_moderation_guardrail, "async_make_request"
) as mock_async_make_request:
with patch.object(azure_text_moderation_guardrail, "async_make_request") as mock_async_make_request:
mock_async_make_request.return_value = {
"blocklistsMatch": [],
"categoriesAnalysis": [
@ -82,6 +80,60 @@ async def test_azure_text_moderation_scans_responses_input() -> None:
assert mock_post.call_args.kwargs["json"]["text"] == "Review this response input"
@pytest.mark.asyncio
async def test_azure_text_moderation_empty_messages_stub_does_not_hide_responses_input() -> None:
guardrail: Final = AzureContentSafetyTextModerationGuardrail(
guardrail_name="azure_text_moderation",
api_key="azure_text_moderation_api_key",
api_base="azure_text_moderation_api_base",
severity_threshold=4,
)
response: Final = Mock()
response.json.return_value = {
"blocklistsMatch": [],
"categoriesAnalysis": [
{"category": "Hate", "severity": 0},
{"category": "Sexual", "severity": 0},
{"category": "SelfHarm", "severity": 0},
{"category": "Violence", "severity": 0},
],
}
with patch.object(guardrail.async_handler, "post", return_value=response) as mock_post:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
cache=None,
data={"messages": [], "input": "Review this response input"},
call_type="aresponses",
)
mock_post.assert_called_once()
assert mock_post.call_args.kwargs["json"]["text"] == "Review this response input"
@pytest.mark.asyncio
async def test_azure_text_moderation_chat_call_type_scans_messages_not_input() -> None:
guardrail: Final = AzureContentSafetyTextModerationGuardrail(
guardrail_name="azure_text_moderation",
api_key="azure_text_moderation_api_key",
api_base="azure_text_moderation_api_base",
)
with patch.object(guardrail, "async_make_request") as mock_async_make_request:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
cache=None,
data={
"messages": [{"role": "user", "content": "chat prompt"}],
"input": "unrelated responses input",
},
call_type="acompletion",
)
mock_async_make_request.assert_called_once()
assert mock_async_make_request.call_args.kwargs["text"] == "chat prompt"
@pytest.mark.asyncio
async def test_azure_text_moderation_guardrail_violation_detected():
"""async_make_request is the single enforcement point — it raises
@ -93,20 +145,14 @@ async def test_azure_text_moderation_guardrail_violation_detected():
api_key="azure_text_moderation_api_key",
api_base="azure_text_moderation_api_base",
)
with patch.object(
azure_text_moderation_guardrail, "async_make_request"
) as mock_async_make_request:
with patch.object(azure_text_moderation_guardrail, "async_make_request") as mock_async_make_request:
mock_async_make_request.side_effect = HTTPException(
status_code=400,
detail={
"error": "Azure Content Safety Guardrail: Hate crossed severity 2, Got severity: 2"
},
detail={"error": "Azure Content Safety Guardrail: Hate crossed severity 2, Got severity: 2"},
)
with pytest.raises(HTTPException):
await azure_text_moderation_guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
api_key="azure_text_moderation_api_key"
),
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
cache=None,
data={
"messages": [
@ -215,9 +261,7 @@ async def test_azure_text_moderation_violation_in_chunk():
):
with pytest.raises(HTTPException):
await azure_text_moderation_guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
api_key="azure_text_moderation_api_key"
),
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
cache=None,
data={
"messages": [
@ -239,9 +283,7 @@ async def test_azure_text_moderation_guardrail_post_call_success_hook():
api_key="azure_text_moderation_api_key",
api_base="azure_text_moderation_api_base",
)
with patch.object(
azure_text_moderation_guardrail, "async_make_request"
) as mock_async_make_request:
with patch.object(azure_text_moderation_guardrail, "async_make_request") as mock_async_make_request:
mock_async_make_request.return_value = {
"blocklistsMatch": [],
"categoriesAnalysis": [
@ -273,9 +315,7 @@ async def test_azure_text_moderation_guardrail_post_call_checks_all_choices():
api_key="azure_text_moderation_api_key",
api_base="azure_text_moderation_api_base",
)
with patch.object(
azure_text_moderation_guardrail, "async_make_request"
) as mock_async_make_request:
with patch.object(azure_text_moderation_guardrail, "async_make_request") as mock_async_make_request:
mock_async_make_request.side_effect = [
{
"blocklistsMatch": [],
@ -290,9 +330,7 @@ async def test_azure_text_moderation_guardrail_post_call_checks_all_choices():
with pytest.raises(HTTPException):
await azure_text_moderation_guardrail.async_post_call_success_hook(
data={},
user_api_key_dict=UserAPIKeyAuth(
api_key="azure_text_moderation_api_key"
),
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
response=ModelResponse(
choices=[
Choices(
@ -307,9 +345,10 @@ async def test_azure_text_moderation_guardrail_post_call_checks_all_choices():
),
)
assert [
call.kwargs["text"] for call in mock_async_make_request.call_args_list
] == ["safe response", "unsafe response"]
assert [call.kwargs["text"] for call in mock_async_make_request.call_args_list] == [
"safe response",
"unsafe response",
]
@pytest.mark.asyncio
@ -320,9 +359,7 @@ async def test_azure_text_moderation_guardrail_post_call_streaming_hook():
api_key="azure_text_moderation_api_key",
api_base="azure_text_moderation_api_base",
)
with patch.object(
azure_text_moderation_guardrail, "async_make_request"
) as mock_async_make_request:
with patch.object(azure_text_moderation_guardrail, "async_make_request") as mock_async_make_request:
mock_async_make_request.return_value = {
"blocklistsMatch": [],
"categoriesAnalysis": [
@ -359,13 +396,7 @@ def test_split_text_by_words():
assert len(chunks) > 1
# Verify no word is broken
for chunk in chunks:
assert (
"word1" in chunk
or "word2" in chunk
or "word3" in chunk
or "word4" in chunk
or "word5" in chunk
)
assert "word1" in chunk or "word2" in chunk or "word3" in chunk or "word4" in chunk or "word5" in chunk
# Test with very long single word (edge case)
long_word = "supercalifragilisticexpialidocious" * 10
@ -464,9 +495,7 @@ async def test_apply_guardrail_scans_every_text():
async def test_apply_guardrail_raises_on_detection_in_any_text():
guardrail = _moderation_guardrail()
with patch.object(
guardrail.async_handler, "post", side_effect=[_moderation_response(0), _moderation_response(6)]
):
with patch.object(guardrail.async_handler, "post", side_effect=[_moderation_response(0), _moderation_response(6)]):
with pytest.raises(HTTPException) as exc_info:
await guardrail.apply_guardrail(
inputs={"texts": ["hello there", "something hateful"]},