mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(guardrails): pick Azure prompt source by call type so a messages stub cannot hide Responses input
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
6788afd267
commit
29fa9f8ac4
6 changed files with 337 additions and 52 deletions
|
|
@ -12,6 +12,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
)
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.types.llms.openai import AllMessageValues, ResponseInputParam
|
||||
from litellm.types.utils import CallTypes, CallTypesLiteral
|
||||
|
||||
# Azure Content Safety APIs have a 10,000 character limit per request.
|
||||
AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH: Final = 10000
|
||||
|
|
@ -23,6 +24,8 @@ AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH: Final = 1000
|
|||
AZURE_CONTENT_SAFETY_DEFAULT_API_VERSION: Final = "2024-09-01"
|
||||
JAVELIN_API_VERSION_STORED_BY_OLDER_RELEASES: Final = "v1"
|
||||
|
||||
_RESPONSES_API_CALL_TYPES: Final = frozenset({CallTypes.responses, CallTypes.aresponses})
|
||||
|
||||
|
||||
def resolve_content_safety_api_version(configured: str | None) -> str:
|
||||
if not configured or configured == JAVELIN_API_VERSION_STORED_BY_OLDER_RELEASES:
|
||||
|
|
@ -131,15 +134,15 @@ class AzureGuardrailBase:
|
|||
|
||||
return chunks
|
||||
|
||||
def get_user_prompt_from_request(self, data: Mapping[str, object]) -> str | None:
|
||||
def get_user_prompt_from_request(self, data: Mapping[str, object], call_type: CallTypesLiteral) -> str | None:
|
||||
if call_type in _RESPONSES_API_CALL_TYPES:
|
||||
responses_input: Final = data.get("input")
|
||||
if not isinstance(responses_input, (str, list)):
|
||||
return None
|
||||
validated_input: Final = cast(ResponseInputParam, responses_input)
|
||||
return get_last_user_message(ResponsesAPIRequestUtils.responses_input_to_chat_messages(validated_input))
|
||||
|
||||
messages: Final = data.get("messages")
|
||||
if isinstance(messages, list):
|
||||
return get_last_user_message(cast(list[AllMessageValues], messages))
|
||||
|
||||
responses_input: Final = data.get("input")
|
||||
if not isinstance(responses_input, (str, list)):
|
||||
if not isinstance(messages, list):
|
||||
return None
|
||||
|
||||
validated_input: Final = cast(ResponseInputParam, responses_input)
|
||||
chat_messages: Final = ResponsesAPIRequestUtils.responses_input_to_chat_messages(validated_input)
|
||||
return get_last_user_message(chat_messages)
|
||||
return get_last_user_message(cast(list[AllMessageValues], messages))
|
||||
|
|
|
|||
|
|
@ -249,7 +249,7 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai
|
|||
"Azure Prompt Shield: Running pre-call prompt scan, on call_type: %s",
|
||||
call_type,
|
||||
)
|
||||
user_prompt: Final = self.get_user_prompt_from_request(data)
|
||||
user_prompt: Final = self.get_user_prompt_from_request(data, call_type)
|
||||
|
||||
if user_prompt:
|
||||
verbose_proxy_logger.debug("Azure Prompt Shield: User prompt: %s", user_prompt)
|
||||
|
|
|
|||
|
|
@ -231,7 +231,7 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr
|
|||
"Azure Text Moderation: Running pre-call prompt scan, on call_type: %s",
|
||||
call_type,
|
||||
)
|
||||
user_prompt: Final = self.get_user_prompt_from_request(data)
|
||||
user_prompt: Final = self.get_user_prompt_from_request(data, call_type)
|
||||
|
||||
if user_prompt:
|
||||
verbose_proxy_logger.info("Azure Text Moderation: User prompt: %s", user_prompt)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,210 @@
|
|||
import json
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from contextlib import ExitStack
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, eventually, gateway_from_environment, object_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
|
||||
_ATTACK_MARKER: Final = "synthetic-attack-marker"
|
||||
|
||||
_AZURE_TARGET_PREFIX: Final = "/contentsafety/text:shieldPrompt?api-version="
|
||||
|
||||
|
||||
def _azure_shield(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
assert request.target.startswith(_AZURE_TARGET_PREFIX), request.target
|
||||
user_prompt: Final = object_value(json.loads(request.body))["userPrompt"]
|
||||
assert isinstance(user_prompt, str)
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"userPromptAnalysis": {"attackDetected": _ATTACK_MARKER in user_prompt},
|
||||
"documentsAnalysis": [],
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
|
||||
def _provider(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
if request.target == "/v1/messages":
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "msg_" + uuid.uuid4().hex,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-4-5-20250929",
|
||||
"content": [{"type": "text", "text": "permitted response"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 11, "output_tokens": 4},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
if request.target == "/v1/responses":
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "resp_" + uuid.uuid4().hex,
|
||||
"object": "response",
|
||||
"created_at": 1700000000,
|
||||
"status": "completed",
|
||||
"model": "gpt-4.1-mini",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_" + uuid.uuid4().hex,
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "permitted response", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
assert request.target == "/v1/chat/completions", request.target
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "chatcmpl-" + uuid.uuid4().hex,
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": "gpt-4.1-mini",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "permitted response"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def azure_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[tuple[Gateway, Wire, Wire]]:
|
||||
directory: Final = tmp_path_factory.mktemp("azure-shield")
|
||||
with ExitStack() as stack:
|
||||
gateway: Final = stack.enter_context(gateway_from_environment())
|
||||
azure: Final = stack.enter_context(wire_server(_azure_shield))
|
||||
provider: Final = stack.enter_context(wire_server(_provider))
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["guardrails"] = [
|
||||
{
|
||||
"guardrail_name": "azure-shield-" + uuid.uuid4().hex,
|
||||
"litellm_params": {
|
||||
"guardrail": "azure/prompt_shield",
|
||||
"mode": "pre_call",
|
||||
"default_on": True,
|
||||
"api_base": azure.url,
|
||||
"api_key": "synthetic-azure-key",
|
||||
"cost_tier": "paid",
|
||||
"price_per_1000_text_records": 0.38,
|
||||
},
|
||||
}
|
||||
]
|
||||
path: Final = directory / "azure-shield.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
candidate: Final = stack.enter_context(owned_proxy(gateway, directory, {}, config=path))
|
||||
yield candidate, azure, provider
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_wires(azure_rig: tuple[Gateway, Wire, Wire]) -> None:
|
||||
azure_rig[1].drain()
|
||||
azure_rig[2].drain()
|
||||
|
||||
|
||||
def _scanned_prompts(azure: Wire) -> list[str]:
|
||||
return [
|
||||
object_value(json.loads(scan.body))["userPrompt"]
|
||||
for scan in azure.drain()
|
||||
if scan.target.startswith(_AZURE_TARGET_PREFIX)
|
||||
]
|
||||
|
||||
|
||||
def _guardrail_entry(model: str) -> dict:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
saved: Final = object_value(rows[0]["metadata"])
|
||||
entries: Final = saved["guardrail_information"]
|
||||
assert isinstance(entries, list) and len(entries) == 1, saved
|
||||
return object_value(entries[0])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("path", "body_shape", "model_provider"),
|
||||
[
|
||||
pytest.param("/v1/chat/completions", "chat", "openai", id="chat-completions-messages"),
|
||||
pytest.param("/v1/messages", "chat", "anthropic", id="anthropic-messages"),
|
||||
pytest.param("/v1/responses", "responses-string", "openai", id="responses-string-input"),
|
||||
pytest.param("/v1/responses", "responses-list", "openai", id="responses-list-input"),
|
||||
pytest.param(
|
||||
"/v1/responses", "responses-string-with-empty-messages", "openai", id="responses-empty-messages-stub"
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_azure_prompt_shield_scans_the_user_prompt_on_every_endpoint(
|
||||
azure_rig: tuple[Gateway, Wire, Wire], path: str, body_shape: str, model_provider: str
|
||||
) -> None:
|
||||
candidate, azure, provider = azure_rig
|
||||
prompt: Final = f"synthetic prompt {body_shape} {uuid.uuid4().hex}"
|
||||
with candidate.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model=("anthropic/claude-sonnet-4-5-20250929" if model_provider == "anthropic" else "openai/gpt-4.1-mini"),
|
||||
api_base=provider.url if model_provider == "anthropic" else provider.url + "/v1",
|
||||
api_key="synthetic-provider-key",
|
||||
)
|
||||
body: Final = {
|
||||
"responses-string": {"input": prompt},
|
||||
"responses-list": {"input": [{"role": "user", "content": [{"type": "input_text", "text": prompt}]}]},
|
||||
"responses-string-with-empty-messages": {"messages": [], "input": prompt},
|
||||
}.get(
|
||||
body_shape,
|
||||
{"messages": [{"role": "user", "content": prompt}], "max_tokens": 16},
|
||||
)
|
||||
response: Final = candidate.request("POST", path, {"model": model, **body})
|
||||
assert response.status_code == 200, response.text
|
||||
assert "permitted response" in response.text
|
||||
assert _scanned_prompts(azure) == [prompt]
|
||||
assert len(provider.drain()) == 1
|
||||
entry: Final = _guardrail_entry(model)
|
||||
assert entry["guardrail_status"] == "success", entry
|
||||
assert entry["guardrail_usage"] == {
|
||||
"requests": 1,
|
||||
"input_characters": len(prompt),
|
||||
"text_records": 1,
|
||||
}, entry
|
||||
assert entry["guardrail_cost"] == pytest.approx(0.38 / 1000), entry
|
||||
|
||||
|
||||
def test_azure_prompt_shield_blocks_attack_in_responses_input(
|
||||
azure_rig: tuple[Gateway, Wire, Wire],
|
||||
) -> None:
|
||||
candidate, azure, provider = azure_rig
|
||||
prompt: Final = f"synthetic prompt {_ATTACK_MARKER} {uuid.uuid4().hex}"
|
||||
with candidate.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-4.1-mini",
|
||||
api_base=provider.url + "/v1",
|
||||
api_key="synthetic-provider-key",
|
||||
)
|
||||
response: Final = candidate.request("POST", "/v1/responses", {"model": model, "input": prompt})
|
||||
assert response.status_code == 400, response.text
|
||||
assert "Violated Azure Prompt Shield guardrail policy" in response.text
|
||||
assert _scanned_prompts(azure) == [prompt]
|
||||
assert provider.drain() == ()
|
||||
|
|
@ -405,6 +405,49 @@ async def test_responses_input_is_scanned_and_billing_is_logged(responses_input:
|
|||
assert entry["guardrail_cost_in_spend"] is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_messages_stub_does_not_hide_responses_input() -> None:
|
||||
"""Cursor sends /v1/responses bodies with an empty messages list plus the real
|
||||
input; a messages-first selector would scan nothing and let the prompt through."""
|
||||
guardrail: Final = _priced_shield_guardrail(cost_tier="paid", price_per_1000_text_records=0.38)
|
||||
prompt: Final = "summarize the thread"
|
||||
data: Final[dict[str, object]] = {"messages": [], "input": prompt}
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", return_value=_shield_response(False)) as mock_post:
|
||||
await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="k"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="aresponses",
|
||||
)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
assert mock_post.call_args.kwargs["json"]["userPrompt"] == prompt
|
||||
entry: Final = _recorded_guardrail_info(data)
|
||||
assert entry["guardrail_usage"] == {"requests": 1, "input_characters": len(prompt), "text_records": 1}
|
||||
assert entry["guardrail_cost"] == pytest.approx(0.00038)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_call_type_scans_messages_not_input() -> None:
|
||||
guardrail: Final = _shield_guardrail()
|
||||
data: Final[dict[str, object]] = {
|
||||
"messages": [{"role": "user", "content": "chat prompt"}],
|
||||
"input": "unrelated responses input",
|
||||
}
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", return_value=_shield_response(False)) as mock_post:
|
||||
await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="k"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="acompletion",
|
||||
)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
assert mock_post.call_args.kwargs["json"]["userPrompt"] == "chat prompt"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_input_attack_detected_raises_http_exception() -> None:
|
||||
guardrail: Final = _priced_shield_guardrail(cost_tier="paid", price_per_1000_text_records=0.38)
|
||||
|
|
|
|||
|
|
@ -20,9 +20,7 @@ async def test_azure_text_moderation_guardrail_pre_call_hook():
|
|||
api_key="azure_text_moderation_api_key",
|
||||
api_base="azure_text_moderation_api_base",
|
||||
)
|
||||
with patch.object(
|
||||
azure_text_moderation_guardrail, "async_make_request"
|
||||
) as mock_async_make_request:
|
||||
with patch.object(azure_text_moderation_guardrail, "async_make_request") as mock_async_make_request:
|
||||
mock_async_make_request.return_value = {
|
||||
"blocklistsMatch": [],
|
||||
"categoriesAnalysis": [
|
||||
|
|
@ -82,6 +80,60 @@ async def test_azure_text_moderation_scans_responses_input() -> None:
|
|||
assert mock_post.call_args.kwargs["json"]["text"] == "Review this response input"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_text_moderation_empty_messages_stub_does_not_hide_responses_input() -> None:
|
||||
guardrail: Final = AzureContentSafetyTextModerationGuardrail(
|
||||
guardrail_name="azure_text_moderation",
|
||||
api_key="azure_text_moderation_api_key",
|
||||
api_base="azure_text_moderation_api_base",
|
||||
severity_threshold=4,
|
||||
)
|
||||
response: Final = Mock()
|
||||
response.json.return_value = {
|
||||
"blocklistsMatch": [],
|
||||
"categoriesAnalysis": [
|
||||
{"category": "Hate", "severity": 0},
|
||||
{"category": "Sexual", "severity": 0},
|
||||
{"category": "SelfHarm", "severity": 0},
|
||||
{"category": "Violence", "severity": 0},
|
||||
],
|
||||
}
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", return_value=response) as mock_post:
|
||||
await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
|
||||
cache=None,
|
||||
data={"messages": [], "input": "Review this response input"},
|
||||
call_type="aresponses",
|
||||
)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
assert mock_post.call_args.kwargs["json"]["text"] == "Review this response input"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_text_moderation_chat_call_type_scans_messages_not_input() -> None:
|
||||
guardrail: Final = AzureContentSafetyTextModerationGuardrail(
|
||||
guardrail_name="azure_text_moderation",
|
||||
api_key="azure_text_moderation_api_key",
|
||||
api_base="azure_text_moderation_api_base",
|
||||
)
|
||||
|
||||
with patch.object(guardrail, "async_make_request") as mock_async_make_request:
|
||||
await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
|
||||
cache=None,
|
||||
data={
|
||||
"messages": [{"role": "user", "content": "chat prompt"}],
|
||||
"input": "unrelated responses input",
|
||||
},
|
||||
call_type="acompletion",
|
||||
)
|
||||
|
||||
mock_async_make_request.assert_called_once()
|
||||
assert mock_async_make_request.call_args.kwargs["text"] == "chat prompt"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_text_moderation_guardrail_violation_detected():
|
||||
"""async_make_request is the single enforcement point — it raises
|
||||
|
|
@ -93,20 +145,14 @@ async def test_azure_text_moderation_guardrail_violation_detected():
|
|||
api_key="azure_text_moderation_api_key",
|
||||
api_base="azure_text_moderation_api_base",
|
||||
)
|
||||
with patch.object(
|
||||
azure_text_moderation_guardrail, "async_make_request"
|
||||
) as mock_async_make_request:
|
||||
with patch.object(azure_text_moderation_guardrail, "async_make_request") as mock_async_make_request:
|
||||
mock_async_make_request.side_effect = HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Azure Content Safety Guardrail: Hate crossed severity 2, Got severity: 2"
|
||||
},
|
||||
detail={"error": "Azure Content Safety Guardrail: Hate crossed severity 2, Got severity: 2"},
|
||||
)
|
||||
with pytest.raises(HTTPException):
|
||||
await azure_text_moderation_guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
api_key="azure_text_moderation_api_key"
|
||||
),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
|
||||
cache=None,
|
||||
data={
|
||||
"messages": [
|
||||
|
|
@ -215,9 +261,7 @@ async def test_azure_text_moderation_violation_in_chunk():
|
|||
):
|
||||
with pytest.raises(HTTPException):
|
||||
await azure_text_moderation_guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
api_key="azure_text_moderation_api_key"
|
||||
),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
|
||||
cache=None,
|
||||
data={
|
||||
"messages": [
|
||||
|
|
@ -239,9 +283,7 @@ async def test_azure_text_moderation_guardrail_post_call_success_hook():
|
|||
api_key="azure_text_moderation_api_key",
|
||||
api_base="azure_text_moderation_api_base",
|
||||
)
|
||||
with patch.object(
|
||||
azure_text_moderation_guardrail, "async_make_request"
|
||||
) as mock_async_make_request:
|
||||
with patch.object(azure_text_moderation_guardrail, "async_make_request") as mock_async_make_request:
|
||||
mock_async_make_request.return_value = {
|
||||
"blocklistsMatch": [],
|
||||
"categoriesAnalysis": [
|
||||
|
|
@ -273,9 +315,7 @@ async def test_azure_text_moderation_guardrail_post_call_checks_all_choices():
|
|||
api_key="azure_text_moderation_api_key",
|
||||
api_base="azure_text_moderation_api_base",
|
||||
)
|
||||
with patch.object(
|
||||
azure_text_moderation_guardrail, "async_make_request"
|
||||
) as mock_async_make_request:
|
||||
with patch.object(azure_text_moderation_guardrail, "async_make_request") as mock_async_make_request:
|
||||
mock_async_make_request.side_effect = [
|
||||
{
|
||||
"blocklistsMatch": [],
|
||||
|
|
@ -290,9 +330,7 @@ async def test_azure_text_moderation_guardrail_post_call_checks_all_choices():
|
|||
with pytest.raises(HTTPException):
|
||||
await azure_text_moderation_guardrail.async_post_call_success_hook(
|
||||
data={},
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
api_key="azure_text_moderation_api_key"
|
||||
),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
|
||||
response=ModelResponse(
|
||||
choices=[
|
||||
Choices(
|
||||
|
|
@ -307,9 +345,10 @@ async def test_azure_text_moderation_guardrail_post_call_checks_all_choices():
|
|||
),
|
||||
)
|
||||
|
||||
assert [
|
||||
call.kwargs["text"] for call in mock_async_make_request.call_args_list
|
||||
] == ["safe response", "unsafe response"]
|
||||
assert [call.kwargs["text"] for call in mock_async_make_request.call_args_list] == [
|
||||
"safe response",
|
||||
"unsafe response",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -320,9 +359,7 @@ async def test_azure_text_moderation_guardrail_post_call_streaming_hook():
|
|||
api_key="azure_text_moderation_api_key",
|
||||
api_base="azure_text_moderation_api_base",
|
||||
)
|
||||
with patch.object(
|
||||
azure_text_moderation_guardrail, "async_make_request"
|
||||
) as mock_async_make_request:
|
||||
with patch.object(azure_text_moderation_guardrail, "async_make_request") as mock_async_make_request:
|
||||
mock_async_make_request.return_value = {
|
||||
"blocklistsMatch": [],
|
||||
"categoriesAnalysis": [
|
||||
|
|
@ -359,13 +396,7 @@ def test_split_text_by_words():
|
|||
assert len(chunks) > 1
|
||||
# Verify no word is broken
|
||||
for chunk in chunks:
|
||||
assert (
|
||||
"word1" in chunk
|
||||
or "word2" in chunk
|
||||
or "word3" in chunk
|
||||
or "word4" in chunk
|
||||
or "word5" in chunk
|
||||
)
|
||||
assert "word1" in chunk or "word2" in chunk or "word3" in chunk or "word4" in chunk or "word5" in chunk
|
||||
|
||||
# Test with very long single word (edge case)
|
||||
long_word = "supercalifragilisticexpialidocious" * 10
|
||||
|
|
@ -464,9 +495,7 @@ async def test_apply_guardrail_scans_every_text():
|
|||
async def test_apply_guardrail_raises_on_detection_in_any_text():
|
||||
guardrail = _moderation_guardrail()
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler, "post", side_effect=[_moderation_response(0), _moderation_response(6)]
|
||||
):
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=[_moderation_response(0), _moderation_response(6)]):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["hello there", "something hateful"]},
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue