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, ...]: 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") output: Final = message.get("output")
blocks: Final = (output,) if _is_mapping(output) else _mappings(output) blocks: Final = (output,) if _is_mapping(output) else _mappings(output)
return tuple(_Block(block, from_tool=True) for block in blocks) 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( return tuple(
entry entry
for block in _mappings(message.get("content")) 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 mime: Final = f"image/{image_format}" if isinstance(image_format, str) else None
return _classify_base64(mime, encoded, "image") return _classify_base64(mime, encoded, "image")
for key in _CONVERSE_UNSCANNABLE_KEYS: for key in _CONVERSE_UNSCANNABLE_KEYS:
if key in block: if block.get(key) is not None:
return _Unscannable(key) return _Unscannable(key)
return None 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 headers_dict: Final = dict(prepared_request.headers) # mutable-ok: the masking helper requires a dict
verbose_proxy_logger.debug( verbose_proxy_logger.debug(
"Bedrock AI request body: %s, url %s, headers: %s", "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, prepared_request.url,
_get_masked_values(headers_dict), _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 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 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 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( attachments: Final = find_request_attachments(
data, data,
call_type, call_type,

View file

@ -4,6 +4,7 @@ Unit tests for Bedrock Guardrails
import json import json
import asyncio import asyncio
import logging
from datetime import datetime, timezone from datetime import datetime, timezone
import sys import sys
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
@ -14,6 +15,7 @@ from fastapi import HTTPException
import litellm import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache from litellm.caching.caching import DualCache
from litellm.exceptions import ModifyResponseException from litellm.exceptions import ModifyResponseException
from litellm.proxy._types import UserAPIKeyAuth 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) logged = " ".join(str(call.args) for call in mock_debug.call_args_list)
assert _ATTACHMENT_PNG_B64 not in logged assert _ATTACHMENT_PNG_B64 not in logged
assert f"<{len(_ATTACHMENT_PNG_B64)} base64 chars>" 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 list(found.images) == [_png_item()]
assert found.unscannable == () 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",)