mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(guardrails): scan Responses API input in Azure Text Moderation (#43965)
* fix(guardrails): scan Responses API input in Azure Text Moderation Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): log Azure Text Moderation prompts at debug and cover streamed Responses blocking 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
19842da059
commit
0980f756bd
4 changed files with 276 additions and 63 deletions
|
|
@ -150,17 +150,3 @@ class AzureGuardrailBase:
|
|||
if not isinstance(messages, list):
|
||||
return None
|
||||
return get_last_user_message(cast(list[AllMessageValues], messages)) # cast-ok: narrowed to list
|
||||
|
||||
def get_user_prompt(self, messages: list[AllMessageValues]) -> str | None:
|
||||
"""
|
||||
Get the last consecutive block of messages from the user.
|
||||
|
||||
Example:
|
||||
messages = [
|
||||
{"role": "user", "content": "Hello, how are you?"},
|
||||
{"role": "assistant", "content": "I'm good, thank you!"},
|
||||
{"role": "user", "content": "What is the weather in Tokyo?"},
|
||||
]
|
||||
get_user_prompt(messages) -> "What is the weather in Tokyo?"
|
||||
"""
|
||||
return get_last_user_message(messages)
|
||||
|
|
|
|||
|
|
@ -21,7 +21,6 @@ from .base import AzureGuardrailBase
|
|||
if TYPE_CHECKING:
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_text_moderation import (
|
||||
AzureTextModerationGuardrailResponse,
|
||||
)
|
||||
|
|
@ -232,14 +231,10 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr
|
|||
"Azure Text Moderation: Running pre-call prompt scan, on call_type: %s",
|
||||
call_type,
|
||||
)
|
||||
new_messages: Final[list[AllMessageValues] | None] = data.get("messages")
|
||||
if new_messages is None:
|
||||
verbose_proxy_logger.warning("Azure Text Moderation: not running guardrail. No messages in data")
|
||||
return data
|
||||
user_prompt: Final = self.get_user_prompt(new_messages)
|
||||
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)
|
||||
verbose_proxy_logger.debug("Azure Text Moderation: User prompt: %s", user_prompt)
|
||||
await self.async_make_request(
|
||||
text=user_prompt,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -17,10 +17,13 @@ from integration._support.wire import Reply, Request, Wire, wire_server
|
|||
from pydantic import JsonValue
|
||||
|
||||
_ATTACK_MARKER: Final = "synthetic-attack-marker"
|
||||
_MODERATION_MARKER: Final = "synthetic-moderation-marker"
|
||||
|
||||
_SHIELD_TARGET_PREFIX: Final = "/contentsafety/text:shieldPrompt?api-version="
|
||||
_ANALYZE_TARGET_PREFIX: Final = "/contentsafety/text:analyze?api-version="
|
||||
|
||||
_OPT_IN_SHIELD: Final = "audit-shield-optin"
|
||||
_TEXT_MODERATION: Final = "audit-text-mod"
|
||||
|
||||
|
||||
def _chat_frame(identity: str, delta: dict[str, JsonValue], finish: str | None = None) -> bytes:
|
||||
|
|
@ -139,7 +142,23 @@ def _azure(outage: threading.Event) -> Callable[[Request], Reply]:
|
|||
}
|
||||
).encode()
|
||||
)
|
||||
return Reply(status=404)
|
||||
assert request.target.startswith(_ANALYZE_TARGET_PREFIX), request.target
|
||||
text: Final = body["text"]
|
||||
assert isinstance(text, str)
|
||||
severity: Final = 4 if _MODERATION_MARKER in text else 0
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"blocklistsMatch": [],
|
||||
"categoriesAnalysis": [
|
||||
{"category": "Hate", "severity": severity},
|
||||
{"category": "Sexual", "severity": 0},
|
||||
{"category": "SelfHarm", "severity": 0},
|
||||
{"category": "Violence", "severity": 0},
|
||||
],
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
return respond
|
||||
|
||||
|
|
@ -182,6 +201,16 @@ def audit_rig(
|
|||
"guardrail_name": "audit-shield",
|
||||
"litellm_params": _shield_params(azure, mode="pre_call", default_on=True),
|
||||
},
|
||||
{
|
||||
"guardrail_name": _TEXT_MODERATION,
|
||||
"litellm_params": {
|
||||
"guardrail": "azure/text_moderations",
|
||||
"mode": "pre_call",
|
||||
"default_on": False,
|
||||
"api_base": azure.url,
|
||||
"api_key": "synthetic-azure-key",
|
||||
},
|
||||
},
|
||||
],
|
||||
)
|
||||
owned: Final = stack.enter_context(owned_proxy_process(gateway, directory, {}, config=config, workers=2))
|
||||
|
|
@ -232,6 +261,16 @@ def chaos_rig(
|
|||
"guardrail_name": "audit-shield",
|
||||
"litellm_params": _shield_params(azure, mode="pre_call", default_on=True),
|
||||
},
|
||||
{
|
||||
"guardrail_name": _TEXT_MODERATION,
|
||||
"litellm_params": {
|
||||
"guardrail": "azure/text_moderations",
|
||||
"mode": "pre_call",
|
||||
"default_on": False,
|
||||
"api_base": azure.url,
|
||||
"api_key": "synthetic-azure-key",
|
||||
},
|
||||
},
|
||||
],
|
||||
)
|
||||
owned: Final = stack.enter_context(owned_proxy_process(gateway, directory, {}, config=config, workers=2))
|
||||
|
|
@ -281,6 +320,14 @@ def _shield_prompts(requests: tuple[Request, ...]) -> tuple[JsonValue, ...]:
|
|||
)
|
||||
|
||||
|
||||
def _analyze_texts(requests: tuple[Request, ...]) -> tuple[JsonValue, ...]:
|
||||
return tuple(
|
||||
object_value(json.loads(scan.body))["text"]
|
||||
for scan in requests
|
||||
if scan.target.startswith(_ANALYZE_TARGET_PREFIX)
|
||||
)
|
||||
|
||||
|
||||
def _guardrail_entries(model: str, count: int = 1) -> list[JsonValue]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
|
||||
|
|
@ -355,6 +402,60 @@ def test_responses_streaming_input_is_scanned_and_billed(
|
|||
assert entry["guardrail_cost"] == pytest.approx(0.38 / 1000), entry
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
[
|
||||
pytest.param(lambda prompt: {"input": prompt}, id="string-input"),
|
||||
pytest.param(
|
||||
lambda prompt: {"input": [{"role": "user", "content": [{"type": "input_text", "text": prompt}]}]},
|
||||
id="list-input",
|
||||
),
|
||||
pytest.param(lambda prompt: {"messages": [], "input": prompt}, id="empty-messages-stub"),
|
||||
],
|
||||
)
|
||||
def test_text_moderation_opt_in_scans_responses_input(
|
||||
audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event],
|
||||
request: pytest.FixtureRequest,
|
||||
body: Callable[[str], dict[str, JsonValue]],
|
||||
) -> None:
|
||||
owned, azure, provider, _ = audit_rig
|
||||
prompt: Final = f"synthetic benign prompt {request.node.callspec.id} {uuid.uuid4().hex}"
|
||||
with owned.gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-4.1-mini", api_base=provider.url + "/v1", api_key="synthetic-provider-key"
|
||||
)
|
||||
response: Final = owned.gateway.request(
|
||||
"POST", "/v1/responses", {"model": model, "guardrails": [_TEXT_MODERATION], **body(prompt)}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
calls: Final = azure.drain()
|
||||
assert _analyze_texts(calls) == (prompt,)
|
||||
assert _shield_prompts(calls) == (prompt,)
|
||||
assert len(_provider_calls(provider)) == 1
|
||||
entries: Final = _guardrail_entries(model, count=2)
|
||||
assert {object_value(entry)["guardrail_name"] for entry in entries} == {"audit-shield", _TEXT_MODERATION}, (
|
||||
entries
|
||||
)
|
||||
|
||||
|
||||
def test_text_moderation_opt_in_scans_chat_messages(audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event]) -> None:
|
||||
owned, azure, provider, _ = audit_rig
|
||||
prompt: Final = "synthetic benign prompt chat-optin " + uuid.uuid4().hex
|
||||
with owned.gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-4.1-mini", api_base=provider.url + "/v1", api_key="synthetic-provider-key"
|
||||
)
|
||||
response: Final = owned.gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "guardrails": [_TEXT_MODERATION], "messages": [{"role": "user", "content": prompt}]},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
calls: Final = azure.drain()
|
||||
assert _analyze_texts(calls) == (prompt,)
|
||||
assert _shield_prompts(calls) == (prompt,)
|
||||
|
||||
|
||||
def test_chat_with_input_key_still_scans_messages_only(
|
||||
audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event],
|
||||
) -> None:
|
||||
|
|
@ -529,6 +630,45 @@ def test_streaming_responses_attack_is_blocked_before_any_stream_bytes(
|
|||
assert _provider_calls(provider) == ()
|
||||
|
||||
|
||||
def test_text_moderation_opt_in_blocks_responses_input_above_threshold(
|
||||
audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event],
|
||||
) -> None:
|
||||
owned, azure, provider, _ = audit_rig
|
||||
prompt: Final = f"synthetic prompt {_MODERATION_MARKER} " + uuid.uuid4().hex
|
||||
with owned.gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-4.1-mini", api_base=provider.url + "/v1", api_key="synthetic-provider-key"
|
||||
)
|
||||
response: Final = owned.gateway.request(
|
||||
"POST", "/v1/responses", {"model": model, "guardrails": [_TEXT_MODERATION], "input": prompt}
|
||||
)
|
||||
assert response.status_code == 400, response.text
|
||||
assert _analyze_texts(azure.drain()) == (prompt,)
|
||||
assert _provider_calls(provider) == ()
|
||||
|
||||
|
||||
def test_text_moderation_opt_in_blocks_streamed_responses_input_above_threshold(
|
||||
audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event],
|
||||
) -> None:
|
||||
owned, azure, provider, _ = audit_rig
|
||||
prompt: Final = f"synthetic prompt {_MODERATION_MARKER} " + uuid.uuid4().hex
|
||||
with owned.gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-4.1-mini", api_base=provider.url + "/v1", api_key="synthetic-provider-key"
|
||||
)
|
||||
with owned.gateway.client.stream(
|
||||
"POST",
|
||||
"/v1/responses",
|
||||
json={"model": model, "guardrails": [_TEXT_MODERATION], "input": prompt, "stream": True},
|
||||
headers={"Authorization": f"Bearer {owned.gateway.key}"},
|
||||
) as response:
|
||||
body: Final = response.read().decode()
|
||||
assert response.status_code == 400, body
|
||||
assert "Prompt Shield" not in body, body
|
||||
assert _analyze_texts(azure.drain()) == (prompt,)
|
||||
assert _provider_calls(provider) == ()
|
||||
|
||||
|
||||
def test_azure_outage_produces_the_same_outcome_on_responses_and_chat(
|
||||
audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event],
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -1,13 +1,15 @@
|
|||
import logging
|
||||
from typing import Final
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_registry import InMemoryGuardrailHandler
|
||||
from litellm.proxy.guardrails.guardrail_hooks.azure.text_moderation import (
|
||||
AzureContentSafetyTextModerationGuardrail,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_registry import InMemoryGuardrailHandler
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
|
||||
|
|
@ -19,9 +21,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": [
|
||||
|
|
@ -49,6 +49,121 @@ async def test_azure_text_moderation_guardrail_pre_call_hook():
|
|||
assert mock_async_make_request.call_args.kwargs["text"] == "Hello, how are you?"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_text_moderation_scans_responses_input() -> None:
|
||||
guardrail: Final = AzureContentSafetyTextModerationGuardrail(
|
||||
guardrail_name="azure_text_moderation",
|
||||
api_key="azure_text_moderation_api_key",
|
||||
api_base="azure_text_moderation_api_base",
|
||||
)
|
||||
response: Final = Mock()
|
||||
response.json.return_value = {
|
||||
"blocklistsMatch": [],
|
||||
"categoriesAnalysis": [
|
||||
{"category": "Hate", "severity": 2},
|
||||
{"category": "Sexual", "severity": 0},
|
||||
{"category": "SelfHarm", "severity": 0},
|
||||
{"category": "Violence", "severity": 0},
|
||||
],
|
||||
}
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", return_value=response) as mock_post:
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
|
||||
cache=None,
|
||||
data={"input": "Review this response input"},
|
||||
call_type="aresponses",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
mock_post.assert_called_once()
|
||||
assert mock_post.call_args.kwargs["json"]["text"] == "Review this response input"
|
||||
|
||||
|
||||
def _moderation_flagging(flagged: str):
|
||||
def azure_by_text(*args: object, **kwargs: object) -> Mock:
|
||||
body = kwargs["json"]
|
||||
assert isinstance(body, dict)
|
||||
return _moderation_response(6 if body["text"] == flagged else 0)
|
||||
|
||||
return azure_by_text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_text_moderation_empty_messages_stub_does_not_hide_responses_input() -> None:
|
||||
guardrail: Final = AzureContentSafetyTextModerationGuardrail(
|
||||
guardrail_name="azure_text_moderation",
|
||||
api_key="azure_text_moderation_api_key",
|
||||
api_base="azure_text_moderation_api_base",
|
||||
severity_threshold=4,
|
||||
)
|
||||
flagged: Final = "flagged responses input"
|
||||
data: Final[dict[str, object]] = {"messages": [], "input": flagged}
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=_moderation_flagging(flagged)):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="aresponses",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_text_moderation_chat_call_type_scans_messages_not_input() -> None:
|
||||
guardrail: Final = AzureContentSafetyTextModerationGuardrail(
|
||||
guardrail_name="azure_text_moderation",
|
||||
api_key="azure_text_moderation_api_key",
|
||||
api_base="azure_text_moderation_api_base",
|
||||
severity_threshold=4,
|
||||
)
|
||||
flagged: Final = "flagged chat prompt"
|
||||
data: Final[dict[str, object]] = {
|
||||
"messages": [{"role": "user", "content": flagged}],
|
||||
"input": "benign responses input",
|
||||
}
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=_moderation_flagging(flagged)):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="acompletion",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_text_moderation_does_not_log_responses_prompt_above_debug(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
guardrail: Final = AzureContentSafetyTextModerationGuardrail(
|
||||
guardrail_name="azure_text_moderation",
|
||||
api_key="azure_text_moderation_api_key",
|
||||
api_base="azure_text_moderation_api_base",
|
||||
)
|
||||
prompt: Final = "unique benign responses prompt e5f8a2c1"
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
|
||||
with patch.object(guardrail.async_handler, "post", return_value=_moderation_response(0)):
|
||||
await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="azure_text_moderation_api_key"),
|
||||
cache=None,
|
||||
data={"input": prompt},
|
||||
call_type="aresponses",
|
||||
)
|
||||
|
||||
assert not any(record.levelno >= logging.INFO and prompt in record.getMessage() for record in caplog.records), [
|
||||
record.getMessage() for record in caplog.records
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_text_moderation_guardrail_violation_detected():
|
||||
"""async_make_request is the single enforcement point — it raises
|
||||
|
|
@ -60,20 +175,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": [
|
||||
|
|
@ -182,9 +291,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": [
|
||||
|
|
@ -206,9 +313,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": [
|
||||
|
|
@ -240,9 +345,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": [],
|
||||
|
|
@ -257,9 +360,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(
|
||||
|
|
@ -274,9 +375,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
|
||||
|
|
@ -287,9 +389,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": [
|
||||
|
|
@ -326,13 +426,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
|
||||
|
|
@ -431,9 +525,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