fix(guardrails): restore Azure guardrail get_user_prompt dispatch and allow logging (#44067)

* fix(guardrails): restore Azure guardrail get_user_prompt dispatch and allow logging

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(guardrails): use transport-level doubles in Azure dispatch regression tests

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(guardrails): add Azure dispatch audit matrix integration cells

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(guardrails): restore global callback lists after Router reset in SDK cells

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(guardrails): make proxy restart cell independent of shutdown timing

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yucheng <yucheng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-01 23:30:32 -07:00 • committed by GitHub
parent 5a3a31ea9a
commit d131c43782
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 3115 additions and 11 deletions

View file

@ -138,6 +138,20 @@ class AzureGuardrailBase:
return chunks
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)
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")
@ -147,6 +161,6 @@ class AzureGuardrailBase:
return get_last_user_message(ResponsesAPIRequestUtils.responses_input_to_chat_messages(validated_input))
messages: Final = data.get("messages")
if not isinstance(messages, list):
if messages is None:
return None
return get_last_user_message(cast(list[AllMessageValues], messages)) # cast-ok: narrowed to list
return self.get_user_prompt(cast(list[AllMessageValues], messages)) # cast-ok: sequence of request messages

View file

@ -27,7 +27,7 @@ from litellm.types.utils import (
GuardrailTracingDetail,
)
from .base import AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH, AzureGuardrailBase
from .base import _RESPONSES_API_CALL_TYPES, AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH, AzureGuardrailBase
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -249,6 +249,9 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai
"Azure Prompt Shield: Running pre-call prompt scan, on call_type: %s",
call_type,
)
if call_type not in _RESPONSES_API_CALL_TYPES and data.get("messages") is None:
verbose_proxy_logger.warning("Azure Prompt Shield: not running guardrail. No messages in data")
return data
user_prompt: Final = self.get_user_prompt_from_request(data, call_type)
if user_prompt:

View file

@ -16,7 +16,7 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import CallTypesLiteral, GenericGuardrailAPIInputs, LLMResponseTypes
from .base import AzureGuardrailBase
from .base import _RESPONSES_API_CALL_TYPES, AzureGuardrailBase
if TYPE_CHECKING:
from litellm.caching.caching import DualCache
@ -231,6 +231,9 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr
"Azure Text Moderation: Running pre-call prompt scan, on call_type: %s",
call_type,
)
if call_type not in _RESPONSES_API_CALL_TYPES and data.get("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_from_request(data, call_type)
if user_prompt:

View file

@ -0,0 +1,551 @@
import json
import threading
import time
import uuid
from collections.abc import Callable, Iterator
from contextlib import ExitStack, contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import yaml
from integration._support.client import Gateway, object_value
from integration._support.process import OwnedProxy, owned_proxy_process
from integration._support.redis_process import OwnedRedis, owned_redis
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"
AZURE_ERROR_MARKERS: Final = ("AZURE_500", "AZURE_403", "AZURE_404")
PROVIDER_401_MARKER: Final = "PROVIDER_401"
OVERSIZED_MARKER: Final = "OVERSIZED_INPUT"
SHIELD_TARGET_PREFIX: Final = "/contentsafety/text:shieldPrompt?api-version="
ANALYZE_TARGET_PREFIX: Final = "/contentsafety/text:analyze?api-version="
HOOKS_SOURCE: Final = """from __future__ import annotations
from typing import Final, cast
from fastapi import HTTPException
from litellm.caching.caching import DualCache
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.azure.prompt_shield import (
AzureContentSafetyPromptShieldGuardrail,
)
from litellm.proxy.guardrails.guardrail_hooks.azure.text_moderation import (
AzureContentSafetyTextModerationGuardrail,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import CallTypesLiteral
class TupleWriter(CustomGuardrail):
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
cache: DualCache,
data: dict[str, object],
call_type: CallTypesLiteral,
) -> dict[str, object]:
messages: Final = data.get("messages")
if isinstance(messages, list):
data["messages"] = tuple(messages) # mutable-ok: the test hook rewrites messages to a tuple
return data
class AllTurnsPromptShield(AzureContentSafetyPromptShieldGuardrail):
def get_user_prompt(self, messages: list[AllMessageValues]) -> str:
return "\\n".join(
message["content"]
for message in messages
if isinstance(message, dict)
and message.get("role") == "user"
and isinstance(message.get("content"), str)
)
class AllTurnsTextModeration(AzureContentSafetyTextModerationGuardrail):
def get_user_prompt(self, messages: list[AllMessageValues]) -> str:
return "\\n".join(
message["content"]
for message in messages
if isinstance(message, dict)
and message.get("role") == "user"
and isinstance(message.get("content"), str)
)
class RequiringPromptShield(AzureContentSafetyPromptShieldGuardrail):
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
cache: DualCache,
data: dict[str, object],
call_type: CallTypesLiteral,
) -> dict[str, object] | None:
messages: Final = data.get("messages")
if not isinstance(messages, list):
raise HTTPException(status_code=400, detail="no user text")
user_prompt: Final = self.get_user_prompt(cast(list[AllMessageValues], messages)) # cast-ok: chat messages
if not user_prompt:
raise HTTPException(status_code=400, detail="no user text")
return await super().async_pre_call_hook(user_api_key_dict, cache, data, call_type)
class RequiringTextModeration(AzureContentSafetyTextModerationGuardrail):
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
cache: DualCache,
data: dict[str, object],
call_type: CallTypesLiteral,
) -> dict[str, object] | None:
messages: Final = data.get("messages")
if not isinstance(messages, list):
raise HTTPException(status_code=400, detail="no user text")
user_prompt: Final = self.get_user_prompt(cast(list[AllMessageValues], messages)) # cast-ok: chat messages
if not user_prompt:
raise HTTPException(status_code=400, detail="no user text")
return await super().async_pre_call_hook(user_api_key_dict, cache, data, call_type)
"""
@dataclass(frozen=True, slots=True)
class AzureBehavior:
delay_seconds: float = 0
down: threading.Event | None = None
entered: threading.Event | None = None
arrived: threading.Semaphore | None = None
release: threading.Event | None = None
barrier_marker: str | None = None
def azure_text(request: Request) -> str:
body: Final = object_value(json.loads(request.body))
if request.target.startswith(SHIELD_TARGET_PREFIX):
prompt: Final = body["userPrompt"]
assert isinstance(prompt, str), body
return prompt
assert request.target.startswith(ANALYZE_TARGET_PREFIX), request.target
text: Final = body["text"]
assert isinstance(text, str), body
return text
def azure_texts(azure: Wire) -> tuple[str, ...]:
return tuple(azure_text(request) for request in azure.drain())
def provider_text(request: Request) -> str:
body: Final = object_value(json.loads(request.body)) if request.body else {}
target: Final = request.target.split("?", 1)[0]
match target:
case "/v1/chat/completions" | "/v1/messages":
messages: Final = body["messages"]
assert isinstance(messages, list), body
return "\n".join(
message_text(message["content"])
for message in messages
if isinstance(message, dict)
and message.get("role") == "user"
and "content" in message
)
case "/v1/responses":
value: Final = body["input"]
if not isinstance(value, list):
return message_text(value)
return "\n".join(
message_text(message["content"])
for message in value
if isinstance(message, dict)
and message.get("role") == "user"
and "content" in message
)
case "/v1/embeddings":
return message_text(body["input"])
case "/v1/completions":
return message_text(body["prompt"])
case _:
raise AssertionError(f"Unexpected provider target {request.target}")
def message_text(value: JsonValue) -> str:
if isinstance(value, str):
return value
if isinstance(value, list):
return "".join(
str(part["text"])
for part in value
if isinstance(part, dict) and isinstance(part.get("text"), str)
)
return ""
def provider_texts(provider: Wire) -> tuple[str, ...]:
requests: Final = tuple(
request
for request in provider.drain()
if request.method != "GET" or request.target.split("?", 1)[0] != "/v1/models"
)
return tuple(provider_text(request) for request in requests)
def provider_messages(provider: Wire) -> tuple[JsonValue, ...]:
requests: Final = tuple(
request
for request in provider.drain()
if request.method != "GET" or request.target.split("?", 1)[0] != "/v1/models"
)
targets: Final = tuple(request.target.split("?", 1)[0] for request in requests)
assert all(target == "/v1/chat/completions" for target in targets), targets
return tuple(object_value(json.loads(request.body))["messages"] for request in requests)
def azure_handler(behavior: AzureBehavior = AzureBehavior()) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert request.method == "POST", request.method
text: Final = azure_text(request)
if behavior.entered is not None and (
behavior.barrier_marker is None or behavior.barrier_marker in text
):
behavior.entered.set()
if behavior.arrived is not None:
behavior.arrived.release()
if behavior.release is not None and (
behavior.barrier_marker is None or behavior.barrier_marker in text
):
assert behavior.release.wait(timeout=30), "Azure barrier was not released"
if behavior.down is not None and behavior.down.is_set():
return Reply(status=503, body=b'{"error":"synthetic Azure outage"}')
if behavior.delay_seconds:
time.sleep(behavior.delay_seconds)
status: Final = next(
(code for marker, code in (("AZURE_500", 500), ("AZURE_403", 403), ("AZURE_404", 404)) if marker in text),
200,
)
if status != 200:
return Reply(
status=status,
body=json.dumps({"error": {"message": f"synthetic Azure error {status}"}}).encode(),
)
if request.target.startswith(SHIELD_TARGET_PREFIX):
return Reply(
body=json.dumps(
{
"userPromptAnalysis": {"attackDetected": ATTACK_MARKER in text},
"documentsAnalysis": [],
}
).encode()
)
assert request.target.startswith(ANALYZE_TARGET_PREFIX), request.target
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 respond
def provider_handler(request: Request) -> Reply:
path: Final = request.target.split("?", 1)[0]
if request.method == "GET" and path == "/v1/models":
return Reply(body=b'{"data":[]}')
assert request.method == "POST", request.method
text: Final = provider_text(request)
if PROVIDER_401_MARKER in text:
return Reply(status=401, body=b'{"error":{"message":"synthetic provider unauthorized"}}')
if OVERSIZED_MARKER in text:
return Reply(
status=400,
body=b'{"error":{"message":"synthetic context length exceeded","code":"context_length_exceeded"}}',
)
identity: Final = uuid.uuid4().hex
match path:
case "/v1/chat/completions":
if b'"stream":true' in request.body.replace(b" ", b""):
return Reply(content_type="text/event-stream", chunks=_chat_chunks(identity))
return Reply(
body=json.dumps(
{
"id": "chatcmpl-" + identity,
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "permitted response"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 2, "completion_tokens": 2, "total_tokens": 4},
}
).encode()
)
case "/v1/messages":
if b'"stream":true' in request.body.replace(b" ", b""):
return Reply(content_type="text/event-stream", chunks=_messages_chunks(identity))
return Reply(
body=json.dumps(
{
"id": "msg_" + identity,
"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": 2, "output_tokens": 2},
}
).encode()
)
case "/v1/responses":
if b'"stream":true' in request.body.replace(b" ", b""):
return Reply(content_type="text/event-stream", chunks=_responses_chunks(identity))
return Reply(
body=json.dumps(
{
"id": "resp_" + identity,
"object": "response",
"created_at": 1700000000,
"status": "completed",
"model": "gpt-4o-mini",
"output": [
{
"type": "message",
"id": "msg_" + identity,
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "permitted response", "annotations": []}],
}
],
"usage": {"input_tokens": 2, "output_tokens": 2, "total_tokens": 4},
}
).encode()
)
case "/v1/embeddings":
return Reply(
body=json.dumps(
{
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}],
"model": "text-embedding-3-small",
"usage": {"prompt_tokens": 1, "total_tokens": 1},
}
).encode()
)
case "/v1/completions":
return Reply(
body=json.dumps(
{
"id": "cmpl-" + identity,
"object": "text_completion",
"created": 1700000000,
"model": "gpt-4o-mini",
"choices": [{"text": "permitted response", "index": 0, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 2, "total_tokens": 3},
}
).encode()
)
case _:
return Reply(status=404, body=json.dumps({"error": "unexpected provider target " + path}).encode())
def _chat_chunks(identity: str) -> tuple[bytes, ...]:
return (
_sse({"id": "chatcmpl-" + identity, "object": "chat.completion.chunk", "choices": [{"index": 0, "delta": {"role": "assistant", "content": "permitted "}, "finish_reason": None}]}),
_sse({"id": "chatcmpl-" + identity, "object": "chat.completion.chunk", "choices": [{"index": 0, "delta": {"content": "response"}, "finish_reason": None}]}),
_sse({"id": "chatcmpl-" + identity, "object": "chat.completion.chunk", "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}),
b"data: [DONE]\n\n",
)
def _messages_chunks(identity: str) -> tuple[bytes, ...]:
return (
_event("message_start", {"type": "message_start", "message": {"id": "msg_" + identity, "type": "message", "role": "assistant", "model": "claude-sonnet-4-5-20250929", "content": [], "stop_reason": None, "stop_sequence": None, "usage": {"input_tokens": 2, "output_tokens": 0}}}),
_event("content_block_start", {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}),
_event("content_block_delta", {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "permitted response"}}),
_event("content_block_stop", {"type": "content_block_stop", "index": 0}),
_event("message_delta", {"type": "message_delta", "delta": {"stop_reason": "end_turn", "stop_sequence": None}, "usage": {"output_tokens": 2}}),
_event("message_stop", {"type": "message_stop"}),
)
def _responses_chunks(identity: str) -> tuple[bytes, ...]:
response: Final = {
"id": "resp_" + identity,
"object": "response",
"created_at": 1700000000,
"status": "completed",
"model": "gpt-4o-mini",
"output": [
{
"type": "message",
"id": "msg_" + identity,
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "permitted response", "annotations": []}],
}
],
"usage": {"input_tokens": 2, "output_tokens": 2, "total_tokens": 4},
}
return (
_event("response.created", {"type": "response.created", "response": response}),
_event("response.output_text.delta", {"type": "response.output_text.delta", "delta": "permitted response"}),
_event("response.completed", {"type": "response.completed", "response": response}),
)
def _sse(value: dict[str, JsonValue]) -> bytes:
return b"data: " + json.dumps(value).encode() + b"\n\n"
def _event(name: str, value: dict[str, JsonValue]) -> bytes:
return f"event: {name}\n".encode() + _sse(value)
def guardrail_configs(azure_url: str, *, default_on: bool = False) -> tuple[dict[str, JsonValue], ...]:
return (
{
"guardrail_name": "tuple-writer",
"litellm_params": {
"guardrail": "azure_dispatch_hooks.TupleWriter",
"mode": "pre_call",
"default_on": default_on,
},
},
{
"guardrail_name": "shield",
"litellm_params": {
"guardrail": "azure/prompt_shield",
"mode": "pre_call",
"default_on": default_on,
"api_base": azure_url,
"api_key": "synthetic-azure-key",
},
},
{
"guardrail_name": "moderation",
"litellm_params": {
"guardrail": "azure/text_moderations",
"mode": "pre_call",
"default_on": False,
"api_base": azure_url,
"api_key": "synthetic-azure-key",
},
},
{
"guardrail_name": "all-turns-shield",
"litellm_params": {
"guardrail": "azure_dispatch_hooks.AllTurnsPromptShield",
"mode": "pre_call",
"default_on": False,
"api_base": azure_url,
"api_key": "synthetic-azure-key",
},
},
{
"guardrail_name": "all-turns-moderation",
"litellm_params": {
"guardrail": "azure_dispatch_hooks.AllTurnsTextModeration",
"mode": "pre_call",
"default_on": False,
"api_base": azure_url,
"api_key": "synthetic-azure-key",
},
},
{
"guardrail_name": "requiring-shield",
"litellm_params": {
"guardrail": "azure_dispatch_hooks.RequiringPromptShield",
"mode": "pre_call",
"default_on": False,
"api_base": azure_url,
"api_key": "synthetic-azure-key",
},
},
{
"guardrail_name": "requiring-moderation",
"litellm_params": {
"guardrail": "azure_dispatch_hooks.RequiringTextModeration",
"mode": "pre_call",
"default_on": False,
"api_base": azure_url,
"api_key": "synthetic-azure-key",
},
},
)
def write_dispatch_config(directory: Path, azure_url: str, *, default_on: bool = False) -> Path:
(directory / "azure_dispatch_hooks.py").write_text(HOOKS_SOURCE)
base_config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
general_settings: Final = {
**base_config["general_settings"],
"store_prompts_in_spend_logs": True,
}
config: Final = {
**base_config,
"guardrails": guardrail_configs(azure_url, default_on=default_on),
"general_settings": general_settings,
}
config_path: Final = directory / "azure-content-safety-dispatch.yaml"
config_path.write_text(yaml.safe_dump(config))
return config_path
@contextmanager
def dispatch_proxy(
gateway: Gateway,
directory: Path,
redis: OwnedRedis,
azure_url: str,
*,
workers: int = 2,
default_on: bool = False,
) -> Iterator[OwnedProxy]:
config: Final = write_dispatch_config(directory, azure_url, default_on=default_on)
with owned_proxy_process(
gateway,
directory,
{
"REDIS_HOST": redis.host,
"REDIS_PORT": str(redis.port),
"LITELLM_DISABLE_NO_REDIS_WARNING": "true",
},
config=config,
workers=workers,
) as owned:
yield owned
@contextmanager
def dispatch_rig(
gateway: Gateway,
directory: Path,
*,
workers: int = 2,
default_on: bool = False,
behavior: AzureBehavior = AzureBehavior(),
) -> Iterator[tuple[OwnedProxy, Wire, Wire, OwnedRedis]]:
with ExitStack() as stack:
redis: Final = stack.enter_context(owned_redis(directory))
azure: Final = stack.enter_context(wire_server(azure_handler(behavior)))
provider: Final = stack.enter_context(wire_server(provider_handler))
owned: Final = stack.enter_context(
dispatch_proxy(gateway, directory, redis, azure.url, workers=workers, default_on=default_on)
)
yield owned, azure, provider, redis

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,514 @@
import concurrent.futures
import socket
import threading
import uuid
from collections.abc import Iterator
from contextlib import ExitStack
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import httpx
import psutil
import pytest
from integration._support.client import Gateway, Scenario, eventually
from integration._support.database import read_rows
from integration._support.process import OwnedProxy
from integration._support.redis_process import owned_redis
from integration._support.wire import Wire, wire_server
from integration.observability.azure_dispatch_support import (
ATTACK_MARKER,
AzureBehavior,
azure_handler,
azure_texts,
dispatch_proxy,
dispatch_rig,
provider_handler,
provider_texts,
write_dispatch_config,
)
from pydantic import JsonValue
@dataclass(frozen=True, slots=True)
class CallSpec:
phase: str
prompt: str
path: str
body: dict[str, JsonValue]
stream: bool
@dataclass(frozen=True, slots=True)
class CallResult:
spec: CallSpec
status: int
text: str
headers: dict[str, str]
def _model(scenario: Scenario, provider: Wire, name: str = "openai/gpt-4o-mini") -> str:
return scenario.model(model=name, api_base=provider.url + "/v1", api_key="synthetic-provider-key")
def _anthropic_model(scenario: Scenario, provider: Wire) -> str:
return scenario.model(
model="anthropic/claude-sonnet-4-5-20250929",
api_base=provider.url,
api_key="synthetic-provider-key",
)
def _assert_spend(request_id: str, response_text: str) -> None:
rows: Final = eventually(
lambda: read_rows('SELECT litellm_call_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (request_id,)),
lambda values: len(values) == 1,
seconds=70,
)
assert len(rows) == 1, response_text
def _request(
candidate: Gateway,
path: str,
body: dict[str, JsonValue],
*,
stream: bool = False,
) -> tuple[int, str, dict[str, str]]:
if not stream:
response: Final = candidate.request("POST", path, body)
return response.status_code, response.text, dict(response.headers)
with candidate.client.stream(
"POST",
path,
json=body,
headers={"Authorization": f"Bearer {candidate.key}"},
) as response:
response.read()
return response.status_code, response.text, dict(response.headers)
def _safe_request(candidate: Gateway, model: str, prompt: str) -> httpx.Response | None:
try:
return candidate.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": prompt}],
"guardrails": ["shield"],
"cache": {"no-cache": True},
},
)
except httpx.HTTPError:
return None
def _worker_processes(owned: OwnedProxy) -> tuple[psutil.Process, ...]:
descendants: Final = tuple(psutil.Process(owned.process.pid).children(recursive=True))
return tuple(
process
for process in descendants
if process.is_running()
and any("spawn_main" in argument for argument in process.cmdline())
)
def _operation_specs(
phase: str,
operation: str,
openai_model: str,
anthropic_model: str,
attack: bool,
) -> tuple[CallSpec, ...]:
marker: Final = ATTACK_MARKER if attack else "synthetic-benign"
if operation == "chat":
return tuple(
CallSpec(
phase,
f"{phase} chat {marker} {index}",
"/v1/chat/completions",
{
"model": openai_model,
"messages": [{"role": "user", "content": f"{phase} chat {marker} {index}"}],
"guardrails": ["tuple-writer", "shield"],
"cache": {"no-cache": True},
},
False,
)
for index in range(2)
)
if operation == "chat-stream":
return tuple(
CallSpec(
phase,
f"{phase} chat stream {marker} {index}",
"/v1/chat/completions",
{
"model": openai_model,
"messages": [{"role": "user", "content": f"{phase} chat stream {marker} {index}"}],
"guardrails": ["tuple-writer", "shield"],
"stream": True,
"cache": {"no-cache": True},
},
True,
)
for index in range(2)
)
if operation == "messages":
return tuple(
CallSpec(
phase,
f"{phase} messages {marker} {index}",
"/v1/messages",
{
"model": anthropic_model,
"max_tokens": 16,
"messages": [{"role": "user", "content": f"{phase} messages {marker} {index}"}],
"guardrails": ["tuple-writer", "shield"],
"cache": {"no-cache": True},
},
False,
)
for index in range(2)
)
if operation == "responses":
return tuple(
CallSpec(
phase,
f"{phase} responses {marker} {index}",
"/v1/responses",
{
"model": openai_model,
"input": f"{phase} responses {marker} {index}",
"guardrails": ["shield"],
"cache": {"no-cache": True},
},
False,
)
for index in range(2)
)
return tuple(
CallSpec(
phase,
f"{phase} list {marker} {index}",
"/v1/chat/completions",
{
"model": openai_model,
"messages": [{"role": "user", "content": f"{phase} list {marker} {index}"}],
"guardrails": ["shield"],
"cache": {"no-cache": True},
},
False,
)
for index in range(2)
)
def _phase_specs(
phase: str, openai_model: str, anthropic_model: str, attack: bool
) -> tuple[CallSpec, ...]:
return (
*_operation_specs(phase, "chat", openai_model, anthropic_model, attack),
*_operation_specs(phase, "chat-stream", openai_model, anthropic_model, attack),
*_operation_specs(phase, "messages", openai_model, anthropic_model, attack),
*_operation_specs(phase, "responses", openai_model, anthropic_model, attack),
*_operation_specs(phase, "list", openai_model, anthropic_model, attack),
)
def _wait_and_request(candidate: Gateway, gate: threading.Event, spec: CallSpec) -> CallResult:
assert gate.wait(timeout=45), f"{spec.phase} request gate was not released"
status, text, headers = _request(candidate, spec.path, spec.body, stream=spec.stream)
return CallResult(spec, status, text, headers)
def test_h20_yaml_default_on_scans_tuple_attack_and_benign(
gateway: Gateway, tmp_path: Path
) -> None:
with dispatch_rig(gateway, tmp_path, default_on=True) as (owned, azure, provider, _):
candidate: Final = owned.gateway
with candidate.scenario() as scenario:
model: Final = _model(scenario, provider)
attack: Final = f"default-on {ATTACK_MARKER} {uuid.uuid4().hex}"
blocked: Final = candidate.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": attack}],
"cache": {"no-cache": True},
},
)
assert blocked.status_code == 400, blocked.text
assert azure_texts(azure) == (attack,), blocked.text
assert provider.drain() == (), blocked.text
benign: Final = f"default-on benign {uuid.uuid4().hex}"
allowed: Final = candidate.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": benign}],
"cache": {"no-cache": True},
},
)
assert allowed.status_code == 200, allowed.text
assert azure_texts(azure) == (benign,), allowed.text
assert provider_texts(provider) == (benign,), allowed.text
_assert_spend(allowed.headers["x-litellm-call-id"], allowed.text)
def test_x1_azure_server_stop_restart_during_mixed_burst(
gateway: Gateway, tmp_path: Path
) -> None:
with ExitStack() as resources:
redis: Final = resources.enter_context(owned_redis(tmp_path))
provider: Final = resources.enter_context(wire_server(provider_handler))
with socket.socket() as reservation:
reservation.bind(("127.0.0.1", 0))
azure_port: Final = int(reservation.getsockname()[1])
first_azure_lifetime: Final = ExitStack()
try:
azure: Final = first_azure_lifetime.enter_context(wire_server(azure_handler(), port=azure_port))
owned: Final = resources.enter_context(dispatch_proxy(gateway, tmp_path, redis, azure.url, workers=2))
candidate: Final = owned.gateway
with candidate.scenario() as scenario:
openai_model: Final = _model(scenario, provider)
anthropic_model: Final = _anthropic_model(scenario, provider)
up_specs: Final = _phase_specs("up", openai_model, anthropic_model, True)
down_specs: Final = _phase_specs("down", openai_model, anthropic_model, True)
recovery_specs: Final = _phase_specs("recovery", openai_model, anthropic_model, False)
up_gate: Final = threading.Event()
down_gate: Final = threading.Event()
recovery_gate: Final = threading.Event()
specs: Final = (*up_specs, *down_specs, *recovery_specs)
gates: Final = (
*((up_gate,) * len(up_specs)),
*((down_gate,) * len(down_specs)),
*((recovery_gate,) * len(recovery_specs)),
)
with concurrent.futures.ThreadPoolExecutor(max_workers=30) as executor:
futures: Final = tuple(
executor.submit(_wait_and_request, candidate, gate, spec)
for gate, spec in zip(gates, specs, strict=True)
)
up_gate.set()
up_results: Final = tuple(future.result(timeout=70) for future in futures[:10])
assert all(result.status == 400 for result in up_results), tuple(
result.text for result in up_results
)
first_azure_texts: Final = azure_texts(azure)
first_azure_lifetime.close()
down_gate.set()
down_results: Final = tuple(future.result(timeout=70) for future in futures[10:20])
assert all(result.status != 200 for result in down_results), tuple(
result.text for result in down_results
)
restarted_azure_lifetime: Final = ExitStack()
try:
restarted_azure: Final = restarted_azure_lifetime.enter_context(
wire_server(azure_handler(), port=azure_port)
)
recovery_gate.set()
recovery_results: Final = tuple(future.result(timeout=70) for future in futures[20:])
assert all(result.status == 200 for result in recovery_results), tuple(
result.text for result in recovery_results
)
restarted_texts: Final = azure_texts(restarted_azure)
assert len(restarted_texts) == len(recovery_results), recovery_results[-1].text
assert set(restarted_texts) == {result.spec.prompt for result in recovery_results}, (
recovery_results[-1].text
)
finally:
restarted_azure_lifetime.close()
assert set(first_azure_texts) == {result.spec.prompt for result in up_results}, up_results[0].text
provider_requests: Final = provider_texts(provider)
assert all(result.spec.prompt not in provider_requests for result in down_results), (
down_results[0].text
)
assert set(provider_requests) == {result.spec.prompt for result in recovery_results}, (
recovery_results[-1].text
)
for result in (*up_results, *down_results, *recovery_results):
if result.status == 200:
_assert_spend(result.headers["x-litellm-call-id"], result.text)
finally:
first_azure_lifetime.close()
def test_x2_slow_azure_edge_completes_twenty_concurrent_requests(
gateway: Gateway, tmp_path: Path
) -> None:
with dispatch_rig(gateway, tmp_path, behavior=AzureBehavior(delay_seconds=0.2)) as (
owned,
azure,
provider,
_,
):
candidate: Final = owned.gateway
with candidate.scenario() as scenario:
model: Final = _model(scenario, provider)
prompts: Final = tuple(f"slow edge {uuid.uuid4().hex}" for _ in range(20))
with concurrent.futures.ThreadPoolExecutor(max_workers=20) as executor:
futures: Final = tuple(
executor.submit(
candidate.request,
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": prompt}],
"guardrails": ["shield"],
"cache": {"no-cache": True},
},
)
for prompt in prompts
)
responses: Final = tuple(future.result(timeout=70) for future in futures)
assert all(response.status_code == 200 for response in responses), tuple(
response.text for response in responses
)
scanned_prompts: Final = azure_texts(azure)
assert len(scanned_prompts) == len(prompts), responses[-1].text
assert set(scanned_prompts) == set(prompts), responses[-1].text
assert set(provider_texts(provider)) == set(prompts), responses[-1].text
for response in responses:
_assert_spend(response.headers["x-litellm-call-id"], response.text)
def test_x3_killing_one_worker_leaves_the_other_serving(
gateway: Gateway, tmp_path: Path
) -> None:
entered: Final = threading.Event()
with dispatch_rig(gateway, tmp_path, behavior=AzureBehavior(delay_seconds=0.15, entered=entered)) as (
owned,
azure,
provider,
_,
):
candidate: Final = owned.gateway
workers: Final = _worker_processes(owned)
assert len(workers) == 2, tuple(process.cmdline() for process in workers)
with candidate.scenario() as scenario:
model: Final = _model(scenario, provider)
with concurrent.futures.ThreadPoolExecutor(max_workers=12) as executor:
inflight: Final = tuple(
executor.submit(_safe_request, candidate, model, f"worker in-flight {uuid.uuid4().hex}")
for _ in range(4)
)
assert entered.wait(timeout=30), "No request reached the slow Azure edge"
killed_worker: Final = workers[0]
survivor: Final = workers[1]
killed_worker.kill()
killed_worker.wait(timeout=10)
assert survivor.is_running(), "The second proxy worker exited with the killed worker"
prompts: Final = tuple(f"worker survivor {uuid.uuid4().hex}" for _ in range(8))
futures: Final = tuple(
executor.submit(_safe_request, candidate, model, prompt)
for prompt in prompts
)
responses: Final = tuple(future.result(timeout=70) for future in (*inflight, *futures))
surviving_responses: Final = tuple(response for response in responses[4:] if response is not None)
assert all(response.status_code == 200 for response in surviving_responses), tuple(
response.text for response in surviving_responses
)
assert len(surviving_responses) == len(prompts), tuple(
response.text for response in surviving_responses if response is not None
)
assert set(azure_texts(azure)).issuperset(prompts), surviving_responses[-1].text
assert set(provider_texts(provider)).issuperset(prompts), surviving_responses[-1].text
for response in surviving_responses:
_assert_spend(response.headers["x-litellm-call-id"], response.text)
def test_x4_proxy_restart_preserves_completed_spend_rows(
gateway: Gateway, tmp_path: Path
) -> None:
entered: Final = threading.Event()
arrived: Final = threading.Semaphore(0)
release: Final = threading.Event()
behavior: Final = AzureBehavior(entered=entered, arrived=arrived, release=release, barrier_marker="X4_WAIT")
with ExitStack() as resources:
redis: Final = resources.enter_context(owned_redis(tmp_path))
azure: Final = resources.enter_context(wire_server(azure_handler(behavior)))
provider: Final = resources.enter_context(wire_server(provider_handler))
first_proxy_lifetime: Final = ExitStack()
try:
first: Final = first_proxy_lifetime.enter_context(
dispatch_proxy(gateway, tmp_path, redis, azure.url, workers=2)
)
with gateway.scenario() as scenario:
model: Final = _model(scenario, provider)
completed_prompts: Final = tuple(f"X4 completed {uuid.uuid4().hex}" for _ in range(5))
completed: Final = tuple(
first.gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": prompt}],
"guardrails": ["shield"],
"cache": {"no-cache": True},
},
)
for prompt in completed_prompts
)
assert all(response.status_code == 200 for response in completed), tuple(
response.text for response in completed
)
for response in completed:
_assert_spend(response.headers["x-litellm-call-id"], response.text)
with concurrent.futures.ThreadPoolExecutor(max_workers=10) as executor:
interrupted_prompts: Final = tuple(f"X4_WAIT {uuid.uuid4().hex}" for _ in range(10))
interrupted: Final = tuple(
executor.submit(_safe_request, first.gateway, model, prompt)
for prompt in interrupted_prompts
)
arrived_count: Final = sum(arrived.acquire(timeout=30) for _ in interrupted_prompts)
assert arrived_count == len(interrupted_prompts), (
f"Only {arrived_count} of {len(interrupted_prompts)} interrupted requests reached Azure"
)
first.process.terminate()
release.set()
first_proxy_lifetime.close()
interrupted_results: Final = tuple(future.result(timeout=70) for future in interrupted)
second_proxy_lifetime: Final = ExitStack()
try:
second: Final = second_proxy_lifetime.enter_context(
dispatch_proxy(gateway, tmp_path, redis, azure.url, workers=2)
)
recovery_prompt: Final = f"X4 restarted benign {uuid.uuid4().hex}"
recovered: Final = second.gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": recovery_prompt}],
"guardrails": ["shield"],
"cache": {"no-cache": True},
},
)
assert recovered.status_code == 200, recovered.text
for response in completed:
_assert_spend(response.headers["x-litellm-call-id"], response.text)
_assert_spend(recovered.headers["x-litellm-call-id"], recovered.text)
azure_received: Final = azure_texts(azure)
azure_prompts: Final = frozenset(azure_received)
provider_received: Final = provider_texts(provider)
provider_prompts: Final = frozenset(provider_received)
expected_provider_prompts: Final = frozenset((*completed_prompts, recovery_prompt))
expected_edge_prompts: Final = tuple(
sorted((*completed_prompts, *interrupted_prompts, recovery_prompt))
)
assert tuple(sorted(azure_received)) == expected_edge_prompts, recovered.text
assert provider_prompts <= azure_prompts, recovered.text
assert expected_provider_prompts <= provider_prompts, recovered.text
for prompt, response in zip(interrupted_prompts, interrupted_results, strict=True):
if response is not None and response.status_code == 200:
assert prompt in provider_prompts, response.text
finally:
second_proxy_lifetime.close()
finally:
release.set()
first_proxy_lifetime.close()

