mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-11 16:26:07 +00:00
fix: openai moderation guardrails (#20718)
* fix: openai moderation guardrails * adds missing import * mv: test file to right place
This commit is contained in:
@@ -5,14 +5,9 @@ OpenAI Moderation Guardrail Integration for LiteLLM
|
||||
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
AsyncGenerator,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Type,
|
||||
Union,
|
||||
)
|
||||
|
||||
from fastapi import HTTPException
|
||||
@@ -20,7 +15,7 @@ from fastapi import HTTPException
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
log_guardrail_information
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
@@ -32,10 +27,8 @@ from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
from .base import OpenAIGuardrailBase
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.llms.openai import OpenAIModerationResponse
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
from litellm.types.utils import ModelResponse, ModelResponseStream
|
||||
|
||||
|
||||
class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
|
||||
@@ -236,108 +229,6 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
|
||||
# Moderation doesn't modify content, just blocks - return inputs unchanged
|
||||
return inputs
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: "UserAPIKeyAuth",
|
||||
response: Any,
|
||||
request_data: Dict[str, Any],
|
||||
) -> AsyncGenerator["ModelResponseStream", None]:
|
||||
"""
|
||||
Process streaming response chunks for OpenAI moderation.
|
||||
|
||||
Collects all chunks from the stream, assembles them into a complete response,
|
||||
and applies moderation check. If content violates moderation policy, raises HTTPException.
|
||||
"""
|
||||
# Import here to avoid circular imports
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
from litellm.main import stream_chunk_builder
|
||||
from litellm.types.utils import TextCompletionResponse
|
||||
|
||||
verbose_proxy_logger.debug("OpenAI Moderation: Running streaming response scan")
|
||||
|
||||
# Collect all chunks to process them together
|
||||
all_chunks: List["ModelResponseStream"] = []
|
||||
async for chunk in response:
|
||||
all_chunks.append(chunk)
|
||||
|
||||
# Assemble the complete response from chunks
|
||||
assembled_model_response: Optional[
|
||||
Union["ModelResponse", TextCompletionResponse]
|
||||
] = stream_chunk_builder(
|
||||
chunks=all_chunks,
|
||||
)
|
||||
|
||||
if isinstance(assembled_model_response, (type(None), TextCompletionResponse)):
|
||||
# If we can't assemble a ModelResponse or it's a text completion,
|
||||
# just yield the original chunks without moderation
|
||||
verbose_proxy_logger.warning(
|
||||
"OpenAI Moderation: Could not assemble ModelResponse from chunks, skipping moderation"
|
||||
)
|
||||
for chunk in all_chunks:
|
||||
yield chunk
|
||||
return
|
||||
|
||||
# Extract response text for moderation
|
||||
response_text = self._extract_response_text(assembled_model_response)
|
||||
if response_text:
|
||||
verbose_proxy_logger.debug(
|
||||
f"OpenAI Moderation: Streaming response text: {response_text[:100]}..." # Log first 100 chars
|
||||
)
|
||||
|
||||
# Make moderation request - this will raise HTTPException if content is flagged
|
||||
moderation_response = await self.async_make_request(
|
||||
input_text=response_text,
|
||||
)
|
||||
|
||||
# Check if content is flagged and raise exception if needed
|
||||
self._check_moderation_result(moderation_response)
|
||||
|
||||
# If we reach here, content passed moderation - yield the original chunks
|
||||
mock_response = MockResponseIterator(model_response=assembled_model_response)
|
||||
|
||||
# Return the reconstructed stream
|
||||
async for chunk in mock_response:
|
||||
yield chunk
|
||||
|
||||
def _extract_response_text(self, response: "ModelResponse") -> Optional[str]:
|
||||
"""
|
||||
Extract text content from the model response for moderation.
|
||||
"""
|
||||
if not hasattr(response, "choices") or not response.choices:
|
||||
return None
|
||||
|
||||
response_texts = []
|
||||
for choice in response.choices:
|
||||
try:
|
||||
# Try to get content from message (chat completion)
|
||||
message = getattr(choice, "message", None)
|
||||
if message:
|
||||
content = getattr(message, "content", None)
|
||||
if content and isinstance(content, str):
|
||||
response_texts.append(content)
|
||||
continue
|
||||
|
||||
# Try to get text (text completion)
|
||||
text = getattr(choice, "text", None)
|
||||
if text and isinstance(text, str):
|
||||
response_texts.append(text)
|
||||
continue
|
||||
|
||||
# Try to get content from delta (streaming)
|
||||
delta = getattr(choice, "delta", None)
|
||||
if delta:
|
||||
content = getattr(delta, "content", None)
|
||||
if content and isinstance(content, str):
|
||||
response_texts.append(content)
|
||||
continue
|
||||
|
||||
except (AttributeError, TypeError):
|
||||
# Skip choices that don't have expected attributes
|
||||
continue
|
||||
|
||||
return "\n".join(response_texts) if response_texts else None
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
|
||||
"""
|
||||
|
||||
@@ -7,7 +7,6 @@ import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../../../.."))
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
@@ -26,7 +25,7 @@ async def test_openai_moderation_guardrail_init():
|
||||
guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
)
|
||||
|
||||
|
||||
assert guardrail.guardrail_name == "test-openai-moderation"
|
||||
assert guardrail.api_key == "test-key"
|
||||
assert guardrail.model == "omni-moderation-latest"
|
||||
@@ -49,27 +48,27 @@ async def test_openai_moderation_guardrail_adds_to_litellm_callbacks():
|
||||
# Clear existing callbacks for clean test
|
||||
original_callbacks = litellm.callbacks.copy()
|
||||
litellm.logging_callback_manager._reset_all_callbacks()
|
||||
|
||||
|
||||
try:
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
guardrail_litellm_params = LitellmParams(
|
||||
guardrail=SupportedGuardrailIntegrations.OPENAI_MODERATION,
|
||||
api_key="test-key",
|
||||
model="omni-moderation-latest",
|
||||
mode="pre_call"
|
||||
mode="pre_call",
|
||||
)
|
||||
guardrail = openai_initialize_guardrail(
|
||||
litellm_params=guardrail_litellm_params,
|
||||
guardrail=Guardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
litellm_params=guardrail_litellm_params
|
||||
)
|
||||
litellm_params=guardrail_litellm_params,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
# Check that the guardrail was added to litellm callbacks
|
||||
assert guardrail in litellm.callbacks
|
||||
assert len(litellm.callbacks) == 1
|
||||
|
||||
|
||||
# Verify it's the correct guardrail
|
||||
callback = litellm.callbacks[0]
|
||||
assert isinstance(callback, OpenAIModerationGuardrail)
|
||||
@@ -85,12 +84,12 @@ async def test_openai_moderation_guardrail_adds_to_litellm_callbacks():
|
||||
async def test_openai_moderation_guardrail_safe_content():
|
||||
"""Test OpenAI moderation guardrail with safe content via apply_guardrail"""
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
)
|
||||
|
||||
|
||||
# Mock safe moderation response
|
||||
mock_response = OpenAIModerationResponse(
|
||||
id="modr-123",
|
||||
@@ -118,25 +117,29 @@ async def test_openai_moderation_guardrail_safe_content():
|
||||
"harassment": [],
|
||||
"self-harm": [],
|
||||
"violence": [],
|
||||
}
|
||||
},
|
||||
)
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
with patch.object(guardrail, 'async_make_request', return_value=mock_response):
|
||||
|
||||
with patch.object(guardrail, "async_make_request", return_value=mock_response):
|
||||
# Test apply_guardrail with safe content using structured_messages
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
structured_messages=[
|
||||
{"role": "user", "content": "Hello, how are you today?"}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={"messages": [{"role": "user", "content": "Hello, how are you today?"}]},
|
||||
input_type="request"
|
||||
request_data={
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello, how are you today?"}
|
||||
]
|
||||
},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
|
||||
# Should return the original inputs unchanged
|
||||
assert result == inputs
|
||||
|
||||
@@ -145,12 +148,12 @@ async def test_openai_moderation_guardrail_safe_content():
|
||||
async def test_openai_moderation_guardrail_apply_guardrail():
|
||||
"""Test OpenAI moderation guardrail apply_guardrail method (unified guardrail interface)"""
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
)
|
||||
|
||||
|
||||
# Mock safe moderation response
|
||||
mock_response = OpenAIModerationResponse(
|
||||
id="modr-123",
|
||||
@@ -178,37 +181,37 @@ async def test_openai_moderation_guardrail_apply_guardrail():
|
||||
"harassment": [],
|
||||
"self-harm": [],
|
||||
"violence": [],
|
||||
}
|
||||
},
|
||||
)
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
with patch.object(guardrail, 'async_make_request', return_value=mock_response):
|
||||
|
||||
with patch.object(guardrail, "async_make_request", return_value=mock_response):
|
||||
# Test apply_guardrail with texts (embeddings-style input)
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=["Hello, how are you?", "What is the weather?"]
|
||||
)
|
||||
|
||||
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
|
||||
# Should return inputs unchanged (moderation doesn't modify, only blocks)
|
||||
assert result == inputs
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_moderation_guardrail_harmful_content():
|
||||
"""Test OpenAI moderation guardrail with harmful content via apply_guardrail"""
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
)
|
||||
|
||||
|
||||
# Mock harmful moderation response
|
||||
mock_response = OpenAIModerationResponse(
|
||||
id="modr-123",
|
||||
@@ -236,40 +239,51 @@ async def test_openai_moderation_guardrail_harmful_content():
|
||||
"harassment": [],
|
||||
"self-harm": [],
|
||||
"violence": [],
|
||||
}
|
||||
},
|
||||
)
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
with patch.object(guardrail, 'async_make_request', return_value=mock_response):
|
||||
|
||||
with patch.object(guardrail, "async_make_request", return_value=mock_response):
|
||||
# Test apply_guardrail with harmful content using structured_messages
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
structured_messages=[
|
||||
{"role": "user", "content": "This is hateful content"}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
# Should raise HTTPException
|
||||
from fastapi import HTTPException
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={"messages": [{"role": "user", "content": "This is hateful content"}]},
|
||||
input_type="request"
|
||||
request_data={
|
||||
"messages": [
|
||||
{"role": "user", "content": "This is hateful content"}
|
||||
]
|
||||
},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "Violated OpenAI moderation policy" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_moderation_guardrail_streaming_safe_content():
|
||||
"""Test OpenAI moderation guardrail with streaming safe content"""
|
||||
"""Test OpenAI moderation guardrail with streaming safe content via UnifiedLLMGuardrails"""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
event_hook="post_call",
|
||||
)
|
||||
|
||||
unified_guardrail = UnifiedLLMGuardrails()
|
||||
|
||||
# Mock safe moderation response
|
||||
mock_response = OpenAIModerationResponse(
|
||||
id="modr-123",
|
||||
@@ -297,72 +311,85 @@ async def test_openai_moderation_guardrail_streaming_safe_content():
|
||||
"harassment": [],
|
||||
"self-harm": [],
|
||||
"violence": [],
|
||||
}
|
||||
},
|
||||
)
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
# Mock streaming chunks
|
||||
async def mock_stream():
|
||||
# Simulate streaming chunks with safe content
|
||||
chunks = [
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="Hello "))]),
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="world"))]),
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="!"))])
|
||||
]
|
||||
for chunk in chunks:
|
||||
chunk1 = MagicMock()
|
||||
chunk1.model = "gpt-4"
|
||||
chunk1.choices = [MagicMock()]
|
||||
chunk1.choices[0].delta = MagicMock()
|
||||
chunk1.choices[0].delta.content = "Hello "
|
||||
chunk1.choices[0].finish_reason = None
|
||||
|
||||
chunk2 = MagicMock()
|
||||
chunk2.model = "gpt-4"
|
||||
chunk2.choices = [MagicMock()]
|
||||
chunk2.choices[0].delta = MagicMock()
|
||||
chunk2.choices[0].delta.content = "world"
|
||||
chunk2.choices[0].finish_reason = None
|
||||
|
||||
# Last chunk with finish_reason
|
||||
chunk3 = MagicMock()
|
||||
chunk3.model = "gpt-4"
|
||||
chunk3.choices = [MagicMock()]
|
||||
chunk3.choices[0].delta = MagicMock()
|
||||
chunk3.choices[0].delta.content = "!"
|
||||
chunk3.choices[0].finish_reason = "stop"
|
||||
|
||||
for chunk in [chunk1, chunk2, chunk3]:
|
||||
yield chunk
|
||||
|
||||
# Mock the stream_chunk_builder to return a proper ModelResponse
|
||||
|
||||
# Mock for stream_chunk_builder
|
||||
mock_model_response = MagicMock()
|
||||
mock_model_response.choices = [
|
||||
MagicMock(message=MagicMock(content="Hello world!"))
|
||||
]
|
||||
|
||||
with patch.object(guardrail, 'async_make_request', return_value=mock_response), \
|
||||
patch('litellm.main.stream_chunk_builder', return_value=mock_model_response), \
|
||||
patch('litellm.llms.base_llm.base_model_iterator.MockResponseIterator') as mock_iterator:
|
||||
|
||||
# Mock the iterator to yield the original chunks
|
||||
async def mock_yield_chunks():
|
||||
chunks = [
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="Hello "))]),
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="world"))]),
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="!"))])
|
||||
]
|
||||
for chunk in chunks:
|
||||
yield chunk
|
||||
|
||||
mock_iterator.return_value.__aiter__ = lambda self: mock_yield_chunks()
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test")
|
||||
mock_model_response.choices = [MagicMock()]
|
||||
mock_model_response.choices[0].message = MagicMock()
|
||||
mock_model_response.choices[0].message.content = "Hello world!"
|
||||
|
||||
with patch.object(guardrail, "async_make_request", return_value=mock_response), patch(
|
||||
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
|
||||
return_value=mock_model_response,
|
||||
):
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test", request_route="/chat/completions"
|
||||
)
|
||||
request_data = {
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello, how are you today?"}
|
||||
]
|
||||
"messages": [{"role": "user", "content": "Hello, how are you today?"}],
|
||||
"guardrail_to_apply": guardrail,
|
||||
"metadata": {"guardrails": ["test-openai-moderation"]},
|
||||
}
|
||||
|
||||
# Test streaming hook with safe content
|
||||
|
||||
# Test streaming hook with safe content via UnifiedLLMGuardrails
|
||||
result_chunks = []
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
async for chunk in unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=mock_stream(),
|
||||
request_data=request_data
|
||||
request_data=request_data,
|
||||
):
|
||||
result_chunks.append(chunk)
|
||||
|
||||
|
||||
# Should return all chunks without blocking
|
||||
assert len(result_chunks) == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_moderation_guardrail_streaming_harmful_content():
|
||||
"""Test OpenAI moderation guardrail with streaming harmful content"""
|
||||
"""Test OpenAI moderation guardrail with streaming harmful content via UnifiedLLMGuardrails"""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
event_hook="post_call",
|
||||
)
|
||||
|
||||
unified_guardrail = UnifiedLLMGuardrails()
|
||||
|
||||
# Mock harmful moderation response
|
||||
mock_response = OpenAIModerationResponse(
|
||||
id="modr-123",
|
||||
@@ -390,46 +417,74 @@ async def test_openai_moderation_guardrail_streaming_harmful_content():
|
||||
"harassment": [],
|
||||
"self-harm": [],
|
||||
"violence": [],
|
||||
}
|
||||
},
|
||||
)
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
# Mock streaming chunks with harmful content
|
||||
async def mock_stream():
|
||||
chunks = [
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="This is "))]),
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="harmful content"))])
|
||||
]
|
||||
for chunk in chunks:
|
||||
# First chunk - no finish_reason
|
||||
chunk1 = MagicMock()
|
||||
chunk1.model = "gpt-4"
|
||||
chunk1.choices = [MagicMock()]
|
||||
chunk1.choices[0].delta = MagicMock()
|
||||
chunk1.choices[0].delta.content = "This is "
|
||||
chunk1.choices[0].finish_reason = None
|
||||
|
||||
# Last chunk - with finish_reason to signal end of stream
|
||||
chunk2 = MagicMock()
|
||||
chunk2.model = "gpt-4"
|
||||
chunk2.choices = [MagicMock()]
|
||||
chunk2.choices[0].delta = MagicMock()
|
||||
chunk2.choices[0].delta.content = "harmful content"
|
||||
chunk2.choices[0].finish_reason = "stop"
|
||||
|
||||
for chunk in [chunk1, chunk2]:
|
||||
yield chunk
|
||||
|
||||
# Mock the stream_chunk_builder to return a ModelResponse with harmful content
|
||||
mock_model_response = MagicMock()
|
||||
mock_model_response.choices = [
|
||||
MagicMock(message=MagicMock(content="This is harmful content"))
|
||||
]
|
||||
|
||||
with patch.object(guardrail, 'async_make_request', return_value=mock_response), \
|
||||
patch('litellm.main.stream_chunk_builder', return_value=mock_model_response):
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test")
|
||||
|
||||
# Mock for stream_chunk_builder - use real litellm types so isinstance checks pass
|
||||
from litellm.types.utils import ModelResponse
|
||||
import litellm
|
||||
mock_model_response = ModelResponse(
|
||||
id="mock-response",
|
||||
model="gpt-4",
|
||||
choices=[
|
||||
litellm.Choices(
|
||||
index=0,
|
||||
message=litellm.Message(
|
||||
role="assistant",
|
||||
content="This is harmful content",
|
||||
),
|
||||
finish_reason="stop",
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
with patch.object(guardrail, "async_make_request", return_value=mock_response), patch(
|
||||
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
|
||||
return_value=mock_model_response,
|
||||
):
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test", request_route="/chat/completions"
|
||||
)
|
||||
request_data = {
|
||||
"messages": [
|
||||
{"role": "user", "content": "Generate harmful content"}
|
||||
]
|
||||
"messages": [{"role": "user", "content": "Generate harmful content"}],
|
||||
"guardrail_to_apply": guardrail,
|
||||
"metadata": {"guardrails": ["test-openai-moderation"]},
|
||||
}
|
||||
|
||||
|
||||
# Should raise HTTPException when processing streaming harmful content
|
||||
from fastapi import HTTPException
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
result_chunks = []
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
async for chunk in unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=mock_stream(),
|
||||
request_data=request_data
|
||||
request_data=request_data,
|
||||
):
|
||||
result_chunks.append(chunk)
|
||||
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "Violated OpenAI moderation policy" in str(exc_info.value.detail)
|
||||
assert "Violated OpenAI moderation policy" in str(exc_info.value.detail)
|
||||
|
||||
+172
@@ -0,0 +1,172 @@
|
||||
import pytest
|
||||
from unittest.mock import MagicMock, patch
|
||||
import os
|
||||
from litellm.proxy.guardrails.guardrail_hooks.openai.moderations import (
|
||||
OpenAIModerationGuardrail,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
from litellm.types.utils import ModelResponseStream, ModelResponse
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_moderation_guardrail_streaming_latency():
|
||||
"""
|
||||
Test that the OpenAI Moderation guardrail, when run via UnifiedLLMGuardrails,
|
||||
supports streaming (fast time-to-first-token) instead of buffering.
|
||||
"""
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
# 1. Initialize the specific guardrail with proper event_hook
|
||||
openai_guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
event_hook="post_call",
|
||||
)
|
||||
|
||||
# 2. Initialize the Unified Guardrail system (which invokes the specific guardrail)
|
||||
unified_guardrail = UnifiedLLMGuardrails()
|
||||
|
||||
# Mock safe moderation response
|
||||
mock_mod_response = MagicMock()
|
||||
mock_mod_response.results = []
|
||||
|
||||
# Mock streaming chunks (no artificial delay - test deterministically)
|
||||
async def mock_stream():
|
||||
chunks_data = ["Hello", " ", "world", "!", " Goodbye"]
|
||||
for i, content in enumerate(chunks_data):
|
||||
chunk = MagicMock(spec=ModelResponseStream)
|
||||
chunk.model = "gpt-4"
|
||||
choice = MagicMock()
|
||||
choice.delta = MagicMock()
|
||||
choice.delta.content = content
|
||||
# Last chunk gets finish_reason
|
||||
choice.finish_reason = "stop" if i == len(chunks_data) - 1 else None
|
||||
chunk.choices = [choice]
|
||||
yield chunk
|
||||
|
||||
# Mock for stream_chunk_builder to return a simple ModelResponse
|
||||
mock_model_response = MagicMock(spec=ModelResponse)
|
||||
mock_model_response.choices = [MagicMock()]
|
||||
mock_model_response.choices[0].message = MagicMock()
|
||||
mock_model_response.choices[0].message.content = "Hello world! Goodbye"
|
||||
|
||||
# Patch the network call in the specific guardrail
|
||||
with patch.object(
|
||||
openai_guardrail, "async_make_request", return_value=mock_mod_response
|
||||
), patch(
|
||||
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
|
||||
return_value=mock_model_response,
|
||||
):
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test", request_route="/chat/completions"
|
||||
)
|
||||
request_data = {
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"guardrail_to_apply": openai_guardrail,
|
||||
"metadata": {
|
||||
"guardrails": ["test-openai-moderation"],
|
||||
"guardrail_config": {"streaming_sampling_rate": 1},
|
||||
}, # Check every chunk for test
|
||||
}
|
||||
|
||||
chunks_received = 0
|
||||
first_chunk_yielded = False
|
||||
|
||||
# Call the hook on UnifiedLLMGuardrails
|
||||
async for chunk in unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=mock_stream(),
|
||||
request_data=request_data,
|
||||
):
|
||||
if not first_chunk_yielded:
|
||||
first_chunk_yielded = True
|
||||
chunks_received += 1
|
||||
|
||||
# Deterministic assertions (no flaky timing checks)
|
||||
assert first_chunk_yielded, "Expected at least one chunk to be yielded"
|
||||
assert chunks_received == 5, f"Expected 5 chunks, got {chunks_received}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_moderation_guardrail_streaming_harmful_content():
|
||||
"""
|
||||
Test that harmful content is caught during streaming via UnifiedLLMGuardrails
|
||||
"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
openai_guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
event_hook="post_call",
|
||||
)
|
||||
unified_guardrail = UnifiedLLMGuardrails()
|
||||
|
||||
# Mock harmful moderation response
|
||||
mock_mod_response = MagicMock()
|
||||
mock_mod_response.results = [
|
||||
MagicMock(
|
||||
flagged=True, categories={"hate": True}, category_scores={"hate": 0.99}
|
||||
)
|
||||
]
|
||||
|
||||
async def mock_stream():
|
||||
chunks_data = ["This ", "is ", "harmful ", "content"]
|
||||
for i, content in enumerate(chunks_data):
|
||||
chunk = MagicMock(spec=ModelResponseStream)
|
||||
chunk.model = "gpt-4"
|
||||
choice = MagicMock()
|
||||
choice.delta = MagicMock()
|
||||
choice.delta.content = content
|
||||
# Last chunk gets finish_reason
|
||||
choice.finish_reason = "stop" if i == len(chunks_data) - 1 else None
|
||||
chunk.choices = [choice]
|
||||
yield chunk
|
||||
|
||||
# Mock for stream_chunk_builder - use real litellm types so isinstance checks pass
|
||||
import litellm
|
||||
|
||||
mock_model_response = ModelResponse(
|
||||
id="mock-response",
|
||||
model="gpt-4",
|
||||
choices=[
|
||||
litellm.Choices(
|
||||
index=0,
|
||||
message=litellm.Message(
|
||||
role="assistant",
|
||||
content="This is harmful content",
|
||||
),
|
||||
finish_reason="stop",
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
openai_guardrail, "async_make_request", return_value=mock_mod_response
|
||||
), patch(
|
||||
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
|
||||
return_value=mock_model_response,
|
||||
):
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test", request_route="/chat/completions"
|
||||
)
|
||||
request_data = {
|
||||
"messages": [{"role": "user", "content": "generate hate"}],
|
||||
"guardrail_to_apply": openai_guardrail,
|
||||
"metadata": {
|
||||
"guardrails": ["test-openai-moderation"],
|
||||
"guardrail_config": {"streaming_sampling_rate": 1},
|
||||
},
|
||||
}
|
||||
|
||||
# Should raise HTTPException
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=mock_stream(),
|
||||
request_data=request_data,
|
||||
):
|
||||
pass
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "Violated OpenAI moderation policy" in str(exc_info.value.detail)
|
||||
Reference in New Issue
Block a user