diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py index 830d125e8ea..597722eb1fb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py index 4312cc283a2..6a9c5aa4fbb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py @@ -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: diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py b/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py index d5d9fec8ff8..d9147cfb62b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py @@ -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: diff --git a/tests/integration/observability/azure_dispatch_support.py b/tests/integration/observability/azure_dispatch_support.py new file mode 100644 index 00000000000..2f1d758ead4 --- /dev/null +++ b/tests/integration/observability/azure_dispatch_support.py @@ -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 diff --git a/tests/integration/observability/test_azure_content_safety_dispatch.py b/tests/integration/observability/test_azure_content_safety_dispatch.py new file mode 100644 index 00000000000..d6cf5e7996e --- /dev/null +++ b/tests/integration/observability/test_azure_content_safety_dispatch.py @@ -0,0 +1,1370 @@ +import asyncio +import concurrent.futures +import json +import threading +import uuid +from collections.abc import Iterator +from pathlib import Path +from typing import Final + +import anthropic +import httpx +import openai +import pytest +from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, object_value +from integration._support.database import read_rows +from integration._support.wire import Wire +from integration.observability.azure_dispatch_support import ( + ATTACK_MARKER, + AzureBehavior, + MODERATION_MARKER, + OVERSIZED_MARKER, + PROVIDER_401_MARKER, + azure_texts, + dispatch_rig, + provider_messages, + provider_texts, +) +from pydantic import JsonValue + +_ATTACK_MARKER: Final = ATTACK_MARKER +_MODERATION_MARKER: Final = MODERATION_MARKER + + +@pytest.fixture(scope="module") +def azure_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[tuple[Gateway, Wire, Wire]]: + directory: Final = tmp_path_factory.mktemp("azure-content-safety-dispatch") + with gateway_from_environment() as gateway: + with dispatch_rig(gateway, directory) as (owned, azure, provider, _): + yield owned.gateway, azure, provider + + +@pytest.fixture +def key_update_rig(tmp_path: Path) -> Iterator[tuple[Gateway, Wire, Wire, AzureBehavior]]: + behavior: Final = AzureBehavior( + entered=threading.Event(), + release=threading.Event(), + barrier_marker="E3_BLOCK", + ) + with gateway_from_environment() as gateway: + with dispatch_rig(gateway, tmp_path, behavior=behavior) as (owned, azure, provider, _): + yield owned.gateway, azure, provider, behavior + + +@pytest.fixture(autouse=True) +def _clear_wires(azure_rig: tuple[Gateway, Wire, Wire]) -> None: + azure_rig[1].drain() + azure_rig[2].drain() + + +def _model(scenario: Scenario, provider: Wire, model: str = "openai/gpt-4o-mini") -> str: + return scenario.model( + model=model, + api_base=provider.url + "/v1", + api_key="synthetic-provider-key", + ) + + +def _guardrail_entry(response: httpx.Response) -> dict[str, JsonValue]: + request_id: Final = response.headers["x-litellm-call-id"] + rows: Final = eventually( + lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (request_id,)), + lambda values: len(values) == 1, + seconds=70, + ) + metadata_value: Final = rows[0]["metadata"] + metadata: Final = object_value(json.loads(metadata_value) if isinstance(metadata_value, str) else metadata_value) + entries: Final = metadata["guardrail_information"] + assert isinstance(entries, list) and len(entries) == 1, f"{response.text}: {metadata}" + return object_value(entries[0]) + + +def _spend_metadata(request_id: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (request_id,)), + lambda values: len(values) == 1, + seconds=70, + ) + metadata_value: Final = rows[0]["metadata"] + return object_value(json.loads(metadata_value) if isinstance(metadata_value, str) else metadata_value) + + +def _chat_body( + model: str, + prompt: str, + guardrails: tuple[str, ...], + *, + stream: bool = False, + no_cache: bool = False, +) -> dict[str, JsonValue]: + return { + "model": model, + "messages": [{"role": "user", "content": prompt}], + **({"guardrails": list(guardrails)} if guardrails else {}), + **({"stream": True} if stream else {}), + **({"cache": {"no-cache": True}} if no_cache else {}), + } + + +def _messages_body(model: str, prompt: str, guardrails: tuple[str, ...], *, stream: bool = False) -> dict[str, JsonValue]: + return { + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": prompt}], + **({"guardrails": list(guardrails)} if guardrails else {}), + **({"stream": True} if stream else {}), + } + + +def _responses_body(model: str, value: JsonValue, guardrails: tuple[str, ...], *, stream: bool = False) -> dict[str, JsonValue]: + return { + "model": model, + "input": value, + **({"guardrails": list(guardrails)} if guardrails else {}), + **({"stream": True} if stream else {}), + } + + +def _stream_request(candidate: Gateway, path: str, body: dict[str, JsonValue]) -> tuple[int, str, dict[str, str]]: + 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 _chat_stream_delta_content(event: dict[str, JsonValue]) -> str: + choices: Final = event["choices"] + assert isinstance(choices, list), event + return "".join( + _chat_stream_choice_content(choice) + for choice in choices + ) + + +def _chat_stream_choice_content(choice: JsonValue) -> str: + delta: Final = object_value(object_value(choice)["delta"]) + content: Final = delta.get("content") + return content if isinstance(content, str) else "" + + +def _chat_stream_content(response_text: str) -> tuple[str, bool]: + events: Final = tuple( + line.removeprefix("data: ") + for line in response_text.splitlines() + if line.startswith("data: ") + ) + content_events: Final = tuple(event for event in events if event != "[DONE]") + chunks: Final = tuple(object_value(json.loads(event)) for event in content_events) + content: Final = "".join(_chat_stream_delta_content(chunk) for chunk in chunks) + return content, bool(events) and events[-1] == "[DONE]" + + +def _openai_chat_sync( + base_url: str, key: str, model: str, prompt: str, guardrails: tuple[str, ...] +) -> None: + client: Final = openai.OpenAI(base_url=base_url, api_key=key, max_retries=0) + with client: + client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": prompt}], + max_tokens=8, + extra_body={"guardrails": list(guardrails)}, + ) + + +async def _openai_chat_async( + base_url: str, key: str, model: str, prompt: str, guardrails: tuple[str, ...] +) -> None: + client: Final = openai.AsyncOpenAI(base_url=base_url, api_key=key, max_retries=0) + async with client: + await client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": prompt}], + max_tokens=8, + extra_body={"guardrails": list(guardrails)}, + ) + + +def _anthropic_messages_sync(base_url: str, key: str, model: str, prompt: str) -> None: + client: Final = anthropic.Anthropic(base_url=base_url, api_key=key, max_retries=0) + with client: + client.messages.create( + model=model, + max_tokens=8, + messages=[{"role": "user", "content": prompt}], + extra_body={"guardrails": ["tuple-writer", "shield"]}, + ) + + +async def _anthropic_messages_async(base_url: str, key: str, model: str, prompt: str) -> None: + client: Final = anthropic.AsyncAnthropic(base_url=base_url, api_key=key, max_retries=0) + async with client: + await client.messages.create( + model=model, + max_tokens=8, + messages=[{"role": "user", "content": prompt}], + extra_body={"guardrails": ["tuple-writer", "shield"]}, + ) + + +def _openai_responses_sync(base_url: str, key: str, model: str, prompt: str) -> None: + client: Final = openai.OpenAI(base_url=base_url, api_key=key, max_retries=0) + with client: + client.responses.create( + model=model, + input=prompt, + extra_body={"guardrails": ["shield"]}, + ) + + +async def _openai_responses_async(base_url: str, key: str, model: str, prompt: str) -> None: + client: Final = openai.AsyncOpenAI(base_url=base_url, api_key=key, max_retries=0) + async with client: + await client.responses.create( + model=model, + input=prompt, + extra_body={"guardrails": ["shield"]}, + ) + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", _ATTACK_MARKER), ("moderation", _MODERATION_MARKER)], + ids=("H1-shield", "H1-moderation"), +) +def test_h1_tuple_attack_is_scanned_for_each_guardrail( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"synthetic prompt {marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": prompt}], + "guardrails": ["tuple-writer", guardrail_name], + }, + ) + assert response.status_code == 400, response.text + assert _azure_texts(azure) == (prompt,), response.text + assert provider.drain() == (), response.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", "benign shield tuple"), ("moderation", "benign moderation tuple")], + ids=("H2-shield", "H2-moderation"), +) +def test_h2_tuple_benign_is_scanned_and_spent( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"{marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _chat_body(model, prompt, ("tuple-writer", guardrail_name)), + ) + assert response.status_code == 200, response.text + assert _azure_texts(azure) == (prompt,), response.text + assert provider_texts(provider) == (prompt,), response.text + assert _guardrail_entry(response)["guardrail_status"] == "success", response.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", _ATTACK_MARKER), ("moderation", _MODERATION_MARKER)], + ids=("H8-shield-attack", "H8-moderation-attack"), +) +def test_h8_list_attack_control_is_scanned_for_each_guardrail( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"synthetic prompt {marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": prompt}], + "guardrails": [guardrail_name], + }, + ) + assert response.status_code == 400, response.text + assert _azure_texts(azure) == (prompt,), response.text + assert provider.drain() == (), response.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", "benign shield list"), ("moderation", "benign moderation list")], + ids=("H8-shield", "H8-moderation"), +) +def test_h8_list_benign_is_scanned_and_served( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"{marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _chat_body(model, prompt, (guardrail_name,)), + ) + assert response.status_code == 200, response.text + assert _azure_texts(azure) == (prompt,), response.text + assert provider_texts(provider) == (prompt,), response.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", _ATTACK_MARKER), ("moderation", _MODERATION_MARKER)], + ids=("H3-shield", "H3-moderation"), +) +def test_h3_tuple_attack_chat_stream_is_blocked( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"stream {marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + status, text, _ = _stream_request( + candidate, + "/v1/chat/completions", + _chat_body(model, prompt, ("tuple-writer", guardrail_name), stream=True), + ) + assert status == 400, text + assert _azure_texts(azure) == (prompt,), text + assert provider.drain() == (), text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", "benign stream shield"), ("moderation", "benign stream moderation")], + ids=("H4-shield", "H4-moderation"), +) +def test_h4_tuple_benign_chat_stream_reaches_provider_and_spend( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"{marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + status, text, headers = _stream_request( + candidate, + "/v1/chat/completions", + _chat_body(model, prompt, ("tuple-writer", guardrail_name), stream=True), + ) + assert status == 200, text + streamed_content, done = _chat_stream_content(text) + assert streamed_content == "permitted response", text + assert done, text + assert _azure_texts(azure) == (prompt,), text + assert provider_texts(provider) == (prompt,), text + spend_metadata: Final = _spend_metadata(headers["x-litellm-call-id"]) + entries: Final = spend_metadata["guardrail_information"] + assert isinstance(entries, list) and len(entries) == 1, text + assert object_value(entries[0])["guardrail_status"] == "success", text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", _ATTACK_MARKER), ("moderation", _MODERATION_MARKER)], + ids=("H5-shield", "H5-moderation"), +) +def test_h5_tuple_attack_anthropic_messages_is_blocked( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"anthropic message {marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=provider.url, + api_key="synthetic-provider-key", + ) + response: Final = candidate.request( + "POST", + "/v1/messages", + _messages_body(model, prompt, ("tuple-writer", guardrail_name)), + ) + assert response.status_code == 400, response.text + assert _azure_texts(azure) == (prompt,), response.text + assert provider.drain() == (), response.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", _ATTACK_MARKER), ("moderation", _MODERATION_MARKER)], + ids=("H6-shield", "H6-moderation"), +) +def test_h6_tuple_attack_anthropic_messages_stream_is_blocked( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"anthropic stream {marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=provider.url, + api_key="synthetic-provider-key", + ) + status, text, _ = _stream_request( + candidate, + "/v1/messages", + _messages_body(model, prompt, ("tuple-writer", guardrail_name), stream=True), + ) + assert status == 400, text + assert _azure_texts(azure) == (prompt,), text + assert provider.drain() == (), text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", _ATTACK_MARKER), ("moderation", _MODERATION_MARKER)], + ids=("H7-shield", "H7-moderation"), +) +def test_h7_list_attack_anthropic_messages_control( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"anthropic list {marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=provider.url, + api_key="synthetic-provider-key", + ) + response: Final = candidate.request( + "POST", + "/v1/messages", + _messages_body(model, prompt, (guardrail_name,)), + ) + assert response.status_code == 400, response.text + assert _azure_texts(azure) == (prompt,), response.text + assert provider.drain() == (), response.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", _ATTACK_MARKER), ("moderation", _MODERATION_MARKER)], + ids=("H9-shield", "H9-moderation"), +) +def test_h9_responses_string_attack_is_scanned( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"synthetic prompt {marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/responses", + _responses_body(model, prompt, (guardrail_name,)), + ) + assert response.status_code == 400, response.text + assert _azure_texts(azure) == (prompt,), response.text + assert provider.drain() == (), response.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", "benign responses stream shield"), ("moderation", "benign responses stream moderation")], + ids=("H9-shield-stream", "H9-moderation-stream"), +) +def test_h9_responses_benign_stream_is_scanned_and_served( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"{marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + status, text, headers = _stream_request( + candidate, + "/v1/responses", + _responses_body(model, prompt, (guardrail_name,), stream=True), + ) + assert status == 200, text + assert "permitted response" in text, text + assert _azure_texts(azure) == (prompt,), text + assert provider_texts(provider) == (prompt,), text + spend_metadata: Final = _spend_metadata(headers["x-litellm-call-id"]) + entries: Final = spend_metadata["guardrail_information"] + assert isinstance(entries, list) and len(entries) == 1, text + assert object_value(entries[0])["guardrail_status"] == "success", text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", _ATTACK_MARKER), ("moderation", _MODERATION_MARKER)], + ids=("H9-shield-list", "H9-moderation-list"), +) +def test_h9_responses_list_input_attack_control( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"responses list {marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/responses", + _responses_body( + model, + [{"role": "user", "content": [{"type": "input_text", "text": prompt}]}], + (guardrail_name,), + ), + ) + assert response.status_code == 400, response.text + assert _azure_texts(azure) == (prompt,), response.text + assert provider.drain() == (), response.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("all-turns-shield", _ATTACK_MARKER), ("all-turns-moderation", _MODERATION_MARKER)], + ids=("H10-shield", "H10-moderation"), +) +def test_h10_subclass_override_scans_every_user_turn( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + first_prompt: Final = f"synthetic prompt {marker} {uuid.uuid4().hex}" + expected_prompt: Final = first_prompt + "\nbenign final user turn" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [ + {"role": "user", "content": first_prompt}, + {"role": "assistant", "content": "ok"}, + {"role": "user", "content": "benign final user turn"}, + ], + "guardrails": [guardrail_name], + }, + ) + assert response.status_code == 400, response.text + assert _azure_texts(azure) == (expected_prompt,), response.text + assert provider.drain() == (), response.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("all-turns-shield", _ATTACK_MARKER), ("all-turns-moderation", _MODERATION_MARKER)], + ids=("H11-shield", "H11-moderation"), +) +def test_h11_tuple_subclass_override_scans_every_user_turn( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + first_prompt: Final = f"tuple override {marker} {uuid.uuid4().hex}" + expected_prompt: Final = first_prompt + "\nbenign final user turn" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [ + {"role": "user", "content": first_prompt}, + {"role": "assistant", "content": "ok"}, + {"role": "user", "content": "benign final user turn"}, + ], + "guardrails": ["tuple-writer", guardrail_name], + }, + ) + assert response.status_code == 400, response.text + assert _azure_texts(azure) == (expected_prompt,), response.text + assert provider.drain() == (), response.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("requiring-shield", "benign shield prompt"), ("requiring-moderation", "benign moderation prompt")], + ids=("H12-shield-benign", "H12-moderation-benign"), +) +def test_h12_guardrail_subclass_can_call_get_user_prompt( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"{marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": prompt}], + "guardrails": [guardrail_name], + }, + ) + assert response.status_code == 200, response.text + assert "permitted response" in response.text, response.text + assert _azure_texts(azure) == (prompt,), response.text + assert provider_texts(provider) == (prompt,), response.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("requiring-shield", _ATTACK_MARKER), ("requiring-moderation", _MODERATION_MARKER)], + ids=("H12-shield-attack", "H12-moderation-attack"), +) +def test_h12_subclass_call_blocks_attack( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"required method {marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _chat_body(model, prompt, (guardrail_name,)), + ) + assert response.status_code == 400, response.text + assert _azure_texts(azure) == (prompt,), response.text + assert provider.drain() == (), response.text + + +@pytest.mark.parametrize("guardrail_name", ["shield", "moderation"], ids=("H13-shield", "H13-moderation")) +def test_h13_messages_less_embeddings_log_allow_without_azure_request( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str +) -> None: + candidate, azure, provider = azure_rig + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider, model="openai/text-embedding-3-small") + response: Final = candidate.request( + "POST", + "/v1/embeddings", + {"model": model, "input": "synthetic benign embedding text", "guardrails": [guardrail_name]}, + ) + assert response.status_code == 200, response.text + assert _azure_texts(azure) == (), response.text + assert provider_texts(provider) == ("synthetic benign embedding text",), response.text + entry: Final = _guardrail_entry(response) + assert entry["guardrail_status"] == "success", response.text + assert entry["guardrail_response"] == "allow", response.text + + +@pytest.mark.parametrize("guardrail_name", ["shield", "moderation"], ids=("H14-shield", "H14-moderation")) +def test_h14_completions_without_messages_log_allow( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"synthetic completion input {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider, model="openai/gpt-3.5-turbo-instruct") + response: Final = candidate.request( + "POST", + "/v1/completions", + {"model": model, "prompt": prompt, "guardrails": [guardrail_name]}, + ) + assert response.status_code == 200, response.text + assert _azure_texts(azure) == (), response.text + assert provider_texts(provider) == (prompt,), response.text + entry: Final = _guardrail_entry(response) + assert entry["guardrail_status"] == "success", response.text + assert entry["guardrail_response"] == "allow", response.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", "cache benign shield"), ("moderation", "cache benign moderation")], + ids=("C1-shield", "C1-moderation"), +) +def test_c1_tuple_benign_cache_twins_scan_each_request( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"{marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + body: Final = _chat_body(model, prompt, ("tuple-writer", guardrail_name)) + first: Final = candidate.request("POST", "/v1/chat/completions", body) + second: Final = candidate.request("POST", "/v1/chat/completions", body) + assert first.status_code == 200, first.text + assert second.status_code == 200, second.text + assert _azure_texts(azure) == (prompt, prompt), second.text + assert provider_texts(provider) == (prompt,), second.text + assert _guardrail_entry(first)["guardrail_status"] == "success", first.text + assert _guardrail_entry(second)["guardrail_status"] == "success", second.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", ATTACK_MARKER), ("moderation", MODERATION_MARKER)], + ids=("C2-shield", "C2-moderation"), +) +def test_c2_tuple_attack_cache_twins_are_both_blocked( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"cache attack {marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + body: Final = _chat_body(model, prompt, ("tuple-writer", guardrail_name)) + first: Final = candidate.request("POST", "/v1/chat/completions", body) + second: Final = candidate.request("POST", "/v1/chat/completions", body) + assert first.status_code == 400, first.text + assert second.status_code == 400, second.text + assert _azure_texts(azure) == (prompt, prompt), second.text + assert provider.drain() == (), second.text + + +@pytest.mark.parametrize("guardrail_name", ["shield", "moderation"], ids=("C3-shield", "C3-moderation")) +def test_c3_embeddings_cache_twins_keep_allow_rows( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"cache embedding {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider, model="openai/text-embedding-3-small") + body: Final = {"model": model, "input": prompt, "guardrails": [guardrail_name]} + first: Final = candidate.request("POST", "/v1/embeddings", body) + second: Final = candidate.request("POST", "/v1/embeddings", body) + assert first.status_code == 200, first.text + assert second.status_code == 200, second.text + assert _azure_texts(azure) == (), second.text + assert provider_texts(provider) == (prompt,), second.text + assert _guardrail_entry(first)["guardrail_response"] == "allow", first.text + assert _guardrail_entry(second)["guardrail_response"] == "allow", second.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", ATTACK_MARKER), ("moderation", MODERATION_MARKER)], + ids=("C4-shield", "C4-moderation"), +) +def test_c4_list_attack_cache_twins_are_both_blocked( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"cache list {marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + body: Final = _chat_body(model, prompt, (guardrail_name,)) + first: Final = candidate.request("POST", "/v1/chat/completions", body) + second: Final = candidate.request("POST", "/v1/chat/completions", body) + assert first.status_code == 400, first.text + assert second.status_code == 400, second.text + assert _azure_texts(azure) == (prompt, prompt), second.text + assert provider.drain() == (), second.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", ATTACK_MARKER), ("moderation", MODERATION_MARKER)], + ids=("C5-shield", "C5-moderation"), +) +def test_c5_anthropic_tuple_attack_cache_twins_are_blocked( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"Anthropic cache {marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=provider.url, + api_key="synthetic-provider-key", + ) + body: Final = _messages_body(model, prompt, ("tuple-writer", guardrail_name)) + first: Final = candidate.request("POST", "/v1/messages", body) + second: Final = candidate.request("POST", "/v1/messages", body) + assert first.status_code == 400, first.text + assert second.status_code == 400, second.text + assert _azure_texts(azure) == (prompt, prompt), second.text + assert provider.drain() == (), second.text + + +@pytest.mark.parametrize("guardrail_name", ["shield", "moderation"], ids=("C6-shield", "C6-moderation")) +def test_c6_responses_benign_cache_twins_scan_each_request( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"Responses cache benign {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + body: Final = _responses_body(model, prompt, (guardrail_name,)) + first: Final = candidate.request("POST", "/v1/responses", body) + second: Final = candidate.request("POST", "/v1/responses", body) + assert first.status_code == 200, first.text + assert second.status_code == 200, second.text + assert _azure_texts(azure) == (prompt, prompt), second.text + assert provider_texts(provider) == (prompt,), second.text + + +def test_s1_integer_messages_fails_without_dispatching_edges(azure_rig: tuple[Gateway, Wire, Wire]) -> None: + candidate, azure, provider = azure_rig + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": 5, "guardrails": ["shield"]}, + ) + assert response.status_code == 500, response.text + assert "error" in response.json(), response.text + assert _azure_texts(azure) == (), response.text + assert provider.drain() == (), response.text + + +@pytest.mark.parametrize( + ("shape", "expected_status"), + [ + ("string", 400), + ("object", 200), + ("empty-list", 200), + ("oversized", 400), + ("null", 400), + ("missing", 400), + ], + ids=( + "S2-string", + "S2-object", + "S2-empty-list", + "S2-oversized", + "S2-null", + "S2-missing", + ), +) +def test_s2_malformed_messages_shapes_match_base_pin( + azure_rig: tuple[Gateway, Wire, Wire], shape: str, expected_status: int +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"shape control {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + messages: Final = { + "string": "malformed messages", + "object": {}, + "empty-list": [], + "oversized": "x" * 5000, + "null": None, + "missing": None, + }[shape] + body: Final = { + "model": model, + "messages": messages, + "guardrails": ["shield"], + } + raw: Final = json.dumps( + {"model": model, "guardrails": ["shield"]} + if shape == "missing" + else body + ) + response: Final = candidate.client.post( + "/v1/chat/completions", + content=raw, + headers={"Authorization": f"Bearer {candidate.key}", "Content-Type": "application/json"}, + ) + assert response.status_code == expected_status, f"{response.status_code}: {response.text}" + assert _azure_texts(azure) == (), response.text + provider.drain() + + +def test_s2d_duplicate_messages_key_matches_single_key_request( + azure_rig: tuple[Gateway, Wire, Wire], +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"duplicate messages key {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + body: Final = _chat_body(model, prompt, ("shield",)) + raw: Final = json.dumps(body) + raw_with_duplicate: Final = ( + raw[:-1] + ',"messages":' + json.dumps(body["messages"]) + "}" + ) + response: Final = candidate.client.post( + "/v1/chat/completions", + content=raw_with_duplicate, + headers={"Authorization": f"Bearer {candidate.key}", "Content-Type": "application/json"}, + ) + assert response.status_code == 200, f"{response.status_code}: {response.text}" + assert _azure_texts(azure) == (prompt,), response.text + assert provider_texts(provider) == (prompt,), response.text + + +@pytest.mark.parametrize("authorization", ["", "Bearer invalid-key"], ids=("S3-missing", "S3-invalid")) +def test_s3_authentication_rejects_before_guardrails( + azure_rig: tuple[Gateway, Wire, Wire], authorization: str +) -> None: + candidate, azure, provider = azure_rig + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _chat_body(model, f"auth {ATTACK_MARKER}", ("tuple-writer", "shield")), + headers={"Authorization": authorization}, + ) + assert response.status_code == 401, response.text + assert _azure_texts(azure) == (), response.text + assert provider.drain() == (), response.text + + +@pytest.mark.parametrize("status", [500, 403, 404], ids=("S4-500", "S4-403", "S4-404")) +@pytest.mark.parametrize("guardrail_name", ["shield", "moderation"], ids=("S4-shield", "S4-moderation")) +def test_s4_azure_error_fails_closed( + azure_rig: tuple[Gateway, Wire, Wire], status: int, guardrail_name: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"AZURE_{status} benign {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _chat_body(model, prompt, ("tuple-writer", guardrail_name)), + ) + assert response.status_code != 200, response.text + assert "synthetic Azure error" in response.text, response.text + assert _azure_texts(azure) == (prompt,), response.text + assert provider.drain() == (), response.text + + +@pytest.mark.parametrize("guardrail_name", ["shield", "moderation"], ids=("S5-shield", "S5-moderation")) +def test_s5_list_guardrail_azure_error_control( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"AZURE_500 list benign {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _chat_body(model, prompt, (guardrail_name,)), + ) + assert response.status_code != 200, response.text + assert "synthetic Azure error" in response.text, response.text + assert _azure_texts(azure) == (prompt,), response.text + assert provider.drain() == (), response.text + + +@pytest.mark.parametrize( + ("marker", "expected_status", "expected_message"), + [ + (PROVIDER_401_MARKER, 401, "synthetic provider unauthorized"), + (OVERSIZED_MARKER, 400, "synthetic context length exceeded"), + ], + ids=("S6-provider-401", "S6-oversized"), +) +def test_s6_provider_errors_follow_azure_scan( + azure_rig: tuple[Gateway, Wire, Wire], marker: str, expected_status: int, expected_message: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"provider error {marker} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _chat_body(model, prompt, ("tuple-writer", "shield")), + ) + assert response.status_code == expected_status, response.text + assert expected_message in response.text, response.text + assert _azure_texts(azure) == (prompt,), response.text + assert provider_texts(provider) == (prompt,), response.text + + +def test_s6_unknown_model_matches_base_pin(azure_rig: tuple[Gateway, Wire, Wire]) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"unknown model benign {uuid.uuid4().hex}" + model: Final = "openai/unknown-audit-model" + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _chat_body(model, prompt, ("tuple-writer", "shield")), + ) + assert response.status_code == 400, f"{response.status_code}: {response.text}" + assert "Invalid model name" in response.text, response.text + scanned: Final = _azure_texts(azure) + assert scanned == (prompt,), f"{response.text}: {scanned!r}" + assert provider_texts(provider) == (), response.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", "last user benign"), ("moderation", "last user benign")], + ids=("S7-shield-benign", "S7-moderation-benign"), +) +def test_s7_tuple_multi_item_text_scans_exact_last_user_block( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"{marker} {uuid.uuid4().hex}" + expected: Final = "part one " + prompt + messages: Final = [ + {"role": "system", "content": "system context"}, + {"role": "user", "content": "earlier user"}, + {"role": "assistant", "content": "assistant response"}, + { + "role": "user", + "content": [ + {"type": "text", "text": "part one "}, + {"type": "text", "text": prompt}, + ], + }, + ] + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": messages, + "guardrails": ["tuple-writer", guardrail_name], + }, + ) + assert response.status_code == 200, response.text + assert _azure_texts(azure) == (expected,), response.text + assert provider_messages(provider) == (messages,), response.text + + +@pytest.mark.parametrize( + ("guardrail_name", "marker"), + [("shield", ATTACK_MARKER), ("moderation", MODERATION_MARKER)], + ids=("S7-shield-attack", "S7-moderation-attack"), +) +def test_s7_tuple_multi_item_attack_is_blocked( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str, marker: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"multi part {marker} {uuid.uuid4().hex}" + expected: Final = "part one " + prompt + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [ + {"role": "system", "content": "system context"}, + {"role": "user", "content": "earlier benign"}, + {"role": "assistant", "content": "assistant response"}, + { + "role": "user", + "content": [ + {"type": "text", "text": "part one "}, + {"type": "text", "text": prompt}, + ], + }, + ], + "guardrails": ["tuple-writer", guardrail_name], + }, + ) + assert response.status_code == 400, response.text + assert _azure_texts(azure) == (expected,), response.text + assert provider.drain() == (), response.text + + +def test_s8_unknown_guardrail_matches_base_pin(azure_rig: tuple[Gateway, Wire, Wire]) -> None: + candidate, azure, provider = azure_rig + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _chat_body(model, "unknown guardrail control", ("unknown-guardrail",)), + ) + assert response.status_code == 200, f"{response.status_code}: {response.text}" + assert _azure_texts(azure) == (), response.text + assert provider_texts(provider) == ("unknown guardrail control",), response.text + + +def test_s9_guardrails_list_contains_azure_guardrails(azure_rig: tuple[Gateway, Wire, Wire]) -> None: + candidate, azure, provider = azure_rig + response: Final = candidate.request("GET", "/guardrails/list") + assert response.status_code == 200, response.text + assert "shield" in response.text and "moderation" in response.text, response.text + assert _azure_texts(azure) == (), response.text + assert provider.drain() == (), response.text + + +def test_s10_malformed_request_does_not_poison_proxy(azure_rig: tuple[Gateway, Wire, Wire]) -> None: + candidate, azure, provider = azure_rig + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + malformed: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": 5, "guardrails": ["shield"]}, + ) + assert malformed.status_code == 500, malformed.text + assert _azure_texts(azure) == (), malformed.text + assert provider.drain() == (), malformed.text + malformed_shape: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": "malformed messages", "guardrails": ["shield"]}, + ) + assert malformed_shape.status_code == 400, malformed_shape.text + assert _azure_texts(azure) == (), malformed_shape.text + provider.drain() + prompt: Final = f"post malformed benign {uuid.uuid4().hex}" + valid: Final = candidate.request( + "POST", + "/v1/chat/completions", + _chat_body(model, prompt, ("shield",)), + ) + health: Final = candidate.request("GET", "/health/liveliness") + assert valid.status_code == 200, valid.text + assert health.status_code == 200, health.text + assert _azure_texts(azure) == (prompt,), valid.text + assert provider_texts(provider) == (prompt,), valid.text + + +@pytest.mark.parametrize("guardrail_name", ["shield", "moderation"], ids=("E1-shield", "E1-moderation")) +def test_e1_empty_messages_matches_base_guardrail_response( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str +) -> None: + candidate, azure, provider = azure_rig + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider, model="openai/text-embedding-3-small") + response: Final = candidate.request( + "POST", + "/v1/embeddings", + { + "model": model, + "input": f"empty messages {uuid.uuid4().hex}", + "messages": [], + "guardrails": [guardrail_name], + }, + ) + assert response.status_code == 200, response.text + assert _azure_texts(azure) == (), response.text + entry: Final = _guardrail_entry(response) + assert entry["guardrail_response"] == {}, f"{response.text}: {entry!r}" + assert len(provider.drain()) == 1, response.text + + +def test_e2_key_and_request_guardrail_precedence_matches_base_pin( + azure_rig: tuple[Gateway, Wire, Wire], +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"precedence {ATTACK_MARKER} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + key: Final = scenario.key(guardrails=["shield"]) + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _chat_body(model, prompt, ("tuple-writer",)), + key=key, + ) + assert response.status_code == 400, f"{response.status_code}: {response.text}" + assert _azure_texts(azure) == (prompt,), response.text + assert provider.drain() == (), response.text + + +def test_e3_key_update_applies_guardrails_during_traffic( + key_update_rig: tuple[Gateway, Wire, Wire, AzureBehavior], +) -> None: + candidate, azure, provider, behavior = key_update_rig + with candidate.scenario() as scenario: + key: Final = scenario.key() + model: Final = _model(scenario, provider) + prompt: Final = f"key update {ATTACK_MARKER} {uuid.uuid4().hex}" + before: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": prompt}], + "cache": {"no-cache": True}, + }, + key=key, + ) + assert before.status_code == 200, before.text + assert _azure_texts(azure) == (), before.text + assert provider_texts(provider) == (prompt,), before.text + active_prompt: Final = f"E3_BLOCK benign {uuid.uuid4().hex}" + entered: Final = behavior.entered + release: Final = behavior.release + assert entered is not None and release is not None, "E3 barrier events were not configured" + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + active_request: Final = executor.submit( + candidate.request, + "POST", + "/v1/chat/completions", + _chat_body(model, active_prompt, ("shield",), no_cache=True), + key=key, + ) + try: + assert entered.wait(timeout=30), "Concurrent request did not reach the Azure responder" + updated: Final = candidate.request( + "POST", + "/key/update", + {"key": key, "guardrails": ["tuple-writer", "shield"]}, + ) + assert updated.status_code == 200, updated.text + finally: + release.set() + active_response: Final = active_request.result(timeout=60) + assert active_response.status_code == 200, active_response.text + assert _azure_texts(azure) == (active_prompt,), active_response.text + assert provider_texts(provider) == (active_prompt,), active_response.text + attacked: Final = f"after key update {ATTACK_MARKER} {uuid.uuid4().hex}" + after: Final = candidate.request( + "POST", + "/v1/chat/completions", + _chat_body(model, attacked, (), no_cache=True), + key=key, + ) + assert after.status_code == 400, after.text + assert _azure_texts(azure) == (attacked,), after.text + assert provider.drain() == (), after.text + + +@pytest.mark.parametrize("guardrail_name", ["shield", "moderation"], ids=("E4-shield", "E4-moderation")) +def test_e4_cache_disabled_twin_requests_have_distinct_spend_rows( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_name: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"cache disabled {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + responses: Final = tuple( + candidate.request( + "POST", + "/v1/chat/completions", + _chat_body(model, prompt, ("tuple-writer", guardrail_name), no_cache=True), + ) + for _ in range(3) + ) + assert all(response.status_code == 200 for response in responses), tuple( + response.text for response in responses + ) + assert len(set(response.headers["x-litellm-call-id"] for response in responses)) == 3 + assert _azure_texts(azure) == (prompt, prompt, prompt), responses[-1].text + assert provider_texts(provider) == (prompt, prompt, prompt), responses[-1].text + assert all(_guardrail_entry(response)["guardrail_status"] == "success" for response in responses), ( + responses[-1].text + ) + + +@pytest.mark.parametrize("async_client", [False, True], ids=("sync", "async")) +def test_h15_openai_sdk_tuple_attack_is_blocked( + azure_rig: tuple[Gateway, Wire, Wire], async_client: bool +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"OpenAI SDK tuple {ATTACK_MARKER} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + base_url: Final = str(candidate.client.base_url).rstrip("/") + "/v1" + if async_client: + with pytest.raises(openai.BadRequestError, match="Azure Prompt Shield"): + asyncio.run(_openai_chat_async(base_url, candidate.key, model, prompt, ("tuple-writer", "shield"))) + else: + with pytest.raises(openai.BadRequestError, match="Azure Prompt Shield"): + _openai_chat_sync(base_url, candidate.key, model, prompt, ("tuple-writer", "shield")) + assert _azure_texts(azure) == (prompt,), "SDK request did not reach Azure with the expected prompt" + assert provider.drain() == (), "Blocked SDK request reached the provider" + + +@pytest.mark.parametrize("async_client", [False, True], ids=("sync", "async")) +def test_h15_openai_sdk_tuple_benign_is_scanned_and_served( + azure_rig: tuple[Gateway, Wire, Wire], async_client: bool +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"OpenAI SDK benign {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + base_url: Final = str(candidate.client.base_url).rstrip("/") + "/v1" + if async_client: + asyncio.run(_openai_chat_async(base_url, candidate.key, model, prompt, ("tuple-writer", "shield"))) + else: + _openai_chat_sync(base_url, candidate.key, model, prompt, ("tuple-writer", "shield")) + assert _azure_texts(azure) == (prompt,), "SDK request did not reach Azure with the expected prompt" + assert provider_texts(provider) == (prompt,), "SDK request did not reach the provider with the expected prompt" + + +@pytest.mark.parametrize("async_client", [False, True], ids=("sync", "async")) +def test_h16_anthropic_sdk_tuple_attack_is_blocked( + azure_rig: tuple[Gateway, Wire, Wire], async_client: bool +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"Anthropic SDK tuple {ATTACK_MARKER} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=provider.url, + api_key="synthetic-provider-key", + ) + base_url: Final = str(candidate.client.base_url).rstrip("/") + if async_client: + with pytest.raises(anthropic.BadRequestError, match="Azure Prompt Shield"): + asyncio.run(_anthropic_messages_async(base_url, candidate.key, model, prompt)) + else: + with pytest.raises(anthropic.BadRequestError, match="Azure Prompt Shield"): + _anthropic_messages_sync(base_url, candidate.key, model, prompt) + assert _azure_texts(azure) == (prompt,), "SDK request did not reach Azure with the expected prompt" + assert provider.drain() == (), "Blocked SDK request reached the provider" + + +@pytest.mark.parametrize("async_client", [False, True], ids=("sync", "async")) +def test_h17_openai_sdk_responses_control_stays_blocked( + azure_rig: tuple[Gateway, Wire, Wire], async_client: bool +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"OpenAI Responses {ATTACK_MARKER} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + model: Final = _model(scenario, provider) + base_url: Final = str(candidate.client.base_url).rstrip("/") + "/v1" + if async_client: + with pytest.raises(openai.BadRequestError, match="Azure Prompt Shield"): + asyncio.run(_openai_responses_async(base_url, candidate.key, model, prompt)) + else: + with pytest.raises(openai.BadRequestError, match="Azure Prompt Shield"): + _openai_responses_sync(base_url, candidate.key, model, prompt) + assert _azure_texts(azure) == (prompt,), "SDK request did not reach Azure with the expected prompt" + assert provider.drain() == (), "Blocked SDK request reached the provider" + + +@pytest.mark.parametrize("guardrail_source", ["key", "team"], ids=("H18-key", "H19-team")) +def test_h18_h19_key_and_team_tuple_guardrails_are_applied( + azure_rig: tuple[Gateway, Wire, Wire], guardrail_source: str +) -> None: + candidate, azure, provider = azure_rig + prompt: Final = f"{guardrail_source} metadata {ATTACK_MARKER} {uuid.uuid4().hex}" + with candidate.scenario() as scenario: + team_id: Final = scenario.team(guardrails=["tuple-writer", "shield"]) if guardrail_source == "team" else "" + key_fields: Final = ( + {"team_id": team_id} + if guardrail_source == "team" + else {"guardrails": ["tuple-writer", "shield"]} + ) + key: Final = scenario.key(**key_fields) + model: Final = _model(scenario, provider) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _chat_body(model, prompt, ()), + key=key, + ) + assert response.status_code == 400, response.text + assert _azure_texts(azure) == (prompt,), response.text + assert provider.drain() == (), response.text + + +def _azure_texts(azure: Wire) -> tuple[str, ...]: + return azure_texts(azure) diff --git a/tests/integration/observability/test_azure_content_safety_dispatch_resilience.py b/tests/integration/observability/test_azure_content_safety_dispatch_resilience.py new file mode 100644 index 00000000000..dc7c8beba66 --- /dev/null +++ b/tests/integration/observability/test_azure_content_safety_dispatch_resilience.py @@ -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() diff --git a/tests/integration/sdk/test_azure_prompt_shield_tuple_messages.py b/tests/integration/sdk/test_azure_prompt_shield_tuple_messages.py new file mode 100644 index 00000000000..52885cfdfb8 --- /dev/null +++ b/tests/integration/sdk/test_azure_prompt_shield_tuple_messages.py @@ -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() == () diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/azure/test_azure_prompt_shield.py b/tests/unit/proxy/guardrails/guardrail_hooks/azure/test_azure_prompt_shield.py index 126d42ec3f6..3784ddb4694 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/azure/test_azure_prompt_shield.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/azure/test_azure_prompt_shield.py @@ -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"), [ diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/azure/test_azure_text_moderation.py b/tests/unit/proxy/guardrails/guardrail_hooks/azure/test_azure_text_moderation.py index 5577c6c2a7c..c57c54afc73 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/azure/test_azure_text_moderation.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/azure/test_azure_text_moderation.py @@ -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