mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): scan document title, context and nested documents
This commit is contained in:
parent
69fb98540f
commit
003e2dede3
2 changed files with 77 additions and 20 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)",)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue