diff --git a/litellm/llms/bedrock/guardrail_attachments.py b/litellm/llms/bedrock/guardrail_attachments.py index fed7bbfa486..c63377d0ec0 100644 --- a/litellm/llms/bedrock/guardrail_attachments.py +++ b/litellm/llms/bedrock/guardrail_attachments.py @@ -9,7 +9,8 @@ request instead of letting the attachment reach the model unread. import base64 import binascii -from collections.abc import Callable, Mapping, Sequence +from collections.abc import Callable, Iterable, Mapping, Sequence +from itertools import chain from types import MappingProxyType from typing import Final, Literal, NamedTuple, TypeGuard @@ -56,6 +57,7 @@ _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 @@ -137,19 +139,31 @@ def _message_blocks(message: Mapping[str, object], nested_tool_blocks: _NestedTo 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) + return _with_document_contents(_Block(block, from_tool=True) for block in blocks) 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")) - for entry in ( - _Block(block, from_tool=from_tool_message), - *(_Block(inner, from_tool=True) for inner in nested_tool_blocks(block)), + return _with_document_contents( + chain.from_iterable( + ( + _Block(block, from_tool=from_tool_message), + *(_Block(inner, from_tool=True) for inner in nested_tool_blocks(block)), + ) + for block in _mappings(message.get("content")) ) ) +def _with_document_contents(entries: Iterable[_Block]) -> tuple[_Block, ...]: + return tuple(chain.from_iterable(_with_document_content(entry) for entry in entries)) + + +def _with_document_content(entry: _Block) -> tuple[_Block, ...]: + source: Final = entry.block.get("source") + if entry.block.get("type") != "document" or not _is_mapping(source) or source.get("type") != "content": + return (entry,) + return (entry, *(_Block(inner, from_tool=entry.from_tool) for inner in _mappings(source.get("content")))) + + def _in_scope(entry: _Block, skip_tool_messages: bool, scan_only_tool_results: bool) -> bool: return not skip_tool_messages if entry.from_tool else not scan_only_tool_results @@ -200,26 +214,8 @@ def _classify_anthropic_block(block: Mapping[str, object]) -> _Classified: def _is_text_document(block: Mapping[str, object]) -> bool: source: Final = block.get("source") - 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) + 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: @@ -232,8 +228,6 @@ 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/tests/unit/llms/bedrock/test_guardrail_attachments.py b/tests/unit/llms/bedrock/test_guardrail_attachments.py index 2a3eb430f80..9d6245a9714 100644 --- a/tests/unit/llms/bedrock/test_guardrail_attachments.py +++ b/tests/unit/llms/bedrock/test_guardrail_attachments.py @@ -449,27 +449,27 @@ def test_oversize_base64_is_refused_by_length_before_decoding(): assert found.unscannable == ("image_url (over 4 MB)",) -def test_content_source_document_with_an_image_is_refused(): +@pytest.mark.parametrize( + "call_type", [CallTypes.anthropic_messages.value, CallTypes.acompletion.value], ids=["messages", "chat"] +) +def test_images_inside_a_content_source_document_are_scanned(call_type): image = {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": PNG_B64}} - block = {"type": "document", "source": {"type": "content", "content": [TEXT, image]}} + pdf = {"type": "document", "source": {"type": "base64", "media_type": "application/pdf", "data": PDF_B64}} + block = {"type": "document", "source": {"type": "content", "content": [TEXT, image, pdf]}} - found = find_request_attachments(_chat(block), CallTypes.anthropic_messages.value, False, False) + found = find_request_attachments(_chat(block), call_type, False, False) + assert list(found.images) == [_png_item()] 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}}) +def test_document_images_inside_a_tool_result_follow_the_tool_scope(): + image = {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": PNG_B64}} + document = {"type": "document", "source": {"type": "content", "content": [image]}} + data = _chat({"type": "tool_result", "tool_use_id": "t1", "content": [document]}) - found = find_request_attachments(data, CallTypes.allm_passthrough_route.value, False, False) + scanned = find_request_attachments(data, CallTypes.anthropic_messages.value, False, False) + skipped = find_request_attachments(data, CallTypes.anthropic_messages.value, True, False) - assert found.unscannable == unscannable + assert list(scanned.images) == [_png_item()] + assert skipped.images == ()