fix(guardrails): keep Azure Text Moderation on messages only so this PR stays Prompt Shield scoped
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-10-01 01:53:54 +00:00
parent 47380ae72d
commit 9bd060238d
4 changed files with 62 additions and 227 deletions

View file

@ -146,3 +146,17 @@ class AzureGuardrailBase:
if not isinstance(messages, list):
return None
return get_last_user_message(cast(list[AllMessageValues], messages)) # cast-ok: narrowed to list
def get_user_prompt(self, messages: list[AllMessageValues]) -> str | None:
"""
Get the last consecutive block of messages from the user.
Example:
messages = [
{"role": "user", "content": "Hello, how are you?"},
{"role": "assistant", "content": "I'm good, thank you!"},
{"role": "user", "content": "What is the weather in Tokyo?"},
]
get_user_prompt(messages) -> "What is the weather in Tokyo?"
"""
return get_last_user_message(messages)

View file

@ -21,6 +21,7 @@ from .base import AzureGuardrailBase
if TYPE_CHECKING:
from litellm.caching.caching import DualCache
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.llms.openai import AllMessageValues
from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_text_moderation import (
AzureTextModerationGuardrailResponse,
)
@ -231,7 +232,11 @@ 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, call_type)
new_messages: Final[list[AllMessageValues] | None] = data.get("messages")
if new_messages is None:
verbose_proxy_logger.warning("Azure Text Moderation: not running guardrail. No messages in data")
return data
user_prompt: Final = self.get_user_prompt(new_messages)
if user_prompt:
verbose_proxy_logger.info("Azure Text Moderation: User prompt: %s", user_prompt)

View file

@ -17,13 +17,10 @@ from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue
_ATTACK_MARKER: Final = "synthetic-attack-marker"
_MODERATION_MARKER: Final = "synthetic-moderation-marker"
_SHIELD_TARGET_PREFIX: Final = "/contentsafety/text:shieldPrompt?api-version="
_ANALYZE_TARGET_PREFIX: Final = "/contentsafety/text:analyze?api-version="
_OPT_IN_SHIELD: Final = "audit-shield-optin"
_TEXT_MODERATION: Final = "audit-text-mod"
def _chat_frame(identity: str, delta: dict[str, JsonValue], finish: str | None = None) -> bytes:
@ -142,23 +139,7 @@ def _azure(outage: threading.Event) -> Callable[[Request], Reply]:
}
).encode()
)
assert request.target.startswith(_ANALYZE_TARGET_PREFIX), request.target
text: Final = body["text"]
assert isinstance(text, str)
severity: Final = 4 if _MODERATION_MARKER in text else 0
return Reply(
body=json.dumps(
{
"blocklistsMatch": [],
"categoriesAnalysis": [
{"category": "Hate", "severity": severity},
{"category": "Sexual", "severity": 0},
{"category": "SelfHarm", "severity": 0},
{"category": "Violence", "severity": 0},
],
}
).encode()
)
return Reply(status=404)
return respond
@ -201,16 +182,6 @@ def audit_rig(
"guardrail_name": "audit-shield",
"litellm_params": _shield_params(azure, mode="pre_call", default_on=True),
},
{
"guardrail_name": _TEXT_MODERATION,
"litellm_params": {
"guardrail": "azure/text_moderations",
"mode": "pre_call",
"default_on": False,
"api_base": azure.url,
"api_key": "synthetic-azure-key",
},
},
],
)
owned: Final = stack.enter_context(owned_proxy_process(gateway, directory, {}, config=config, workers=2))
@ -261,16 +232,6 @@ def chaos_rig(
"guardrail_name": "audit-shield",
"litellm_params": _shield_params(azure, mode="pre_call", default_on=True),
},
{
"guardrail_name": _TEXT_MODERATION,
"litellm_params": {
"guardrail": "azure/text_moderations",
"mode": "pre_call",
"default_on": False,
"api_base": azure.url,
"api_key": "synthetic-azure-key",
},
},
],
)
owned: Final = stack.enter_context(owned_proxy_process(gateway, directory, {}, config=config, workers=2))
@ -320,14 +281,6 @@ def _shield_prompts(requests: tuple[Request, ...]) -> tuple[JsonValue, ...]:
)
def _analyze_texts(requests: tuple[Request, ...]) -> tuple[JsonValue, ...]:
return tuple(
object_value(json.loads(scan.body))["text"]
for scan in requests
if scan.target.startswith(_ANALYZE_TARGET_PREFIX)
)
def _guardrail_entries(model: str, count: int = 1) -> list[JsonValue]:
rows: Final = eventually(
lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
@ -402,60 +355,6 @@ def test_responses_streaming_input_is_scanned_and_billed(
assert entry["guardrail_cost"] == pytest.approx(0.38 / 1000), entry
@pytest.mark.parametrize(
"body",
[
pytest.param(lambda prompt: {"input": prompt}, id="string-input"),
pytest.param(
lambda prompt: {"input": [{"role": "user", "content": [{"type": "input_text", "text": prompt}]}]},
id="list-input",
),
pytest.param(lambda prompt: {"messages": [], "input": prompt}, id="empty-messages-stub"),
],
)
def test_text_moderation_opt_in_scans_responses_input(
audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event],
request: pytest.FixtureRequest,
body: Callable[[str], dict[str, JsonValue]],
) -> None:
owned, azure, provider, _ = audit_rig
prompt: Final = f"synthetic benign prompt {request.node.callspec.id} {uuid.uuid4().hex}"
with owned.gateway.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 = owned.gateway.request(
"POST", "/v1/responses", {"model": model, "guardrails": [_TEXT_MODERATION], **body(prompt)}
)
assert response.status_code == 200, response.text
calls: Final = azure.drain()
assert _analyze_texts(calls) == (prompt,)
assert _shield_prompts(calls) == (prompt,)
assert len(_provider_calls(provider)) == 1
entries: Final = _guardrail_entries(model, count=2)
assert {object_value(entry)["guardrail_name"] for entry in entries} == {"audit-shield", _TEXT_MODERATION}, (
entries
)
def test_text_moderation_opt_in_scans_chat_messages(audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event]) -> None:
owned, azure, provider, _ = audit_rig
prompt: Final = "synthetic benign prompt chat-optin " + uuid.uuid4().hex
with owned.gateway.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 = owned.gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "guardrails": [_TEXT_MODERATION], "messages": [{"role": "user", "content": prompt}]},
)
assert response.status_code == 200, response.text
calls: Final = azure.drain()
assert _analyze_texts(calls) == (prompt,)
assert _shield_prompts(calls) == (prompt,)
def test_chat_with_input_key_still_scans_messages_only(
audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event],
) -> None:
@ -630,23 +529,6 @@ def test_streaming_responses_attack_is_blocked_before_any_stream_bytes(
assert _provider_calls(provider) == ()
def test_text_moderation_opt_in_blocks_responses_input_above_threshold(
audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event],
) -> None:
owned, azure, provider, _ = audit_rig
prompt: Final = f"synthetic prompt {_MODERATION_MARKER} " + uuid.uuid4().hex
with owned.gateway.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 = owned.gateway.request(
"POST", "/v1/responses", {"model": model, "guardrails": [_TEXT_MODERATION], "input": prompt}
)
assert response.status_code == 400, response.text
assert _analyze_texts(azure.drain()) == (prompt,)
assert _provider_calls(provider) == ()
def test_azure_outage_produces_the_same_outcome_on_responses_and_chat(
audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event],
) -> None:

