mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
e0284b080a
commit
f4206f740d
1 changed files with 69 additions and 14 deletions
|
|
@ -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),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue