From 6bce3dce0dd1ebcf1e0ea9d5a7bb206568d9e3d9 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Fri, 21 Aug 2026 21:19:35 -0700 Subject: [PATCH] test(callbacks): unwind the callbacks global the policy engine and realtime tests scaffold around (#37826) * 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 --- test-quality-budget.json | 4 +- .../test_realtime_streaming.py | 134 ++-- .../policy_engine/test_pipeline_executor.py | 570 ++++++++---------- 3 files changed, 310 insertions(+), 398 deletions(-) diff --git a/test-quality-budget.json b/test-quality-budget.json index 91e8f39d195..46e368a495b 100644 --- a/test-quality-budget.json +++ b/test-quality-budget.json @@ -6,13 +6,13 @@ "limit": 742 }, "TQ003": { - "limit": 1074 + "limit": 1073 }, "TQ004": { "limit": 469 }, "TQ005": { - "limit": 2621 + "limit": 2562 }, "TQ006": { "limit": 34 diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index ccf353b1b6c..61b63e2b917 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -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 diff --git a/tests/test_litellm/proxy/policy_engine/test_pipeline_executor.py b/tests/test_litellm/proxy/policy_engine/test_pipeline_executor.py index 840d93eb12c..22d212dd8ae 100644 --- a/tests/test_litellm/proxy/policy_engine/test_pipeline_executor.py +++ b/tests/test_litellm/proxy/policy_engine/test_pipeline_executor.py @@ -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):