diff --git a/tests/integration/observability/test_azure_content_safety_endpoints.py b/tests/integration/observability/test_azure_content_safety_endpoints.py index 4d6b32804ba..ed2e43fbcb9 100644 --- a/tests/integration/observability/test_azure_content_safety_endpoints.py +++ b/tests/integration/observability/test_azure_content_safety_endpoints.py @@ -1,6 +1,6 @@ import json import uuid -from collections.abc import Iterator +from collections.abc import Callable, Iterator from contextlib import ExitStack from pathlib import Path from typing import Final @@ -11,6 +11,7 @@ from integration._support.client import Gateway, eventually, gateway_from_enviro from integration._support.database import read_rows from integration._support.process import owned_proxy from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue _ATTACK_MARKER: Final = "synthetic-attack-marker" @@ -126,15 +127,15 @@ def _clear_wires(azure_rig: tuple[Gateway, Wire, Wire]) -> None: azure_rig[2].drain() -def _scanned_prompts(azure: Wire) -> list[str]: - return [ +def _scanned_prompts(azure: Wire) -> tuple[JsonValue, ...]: + return tuple( 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: +def _guardrail_entry(model: str) -> dict[str, JsonValue]: rows: Final = eventually( lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), lambda values: len(values) == 1, @@ -147,40 +148,59 @@ def _guardrail_entry(model: str) -> dict: @pytest.mark.parametrize( - ("path", "body_shape", "model_provider"), + ("path", "body", "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" + "/v1/chat/completions", + lambda prompt: {"messages": [{"role": "user", "content": prompt}], "max_tokens": 16}, + "openai", + id="chat-completions-messages", + ), + pytest.param( + "/v1/messages", + lambda prompt: {"messages": [{"role": "user", "content": prompt}], "max_tokens": 16}, + "anthropic", + id="anthropic-messages", + ), + pytest.param( + "/v1/responses", + lambda prompt: {"input": prompt}, + "openai", + id="responses-string-input", + ), + pytest.param( + "/v1/responses", + lambda prompt: {"input": [{"role": "user", "content": [{"type": "input_text", "text": prompt}]}]}, + "openai", + id="responses-list-input", + ), + pytest.param( + "/v1/responses", + lambda prompt: {"messages": [], "input": prompt}, + "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 + request: pytest.FixtureRequest, + azure_rig: tuple[Gateway, Wire, Wire], + path: str, + body: Callable[[str], dict[str, JsonValue]], + model_provider: str, ) -> None: candidate, azure, provider = azure_rig - prompt: Final = f"synthetic prompt {body_shape} {uuid.uuid4().hex}" + prompt: Final = f"synthetic prompt {request.node.callspec.id} {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}) + response: Final = candidate.request("POST", path, {"model": model, **body(prompt)}) assert response.status_code == 200, response.text assert "permitted response" in response.text - assert _scanned_prompts(azure) == [prompt] + assert _scanned_prompts(azure) == (prompt,) assert len(provider.drain()) == 1 entry: Final = _guardrail_entry(model) assert entry["guardrail_status"] == "success", entry @@ -206,5 +226,5 @@ def test_azure_prompt_shield_blocks_attack_in_responses_input( 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 _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 9d1eb623068..e9696f8807e 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 @@ -407,8 +407,6 @@ async def test_responses_input_is_scanned_and_billing_is_logged(responses_input: @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}