fixes for mock tests

This commit is contained in:
Ishaan Jaffer 2025-10-28 17:31:54 -07:00
parent 23b1f1afda
commit 74e4d3f6da
2 changed files with 24 additions and 20 deletions

View File

@ -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")

View File

@ -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