fix(guardrails): scan images inside content documents instead of refusing them

This commit is contained in:
Yucheng He 2026-09-30 01:26:28 -07:00
parent c584e2e167
commit 561214afff
2 changed files with 40 additions and 46 deletions

View file

@ -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)

View file

@ -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 == ()