test(callbacks): unwind the callbacks global the policy engine and realtime tests scaffold around (#37826)
Some checks are pending
Unit Tests: Proxy DB Operations / proxy-runtime (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-server-core (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-utils (push) Blocked by required conditions
Unit Tests / core-utils (push) Waiting to run
Unit Tests / enterprise-routing (push) Waiting to run
Unit Tests / integrations (push) Waiting to run
Unit Tests / All Other Providers (push) Waiting to run
Unit Tests / Vertex AI (push) Waiting to run
Unit Tests / misc (push) Waiting to run
Unit Tests / proxy-auth (push) Waiting to run
Unit Tests / proxy-endpoints (push) Waiting to run
Unit Tests / proxy-infra (push) Waiting to run
Unit Tests / proxy-server (push) Waiting to run
Unit Tests / responses-caching-types (push) Waiting to run
GitHub Actions Security Analysis / zizmor (push) Waiting to run
CI Coverage / assert-ci-coverage (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Publish basedpyright base counts / publish (push) Waiting to run
Code Quality Checks / code-quality (push) Waiting to run
UI Unit Tests / ui-unit-tests (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Waiting to run
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Waiting to run
Unit Tests: Documentation Validation / documentation (push) Waiting to run
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Waiting to run
Unit Tests: Proxy DB Operations / auth-checks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / budgets (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / custom-logging (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / db-and-spend (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / key-generation (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / logging-misc (push) Blocked by required conditions

* test(policy-engine): unwind the callback global the pipeline tests scaffold around

Every one of the 16 tests in this file set litellm.callbacks by hand, each
wrapping its body in a try/finally to put the old value back, and each capturing
that old value with a .copy() first. That is 32 TQ005 violations and about 70
lines of scaffolding to say what monkeypatch.setattr says in one.

The write also sat outside the try, so the block that restores it did not cover
the statement that changed it.

16 tests pass either way, and litellm.callbacks reads restored on both sides,
because the conftest snapshot already lists it. The point is that these tests
stop depending on that snapshot to clean up after them.

* test(realtime): unwind the same callback global in the realtime streaming tests

Same global, same shape as the previous commit. 25 writes to litellm.callbacks,
2 of them wrapped in a try/finally that resets to [] rather than to the old
value, and 12 tests that write it with no protection at all.

monkeypatch.setattr replaces all of them, and the sys.path.insert with its
now-unused os and sys imports goes too.

Both sides read restored here as well, for the same reason as the previous
commit: litellm.callbacks is in the conftest snapshot. What changes is that
these tests no longer lean on it.

101 tests pass in this file, 16 in the policy engine one.

* style(realtime): wrap the one signature the monkeypatch param pushed past 120
This commit is contained in:
yuneng-jiang 2026-08-21 21:19:35 -07:00 • committed by GitHub
parent 73307070c2
commit 6bce3dce0d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 310 additions and 398 deletions

View file

@ -6,13 +6,13 @@
"limit": 742
},
"TQ003": {
"limit": 1074
"limit": 1073
},
"TQ004": {
"limit": 469
},
"TQ005": {
"limit": 2621
"limit": 2562
},
"TQ006": {
"limit": 34

View file

@ -1,6 +1,4 @@
import json
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -8,7 +6,6 @@ from websockets.exceptions import ConnectionClosed
import litellm
sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.realtime_streaming import (
@ -1326,7 +1323,7 @@ async def test_log_messages_includes_tools_in_model_call_details():
@pytest.mark.asyncio
async def test_realtime_guardrail_blocks_prompt_injection():
async def test_realtime_guardrail_blocks_prompt_injection(monkeypatch: pytest.MonkeyPatch):
"""
Test that when a transcription event containing prompt injection arrives from the
backend, a registered guardrail blocks it — sending a warning to the client
@ -1350,7 +1347,7 @@ async def test_realtime_guardrail_blocks_prompt_injection():
event_hook=GuardrailEventHooks.realtime_input_transcription,
default_on=True,
)
litellm.callbacks = [guardrail]
monkeypatch.setattr(litellm, "callbacks", [guardrail])
# --- client websocket mock ---
client_ws = MagicMock()
@ -1405,11 +1402,10 @@ async def test_realtime_guardrail_blocks_prompt_injection():
f"Expected guardrail_violation error type, got: {error_events[0]}"
)
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_realtime_guardrail_allows_clean_transcript():
async def test_realtime_guardrail_allows_clean_transcript(monkeypatch: pytest.MonkeyPatch):
"""
Test that a clean transcript passes through the guardrail and triggers
response.create to the backend.
@ -1430,7 +1426,7 @@ async def test_realtime_guardrail_allows_clean_transcript():
event_hook=GuardrailEventHooks.realtime_input_transcription,
default_on=True,
)
litellm.callbacks = [guardrail]
monkeypatch.setattr(litellm, "callbacks", [guardrail])
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
@ -1463,11 +1459,10 @@ async def test_realtime_guardrail_allows_clean_transcript():
response_creates = [e for e in sent_to_backend if e.get("type") == "response.create"]
assert len(response_creates) == 1, f"Clean transcript should trigger response.create, got: {sent_to_backend}"
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_realtime_text_input_guardrail_blocks_and_returns_error():
async def test_realtime_text_input_guardrail_blocks_and_returns_error(monkeypatch: pytest.MonkeyPatch):
"""
Test that when conversation.item.create arrives with text that triggers a guardrail,
the proxy blocks it (doesn't forward to backend) and returns an error event directly
@ -1495,7 +1490,7 @@ async def test_realtime_text_input_guardrail_blocks_and_returns_error():
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
)
litellm.callbacks = [guardrail]
monkeypatch.setattr(litellm, "callbacks", [guardrail])
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
@ -1558,11 +1553,10 @@ async def test_realtime_text_input_guardrail_blocks_and_returns_error():
]
assert len(original_items) == 0, f"Blocked item should not be forwarded to backend, got: {original_items}"
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_realtime_function_call_output_guardrail_blocks_and_returns_error():
async def test_realtime_function_call_output_guardrail_blocks_and_returns_error(monkeypatch: pytest.MonkeyPatch):
"""
Test that a client-supplied function_call_output whose content triggers a
guardrail is blocked: it is not forwarded to the backend, and an error
@ -1590,7 +1584,7 @@ async def test_realtime_function_call_output_guardrail_blocks_and_returns_error(
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
)
litellm.callbacks = [guardrail]
monkeypatch.setattr(litellm, "callbacks", [guardrail])
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
@ -1648,11 +1642,10 @@ async def test_realtime_function_call_output_guardrail_blocks_and_returns_error(
assert sanitized_item["call_id"] == "call_123"
assert "test@example.com" not in sanitized_item["output"]
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_realtime_function_call_output_guardrail_allows_clean_output():
async def test_realtime_function_call_output_guardrail_allows_clean_output(monkeypatch: pytest.MonkeyPatch):
"""
Test that a clean function_call_output passes through and reaches the backend
when guardrails are configured.
@ -1670,7 +1663,7 @@ async def test_realtime_function_call_output_guardrail_allows_clean_output():
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
)
litellm.callbacks = [guardrail]
monkeypatch.setattr(litellm, "callbacks", [guardrail])
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
@ -1714,11 +1707,10 @@ async def test_realtime_function_call_output_guardrail_allows_clean_output():
]
assert len(forwarded) == 1, f"Clean function_call_output should be forwarded, got: {forwarded}"
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_realtime_text_input_guardrail_uses_pre_call_mode():
async def test_realtime_text_input_guardrail_uses_pre_call_mode(monkeypatch: pytest.MonkeyPatch):
"""
Test that _has_realtime_guardrails returns True for a guardrail configured with
pre_call mode (not just realtime_input_transcription).
@ -1736,7 +1728,7 @@ async def test_realtime_text_input_guardrail_uses_pre_call_mode():
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
)
litellm.callbacks = [guardrail]
monkeypatch.setattr(litellm, "callbacks", [guardrail])
client_ws = MagicMock()
backend_ws = MagicMock()
@ -1751,11 +1743,10 @@ async def test_realtime_text_input_guardrail_uses_pre_call_mode():
"pre_call-only guardrail must not disable server_vad auto-response"
)
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_realtime_session_created_injects_session_update_for_audio_guardrail():
async def test_realtime_session_created_injects_session_update_for_audio_guardrail(monkeypatch: pytest.MonkeyPatch):
"""
Test that when an audio transcription guardrail is configured, a session.created
event from the backend triggers a session.update injection (create_response: false)
@ -1775,7 +1766,7 @@ async def test_realtime_session_created_injects_session_update_for_audio_guardra
event_hook=GuardrailEventHooks.realtime_input_transcription,
default_on=True,
)
litellm.callbacks = [guardrail]
monkeypatch.setattr(litellm, "callbacks", [guardrail])
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
@ -1809,11 +1800,12 @@ async def test_realtime_session_created_injects_session_update_for_audio_guardra
"GA session.update must nest turn_detection under audio.input"
)
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_realtime_session_created_does_not_inject_session_update_for_pre_call_only():
async def test_realtime_session_created_does_not_inject_session_update_for_pre_call_only(
monkeypatch: pytest.MonkeyPatch,
):
"""
pre_call-only guardrails must not inject create_response:false on realtime
sessions — that breaks server_vad for audio-only voice agents (e.g. Model Armor).
@ -1831,7 +1823,7 @@ async def test_realtime_session_created_does_not_inject_session_update_for_pre_c
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
)
litellm.callbacks = [guardrail]
monkeypatch.setattr(litellm, "callbacks", [guardrail])
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
@ -1853,11 +1845,10 @@ async def test_realtime_session_created_does_not_inject_session_update_for_pre_c
session_updates = [e for e in sent_to_backend if e.get("type") == "session.update"]
assert len(session_updates) == 0, f"pre_call-only guardrail must not inject session.update, got: {sent_to_backend}"
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_pre_call_and_post_call_guardrails_do_not_disable_server_vad():
async def test_pre_call_and_post_call_guardrails_do_not_disable_server_vad(monkeypatch: pytest.MonkeyPatch):
"""Model Armor-style pre_call + post_call must not gate audio VAD."""
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
@ -1867,18 +1858,22 @@ async def test_pre_call_and_post_call_guardrails_do_not_disable_server_vad():
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
return inputs
litellm.callbacks = [
ModelArmorStyleGuardrail(
guardrail_name="model_armor_all_pre_call",
event_hook=GuardrailEventHooks.pre_call,
default_on=False,
),
ModelArmorStyleGuardrail(
guardrail_name="model_armor_all_post_call",
event_hook=GuardrailEventHooks.post_call,
default_on=False,
),
]
monkeypatch.setattr(
litellm,
"callbacks",
[
ModelArmorStyleGuardrail(
guardrail_name="model_armor_all_pre_call",
event_hook=GuardrailEventHooks.pre_call,
default_on=False,
),
ModelArmorStyleGuardrail(
guardrail_name="model_armor_all_post_call",
event_hook=GuardrailEventHooks.post_call,
default_on=False,
),
],
)
client_ws = MagicMock()
backend_ws = MagicMock()
@ -1900,11 +1895,10 @@ async def test_pre_call_and_post_call_guardrails_do_not_disable_server_vad():
assert streaming._has_realtime_guardrails() is True
assert streaming._has_audio_transcription_guardrails() is False
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_end_session_after_n_fails_closes_connection():
async def test_end_session_after_n_fails_closes_connection(monkeypatch: pytest.MonkeyPatch):
"""
Test that end_session_after_n_fails=2 closes the backend websocket after
the second guardrail violation in a session.
@ -1923,7 +1917,7 @@ async def test_end_session_after_n_fails_closes_connection():
default_on=True,
end_session_after_n_fails=2,
)
litellm.callbacks = [guardrail]
monkeypatch.setattr(litellm, "callbacks", [guardrail])
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
@ -1948,11 +1942,10 @@ async def test_end_session_after_n_fails_closes_connection():
assert backend_ws.close.called, "Expected backend_ws.close() to be called after 2 violations"
assert streaming._violation_count == 2
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_on_violation_end_session_closes_on_first_fail():
async def test_on_violation_end_session_closes_on_first_fail(monkeypatch: pytest.MonkeyPatch):
"""
Test that on_violation='end_session' closes the session immediately on the
first violation, regardless of end_session_after_n_fails.
@ -1971,7 +1964,7 @@ async def test_on_violation_end_session_closes_on_first_fail():
default_on=True,
on_violation="end_session",
)
litellm.callbacks = [guardrail]
monkeypatch.setattr(litellm, "callbacks", [guardrail])
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
@ -1995,7 +1988,6 @@ async def test_on_violation_end_session_closes_on_first_fail():
assert backend_ws.close.called, "Expected session to close immediately with on_violation=end_session"
assert streaming._violation_count == 1
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
@ -2898,53 +2890,47 @@ def _transcription_guardrail():
)
def test_setup_folds_in_auto_response_disable_when_transcription_guardrail_active():
def test_setup_folds_in_auto_response_disable_when_transcription_guardrail_active(monkeypatch: pytest.MonkeyPatch):
"""Gemini rejects a second setup, so a transcription guardrail's auto-response
disable must be folded into the one-and-only setup; otherwise the model
auto-responds and the guardrail is bypassed."""
import litellm
litellm.callbacks = [_transcription_guardrail()]
try:
streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
setup = json.dumps(
{
"setup": {
"model": "models/gemini-3.1-flash-live-preview",
"generationConfig": {"responseModalities": ["AUDIO"]},
"inputAudioTranscription": {},
}
monkeypatch.setattr(litellm, "callbacks", [_transcription_guardrail()])
streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
setup = json.dumps(
{
"setup": {
"model": "models/gemini-3.1-flash-live-preview",
"generationConfig": {"responseModalities": ["AUDIO"]},
"inputAudioTranscription": {},
}
)
out = json.loads(streaming._maybe_inject_guardrail_auto_response_disable(setup))
aad = out["setup"]["realtimeInputConfig"]["automaticActivityDetection"]
assert aad["disabled"] is True
finally:
litellm.callbacks = []
}
)
out = json.loads(streaming._maybe_inject_guardrail_auto_response_disable(setup))
aad = out["setup"]["realtimeInputConfig"]["automaticActivityDetection"]
assert aad["disabled"] is True
def test_setup_unchanged_without_transcription_guardrail():
def test_setup_unchanged_without_transcription_guardrail(monkeypatch: pytest.MonkeyPatch):
import litellm
litellm.callbacks = []
monkeypatch.setattr(litellm, "callbacks", [])
streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
setup = json.dumps({"setup": {"model": "x", "generationConfig": {"responseModalities": ["AUDIO"]}}})
out = streaming._maybe_inject_guardrail_auto_response_disable(setup)
assert json.loads(out) == json.loads(setup)
def test_non_bidi_setup_left_untouched_for_followup_capable_providers():
def test_non_bidi_setup_left_untouched_for_followup_capable_providers(monkeypatch: pytest.MonkeyPatch):
"""OpenAI realtime accepts a follow-up session.update, so a non-bidi message
(no top-level 'setup' key) must be left untouched even with a guardrail on."""
import litellm
litellm.callbacks = [_transcription_guardrail()]
try:
streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
msg = json.dumps({"type": "session.update", "session": {"instructions": "hi"}})
assert streaming._maybe_inject_guardrail_auto_response_disable(msg) == msg
finally:
litellm.callbacks = []
monkeypatch.setattr(litellm, "callbacks", [_transcription_guardrail()])
streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
msg = json.dumps({"type": "session.update", "session": {"instructions": "hi"}})
assert streaming._maybe_inject_guardrail_auto_response_disable(msg) == msg
@pytest.mark.asyncio

View file

@ -165,7 +165,7 @@ class ContentCheckGuardrail(CustomGuardrail):
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
@pytest.mark.asyncio
async def test_escalation_step1_fails_step2_blocks():
async def test_escalation_step1_fails_step2_blocks(monkeypatch):
"""
Pipeline: simple-filter (on_fail: next) -> advanced-filter (on_fail: block)
Input: request that fails simple-filter
@ -182,36 +182,32 @@ async def test_escalation_step1_fails_step2_blocks():
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [simple_guard, advanced_guard]
monkeypatch.setattr(litellm, "callbacks", [simple_guard, advanced_guard])
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "bad content"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="content-safety",
)
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "bad content"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="content-safety",
)
assert simple_guard.calls == 1
assert advanced_guard.calls == 1
assert result.terminal_action == "block"
assert len(result.step_results) == 2
assert result.step_results[0].guardrail_name == "simple-filter"
assert result.step_results[0].outcome == "fail"
assert result.step_results[0].action_taken == "next"
assert result.step_results[1].guardrail_name == "advanced-filter"
assert result.step_results[1].outcome == "fail"
assert result.step_results[1].action_taken == "block"
finally:
litellm.callbacks = original_callbacks
assert simple_guard.calls == 1
assert advanced_guard.calls == 1
assert result.terminal_action == "block"
assert len(result.step_results) == 2
assert result.step_results[0].guardrail_name == "simple-filter"
assert result.step_results[0].outcome == "fail"
assert result.step_results[0].action_taken == "next"
assert result.step_results[1].guardrail_name == "advanced-filter"
assert result.step_results[1].outcome == "fail"
assert result.step_results[1].action_taken == "block"
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
@pytest.mark.asyncio
async def test_block_carries_original_guardrail_exception():
async def test_block_carries_original_guardrail_exception(monkeypatch):
"""A blocking step must expose the guardrail's own raised exception on the
result so the caller can re-raise it verbatim, giving the policy path the
same response/trace as a direct guardrail attachment."""
@ -219,67 +215,52 @@ async def test_block_carries_original_guardrail_exception():
pipeline = GuardrailPipeline(
mode="pre_call",
steps=[
PipelineStep(
guardrail="moderation-filter", on_fail="block", on_pass="allow"
)
],
steps=[PipelineStep(guardrail="moderation-filter", on_fail="block", on_pass="allow")],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [guard]
monkeypatch.setattr(litellm, "callbacks", [guard])
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "bad content"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="content-safety",
)
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "bad content"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="content-safety",
)
assert result.terminal_action == "block"
assert isinstance(result.original_exception, HTTPException)
assert result.original_exception.status_code == 400
assert result.original_exception.detail == "Content policy violation"
finally:
litellm.callbacks = original_callbacks
assert result.terminal_action == "block"
assert isinstance(result.original_exception, HTTPException)
assert result.original_exception.status_code == 400
assert result.original_exception.detail == "Content policy violation"
@pytest.mark.asyncio
async def test_unsupported_mode_yields_error_outcome_without_exception():
async def test_unsupported_mode_yields_error_outcome_without_exception(monkeypatch):
"""An unexpected hook mode must surface as an error outcome (carrying no
original exception), not crash or run the guardrail."""
guard = AlwaysPassGuardrail(guardrail_name="filter")
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [guard]
monkeypatch.setattr(litellm, "callbacks", [guard])
try:
result = await PipelineExecutor.execute_steps(
steps=[PipelineStep(guardrail="filter", on_error="block", on_fail="block")],
mode="during_call",
data={"messages": [{"role": "user", "content": "hi"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="content-safety",
)
result = await PipelineExecutor.execute_steps(
steps=[PipelineStep(guardrail="filter", on_error="block", on_fail="block")],
mode="during_call",
data={"messages": [{"role": "user", "content": "hi"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="content-safety",
)
assert guard.calls == 0
assert result.terminal_action == "block"
assert result.step_results[0].outcome == "error"
assert (
"Unsupported pipeline mode: during_call"
in result.step_results[0].error_detail
)
assert result.original_exception is None
finally:
litellm.callbacks = original_callbacks
assert guard.calls == 0
assert result.terminal_action == "block"
assert result.step_results[0].outcome == "error"
assert "Unsupported pipeline mode: during_call" in result.step_results[0].error_detail
assert result.original_exception is None
@pytest.mark.asyncio
async def test_passthrough_guardrail_failure_can_pipeline_block():
async def test_passthrough_guardrail_failure_can_pipeline_block(monkeypatch):
"""
Pipeline: passthrough guardrail (on_fail: block)
Expected: passthrough ModifyResponseException is treated as policy fail,
@ -298,35 +279,31 @@ async def test_passthrough_guardrail_failure_can_pipeline_block():
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [passthrough_guard]
monkeypatch.setattr(litellm, "callbacks", [passthrough_guard])
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={
"model": "fake-model",
"messages": [{"role": "user", "content": "bad content"}],
},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="content-safety",
)
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={
"model": "fake-model",
"messages": [{"role": "user", "content": "bad content"}],
},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="content-safety",
)
assert passthrough_guard.calls == 1
assert result.terminal_action == "block"
assert len(result.step_results) == 1
assert result.step_results[0].guardrail_name == "passthrough-filter"
assert result.step_results[0].outcome == "fail"
assert result.step_results[0].action_taken == "block"
assert result.error_message == "Content policy violation"
finally:
litellm.callbacks = original_callbacks
assert passthrough_guard.calls == 1
assert result.terminal_action == "block"
assert len(result.step_results) == 1
assert result.step_results[0].guardrail_name == "passthrough-filter"
assert result.step_results[0].outcome == "fail"
assert result.step_results[0].action_taken == "block"
assert result.error_message == "Content policy violation"
@pytest.mark.asyncio
async def test_custom_code_guardrail_failure_can_pipeline_block():
async def test_custom_code_guardrail_failure_can_pipeline_block(monkeypatch):
"""
Pipeline: custom code guardrail (on_fail: block)
Expected: custom code keeps its standalone passthrough block behavior, and
@ -334,10 +311,7 @@ async def test_custom_code_guardrail_failure_can_pipeline_block():
"""
custom_guard = CustomCodeGuardrail(
guardrail_name="custom-code-filter",
custom_code=(
"def apply_guardrail(inputs, request_data, input_type):\n"
' return block("SSN detected")\n'
),
custom_code=('def apply_guardrail(inputs, request_data, input_type):\n return block("SSN detected")\n'),
)
pipeline = GuardrailPipeline(
@ -351,35 +325,31 @@ async def test_custom_code_guardrail_failure_can_pipeline_block():
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [custom_guard]
monkeypatch.setattr(litellm, "callbacks", [custom_guard])
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={
"model": "fake-model",
"messages": [{"role": "user", "content": "123-45-6789"}],
},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="content-safety",
)
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={
"model": "fake-model",
"messages": [{"role": "user", "content": "123-45-6789"}],
},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="content-safety",
)
assert result.terminal_action == "block"
assert len(result.step_results) == 1
assert result.step_results[0].guardrail_name == "custom-code-filter"
assert result.step_results[0].outcome == "fail"
assert result.step_results[0].action_taken == "block"
assert result.error_message == "SSN detected"
finally:
litellm.callbacks = original_callbacks
assert result.terminal_action == "block"
assert len(result.step_results) == 1
assert result.step_results[0].guardrail_name == "custom-code-filter"
assert result.step_results[0].outcome == "fail"
assert result.step_results[0].action_taken == "block"
assert result.error_message == "SSN detected"
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
@pytest.mark.asyncio
async def test_early_allow_step1_passes_step2_skipped():
async def test_early_allow_step1_passes_step2_skipped(monkeypatch):
"""
Pipeline: simple-filter (on_pass: allow) -> advanced-filter
Input: clean request that passes simple-filter
@ -396,32 +366,28 @@ async def test_early_allow_step1_passes_step2_skipped():
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [simple_guard, advanced_guard]
monkeypatch.setattr(litellm, "callbacks", [simple_guard, advanced_guard])
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "clean content"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="content-safety",
)
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "clean content"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="content-safety",
)
assert simple_guard.calls == 1
assert advanced_guard.calls == 0
assert result.terminal_action == "allow"
assert len(result.step_results) == 1
assert result.step_results[0].outcome == "pass"
assert result.step_results[0].action_taken == "allow"
finally:
litellm.callbacks = original_callbacks
assert simple_guard.calls == 1
assert advanced_guard.calls == 0
assert result.terminal_action == "allow"
assert len(result.step_results) == 1
assert result.step_results[0].outcome == "pass"
assert result.step_results[0].action_taken == "allow"
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
@pytest.mark.asyncio
async def test_escalation_step1_fails_step2_passes():
async def test_escalation_step1_fails_step2_passes(monkeypatch):
"""
Pipeline: simple-filter (on_fail: next) -> advanced-filter (on_pass: allow)
Input: request that fails simple but passes advanced
@ -438,34 +404,30 @@ async def test_escalation_step1_fails_step2_passes():
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [simple_guard, advanced_guard]
monkeypatch.setattr(litellm, "callbacks", [simple_guard, advanced_guard])
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "borderline content"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="content-safety",
)
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "borderline content"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="content-safety",
)
assert simple_guard.calls == 1
assert advanced_guard.calls == 1
assert result.terminal_action == "allow"
assert len(result.step_results) == 2
assert result.step_results[0].outcome == "fail"
assert result.step_results[0].action_taken == "next"
assert result.step_results[1].outcome == "pass"
assert result.step_results[1].action_taken == "allow"
finally:
litellm.callbacks = original_callbacks
assert simple_guard.calls == 1
assert advanced_guard.calls == 1
assert result.terminal_action == "allow"
assert len(result.step_results) == 2
assert result.step_results[0].outcome == "fail"
assert result.step_results[0].action_taken == "next"
assert result.step_results[1].outcome == "pass"
assert result.step_results[1].action_taken == "allow"
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
@pytest.mark.asyncio
async def test_data_forwarding_pii_masking():
async def test_data_forwarding_pii_masking(monkeypatch):
"""
Pipeline: pii-masker (pass_data: true, on_pass: next) -> content-check (on_pass: allow)
Input: "Hello John Smith"
@ -487,31 +449,27 @@ async def test_data_forwarding_pii_masking():
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [pii_guard, content_guard]
monkeypatch.setattr(litellm, "callbacks", [pii_guard, content_guard])
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "Hello John Smith"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="pii-then-safety",
)
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "Hello John Smith"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="pii-then-safety",
)
assert pii_guard.calls == 1
assert content_guard.calls == 1
assert content_guard.received_messages[0]["content"] == "Hello [REDACTED]"
assert result.terminal_action == "allow"
assert result.modified_data is not None
assert result.modified_data["messages"][0]["content"] == "Hello [REDACTED]"
finally:
litellm.callbacks = original_callbacks
assert pii_guard.calls == 1
assert content_guard.calls == 1
assert content_guard.received_messages[0]["content"] == "Hello [REDACTED]"
assert result.terminal_action == "allow"
assert result.modified_data is not None
assert result.modified_data["messages"][0]["content"] == "Hello [REDACTED]"
@pytest.mark.asyncio
async def test_guardrail_not_found_uses_on_fail():
async def test_guardrail_not_found_uses_on_fail(monkeypatch):
"""
If a guardrail is not found, treat as error and use on_fail action.
"""
@ -526,29 +484,25 @@ async def test_guardrail_not_found_uses_on_fail():
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = []
monkeypatch.setattr(litellm, "callbacks", [])
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="test-policy",
)
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="test-policy",
)
assert result.terminal_action == "block"
assert result.step_results[0].outcome == "error"
assert "not found" in result.step_results[0].error_detail
finally:
litellm.callbacks = original_callbacks
assert result.terminal_action == "block"
assert result.step_results[0].outcome == "error"
assert "not found" in result.step_results[0].error_detail
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
@pytest.mark.asyncio
async def test_on_error_next_fallback_on_api_outage_on_fail_blocks_content():
async def test_on_error_next_fallback_on_api_outage_on_fail_blocks_content(monkeypatch):
"""
Policy intervention (400) uses on_fail; technical error (503) uses on_error.
@ -574,32 +528,28 @@ async def test_on_error_next_fallback_on_api_outage_on_fail_blocks_content():
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [primary, fallback]
monkeypatch.setattr(litellm, "callbacks", [primary, fallback])
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "any"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="mod-fallback",
)
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "any"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="mod-fallback",
)
assert primary.calls == 1
assert fallback.calls == 1
assert result.terminal_action == "allow"
assert result.step_results[0].outcome == "error"
assert result.step_results[0].action_taken == "next"
assert result.step_results[1].outcome == "pass"
finally:
litellm.callbacks = original_callbacks
assert primary.calls == 1
assert fallback.calls == 1
assert result.terminal_action == "allow"
assert result.step_results[0].outcome == "error"
assert result.step_results[0].action_taken == "next"
assert result.step_results[1].outcome == "pass"
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
@pytest.mark.asyncio
async def test_on_fail_next_on_content_on_error_block_stops_api_fallback():
async def test_on_fail_next_on_content_on_error_block_stops_api_fallback(monkeypatch):
"""
Content policy fail (400) uses on_fail: next; API error uses on_error: block (no second step).
"""
@ -625,48 +575,40 @@ async def test_on_fail_next_on_content_on_error_block_stops_api_fallback():
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [primary_content, fallback]
monkeypatch.setattr(litellm, "callbacks", [primary_content, fallback])
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline_content.steps,
mode=pipeline_content.mode,
data={"messages": [{"role": "user", "content": "bad"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="test",
)
assert result.terminal_action == "allow"
assert primary_content.calls == 1
assert fallback.calls == 1
finally:
litellm.callbacks = original_callbacks
result = await PipelineExecutor.execute_steps(
steps=pipeline_content.steps,
mode=pipeline_content.mode,
data={"messages": [{"role": "user", "content": "bad"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="test",
)
assert result.terminal_action == "allow"
assert primary_content.calls == 1
assert fallback.calls == 1
# API outage: on_error block -> do not run fallback
fallback.calls = 0
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [primary_api, fallback]
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline_content.steps,
mode=pipeline_content.mode,
data={"messages": [{"role": "user", "content": "ok"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="test",
)
assert result.terminal_action == "block"
assert primary_api.calls == 1
assert fallback.calls == 0
assert result.step_results[0].outcome == "error"
assert result.step_results[0].action_taken == "block"
finally:
litellm.callbacks = original_callbacks
monkeypatch.setattr(litellm, "callbacks", [primary_api, fallback])
result = await PipelineExecutor.execute_steps(
steps=pipeline_content.steps,
mode=pipeline_content.mode,
data={"messages": [{"role": "user", "content": "ok"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="test",
)
assert result.terminal_action == "block"
assert primary_api.calls == 1
assert fallback.calls == 0
assert result.step_results[0].outcome == "error"
assert result.step_results[0].action_taken == "block"
@pytest.mark.asyncio
async def test_guardrail_not_found_with_next_continues():
async def test_guardrail_not_found_with_next_continues(monkeypatch):
"""
If a guardrail is not found and on_fail is 'next', continue to next step.
"""
@ -688,32 +630,28 @@ async def test_guardrail_not_found_with_next_continues():
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [pass_guard]
monkeypatch.setattr(litellm, "callbacks", [pass_guard])
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="test-policy",
)
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="test-policy",
)
assert result.terminal_action == "allow"
assert len(result.step_results) == 2
assert result.step_results[0].outcome == "error"
assert result.step_results[0].action_taken == "next"
assert result.step_results[1].outcome == "pass"
assert pass_guard.calls == 1
finally:
litellm.callbacks = original_callbacks
assert result.terminal_action == "allow"
assert len(result.step_results) == 2
assert result.step_results[0].outcome == "error"
assert result.step_results[0].action_taken == "next"
assert result.step_results[1].outcome == "pass"
assert pass_guard.calls == 1
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
@pytest.mark.asyncio
async def test_single_step_pipeline_block():
async def test_single_step_pipeline_block(monkeypatch):
"""Single step pipeline that blocks."""
guard = AlwaysFailGuardrail(guardrail_name="blocker")
@ -722,27 +660,23 @@ async def test_single_step_pipeline_block():
steps=[PipelineStep(guardrail="blocker", on_fail="block")],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [guard]
monkeypatch.setattr(litellm, "callbacks", [guard])
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="test",
)
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="test",
)
assert result.terminal_action == "block"
assert guard.calls == 1
finally:
litellm.callbacks = original_callbacks
assert result.terminal_action == "block"
assert guard.calls == 1
@pytest.mark.asyncio
async def test_single_step_pipeline_allow():
async def test_single_step_pipeline_allow(monkeypatch):
"""Single step pipeline that allows."""
guard = AlwaysPassGuardrail(guardrail_name="passer")
@ -751,27 +685,23 @@ async def test_single_step_pipeline_allow():
steps=[PipelineStep(guardrail="passer", on_pass="allow")],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [guard]
monkeypatch.setattr(litellm, "callbacks", [guard])
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="test",
)
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="test",
)
assert result.terminal_action == "allow"
assert guard.calls == 1
finally:
litellm.callbacks = original_callbacks
assert result.terminal_action == "allow"
assert guard.calls == 1
@pytest.mark.asyncio
async def test_step_results_include_duration():
async def test_step_results_include_duration(monkeypatch):
"""Step results should include timing information."""
guard = AlwaysPassGuardrail(guardrail_name="timed")
@ -780,23 +710,19 @@ async def test_step_results_include_duration():
steps=[PipelineStep(guardrail="timed")],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [guard]
monkeypatch.setattr(litellm, "callbacks", [guard])
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="test",
)
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="test",
)
assert result.step_results[0].duration_seconds is not None
assert result.step_results[0].duration_seconds >= 0
finally:
litellm.callbacks = original_callbacks
assert result.step_results[0].duration_seconds is not None
assert result.step_results[0].duration_seconds >= 0
class _PolicyOptOutGuardrail(CustomGuardrail):