mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(guardrails): restore Azure guardrail get_user_prompt dispatch and allow logging (#44067)
* fix(guardrails): restore Azure guardrail get_user_prompt dispatch and allow logging Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): use transport-level doubles in Azure dispatch regression tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): add Azure dispatch audit matrix integration cells Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): restore global callback lists after Router reset in SDK cells Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): make proxy restart cell independent of shutdown timing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng <yucheng@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
5a3a31ea9a
commit
d131c43782
9 changed files with 3115 additions and 11 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
551
tests/integration/observability/azure_dispatch_support.py
Normal file
551
tests/integration/observability/azure_dispatch_support.py
Normal file
|
|
@ -0,0 +1,551 @@
|
|||
import json
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterator
|
||||
from contextlib import ExitStack, contextmanager
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import yaml
|
||||
from integration._support.client import Gateway, object_value
|
||||
from integration._support.process import OwnedProxy, owned_proxy_process
|
||||
from integration._support.redis_process import OwnedRedis, owned_redis
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue
|
||||
|
||||
ATTACK_MARKER: Final = "synthetic-attack-marker"
|
||||
MODERATION_MARKER: Final = "synthetic-moderation-marker"
|
||||
AZURE_ERROR_MARKERS: Final = ("AZURE_500", "AZURE_403", "AZURE_404")
|
||||
PROVIDER_401_MARKER: Final = "PROVIDER_401"
|
||||
OVERSIZED_MARKER: Final = "OVERSIZED_INPUT"
|
||||
SHIELD_TARGET_PREFIX: Final = "/contentsafety/text:shieldPrompt?api-version="
|
||||
ANALYZE_TARGET_PREFIX: Final = "/contentsafety/text:analyze?api-version="
|
||||
|
||||
HOOKS_SOURCE: Final = """from __future__ import annotations
|
||||
|
||||
from typing import Final, cast
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.azure.prompt_shield import (
|
||||
AzureContentSafetyPromptShieldGuardrail,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.azure.text_moderation import (
|
||||
AzureContentSafetyTextModerationGuardrail,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import CallTypesLiteral
|
||||
|
||||
|
||||
class TupleWriter(CustomGuardrail):
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: DualCache,
|
||||
data: dict[str, object],
|
||||
call_type: CallTypesLiteral,
|
||||
) -> dict[str, object]:
|
||||
messages: Final = data.get("messages")
|
||||
if isinstance(messages, list):
|
||||
data["messages"] = tuple(messages) # mutable-ok: the test hook rewrites messages to a tuple
|
||||
return data
|
||||
|
||||
|
||||
class AllTurnsPromptShield(AzureContentSafetyPromptShieldGuardrail):
|
||||
def get_user_prompt(self, messages: list[AllMessageValues]) -> str:
|
||||
return "\\n".join(
|
||||
message["content"]
|
||||
for message in messages
|
||||
if isinstance(message, dict)
|
||||
and message.get("role") == "user"
|
||||
and isinstance(message.get("content"), str)
|
||||
)
|
||||
|
||||
|
||||
class AllTurnsTextModeration(AzureContentSafetyTextModerationGuardrail):
|
||||
def get_user_prompt(self, messages: list[AllMessageValues]) -> str:
|
||||
return "\\n".join(
|
||||
message["content"]
|
||||
for message in messages
|
||||
if isinstance(message, dict)
|
||||
and message.get("role") == "user"
|
||||
and isinstance(message.get("content"), str)
|
||||
)
|
||||
|
||||
|
||||
class RequiringPromptShield(AzureContentSafetyPromptShieldGuardrail):
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: DualCache,
|
||||
data: dict[str, object],
|
||||
call_type: CallTypesLiteral,
|
||||
) -> dict[str, object] | None:
|
||||
messages: Final = data.get("messages")
|
||||
if not isinstance(messages, list):
|
||||
raise HTTPException(status_code=400, detail="no user text")
|
||||
user_prompt: Final = self.get_user_prompt(cast(list[AllMessageValues], messages)) # cast-ok: chat messages
|
||||
if not user_prompt:
|
||||
raise HTTPException(status_code=400, detail="no user text")
|
||||
return await super().async_pre_call_hook(user_api_key_dict, cache, data, call_type)
|
||||
|
||||
|
||||
class RequiringTextModeration(AzureContentSafetyTextModerationGuardrail):
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: DualCache,
|
||||
data: dict[str, object],
|
||||
call_type: CallTypesLiteral,
|
||||
) -> dict[str, object] | None:
|
||||
messages: Final = data.get("messages")
|
||||
if not isinstance(messages, list):
|
||||
raise HTTPException(status_code=400, detail="no user text")
|
||||
user_prompt: Final = self.get_user_prompt(cast(list[AllMessageValues], messages)) # cast-ok: chat messages
|
||||
if not user_prompt:
|
||||
raise HTTPException(status_code=400, detail="no user text")
|
||||
return await super().async_pre_call_hook(user_api_key_dict, cache, data, call_type)
|
||||
"""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AzureBehavior:
|
||||
delay_seconds: float = 0
|
||||
down: threading.Event | None = None
|
||||
entered: threading.Event | None = None
|
||||
arrived: threading.Semaphore | None = None
|
||||
release: threading.Event | None = None
|
||||
barrier_marker: str | None = None
|
||||
|
||||
|
||||
def azure_text(request: Request) -> str:
|
||||
body: Final = object_value(json.loads(request.body))
|
||||
if request.target.startswith(SHIELD_TARGET_PREFIX):
|
||||
prompt: Final = body["userPrompt"]
|
||||
assert isinstance(prompt, str), body
|
||||
return prompt
|
||||
assert request.target.startswith(ANALYZE_TARGET_PREFIX), request.target
|
||||
text: Final = body["text"]
|
||||
assert isinstance(text, str), body
|
||||
return text
|
||||
|
||||
|
||||
def azure_texts(azure: Wire) -> tuple[str, ...]:
|
||||
return tuple(azure_text(request) for request in azure.drain())
|
||||
|
||||
|
||||
def provider_text(request: Request) -> str:
|
||||
body: Final = object_value(json.loads(request.body)) if request.body else {}
|
||||
target: Final = request.target.split("?", 1)[0]
|
||||
match target:
|
||||
case "/v1/chat/completions" | "/v1/messages":
|
||||
messages: Final = body["messages"]
|
||||
assert isinstance(messages, list), body
|
||||
return "\n".join(
|
||||
message_text(message["content"])
|
||||
for message in messages
|
||||
if isinstance(message, dict)
|
||||
and message.get("role") == "user"
|
||||
and "content" in message
|
||||
)
|
||||
case "/v1/responses":
|
||||
value: Final = body["input"]
|
||||
if not isinstance(value, list):
|
||||
return message_text(value)
|
||||
return "\n".join(
|
||||
message_text(message["content"])
|
||||
for message in value
|
||||
if isinstance(message, dict)
|
||||
and message.get("role") == "user"
|
||||
and "content" in message
|
||||
)
|
||||
case "/v1/embeddings":
|
||||
return message_text(body["input"])
|
||||
case "/v1/completions":
|
||||
return message_text(body["prompt"])
|
||||
case _:
|
||||
raise AssertionError(f"Unexpected provider target {request.target}")
|
||||
|
||||
|
||||
def message_text(value: JsonValue) -> str:
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
if isinstance(value, list):
|
||||
return "".join(
|
||||
str(part["text"])
|
||||
for part in value
|
||||
if isinstance(part, dict) and isinstance(part.get("text"), str)
|
||||
)
|
||||
return ""
|
||||
|
||||
|
||||
def provider_texts(provider: Wire) -> tuple[str, ...]:
|
||||
requests: Final = tuple(
|
||||
request
|
||||
for request in provider.drain()
|
||||
if request.method != "GET" or request.target.split("?", 1)[0] != "/v1/models"
|
||||
)
|
||||
return tuple(provider_text(request) for request in requests)
|
||||
|
||||
|
||||
def provider_messages(provider: Wire) -> tuple[JsonValue, ...]:
|
||||
requests: Final = tuple(
|
||||
request
|
||||
for request in provider.drain()
|
||||
if request.method != "GET" or request.target.split("?", 1)[0] != "/v1/models"
|
||||
)
|
||||
targets: Final = tuple(request.target.split("?", 1)[0] for request in requests)
|
||||
assert all(target == "/v1/chat/completions" for target in targets), targets
|
||||
return tuple(object_value(json.loads(request.body))["messages"] for request in requests)
|
||||
|
||||
|
||||
def azure_handler(behavior: AzureBehavior = AzureBehavior()) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST", request.method
|
||||
text: Final = azure_text(request)
|
||||
if behavior.entered is not None and (
|
||||
behavior.barrier_marker is None or behavior.barrier_marker in text
|
||||
):
|
||||
behavior.entered.set()
|
||||
if behavior.arrived is not None:
|
||||
behavior.arrived.release()
|
||||
if behavior.release is not None and (
|
||||
behavior.barrier_marker is None or behavior.barrier_marker in text
|
||||
):
|
||||
assert behavior.release.wait(timeout=30), "Azure barrier was not released"
|
||||
if behavior.down is not None and behavior.down.is_set():
|
||||
return Reply(status=503, body=b'{"error":"synthetic Azure outage"}')
|
||||
if behavior.delay_seconds:
|
||||
time.sleep(behavior.delay_seconds)
|
||||
status: Final = next(
|
||||
(code for marker, code in (("AZURE_500", 500), ("AZURE_403", 403), ("AZURE_404", 404)) if marker in text),
|
||||
200,
|
||||
)
|
||||
if status != 200:
|
||||
return Reply(
|
||||
status=status,
|
||||
body=json.dumps({"error": {"message": f"synthetic Azure error {status}"}}).encode(),
|
||||
)
|
||||
if request.target.startswith(SHIELD_TARGET_PREFIX):
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"userPromptAnalysis": {"attackDetected": ATTACK_MARKER in text},
|
||||
"documentsAnalysis": [],
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
assert request.target.startswith(ANALYZE_TARGET_PREFIX), request.target
|
||||
severity: Final = 4 if MODERATION_MARKER in text else 0
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"blocklistsMatch": [],
|
||||
"categoriesAnalysis": [
|
||||
{"category": "Hate", "severity": severity},
|
||||
{"category": "Sexual", "severity": 0},
|
||||
{"category": "SelfHarm", "severity": 0},
|
||||
{"category": "Violence", "severity": 0},
|
||||
],
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def provider_handler(request: Request) -> Reply:
|
||||
path: Final = request.target.split("?", 1)[0]
|
||||
if request.method == "GET" and path == "/v1/models":
|
||||
return Reply(body=b'{"data":[]}')
|
||||
assert request.method == "POST", request.method
|
||||
text: Final = provider_text(request)
|
||||
if PROVIDER_401_MARKER in text:
|
||||
return Reply(status=401, body=b'{"error":{"message":"synthetic provider unauthorized"}}')
|
||||
if OVERSIZED_MARKER in text:
|
||||
return Reply(
|
||||
status=400,
|
||||
body=b'{"error":{"message":"synthetic context length exceeded","code":"context_length_exceeded"}}',
|
||||
)
|
||||
identity: Final = uuid.uuid4().hex
|
||||
match path:
|
||||
case "/v1/chat/completions":
|
||||
if b'"stream":true' in request.body.replace(b" ", b""):
|
||||
return Reply(content_type="text/event-stream", chunks=_chat_chunks(identity))
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "chatcmpl-" + identity,
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "permitted response"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 2, "completion_tokens": 2, "total_tokens": 4},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
case "/v1/messages":
|
||||
if b'"stream":true' in request.body.replace(b" ", b""):
|
||||
return Reply(content_type="text/event-stream", chunks=_messages_chunks(identity))
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "msg_" + identity,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-4-5-20250929",
|
||||
"content": [{"type": "text", "text": "permitted response"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 2, "output_tokens": 2},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
case "/v1/responses":
|
||||
if b'"stream":true' in request.body.replace(b" ", b""):
|
||||
return Reply(content_type="text/event-stream", chunks=_responses_chunks(identity))
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "resp_" + identity,
|
||||
"object": "response",
|
||||
"created_at": 1700000000,
|
||||
"status": "completed",
|
||||
"model": "gpt-4o-mini",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_" + identity,
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "permitted response", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 2, "output_tokens": 2, "total_tokens": 4},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
case "/v1/embeddings":
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}],
|
||||
"model": "text-embedding-3-small",
|
||||
"usage": {"prompt_tokens": 1, "total_tokens": 1},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
case "/v1/completions":
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "cmpl-" + identity,
|
||||
"object": "text_completion",
|
||||
"created": 1700000000,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [{"text": "permitted response", "index": 0, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 2, "total_tokens": 3},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
case _:
|
||||
return Reply(status=404, body=json.dumps({"error": "unexpected provider target " + path}).encode())
|
||||
|
||||
|
||||
def _chat_chunks(identity: str) -> tuple[bytes, ...]:
|
||||
return (
|
||||
_sse({"id": "chatcmpl-" + identity, "object": "chat.completion.chunk", "choices": [{"index": 0, "delta": {"role": "assistant", "content": "permitted "}, "finish_reason": None}]}),
|
||||
_sse({"id": "chatcmpl-" + identity, "object": "chat.completion.chunk", "choices": [{"index": 0, "delta": {"content": "response"}, "finish_reason": None}]}),
|
||||
_sse({"id": "chatcmpl-" + identity, "object": "chat.completion.chunk", "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}),
|
||||
b"data: [DONE]\n\n",
|
||||
)
|
||||
|
||||
|
||||
def _messages_chunks(identity: str) -> tuple[bytes, ...]:
|
||||
return (
|
||||
_event("message_start", {"type": "message_start", "message": {"id": "msg_" + identity, "type": "message", "role": "assistant", "model": "claude-sonnet-4-5-20250929", "content": [], "stop_reason": None, "stop_sequence": None, "usage": {"input_tokens": 2, "output_tokens": 0}}}),
|
||||
_event("content_block_start", {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}),
|
||||
_event("content_block_delta", {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "permitted response"}}),
|
||||
_event("content_block_stop", {"type": "content_block_stop", "index": 0}),
|
||||
_event("message_delta", {"type": "message_delta", "delta": {"stop_reason": "end_turn", "stop_sequence": None}, "usage": {"output_tokens": 2}}),
|
||||
_event("message_stop", {"type": "message_stop"}),
|
||||
)
|
||||
|
||||
|
||||
def _responses_chunks(identity: str) -> tuple[bytes, ...]:
|
||||
response: Final = {
|
||||
"id": "resp_" + identity,
|
||||
"object": "response",
|
||||
"created_at": 1700000000,
|
||||
"status": "completed",
|
||||
"model": "gpt-4o-mini",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_" + identity,
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "permitted response", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 2, "output_tokens": 2, "total_tokens": 4},
|
||||
}
|
||||
return (
|
||||
_event("response.created", {"type": "response.created", "response": response}),
|
||||
_event("response.output_text.delta", {"type": "response.output_text.delta", "delta": "permitted response"}),
|
||||
_event("response.completed", {"type": "response.completed", "response": response}),
|
||||
)
|
||||
|
||||
|
||||
def _sse(value: dict[str, JsonValue]) -> bytes:
|
||||
return b"data: " + json.dumps(value).encode() + b"\n\n"
|
||||
|
||||
|
||||
def _event(name: str, value: dict[str, JsonValue]) -> bytes:
|
||||
return f"event: {name}\n".encode() + _sse(value)
|
||||
|
||||
|
||||
def guardrail_configs(azure_url: str, *, default_on: bool = False) -> tuple[dict[str, JsonValue], ...]:
|
||||
return (
|
||||
{
|
||||
"guardrail_name": "tuple-writer",
|
||||
"litellm_params": {
|
||||
"guardrail": "azure_dispatch_hooks.TupleWriter",
|
||||
"mode": "pre_call",
|
||||
"default_on": default_on,
|
||||
},
|
||||
},
|
||||
{
|
||||
"guardrail_name": "shield",
|
||||
"litellm_params": {
|
||||
"guardrail": "azure/prompt_shield",
|
||||
"mode": "pre_call",
|
||||
"default_on": default_on,
|
||||
"api_base": azure_url,
|
||||
"api_key": "synthetic-azure-key",
|
||||
},
|
||||
},
|
||||
{
|
||||
"guardrail_name": "moderation",
|
||||
"litellm_params": {
|
||||
"guardrail": "azure/text_moderations",
|
||||
"mode": "pre_call",
|
||||
"default_on": False,
|
||||
"api_base": azure_url,
|
||||
"api_key": "synthetic-azure-key",
|
||||
},
|
||||
},
|
||||
{
|
||||
"guardrail_name": "all-turns-shield",
|
||||
"litellm_params": {
|
||||
"guardrail": "azure_dispatch_hooks.AllTurnsPromptShield",
|
||||
"mode": "pre_call",
|
||||
"default_on": False,
|
||||
"api_base": azure_url,
|
||||
"api_key": "synthetic-azure-key",
|
||||
},
|
||||
},
|
||||
{
|
||||
"guardrail_name": "all-turns-moderation",
|
||||
"litellm_params": {
|
||||
"guardrail": "azure_dispatch_hooks.AllTurnsTextModeration",
|
||||
"mode": "pre_call",
|
||||
"default_on": False,
|
||||
"api_base": azure_url,
|
||||
"api_key": "synthetic-azure-key",
|
||||
},
|
||||
},
|
||||
{
|
||||
"guardrail_name": "requiring-shield",
|
||||
"litellm_params": {
|
||||
"guardrail": "azure_dispatch_hooks.RequiringPromptShield",
|
||||
"mode": "pre_call",
|
||||
"default_on": False,
|
||||
"api_base": azure_url,
|
||||
"api_key": "synthetic-azure-key",
|
||||
},
|
||||
},
|
||||
{
|
||||
"guardrail_name": "requiring-moderation",
|
||||
"litellm_params": {
|
||||
"guardrail": "azure_dispatch_hooks.RequiringTextModeration",
|
||||
"mode": "pre_call",
|
||||
"default_on": False,
|
||||
"api_base": azure_url,
|
||||
"api_key": "synthetic-azure-key",
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def write_dispatch_config(directory: Path, azure_url: str, *, default_on: bool = False) -> Path:
|
||||
(directory / "azure_dispatch_hooks.py").write_text(HOOKS_SOURCE)
|
||||
base_config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
general_settings: Final = {
|
||||
**base_config["general_settings"],
|
||||
"store_prompts_in_spend_logs": True,
|
||||
}
|
||||
config: Final = {
|
||||
**base_config,
|
||||
"guardrails": guardrail_configs(azure_url, default_on=default_on),
|
||||
"general_settings": general_settings,
|
||||
}
|
||||
config_path: Final = directory / "azure-content-safety-dispatch.yaml"
|
||||
config_path.write_text(yaml.safe_dump(config))
|
||||
return config_path
|
||||
|
||||
|
||||
@contextmanager
|
||||
def dispatch_proxy(
|
||||
gateway: Gateway,
|
||||
directory: Path,
|
||||
redis: OwnedRedis,
|
||||
azure_url: str,
|
||||
*,
|
||||
workers: int = 2,
|
||||
default_on: bool = False,
|
||||
) -> Iterator[OwnedProxy]:
|
||||
config: Final = write_dispatch_config(directory, azure_url, default_on=default_on)
|
||||
with owned_proxy_process(
|
||||
gateway,
|
||||
directory,
|
||||
{
|
||||
"REDIS_HOST": redis.host,
|
||||
"REDIS_PORT": str(redis.port),
|
||||
"LITELLM_DISABLE_NO_REDIS_WARNING": "true",
|
||||
},
|
||||
config=config,
|
||||
workers=workers,
|
||||
) as owned:
|
||||
yield owned
|
||||
|
||||
|
||||
@contextmanager
|
||||
def dispatch_rig(
|
||||
gateway: Gateway,
|
||||
directory: Path,
|
||||
*,
|
||||
workers: int = 2,
|
||||
default_on: bool = False,
|
||||
behavior: AzureBehavior = AzureBehavior(),
|
||||
) -> Iterator[tuple[OwnedProxy, Wire, Wire, OwnedRedis]]:
|
||||
with ExitStack() as stack:
|
||||
redis: Final = stack.enter_context(owned_redis(directory))
|
||||
azure: Final = stack.enter_context(wire_server(azure_handler(behavior)))
|
||||
provider: Final = stack.enter_context(wire_server(provider_handler))
|
||||
owned: Final = stack.enter_context(
|
||||
dispatch_proxy(gateway, directory, redis, azure.url, workers=workers, default_on=default_on)
|
||||
)
|
||||
yield owned, azure, provider, redis
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -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()
|
||||
351
tests/integration/sdk/test_azure_prompt_shield_tuple_messages.py
Normal file
351
tests/integration/sdk/test_azure_prompt_shield_tuple_messages.py
Normal file
|
|
@ -0,0 +1,351 @@
|
|||
import asyncio
|
||||
import json
|
||||
from collections.abc import Iterator
|
||||
from contextlib import ExitStack
|
||||
from typing import Final
|
||||
|
||||
import litellm
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from integration._support.client import object_value
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from litellm import Router
|
||||
from litellm.proxy.guardrails.guardrail_hooks.azure.prompt_shield import (
|
||||
AzureContentSafetyPromptShieldGuardrail,
|
||||
)
|
||||
|
||||
_ATTACK_MARKER: Final = "synthetic-sdk-attack-marker"
|
||||
_ATTACK_PROMPT: Final = f"synthetic sdk prompt {_ATTACK_MARKER}"
|
||||
_SHIELD_TARGET_PREFIX: Final = "/contentsafety/text:shieldPrompt?api-version="
|
||||
_PROVIDER_KEY: Final = "synthetic-provider-key"
|
||||
_AZURE_KEY: Final = "synthetic-azure-key"
|
||||
_GUARDRAIL_NAME: Final = "sdk-azure-shield"
|
||||
|
||||
|
||||
def _azure(request: Request) -> Reply:
|
||||
assert request.method == "POST", request.method
|
||||
assert request.target.startswith(_SHIELD_TARGET_PREFIX), request.target
|
||||
body: Final = object_value(json.loads(request.body))
|
||||
prompt: Final = body["userPrompt"]
|
||||
assert isinstance(prompt, str), body
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"userPromptAnalysis": {"attackDetected": _ATTACK_MARKER in prompt},
|
||||
"documentsAnalysis": [],
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
|
||||
def _provider(request: Request) -> Reply:
|
||||
assert request.method == "POST", request.method
|
||||
assert request.target == "/v1/chat/completions", request.target
|
||||
if b'"stream":true' in request.body.replace(b" ", b""):
|
||||
chunk: Final = {
|
||||
"id": "chatcmpl-sdk-azure-guardrail",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1700000000,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [{"index": 0, "delta": {"content": "permitted response"}, "finish_reason": None}],
|
||||
}
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=(b"data: " + json.dumps(chunk).encode() + b"\n\n", b"data: [DONE]\n\n"),
|
||||
)
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "chatcmpl-sdk-azure-guardrail",
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "permitted response"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sdk_rig(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> Iterator[tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire]]:
|
||||
with ExitStack() as stack:
|
||||
azure: Final = stack.enter_context(wire_server(_azure))
|
||||
provider: Final = stack.enter_context(wire_server(_provider))
|
||||
guardrail: Final = AzureContentSafetyPromptShieldGuardrail(
|
||||
guardrail_name=_GUARDRAIL_NAME,
|
||||
api_key=_AZURE_KEY,
|
||||
api_base=azure.url,
|
||||
event_hook="pre_call",
|
||||
default_on=False,
|
||||
)
|
||||
monkeypatch.setattr(litellm, "callbacks", [guardrail])
|
||||
monkeypatch.setattr(litellm, "success_callback", list(litellm.success_callback))
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", list(litellm._async_success_callback))
|
||||
monkeypatch.setattr(litellm, "failure_callback", list(litellm.failure_callback))
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", list(litellm._async_failure_callback))
|
||||
try:
|
||||
yield guardrail, azure, provider
|
||||
finally:
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(guardrail)
|
||||
|
||||
|
||||
def _azure_prompts(azure: Wire) -> tuple[str, ...]:
|
||||
return tuple(_azure_prompt(request) for request in azure.drain())
|
||||
|
||||
|
||||
def _azure_prompt(request: Request) -> str:
|
||||
body: Final = object_value(json.loads(request.body))
|
||||
prompt: Final = body["userPrompt"]
|
||||
assert isinstance(prompt, str), body
|
||||
return prompt
|
||||
|
||||
|
||||
def _provider_prompts(provider: Wire) -> tuple[str, ...]:
|
||||
return tuple(_provider_prompt(request) for request in provider.drain())
|
||||
|
||||
|
||||
def _provider_prompt(request: Request) -> str:
|
||||
body: Final = object_value(json.loads(request.body))
|
||||
messages: Final = body["messages"]
|
||||
assert isinstance(messages, list), body
|
||||
prompt: Final = object_value(messages[-1])["content"]
|
||||
assert isinstance(prompt, str), body
|
||||
return prompt
|
||||
|
||||
|
||||
def test_k1_acompletion_scans_tuple_messages(sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire]) -> None:
|
||||
_, azure, provider = sdk_rig
|
||||
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
|
||||
asyncio.run(
|
||||
litellm.acompletion(
|
||||
model="openai/gpt-4o-mini",
|
||||
api_base=provider.url + "/v1",
|
||||
api_key=_PROVIDER_KEY,
|
||||
messages=({"role": "user", "content": _ATTACK_PROMPT},),
|
||||
guardrails=[_GUARDRAIL_NAME],
|
||||
max_tokens=8,
|
||||
)
|
||||
)
|
||||
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
|
||||
assert provider.drain() == ()
|
||||
|
||||
|
||||
def test_k1_acompletion_scans_system_and_user_tuple_messages(
|
||||
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire],
|
||||
) -> None:
|
||||
_, azure, provider = sdk_rig
|
||||
messages: Final = (
|
||||
{"role": "system", "content": "system context"},
|
||||
{"role": "user", "content": _ATTACK_PROMPT},
|
||||
)
|
||||
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
|
||||
asyncio.run(
|
||||
litellm.acompletion(
|
||||
model="openai/gpt-4o-mini",
|
||||
api_base=provider.url + "/v1",
|
||||
api_key=_PROVIDER_KEY,
|
||||
messages=messages,
|
||||
guardrails=[_GUARDRAIL_NAME],
|
||||
max_tokens=8,
|
||||
)
|
||||
)
|
||||
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
|
||||
assert provider.drain() == ()
|
||||
|
||||
|
||||
def test_k1_acompletion_stream_scans_tuple_messages(
|
||||
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire],
|
||||
) -> None:
|
||||
_, azure, provider = sdk_rig
|
||||
|
||||
async def consume() -> None:
|
||||
stream: Final = await litellm.acompletion(
|
||||
model="openai/gpt-4o-mini",
|
||||
api_base=provider.url + "/v1",
|
||||
api_key=_PROVIDER_KEY,
|
||||
messages=({"role": "user", "content": _ATTACK_PROMPT},),
|
||||
guardrails=[_GUARDRAIL_NAME],
|
||||
max_tokens=8,
|
||||
stream=True,
|
||||
)
|
||||
async for _chunk in stream:
|
||||
pass
|
||||
|
||||
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
|
||||
asyncio.run(consume())
|
||||
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
|
||||
assert provider.drain() == ()
|
||||
|
||||
|
||||
def test_k1_acompletion_list_messages_control(
|
||||
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire]
|
||||
) -> None:
|
||||
_, azure, provider = sdk_rig
|
||||
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
|
||||
asyncio.run(
|
||||
litellm.acompletion(
|
||||
model="openai/gpt-4o-mini",
|
||||
api_base=provider.url + "/v1",
|
||||
api_key=_PROVIDER_KEY,
|
||||
messages=[{"role": "user", "content": _ATTACK_PROMPT}],
|
||||
guardrails=[_GUARDRAIL_NAME],
|
||||
max_tokens=8,
|
||||
)
|
||||
)
|
||||
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
|
||||
assert provider.drain() == ()
|
||||
|
||||
|
||||
def _router(provider: Wire) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "sdk-guardrail-model",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_base": provider.url + "/v1",
|
||||
"api_key": _PROVIDER_KEY,
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def test_k3_router_acompletion_guardrails_kwarg_scans_tuple_messages(
|
||||
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire],
|
||||
) -> None:
|
||||
_, azure, provider = sdk_rig
|
||||
router: Final = _router(provider)
|
||||
try:
|
||||
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
|
||||
asyncio.run(
|
||||
router.acompletion(
|
||||
model="sdk-guardrail-model",
|
||||
messages=({"role": "user", "content": _ATTACK_PROMPT},),
|
||||
guardrails=[_GUARDRAIL_NAME],
|
||||
max_tokens=8,
|
||||
)
|
||||
)
|
||||
finally:
|
||||
router.reset()
|
||||
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
|
||||
assert provider.drain() == ()
|
||||
|
||||
|
||||
def test_k3_router_acompletion_guardrails_kwarg_list_control(
|
||||
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire],
|
||||
) -> None:
|
||||
_, azure, provider = sdk_rig
|
||||
router: Final = _router(provider)
|
||||
try:
|
||||
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
|
||||
asyncio.run(
|
||||
router.acompletion(
|
||||
model="sdk-guardrail-model",
|
||||
messages=[{"role": "user", "content": _ATTACK_PROMPT}],
|
||||
guardrails=[_GUARDRAIL_NAME],
|
||||
max_tokens=8,
|
||||
)
|
||||
)
|
||||
finally:
|
||||
router.reset()
|
||||
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
|
||||
assert provider.drain() == ()
|
||||
|
||||
|
||||
def test_k2_litellm_completion_tuple_behavior_pinned_from_base(
|
||||
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire],
|
||||
) -> None:
|
||||
_, azure, provider = sdk_rig
|
||||
expected_block: Final = False
|
||||
messages: Final = ({"role": "user", "content": _ATTACK_PROMPT},)
|
||||
if expected_block:
|
||||
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
|
||||
litellm.completion(
|
||||
model="openai/gpt-4o-mini",
|
||||
api_base=provider.url + "/v1",
|
||||
api_key=_PROVIDER_KEY,
|
||||
messages=messages,
|
||||
guardrails=[_GUARDRAIL_NAME],
|
||||
max_tokens=8,
|
||||
)
|
||||
else:
|
||||
response: Final = litellm.completion(
|
||||
model="openai/gpt-4o-mini",
|
||||
api_base=provider.url + "/v1",
|
||||
api_key=_PROVIDER_KEY,
|
||||
messages=messages,
|
||||
guardrails=[_GUARDRAIL_NAME],
|
||||
max_tokens=8,
|
||||
)
|
||||
assert response.choices
|
||||
assert _azure_prompts(azure) == ((_ATTACK_PROMPT,) if expected_block else ())
|
||||
if expected_block:
|
||||
assert provider.drain() == ()
|
||||
else:
|
||||
assert _provider_prompts(provider) == (_ATTACK_PROMPT,)
|
||||
|
||||
|
||||
def _router_with_guardrails(provider: Wire) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "sdk-guardrail-model",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_base": provider.url + "/v1",
|
||||
"api_key": _PROVIDER_KEY,
|
||||
"guardrails": [_GUARDRAIL_NAME],
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def test_k4_router_deployment_guardrails_scan_tuple_messages(
|
||||
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire],
|
||||
) -> None:
|
||||
_, azure, provider = sdk_rig
|
||||
router: Final = _router_with_guardrails(provider)
|
||||
try:
|
||||
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
|
||||
asyncio.run(
|
||||
router.acompletion(
|
||||
model="sdk-guardrail-model",
|
||||
messages=({"role": "user", "content": _ATTACK_PROMPT},),
|
||||
max_tokens=8,
|
||||
)
|
||||
)
|
||||
finally:
|
||||
router.reset()
|
||||
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
|
||||
assert provider.drain() == ()
|
||||
|
||||
|
||||
def test_k4_router_deployment_guardrails_scan_list_messages(
|
||||
sdk_rig: tuple[AzureContentSafetyPromptShieldGuardrail, Wire, Wire],
|
||||
) -> None:
|
||||
_, azure, provider = sdk_rig
|
||||
router: Final = _router_with_guardrails(provider)
|
||||
try:
|
||||
with pytest.raises((HTTPException, litellm.BadRequestError), match="Violated Azure Prompt Shield"):
|
||||
asyncio.run(
|
||||
router.acompletion(
|
||||
model="sdk-guardrail-model",
|
||||
messages=[{"role": "user", "content": _ATTACK_PROMPT}],
|
||||
max_tokens=8,
|
||||
)
|
||||
)
|
||||
finally:
|
||||
router.reset()
|
||||
assert _azure_prompts(azure) == (_ATTACK_PROMPT,)
|
||||
assert provider.drain() == ()
|
||||
|
|
@ -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"),
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue