From c584e2e1675a114f77fb273666796c734027eee8 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Wed, 30 Sep 2026 01:16:16 -0700 Subject: [PATCH] fix(guardrails): let Converse text documents through and refuse documents that carry images --- litellm/llms/bedrock/guardrail_attachments.py | 25 +++++++++++++++--- .../guardrail_hooks/bedrock_guardrails.py | 9 ++++--- .../test_bedrock_guardrails.py | 16 ++++++++++++ .../bedrock/test_guardrail_attachments.py | 26 +++++++++++++++++++ 4 files changed, 69 insertions(+), 7 deletions(-) diff --git a/litellm/llms/bedrock/guardrail_attachments.py b/litellm/llms/bedrock/guardrail_attachments.py index eb78d98a197..fed7bbfa486 100644 --- a/litellm/llms/bedrock/guardrail_attachments.py +++ b/litellm/llms/bedrock/guardrail_attachments.py @@ -56,7 +56,6 @@ _OPENAI_UNSCANNABLE_TYPES: Final = frozenset( {"file", "input_file", "input_audio", "video_url", "audio_url", "document", "container_upload"} ) _ANTHROPIC_UNSCANNABLE_TYPES: Final = frozenset({"document", "container_upload"}) -_TEXT_DOCUMENT_SOURCE_TYPES: Final = frozenset({"text", "content"}) _CONVERSE_UNSCANNABLE_KEYS: Final = ("document", "video", "audio") _MAX_IMAGE_BYTES: Final = 4 * 1024 * 1024 _MAX_IMAGE_BASE64_CHARS: Final = -(-_MAX_IMAGE_BYTES // 3) * 4 @@ -201,8 +200,26 @@ def _classify_anthropic_block(block: Mapping[str, object]) -> _Classified: def _is_text_document(block: Mapping[str, object]) -> bool: source: Final = block.get("source") - source_type: Final = source.get("type") if _is_mapping(source) else None - return isinstance(source_type, str) and source_type in _TEXT_DOCUMENT_SOURCE_TYPES + if not _is_mapping(source): + return False + content: Final = source.get("content") + return source.get("type") == "text" or ( + source.get("type") == "content" + and (isinstance(content, str) or _all_blocks(content, lambda inner: inner.get("type") == "text")) + ) + + +def _is_converse_text_document(document: object) -> bool: + source: Final = document.get("source") if _is_mapping(document) else None + if not _is_mapping(source): + return False + return isinstance(source.get("text"), str) or _all_blocks( + source.get("content"), lambda inner: set(inner) == {"text"} + ) + + +def _all_blocks(value: object, predicate: Callable[[Mapping[str, object]], bool]) -> bool: + return _is_list(value) and all(_is_mapping(inner) and predicate(inner) for inner in value) def _classify_converse_block(block: Mapping[str, object]) -> _Classified: @@ -215,6 +232,8 @@ def _classify_converse_block(block: Mapping[str, object]) -> _Classified: return _Unscannable("image (no inline bytes)") mime: Final = f"image/{image_format}" if isinstance(image_format, str) else None return _classify_base64(mime, encoded, "image") + if _is_converse_text_document(block.get("document")): + return None for key in _CONVERSE_UNSCANNABLE_KEYS: if block.get(key) is not None: return _Unscannable(key) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index dbce5c8fa6c..30e2b60000d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -2566,12 +2566,13 @@ 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. A subclass that overrides - ``apply_guardrail`` or ``make_bedrock_api_request`` skips this scan. + call raises the same error the text scan of the same hook raises. A guardrail whose + ``apply_guardrail`` or ``make_bedrock_api_request`` is replaced, by a subclass or on the + instance, skips this scan. """ if ( - type(self).apply_guardrail is not BedrockGuardrail.apply_guardrail - or type(self).make_bedrock_api_request is not BedrockGuardrail.make_bedrock_api_request + getattr(self.apply_guardrail, "__func__", None) is not BedrockGuardrail.apply_guardrail + or getattr(self.make_bedrock_api_request, "__func__", None) is not BedrockGuardrail.make_bedrock_api_request ): return attachments: Final = find_request_attachments( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index bf447093c73..3f8e85dcaff 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -6434,6 +6434,22 @@ async def test_attachment_scan_skipped_when_subclass_overrides_apply_guardrail() assert mock_post.await_count == 0 +@pytest.mark.asyncio +async def test_attachment_scan_skipped_when_instance_replaces_make_bedrock_api_request(): + guardrail = BedrockGuardrail( + guardrail_name="bedrock-attachments", guardrailIdentifier="gid", guardrailVersion="DRAFT" + ) + guardrail.make_bedrock_api_request = AsyncMock(return_value={"action": "NONE"}) + + 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 + ) + + assert pdf_result is None + assert mock_post.await_count == 0 + + class _CustomRequestGuardrail(BedrockGuardrail): async def make_bedrock_api_request(self, source, messages=None, response=None, request_data=None, **kwargs): return {"action": "NONE"} diff --git a/tests/unit/llms/bedrock/test_guardrail_attachments.py b/tests/unit/llms/bedrock/test_guardrail_attachments.py index 7c0fb7784a8..2a3eb430f80 100644 --- a/tests/unit/llms/bedrock/test_guardrail_attachments.py +++ b/tests/unit/llms/bedrock/test_guardrail_attachments.py @@ -447,3 +447,29 @@ def test_oversize_base64_is_refused_by_length_before_decoding(): found = find_request_attachments(data, CallTypes.acompletion.value, False, False) assert found.unscannable == ("image_url (over 4 MB)",) + + +def test_content_source_document_with_an_image_is_refused(): + image = {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": PNG_B64}} + block = {"type": "document", "source": {"type": "content", "content": [TEXT, image]}} + + found = find_request_attachments(_chat(block), CallTypes.anthropic_messages.value, False, False) + + assert found.unscannable == ("document",) + + +@pytest.mark.parametrize( + "source, unscannable", + [ + pytest.param({"text": "The grass is purple."}, (), id="text"), + pytest.param({"content": [{"text": "a"}, {"text": "b"}]}, (), id="content"), + pytest.param({"content": [{"text": "a"}, {"image": {}}]}, ("document",), id="content-with-image"), + pytest.param({"bytes": PDF_B64}, ("document",), id="bytes"), + ], +) +def test_converse_text_source_document_is_not_an_attachment(source, unscannable): + data = _converse({"document": {"format": "txt", "name": "doc", "source": source}}) + + found = find_request_attachments(data, CallTypes.allm_passthrough_route.value, False, False) + + assert found.unscannable == unscannable