fix(guardrails): let text documents, unpadded base64 images and request-override subclasses through as before

This commit is contained in:
Yucheng He 2026-09-29 20:25:14 -07:00
parent 73431b0340
commit 4f902493d6
4 changed files with 107 additions and 9 deletions

View file

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

View file

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

View file

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

View file

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