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>
This commit is contained in:
yucheng 2026-09-30 21:35:25 +00:00
parent e0284b080a
commit f4206f740d

View file

@ -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),