fix(guardrails): scan or refuse attachments in the bedrock guardrail

This commit is contained in:
Yucheng He 2026-09-26 16:24:06 -07:00
parent 69d2a3c24f
commit 567a1f78e6
8 changed files with 1004 additions and 17 deletions

View file

@ -0,0 +1,235 @@
"""
Find the attachments in a raw request that the Bedrock guardrail has to scan or refuse.
ApplyGuardrail scans inline PNG and JPEG images of up to 4 MB. Every other attachment
(documents, files, audio, video, and images sent as a remote URL, a file id, in another
format or over the size limit) is reported as unscannable so the guardrail can block the
request instead of letting the attachment reach the model unread.
"""
import base64
import binascii
from collections.abc import Callable, Mapping, Sequence
from types import MappingProxyType
from typing import Final, Literal, NamedTuple
from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockContentItem,
BedrockImageContent,
)
from litellm.types.utils import CallTypes
BedrockImageFormat = Literal["png", "jpeg"]
class RequestAttachments(NamedTuple):
images: tuple[BedrockContentItem, ...]
unscannable: tuple[str, ...]
class _Image(NamedTuple):
item: BedrockContentItem
class _Unscannable(NamedTuple):
label: str
class _Block(NamedTuple):
block: Mapping[str, object]
from_tool: bool
_Classified = _Image | _Unscannable | None
_BlockClassifier = Callable[[Mapping[str, object]], _Classified]
_NestedToolBlocks = Callable[[Mapping[str, object]], tuple[Mapping[str, object], ...]]
NO_ATTACHMENTS: Final = RequestAttachments(images=(), unscannable=())
_IMAGE_FORMAT_BY_MIME: Final[Mapping[str, BedrockImageFormat]] = MappingProxyType(
{"image/png": "png", "image/jpeg": "jpeg", "image/jpg": "jpeg"}
)
_CHAT_CALL_TYPES: Final = frozenset({CallTypes.completion.value, CallTypes.acompletion.value})
_RESPONSES_CALL_TYPES: Final = frozenset({CallTypes.responses.value, CallTypes.aresponses.value})
_CHAT_UNSCANNABLE_TYPES: Final = frozenset({"file", "input_audio", "video_url", "audio_url", "document"})
_ANTHROPIC_UNSCANNABLE_TYPES: Final = frozenset({"document", "container_upload"})
_RESPONSES_UNSCANNABLE_TYPES: Final = frozenset({"input_file", "input_audio"})
_CONVERSE_UNSCANNABLE_KEYS: Final = ("document", "video")
_MAX_IMAGE_BYTES: Final = 4 * 1024 * 1024
_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"})
_MAX_LABEL_MIME_CHARS: Final = 40
def find_request_attachments(
data: Mapping[str, object],
call_type: str,
skip_tool_messages: bool,
latest_user_message_only: bool,
scan_only_tool_results: bool = False,
) -> RequestAttachments:
"""List the scannable images and the unscannable attachments in the message content and tool results."""
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.block)) is not None
)
return RequestAttachments(
images=tuple(result.item for result in classified if isinstance(result, _Image)),
unscannable=tuple(result.label for result in classified if isinstance(result, _Unscannable)),
)
def _messages_and_classifier(
data: Mapping[str, object], call_type: str
) -> tuple[Sequence[Mapping[str, object]], _BlockClassifier, _NestedToolBlocks]:
if call_type in _CHAT_CALL_TYPES:
return _mappings(data.get("messages")), _classify_chat_block, _no_nested_blocks
if call_type == CallTypes.anthropic_messages.value:
return _mappings(data.get("messages")), _classify_anthropic_block, _anthropic_tool_result_blocks
if call_type in _RESPONSES_CALL_TYPES:
return _mappings(data.get("input")), _classify_responses_block, _no_nested_blocks
if call_type == CallTypes.allm_passthrough_route.value and _is_bedrock_converse(data):
body: Final = data.get("data")
messages: Final = _mappings(body.get("messages") if isinstance(body, Mapping) else None)
return messages, _classify_converse_block, _converse_tool_result_blocks
return (), _classify_nothing, _no_nested_blocks
def _is_bedrock_converse(data: Mapping[str, object]) -> bool:
endpoint: Final = data.get("endpoint")
return (
data.get("custom_llm_provider") == "bedrock"
and isinstance(endpoint, str)
and endpoint.rstrip("/").rsplit("/", 1)[-1] in _CONVERSE_ACTIONS
)
def _mappings(value: object) -> tuple[Mapping[str, object], ...]:
if not isinstance(value, list):
return ()
return tuple(item for item in value if isinstance(item, Mapping))
def _latest_user_message(messages: Sequence[Mapping[str, object]]) -> tuple[Mapping[str, object], ...]:
return tuple(message for message in messages if message.get("role") == "user")[-1:]
def _message_blocks(message: Mapping[str, object], nested_tool_blocks: _NestedToolBlocks) -> tuple[_Block, ...]:
if message.get("type") in _TOOL_OUTPUT_ITEM_TYPES:
return tuple(_Block(block, from_tool=True) for block in _mappings(message.get("output")))
from_tool_message: Final = message.get("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)),
)
)
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
def _no_nested_blocks(block: Mapping[str, object]) -> tuple[Mapping[str, object], ...]:
return ()
def _anthropic_tool_result_blocks(block: Mapping[str, object]) -> tuple[Mapping[str, object], ...]:
return _mappings(block.get("content")) if block.get("type") == "tool_result" else ()
def _converse_tool_result_blocks(block: Mapping[str, object]) -> tuple[Mapping[str, object], ...]:
tool_result: Final = block.get("toolResult")
return _mappings(tool_result.get("content")) if isinstance(tool_result, Mapping) else ()
def _classify_nothing(block: Mapping[str, object]) -> _Classified:
return None
def _classify_chat_block(block: Mapping[str, object]) -> _Classified:
block_type: Final = block.get("type")
if block_type == "image_url":
image_url: Final = block.get("image_url")
url: Final = image_url.get("url") if isinstance(image_url, Mapping) else image_url
return _classify_data_uri(url, "image_url")
if isinstance(block_type, str) and block_type in _CHAT_UNSCANNABLE_TYPES:
return _Unscannable(block_type)
return None
def _classify_anthropic_block(block: Mapping[str, object]) -> _Classified:
block_type: Final = block.get("type")
if block_type == "image":
source: Final = block.get("source")
if isinstance(source, Mapping) and source.get("type") == "base64":
return _classify_base64(source.get("media_type"), source.get("data"), "image")
return _Unscannable("image (url or file source)")
if isinstance(block_type, str) and block_type in _ANTHROPIC_UNSCANNABLE_TYPES:
return _Unscannable(block_type)
return None
def _classify_responses_block(block: Mapping[str, object]) -> _Classified:
block_type: Final = block.get("type")
if block_type == "input_image":
return _classify_data_uri(block.get("image_url"), "input_image")
if isinstance(block_type, str) and block_type in _RESPONSES_UNSCANNABLE_TYPES:
return _Unscannable(block_type)
return None
def _classify_converse_block(block: Mapping[str, object]) -> _Classified:
image: Final = block.get("image")
if isinstance(image, Mapping):
image_format: Final = image.get("format")
source: Final = image.get("source")
encoded: Final = source.get("bytes") if isinstance(source, Mapping) else None
if encoded is None:
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")
for key in _CONVERSE_UNSCANNABLE_KEYS:
if key in block:
return _Unscannable(key)
return None
def _classify_data_uri(url: object, label: str) -> _Classified:
if not isinstance(url, str) or not url.startswith("data:"):
return _Unscannable(f"{label} (remote URL or file id)")
header, _, payload = url.partition(",")
params: Final = header.removeprefix("data:").lower().split(";")
if "base64" not in params[1:]:
return _Unscannable(f"{label} (not base64)")
return _classify_base64(params[0], payload, label)
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)
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})))
def _decoded_size(compact: str) -> int | None:
if not compact:
return None
try:
return len(base64.b64decode(compact, validate=True))
except (binascii.Error, ValueError):
return None

View file

@ -43,6 +43,7 @@ from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import (
from litellm.llms.anthropic.chat.guardrail_translation.handler import AnthropicMessagesHandler
from litellm.llms.base_llm.guardrail_translation.utils import (
effective_scan_only_tool_results_for_guardrail,
effective_skip_tool_message_for_guardrail,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, bedrock_bearer_token, run_aws_signing
from litellm.llms.custom_httpx.http_handler import (
@ -59,6 +60,7 @@ from litellm.proxy.guardrails.anthropic_sse import (
is_raw_sse_stream,
model_response_text,
)
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrail_attachments import find_request_attachments
from litellm.types.guardrails import (
BedrockChecksConfigModel,
BedrockGuardrailStreamingParams,
@ -119,6 +121,7 @@ _BEDROCK_INVOKE_GUARDRAIL_CHECKS_PATH: Final = "/guardrail-checks/invoke"
# more text blocks is split across multiple messages so ALL content is scanned --
# never truncated (truncation would let a user hide content past the limit).
_BEDROCK_CHECKS_MAX_CONTENT_BLOCKS: Final = 10
_BEDROCK_APPLY_GUARDRAIL_MAX_IMAGES: Final = 20
_BEDROCK_CHECKS_KNOWN_KEYS: Final = frozenset({"contentFilter", "promptAttack", "sensitiveInformation"})
# Keys in a sensitiveInformation result that pinpoint the PII location. They are
# stripped before the response is handed to standard logging / telemetry so the
@ -924,21 +927,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
guardrail call and is logged exactly once here.
"""
start_time: Final = datetime.now(timezone.utc)
bedrock_request_data: Final[dict] = dict(
self.convert_to_bedrock_format(source=source, messages=messages, response=response)
bedrock_request_data, api_key = self._bedrock_request_body(
self.convert_to_bedrock_format(source=source, messages=messages, response=response),
request_data,
)
api_key: str | None = None
if request_data:
dynamic_request_body_params = self.get_guardrail_dynamic_request_body_params(request_data=request_data)
bedrock_request_data.update(
{
key: value
for key, value in dynamic_request_body_params.items()
if key not in _BEDROCK_DYNAMIC_BODY_DENYLIST
}
)
if request_data.get("api_key") is not None:
api_key = request_data["api_key"]
event_type: Final = (
logging_event_type
@ -956,10 +948,51 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
source,
)
return BedrockGuardrailResponse()
return await self._apply_guardrail_to_content(
content=content,
bedrock_request_data=bedrock_request_data,
api_key=api_key,
request_data=request_data,
event_type=event_type,
start_time=start_time,
allow_chunking=not self._content_uses_contextual_grounding(content),
)
def _bedrock_request_body(
self,
base_request: BedrockRequest,
request_data: dict | None, # mutable-ok: proxy request body dict, read by the dynamic-params helper
) -> tuple[dict, str | None]:
"""Merge the request's dynamic ApplyGuardrail params into `base_request` and pick up its api_key."""
bedrock_request_data: Final[dict] = dict(base_request)
api_key: str | None = None
if request_data:
dynamic_request_body_params = self.get_guardrail_dynamic_request_body_params(request_data=request_data)
bedrock_request_data.update(
{
key: value
for key, value in dynamic_request_body_params.items()
if key not in _BEDROCK_DYNAMIC_BODY_DENYLIST
}
)
if request_data.get("api_key") is not None:
api_key = request_data["api_key"]
return bedrock_request_data, api_key
async def _apply_guardrail_to_content(
self,
content: Sequence[BedrockContentItem],
bedrock_request_data: Mapping[str, object],
api_key: str | None,
request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper
event_type: GuardrailEventHooks,
start_time: "datetime",
allow_chunking: bool,
) -> BedrockGuardrailResponse:
"""Post prepared ApplyGuardrail content and log the call once, as a success or a failure."""
credentials, aws_region_name = await run_aws_signing(
self._load_credentials, bearer_token=bedrock_bearer_token(api_key)
)
allow_chunking: Final = not self._content_uses_contextual_grounding(content)
completed_chunk_usages: Final[list[BedrockGuardrailUsage]] = [] # mutable-ok: billed-chunk usage accumulator
try:
@ -2504,6 +2537,95 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
# This means all actions were ANONYMIZED or NONE, so don't raise exception
return False
async def async_scan_request_attachments(
self,
data: dict,
call_type: CallTypesLiteral,
event_type: GuardrailEventHooks = GuardrailEventHooks.pre_call,
) -> None:
"""Scan the request's inline PNG/JPEG images with ApplyGuardrail, 20 per call, and block attachments it cannot scan.
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.
"""
attachments: Final = find_request_attachments(
data,
call_type,
skip_tool_messages=effective_skip_tool_message_for_guardrail(self),
latest_user_message_only=self.experimental_use_latest_role_message_only,
scan_only_tool_results=effective_scan_only_tool_results_for_guardrail(self),
)
images: Final = attachments.images if self.checks is None else ()
unscannable: Final = attachments.unscannable + tuple("image" for _ in attachments.images if self.checks)
if unscannable:
if self.optional_params.get("skip_unscannable_attachments") is not True:
raise self._unscannable_attachments_exception(unscannable, request_data=data, event_type=event_type)
verbose_proxy_logger.warning(
"Bedrock Guardrail %s: letting %d unscannable attachment(s) through unscanned because "
"skip_unscannable_attachments is enabled",
self.guardrail_name,
len(unscannable),
)
if not images:
return
bedrock_request_data, api_key = self._bedrock_request_body(BedrockRequest(source="INPUT"), data)
try:
for start in range(0, len(images), _BEDROCK_APPLY_GUARDRAIL_MAX_IMAGES):
await self._apply_guardrail_to_content(
content=images[start : start + _BEDROCK_APPLY_GUARDRAIL_MAX_IMAGES],
bedrock_request_data=bedrock_request_data,
api_key=api_key,
request_data=data,
event_type=event_type,
start_time=datetime.now(timezone.utc),
allow_chunking=False,
)
except (HTTPException, ModifyResponseException):
raise
except Exception as e:
if event_type != GuardrailEventHooks.pre_call:
raise
verbose_proxy_logger.error("Bedrock Guardrail: Failed to scan attachments: %s", str(e))
raise Exception(f"Bedrock guardrail failed: {e}") from e
def _unscannable_attachments_exception(
self,
unscannable: Sequence[str],
request_data: dict, # mutable-ok: proxy request body dict, mutated by the logging helper
event_type: GuardrailEventHooks,
) -> HTTPException | ModifyResponseException:
"""Log the refusal and build the block error for attachments ApplyGuardrail cannot scan."""
reason: Final = (
f"Bedrock guardrail cannot scan {len(unscannable)} attachment(s) "
f"({', '.join(sorted(set(unscannable)))}), so the request was blocked"
)
now: Final = datetime.now(timezone.utc).timestamp()
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
guardrail_json_response={"unscannable_attachments": list(unscannable)},
request_data=request_data,
guardrail_status="guardrail_intervened",
start_time=now,
end_time=now,
duration=0.0,
event_type=event_type,
)
if self.disable_exception_on_block is True:
return ModifyResponseException(
message=reason,
model=request_data.get("model", "bedrock-guardrail"),
request_data=request_data,
guardrail_name=self.guardrail_name,
)
detail: Final[dict[str, object]] = {"error": "Violated guardrail policy", "reason": reason}
if self.guardrailIdentifier:
detail["guardrailIdentifier"] = self.guardrailIdentifier
if self.guardrailVersion:
detail["guardrailVersion"] = self.guardrailVersion
return HTTPException(status_code=400, detail=detail)
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
@ -2587,6 +2709,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
return
await self.async_scan_request_attachments(data=data, call_type=call_type, event_type=event_type)
new_messages: Final = self.get_guardrails_messages_for_call_type(
call_type=cast(CallTypes, call_type),
data=data,

View file

@ -9,7 +9,7 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint
import copy
import json
from collections.abc import AsyncGenerator, AsyncIterable, Awaitable, Callable, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Protocol
from typing import TYPE_CHECKING, Any, Final, Protocol, runtime_checkable
from fastapi import HTTPException
@ -45,6 +45,13 @@ A2A_CALL_TYPES: Final = (CallTypes.asend_message, CallTypes.send_message)
GUARDRAIL_NAME: Final = "unified_llm_guardrails"
@runtime_checkable
class RequestAttachmentScanner(Protocol):
"""A guardrail that scans the raw request's attachments before its text is extracted."""
async def async_scan_request_attachments(self, data: dict, call_type: CallTypesLiteral) -> None: ...
class _EndpointTranslation(Protocol):
@property
def process_input_messages(self) -> "Callable[..., Awaitable[dict[str, object]]]": ...
@ -228,6 +235,11 @@ class UnifiedLLMGuardrails(CustomLogger):
_ensure_litellm_metadata(data, user_api_key_dict)
if isinstance(guardrail_to_apply, RequestAttachmentScanner) and hasattr(
type(guardrail_to_apply), "async_scan_request_attachments"
):
await guardrail_to_apply.async_scan_request_attachments(data=data, call_type=call_type)
data = await endpoint_translation.process_input_messages(
data=data,
guardrail_to_apply=guardrail_to_apply,

View file

@ -41,6 +41,7 @@ def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail):
aws_bedrock_runtime_endpoint=litellm_params.aws_bedrock_runtime_endpoint,
experimental_use_latest_role_message_only=litellm_params.experimental_use_latest_role_message_only,
only_scan_new_messages=litellm_params.only_scan_new_messages or False,
skip_unscannable_attachments=litellm_params.skip_unscannable_attachments,
streaming_buffer_until_moderated=streaming_params.streaming_buffer_until_moderated,
streaming_sampling_rate=streaming_params.streaming_sampling_rate,
streaming_end_of_stream_only=streaming_params.streaming_end_of_stream_only,

View file

@ -12,8 +12,18 @@ class BedrockTextContent(TypedDict, total=False):
qualifiers: list[BedrockGuardrailQualifier]
class BedrockImageSource(TypedDict):
bytes: str
class BedrockImageContent(TypedDict):
format: Literal["png", "jpeg"]
source: BedrockImageSource
class BedrockContentItem(TypedDict, total=False):
text: BedrockTextContent
image: BedrockImageContent
class BedrockRequest(TypedDict, total=False):

View file

@ -0,0 +1,330 @@
import base64
import pytest
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrail_attachments import find_request_attachments
from litellm.types.utils import CallTypes
PNG_B64 = base64.b64encode(b"\x89PNG\r\n\x1a\nfake-png").decode()
JPEG_B64 = base64.b64encode(b"\xff\xd8\xfffake-jpeg").decode()
OVERSIZE_PNG_B64 = base64.b64encode(b"\x89PNG" + b"\0" * (4 * 1024 * 1024)).decode()
PDF_B64 = base64.b64encode(b"%PDF-1.4 fake").decode()
def _png_item(encoded: str = PNG_B64) -> dict:
return {"image": {"format": "png", "source": {"bytes": encoded}}}
def _jpeg_item() -> dict:
return {"image": {"format": "jpeg", "source": {"bytes": JPEG_B64}}}
def _chat(*blocks: object, role: str = "user") -> dict:
return {"messages": [{"role": role, "content": list(blocks)}]}
def _converse(*blocks: object, endpoint: str = "model/haiku/converse", provider: str = "bedrock") -> dict:
return {
"endpoint": endpoint,
"custom_llm_provider": provider,
"data": {"messages": [{"role": "user", "content": list(blocks)}]},
}
TEXT = {"type": "text", "text": "hello"}
@pytest.mark.parametrize(
"data, call_type, images, unscannable",
[
pytest.param(_chat(TEXT), CallTypes.acompletion.value, [], [], id="chat-text-only"),
pytest.param(
{"messages": [{"role": "user", "content": "hi"}]},
CallTypes.acompletion.value,
[],
[],
id="chat-string-content",
),
pytest.param(
_chat(TEXT, {"type": "image_url", "image_url": {"url": f"data:image/png;base64,{PNG_B64}"}}),
CallTypes.acompletion.value,
[_png_item()],
[],
id="chat-png-data-uri-after-text",
),
pytest.param(
_chat({"type": "image_url", "image_url": f"data:image/jpeg;base64,{JPEG_B64}"}),
CallTypes.completion.value,
[_jpeg_item()],
[],
id="chat-jpeg-string-url",
),
pytest.param(
_chat(
{"type": "image_url", "image_url": {"url": f"data:image/png;base64,{PNG_B64}"}},
TEXT,
{"type": "image_url", "image_url": {"url": f"data:image/jpg;base64,{JPEG_B64}"}},
),
CallTypes.acompletion.value,
[_png_item(), _jpeg_item()],
[],
id="chat-two-images",
),
pytest.param(
_chat({"type": "image_url", "image_url": {"url": "https://example.com/cat.png"}}),
CallTypes.acompletion.value,
[],
["image_url (remote URL or file id)"],
id="chat-remote-image",
),
pytest.param(
_chat({"type": "image_url", "image_url": {"url": f"data:image/gif;base64,{PNG_B64}"}}),
CallTypes.acompletion.value,
[],
["image_url (image/gif)"],
id="chat-gif",
),
pytest.param(
_chat({"type": "image_url", "image_url": {"url": "data:image/png;base64,***not-base64***"}}),
CallTypes.acompletion.value,
[],
["image_url (invalid base64)"],
id="chat-malformed-base64",
),
pytest.param(
_chat({"type": "image_url", "image_url": {"url": f"data:image/png;base64,{OVERSIZE_PNG_B64}"}}),
CallTypes.acompletion.value,
[],
["image_url (over 4 MB)"],
id="chat-image-over-4mb",
),
pytest.param(
_chat({"type": "image_url", "image_url": {"url": "data:image/png;base64,"}}),
CallTypes.acompletion.value,
[],
["image_url (invalid base64)"],
id="chat-empty-base64",
),
pytest.param(
_chat(TEXT, {"type": "file", "file": {"file_data": f"data:application/pdf;base64,{PDF_B64}"}}),
CallTypes.acompletion.value,
[],
["file"],
id="chat-pdf-file",
),
pytest.param(
_chat(
{"type": "input_audio", "input_audio": {"data": "AAAA", "format": "wav"}},
{"type": "video_url", "video_url": {"url": "https://example.com/v.mp4"}},
),
CallTypes.acompletion.value,
[],
["input_audio", "video_url"],
id="chat-audio-and-video",
),
pytest.param(
_chat(
{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": PNG_B64}},
{"type": "document", "source": {"type": "base64", "media_type": "application/pdf", "data": PDF_B64}},
),
CallTypes.anthropic_messages.value,
[_png_item()],
["document"],
id="anthropic-image-and-document",
),
pytest.param(
_chat({"type": "image", "source": {"type": "url", "url": "https://example.com/cat.png"}}),
CallTypes.anthropic_messages.value,
[],
["image (url or file source)"],
id="anthropic-url-image",
),
pytest.param(
{
"input": [
{
"role": "user",
"content": [
{"type": "input_text", "text": "hi"},
{"type": "input_image", "image_url": f"data:image/png;base64,{PNG_B64}"},
{"type": "input_file", "file_data": f"data:application/pdf;base64,{PDF_B64}"},
],
}
]
},
CallTypes.aresponses.value,
[_png_item()],
["input_file"],
id="responses-image-and-file",
),
pytest.param({"input": "just text"}, CallTypes.responses.value, [], [], id="responses-string-input"),
pytest.param(
_converse({"text": "hi"}, _png_item(), {"document": {"format": "pdf", "source": {"bytes": PDF_B64}}}),
CallTypes.allm_passthrough_route.value,
[_png_item()],
["document"],
id="converse-image-and-document",
),
pytest.param(
_converse({"image": {"format": "png", "source": {"s3Location": {"uri": "s3://b/k.png"}}}}),
CallTypes.allm_passthrough_route.value,
[],
["image (no inline bytes)"],
id="converse-s3-image",
),
pytest.param(
_converse(
{"video": {"format": "mp4", "source": {"bytes": PDF_B64}}}, endpoint="model/haiku/converse-stream"
),
CallTypes.allm_passthrough_route.value,
[],
["video"],
id="converse-stream-video",
),
pytest.param(
_converse(_png_item(), endpoint="model/haiku/invoke"),
CallTypes.allm_passthrough_route.value,
[],
[],
id="bedrock-invoke-not-converse",
),
pytest.param(
_converse(_png_item(), provider="vertex_ai"),
CallTypes.allm_passthrough_route.value,
[],
[],
id="non-bedrock-passthrough",
),
pytest.param(
_chat({"type": "image_url", "image_url": {"url": f"data:image/PNG;BASE64,{PNG_B64}"}}),
CallTypes.acompletion.value,
[_png_item()],
[],
id="chat-uppercase-data-uri",
),
pytest.param(
_chat(
{
"type": "tool_result",
"tool_use_id": "toolu_1",
"content": [
{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": PNG_B64}},
{"type": "document", "source": {"type": "url", "url": "https://example.com/r.pdf"}},
],
}
),
CallTypes.anthropic_messages.value,
[_png_item()],
["document"],
id="anthropic-tool-result-image-and-document",
),
pytest.param(
{
"input": [
{
"type": "function_call_output",
"call_id": "call_1",
"output": [
{"type": "input_image", "image_url": f"data:image/png;base64,{PNG_B64}"},
{"type": "input_file", "file_id": "file-1"},
],
},
{"type": "function_call_output", "call_id": "call_2", "output": "plain text"},
]
},
CallTypes.aresponses.value,
[_png_item()],
["input_file"],
id="responses-function-call-output-image-and-file",
),
pytest.param(
_converse(
{
"toolResult": {
"toolUseId": "t1",
"content": [{"document": {"format": "pdf", "source": {"bytes": PDF_B64}}}, _png_item()],
}
}
),
CallTypes.allm_passthrough_route.value,
[_png_item()],
["document"],
id="converse-tool-result-document-and-image",
),
pytest.param(
_chat({"type": "image_url", "image_url": {"url": f"data:image/png;base64,{PNG_B64}"}}),
CallTypes.call_mcp_tool.value,
[],
[],
id="mcp-call-type-ignored",
),
],
)
def test_find_request_attachments(data, call_type, images, unscannable):
found = find_request_attachments(data, call_type, skip_tool_messages=False, latest_user_message_only=False)
assert list(found.images) == images
assert list(found.unscannable) == unscannable
def test_whitespace_in_base64_is_stripped_before_sending():
wrapped = PNG_B64[:8] + "\n" + PNG_B64[8:]
data = _chat({"type": "image_url", "image_url": {"url": f"data:image/png;base64,{wrapped}"}})
found = find_request_attachments(data, CallTypes.acompletion.value, False, False)
assert list(found.images) == [_png_item()]
def test_tool_messages_skipped_only_when_asked():
data = {
"messages": [
{"role": "user", "content": [TEXT]},
{"role": "tool", "content": [{"type": "file", "file": {"file_id": "file-1"}}]},
]
}
scanned = find_request_attachments(data, CallTypes.acompletion.value, False, False)
skipped = find_request_attachments(data, CallTypes.acompletion.value, True, False)
assert scanned.unscannable == ("file",)
assert skipped.unscannable == ()
def test_scan_only_tool_results_keeps_only_tool_output():
data = _chat(
{"type": "document", "source": {"type": "url", "url": "https://example.com/user.pdf"}},
{"type": "tool_result", "tool_use_id": "toolu_1", "content": [{"type": "document", "source": {}}]},
)
everything = find_request_attachments(data, CallTypes.anthropic_messages.value, False, False)
tool_only = find_request_attachments(data, CallTypes.anthropic_messages.value, False, False, True)
skip_tool = find_request_attachments(data, CallTypes.anthropic_messages.value, True, False)
assert everything.unscannable == ("document", "document")
assert tool_only.unscannable == ("document",)
assert skip_tool.unscannable == ("document",)
assert find_request_attachments(data, CallTypes.anthropic_messages.value, True, False, True).unscannable == ()
def test_long_mime_is_truncated_in_label():
data = _chat({"type": "image_url", "image_url": {"url": f"data:image/{'x' * 500};base64,{PNG_B64}"}})
found = find_request_attachments(data, CallTypes.acompletion.value, False, False)
assert found.unscannable == (f"image_url (image/{'x' * 34})",)
def test_latest_user_message_only():
data = {
"messages": [
{"role": "user", "content": [{"type": "file", "file": {"file_id": "old"}}]},
{"role": "assistant", "content": "ok"},
{"role": "user", "content": [{"type": "image_url", "image_url": f"data:image/png;base64,{PNG_B64}"}]},
]
}
found = find_request_attachments(data, CallTypes.acompletion.value, False, True)
assert list(found.images) == [_png_item()]
assert found.unscannable == ()

View file

@ -6063,7 +6063,9 @@ async def test_bearer_token_never_runs_the_sigv4_credential_chain(monkeypatch):
mock_response.json.return_value = {"action": "NONE", "assessments": []}
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock, return_value=mock_response) as mock_post:
response = await guardrail.make_bedrock_api_request(source="INPUT", messages=[{"role": "user", "content": "hello"}])
response = await guardrail.make_bedrock_api_request(
source="INPUT", messages=[{"role": "user", "content": "hello"}]
)
assert response["action"] == "NONE"
assert mock_post.call_args.kwargs["headers"]["Authorization"] == "Bearer env-bearer-token-12345"
@ -6100,3 +6102,235 @@ async def test_apply_guardrail_signs_off_the_event_loop(monkeypatch):
assert response["action"] == "NONE"
assert probe.served_during_refresh is True
_ATTACHMENT_PNG_B64 = "iVBORw0KGgpmYWtlLXBuZw=="
_ATTACHMENT_PDF_URI = "data:application/pdf;base64,JVBERi0xLjQgZmFrZQ=="
def _image_only_chat_request() -> dict:
return {
"model": "claude",
"messages": [
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": f"data:image/png;base64,{_ATTACHMENT_PNG_B64}"}}
],
}
],
}
def _pdf_chat_request() -> dict:
return {
"model": "claude",
"messages": [
{
"role": "user",
"content": [
{"type": "text", "text": "summarize"},
{"type": "file", "file": {"file_data": _ATTACHMENT_PDF_URI}},
],
}
],
}
def _attachment_guardrail(**kwargs) -> BedrockGuardrail:
return BedrockGuardrail(
guardrail_name="bedrock-attachments",
guardrailIdentifier="gid",
guardrailVersion="DRAFT",
**kwargs,
)
def _patched_bedrock_post(guardrail: BedrockGuardrail, response: MagicMock):
credentials = MagicMock(access_key="k", secret_key="s", token=None)
return (
patch.object(guardrail.async_handler, "post", new_callable=AsyncMock, return_value=response),
patch.object(guardrail, "_load_credentials", return_value=(credentials, "us-east-1")),
patch.object(guardrail, "_prepare_request", return_value=MagicMock()),
)
@pytest.mark.asyncio
async def test_attachment_scan_sends_image_only_turn_to_apply_guardrail_and_blocks():
guardrail = _attachment_guardrail()
post_patch, credentials_patch, prepare_patch = _patched_bedrock_post(
guardrail, _blocking_bedrock_httpx_response("violence")
)
with post_patch as mock_post, credentials_patch, prepare_patch as mock_prepare:
with pytest.raises(HTTPException) as exc_info:
await guardrail.async_scan_request_attachments(
data=_image_only_chat_request(), call_type=CallTypes.acompletion.value
)
assert exc_info.value.status_code == 400
assert mock_post.await_count == 1
sent_body = mock_prepare.call_args.kwargs["data"]
assert sent_body["source"] == "INPUT"
assert sent_body["content"] == ({"image": {"format": "png", "source": {"bytes": _ATTACHMENT_PNG_B64}}},)
@pytest.mark.asyncio
async def test_attachment_scan_lets_a_passing_image_through():
guardrail = _attachment_guardrail()
data = _image_only_chat_request()
post_patch, credentials_patch, prepare_patch = _patched_bedrock_post(
guardrail, _passing_bedrock_httpx_response("ok")
)
with post_patch as mock_post, credentials_patch, prepare_patch:
await guardrail.async_scan_request_attachments(data=data, call_type=CallTypes.acompletion.value)
assert mock_post.await_count == 1
logged = data["metadata"]["standard_logging_guardrail_information"]
assert [entry["guardrail_status"] for entry in logged] == ["success"]
@pytest.mark.asyncio
async def test_attachment_scan_blocks_unscannable_attachment_without_calling_bedrock():
guardrail = _attachment_guardrail()
data = _pdf_chat_request()
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
with pytest.raises(HTTPException) as exc_info:
await guardrail.async_scan_request_attachments(data=data, call_type=CallTypes.acompletion.value)
assert exc_info.value.status_code == 400
assert exc_info.value.detail["error"] == "Violated guardrail policy"
assert "cannot scan 1 attachment(s) (file)" in exc_info.value.detail["reason"]
assert exc_info.value.detail["guardrailIdentifier"] == "gid"
mock_post.assert_not_awaited()
logged = data["metadata"]["standard_logging_guardrail_information"]
assert logged[0]["guardrail_status"] == "guardrail_intervened"
assert logged[0]["guardrail_response"] == {"unscannable_attachments": ["file"]}
@pytest.mark.asyncio
async def test_attachment_scan_skip_unscannable_attachments_lets_documents_through():
guardrail = _attachment_guardrail(skip_unscannable_attachments=True)
data = _pdf_chat_request()
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
result = await guardrail.async_scan_request_attachments(data=data, call_type=CallTypes.acompletion.value)
assert result is None
assert data == _pdf_chat_request()
mock_post.assert_not_awaited()
@pytest.mark.asyncio
async def test_attachment_scan_block_returns_a_response_when_exceptions_are_disabled():
guardrail = _attachment_guardrail(disable_exception_on_block=True)
with pytest.raises(ModifyResponseException) as exc_info:
await guardrail.async_scan_request_attachments(data=_pdf_chat_request(), call_type=CallTypes.acompletion.value)
assert "cannot scan 1 attachment(s) (file)" in exc_info.value.message
@pytest.mark.asyncio
async def test_attachment_scan_treats_images_as_unscannable_in_checks_mode():
guardrail = BedrockGuardrail(guardrail_name="bedrock-checks", checks={"contentFilter": {}})
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
with pytest.raises(HTTPException) as exc_info:
await guardrail.async_scan_request_attachments(
data=_image_only_chat_request(), call_type=CallTypes.acompletion.value
)
assert "cannot scan 1 attachment(s) (image)" in exc_info.value.detail["reason"]
mock_post.assert_not_awaited()
@pytest.mark.asyncio
async def test_attachment_scan_does_nothing_for_text_only_requests():
guardrail = _attachment_guardrail()
data = {"model": "claude", "messages": [{"role": "user", "content": [{"type": "text", "text": "hi"}]}]}
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
await guardrail.async_scan_request_attachments(data=data, call_type=CallTypes.acompletion.value)
mock_post.assert_not_awaited()
assert "metadata" not in data
@pytest.mark.asyncio
async def test_during_call_hook_scans_image_only_turn():
guardrail = _attachment_guardrail(event_hook=GuardrailEventHooks.during_call, default_on=True)
post_patch, credentials_patch, prepare_patch = _patched_bedrock_post(
guardrail, _blocking_bedrock_httpx_response("violence")
)
with post_patch as mock_post, credentials_patch, prepare_patch:
with pytest.raises(HTTPException):
await guardrail.async_moderation_hook(
data=_image_only_chat_request(),
user_api_key_dict=UserAPIKeyAuth(),
call_type=CallTypes.acompletion.value,
)
assert mock_post.await_count == 1
@pytest.mark.asyncio
async def test_attachment_scan_in_checks_mode_never_posts_images_even_when_skipping():
guardrail = BedrockGuardrail(
guardrail_name="bedrock-checks", checks={"contentFilter": {}}, skip_unscannable_attachments=True
)
data = _image_only_chat_request()
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
result = await guardrail.async_scan_request_attachments(data=data, call_type=CallTypes.acompletion.value)
assert result is None
assert data == _image_only_chat_request()
mock_post.assert_not_awaited()
@pytest.mark.asyncio
async def test_attachment_scan_sends_at_most_twenty_images_per_call():
guardrail = _attachment_guardrail()
data = _image_only_chat_request()
data["messages"][0]["content"] = data["messages"][0]["content"] * 25
post_patch, credentials_patch, prepare_patch = _patched_bedrock_post(
guardrail, _passing_bedrock_httpx_response("ok")
)
with post_patch as mock_post, credentials_patch, prepare_patch as mock_prepare:
await guardrail.async_scan_request_attachments(data=data, call_type=CallTypes.acompletion.value)
assert mock_post.await_count == 2
assert [len(call.kwargs["data"]["content"]) for call in mock_prepare.call_args_list] == [20, 5]
@pytest.mark.asyncio
async def test_during_call_hook_skips_attachment_scan_when_guardrail_is_not_requested():
guardrail = _attachment_guardrail(event_hook=GuardrailEventHooks.during_call, default_on=False)
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
result = await guardrail.async_moderation_hook(
data=_pdf_chat_request(), user_api_key_dict=UserAPIKeyAuth(), call_type="acompletion"
)
assert result is None
assert mock_post.await_count == 0
@pytest.mark.asyncio
async def test_attachment_scan_honors_scan_only_tool_results():
guardrail = _attachment_guardrail()
guardrail.scan_only_tool_results = True
data = _pdf_chat_request()
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
result = await guardrail.async_scan_request_attachments(data=data, call_type=CallTypes.acompletion.value)
assert result is None
assert mock_post.await_count == 0
assert data == _pdf_chat_request()

View file

@ -2395,3 +2395,44 @@ class TestTranslationMappingsAreReadLive:
assert not [
name for name, value in vars(unified_module).items() if isinstance(value, dict) and CallTypes.aocr in value
]
class AttachmentScanningGuardrail(RecordingGuardrail):
"""Records the raw request it is handed for an attachment scan."""
def __init__(self):
super().__init__()
self.attachment_scans = []
async def async_scan_request_attachments(self, data, call_type):
self.attachment_scans.append({"call_type": call_type, "apply_calls_so_far": len(self.apply_calls)})
@pytest.mark.asyncio
async def test_pre_call_hook_scans_attachments_before_text_extraction():
guardrail = AttachmentScanningGuardrail()
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
cache=DualCache(),
data={"guardrail_to_apply": guardrail, "model": "gpt-4o", "input": "hello"},
call_type=CallTypes.aresponses.value,
)
assert guardrail.attachment_scans == [{"call_type": CallTypes.aresponses.value, "apply_calls_so_far": 0}]
assert len(guardrail.apply_calls) == 1
@pytest.mark.asyncio
async def test_pre_call_hook_skips_attachment_scan_when_guardrail_has_none():
guardrail = RecordingGuardrail()
guardrail.async_scan_request_attachments = None # instance attribute, not a method on the class
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
cache=DualCache(),
data={"guardrail_to_apply": guardrail, "model": "gpt-4o", "input": "hello"},
call_type=CallTypes.aresponses.value,
)
assert len(guardrail.apply_calls) == 1