diff --git a/litellm/llms/bedrock/guardrail_attachments.py b/litellm/llms/bedrock/guardrail_attachments.py index a66f1d815c0..b28e524ca4d 100644 --- a/litellm/llms/bedrock/guardrail_attachments.py +++ b/litellm/llms/bedrock/guardrail_attachments.py @@ -44,7 +44,7 @@ class _DocumentText(NamedTuple): class _Block(NamedTuple): block: Mapping[str, object] from_tool: bool - in_document: bool = False + document_depth: int = 0 _Classified = _Image | _Unscannable | _DocumentText | None @@ -64,6 +64,7 @@ _OPENAI_UNSCANNABLE_TYPES: Final = frozenset( ) _ANTHROPIC_UNSCANNABLE_TYPES: Final = frozenset({"document", "container_upload"}) _TEXT_DOCUMENT_SOURCE_TYPES: Final = frozenset({"text", "content"}) +_MAX_DOCUMENT_DEPTH: Final = 3 _CONVERSE_UNSCANNABLE_KEYS: Final = ("document", "video", "audio") _MAX_IMAGE_BYTES: Final = 4 * 1024 * 1024 _MAX_IMAGE_BASE64_CHARS: Final = -(-_MAX_IMAGE_BYTES // 3) * 4 @@ -85,11 +86,12 @@ def find_request_attachments( messages, classify, nested_tool_blocks = _messages_and_classifier(data, call_type) selected: Final = messages if not latest_user_message_only else _latest_user_message(messages) classified: Final = tuple( - result - for message in selected - for entry in _message_blocks(message, nested_tool_blocks) - if _in_scope(entry, skip_tool_messages, scan_only_tool_results) - and (result := _classify_entry(entry, classify)) is not None + chain.from_iterable( + _classify_entry(entry, classify) + for message in selected + for entry in _message_blocks(message, nested_tool_blocks) + if _in_scope(entry, skip_tool_messages, scan_only_tool_results) + ) ) return RequestAttachments( images=tuple(result.item for result in classified if isinstance(result, _Image)), @@ -98,21 +100,32 @@ def find_request_attachments( ) -def _classify_entry(entry: _Block, classify: _BlockClassifier) -> _Classified: +def _classify_entry(entry: _Block, classify: _BlockClassifier) -> tuple[_Image | _Unscannable | _DocumentText, ...]: + if entry.document_depth >= _MAX_DOCUMENT_DEPTH and _content_document_source(entry.block) is not None: + return (_Unscannable("document (nested too deep)"),) text: Final = _document_text(entry) - return _DocumentText(text) if text else classify(entry.block) + result: Final = classify(entry.block) + return tuple(item for item in (_DocumentText(text) if text else None, result) if item is not None) -def _document_text(entry: _Block) -> str | None: +def _document_text(entry: _Block) -> str: + """Return the text a block carries as part of a document: a nested text block, or a document's title, context and text source.""" block: Final = entry.block - if entry.in_document: - text: Final = block.get("text") if block.get("type") == "text" else None - return text if isinstance(text, str) else None - source: Final = block.get("source") if block.get("type") == "document" else None - if not _is_mapping(source): - return None - source_text: Final = source.get("data") if source.get("type") == "text" else source.get("content") - return source_text if source.get("type") in _TEXT_DOCUMENT_SOURCE_TYPES and isinstance(source_text, str) else None + if entry.document_depth > 0 and block.get("type") == "text": + text: Final = block.get("text") + return text if isinstance(text, str) else "" + if block.get("type") != "document": + return "" + source: Final = block.get("source") + source_type: Final = source.get("type") if _is_mapping(source) else None + source_text: Final = ( + (source.get("data") if source_type == "text" else source.get("content")) + if _is_mapping(source) and isinstance(source_type, str) and source_type in _TEXT_DOCUMENT_SOURCE_TYPES + else None + ) + return "\n".join( + part for part in (block.get("title"), block.get("context"), source_text) if isinstance(part, str) and part + ) def _messages_and_classifier( @@ -182,15 +195,25 @@ def _with_document_contents(entries: Iterable[_Block]) -> tuple[_Block, ...]: 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": + source: Final = _content_document_source(entry.block) + if source is None or entry.document_depth >= _MAX_DOCUMENT_DEPTH: return (entry,) return ( entry, - *(_Block(inner, from_tool=entry.from_tool, in_document=True) for inner in _mappings(source.get("content"))), + *chain.from_iterable( + _with_document_content(_Block(inner, from_tool=entry.from_tool, document_depth=entry.document_depth + 1)) + for inner in _mappings(source.get("content")) + ), ) +def _content_document_source(block: Mapping[str, object]) -> Mapping[str, object] | None: + source: Final = block.get("source") + if block.get("type") != "document" or not _is_mapping(source) or source.get("type") != "content": + return None + return source + + 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 diff --git a/tests/unit/llms/bedrock/test_guardrail_attachments.py b/tests/unit/llms/bedrock/test_guardrail_attachments.py index c2657ce7077..9e571d4a5d9 100644 --- a/tests/unit/llms/bedrock/test_guardrail_attachments.py +++ b/tests/unit/llms/bedrock/test_guardrail_attachments.py @@ -479,3 +479,37 @@ def test_document_images_inside_a_tool_result_follow_the_tool_scope(): assert list(scanned.images) == [_png_item()] assert skipped.images == () + + +def test_document_title_and_context_are_scanned_as_text(): + text_doc = { + "type": "document", + "title": "record", + "context": "SSN 123-45-6789", + "source": {"type": "text", "media_type": "text/plain", "data": "body"}, + } + pdf = { + "type": "document", + "context": "note", + "source": {"type": "base64", "media_type": "application/pdf", "data": PDF_B64}, + } + + found = find_request_attachments(_chat(text_doc, pdf), CallTypes.anthropic_messages.value, False, False) + + assert found.document_texts == ("record\nSSN 123-45-6789\nbody", "note") + assert found.unscannable == ("document",) + + +def _content_document(*inner: object) -> dict: + return {"type": "document", "source": {"type": "content", "content": list(inner)}} + + +def test_nested_documents_are_scanned_and_deep_nesting_is_refused(): + inner_text = {"type": "document", "source": {"type": "text", "media_type": "text/plain", "data": "inner"}} + nested = _content_document(inner_text, _content_document({"type": "text", "text": "deeper"})) + too_deep = _content_document(_content_document(_content_document(_content_document(TEXT)))) + + found = find_request_attachments(_chat(nested, too_deep), CallTypes.anthropic_messages.value, False, False) + + assert found.document_texts == ("inner", "deeper") + assert found.unscannable == ("document (nested too deep)",)