mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Refactor Noma guardrail to use shared Responses transformation and include system instructions (#17315)
* Support system prompts in noma guardrails * Use litellm util to covert chat completions to responses api
This commit is contained in:
parent
98a244450e
commit
71efcb7115
3 changed files with 263 additions and 156 deletions
|
|
@ -148,7 +148,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if role == "system":
|
||||
# Extract system message as instructions
|
||||
if isinstance(content, str):
|
||||
instructions = content
|
||||
if instructions:
|
||||
# Concatenate multiple system prompts with a space
|
||||
instructions = f"{instructions} {content}"
|
||||
else:
|
||||
instructions = content
|
||||
else:
|
||||
input_items.append(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -28,6 +28,9 @@ from fastapi import HTTPException
|
|||
import litellm
|
||||
from litellm import DualCache, ModelResponse
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
)
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
|
|
@ -111,6 +114,7 @@ class NomaGuardrail(CustomGuardrail):
|
|||
self.async_handler = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.GuardrailCallback
|
||||
)
|
||||
self._responses_transform_handler = LiteLLMResponsesTransformationHandler()
|
||||
self.api_key = api_key or os.environ.get("NOMA_API_KEY")
|
||||
self.api_base = api_base or os.environ.get(
|
||||
"NOMA_API_BASE", NomaGuardrail._DEFAULT_API_BASE
|
||||
|
|
@ -164,13 +168,28 @@ class NomaGuardrail(CustomGuardrail):
|
|||
start_time = datetime.now()
|
||||
extra_data = self.get_guardrail_dynamic_request_body_params(request_data)
|
||||
|
||||
user_message = await self._extract_user_message(request_data)
|
||||
if not user_message:
|
||||
messages = request_data.get("messages") or []
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
payload = {
|
||||
"input": [{"type": "message", "role": "user", "content": user_message}]
|
||||
}
|
||||
input_items, instructions = self._responses_transform_handler.convert_chat_completion_messages_to_responses_api( # type: ignore[arg-type]
|
||||
messages
|
||||
)
|
||||
|
||||
if instructions:
|
||||
system_message = {
|
||||
"type": "message",
|
||||
"role": "system",
|
||||
"content": [
|
||||
{"type": "input_text", "text": instructions},
|
||||
],
|
||||
}
|
||||
input_items.insert(0, system_message)
|
||||
|
||||
if not input_items:
|
||||
return None
|
||||
|
||||
payload = {"input": input_items}
|
||||
response_json = await self._call_noma_api(
|
||||
payload=payload,
|
||||
llm_request_id=None,
|
||||
|
|
@ -198,9 +217,9 @@ class NomaGuardrail(CustomGuardrail):
|
|||
|
||||
if self.monitor_mode:
|
||||
await self._handle_verdict_background(
|
||||
USER_ROLE, json.dumps(user_message), response_json
|
||||
USER_ROLE, json.dumps(input_items), response_json
|
||||
)
|
||||
return json.dumps(user_message)
|
||||
return json.dumps(input_items)
|
||||
|
||||
# Check if we should anonymize content
|
||||
if self._should_anonymize(response_json, USER_ROLE):
|
||||
|
|
@ -215,8 +234,8 @@ class NomaGuardrail(CustomGuardrail):
|
|||
)
|
||||
return anonymized_content
|
||||
|
||||
await self._check_verdict(USER_ROLE, json.dumps(user_message), response_json)
|
||||
return json.dumps(user_message)
|
||||
await self._check_verdict(USER_ROLE, json.dumps(input_items), response_json)
|
||||
return json.dumps(input_items)
|
||||
|
||||
async def _process_llm_response_check(
|
||||
self,
|
||||
|
|
@ -732,47 +751,6 @@ class NomaGuardrail(CustomGuardrail):
|
|||
|
||||
return response
|
||||
|
||||
async def _extract_user_message(self, data: dict) -> Optional[List[dict]]:
|
||||
"""Extract the last user message from request data"""
|
||||
messages = data.get("messages", [])
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
# Get the last user message
|
||||
user_messages = [msg for msg in messages if msg.get("role") == USER_ROLE]
|
||||
if not user_messages:
|
||||
return None
|
||||
|
||||
last_user_message = user_messages[-1].get("content", "")
|
||||
if isinstance(last_user_message, str):
|
||||
return [{"type": "input_text", "text": last_user_message}]
|
||||
elif isinstance(last_user_message, list):
|
||||
converted_messages = []
|
||||
for message in last_user_message:
|
||||
converted_message = self._convert_single_user_message_to_payload(
|
||||
message
|
||||
)
|
||||
if converted_message is not None:
|
||||
converted_messages.append(converted_message)
|
||||
return converted_messages
|
||||
else:
|
||||
return None
|
||||
|
||||
def _convert_single_user_message_to_payload(
|
||||
self, user_message: Any
|
||||
) -> Optional[dict]:
|
||||
if isinstance(user_message, str):
|
||||
return {"type": "input_text", "text": user_message}
|
||||
elif user_message.get("type", "") == "image_url":
|
||||
return {
|
||||
"type": "input_image",
|
||||
"image_url": user_message.get("image_url", {}).get("url", ""),
|
||||
}
|
||||
elif user_message.get("type", "") == "text":
|
||||
return {"type": "input_text", "text": user_message.get("text", "")}
|
||||
else:
|
||||
return None
|
||||
|
||||
async def _call_noma_api(
|
||||
self,
|
||||
payload: dict,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import copy
|
||||
import os
|
||||
from typing import cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
|
|
@ -14,6 +15,7 @@ from litellm.proxy.guardrails.guardrail_hooks.noma import (
|
|||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.noma.noma import NomaBlockedMessage
|
||||
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import Choices, Message
|
||||
|
||||
|
||||
|
|
@ -413,7 +415,7 @@ class TestNomaGuardrailHooks:
|
|||
|
||||
# Verify API call details
|
||||
call_args = mock_post.call_args
|
||||
# Verify the URL endpoint
|
||||
# Verify the URL endpoint
|
||||
assert call_args.args[0].endswith("/ai-dr/v2/prompt/scan")
|
||||
# Verify headers and JSON payload
|
||||
if "headers" in call_args.kwargs:
|
||||
|
|
@ -426,6 +428,130 @@ class TestNomaGuardrailHooks:
|
|||
assert "x-noma-context" in json_payload
|
||||
assert json_payload["x-noma-context"]["applicationId"] == "test-app"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_hook_with_system_prompt(
|
||||
self, noma_guardrail, mock_user_api_key_dict
|
||||
):
|
||||
"""Test pre-call hook includes system prompt in Noma API request"""
|
||||
request_data = {
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{"role": "user", "content": "Hello, how are you?"},
|
||||
],
|
||||
"litellm_call_id": "test-call-id",
|
||||
"metadata": {"requester_ip_address": "192.168.1.1"},
|
||||
}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"aggregatedScanResult": False, # False means safe
|
||||
"scanResult": [
|
||||
{
|
||||
"role": "system",
|
||||
"type": "message",
|
||||
"results": {}
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"type": "message",
|
||||
"results": {}
|
||||
}
|
||||
]
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
with patch.object(
|
||||
noma_guardrail.async_handler, "post", return_value=mock_response
|
||||
) as mock_post:
|
||||
result = await noma_guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
cache=MagicMock(),
|
||||
data=request_data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert result == request_data
|
||||
mock_post.assert_called_once()
|
||||
|
||||
# Verify the payload includes both system and user messages
|
||||
call_args = mock_post.call_args
|
||||
json_payload = call_args.kwargs["json"]
|
||||
assert "input" in json_payload
|
||||
messages = json_payload["input"]
|
||||
|
||||
# Should have 2 messages: system and user
|
||||
assert len(messages) == 2
|
||||
|
||||
# First message should be system
|
||||
assert messages[0]["type"] == "message"
|
||||
assert messages[0]["role"] == "system"
|
||||
assert messages[0]["content"][0]["type"] == "input_text"
|
||||
assert messages[0]["content"][0]["text"] == "You are a helpful assistant"
|
||||
|
||||
# Second message should be user
|
||||
assert messages[1]["type"] == "message"
|
||||
assert messages[1]["role"] == "user"
|
||||
assert messages[1]["content"][0]["type"] == "input_text"
|
||||
assert messages[1]["content"][0]["text"] == "Hello, how are you?"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_hook_with_multiple_system_prompts(
|
||||
self, noma_guardrail, mock_user_api_key_dict
|
||||
):
|
||||
"""Test pre-call hook combines multiple system prompts into single message"""
|
||||
request_data = {
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{"role": "system", "content": "You should be polite and respectful"},
|
||||
{"role": "user", "content": "Hello, how are you?"},
|
||||
],
|
||||
"litellm_call_id": "test-call-id",
|
||||
}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"aggregatedScanResult": False,
|
||||
"scanResult": [
|
||||
{"role": "system", "type": "message", "results": {}},
|
||||
{"role": "user", "type": "message", "results": {}}
|
||||
]
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
with patch.object(
|
||||
noma_guardrail.async_handler, "post", return_value=mock_response
|
||||
) as mock_post:
|
||||
result = await noma_guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
cache=MagicMock(),
|
||||
data=request_data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert result == request_data
|
||||
mock_post.assert_called_once()
|
||||
|
||||
# Verify the payload combines system prompts into single message
|
||||
call_args = mock_post.call_args
|
||||
json_payload = call_args.kwargs["json"]
|
||||
messages = json_payload["input"]
|
||||
|
||||
# Should have 2 messages: 1 combined system and 1 user
|
||||
assert len(messages) == 2
|
||||
|
||||
# First message should be system with combined content
|
||||
assert messages[0]["role"] == "system"
|
||||
assert messages[0]["content"][0]["type"] == "input_text"
|
||||
assert (
|
||||
messages[0]["content"][0]["text"]
|
||||
== "You are a helpful assistant You should be polite and respectful"
|
||||
)
|
||||
|
||||
# Second message should be user
|
||||
assert messages[1]["role"] == "user"
|
||||
assert messages[1]["content"][0]["type"] == "input_text"
|
||||
assert messages[1]["content"][0]["text"] == "Hello, how are you?"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_hook_blocked(
|
||||
self, noma_guardrail, mock_user_api_key_dict, mock_request_data
|
||||
|
|
@ -644,34 +770,6 @@ class TestNomaGuardrailHooks:
|
|||
|
||||
assert result == mock_request_data
|
||||
|
||||
def test_extract_user_message(self, noma_guardrail):
|
||||
data = {
|
||||
"messages": [
|
||||
{"role": "system", "content": "System prompt"},
|
||||
{"role": "user", "content": "First user message"},
|
||||
{"role": "assistant", "content": "Assistant response"},
|
||||
{"role": "user", "content": "Second user message"},
|
||||
]
|
||||
}
|
||||
|
||||
import asyncio
|
||||
|
||||
message = asyncio.run(noma_guardrail._extract_user_message(data))
|
||||
assert message == [{"type": "input_text", "text": "Second user message"}]
|
||||
|
||||
data = {"messages": [{"role": "system", "content": "System prompt"}]}
|
||||
message = asyncio.run(noma_guardrail._extract_user_message(data))
|
||||
assert message is None
|
||||
|
||||
data = {"messages": []}
|
||||
message = asyncio.run(noma_guardrail._extract_user_message(data))
|
||||
assert message is None
|
||||
|
||||
data = {}
|
||||
message = asyncio.run(noma_guardrail._extract_user_message(data))
|
||||
assert message is None
|
||||
|
||||
|
||||
class TestBackgroundProcessing:
|
||||
"""Test the new background processing functionality"""
|
||||
|
||||
|
|
@ -1025,57 +1123,66 @@ class TestNomaImageProcessing:
|
|||
metadata={},
|
||||
)
|
||||
|
||||
def test_extract_user_message_with_image_url(self, noma_guardrail):
|
||||
"""Test extracting user message with image_url content"""
|
||||
import asyncio
|
||||
|
||||
data = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "https://example.com/image.jpg"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
def test_extract_user_message_with_image_url(self):
|
||||
"""User message with only image_url becomes a single input_image content item."""
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
)
|
||||
|
||||
message = asyncio.run(noma_guardrail._extract_user_message(data))
|
||||
handler = LiteLLMResponsesTransformationHandler()
|
||||
messages: list[AllMessageValues] = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "https://example.com/image.jpg"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
|
||||
input_items, _ = handler.convert_chat_completion_messages_to_responses_api(messages)
|
||||
|
||||
assert len(input_items) == 1
|
||||
message = input_items[0]["content"]
|
||||
assert message is not None
|
||||
assert len(message) == 1
|
||||
assert message[0]["type"] == "input_image"
|
||||
assert message[0]["image_url"] == "https://example.com/image.jpg"
|
||||
|
||||
def test_extract_user_message_with_mixed_content(self, noma_guardrail):
|
||||
"""Test extracting user message with mixed text and image content"""
|
||||
import asyncio
|
||||
|
||||
data = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "What's in this image?"
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "https://example.com/image.jpg"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
def test_extract_user_message_with_mixed_content(self):
|
||||
"""User message with text + image becomes input_text then input_image in content list."""
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
)
|
||||
|
||||
message = asyncio.run(noma_guardrail._extract_user_message(data))
|
||||
handler = LiteLLMResponsesTransformationHandler()
|
||||
messages: list[AllMessageValues] = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "What's in this image?",
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "https://example.com/image.jpg"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
|
||||
input_items, _ = handler.convert_chat_completion_messages_to_responses_api(messages)
|
||||
|
||||
# Match the original assertions: `message` is the content list
|
||||
assert len(input_items) == 1
|
||||
message = input_items[0]["content"]
|
||||
assert message is not None
|
||||
assert len(message) == 2
|
||||
# First item should be text
|
||||
|
|
@ -1085,37 +1192,43 @@ class TestNomaImageProcessing:
|
|||
assert message[1]["type"] == "input_image"
|
||||
assert message[1]["image_url"] == "https://example.com/image.jpg"
|
||||
|
||||
def test_extract_user_message_with_multiple_images(self, noma_guardrail):
|
||||
"""Test extracting user message with multiple images"""
|
||||
import asyncio
|
||||
|
||||
data = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Compare these images"
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "https://example.com/image1.jpg"
|
||||
}
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "https://example.com/image2.jpg"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
def test_extract_user_message_with_multiple_images(self):
|
||||
"""User message with multiple images becomes multiple input_image items."""
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
)
|
||||
|
||||
message = asyncio.run(noma_guardrail._extract_user_message(data))
|
||||
handler = LiteLLMResponsesTransformationHandler()
|
||||
|
||||
messages: list[AllMessageValues] = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Compare these images",
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "https://example.com/image1.jpg"
|
||||
}
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "https://example.com/image2.jpg"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
|
||||
input_items, _ = handler.convert_chat_completion_messages_to_responses_api(messages)
|
||||
|
||||
# Match the original assertions
|
||||
assert len(input_items) == 1
|
||||
message = input_items[0]["content"]
|
||||
assert message is not None
|
||||
assert len(message) == 3
|
||||
assert message[0]["type"] == "input_text"
|
||||
|
|
@ -1301,8 +1414,14 @@ class TestNomaImageProcessing:
|
|||
assert exc_info.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_image_with_base64_data(self, noma_guardrail):
|
||||
async def test_image_with_base64_data(
|
||||
self, noma_guardrail, mock_user_api_key_dict
|
||||
):
|
||||
"""Test extracting image with base64 data URL"""
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
)
|
||||
|
||||
data = {
|
||||
"messages": [
|
||||
{
|
||||
|
|
@ -1319,7 +1438,13 @@ class TestNomaImageProcessing:
|
|||
]
|
||||
}
|
||||
|
||||
message = await noma_guardrail._extract_user_message(data)
|
||||
handler = LiteLLMResponsesTransformationHandler()
|
||||
messages = cast(list[AllMessageValues], data["messages"])
|
||||
|
||||
input_items, _ = handler.convert_chat_completion_messages_to_responses_api(messages)
|
||||
|
||||
assert len(input_items) == 1
|
||||
message = input_items[0]["content"]
|
||||
assert message is not None
|
||||
assert len(message) == 1
|
||||
assert message[0]["type"] == "input_image"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue