From f4206f740df636480475beb05dbbfd146cdd0b90 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 21:35:25 +0000 Subject: [PATCH] test(guardrails): isolate worker-kill audit rig and cover during_call on chat Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_azure_content_safety_audit.py | 83 +++++++++++++++---- 1 file changed, 69 insertions(+), 14 deletions(-) diff --git a/tests/integration/observability/test_azure_content_safety_audit.py b/tests/integration/observability/test_azure_content_safety_audit.py index 7e2a7da119f..376412b04d2 100644 --- a/tests/integration/observability/test_azure_content_safety_audit.py +++ b/tests/integration/observability/test_azure_content_safety_audit.py @@ -243,6 +243,40 @@ def optin_rig( ) +@pytest.fixture(scope="module") +def chaos_rig( + tmp_path_factory: pytest.TempPathFactory, +) -> Iterator[tuple[OwnedProxy, Wire, Wire, threading.Event]]: + directory: Final = tmp_path_factory.mktemp("azure-chaos") + outage: Final = threading.Event() + with ExitStack() as stack: + gateway: Final = stack.enter_context(gateway_from_environment()) + azure: Final = stack.enter_context(wire_server(_azure(outage))) + provider: Final = stack.enter_context(wire_server(_provider)) + config: Final = _config( + directory, + azure, + [ + { + "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)) + yield owned, azure, provider, outage + + @pytest.fixture(scope="module") def during_rig( tmp_path_factory: pytest.TempPathFactory, @@ -271,7 +305,7 @@ def during_rig( @pytest.fixture(autouse=True) def _clear_wires(request: pytest.FixtureRequest) -> None: - for name in ("audit_rig", "optin_rig", "during_rig"): + for name in ("audit_rig", "optin_rig", "during_rig", "chaos_rig"): if name in request.fixturenames: rig: Final = request.getfixturevalue(name) rig[1].drain() @@ -503,7 +537,7 @@ def test_openai_sdk_responses_calls_are_scanned_and_billed( @pytest.mark.parametrize( - ("bad_input", "expected_status", "provider_calls"), + ("bad_input", "expected_status", "max_provider_calls"), [ pytest.param(123, 500, 0, id="int-input"), pytest.param({"a": 1}, 200, 1, id="dict-input"), @@ -514,17 +548,21 @@ def test_unscannable_responses_input_matches_base_behavior( audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event], bad_input: JsonValue, expected_status: int, - provider_calls: int, + max_provider_calls: int, ) -> None: owned, azure, provider, _ = audit_rig 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, "input": bad_input}) + response: Final = owned.gateway.request( + "POST", + "/v1/responses", + {"model": model, "input": bad_input, "metadata": {"cell": uuid.uuid4().hex}}, + ) assert response.status_code == expected_status, response.text assert _shield_prompts(azure.drain()) == () - assert len(_provider_calls(provider)) == provider_calls + assert len(_provider_calls(provider)) <= max_provider_calls def test_long_responses_input_is_chunked_and_billed(audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event]) -> None: @@ -730,17 +768,31 @@ def test_opt_in_shield_scans_responses_input_exactly_once( assert entry["guardrail_name"] == _OPT_IN_SHIELD, entry -def test_during_call_shield_behavior_on_responses(during_rig: tuple[Gateway, Wire, Wire]) -> None: +def test_during_call_shield_does_not_scan_any_endpoint(during_rig: tuple[Gateway, Wire, Wire]) -> None: gateway, azure, provider = during_rig - prompt: Final = "synthetic prompt during " + uuid.uuid4().hex + chat_prompt: Final = "synthetic prompt during-chat " + uuid.uuid4().hex + responses_prompt: Final = "synthetic prompt during-responses " + uuid.uuid4().hex with gateway.scenario() as scenario: - model: Final = scenario.model( + chat_model: Final = scenario.model( model="openai/gpt-4.1-mini", api_base=provider.url + "/v1", api_key="synthetic-provider-key" ) - response: Final = gateway.request("POST", "/v1/responses", {"model": model, "input": prompt}) - assert response.status_code == 200, response.text + responses_model: Final = scenario.model( + model="openai/gpt-4.1-mini", api_base=provider.url + "/v1", api_key="synthetic-provider-key" + ) + chat_response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": chat_model, "messages": [{"role": "user", "content": chat_prompt}]}, + ) + responses_response: Final = gateway.request( + "POST", "/v1/responses", {"model": responses_model, "input": responses_prompt} + ) + assert chat_response.status_code == responses_response.status_code == 200, ( + chat_response.text, + responses_response.text, + ) assert _shield_prompts(azure.drain()) == () - assert len(_provider_calls(provider)) == 1 + assert len(_provider_calls(provider)) == 2 def test_concurrent_mixed_requests_scan_each_prompt_once( @@ -877,9 +929,9 @@ def test_azure_outage_burst_then_recovery_bills_fresh_requests_once( def test_killing_a_worker_mid_burst_leaves_no_duplicate_rows( - audit_rig: tuple[OwnedProxy, Wire, Wire, threading.Event], + chaos_rig: tuple[OwnedProxy, Wire, Wire, threading.Event], ) -> None: - owned, azure, provider, _ = audit_rig + owned, azure, provider, _ = chaos_rig port: Final = owned.gateway.client.base_url.port workers: Final = tuple( child @@ -906,7 +958,10 @@ def test_killing_a_worker_mid_burst_leaves_no_duplicate_rows( ) survivors: Final = tuple(status for _, status in outcomes if status != -1) assert survivors and {status for status in survivors} == {200}, outcomes - assert set(_shield_prompts(azure.drain())) == {identity for identity, status in outcomes if status == 200} + scans: Final = _shield_prompts(azure.drain()) + assert len(scans) == len(set(scans)), scans + assert set(scans) <= set(identities), scans + assert {identity for identity, status in outcomes if status == 200} <= set(scans), (outcomes, scans) rows: Final = eventually( lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), lambda values: len(values) == len(survivors),