diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py index ec183ddf7e7..0b38d76ef03 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py @@ -146,3 +146,17 @@ 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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py b/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py index 370d8e6307c..0dca8be3307 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py @@ -21,6 +21,7 @@ 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, ) @@ -231,7 +232,11 @@ 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, 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) if user_prompt: verbose_proxy_logger.info("Azure Text Moderation: User prompt: %s", user_prompt) diff --git a/tests/integration/observability/test_azure_content_safety_audit.py b/tests/integration/observability/test_azure_content_safety_audit.py index 711268d3dbd..6639115b484 100644 --- a/tests/integration/observability/test_azure_content_safety_audit.py +++ b/tests/integration/observability/test_azure_content_safety_audit.py @@ -17,13 +17,10 @@ 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: @@ -142,23 +139,7 @@ def _azure(outage: threading.Event) -> Callable[[Request], Reply]: } ).encode() ) - 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 Reply(status=404) return respond @@ -201,16 +182,6 @@ 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)) @@ -261,16 +232,6 @@ 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)) @@ -320,14 +281,6 @@ 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,)), @@ -402,60 +355,6 @@ 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: @@ -630,23 +529,6 @@ 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_azure_outage_produces_the_same_outcome_on_responses_and_chat( audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event], ) -> None: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/azure/test_azure_text_moderation.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/azure/test_azure_text_moderation.py index accc72b4b39..4fbc33edcd6 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/azure/test_azure_text_moderation.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/azure/test_azure_text_moderation.py @@ -1,14 +1,13 @@ -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 @@ -20,7 +19,9 @@ 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": [ @@ -48,96 +49,6 @@ 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_guardrail_violation_detected(): """async_make_request is the single enforcement point — it raises @@ -149,14 +60,20 @@ 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": [ @@ -265,7 +182,9 @@ 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": [ @@ -287,7 +206,9 @@ 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": [ @@ -319,7 +240,9 @@ 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": [], @@ -334,7 +257,9 @@ 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( @@ -349,10 +274,9 @@ 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 @@ -363,7 +287,9 @@ 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": [ @@ -400,7 +326,13 @@ 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 @@ -499,7 +431,9 @@ 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"]},