diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py index 5b1474467b5..ab4c3d54abf 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py @@ -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)) diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py index 91f5e9c7a31..e9516e4633a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py b/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py index 6ca7f42300a..370d8e6307c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py @@ -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) diff --git a/tests/integration/observability/test_azure_content_safety_endpoints.py b/tests/integration/observability/test_azure_content_safety_endpoints.py new file mode 100644 index 00000000000..4d6b32804ba --- /dev/null +++ b/tests/integration/observability/test_azure_content_safety_endpoints.py @@ -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() == () diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/azure/test_azure_prompt_shield.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/azure/test_azure_prompt_shield.py index 46bde80fa3f..9d1eb623068 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/azure/test_azure_prompt_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/azure/test_azure_prompt_shield.py @@ -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) 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 0d815a7fb3c..6c1657f934f 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 @@ -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"]},