fix(guardrails): address Resemble review findings

This commit is contained in:
devshahofficial 2026-05-26 16:39:00 -07:00
parent cf30a36447
commit 310bbbfff5
8 changed files with 209 additions and 7 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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