fix(guardrails): scan document title, context and nested documents

This commit is contained in:
Yucheng He 2026-09-30 13:40:39 -07:00
parent 69fb98540f
commit 003e2dede3
2 changed files with 77 additions and 20 deletions

View file

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

View file

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