test(guardrails): stop five guardrail test files leaking env vars on failure (#37828)

Onyx, prompt security, hiddenlayer, repelloai and deepkeep all write straight to
os.environ and unset again at the bottom of each test. None of the five has a
try/finally, so the moment a test fails it returns to the runner with the keys
still set and whatever runs next in that worker inherits them.

Raising inside test_onyx_guard_with_custom_timeout_from_kwargs on the current
files leaves ONYX_API_BASE and ONYX_API_KEY behind; doing the same in
test_hiddenlayer_config_saas leaves HIDDENLAYER_API_BASE. Both come back clean
after this.

89 raw writes and the hand-rolled deletes become monkeypatch calls. The
class-level setup_method and teardown_method pair in the onyx file, sweeping the
same three keys twice, becomes one autouse fixture. The sys.path.insert lines
and their now-unused imports go too, and litellm.set_verbose = True, which only
turned global debug logging on for whatever ran next, is dropped rather than
restored.

test_onyx_guard_config and test_prompt_security_guard_config asserted nothing at
all, so they could only fail by raising. Each now pins what init_guardrails_v2
produces: exactly one guardrail of the right class on litellm.callbacks,
carrying the configured name, default_on and hook. The zero-assert tests in the
other three are left alone; those are a judgement about each guardrail rather
than a mechanical sweep.

tests/test_litellm/proxy/guardrails passes at 2873.
This commit is contained in:
yuneng-jiang 2026-08-21 21:29:31 -07:00 • committed by GitHub
parent 6bce3dce0d
commit 39a580aa91
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 99 additions and 223 deletions

View file

@ -1,18 +1,18 @@
{
"TQ001": {
"limit": 746
"limit": 744
},
"TQ002": {
"limit": 742
},
"TQ003": {
"limit": 1073
"limit": 1068
},
"TQ004": {
"limit": 469
},
"TQ005": {
"limit": 2562
"limit": 2549
},
"TQ006": {
"limit": 34

View file

@ -1,10 +1,8 @@
import os
import sys
import pytest
from unittest.mock import patch, MagicMock, AsyncMock
from httpx import Response, Request
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm.proxy.guardrails.guardrail_hooks.deepkeep.deepkeep import (
@ -17,10 +15,9 @@ from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
from litellm.exceptions import GuardrailRaisedException
def test_deepkeep_guard_config(monkeypatch):
def test_deepkeep_guard_config(monkeypatch: pytest.MonkeyPatch):
"""Test DeepKeep guard configuration with init_guardrails_v2."""
litellm.set_verbose = True
litellm.guardrail_name_config_map = {}
monkeypatch.setattr(litellm, "guardrail_name_config_map", {})
monkeypatch.setenv("DEEPKEEP_API_KEY", "test-key")
monkeypatch.setenv("DEEPKEEP_API_BASE", "https://test.deepkeep.ai")
@ -42,9 +39,6 @@ def test_deepkeep_guard_config(monkeypatch):
)
# Clean up
del os.environ["DEEPKEEP_API_KEY"]
del os.environ["DEEPKEEP_API_BASE"]
del os.environ["DEEPKEEP_FIREWALL_ID"]
class TestDeepKeepGuardrail:
@ -108,7 +102,7 @@ class TestDeepKeepGuardrail:
== "https://test.deepkeep.ai/v3/openai/beta/litellm_basic_guardrail_api"
)
def test_initialization_with_env_vars(self, monkeypatch):
def test_initialization_with_env_vars(self, monkeypatch: pytest.MonkeyPatch):
"""should initialize successfully using environment variables."""
monkeypatch.setenv("DEEPKEEP_API_KEY", "env-key")
monkeypatch.setenv("DEEPKEEP_API_BASE", "https://env.deepkeep.ai")

View file

@ -1,5 +1,4 @@
import os
import sys
import uuid
from typing import List, cast
from unittest.mock import AsyncMock, MagicMock, patch
@ -8,7 +7,6 @@ import pytest
from fastapi import HTTPException
from httpx import Request, Response
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm import ModelResponse
@ -26,10 +24,9 @@ from litellm.types.utils import (
)
def test_hiddenlayer_config_saas(monkeypatch):
def test_hiddenlayer_config_saas(monkeypatch: pytest.MonkeyPatch):
"""Test Hiddenlayer SaaS configuration with init_guardrails_v2."""
litellm.set_verbose = True
litellm.guardrail_name_config_map = {}
monkeypatch.setattr(litellm, "guardrail_name_config_map", {})
# Set environment variables for testing
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -50,8 +47,6 @@ def test_hiddenlayer_config_saas(monkeypatch):
)
# Clean up
if "HIDDENLAYER_API_BASE" in os.environ:
del os.environ["HIDDENLAYER_API_BASE"]
class TestHiddenlayerGuardrail:
@ -71,7 +66,7 @@ class TestHiddenlayerGuardrail:
if key in os.environ:
del os.environ[key]
def test_initialization(self, monkeypatch):
def test_initialization(self, monkeypatch: pytest.MonkeyPatch):
"""Test successful initialization with default values."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -84,17 +79,16 @@ class TestHiddenlayerGuardrail:
assert guardrail.guardrail_name == "hiddenlayer"
assert guardrail.event_hook == "pre_call"
def test_initialization_fails_when_api_key_missing(self):
def test_initialization_fails_when_api_key_missing(self, monkeypatch: pytest.MonkeyPatch):
"""Test that initialization fails when API key is not set."""
# Ensure API key is not set
if "HIDDENLAYER_CLIENT_SECRET" in os.environ:
del os.environ["HIDDENLAYER_CLIENT_SECRET"]
monkeypatch.delenv("HIDDENLAYER_CLIENT_SECRET", raising=False)
with pytest.raises(RuntimeError):
HiddenlayerGuardrail(guardrail_name="hiddenlayer", event_hook="pre_call")
@pytest.mark.asyncio
async def test_apply_guardrail_request_no_violations(self, monkeypatch):
async def test_apply_guardrail_request_no_violations(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail for request with no violations detected."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -151,7 +145,7 @@ class TestHiddenlayerGuardrail:
assert call_args.args[0] == f"{guardrail.api_base}/detection/v1/interactions"
@pytest.mark.asyncio
async def test_apply_guardrail_request_with_violations(self, monkeypatch):
async def test_apply_guardrail_request_with_violations(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail for request with violations detected."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -209,7 +203,7 @@ class TestHiddenlayerGuardrail:
assert "Blocked by Hiddenlayer" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_apply_guardrail_response_no_violations(self, monkeypatch):
async def test_apply_guardrail_response_no_violations(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail for response with no violations detected."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -279,7 +273,7 @@ class TestHiddenlayerGuardrail:
mock_post.assert_called_once()
@pytest.mark.asyncio
async def test_apply_guardrail_response_with_violations(self, monkeypatch):
async def test_apply_guardrail_response_with_violations(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail for response with violations detected."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -348,7 +342,7 @@ class TestHiddenlayerGuardrail:
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_apply_guardrail_api_error_handling(self, monkeypatch):
async def test_apply_guardrail_api_error_handling(self, monkeypatch: pytest.MonkeyPatch):
"""Test handling of API errors in apply_guardrail."""
# Set required API key
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -391,7 +385,7 @@ class TestHiddenlayerGuardrail:
assert result == inputs
@pytest.mark.asyncio
async def test_validate_with_call_hiddenlayer_method(self, monkeypatch):
async def test_validate_with_call_hiddenlayer_method(self, monkeypatch: pytest.MonkeyPatch):
"""Test the _validate_with_guard_server internal method."""
# Set required API key
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -433,7 +427,7 @@ class TestHiddenlayerGuardrail:
)
@pytest.mark.asyncio
async def test_apply_guardrail_request_with_image(self, monkeypatch):
async def test_apply_guardrail_request_with_image(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail sends multimodal content (image) to HiddenLayer v1."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -498,7 +492,7 @@ class TestHiddenlayerGuardrail:
assert result is not None
@pytest.mark.asyncio
async def test_apply_guardrail_redact_with_image_content(self, monkeypatch):
async def test_apply_guardrail_redact_with_image_content(self, monkeypatch: pytest.MonkeyPatch):
"""Test that REDACT action with multimodal content extracts text properly into inputs['texts']."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -570,10 +564,9 @@ class TestHiddenlayerGuardrail:
assert config_model.__name__ == "HiddenlayerGuardrailConfigModel"
def test_hiddenlayer_config_v2(monkeypatch):
def test_hiddenlayer_config_v2(monkeypatch: pytest.MonkeyPatch):
"""Test HiddenLayer V2 configuration with init_guardrails_v2."""
litellm.set_verbose = True
litellm.guardrail_name_config_map = {}
monkeypatch.setattr(litellm, "guardrail_name_config_map", {})
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -593,8 +586,6 @@ def test_hiddenlayer_config_v2(monkeypatch):
config_file_path="",
)
if "HIDDENLAYER_API_BASE" in os.environ:
del os.environ["HIDDENLAYER_API_BASE"]
class TestHiddenlayerGuardrailV2:
@ -612,7 +603,7 @@ class TestHiddenlayerGuardrailV2:
if key in os.environ:
del os.environ[key]
def test_initialization(self, monkeypatch):
def test_initialization(self, monkeypatch: pytest.MonkeyPatch):
"""Test successful initialization with default values."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -624,16 +615,15 @@ class TestHiddenlayerGuardrailV2:
assert guardrail.guardrail_name == "hiddenlayer"
assert guardrail.event_hook == "pre_call"
def test_initialization_fails_when_api_key_missing(self):
def test_initialization_fails_when_api_key_missing(self, monkeypatch: pytest.MonkeyPatch):
"""Test that initialization fails when API key is not set for SaaS."""
if "HIDDENLAYER_CLIENT_SECRET" in os.environ:
del os.environ["HIDDENLAYER_CLIENT_SECRET"]
monkeypatch.delenv("HIDDENLAYER_CLIENT_SECRET", raising=False)
with pytest.raises(RuntimeError):
HiddenlayerGuardrailV2(guardrail_name="hiddenlayer", event_hook="pre_call")
@pytest.mark.asyncio
async def test_apply_guardrail_request_no_violations(self, monkeypatch):
async def test_apply_guardrail_request_no_violations(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail for request with no violations detected."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -691,7 +681,7 @@ class TestHiddenlayerGuardrailV2:
assert "detection/v2/request-evaluations" in call_args.args[0]
@pytest.mark.asyncio
async def test_apply_guardrail_request_with_violations(self, monkeypatch):
async def test_apply_guardrail_request_with_violations(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail for request with violations detected (block via header)."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -751,7 +741,7 @@ class TestHiddenlayerGuardrailV2:
assert "Blocked by Hiddenlayer" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_apply_guardrail_response_no_violations(self, monkeypatch):
async def test_apply_guardrail_response_no_violations(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail for response with no violations detected."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -816,7 +806,7 @@ class TestHiddenlayerGuardrailV2:
assert "detection/v2/response-evaluations" in call_args.args[0]
@pytest.mark.asyncio
async def test_apply_guardrail_response_with_violations(self, monkeypatch):
async def test_apply_guardrail_response_with_violations(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail for response with violations detected (block via header)."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -863,7 +853,7 @@ class TestHiddenlayerGuardrailV2:
assert "Blocked by Hiddenlayer" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_apply_guardrail_response_with_tool_calls(self, monkeypatch):
async def test_apply_guardrail_response_with_tool_calls(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail for response containing tool calls."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -924,7 +914,7 @@ class TestHiddenlayerGuardrailV2:
assert "detection/v2/response-evaluations" in call_args.args[0]
@pytest.mark.asyncio
async def test_call_hiddenlayer_uses_correct_endpoints(self, monkeypatch):
async def test_call_hiddenlayer_uses_correct_endpoints(self, monkeypatch: pytest.MonkeyPatch):
"""Test that _call_hiddenlayer uses the v2 request/response evaluation endpoints."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -959,7 +949,7 @@ class TestHiddenlayerGuardrailV2:
assert "detection/v2/response-evaluations" in mock_post.call_args.args[0]
@pytest.mark.asyncio
async def test_apply_guardrail_request_with_image(self, monkeypatch):
async def test_apply_guardrail_request_with_image(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail sends multimodal content (image) to HiddenLayer v2."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
@ -1030,7 +1020,7 @@ class TestHiddenlayerGuardrailV2:
assert texts == ["how much is on this receipt?"]
@pytest.mark.asyncio
async def test_apply_guardrail_request_with_image_multimodal_response(self, monkeypatch):
async def test_apply_guardrail_request_with_image_multimodal_response(self, monkeypatch: pytest.MonkeyPatch):
"""Test that new_texts extraction handles multimodal content (list) returned by HiddenLayer v2."""
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")

View file

@ -1,5 +1,3 @@
import os
import sys
import uuid
from unittest.mock import AsyncMock, MagicMock, patch
@ -8,8 +6,6 @@ import pytest
from fastapi import HTTPException
from httpx import Request, Response
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm import ModelResponse
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -18,12 +14,11 @@ from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
from litellm.types.utils import Choices, GenericGuardrailAPIInputs, Message
def test_onyx_guard_config(monkeypatch):
def test_onyx_guard_config(monkeypatch: pytest.MonkeyPatch):
"""Test Onyx guard configuration with init_guardrails_v2."""
litellm.set_verbose = True
litellm.guardrail_name_config_map = {}
monkeypatch.setattr(litellm, "guardrail_name_config_map", {})
monkeypatch.setattr(litellm, "callbacks", [])
# Set environment variables for testing
monkeypatch.setenv("ONYX_API_BASE", "https://test.onyx.security")
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
@ -41,16 +36,15 @@ def test_onyx_guard_config(monkeypatch):
config_file_path="",
)
# Clean up
if "ONYX_API_BASE" in os.environ:
del os.environ["ONYX_API_BASE"]
if "ONYX_API_KEY" in os.environ:
del os.environ["ONYX_API_KEY"]
registered = [c for c in litellm.callbacks if isinstance(c, OnyxGuardrail)]
assert len(registered) == 1
assert registered[0].guardrail_name == "onyx-guard"
assert registered[0].default_on is True
assert registered[0].event_hook == "pre_call"
def test_onyx_guard_with_custom_timeout_from_kwargs(monkeypatch):
def test_onyx_guard_with_custom_timeout_from_kwargs(monkeypatch: pytest.MonkeyPatch):
"""Test Onyx guard instantiation with custom timeout passed via kwargs."""
# Set environment variables for testing
monkeypatch.setenv("ONYX_API_BASE", "https://test.onyx.security")
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
@ -74,20 +68,13 @@ def test_onyx_guard_with_custom_timeout_from_kwargs(monkeypatch):
assert timeout_param.read == 45.0
assert timeout_param.connect == 5.0
# Clean up
if "ONYX_API_BASE" in os.environ:
del os.environ["ONYX_API_BASE"]
if "ONYX_API_KEY" in os.environ:
del os.environ["ONYX_API_KEY"]
def test_onyx_guard_with_timeout_none_uses_env_var(monkeypatch):
def test_onyx_guard_with_timeout_none_uses_env_var(monkeypatch: pytest.MonkeyPatch):
"""Test Onyx guard with timeout=None uses ONYX_TIMEOUT env var.
When timeout=None is passed (as it would be from config model with default None),
the ONYX_TIMEOUT environment variable should be used.
"""
# Set environment variables for testing
monkeypatch.setenv("ONYX_API_BASE", "https://test.onyx.security")
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
monkeypatch.setenv("ONYX_TIMEOUT", "60")
@ -112,23 +99,13 @@ def test_onyx_guard_with_timeout_none_uses_env_var(monkeypatch):
assert timeout_param.read == 60.0
assert timeout_param.connect == 5.0
# Clean up
if "ONYX_API_BASE" in os.environ:
del os.environ["ONYX_API_BASE"]
if "ONYX_API_KEY" in os.environ:
del os.environ["ONYX_API_KEY"]
if "ONYX_TIMEOUT" in os.environ:
del os.environ["ONYX_TIMEOUT"]
def test_onyx_guard_with_timeout_none_defaults_to_10(monkeypatch):
def test_onyx_guard_with_timeout_none_defaults_to_10(monkeypatch: pytest.MonkeyPatch):
"""Test Onyx guard with timeout=None and no env var defaults to 10 seconds."""
# Set environment variables for testing
monkeypatch.setenv("ONYX_API_BASE", "https://test.onyx.security")
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
# Ensure ONYX_TIMEOUT is not set
if "ONYX_TIMEOUT" in os.environ:
del os.environ["ONYX_TIMEOUT"]
monkeypatch.delenv("ONYX_TIMEOUT", raising=False)
with patch(
"litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client"
@ -150,33 +127,17 @@ def test_onyx_guard_with_timeout_none_defaults_to_10(monkeypatch):
assert timeout_param.read == 10.0
assert timeout_param.connect == 5.0
# Clean up
if "ONYX_API_BASE" in os.environ:
del os.environ["ONYX_API_BASE"]
if "ONYX_API_KEY" in os.environ:
del os.environ["ONYX_API_KEY"]
class TestOnyxGuardrail:
"""Test suite for Onyx Security Guardrail integration."""
def setup_method(self):
"""Setup test environment."""
# Clean up any existing environment variables
for key in ["ONYX_API_BASE", "ONYX_API_KEY", "ONYX_TIMEOUT"]:
if key in os.environ:
del os.environ[key]
@pytest.fixture(autouse=True)
def clear_onyx_env(self, monkeypatch: pytest.MonkeyPatch) -> None:
for key in ("ONYX_API_BASE", "ONYX_API_KEY", "ONYX_TIMEOUT"):
monkeypatch.delenv(key, raising=False)
def teardown_method(self):
"""Clean up test environment."""
# Clean up any environment variables set during tests
for key in ["ONYX_API_BASE", "ONYX_API_KEY", "ONYX_TIMEOUT"]:
if key in os.environ:
del os.environ[key]
def test_initialization_with_defaults(self, monkeypatch):
def test_initialization_with_defaults(self, monkeypatch: pytest.MonkeyPatch):
"""Test successful initialization with default values."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
guardrail = OnyxGuardrail(
@ -189,7 +150,7 @@ class TestOnyxGuardrail:
assert guardrail.guardrail_name == "test-guard"
assert guardrail.event_hook == "pre_call"
def test_initialization_with_env_vars(self, monkeypatch):
def test_initialization_with_env_vars(self, monkeypatch: pytest.MonkeyPatch):
"""Test initialization with environment variables."""
monkeypatch.setenv("ONYX_API_BASE", "https://custom.onyx.security")
monkeypatch.setenv("ONYX_API_KEY", "custom-api-key")
@ -202,18 +163,17 @@ class TestOnyxGuardrail:
assert guardrail.api_key == "custom-api-key"
assert guardrail.event_hook == "post_call"
def test_initialization_fails_when_api_key_missing(self):
def test_initialization_fails_when_api_key_missing(self, monkeypatch: pytest.MonkeyPatch):
"""Test that initialization fails when API key is not set."""
# Ensure API key is not set
if "ONYX_API_KEY" in os.environ:
del os.environ["ONYX_API_KEY"]
monkeypatch.delenv("ONYX_API_KEY", raising=False)
with pytest.raises(
ValueError, match="ONYX_API_KEY environment variable is not set"
):
OnyxGuardrail(guardrail_name="test-guard", event_hook="pre_call")
def test_initialization_with_default_timeout(self, monkeypatch):
def test_initialization_with_default_timeout(self, monkeypatch: pytest.MonkeyPatch):
"""Test that default timeout is 10.0 seconds."""
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
@ -232,7 +192,7 @@ class TestOnyxGuardrail:
assert timeout_param.read == 10.0
assert timeout_param.connect == 5.0
def test_initialization_with_custom_timeout_parameter(self, monkeypatch):
def test_initialization_with_custom_timeout_parameter(self, monkeypatch: pytest.MonkeyPatch):
"""Test initialization with custom timeout parameter."""
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
@ -254,7 +214,7 @@ class TestOnyxGuardrail:
assert timeout_param.read == 30.0
assert timeout_param.connect == 5.0
def test_initialization_with_timeout_from_env_var(self, monkeypatch):
def test_initialization_with_timeout_from_env_var(self, monkeypatch: pytest.MonkeyPatch):
"""Test initialization with timeout from ONYX_TIMEOUT environment variable.
Note: The env var is only used when timeout=None is explicitly passed,
@ -282,7 +242,7 @@ class TestOnyxGuardrail:
assert timeout_param.read == 25.0
assert timeout_param.connect == 5.0
def test_initialization_timeout_parameter_overrides_env_var(self, monkeypatch):
def test_initialization_timeout_parameter_overrides_env_var(self, monkeypatch: pytest.MonkeyPatch):
"""Test that timeout parameter overrides ONYX_TIMEOUT environment variable."""
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
monkeypatch.setenv("ONYX_TIMEOUT", "25")
@ -306,9 +266,8 @@ class TestOnyxGuardrail:
assert timeout_param.connect == 5.0
@pytest.mark.asyncio
async def test_apply_guardrail_request_no_violations(self, monkeypatch):
async def test_apply_guardrail_request_no_violations(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail for request with no violations detected."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
# Setup guardrail
@ -372,9 +331,8 @@ class TestOnyxGuardrail:
assert call_args.kwargs["json"]["conversation_id"] == "test-call-id"
@pytest.mark.asyncio
async def test_apply_guardrail_request_with_violations(self, monkeypatch):
async def test_apply_guardrail_request_with_violations(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail for request with violations detected."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
# Setup guardrail
@ -423,9 +381,8 @@ class TestOnyxGuardrail:
assert "prompt_injection" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_apply_guardrail_response_no_violations(self, monkeypatch):
async def test_apply_guardrail_response_no_violations(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail for response with no violations detected."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
# Setup guardrail
@ -497,9 +454,8 @@ class TestOnyxGuardrail:
assert call_args.kwargs["json"]["conversation_id"] == "test-call-id-2"
@pytest.mark.asyncio
async def test_apply_guardrail_response_with_violations(self, monkeypatch):
async def test_apply_guardrail_response_with_violations(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail for response with violations detected."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
# Setup guardrail
@ -558,9 +514,8 @@ class TestOnyxGuardrail:
assert "illegal_instructions" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_apply_guardrail_api_error_handling(self, monkeypatch):
async def test_apply_guardrail_api_error_handling(self, monkeypatch: pytest.MonkeyPatch):
"""Test handling of API errors in apply_guardrail."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
guardrail = OnyxGuardrail(
@ -591,9 +546,8 @@ class TestOnyxGuardrail:
assert result == inputs
@pytest.mark.asyncio
async def test_apply_guardrail_timeout_error_handling(self, monkeypatch):
async def test_apply_guardrail_timeout_error_handling(self, monkeypatch: pytest.MonkeyPatch):
"""Test handling of timeout errors in apply_guardrail (graceful degradation)."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
guardrail = OnyxGuardrail(
@ -629,9 +583,8 @@ class TestOnyxGuardrail:
assert result == inputs
@pytest.mark.asyncio
async def test_apply_guardrail_read_timeout_error_handling(self, monkeypatch):
async def test_apply_guardrail_read_timeout_error_handling(self, monkeypatch: pytest.MonkeyPatch):
"""Test handling of read timeout errors in apply_guardrail."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
guardrail = OnyxGuardrail(
@ -667,9 +620,8 @@ class TestOnyxGuardrail:
assert result == inputs
@pytest.mark.asyncio
async def test_apply_guardrail_connect_timeout_error_handling(self, monkeypatch):
async def test_apply_guardrail_connect_timeout_error_handling(self, monkeypatch: pytest.MonkeyPatch):
"""Test handling of connect timeout errors in apply_guardrail."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
guardrail = OnyxGuardrail(
@ -705,9 +657,8 @@ class TestOnyxGuardrail:
assert result == inputs
@pytest.mark.asyncio
async def test_apply_guardrail_no_logging_obj(self, monkeypatch):
async def test_apply_guardrail_no_logging_obj(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail without logging object (uses UUID)."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
guardrail = OnyxGuardrail(
@ -747,9 +698,8 @@ class TestOnyxGuardrail:
assert call_args.kwargs["json"]["conversation_id"] == "test-uuid"
@pytest.mark.asyncio
async def test_validate_with_guard_server_method(self, monkeypatch):
async def test_validate_with_guard_server_method(self, monkeypatch: pytest.MonkeyPatch):
"""Test the _validate_with_guard_server internal method."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
guardrail = OnyxGuardrail(
@ -788,9 +738,8 @@ class TestOnyxGuardrail:
)
@pytest.mark.asyncio
async def test_validate_with_guard_server_blocked(self, monkeypatch):
async def test_validate_with_guard_server_blocked(self, monkeypatch: pytest.MonkeyPatch):
"""Test _validate_with_guard_server when request is blocked."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
guardrail = OnyxGuardrail(
@ -825,9 +774,8 @@ class TestOnyxGuardrail:
assert config_model.__name__ == "OnyxGuardrailConfigModel"
@pytest.mark.asyncio
async def test_apply_guardrail_with_modelresponse(self, monkeypatch):
async def test_apply_guardrail_with_modelresponse(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail with ModelResponse object for response type."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
guardrail = OnyxGuardrail(
@ -880,9 +828,8 @@ class TestOnyxGuardrail:
assert "payload" in call_args.kwargs["json"]
@pytest.mark.asyncio
async def test_apply_guardrail_response_error_handling(self, monkeypatch):
async def test_apply_guardrail_response_error_handling(self, monkeypatch: pytest.MonkeyPatch):
"""Test error handling when processing response data."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
guardrail = OnyxGuardrail(
@ -925,7 +872,7 @@ class TestOnyxIntegration:
"""Test integration scenarios."""
@pytest.mark.asyncio
async def test_full_guardrail_flow(self, monkeypatch):
async def test_full_guardrail_flow(self, monkeypatch: pytest.MonkeyPatch):
"""Test full guardrail flow with multiple hooks."""
# Set environment variables
monkeypatch.setenv("ONYX_API_BASE", "https://test.onyx.security")
@ -966,16 +913,10 @@ class TestOnyxIntegration:
)
assert len(custom_loggers) >= 3
# Clean up
if "ONYX_API_BASE" in os.environ:
del os.environ["ONYX_API_BASE"]
if "ONYX_API_KEY" in os.environ:
del os.environ["ONYX_API_KEY"]
@pytest.mark.asyncio
async def test_apply_guardrail_empty_request_data(self, monkeypatch):
async def test_apply_guardrail_empty_request_data(self, monkeypatch: pytest.MonkeyPatch):
"""Test apply_guardrail with empty request data."""
# Set required API key
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
guardrail = OnyxGuardrail(

View file

@ -1,11 +1,9 @@
import os
import sys
import pytest
from fastapi import HTTPException
from httpx import ConnectError, Request, Response
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm import DualCache
@ -93,23 +91,23 @@ class TestRepelloAIInitialization:
with pytest.raises(ValueError, match="asset_id"):
RepelloAIGuardrail(api_key="test-api-key", guardrail_name="t")
def test_api_key_from_env(self, monkeypatch):
def test_api_key_from_env(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("REPELLOAI_API_KEY", "env-key")
guardrail = RepelloAIGuardrail(asset_id="asset-123", guardrail_name="t")
assert guardrail.repelloai_api_key == "env-key"
def test_api_key_from_argus_env(self, monkeypatch):
def test_api_key_from_argus_env(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("ARGUS_API_KEY", "argus-key")
guardrail = RepelloAIGuardrail(asset_id="asset-123", guardrail_name="t")
assert guardrail.repelloai_api_key == "argus-key"
def test_argus_env_preferred_over_legacy(self, monkeypatch):
def test_argus_env_preferred_over_legacy(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("ARGUS_API_KEY", "argus-key")
monkeypatch.setenv("REPELLOAI_API_KEY", "legacy-key")
guardrail = RepelloAIGuardrail(asset_id="asset-123", guardrail_name="t")
assert guardrail.repelloai_api_key == "argus-key"
def test_explicit_api_key_preferred_over_env(self, monkeypatch):
def test_explicit_api_key_preferred_over_env(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("ARGUS_API_KEY", "argus-key")
guardrail = RepelloAIGuardrail(
api_key="explicit-key", asset_id="asset-123", guardrail_name="t"
@ -145,9 +143,9 @@ class TestRepelloAIInitialization:
assert guardrail.api_base == DEFAULT_REPELLOAI_API_BASE
assert guardrail.unreachable_fallback == "fail_closed"
def test_init_guardrails_v2_wiring(self, monkeypatch):
def test_init_guardrails_v2_wiring(self, monkeypatch: pytest.MonkeyPatch):
"""The guardrail registers and constructs via the config.yaml path."""
litellm.guardrail_name_config_map = {}
monkeypatch.setattr(litellm, "guardrail_name_config_map", {})
monkeypatch.setenv("REPELLOAI_API_KEY", "test-key")
init_guardrails_v2(
all_guardrails=[

View file

@ -1,5 +1,3 @@
import os
import sys
from fastapi.exceptions import HTTPException
from unittest.mock import patch, AsyncMock
from httpx import Response, Request
@ -12,19 +10,15 @@ from litellm.proxy.guardrails.guardrail_hooks.prompt_security.prompt_security im
PromptSecurityGuardrail,
)
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import litellm
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
def test_prompt_security_guard_config(monkeypatch):
def test_prompt_security_guard_config(monkeypatch: pytest.MonkeyPatch):
"""Test guardrail initialization with proper configuration"""
litellm.set_verbose = True
litellm.guardrail_name_config_map = {}
monkeypatch.setattr(litellm, "guardrail_name_config_map", {})
monkeypatch.setattr(litellm, "callbacks", [])
# Set environment variables for testing
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
@ -42,21 +36,19 @@ def test_prompt_security_guard_config(monkeypatch):
config_file_path="",
)
# Clean up
del os.environ["PROMPT_SECURITY_API_KEY"]
del os.environ["PROMPT_SECURITY_API_BASE"]
registered = [c for c in litellm.callbacks if isinstance(c, PromptSecurityGuardrail)]
assert len(registered) == 1
assert registered[0].guardrail_name == "prompt_security"
assert registered[0].default_on is True
assert registered[0].event_hook == "during_call"
def test_prompt_security_guard_config_no_api_key():
def test_prompt_security_guard_config_no_api_key(monkeypatch: pytest.MonkeyPatch):
"""Test that initialization fails when API key is missing"""
litellm.set_verbose = True
litellm.guardrail_name_config_map = {}
monkeypatch.setattr(litellm, "guardrail_name_config_map", {})
# Ensure API key is not in environment
if "PROMPT_SECURITY_API_KEY" in os.environ:
del os.environ["PROMPT_SECURITY_API_KEY"]
if "PROMPT_SECURITY_API_BASE" in os.environ:
del os.environ["PROMPT_SECURITY_API_BASE"]
monkeypatch.delenv("PROMPT_SECURITY_API_KEY", raising=False)
monkeypatch.delenv("PROMPT_SECURITY_API_BASE", raising=False)
with pytest.raises(
PromptSecurityGuardrailMissingSecrets,
@ -78,7 +70,7 @@ def test_prompt_security_guard_config_no_api_key():
@pytest.mark.asyncio
async def test_apply_guardrail_block_request(monkeypatch):
async def test_apply_guardrail_block_request(monkeypatch: pytest.MonkeyPatch):
"""Test that apply_guardrail blocks malicious prompts"""
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
@ -126,13 +118,9 @@ async def test_apply_guardrail_block_request(monkeypatch):
assert "prompt_injection" in str(excinfo.value.detail)
assert "jailbreak" in str(excinfo.value.detail)
# Clean up
del os.environ["PROMPT_SECURITY_API_KEY"]
del os.environ["PROMPT_SECURITY_API_BASE"]
@pytest.mark.asyncio
async def test_apply_guardrail_modify_request(monkeypatch):
async def test_apply_guardrail_modify_request(monkeypatch: pytest.MonkeyPatch):
"""Test that apply_guardrail modifies prompts when needed"""
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
@ -177,13 +165,9 @@ async def test_apply_guardrail_modify_request(monkeypatch):
assert result["texts"] == ["User prompt with PII: SSN [REDACTED]"]
# Clean up
del os.environ["PROMPT_SECURITY_API_KEY"]
del os.environ["PROMPT_SECURITY_API_BASE"]
@pytest.mark.asyncio
async def test_apply_guardrail_allow_request(monkeypatch):
async def test_apply_guardrail_allow_request(monkeypatch: pytest.MonkeyPatch):
"""Test that apply_guardrail allows safe prompts"""
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
@ -220,13 +204,9 @@ async def test_apply_guardrail_allow_request(monkeypatch):
assert result == inputs
# Clean up
del os.environ["PROMPT_SECURITY_API_KEY"]
del os.environ["PROMPT_SECURITY_API_BASE"]
@pytest.mark.asyncio
async def test_apply_guardrail_block_response(monkeypatch):
async def test_apply_guardrail_block_response(monkeypatch: pytest.MonkeyPatch):
"""Test that apply_guardrail blocks malicious responses"""
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
@ -267,13 +247,9 @@ async def test_apply_guardrail_block_response(monkeypatch):
assert "Blocked by Prompt Security" in str(excinfo.value.detail)
assert "pii_exposure" in str(excinfo.value.detail)
# Clean up
del os.environ["PROMPT_SECURITY_API_KEY"]
del os.environ["PROMPT_SECURITY_API_BASE"]
@pytest.mark.asyncio
async def test_apply_guardrail_modify_response(monkeypatch):
async def test_apply_guardrail_modify_response(monkeypatch: pytest.MonkeyPatch):
"""Test that apply_guardrail modifies responses when needed"""
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
@ -311,13 +287,9 @@ async def test_apply_guardrail_modify_response(monkeypatch):
assert result["texts"] == ["Your SSN is [REDACTED]"]
# Clean up
del os.environ["PROMPT_SECURITY_API_KEY"]
del os.environ["PROMPT_SECURITY_API_BASE"]
@pytest.mark.asyncio
async def test_file_sanitization(monkeypatch):
async def test_file_sanitization(monkeypatch: pytest.MonkeyPatch):
"""Test file sanitization for images"""
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
@ -401,13 +373,9 @@ async def test_file_sanitization(monkeypatch):
# Should complete without errors and return the data
assert result is not None
# Clean up
del os.environ["PROMPT_SECURITY_API_KEY"]
del os.environ["PROMPT_SECURITY_API_BASE"]
@pytest.mark.asyncio
async def test_file_sanitization_block(monkeypatch):
async def test_file_sanitization_block(monkeypatch: pytest.MonkeyPatch):
"""Test that file sanitization blocks malicious files"""
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
@ -485,13 +453,9 @@ async def test_file_sanitization_block(monkeypatch):
assert "File blocked by Prompt Security" in str(excinfo.value.detail)
assert "malware_detected" in str(excinfo.value.detail)
# Clean up
del os.environ["PROMPT_SECURITY_API_KEY"]
del os.environ["PROMPT_SECURITY_API_BASE"]
@pytest.mark.asyncio
async def test_user_api_key_alias_forwarding(monkeypatch):
async def test_user_api_key_alias_forwarding(monkeypatch: pytest.MonkeyPatch):
"""Test that user API key alias is properly sent via headers and payload"""
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
@ -530,12 +494,9 @@ async def test_user_api_key_alias_forwarding(monkeypatch):
payload = call_kwargs["json"]
assert payload["user"] == "vk-alias"
del os.environ["PROMPT_SECURITY_API_KEY"]
del os.environ["PROMPT_SECURITY_API_BASE"]
@pytest.mark.asyncio
async def test_role_filtering(monkeypatch):
async def test_role_filtering(monkeypatch: pytest.MonkeyPatch):
"""Test that tool/function messages are filtered out by default"""
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
@ -594,13 +555,9 @@ async def test_role_filtering(monkeypatch):
assert len(sent_messages) == 3
assert all(msg["role"] in ["system", "user", "assistant"] for msg in sent_messages)
# Clean up
del os.environ["PROMPT_SECURITY_API_KEY"]
del os.environ["PROMPT_SECURITY_API_BASE"]
@pytest.mark.asyncio
async def test_check_tool_results_enabled(monkeypatch):
async def test_check_tool_results_enabled(monkeypatch: pytest.MonkeyPatch):
"""Test with check_tool_results=True: transforms tool/function to 'other' role"""
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
@ -680,7 +637,3 @@ async def test_check_tool_results_enabled(monkeypatch):
assert "indirect_prompt_injection" in str(excinfo.value.detail)
# Clean up
del os.environ["PROMPT_SECURITY_API_KEY"]
del os.environ["PROMPT_SECURITY_API_BASE"]
del os.environ["PROMPT_SECURITY_CHECK_TOOL_RESULTS"]