mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): let Converse text documents through and refuse documents that carry images
This commit is contained in:
parent
4f902493d6
commit
c584e2e167
4 changed files with 69 additions and 7 deletions
|
|
@ -56,7 +56,6 @@ _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
|
||||
|
|
@ -201,8 +200,26 @@ def _classify_anthropic_block(block: Mapping[str, object]) -> _Classified:
|
|||
|
||||
def _is_text_document(block: Mapping[str, object]) -> bool:
|
||||
source: Final = block.get("source")
|
||||
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
|
||||
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)
|
||||
|
||||
|
||||
def _classify_converse_block(block: Mapping[str, object]) -> _Classified:
|
||||
|
|
@ -215,6 +232,8 @@ 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)
|
||||
|
|
|
|||
|
|
@ -2566,12 +2566,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
Documents, files, audio, video, and images sent by URL, by file id or in another format block the
|
||||
request unless ``skip_unscannable_attachments`` is set. ``checks`` mode calls the text-only
|
||||
InvokeGuardrailChecks API, so there every image counts as unscannable too. A failed ApplyGuardrail
|
||||
call raises the same error the text scan of the same hook raises. A subclass that overrides
|
||||
``apply_guardrail`` or ``make_bedrock_api_request`` skips this scan.
|
||||
call raises the same error the text scan of the same hook raises. A guardrail whose
|
||||
``apply_guardrail`` or ``make_bedrock_api_request`` is replaced, by a subclass or on the
|
||||
instance, skips this scan.
|
||||
"""
|
||||
if (
|
||||
type(self).apply_guardrail is not BedrockGuardrail.apply_guardrail
|
||||
or type(self).make_bedrock_api_request is not BedrockGuardrail.make_bedrock_api_request
|
||||
getattr(self.apply_guardrail, "__func__", None) is not BedrockGuardrail.apply_guardrail
|
||||
or getattr(self.make_bedrock_api_request, "__func__", None) is not BedrockGuardrail.make_bedrock_api_request
|
||||
):
|
||||
return
|
||||
attachments: Final = find_request_attachments(
|
||||
|
|
|
|||
|
|
@ -6434,6 +6434,22 @@ async def test_attachment_scan_skipped_when_subclass_overrides_apply_guardrail()
|
|||
assert mock_post.await_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_attachment_scan_skipped_when_instance_replaces_make_bedrock_api_request():
|
||||
guardrail = BedrockGuardrail(
|
||||
guardrail_name="bedrock-attachments", guardrailIdentifier="gid", guardrailVersion="DRAFT"
|
||||
)
|
||||
guardrail.make_bedrock_api_request = AsyncMock(return_value={"action": "NONE"})
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
|
||||
pdf_result = await guardrail.async_scan_request_attachments(
|
||||
data=_pdf_chat_request(), call_type=CallTypes.acompletion.value
|
||||
)
|
||||
|
||||
assert pdf_result is None
|
||||
assert mock_post.await_count == 0
|
||||
|
||||
|
||||
class _CustomRequestGuardrail(BedrockGuardrail):
|
||||
async def make_bedrock_api_request(self, source, messages=None, response=None, request_data=None, **kwargs):
|
||||
return {"action": "NONE"}
|
||||
|
|
|
|||
|
|
@ -447,3 +447,29 @@ def test_oversize_base64_is_refused_by_length_before_decoding():
|
|||
found = find_request_attachments(data, CallTypes.acompletion.value, False, False)
|
||||
|
||||
assert found.unscannable == ("image_url (over 4 MB)",)
|
||||
|
||||
|
||||
def test_content_source_document_with_an_image_is_refused():
|
||||
image = {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": PNG_B64}}
|
||||
block = {"type": "document", "source": {"type": "content", "content": [TEXT, image]}}
|
||||
|
||||
found = find_request_attachments(_chat(block), CallTypes.anthropic_messages.value, False, False)
|
||||
|
||||
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}})
|
||||
|
||||
found = find_request_attachments(data, CallTypes.allm_passthrough_route.value, False, False)
|
||||
|
||||
assert found.unscannable == unscannable
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue