mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): scan images inside content documents instead of refusing them
This commit is contained in:
parent
c584e2e167
commit
561214afff
2 changed files with 40 additions and 46 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 == ()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue