diff --git a/litellm/llms/bedrock/guardrail_attachments.py b/litellm/llms/bedrock/guardrail_attachments.py index ef0331e862b..eb78d98a197 100644 --- a/litellm/llms/bedrock/guardrail_attachments.py +++ b/litellm/llms/bedrock/guardrail_attachments.py @@ -56,8 +56,11 @@ _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 +_URL_SAFE_TO_STANDARD_BASE64: Final = str.maketrans("-_", "+/") _CONVERSE_ACTIONS: Final = frozenset({"converse", "converse-stream"}) _TOOL_ROLES: Final = frozenset({"tool", "function"}) _TOOL_OUTPUT_ITEM_TYPES: Final = frozenset({"function_call_output", "custom_tool_call_output", "computer_call_output"}) @@ -171,7 +174,7 @@ def _classify_nothing(block: Mapping[str, object]) -> _Classified: def _classify_openai_block(block: Mapping[str, object]) -> _Classified: block_type: Final = block.get("type") - if block_type == "image": + if block_type in ("image", "document"): return _classify_anthropic_block(block) if isinstance(block_type, str) and block_type in _OPENAI_IMAGE_TYPES: image_url: Final = block.get("image_url") @@ -189,11 +192,19 @@ def _classify_anthropic_block(block: Mapping[str, object]) -> _Classified: if _is_mapping(source) and source.get("type") == "base64": return _classify_base64(source.get("media_type"), source.get("data"), "image") return _Unscannable("image (url or file source)") + if block_type == "document" and _is_text_document(block): + return None if isinstance(block_type, str) and block_type in _ANTHROPIC_UNSCANNABLE_TYPES: return _Unscannable(block_type) return None +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 + + def _classify_converse_block(block: Mapping[str, object]) -> _Classified: image: Final = block.get("image") if _is_mapping(image): @@ -224,19 +235,29 @@ def _classify_base64(mime: object, encoded: object, label: str) -> _Classified: image_format: Final = _IMAGE_FORMAT_BY_MIME.get(mime.lower()) if isinstance(mime, str) else None if image_format is None: return _Unscannable(f"{label} ({mime[:_MAX_LABEL_MIME_CHARS]})" if isinstance(mime, str) and mime else label) - compact: Final = "".join(encoded.split()) if isinstance(encoded, str) else "" - decoded_size: Final = _decoded_size(compact) + standard: Final = _standard_base64(encoded) + if len(standard) > _MAX_IMAGE_BASE64_CHARS: + return _Unscannable(f"{label} (over 4 MB)") + decoded_size: Final = _decoded_size(standard) if decoded_size is None: return _Unscannable(f"{label} (invalid base64)") if decoded_size > _MAX_IMAGE_BYTES: return _Unscannable(f"{label} (over 4 MB)") - return _Image(BedrockContentItem(image=BedrockImageContent(format=image_format, source={"bytes": compact}))) + return _Image(BedrockContentItem(image=BedrockImageContent(format=image_format, source={"bytes": standard}))) -def _decoded_size(compact: str) -> int | None: - if not compact: +def _standard_base64(encoded: object) -> str: + """Return the payload as padded standard base64, accepting whitespace, missing padding and the URL-safe alphabet.""" + compact: Final = ( + "".join(encoded.split()).translate(_URL_SAFE_TO_STANDARD_BASE64) if isinstance(encoded, str) else "" + ) + return compact + "=" * (-len(compact) % 4) + + +def _decoded_size(standard: str) -> int | None: + if not standard: return None try: - return len(base64.b64decode(compact, validate=True)) + return len(base64.b64decode(standard, validate=True)) except (binascii.Error, ValueError): return None diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index ab1fb28714f..dbce5c8fa6c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -2567,9 +2567,12 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): 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`` skips this scan. + ``apply_guardrail`` or ``make_bedrock_api_request`` skips this scan. """ - if type(self).apply_guardrail is not BedrockGuardrail.apply_guardrail: + if ( + type(self).apply_guardrail is not BedrockGuardrail.apply_guardrail + or type(self).make_bedrock_api_request is not BedrockGuardrail.make_bedrock_api_request + ): return attachments: Final = find_request_attachments( data, 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 7cf6a1774aa..bf447093c73 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 @@ -6432,3 +6432,27 @@ async def test_attachment_scan_skipped_when_subclass_overrides_apply_guardrail() assert pdf_result is None assert image_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"} + + +@pytest.mark.asyncio +async def test_attachment_scan_skipped_when_subclass_overrides_make_bedrock_api_request(): + guardrail = _CustomRequestGuardrail( + 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 diff --git a/tests/unit/llms/bedrock/test_guardrail_attachments.py b/tests/unit/llms/bedrock/test_guardrail_attachments.py index 032d1b05273..7c0fb7784a8 100644 --- a/tests/unit/llms/bedrock/test_guardrail_attachments.py +++ b/tests/unit/llms/bedrock/test_guardrail_attachments.py @@ -397,3 +397,53 @@ def test_converse_null_document_is_not_an_attachment(): assert list(found.images) == [_png_item()] assert found.unscannable == ("audio",) + + +@pytest.mark.parametrize( + "block", + [ + pytest.param( + {"type": "document", "source": {"type": "text", "media_type": "text/plain", "data": "hi"}}, id="text" + ), + pytest.param({"type": "document", "source": {"type": "content", "content": [TEXT]}}, id="content"), + ], +) +@pytest.mark.parametrize( + "call_type", [CallTypes.anthropic_messages.value, CallTypes.acompletion.value], ids=["messages", "chat"] +) +def test_text_source_document_is_not_an_attachment(block, call_type): + pdf = {"type": "document", "source": {"type": "base64", "media_type": "application/pdf", "data": PDF_B64}} + + found = find_request_attachments(_chat(TEXT, block, pdf), call_type, False, False) + + assert found.images == () + assert found.unscannable == ("document",) + + +def test_unpadded_and_url_safe_base64_are_sent_as_standard_base64(): + raw = b"\x89PNG\r\n\x1a\n\xfb\xff\xfe-fake" + standard = base64.b64encode(raw).decode() + unpadded = standard.rstrip("=") + url_safe = base64.urlsafe_b64encode(raw).decode().rstrip("=") + data = _chat( + *( + {"type": "image_url", "image_url": {"url": f"data:image/png;base64,{encoded}"}} + for encoded in (unpadded, url_safe) + ) + ) + + found = find_request_attachments(data, CallTypes.acompletion.value, False, False) + + assert unpadded != standard + assert set("-_") & set(url_safe) + assert list(found.images) == [_png_item(standard), _png_item(standard)] + assert found.unscannable == () + + +def test_oversize_base64_is_refused_by_length_before_decoding(): + encoded = "!" * (len(OVERSIZE_PNG_B64) + 4) + data = _chat({"type": "image_url", "image_url": {"url": f"data:image/png;base64,{encoded}"}}) + + found = find_request_attachments(data, CallTypes.acompletion.value, False, False) + + assert found.unscannable == ("image_url (over 4 MB)",)