mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): let text documents, unpadded base64 images and request-override subclasses through as before
This commit is contained in:
parent
73431b0340
commit
4f902493d6
4 changed files with 107 additions and 9 deletions
|
|
@ -56,8 +56,11 @@ _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
|
||||
_URL_SAFE_TO_STANDARD_BASE64: Final = str.maketrans("-_", "+/")
|
||||
_CONVERSE_ACTIONS: Final = frozenset({"converse", "converse-stream"})
|
||||
_TOOL_ROLES: Final = frozenset({"tool", "function"})
|
||||
_TOOL_OUTPUT_ITEM_TYPES: Final = frozenset({"function_call_output", "custom_tool_call_output", "computer_call_output"})
|
||||
|
|
@ -171,7 +174,7 @@ def _classify_nothing(block: Mapping[str, object]) -> _Classified:
|
|||
|
||||
def _classify_openai_block(block: Mapping[str, object]) -> _Classified:
|
||||
block_type: Final = block.get("type")
|
||||
if block_type == "image":
|
||||
if block_type in ("image", "document"):
|
||||
return _classify_anthropic_block(block)
|
||||
if isinstance(block_type, str) and block_type in _OPENAI_IMAGE_TYPES:
|
||||
image_url: Final = block.get("image_url")
|
||||
|
|
@ -189,11 +192,19 @@ def _classify_anthropic_block(block: Mapping[str, object]) -> _Classified:
|
|||
if _is_mapping(source) and source.get("type") == "base64":
|
||||
return _classify_base64(source.get("media_type"), source.get("data"), "image")
|
||||
return _Unscannable("image (url or file source)")
|
||||
if block_type == "document" and _is_text_document(block):
|
||||
return None
|
||||
if isinstance(block_type, str) and block_type in _ANTHROPIC_UNSCANNABLE_TYPES:
|
||||
return _Unscannable(block_type)
|
||||
return None
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
def _classify_converse_block(block: Mapping[str, object]) -> _Classified:
|
||||
image: Final = block.get("image")
|
||||
if _is_mapping(image):
|
||||
|
|
@ -224,19 +235,29 @@ def _classify_base64(mime: object, encoded: object, label: str) -> _Classified:
|
|||
image_format: Final = _IMAGE_FORMAT_BY_MIME.get(mime.lower()) if isinstance(mime, str) else None
|
||||
if image_format is None:
|
||||
return _Unscannable(f"{label} ({mime[:_MAX_LABEL_MIME_CHARS]})" if isinstance(mime, str) and mime else label)
|
||||
compact: Final = "".join(encoded.split()) if isinstance(encoded, str) else ""
|
||||
decoded_size: Final = _decoded_size(compact)
|
||||
standard: Final = _standard_base64(encoded)
|
||||
if len(standard) > _MAX_IMAGE_BASE64_CHARS:
|
||||
return _Unscannable(f"{label} (over 4 MB)")
|
||||
decoded_size: Final = _decoded_size(standard)
|
||||
if decoded_size is None:
|
||||
return _Unscannable(f"{label} (invalid base64)")
|
||||
if decoded_size > _MAX_IMAGE_BYTES:
|
||||
return _Unscannable(f"{label} (over 4 MB)")
|
||||
return _Image(BedrockContentItem(image=BedrockImageContent(format=image_format, source={"bytes": compact})))
|
||||
return _Image(BedrockContentItem(image=BedrockImageContent(format=image_format, source={"bytes": standard})))
|
||||
|
||||
|
||||
def _decoded_size(compact: str) -> int | None:
|
||||
if not compact:
|
||||
def _standard_base64(encoded: object) -> str:
|
||||
"""Return the payload as padded standard base64, accepting whitespace, missing padding and the URL-safe alphabet."""
|
||||
compact: Final = (
|
||||
"".join(encoded.split()).translate(_URL_SAFE_TO_STANDARD_BASE64) if isinstance(encoded, str) else ""
|
||||
)
|
||||
return compact + "=" * (-len(compact) % 4)
|
||||
|
||||
|
||||
def _decoded_size(standard: str) -> int | None:
|
||||
if not standard:
|
||||
return None
|
||||
try:
|
||||
return len(base64.b64decode(compact, validate=True))
|
||||
return len(base64.b64decode(standard, validate=True))
|
||||
except (binascii.Error, ValueError):
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -2567,9 +2567,12 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
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`` skips this scan.
|
||||
``apply_guardrail`` or ``make_bedrock_api_request`` skips this scan.
|
||||
"""
|
||||
if type(self).apply_guardrail is not BedrockGuardrail.apply_guardrail:
|
||||
if (
|
||||
type(self).apply_guardrail is not BedrockGuardrail.apply_guardrail
|
||||
or type(self).make_bedrock_api_request is not BedrockGuardrail.make_bedrock_api_request
|
||||
):
|
||||
return
|
||||
attachments: Final = find_request_attachments(
|
||||
data,
|
||||
|
|
|
|||
|
|
@ -6432,3 +6432,27 @@ async def test_attachment_scan_skipped_when_subclass_overrides_apply_guardrail()
|
|||
assert pdf_result is None
|
||||
assert image_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"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_attachment_scan_skipped_when_subclass_overrides_make_bedrock_api_request():
|
||||
guardrail = _CustomRequestGuardrail(
|
||||
guardrail_name="bedrock-attachments", guardrailIdentifier="gid", guardrailVersion="DRAFT"
|
||||
)
|
||||
|
||||
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
|
||||
)
|
||||
image_result = await guardrail.async_scan_request_attachments(
|
||||
data=_image_only_chat_request(), call_type=CallTypes.acompletion.value
|
||||
)
|
||||
|
||||
assert pdf_result is None
|
||||
assert image_result is None
|
||||
assert mock_post.await_count == 0
|
||||
|
|
|
|||
|
|
@ -397,3 +397,53 @@ def test_converse_null_document_is_not_an_attachment():
|
|||
|
||||
assert list(found.images) == [_png_item()]
|
||||
assert found.unscannable == ("audio",)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"block",
|
||||
[
|
||||
pytest.param(
|
||||
{"type": "document", "source": {"type": "text", "media_type": "text/plain", "data": "hi"}}, id="text"
|
||||
),
|
||||
pytest.param({"type": "document", "source": {"type": "content", "content": [TEXT]}}, id="content"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"call_type", [CallTypes.anthropic_messages.value, CallTypes.acompletion.value], ids=["messages", "chat"]
|
||||
)
|
||||
def test_text_source_document_is_not_an_attachment(block, call_type):
|
||||
pdf = {"type": "document", "source": {"type": "base64", "media_type": "application/pdf", "data": PDF_B64}}
|
||||
|
||||
found = find_request_attachments(_chat(TEXT, block, pdf), call_type, False, False)
|
||||
|
||||
assert found.images == ()
|
||||
assert found.unscannable == ("document",)
|
||||
|
||||
|
||||
def test_unpadded_and_url_safe_base64_are_sent_as_standard_base64():
|
||||
raw = b"\x89PNG\r\n\x1a\n\xfb\xff\xfe-fake"
|
||||
standard = base64.b64encode(raw).decode()
|
||||
unpadded = standard.rstrip("=")
|
||||
url_safe = base64.urlsafe_b64encode(raw).decode().rstrip("=")
|
||||
data = _chat(
|
||||
*(
|
||||
{"type": "image_url", "image_url": {"url": f"data:image/png;base64,{encoded}"}}
|
||||
for encoded in (unpadded, url_safe)
|
||||
)
|
||||
)
|
||||
|
||||
found = find_request_attachments(data, CallTypes.acompletion.value, False, False)
|
||||
|
||||
assert unpadded != standard
|
||||
assert set("-_") & set(url_safe)
|
||||
assert list(found.images) == [_png_item(standard), _png_item(standard)]
|
||||
assert found.unscannable == ()
|
||||
|
||||
|
||||
def test_oversize_base64_is_refused_by_length_before_decoding():
|
||||
encoded = "!" * (len(OVERSIZE_PNG_B64) + 4)
|
||||
data = _chat({"type": "image_url", "image_url": {"url": f"data:image/png;base64,{encoded}"}})
|
||||
|
||||
found = find_request_attachments(data, CallTypes.acompletion.value, False, False)
|
||||
|
||||
assert found.unscannable == ("image_url (over 4 MB)",)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue