mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
6bce3dce0d
commit
39a580aa91
6 changed files with 99 additions and 223 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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=[
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue