diff --git a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py b/tests/enterprise/litellm_enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py index 15b59724a3..186056dac9 100644 --- a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py +++ b/tests/enterprise/litellm_enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py @@ -1,26 +1,28 @@ """ Test the /guardrails/apply_guardrail endpoint """ -import sys import os -import pytest +import sys from unittest.mock import AsyncMock, Mock, patch +import pytest + sys.path.insert(0, os.path.abspath("../../../../..")) from fastapi import HTTPException -from litellm.types.guardrails import ApplyGuardrailRequest, ApplyGuardrailResponse -from litellm.proxy._types import UserAPIKeyAuth + from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.guardrails import ApplyGuardrailRequest, ApplyGuardrailResponse @pytest.mark.asyncio async def test_apply_guardrail_endpoint_returns_correct_response(): """Test that apply_guardrail endpoint returns ApplyGuardrailResponse object""" - from enterprise.litellm_enterprise.proxy.guardrails.endpoints import apply_guardrail - + from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail + # Mock the guardrail registry - with patch("enterprise.litellm_enterprise.proxy.guardrails.endpoints.GUARDRAIL_REGISTRY") as mock_registry: + with patch("litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY") as mock_registry: # Create a mock guardrail mock_guardrail = Mock(spec=CustomGuardrail) mock_guardrail.apply_guardrail = AsyncMock(return_value="Redacted text: [REDACTED] and [REDACTED]") @@ -57,10 +59,11 @@ async def test_apply_guardrail_endpoint_returns_correct_response(): @pytest.mark.asyncio async def test_apply_guardrail_endpoint_guardrail_not_found(): """Test that apply_guardrail endpoint raises exception when guardrail not found""" - from enterprise.litellm_enterprise.proxy.guardrails.endpoints import apply_guardrail - + from litellm.proxy._types import ProxyException + from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail + # Mock the guardrail registry to return None - with patch("enterprise.litellm_enterprise.proxy.guardrails.endpoints.GUARDRAIL_REGISTRY") as mock_registry: + with patch("litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY") as mock_registry: mock_registry.get_initialized_guardrail_callback.return_value = None # Create the request @@ -74,19 +77,20 @@ async def test_apply_guardrail_endpoint_guardrail_not_found(): user_api_key_dict = UserAPIKeyAuth(api_key="test-key") # Verify exception is raised - with pytest.raises(Exception) as exc_info: + with pytest.raises(ProxyException) as exc_info: await apply_guardrail(request=request, user_api_key_dict=user_api_key_dict) - assert "Guardrail non-existent-guardrail not found" in str(exc_info.value) + assert "non-existent-guardrail" in exc_info.value.message + assert "not found" in exc_info.value.message @pytest.mark.asyncio async def test_apply_guardrail_endpoint_with_presidio_guardrail(): """Test apply_guardrail endpoint with a Presidio-like guardrail""" - from enterprise.litellm_enterprise.proxy.guardrails.endpoints import apply_guardrail - + from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail + # Mock the guardrail registry - with patch("enterprise.litellm_enterprise.proxy.guardrails.endpoints.GUARDRAIL_REGISTRY") as mock_registry: + with patch("litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY") as mock_registry: # Create a mock guardrail that simulates Presidio behavior mock_guardrail = Mock(spec=CustomGuardrail) # Simulate masking PII entities @@ -121,10 +125,10 @@ async def test_apply_guardrail_endpoint_with_presidio_guardrail(): @pytest.mark.asyncio async def test_apply_guardrail_endpoint_without_optional_params(): """Test apply_guardrail endpoint without optional language and entities parameters""" - from enterprise.litellm_enterprise.proxy.guardrails.endpoints import apply_guardrail - + from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail + # Mock the guardrail registry - with patch("enterprise.litellm_enterprise.proxy.guardrails.endpoints.GUARDRAIL_REGISTRY") as mock_registry: + with patch("litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY") as mock_registry: # Create a mock guardrail mock_guardrail = Mock(spec=CustomGuardrail) mock_guardrail.apply_guardrail = AsyncMock(return_value="Processed text") diff --git a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py b/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py index b30f0bf1e8..2cf6799662 100644 --- a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py +++ b/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py @@ -144,7 +144,7 @@ async def test_bedrock_apply_guardrail_api_failure(): @pytest.mark.asyncio async def test_bedrock_apply_guardrail_endpoint_integration(): """Test the full endpoint integration with Bedrock guardrail""" - from enterprise.litellm_enterprise.proxy.guardrails.endpoints import apply_guardrail + from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail # Create a real BedrockGuardrail instance guardrail = BedrockGuardrail( @@ -154,7 +154,7 @@ async def test_bedrock_apply_guardrail_endpoint_integration(): ) # Mock the guardrail registry - with patch("enterprise.litellm_enterprise.proxy.guardrails.endpoints.GUARDRAIL_REGISTRY") as mock_registry: + with patch("litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY") as mock_registry: # Mock the make_bedrock_api_request method with patch.object(guardrail, 'make_bedrock_api_request', new_callable=AsyncMock) as mock_api_request: # Mock a successful response from Bedrock