test(grayswan): assert the full generic guardrail body and kill a real serving worker

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-30 00:46:48 +00:00
parent ce933b35b8
commit 1cfac8e8c1
2 changed files with 66 additions and 18 deletions

View file

@ -143,6 +143,20 @@ def _chat_provider(message: dict[str, JsonValue]):
return _serving_model_probe(respond)
def _normalized_generic_body(body: dict[str, JsonValue]) -> dict[str, JsonValue]:
headers: Final = body.get("request_headers")
normalized_headers: Final = (
{**headers, "host": "<host>", "content-length": "<length>"} if isinstance(headers, dict) else headers
)
return {
**body,
"litellm_call_id": "<call-id>",
"litellm_trace_id": "<trace-id>",
"litellm_version": "<version>",
"request_headers": normalized_headers,
}
def _monitor_bodies(vendor: Wire, expected: int = 1, seconds: float = 30) -> tuple[dict[str, JsonValue], ...]:
collected: tuple[dict[str, JsonValue], ...] = ()
@ -1092,25 +1106,21 @@ def test_post_call_generic_guardrail_inputs_unchanged(gateway: Gateway, tmp_path
assert request.target == "/beta/litellm_basic_guardrail_api", request.target
return Reply(body=json.dumps({"action": "NONE"}).encode())
generic_entry: Final = {
"guardrail_name": generic_name,
"litellm_params": {
"guardrail": "generic_guardrail_api",
"mode": "post_call",
"default_on": True,
"api_base": None,
"api_key": "synthetic-guardrail-key",
},
}
with (
wire_server(_vendor()) as vendor,
wire_server(generic_policy) as policy,
wire_server(_chat_provider({"role": "assistant", "content": response_text})) as upstream,
):
generic_entry["litellm_params"]["api_base"] = (
policy.url
) # writable-ok: wire port only exists inside the context
generic_entry: Final = {
"guardrail_name": generic_name,
"litellm_params": {
"guardrail": "generic_guardrail_api",
"mode": "post_call",
"default_on": True,
"api_base": policy.url,
"api_key": "synthetic-guardrail-key",
},
}
config_path: Final = _grayswan_config(
tmp_path, identity, vendor.url, "post_call", extra_guardrails=(generic_entry,)
)
@ -1138,7 +1148,32 @@ def test_post_call_generic_guardrail_inputs_unchanged(gateway: Gateway, tmp_path
seconds=30,
)
generic_body: Final = generic_bodies[0]
assert generic_body["texts"] == [response_text], generic_body
assert _normalized_generic_body(generic_body) == {
"additional_provider_specific_params": {},
"images": None,
"input_type": "response",
"litellm_call_id": "<call-id>",
"litellm_trace_id": "<trace-id>",
"litellm_version": "<version>",
"model": "gpt-4o-mini",
"request_data": {
"user_api_key_hash": "litellm_proxy_master_key",
"user_api_key_user_id": "default_user_id",
},
"request_headers": {
"accept": "*/*",
"accept-encoding": "gzip, deflate, br",
"connection": "keep-alive",
"content-length": "<length>",
"content-type": "application/json",
"host": "<host>",
"user-agent": "python-httpx/0.28.1",
},
"structured_messages": None,
"texts": [response_text],
"tool_calls": None,
"tools": None,
}, generic_body
assert grayswan_body["messages"][:-1] == [dict(message) for message in _REQUEST_MESSAGES], grayswan_body

View file

@ -8,6 +8,7 @@ from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from typing import Final
import psutil
import yaml
from integration._support.client import Gateway
from integration._support.process import group_members, owned_proxy_process
@ -206,9 +207,21 @@ def test_worker_kill_mid_burst_survivor_keeps_serving(gateway: Gateway, tmp_path
warm: Final = _fire(candidate, model, "marker-warm", False)
assert warm == 200
members: Final = group_members(owned.process.pid)
children: Final = tuple(member for member in members if member.pid != owned.process.pid)
assert len(children) >= 2, [member.pid for member in members]
os.kill(children[0].pid, signal.SIGKILL)
candidate_port: Final = candidate.client.base_url.port
workers_listening: Final = tuple(
member
for member in members
if member.pid != owned.process.pid
and any(
connection.laddr.port == candidate_port and connection.status == "LISTEN"
for connection in member.net_connections(kind="inet")
)
)
assert len(workers_listening) == 2, [member.pid for member in members]
victim: Final = workers_listening[0]
os.kill(victim.pid, signal.SIGKILL)
psutil.wait_procs((victim,), timeout=10)
assert not psutil.pid_exists(victim.pid), victim.pid
statuses: Final = tuple(_fire(candidate, model, f"marker-kill-{index}", False) for index in range(6))
assert all(status == 200 for status in statuses), statuses
bodies: Final = _monitor_bodies(vendor_wire, expected=7)