Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
mask_request_content=litellm_params.mask_request_content,
mask_response_content=litellm_params.mask_response_content,
fail_on_error=litellm_params.fail_on_error,
skip_unscannable_attachments=litellm_params.skip_unscannable_attachments,
)
litellm.logging_callback_manager.add_litellm_callback(_model_armor_callback)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -25,10 +25,6 @@

MODEL_ARMOR_MAX_FILE_SIZE_BYTES = 4 * 1024 * 1024

# Hard cap on how many attachments a single request may submit to Model Armor, to bound
# per-request fan-out (latency and quota).
MAX_FILE_ATTACHMENTS_PER_REQUEST = 10

_REMOTE_URI_SCHEMES = ("gs://", "http://", "https://")

ModelArmorByteDataType = Literal["PDF", "WORD_DOCUMENT", "EXCEL_DOCUMENT", "POWERPOINT_DOCUMENT", "CSV", "TXT"]
Expand Down
30 changes: 16 additions & 14 deletions litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,6 @@
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import (
MAX_FILE_ATTACHMENTS_PER_REQUEST,
MODEL_ARMOR_MAX_FILE_SIZE_BYTES,
plan_file_scans,
)
Expand Down Expand Up @@ -383,10 +382,14 @@ async def _scan_request_files(self, messages: Sequence[AllMessageValues], data:

Each attachment is sent through the byte API and a MATCH_FOUND raises a 400 before the
request reaches the LLM. File scanning does not support masking (Model Armor returns
findings, not a sanitized document), so it only blocks. Anything the guardrail cannot
scan - a file_id or remote URL reference with no inline bytes, a document over the 4 MB
byte limit, or more attachments than the per-request cap - is a guardrail failure and
blocks unless the operator has opted into fail-open via fail_on_error=False.
findings, not a sanitized document), so it only blocks. A file_id or remote URL reference
with no inline bytes and a document over the 4 MB byte limit are guardrail failures that
block unless the operator has opted into fail-open via fail_on_error=False.

skip_unscannable_attachments decouples reference-only attachments from fail_on_error: when
enabled, attachments Model Armor cannot scan (file_id, gs://, or http(s) references with no
inline bytes, and inline content whose base64 will not decode) pass through instead of
blocking, while fail_on_error still governs real Model Armor API errors.
"""
from litellm.proxy.common_utils.callback_utils import (
_get_or_create_proxy_metadata_bucket,
Expand All @@ -395,7 +398,14 @@ async def _scan_request_files(self, messages: Sequence[AllMessageValues], data:

plan = plan_file_scans(messages)
attachments = plan.attachments
unscannable_references = plan.unscannable_count
skip_unscannable = bool(self.optional_params.get("skip_unscannable_attachments", False))
if skip_unscannable and plan.unscannable_count > 0:
verbose_proxy_logger.warning(
"Model Armor: allowing %d unscannable attachment(s) through because "
"skip_unscannable_attachments is enabled",
plan.unscannable_count,
)
unscannable_references = 0 if skip_unscannable else plan.unscannable_count
if not attachments and unscannable_references == 0:
return

Expand All @@ -415,14 +425,6 @@ async def _scan_request_files(self, messages: Sequence[AllMessageValues], data:
metadata["_model_armor_status"] = "blocked"
raise self._unscannable_block_error(reason)

if len(attachments) > MAX_FILE_ATTACHMENTS_PER_REQUEST:
reason = f"{len(attachments)} attachments exceed the per-request scan limit of {MAX_FILE_ATTACHMENTS_PER_REQUEST}"
verbose_proxy_logger.warning("Model Armor: %s", reason)
if fail_on_error:
metadata["_model_armor_status"] = "blocked"
raise self._unscannable_block_error(reason)
attachments = attachments[:MAX_FILE_ATTACHMENTS_PER_REQUEST]

for attachment in attachments:
if len(attachment.file_bytes) > MODEL_ARMOR_MAX_FILE_SIZE_BYTES:
reason = (
Expand Down
8 changes: 8 additions & 0 deletions litellm/types/guardrails.py
Original file line number Diff line number Diff line change
Expand Up @@ -800,6 +800,14 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
"so only a valid guardrail response can block or modify it."
),
)
skip_unscannable_attachments: Optional[bool] = Field(
default=False,
description=(
"Implemented by guardrail='model_armor'. When True, attachment references that carry no "
"inline bytes (file_id, gs://, or http(s) URLs) pass through unscanned instead of blocking, "
"while fail_on_error still governs real Model Armor API errors. Default False blocks them."
),
)

additional_provider_specific_params: Optional[Dict[str, Any]] = Field(
default=None,
Expand Down
155 changes: 141 additions & 14 deletions tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py
Original file line number Diff line number Diff line change
Expand Up @@ -2205,39 +2205,166 @@ async def test_pre_call_file_id_reference_skipped_when_fail_open():


@pytest.mark.asyncio
async def test_pre_call_blocks_when_attachment_count_exceeds_cap():
"""More attachments than the per-request cap fail closed by default to bound scan fan-out."""
from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import (
MAX_FILE_ATTACHMENTS_PER_REQUEST,
async def test_pre_call_file_id_reference_passthrough_when_skip_unscannable_enabled():
"""skip_unscannable_attachments lets a file_id reference through even with fail_on_error=True."""
guardrail = _make_guardrail(skip_unscannable_attachments=True)
request_data = {
"model": "gpt-4",
"messages": [
{
"role": "user",
"content": [
{"type": "text", "text": "summarize this"},
{"type": "file", "file": {"file_id": "file-abc123"}},
],
}
],
"metadata": {"guardrails": ["model-armor-test"]},
}

with patch.object(
guardrail.async_handler,
"post",
AsyncMock(return_value=_armor_response(blocked=False)),
) as mock_post:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=MagicMock(spec=DualCache),
data=request_data,
call_type="completion",
)

assert _byte_items_sent(mock_post) == []
assert _text_payloads_sent(mock_post) == ["summarize this"]


@pytest.mark.asyncio
async def test_pre_call_gs_uri_reference_passthrough_when_skip_unscannable_enabled():
"""A gs:// document reference passes through when skip_unscannable_attachments is enabled."""
guardrail = _make_guardrail(skip_unscannable_attachments=True)
request_data = {
"model": "gpt-4",
"messages": [
{
"role": "user",
"content": [
{
"type": "file",
"file": {"file_data": "gs://my-bucket/report.pdf", "filename": "report.pdf"},
}
],
}
],
"metadata": {"guardrails": ["model-armor-test"]},
}

with patch.object(
guardrail.async_handler,
"post",
AsyncMock(return_value=_armor_response(blocked=False)),
) as mock_post:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=MagicMock(spec=DualCache),
data=request_data,
call_type="completion",
)

assert _byte_items_sent(mock_post) == []


def test_initialize_guardrail_forwards_skip_unscannable_attachments():
"""skip_unscannable_attachments configured in litellm_params reaches the guardrail instance."""
from litellm.proxy.guardrails.guardrail_hooks.model_armor import initialize_guardrail
from litellm.types.guardrails import Guardrail, LitellmParams

litellm_params = LitellmParams(
guardrail="model_armor",
mode="pre_call",
template_id="demo-template",
project_id="demo-project",
skip_unscannable_attachments=True,
)
guardrail = initialize_guardrail(
litellm_params=litellm_params,
guardrail=Guardrail(guardrail_name="model-armor-config-test"),
)

guardrail = _make_guardrail()
assert guardrail.optional_params.get("skip_unscannable_attachments") is True


def test_initialize_guardrail_skip_unscannable_defaults_false():
"""A config that omits skip_unscannable_attachments keeps the secure default (block)."""
from litellm.proxy.guardrails.guardrail_hooks.model_armor import initialize_guardrail
from litellm.types.guardrails import Guardrail, LitellmParams

litellm_params = LitellmParams(
guardrail="model_armor",
mode="pre_call",
template_id="demo-template",
project_id="demo-project",
)
guardrail = initialize_guardrail(
litellm_params=litellm_params,
guardrail=Guardrail(guardrail_name="model-armor-config-default"),
)

assert guardrail.optional_params.get("skip_unscannable_attachments") is False


@pytest.mark.asyncio
async def test_skip_unscannable_still_fails_closed_on_api_error():
"""skip_unscannable_attachments only affects references; a real API error still fails closed."""
guardrail = _make_guardrail(skip_unscannable_attachments=True, fail_on_error=True)
pdf_b64 = base64.b64encode(PDF_BYTES).decode("utf-8")
block = {
"type": "file",
"file": {"file_data": f"data:application/pdf;base64,{pdf_b64}"},
}
request_data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": [block] * (MAX_FILE_ATTACHMENTS_PER_REQUEST + 1)}],
"messages": [_file_message(pdf_b64)],
"metadata": {"guardrails": ["model-armor-test"]},
}

with patch.object(
guardrail.async_handler,
"post",
AsyncMock(return_value=_armor_response(blocked=False)),
AsyncMock(side_effect=Exception("model armor upstream 500")),
):
with pytest.raises(HTTPException) as exc_info:
with pytest.raises(Exception) as exc_info:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=MagicMock(spec=DualCache),
data=request_data,
call_type="completion",
)

assert exc_info.value.status_code == 400
assert "per-request scan limit" in str(exc_info.value.detail)
assert "model armor upstream 500" in str(exc_info.value)


@pytest.mark.asyncio
async def test_pre_call_scans_every_attachment_without_a_count_cap():
"""There is no per-request attachment cap: every scannable attachment is submitted to Model Armor."""
guardrail = _make_guardrail()
pdf_b64 = base64.b64encode(PDF_BYTES).decode("utf-8")
block = {
"type": "file",
"file": {"file_data": f"data:application/pdf;base64,{pdf_b64}"},
}
count = 25
request_data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": [block] * count}],
"metadata": {"guardrails": ["model-armor-test"]},
}

mock_post = AsyncMock(return_value=_armor_response(blocked=False))
with patch.object(guardrail.async_handler, "post", mock_post):
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=MagicMock(spec=DualCache),
data=request_data,
call_type="completion",
)

assert len(_byte_items_sent(mock_post)) == count


@pytest.mark.asyncio
Expand Down
Loading