mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
fix(guardrails): address Resemble review findings
This commit is contained in:
parent
cf30a36447
commit
310bbbfff5
8 changed files with 209 additions and 7 deletions
|
|
@ -43,6 +43,8 @@ guardrails:
|
|||
resemble_audio_source_tracing: true
|
||||
# Do not persist media on Resemble after the scan
|
||||
resemble_zero_retention_mode: true
|
||||
# Maximum distinct media URLs to scan per request
|
||||
resemble_max_media_urls: 10
|
||||
# Block the request if Resemble is unreachable (default: fail open)
|
||||
resemble_fail_closed: false
|
||||
```
|
||||
|
|
@ -160,6 +162,7 @@ curl -i http://0.0.0.0:4000/v1/chat/completions \
|
|||
| `resemble_use_reverse_search` | bool | `false` | (Image only) search the web for matching images to improve accuracy. |
|
||||
| `resemble_zero_retention_mode` | bool | `false` | Automatically delete submitted media after detection. URLs are redacted and filenames are tokenized. |
|
||||
| `resemble_metadata_key` | string | `"mediaUrl"` | Key under request `metadata` to read the media URL from when it is not present in the message content. |
|
||||
| `resemble_max_media_urls` | integer | `10` | Maximum number of distinct media URLs to scan per request. Requests above this limit are rejected first. |
|
||||
| `resemble_poll_interval_seconds` | number | `2.0` | How often to poll Resemble for the detection result. |
|
||||
| `resemble_poll_timeout_seconds` | number | `60.0` | Maximum total time to wait for a detection result before failing. |
|
||||
| `resemble_fail_closed` | bool | `false` | If `true`, Resemble API errors **block** the request. If `false` (default), errors are logged and ignored. |
|
||||
|
|
|
|||
|
|
@ -93,14 +93,19 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
skip_system_message=skip_system,
|
||||
)
|
||||
|
||||
# Step 2: Apply guardrail to all texts and tool calls in batch
|
||||
if texts_to_check or tool_calls_to_check:
|
||||
structured_messages = self.get_structured_messages(data)
|
||||
guardrail_input_has_media = bool(images_to_check) or self._has_media_content(
|
||||
messages=structured_messages or messages,
|
||||
skip_system_message=skip_system,
|
||||
)
|
||||
|
||||
# Step 2: Apply guardrail to all text, media, and tool calls in batch
|
||||
if texts_to_check or tool_calls_to_check or guardrail_input_has_media:
|
||||
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
|
||||
if images_to_check:
|
||||
inputs["images"] = images_to_check
|
||||
if tool_calls_to_check:
|
||||
inputs["tool_calls"] = tool_calls_to_check # type: ignore
|
||||
structured_messages = self.get_structured_messages(data)
|
||||
if structured_messages:
|
||||
inputs["structured_messages"] = (
|
||||
openai_messages_without_system(structured_messages)
|
||||
|
|
@ -220,6 +225,43 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
tool_calls_to_check.append(cast(ChatCompletionToolParam, tool_call))
|
||||
tool_call_task_mappings.append((msg_idx, int(tool_call_idx)))
|
||||
|
||||
def _has_media_content(
|
||||
self,
|
||||
messages: Optional[List[Any]],
|
||||
skip_system_message: bool = False,
|
||||
) -> bool:
|
||||
if not messages:
|
||||
return False
|
||||
|
||||
for message in messages:
|
||||
if (
|
||||
skip_system_message
|
||||
and str(message.get("role") or "").lower() == "system"
|
||||
):
|
||||
continue
|
||||
|
||||
content = message.get("content")
|
||||
if not isinstance(content, list):
|
||||
continue
|
||||
|
||||
for content_item in content:
|
||||
if not isinstance(content_item, dict):
|
||||
continue
|
||||
|
||||
if content_item.get("type") == "image_url":
|
||||
image_url = content_item.get("image_url")
|
||||
if isinstance(image_url, str):
|
||||
return True
|
||||
if isinstance(image_url, dict) and image_url.get("url"):
|
||||
return True
|
||||
|
||||
if content_item.get("type") == "input_audio":
|
||||
audio = content_item.get("input_audio")
|
||||
if isinstance(audio, dict) and audio.get("url"):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
async def _apply_guardrail_responses_to_input_texts(
|
||||
self,
|
||||
messages: List[Dict[str, Any]],
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
litellm_params, "resemble_zero_retention_mode", None
|
||||
),
|
||||
metadata_key=getattr(litellm_params, "resemble_metadata_key", None),
|
||||
max_media_urls=getattr(litellm_params, "resemble_max_media_urls", None),
|
||||
poll_interval_seconds=getattr(
|
||||
litellm_params, "resemble_poll_interval_seconds", None
|
||||
),
|
||||
|
|
|
|||
|
|
@ -91,6 +91,7 @@ class ResembleGuardrail(CustomGuardrail):
|
|||
use_reverse_search: Optional[bool] = None,
|
||||
zero_retention_mode: Optional[bool] = None,
|
||||
metadata_key: Optional[str] = None,
|
||||
max_media_urls: Optional[int] = None,
|
||||
poll_interval_seconds: Optional[float] = None,
|
||||
poll_timeout_seconds: Optional[float] = None,
|
||||
fail_closed: Optional[bool] = None,
|
||||
|
|
@ -117,6 +118,7 @@ class ResembleGuardrail(CustomGuardrail):
|
|||
self.use_reverse_search: bool = bool(use_reverse_search)
|
||||
self.zero_retention_mode: bool = bool(zero_retention_mode)
|
||||
self.metadata_key: str = metadata_key or "mediaUrl"
|
||||
self.max_media_urls: int = max_media_urls if max_media_urls is not None else 10
|
||||
self.poll_interval_seconds: float = (
|
||||
poll_interval_seconds if poll_interval_seconds is not None else 2.0
|
||||
)
|
||||
|
|
@ -128,12 +130,13 @@ class ResembleGuardrail(CustomGuardrail):
|
|||
verbose_proxy_logger.debug(
|
||||
"Resemble guardrail initialized: name=%s threshold=%s "
|
||||
"audio_source_tracing=%s use_reverse_search=%s "
|
||||
"zero_retention_mode=%s fail_closed=%s",
|
||||
"zero_retention_mode=%s max_media_urls=%s fail_closed=%s",
|
||||
kwargs.get("guardrail_name", "unknown"),
|
||||
self.threshold,
|
||||
self.audio_source_tracing,
|
||||
self.use_reverse_search,
|
||||
self.zero_retention_mode,
|
||||
self.max_media_urls,
|
||||
self.fail_closed,
|
||||
)
|
||||
|
||||
|
|
@ -218,6 +221,7 @@ class ResembleGuardrail(CustomGuardrail):
|
|||
)
|
||||
return inputs
|
||||
|
||||
self._enforce_media_url_limit(media_urls)
|
||||
for media_url in media_urls:
|
||||
await self._scan_single_url(media_url)
|
||||
return inputs
|
||||
|
|
@ -234,9 +238,29 @@ class ResembleGuardrail(CustomGuardrail):
|
|||
)
|
||||
return
|
||||
|
||||
self._enforce_media_url_limit(media_urls)
|
||||
for media_url in media_urls:
|
||||
await self._scan_single_url(media_url)
|
||||
|
||||
def _enforce_media_url_limit(self, media_urls: List[str]) -> None:
|
||||
if len(media_urls) <= self.max_media_urls:
|
||||
return
|
||||
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Too many media URLs for Resemble Detect guardrail",
|
||||
"resemble": {
|
||||
"media_url_count": len(media_urls),
|
||||
"max_media_urls": self.max_media_urls,
|
||||
"reason": (
|
||||
"Request includes more media URLs than this Resemble "
|
||||
"guardrail is configured to scan."
|
||||
),
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def _scan_single_url(self, media_url: str) -> None:
|
||||
try:
|
||||
item = await self._create_and_poll_detection(media_url)
|
||||
|
|
|
|||
|
|
@ -479,6 +479,15 @@ class ResembleGuardrailParamsConfigModel(BaseModel):
|
|||
"present in message content. Default: mediaUrl."
|
||||
),
|
||||
)
|
||||
resemble_max_media_urls: Optional[int] = Field(
|
||||
default=10,
|
||||
ge=1,
|
||||
description=(
|
||||
"Maximum number of distinct media URLs Resemble will scan per request. "
|
||||
"Requests above this limit are rejected before any Resemble API call. "
|
||||
"Default 10."
|
||||
),
|
||||
)
|
||||
resemble_poll_interval_seconds: Optional[float] = Field(
|
||||
default=2.0,
|
||||
description="How often to poll Resemble for the detection result. Default 2s.",
|
||||
|
|
|
|||
|
|
@ -56,6 +56,15 @@ class ResembleGuardrailConfigModel(GuardrailConfigModel):
|
|||
"not present in the message content. Default `mediaUrl`."
|
||||
),
|
||||
)
|
||||
resemble_max_media_urls: Optional[int] = Field(
|
||||
default=10,
|
||||
ge=1,
|
||||
description=(
|
||||
"Maximum number of distinct media URLs Resemble will scan per request. "
|
||||
"Requests above this limit are rejected before any Resemble API call. "
|
||||
"Default 10."
|
||||
),
|
||||
)
|
||||
resemble_poll_interval_seconds: Optional[float] = Field(
|
||||
default=2.0,
|
||||
description=(
|
||||
|
|
|
|||
|
|
@ -8,8 +8,7 @@ with guardrail transformations, including tool calls.
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from typing import Any, List, Literal, Optional, Tuple
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from typing import Any, Literal, Optional
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -84,6 +83,68 @@ class MockGuardrail(CustomGuardrail):
|
|||
return result
|
||||
|
||||
|
||||
class TestOpenAIChatCompletionsHandlerMediaInput:
|
||||
"""Test input processing for media-only chat messages."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_image_only_message_calls_guardrail(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = MockGuardrail()
|
||||
data = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "https://cdn.example.com/synthetic.jpg"
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert result == data
|
||||
assert guardrail.last_inputs is not None
|
||||
assert guardrail.last_inputs["texts"] == []
|
||||
assert guardrail.last_inputs["images"] == [
|
||||
"https://cdn.example.com/synthetic.jpg"
|
||||
]
|
||||
assert guardrail.last_inputs["structured_messages"] == data["messages"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_audio_only_message_calls_guardrail(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = MockGuardrail()
|
||||
data = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "input_audio",
|
||||
"input_audio": {
|
||||
"url": "https://cdn.example.com/synthetic.wav"
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert result == data
|
||||
assert guardrail.last_inputs is not None
|
||||
assert guardrail.last_inputs["texts"] == []
|
||||
assert "images" not in guardrail.last_inputs
|
||||
assert guardrail.last_inputs["structured_messages"] == data["messages"]
|
||||
|
||||
|
||||
class TestOpenAIChatCompletionsHandlerToolsInput:
|
||||
"""Test input processing with tools (function definitions)"""
|
||||
|
||||
|
|
@ -765,7 +826,7 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
|
|||
This test verifies the fix for the bug where accessing chunk.choices[0]
|
||||
would raise IndexError when a streaming chunk has an empty choices list.
|
||||
"""
|
||||
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
|
||||
from litellm.types.utils import ModelResponseStream
|
||||
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = MockPassThroughGuardrail(guardrail_name="test")
|
||||
|
|
|
|||
|
|
@ -467,6 +467,59 @@ async def test_apply_guardrail_blocks_generic_text_media_url():
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scan_request_rejects_too_many_media_urls_before_scanning():
|
||||
guard = _make_guardrail(max_media_urls=1)
|
||||
data = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": (
|
||||
"check https://cdn.example.com/one.wav and "
|
||||
"https://cdn.example.com/two.wav"
|
||||
),
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
scan_mock = AsyncMock()
|
||||
with patch.object(guard, "_scan_single_url", new=scan_mock):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guard._scan_request(data)
|
||||
|
||||
scan_mock.assert_not_called()
|
||||
assert exc_info.value.status_code == 400
|
||||
detail = exc_info.value.detail
|
||||
assert detail["resemble"]["media_url_count"] == 2
|
||||
assert detail["resemble"]["max_media_urls"] == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_rejects_too_many_media_urls_before_scanning():
|
||||
guard = _make_guardrail(max_media_urls=1)
|
||||
inputs = {
|
||||
"images": [
|
||||
"https://cdn.example.com/one.jpg",
|
||||
"https://cdn.example.com/two.jpg",
|
||||
]
|
||||
}
|
||||
|
||||
scan_mock = AsyncMock()
|
||||
with patch.object(guard, "_scan_single_url", new=scan_mock):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guard.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
scan_mock.assert_not_called()
|
||||
assert exc_info.value.status_code == 400
|
||||
detail = exc_info.value.detail
|
||||
assert detail["resemble"]["media_url_count"] == 2
|
||||
assert detail["resemble"]["max_media_urls"] == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_passes_when_audio_is_real():
|
||||
guard = _make_guardrail()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue