From ab85eb8f43458af6b355506618596000b97d53d2 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 12 Jan 2026 17:27:13 +0530 Subject: [PATCH] Fix: gaurdrail moderation support with responses API --- .../guardrail_hooks/openai/moderations.py | 44 +++++++-- .../test_openai_moderations_hook.py | 98 +++++++++++++++++++ 2 files changed, 134 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py index aad2a69ca52..852cf01bc09 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py +++ b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py @@ -178,6 +178,33 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): }, ) + def _extract_user_content_from_data(self, data: Dict[str, Any]) -> Optional[str]: + """ + Extract user content from request data, supporting both Chat Completions and Responses API. + + For Chat Completions: extracts from 'messages' field + For Responses API: extracts from 'input' field + + Returns: + The extracted user content string, or None if no content found + """ + # Try to get messages first (Chat Completions API) + messages: Optional[List["AllMessageValues"]] = data.get("messages") + if messages is not None: + return self.get_user_prompt(messages) + + # Try to get input (Responses API) + input_data = data.get("input") + if input_data is not None: + # input can be a string or a list of message-like objects + if isinstance(input_data, str): + return input_data + elif isinstance(input_data, list): + # Treat input as messages and extract user content + return self.get_user_prompt(input_data) + + return None + @log_guardrail_information async def async_pre_call_hook( self, @@ -210,14 +237,14 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): if call_type == "moderation": return data - new_messages: Optional[List["AllMessageValues"]] = data.get("messages") - if new_messages is None: + user_prompt = self._extract_user_content_from_data(data) + + if user_prompt is None: verbose_proxy_logger.warning( - "OpenAI Moderation: not running guardrail. No messages in data" + "OpenAI Moderation: not running guardrail. No messages or input in data" ) return data - user_prompt = self.get_user_prompt(new_messages) if user_prompt: verbose_proxy_logger.debug( f"OpenAI Moderation: User prompt: {user_prompt[:100]}..." # Log first 100 chars for debugging @@ -265,14 +292,15 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): if call_type == "moderation": return data - new_messages: Optional[List["AllMessageValues"]] = data.get("messages") - if new_messages is None: + # Extract user content from either messages or input field + user_prompt = self._extract_user_content_from_data(data) + + if user_prompt is None: verbose_proxy_logger.warning( - "OpenAI Moderation: not running guardrail. No messages in data" + "OpenAI Moderation: not running guardrail. No messages or input in data" ) return data - user_prompt = self.get_user_prompt(new_messages) if user_prompt: moderation_response = await self.async_make_request( input_text=user_prompt, diff --git a/tests/local_testing/test_openai_moderations_hook.py b/tests/local_testing/test_openai_moderations_hook.py index 5d2cab38b33..3acd36f32f0 100644 --- a/tests/local_testing/test_openai_moderations_hook.py +++ b/tests/local_testing/test_openai_moderations_hook.py @@ -85,3 +85,101 @@ async def test_openai_moderation_error_raising(monkeypatch): print("Got exception: ", e) assert "Violated content safety policy" in str(e) pass + + +@pytest.mark.asyncio +async def test_openai_moderation_responses_api_input_field(): + """ + Tests that OpenAI Moderation works with Responses API input field. + + This test verifies the fix for the issue where moderation was skipped + for Responses API because it only checked for 'messages' field but + Responses API uses 'input' field instead. + """ + from unittest.mock import AsyncMock, MagicMock, patch + from litellm.types.llms.openai import ( + OpenAIModerationResponse, + OpenAIModerationResult, + ) + from litellm.proxy.guardrails.guardrail_hooks.openai.moderations import ( + OpenAIModerationGuardrail, + ) + + # Initialize the open-source OpenAI Moderation guardrail + openai_mod = OpenAIModerationGuardrail( + guardrail_name="openai-moderation-test", + api_key="fake-key-for-testing", + model="omni-moderation-latest", + ) + + _api_key = "sk-12345" + _api_key = hash_token("sk-12345") + user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) + + # Mock the async_make_request to return a flagged response + mock_moderation_response = OpenAIModerationResponse( + id="modr-123", + model="omni-moderation-latest", + results=[ + OpenAIModerationResult( + flagged=True, + categories={"violence": True, "hate": False}, + category_scores={"violence": 0.95, "hate": 0.1}, + category_applied_input_types=None, + ) + ], + ) + + with patch.object( + openai_mod, "async_make_request", return_value=mock_moderation_response + ): + # Test 1: Responses API with input as string + try: + await openai_mod.async_moderation_hook( + data={ + "model": "gpt-4o", + "input": "I want to hurt people", + }, + user_api_key_dict=user_api_key_dict, + call_type="responses", + ) + pytest.fail("Should have raised HTTPException for flagged content") + except Exception as e: + print("Got exception for string input: ", e) + assert "Violated OpenAI moderation policy" in str(e) + + # Test 2: Responses API with input as list of messages + try: + await openai_mod.async_moderation_hook( + data={ + "model": "gpt-4o", + "input": [ + {"role": "user", "content": "I want to hurt people"} + ], + }, + user_api_key_dict=user_api_key_dict, + call_type="responses", + ) + pytest.fail("Should have raised HTTPException for flagged content") + except Exception as e: + print("Got exception for list input: ", e) + assert "Violated OpenAI moderation policy" in str(e) + + # Test 3: Verify it still works with messages field (Chat Completions) + try: + await openai_mod.async_moderation_hook( + data={ + "model": "gpt-4o", + "messages": [ + {"role": "user", "content": "I want to hurt people"} + ], + }, + user_api_key_dict=user_api_key_dict, + call_type="completion", + ) + pytest.fail("Should have raised HTTPException for flagged content") + except Exception as e: + print("Got exception for messages field: ", e) + assert "Violated OpenAI moderation policy" in str(e) + + print("✓ All Responses API moderation tests passed!")