View file

@ -1,14 +1,13 @@
from typing import Final
from unittest.mock import Mock, patch
import pytest
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_registry import InMemoryGuardrailHandler
from litellm.proxy.guardrails.guardrail_hooks.azure.text_moderation import (
AzureContentSafetyTextModerationGuardrail,
)
from litellm.proxy.guardrails.guardrail_registry import InMemoryGuardrailHandler
from litellm.types.utils import Choices, Message, ModelResponse
@ -20,7 +19,9 @@ 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": [
@ -48,96 +49,6 @@ async def test_azure_text_moderation_guardrail_pre_call_hook():
assert mock_async_make_request.call_args.kwargs["text"] == "Hello, how are you?"
@pytest.mark.asyncio
async def test_azure_text_moderation_scans_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",
)
response: Final = Mock()
response.json.return_value = {
"blocklistsMatch": [],
"categoriesAnalysis": [
{"category": "Hate", "severity": 2},
{"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:
with pytest.raises(HTTPException) as exc_info:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
cache=None,
data={"input": "Review this response input"},
call_type="aresponses",
)
assert exc_info.value.status_code == 400
mock_post.assert_called_once()
assert mock_post.call_args.kwargs["json"]["text"] == "Review this response input"
def _moderation_flagging(flagged: str):
def azure_by_text(*args: object, **kwargs: object) -> Mock:
body = kwargs["json"]
assert isinstance(body, dict)
return _moderation_response(6 if body["text"] == flagged else 0)
return azure_by_text
@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,
)
flagged: Final = "flagged responses input"
data: Final[dict[str, object]] = {"messages": [], "input": flagged}
with patch.object(guardrail.async_handler, "post", side_effect=_moderation_flagging(flagged)):
with pytest.raises(HTTPException) as exc_info:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
cache=None,
data=data,
call_type="aresponses",
)
assert exc_info.value.status_code == 400
@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",
severity_threshold=4,
)
flagged: Final = "flagged chat prompt"
data: Final[dict[str, object]] = {
"messages": [{"role": "user", "content": flagged}],
"input": "benign responses input",
}
with patch.object(guardrail.async_handler, "post", side_effect=_moderation_flagging(flagged)):
with pytest.raises(HTTPException) as exc_info:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
cache=None,
data=data,
call_type="acompletion",
)
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_azure_text_moderation_guardrail_violation_detected():
"""async_make_request is the single enforcement point — it raises
@ -149,14 +60,20 @@ 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": [
@ -265,7 +182,9 @@ 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": [
@ -287,7 +206,9 @@ 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": [
@ -319,7 +240,9 @@ 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": [],
@ -334,7 +257,9 @@ 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(
@ -349,10 +274,9 @@ 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
@ -363,7 +287,9 @@ 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": [
@ -400,7 +326,13 @@ 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
@ -499,7 +431,9 @@ 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"]},