fix(guardrails): keep base behavior for Bedrock subclasses and malformed messages
Some checks are pending
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run

- skip the attachment scan when a subclass overrides apply_guardrail
- ignore non-string message role/type instead of raising
- treat a null Converse document/video/audio block as absent
- keep the base request-body debug log for text-only calls
This commit is contained in:
Yucheng He 2026-09-29 10:54:02 -07:00
parent eeda4cb61d
commit 73431b0340
4 changed files with 94 additions and 5 deletions

View file

@ -131,11 +131,13 @@ def _latest_user_message(messages: Sequence[Mapping[str, object]]) -> tuple[Mapp
def _message_blocks(message: Mapping[str, object], nested_tool_blocks: _NestedToolBlocks) -> tuple[_Block, ...]:
if message.get("type") in _TOOL_OUTPUT_ITEM_TYPES:
message_type: Final = message.get("type")
if isinstance(message_type, str) and message_type in _TOOL_OUTPUT_ITEM_TYPES:
output: Final = message.get("output")
blocks: Final = (output,) if _is_mapping(output) else _mappings(output)
return tuple(_Block(block, from_tool=True) for block in blocks)
from_tool_message: Final = message.get("role") in _TOOL_ROLES
role: Final = message.get("role")
from_tool_message: Final = isinstance(role, str) and role in _TOOL_ROLES
return tuple(
entry
for block in _mappings(message.get("content"))
@ -203,7 +205,7 @@ def _classify_converse_block(block: Mapping[str, object]) -> _Classified:
mime: Final = f"image/{image_format}" if isinstance(image_format, str) else None
return _classify_base64(mime, encoded, "image")
for key in _CONVERSE_UNSCANNABLE_KEYS:
if key in block:
if block.get(key) is not None:
return _Unscannable(key)
return None

View file

@ -1279,7 +1279,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
headers_dict: Final = dict(prepared_request.headers) # mutable-ok: the masking helper requires a dict
verbose_proxy_logger.debug(
"Bedrock AI request body: %s, url %s, headers: %s",
{**bedrock_request_data, "content": _without_image_bytes(content)},
{**bedrock_request_data, "content": _without_image_bytes(content)}
if any("image" in item for item in content)
else bedrock_request_data,
prepared_request.url,
_get_masked_values(headers_dict),
)
@ -2564,8 +2566,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
Documents, files, audio, video, and images sent by URL, by file id or in another format block the
request unless ``skip_unscannable_attachments`` is set. ``checks`` mode calls the text-only
InvokeGuardrailChecks API, so there every image counts as unscannable too. A failed ApplyGuardrail
call raises the same error the text scan of the same hook raises.
call raises the same error the text scan of the same hook raises. A subclass that overrides
``apply_guardrail`` skips this scan.
"""
if type(self).apply_guardrail is not BedrockGuardrail.apply_guardrail:
return
attachments: Final = find_request_attachments(
data,
call_type,

View file

@ -4,6 +4,7 @@ Unit tests for Bedrock Guardrails
import json
import asyncio
import logging
from datetime import datetime, timezone
import sys
from unittest.mock import AsyncMock, MagicMock, patch
@ -14,6 +15,7 @@ from fastapi import HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.exceptions import ModifyResponseException
from litellm.proxy._types import UserAPIKeyAuth
@ -6374,3 +6376,59 @@ async def test_attachment_scan_debug_log_omits_image_bytes():
logged = " ".join(str(call.args) for call in mock_debug.call_args_list)
assert _ATTACHMENT_PNG_B64 not in logged
assert f"<{len(_ATTACHMENT_PNG_B64)} base64 chars>" in logged
@pytest.mark.asyncio
async def test_text_only_debug_log_prints_the_signed_request_body():
guardrail = _attachment_guardrail()
post_patch, credentials_patch, prepare_patch = _patched_bedrock_post(
guardrail, _passing_bedrock_httpx_response("ok")
)
captured_records: list[logging.LogRecord] = []
class _RecordingHandler(logging.Handler):
def emit(self, record: logging.LogRecord) -> None:
captured_records.append(record)
handler = _RecordingHandler(level=logging.DEBUG)
previous_level = verbose_proxy_logger.level
verbose_proxy_logger.addHandler(handler)
verbose_proxy_logger.setLevel(logging.DEBUG)
try:
with post_patch, credentials_patch, prepare_patch:
await guardrail.make_bedrock_api_request(
source="INPUT", messages=[{"role": "user", "content": "hello"}], request_data={}
)
finally:
verbose_proxy_logger.removeHandler(handler)
verbose_proxy_logger.setLevel(previous_level)
body_lines = [
record.getMessage() for record in captured_records if record.getMessage().startswith("Bedrock AI request body")
]
assert len(body_lines) == 1
assert "'content': ({'text': {'text': 'hello'}},)" in body_lines[0]
class _CustomApplyGuardrail(BedrockGuardrail):
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
return inputs
@pytest.mark.asyncio
async def test_attachment_scan_skipped_when_subclass_overrides_apply_guardrail():
guardrail = _CustomApplyGuardrail(
guardrail_name="bedrock-attachments", guardrailIdentifier="gid", guardrailVersion="DRAFT"
)
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
pdf_result = await guardrail.async_scan_request_attachments(
data=_pdf_chat_request(), call_type=CallTypes.acompletion.value
)
image_result = await guardrail.async_scan_request_attachments(
data=_image_only_chat_request(), call_type=CallTypes.acompletion.value
)
assert pdf_result is None
assert image_result is None
assert mock_post.await_count == 0

View file

@ -373,3 +373,27 @@ def test_latest_user_message_only():
assert list(found.images) == [_png_item()]
assert found.unscannable == ()
def test_non_string_role_and_type_do_not_raise():
data = {
"messages": [
{"role": "user", "content": [{"type": "image_url", "image_url": f"data:image/png;base64,{PNG_B64}"}]},
{"role": ["tool"], "type": {"x": 1}, "content": [{"type": "file", "file": {"file_id": "f"}}]},
{"role": {"r": "tool"}, "type": ["function_call_output"], "content": [{"type": "file", "file": {}}]},
]
}
found = find_request_attachments(data, CallTypes.acompletion.value, True, False)
assert list(found.images) == [_png_item()]
assert found.unscannable == ("file", "file")
def test_converse_null_document_is_not_an_attachment():
data = _converse(_png_item(), {"document": None}, {"video": None, "audio": None}, {"audio": {"format": "mp3"}})
found = find_request_attachments(data, CallTypes.allm_passthrough_route.value, False, False)
assert list(found.images) == [_png_item()]
assert found.unscannable == ("audio",)