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:
idola9 2025-12-02 05:56:14 +02:00 • committed by GitHub
parent 98a244450e
commit 71efcb7115
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 263 additions and 156 deletions

View file

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

View file

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

View file

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