View file

@ -0,0 +1,351 @@
import asyncio
import json
from collections.abc import Iterator
from contextlib import ExitStack
from typing import Final
import litellm
import pytest
from fastapi import HTTPException
from integration._support.client import object_value
from integration._support.wire import Reply, Request, Wire, wire_server
from litellm import Router
from litellm.proxy.guardrails.guardrail_hooks.azure.prompt_shield import (
AzureContentSafetyPromptShieldGuardrail,
)
_ATTACK_MARKER: Final = "synthetic-sdk-attack-marker"
_ATTACK_PROMPT: Final = f"synthetic sdk prompt {_ATTACK_MARKER}"
_SHIELD_TARGET_PREFIX: Final = "/contentsafety/text:shieldPrompt?api-version="
_PROVIDER_KEY: Final = "synthetic-provider-key"
_AZURE_KEY: Final = "synthetic-azure-key"
_GUARDRAIL_NAME: Final = "sdk-azure-shield"
def _azure(request: Request) -> Reply:
assert request.method == "POST", request.method
assert request.target.startswith(_SHIELD_TARGET_PREFIX), request.target
body: Final = object_value(json.loads(request.body))
prompt: Final = body["userPrompt"]
assert isinstance(prompt, str), body
return Reply(
body=json.dumps(
{
"userPromptAnalysis": {"attackDetected": _ATTACK_MARKER in prompt},
"documentsAnalysis": [],
}
).encode()
)
def _provider(request: Request) -> Reply:
assert request.method == "POST", request.method
assert request.target == "/v1/chat/completions", request.target
if b'"stream":true' in request.body.replace(b" ", b""):
chunk: Final = {
"id": "chatcmpl-sdk-azure-guardrail",
"object": "chat.completion.chunk",
"created": 1700000000,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "delta": {"content": "permitted response"}, "finish_reason": None}],
}
return Reply(
content_type="text/event-stream",
chunks=(b"data: " + json.dumps(chunk).encode() + b"\n\n", b"data: [DONE]\n\n"),
)
return Reply(
body=json.dumps(
{
"id": "chatcmpl-sdk-azure-guardrail",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "permitted response"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
}
).encode()
)
@pytest.fixture
def sdk_rig(
monkeypatch: pytest.MonkeyPatch,
) -> Iterator[tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire]]:
with ExitStack() as stack:
azure: Final = stack.enter_context(wire_server(_azure))
provider: Final = stack.enter_context(wire_server(_provider))
guardrail: Final = AzureContentSafetyPromptShieldGuardrail(
guardrail_name=_GUARDRAIL_NAME,
api_key=_AZURE_KEY,
api_base=azure.url,
event_hook="pre_call",
default_on=False,
)
monkeypatch.setattr(litellm, "callbacks", [guardrail])
monkeypatch.setattr(litellm, "success_callback", list(litellm.success_callback))
monkeypatch.setattr(litellm, "_async_success_callback", list(litellm._async_success_callback))
monkeypatch.setattr(litellm, "failure_callback", list(litellm.failure_callback))
monkeypatch.setattr(litellm, "_async_failure_callback", list(litellm._async_failure_callback))
try:
yield guardrail, azure, provider
finally:
litellm.logging_callback_manager.remove_callback_from_all_lists(guardrail)
def _azure_prompts(azure: Wire) -> tuple[str, ...]:
return tuple(_azure_prompt(request) for request in azure.drain())
def _azure_prompt(request: Request) -> str:
body: Final = object_value(json.loads(request.body))
prompt: Final = body["userPrompt"]
assert isinstance(prompt, str), body
return prompt
def _provider_prompts(provider: Wire) -> tuple[str, ...]:
return tuple(_provider_prompt(request) for request in provider.drain())
def _provider_prompt(request: Request) -> str:
body: Final = object_value(json.loads(request.body))
messages: Final = body["messages"]
assert isinstance(messages, list), body
prompt: Final = object_value(messages[-1])["content"]
assert isinstance(prompt, str), body
return prompt
def test_k1_acompletion_scans_tuple_messages(sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire]) -> None:
_, azure, provider = sdk_rig
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
asyncio.run(
litellm.acompletion(
model="openai/gpt-4o-mini",
api_base=provider.url + "/v1",
api_key=_PROVIDER_KEY,
messages=({"role": "user", "content": _ATTACK_PROMPT},),
guardrails=[_GUARDRAIL_NAME],
max_tokens=8,
)
)
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
assert provider.drain() == ()
def test_k1_acompletion_scans_system_and_user_tuple_messages(
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire],
) -> None:
_, azure, provider = sdk_rig
messages: Final = (
{"role": "system", "content": "system context"},
{"role": "user", "content": _ATTACK_PROMPT},
)
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
asyncio.run(
litellm.acompletion(
model="openai/gpt-4o-mini",
api_base=provider.url + "/v1",
api_key=_PROVIDER_KEY,
messages=messages,
guardrails=[_GUARDRAIL_NAME],
max_tokens=8,
)
)
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
assert provider.drain() == ()
def test_k1_acompletion_stream_scans_tuple_messages(
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire],
) -> None:
_, azure, provider = sdk_rig
async def consume() -> None:
stream: Final = await litellm.acompletion(
model="openai/gpt-4o-mini",
api_base=provider.url + "/v1",
api_key=_PROVIDER_KEY,
messages=({"role": "user", "content": _ATTACK_PROMPT},),
guardrails=[_GUARDRAIL_NAME],
max_tokens=8,
stream=True,
)
async for _chunk in stream:
pass
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
asyncio.run(consume())
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
assert provider.drain() == ()
def test_k1_acompletion_list_messages_control(
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire]
) -> None:
_, azure, provider = sdk_rig
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
asyncio.run(
litellm.acompletion(
model="openai/gpt-4o-mini",
api_base=provider.url + "/v1",
api_key=_PROVIDER_KEY,
messages=[{"role": "user", "content": _ATTACK_PROMPT}],
guardrails=[_GUARDRAIL_NAME],
max_tokens=8,
)
)
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
assert provider.drain() == ()
def _router(provider: Wire) -> Router:
return Router(
model_list=[
{
"model_name": "sdk-guardrail-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_base": provider.url + "/v1",
"api_key": _PROVIDER_KEY,
},
}
]
)
def test_k3_router_acompletion_guardrails_kwarg_scans_tuple_messages(
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire],
) -> None:
_, azure, provider = sdk_rig
router: Final = _router(provider)
try:
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
asyncio.run(
router.acompletion(
model="sdk-guardrail-model",
messages=({"role": "user", "content": _ATTACK_PROMPT},),
guardrails=[_GUARDRAIL_NAME],
max_tokens=8,
)
)
finally:
router.reset()
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
assert provider.drain() == ()
def test_k3_router_acompletion_guardrails_kwarg_list_control(
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire],
) -> None:
_, azure, provider = sdk_rig
router: Final = _router(provider)
try:
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
asyncio.run(
router.acompletion(
model="sdk-guardrail-model",
messages=[{"role": "user", "content": _ATTACK_PROMPT}],
guardrails=[_GUARDRAIL_NAME],
max_tokens=8,
)
)
finally:
router.reset()
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
assert provider.drain() == ()
def test_k2_litellm_completion_tuple_behavior_pinned_from_base(
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire],
) -> None:
_, azure, provider = sdk_rig
expected_block: Final = False
messages: Final = ({"role": "user", "content": _ATTACK_PROMPT},)
if expected_block:
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
litellm.completion(
model="openai/gpt-4o-mini",
api_base=provider.url + "/v1",
api_key=_PROVIDER_KEY,
messages=messages,
guardrails=[_GUARDRAIL_NAME],
max_tokens=8,
)
else:
response: Final = litellm.completion(
model="openai/gpt-4o-mini",
api_base=provider.url + "/v1",
api_key=_PROVIDER_KEY,
messages=messages,
guardrails=[_GUARDRAIL_NAME],
max_tokens=8,
)
assert response.choices
assert _azure_prompts(azure) == ((_ATTACK_PROMPT,) if expected_block else ())
if expected_block:
assert provider.drain() == ()
else:
assert _provider_prompts(provider) == (_ATTACK_PROMPT,)
def _router_with_guardrails(provider: Wire) -> Router:
return Router(
model_list=[
{
"model_name": "sdk-guardrail-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_base": provider.url + "/v1",
"api_key": _PROVIDER_KEY,
"guardrails": [_GUARDRAIL_NAME],
},
}
]
)
def test_k4_router_deployment_guardrails_scan_tuple_messages(
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire],
) -> None:
_, azure, provider = sdk_rig
router: Final = _router_with_guardrails(provider)
try:
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
asyncio.run(
router.acompletion(
model="sdk-guardrail-model",
messages=({"role": "user", "content": _ATTACK_PROMPT},),
max_tokens=8,
)
)
finally:
router.reset()
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
assert provider.drain() == ()
def test_k4_router_deployment_guardrails_scan_list_messages(
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire],
) -> None:
_, azure, provider = sdk_rig
router: Final = _router_with_guardrails(provider)
try:
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
asyncio.run(
router.acompletion(
model="sdk-guardrail-model",
messages=[{"role": "user", "content": _ATTACK_PROMPT}],
max_tokens=8,
)
)
finally:
router.reset()
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
assert provider.drain() == ()

View file

@ -1,15 +1,21 @@
from typing import Final
from typing import Final, cast
from unittest.mock import Mock, patch
import httpx
import pytest
from fastapi import HTTPException
from pydantic import JsonValue, TypeAdapter
from litellm import DualCache
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.azure.prompt_shield import (
AzureContentSafetyPromptShieldGuardrail,
)
from litellm.proxy.guardrails.guardrail_registry import InMemoryGuardrailHandler
from litellm.types.guardrails import LitellmParams
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import CallTypesLiteral
@pytest.mark.asyncio
@ -274,11 +280,18 @@ def _shield_response(attack_detected):
return response
def _shield_guardrail():
def _shield_guardrail(api_base: str = "azure_prompt_shield_api_base"):
return AzureContentSafetyPromptShieldGuardrail(
guardrail_name="azure_prompt_shield",
api_key="azure_prompt_shield_api_key",
api_base="azure_prompt_shield_api_base",
api_base=api_base,
)
def _shield_http_response(attack_detected: bool) -> httpx.Response:
return httpx.Response(
200,
json={"userPromptAnalysis": {"attackDetected": attack_detected}, "documentsAnalysis": []},
)
@ -359,6 +372,139 @@ def _recorded_guardrail_info(container):
return entries[0]
@pytest.mark.asyncio
async def test_prompt_shield_scans_tuple_messages() -> None:
guardrail: Final = _shield_guardrail("https://azure-content-safety.example")
prompt: Final = "synthetic tuple prompt"
data: Final[dict[str, object]] = {"messages": ({"role": "user", "content": prompt},)}
azure_response: Final = _shield_http_response(False)
azure_http_handler: Final = AsyncHTTPHandler(transport=httpx.MockTransport(lambda _request: azure_response))
guardrail.async_handler = azure_http_handler
try:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
cache=DualCache(),
data=data,
call_type="completion",
)
request_body: Final = TypeAdapter(dict[str, JsonValue]).validate_json(azure_response.request.read())
finally:
await azure_http_handler.close()
assert request_body["userPrompt"] == prompt
@pytest.mark.asyncio
async def test_prompt_shield_dispatches_to_subclass_get_user_prompt_override() -> None:
class AllTurnsPromptShield(AzureContentSafetyPromptShieldGuardrail):
def get_user_prompt(self, messages: list[AllMessageValues]) -> str:
return "\n".join(
message["content"]
for message in messages
if isinstance(message, dict)
and message.get("role") == "user"
and isinstance(message.get("content"), str)
)
guardrail: Final = AllTurnsPromptShield(
guardrail_name="azure_prompt_shield",
api_key="azure_prompt_shield_api_key",
api_base="https://azure-content-safety.example",
)
first_prompt: Final = "synthetic first user turn"
expected_prompt: Final = first_prompt + "\nbenign final user turn"
data: Final[dict[str, object]] = {
"messages": [
{"role": "user", "content": first_prompt},
{"role": "assistant", "content": "ok"},
{"role": "user", "content": "benign final user turn"},
]
}
azure_response: Final = _shield_http_response(False)
azure_http_handler: Final = AsyncHTTPHandler(transport=httpx.MockTransport(lambda _request: azure_response))
guardrail.async_handler = azure_http_handler
try:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
cache=DualCache(),
data=data,
call_type="completion",
)
request_body: Final = TypeAdapter(dict[str, JsonValue]).validate_json(azure_response.request.read())
finally:
await azure_http_handler.close()
assert request_body["userPrompt"] == expected_prompt
@pytest.mark.asyncio
async def test_prompt_shield_subclass_can_call_get_user_prompt() -> None:
class RequiringPromptShield(AzureContentSafetyPromptShieldGuardrail):
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
cache: DualCache,
data: dict[str, object],
call_type: CallTypesLiteral,
) -> dict[str, object] | None:
messages: Final = cast(list[AllMessageValues], data["messages"]) # cast-ok: chat input
user_prompt: Final = self.get_user_prompt(messages)
assert user_prompt
return await super().async_pre_call_hook(user_api_key_dict, cache, data, call_type)
guardrail: Final = RequiringPromptShield(
guardrail_name="azure_prompt_shield",
api_key="azure_prompt_shield_api_key",
api_base="https://azure-content-safety.example",
)
prompt: Final = "synthetic direct method prompt"
data: Final[dict[str, object]] = {"messages": [{"role": "user", "content": prompt}]}
azure_response: Final = _shield_http_response(False)
azure_http_handler: Final = AsyncHTTPHandler(transport=httpx.MockTransport(lambda _request: azure_response))
guardrail.async_handler = azure_http_handler
try:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
cache=DualCache(),
data=data,
call_type="completion",
)
request_body: Final = TypeAdapter(dict[str, JsonValue]).validate_json(azure_response.request.read())
finally:
await azure_http_handler.close()
assert request_body["userPrompt"] == prompt
@pytest.mark.asyncio
async def test_prompt_shield_messages_less_embeddings_return_data_and_log_allow() -> None:
guardrail: Final = _shield_guardrail()
data: Final[dict[str, object]] = {"input": "synthetic embedding input", "metadata": {}}
def fail_on_azure_request(_request: httpx.Request) -> httpx.Response:
raise AssertionError("unexpected Azure request")
azure_http_handler: Final = AsyncHTTPHandler(transport=httpx.MockTransport(fail_on_azure_request))
guardrail.async_handler = azure_http_handler
try:
result: Final = await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
cache=DualCache(),
data=data,
call_type="embedding",
)
finally:
await azure_http_handler.close()
entry: Final = _recorded_guardrail_info(data)
assert entry["guardrail_response"] == "allow"
assert result is data
@pytest.mark.parametrize(
("responses_input", "expected_prompt"),
[

View file

@ -1,16 +1,21 @@
import logging
from typing import Final
from typing import Final, cast
from unittest.mock import Mock, patch
import httpx
import pytest
from fastapi import HTTPException
from pydantic import JsonValue, TypeAdapter
from litellm import DualCache
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import UserAPIKeyAuth
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
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import CallTypesLiteral, Choices, Message, ModelResponse
@pytest.mark.asyncio
@ -494,14 +499,161 @@ def _moderation_response(severity):
return response
def _moderation_guardrail():
def _moderation_guardrail(api_base: str = "azure_text_moderation_api_base"):
return AzureContentSafetyTextModerationGuardrail(
guardrail_name="azure_text_moderation",
api_key="azure_text_moderation_api_key",
api_base="azure_text_moderation_api_base",
api_base=api_base,
)
def _moderation_http_response(severity: int) -> httpx.Response:
return httpx.Response(
200,
json={"blocklistsMatch": [], "categoriesAnalysis": [{"category": "Hate", "severity": severity}]},
)
def _standard_guardrail_entry(data: dict[str, object]) -> dict[str, JsonValue]:
metadata: Final = TypeAdapter(dict[str, JsonValue]).validate_python(data["metadata"])
entries: Final = metadata["standard_logging_guardrail_information"]
assert isinstance(entries, list) and len(entries) == 1
return TypeAdapter(dict[str, JsonValue]).validate_python(entries[0])
@pytest.mark.asyncio
async def test_text_moderation_scans_tuple_messages() -> None:
guardrail: Final = _moderation_guardrail("https://azure-content-safety.example")
prompt: Final = "synthetic tuple prompt"
data: Final[dict[str, object]] = {"messages": ({"role": "user", "content": prompt},)}
azure_response: Final = _moderation_http_response(0)
azure_http_handler: Final = AsyncHTTPHandler(transport=httpx.MockTransport(lambda _request: azure_response))
guardrail.async_handler = azure_http_handler
try:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
cache=DualCache(),
data=data,
call_type="completion",
)
request_body: Final = TypeAdapter(dict[str, JsonValue]).validate_json(azure_response.request.read())
finally:
await azure_http_handler.close()
assert request_body["text"] == prompt
@pytest.mark.asyncio
async def test_text_moderation_dispatches_to_subclass_get_user_prompt_override() -> None:
class AllTurnsTextModeration(AzureContentSafetyTextModerationGuardrail):
def get_user_prompt(self, messages: list[AllMessageValues]) -> str:
return "\n".join(
message["content"]
for message in messages
if isinstance(message, dict)
and message.get("role") == "user"
and isinstance(message.get("content"), str)
)
guardrail: Final = AllTurnsTextModeration(
guardrail_name="azure_text_moderation",
api_key="azure_text_moderation_api_key",
api_base="https://azure-content-safety.example",
)
first_prompt: Final = "synthetic first user turn"
expected_prompt: Final = first_prompt + "\nbenign final user turn"
data: Final[dict[str, object]] = {
"messages": [
{"role": "user", "content": first_prompt},
{"role": "assistant", "content": "ok"},
{"role": "user", "content": "benign final user turn"},
]
}
azure_response: Final = _moderation_http_response(0)
azure_http_handler: Final = AsyncHTTPHandler(transport=httpx.MockTransport(lambda _request: azure_response))
guardrail.async_handler = azure_http_handler
try:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
cache=DualCache(),
data=data,
call_type="completion",
)
request_body: Final = TypeAdapter(dict[str, JsonValue]).validate_json(azure_response.request.read())
finally:
await azure_http_handler.close()
assert request_body["text"] == expected_prompt
@pytest.mark.asyncio
async def test_text_moderation_subclass_can_call_get_user_prompt() -> None:
class RequiringTextModeration(AzureContentSafetyTextModerationGuardrail):
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
cache: DualCache,
data: dict[str, object],
call_type: CallTypesLiteral,
) -> dict[str, object] | None:
messages: Final = cast(list[AllMessageValues], data["messages"]) # cast-ok: chat input
user_prompt: Final = self.get_user_prompt(messages)
assert user_prompt
return await super().async_pre_call_hook(user_api_key_dict, cache, data, call_type)
guardrail: Final = RequiringTextModeration(
guardrail_name="azure_text_moderation",
api_key="azure_text_moderation_api_key",
api_base="https://azure-content-safety.example",
)
prompt: Final = "synthetic direct method prompt"
data: Final[dict[str, object]] = {"messages": [{"role": "user", "content": prompt}]}
azure_response: Final = _moderation_http_response(0)
azure_http_handler: Final = AsyncHTTPHandler(transport=httpx.MockTransport(lambda _request: azure_response))
guardrail.async_handler = azure_http_handler
try:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
cache=DualCache(),
data=data,
call_type="completion",
)
request_body: Final = TypeAdapter(dict[str, JsonValue]).validate_json(azure_response.request.read())
finally:
await azure_http_handler.close()
assert request_body["text"] == prompt
@pytest.mark.asyncio
async def test_text_moderation_messages_less_embeddings_return_data_and_log_allow() -> None:
guardrail: Final = _moderation_guardrail()
data: Final[dict[str, object]] = {"input": "synthetic embedding input", "metadata": {}}
def fail_on_azure_request(_request: httpx.Request) -> httpx.Response:
raise AssertionError("unexpected Azure request")
azure_http_handler: Final = AsyncHTTPHandler(transport=httpx.MockTransport(fail_on_azure_request))
guardrail.async_handler = azure_http_handler
try:
result: Final = await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
cache=DualCache(),
data=data,
call_type="embedding",
)
finally:
await azure_http_handler.close()
entry: Final = _standard_guardrail_entry(data)
assert entry["guardrail_response"] == "allow"
assert result is data
@pytest.mark.asyncio
async def test_apply_guardrail_scans_every_text():
"""/guardrails/apply_guardrail reaches this method directly. Inheriting the base