This commit is contained in:
yucheng-berri 2026-09-30 21:24:25 +00:00 • committed by GitHub
commit cc252e0e6e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 1574 additions and 27 deletions

View file

@ -0,0 +1,341 @@
"""
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, Iterable, Mapping, Sequence
from functools import reduce
from itertools import chain
from types import MappingProxyType
from typing import Final, Literal, NamedTuple, TypeGuard
from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockContentItem,
BedrockImageContent,
BedrockImageSource,
)
from litellm.types.utils import CallTypes
BedrockImageFormat = Literal["png", "jpeg"]
class RequestAttachments(NamedTuple):
images: tuple[BedrockContentItem, ...]
unscannable: tuple[str, ...]
document_texts: tuple[str, ...] = ()
class _Image(NamedTuple):
item: BedrockContentItem
class _Unscannable(NamedTuple):
label: str
class _DocumentText(NamedTuple):
text: str
class _Block(NamedTuple):
block: Mapping[str, object]
from_tool: bool
document_depth: int = 0
_Classified = _Image | _Unscannable | _DocumentText | None
_BlockClassifier = Callable[[Mapping[str, object]], _Classified] # mutable-ok: Callable parameter list
_NestedToolBlocks = Callable[
[Mapping[str, object]], # mutable-ok: Callable parameter list
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})
_OPENAI_IMAGE_TYPES: Final = frozenset({"image_url", "input_image", "computer_screenshot"})
_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"})
_MAX_DOCUMENT_DEPTH: Final = 3
_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"})
_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, the document text 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)
entries: Final = chain.from_iterable(_message_blocks(message, nested_tool_blocks) for message in selected)
classified: Final = tuple(
chain.from_iterable(
_classify_entry(entry, classify)
for entry in entries
if _in_scope(entry, skip_tool_messages, scan_only_tool_results)
)
)
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)),
document_texts=tuple(result.text for result in classified if isinstance(result, _DocumentText)),
)
def _classify_entry(entry: _Block, classify: _BlockClassifier) -> tuple[_Image | _Unscannable | _DocumentText, ...]:
if entry.document_depth >= _MAX_DOCUMENT_DEPTH and _content_document_source(entry.block) is not None:
return (_Unscannable("document (nested too deep)"),)
text: Final = _document_text(entry)
result: Final = classify(entry.block)
return tuple(item for item in (_DocumentText(text) if text else None, result) if item is not None)
def _document_text(entry: _Block) -> str:
"""Return the text a block carries as part of a document: a nested text block, or a document's title, context and text source."""
block: Final = entry.block
if entry.document_depth > 0 and block.get("type") == "text":
text: Final = block.get("text")
return text if isinstance(text, str) else ""
if block.get("type") != "document":
return ""
source: Final = block.get("source")
source_type: Final = source.get("type") if _is_mapping(source) else None
source_text: Final = (
(source.get("data") if source_type == "text" else source.get("content"))
if _is_mapping(source) and isinstance(source_type, str) and source_type in _TEXT_DOCUMENT_SOURCE_TYPES
else None
)
return "\n".join(
part for part in (block.get("title"), block.get("context"), source_text) if isinstance(part, str) and part
)
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_openai_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_openai_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 _is_mapping(body) 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 _is_mapping(value: object) -> TypeGuard[Mapping[str, object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip
return isinstance(value, Mapping)
def _is_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip
return isinstance(value, list)
def _mappings(value: object) -> tuple[Mapping[str, object], ...]:
if not _is_list(value):
return ()
return tuple(item for item in value if _is_mapping(item))
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, ...]:
message_type: Final = message.get("type")
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 _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 _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 reduce(_with_document_level_expanded, range(_MAX_DOCUMENT_DEPTH), tuple(entries))
def _with_document_level_expanded(entries: tuple[_Block, ...], depth: int) -> tuple[_Block, ...]:
return tuple(
chain.from_iterable(
_with_document_children(entry) if entry.document_depth == depth else (entry,) for entry in entries
)
)
def _with_document_children(entry: _Block) -> tuple[_Block, ...]:
source: Final = _content_document_source(entry.block)
if source is None:
return (entry,)
return (
entry,
*(
_Block(inner, from_tool=entry.from_tool, document_depth=entry.document_depth + 1)
for inner in _mappings(source.get("content"))
),
)
def _content_document_source(block: Mapping[str, object]) -> Mapping[str, object] | None:
source: Final = block.get("source")
if block.get("type") != "document" or not _is_mapping(source) or source.get("type") != "content":
return None
return source
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 _is_mapping(tool_result) else ()
def _classify_nothing(block: Mapping[str, object]) -> _Classified:
return None
def _classify_openai_block(block: Mapping[str, object]) -> _Classified:
block_type: Final = block.get("type")
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")
url: Final = image_url.get("url") if _is_mapping(image_url) else image_url
return _classify_data_uri(url, block_type)
if isinstance(block_type, str) and block_type in _OPENAI_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 _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):
image_format: Final = image.get("format")
source: Final = image.get("source")
encoded: Final = source.get("bytes") if _is_mapping(source) 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 block.get(key) is not None:
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)
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=BedrockImageSource(bytes=standard)))
)
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(standard, validate=True))
except (binascii.Error, ValueError):
return None

View file

@ -10460,7 +10460,7 @@
}
],
"default": false,
"description": "Implemented by guardrail='model_armor'. When True, attachment references that carry no inline bytes (file_id, gs://, or http(s) URLs) pass through unscanned instead of blocking, while fail_on_error still governs real Model Armor API errors. Default False blocks them.",
"description": "Implemented by guardrail='model_armor' and guardrail='bedrock'. When True, attachments the guardrail cannot scan pass through unscanned instead of blocking. For Model Armor these are references with no inline bytes (file_id, gs://, or http(s) URLs), and fail_on_error still governs real Model Armor API errors. For Bedrock these are documents, files, audio, video, and images that are not inline PNG or JPEG up to 4 MB. Default False blocks them.",
"title": "Skip Unscannable Attachments"
},
"sticky_session_routing": {
@ -13256,7 +13256,7 @@
}
],
"default": false,
"description": "Implemented by guardrail='model_armor'. When True, attachment references that carry no inline bytes (file_id, gs://, or http(s) URLs) pass through unscanned instead of blocking, while fail_on_error still governs real Model Armor API errors. Default False blocks them.",
"description": "Implemented by guardrail='model_armor' and guardrail='bedrock'. When True, attachments the guardrail cannot scan pass through unscanned instead of blocking. For Model Armor these are references with no inline bytes (file_id, gs://, or http(s) URLs), and fail_on_error still governs real Model Armor API errors. For Bedrock these are documents, files, audio, video, and images that are not inline PNG or JPEG up to 4 MB. Default False blocks them.",
"title": "Skip Unscannable Attachments"
},
"sticky_session_routing": {

View file

@ -43,8 +43,10 @@ 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.bedrock.guardrail_attachments import find_request_attachments
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -75,6 +77,8 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockGuardrailQualifier,
BedrockGuardrailResponse,
BedrockGuardrailUsage,
BedrockImageContent,
BedrockImageSource,
BedrockRequest,
BedrockTextContent,
)
@ -119,6 +123,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
@ -222,6 +227,20 @@ def _redact_assessment_match_fields(assessments: list[dict]) -> list[dict]:
_RESPONSES_API_CALL_TYPES: Final = frozenset({CallTypes.responses, CallTypes.aresponses})
def _without_image_bytes(content: Sequence[BedrockContentItem]) -> tuple[BedrockContentItem, ...]:
return tuple(
BedrockContentItem(
image=BedrockImageContent(
format=item["image"]["format"],
source=BedrockImageSource(bytes=f"<{len(item['image']['source']['bytes'])} base64 chars>"),
)
)
if "image" in item
else item
for item in content
)
def _is_responses_api_route(request_route: str | None) -> bool:
if request_route is None:
return False
@ -924,21 +943,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 +964,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) # mutable-ok: JSON request body
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(
{ # mutable-ok: JSON request body
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:
@ -1230,7 +1279,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
headers_dict: Final = dict(prepared_request.headers) # mutable-ok: the masking helper requires a dict
verbose_proxy_logger.debug(
"Bedrock AI request body: %s, url %s, headers: %s",
bedrock_request_data,
{**bedrock_request_data, "content": _without_image_bytes(content)} # mutable-ok: JSON request body
if any("image" in item for item in content)
else bedrock_request_data,
prepared_request.url,
_get_masked_values(headers_dict),
)
@ -2504,6 +2555,131 @@ 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 and text documents, and block attachments it cannot scan.
Images go to ApplyGuardrail 20 per call. Text documents go through ``make_bedrock_api_request``
as one user turn, and any intervention on them blocks the request, since their text cannot be
rewritten with a mask.
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 guardrail whose
``apply_guardrail`` or ``make_bedrock_api_request`` is replaced, by a subclass or on the
instance, skips this scan.
"""
if (
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(
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 attachments.document_texts:
await self._scan_document_texts(attachments.document_texts, request_data=data, event_type=event_type)
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
async def _scan_document_texts(
self,
document_texts: Sequence[str],
request_data: dict, # mutable-ok: proxy request body dict, mutated by the logging helper
event_type: GuardrailEventHooks,
) -> None:
"""Scan text documents like a user turn and block when the guardrail intervenes on them."""
response: Final = await self.make_bedrock_api_request(
source="INPUT",
messages=[{"role": "user", "content": text} for text in document_texts], # mutable-ok: API message payload
request_data=request_data,
logging_event_type=event_type,
)
if response.get("action") == "GUARDRAIL_INTERVENED":
raise self._unscannable_attachments_exception(
tuple("document (text the guardrail would mask)" for _ in document_texts),
request_data=request_data,
event_type=event_type,
)
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(frozenset(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)}, # mutable-ok: JSON wire format
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]] = { # mutable-ok: HTTPException detail
"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,
@ -2520,6 +2696,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
event_type: Final[GuardrailEventHooks] = GuardrailEventHooks.pre_call
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
return data
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),
@ -2587,6 +2764,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
@ -46,6 +46,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]]]": ...
@ -229,6 +236,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

@ -1041,9 +1041,11 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
skip_unscannable_attachments: bool | None = Field(
default=False,
description=(
"Implemented by guardrail='model_armor'. When True, attachment references that carry no "
"inline bytes (file_id, gs://, or http(s) URLs) pass through unscanned instead of blocking, "
"while fail_on_error still governs real Model Armor API errors. Default False blocks them."
"Implemented by guardrail='model_armor' and guardrail='bedrock'. When True, attachments the "
"guardrail cannot scan pass through unscanned instead of blocking. For Model Armor these are "
"references with no inline bytes (file_id, gs://, or http(s) URLs), and fail_on_error still "
"governs real Model Armor API errors. For Bedrock these are documents, files, audio, video, "
"and images that are not inline PNG or JPEG up to 4 MB. Default False blocks them."
),
)
sanitize_error_detail: bool | None = Field(

View file

@ -1,6 +1,6 @@
from typing import Literal
from typing_extensions import TypedDict
from typing_extensions import ReadOnly, TypedDict
# Bedrock contextual grounding tags each content block so the guardrail knows
# which text is the reference source, the user question, and the content to grade.
@ -12,8 +12,18 @@ class BedrockTextContent(TypedDict, total=False):
qualifiers: list[BedrockGuardrailQualifier]
class BedrockImageSource(TypedDict):
bytes: ReadOnly[str]
class BedrockImageContent(TypedDict):
format: ReadOnly[Literal["png", "jpeg"]]
source: ReadOnly[BedrockImageSource]
class BedrockContentItem(TypedDict, total=False):
text: BedrockTextContent
image: BedrockImageContent
class BedrockRequest(TypedDict, total=False):

View file

@ -4,6 +4,7 @@ Unit tests for Bedrock Guardrails
import json
import asyncio
import logging
from datetime import datetime, timezone
import sys
from unittest.mock import AsyncMock, MagicMock, patch
@ -14,6 +15,7 @@ from fastapi import HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.exceptions import ModifyResponseException
from litellm.proxy._types import UserAPIKeyAuth
@ -6063,7 +6065,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 +6104,438 @@ 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()
@pytest.mark.asyncio
async def test_native_pre_call_hook_blocks_unscannable_attachment():
guardrail = _attachment_guardrail(event_hook=GuardrailEventHooks.pre_call, default_on=True)
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
with pytest.raises(HTTPException) as exc_info:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=MagicMock(),
data=_pdf_chat_request(),
call_type="acompletion",
)
assert exc_info.value.status_code == 400
assert "cannot scan 1 attachment(s) (file)" in exc_info.value.detail["reason"]
assert mock_post.await_count == 0
@pytest.mark.asyncio
async def test_attachment_scan_debug_log_omits_image_bytes():
guardrail = _attachment_guardrail()
post_patch, credentials_patch, prepare_patch = _patched_bedrock_post(
guardrail, _passing_bedrock_httpx_response("ok")
)
with (
post_patch,
credentials_patch,
prepare_patch,
patch("litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails.verbose_proxy_logger.debug") as mock_debug,
):
await guardrail.async_scan_request_attachments(
data=_image_only_chat_request(), call_type=CallTypes.acompletion.value
)
logged = " ".join(str(call.args) for call in mock_debug.call_args_list)
assert _ATTACHMENT_PNG_B64 not in logged
assert f"<{len(_ATTACHMENT_PNG_B64)} base64 chars>" in logged
@pytest.mark.asyncio
async def test_text_only_debug_log_prints_the_signed_request_body():
guardrail = _attachment_guardrail()
post_patch, credentials_patch, prepare_patch = _patched_bedrock_post(
guardrail, _passing_bedrock_httpx_response("ok")
)
captured_records: list[logging.LogRecord] = []
class _RecordingHandler(logging.Handler):
def emit(self, record: logging.LogRecord) -> None:
captured_records.append(record)
handler = _RecordingHandler(level=logging.DEBUG)
previous_level = verbose_proxy_logger.level
verbose_proxy_logger.addHandler(handler)
verbose_proxy_logger.setLevel(logging.DEBUG)
try:
with post_patch, credentials_patch, prepare_patch:
await guardrail.make_bedrock_api_request(
source="INPUT", messages=[{"role": "user", "content": "hello"}], request_data={}
)
finally:
verbose_proxy_logger.removeHandler(handler)
verbose_proxy_logger.setLevel(previous_level)
body_lines = [
record.getMessage() for record in captured_records if record.getMessage().startswith("Bedrock AI request body")
]
assert len(body_lines) == 1
assert "'content': ({'text': {'text': 'hello'}},)" in body_lines[0]
def _text_document_messages_request(text: str) -> dict:
return {
"model": "claude",
"messages": [
{
"role": "user",
"content": [
{"type": "document", "source": {"type": "text", "media_type": "text/plain", "data": text}},
{"type": "text", "text": "summarize"},
],
}
],
}
def _anonymizing_bedrock_httpx_response() -> MagicMock:
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"action": "GUARDRAIL_INTERVENED",
"outputs": [{"text": "SSN {US_SOCIAL_SECURITY_NUMBER}"}],
"assessments": [
{
"sensitiveInformationPolicy": {
"piiEntities": [{"type": "US_SOCIAL_SECURITY_NUMBER", "action": "ANONYMIZED"}]
}
}
],
"usage": {"contentPolicyUnits": 1},
}
return response
@pytest.mark.asyncio
@pytest.mark.parametrize(
"response, blocked",
[
pytest.param(_passing_bedrock_httpx_response("ok"), False, id="pass"),
pytest.param(_blocking_bedrock_httpx_response("deny"), True, id="block"),
pytest.param(_anonymizing_bedrock_httpx_response(), True, id="anonymize"),
],
)
async def test_attachment_scan_sends_text_document_as_text(response, blocked):
guardrail = _attachment_guardrail()
post_patch, credentials_patch, prepare_patch = _patched_bedrock_post(guardrail, response)
with post_patch as mock_post, credentials_patch, prepare_patch as mock_prepare:
if blocked:
with pytest.raises(HTTPException) as exc_info:
await guardrail.async_scan_request_attachments(
data=_text_document_messages_request("SSN 123-45-6789"),
call_type=CallTypes.anthropic_messages.value,
)
assert exc_info.value.status_code == 400
else:
result = await guardrail.async_scan_request_attachments(
data=_text_document_messages_request("SSN 123-45-6789"),
call_type=CallTypes.anthropic_messages.value,
)
assert result is None
assert mock_post.await_count == 1
sent_body = mock_prepare.call_args.kwargs["data"]
assert sent_body["source"] == "INPUT"
assert sent_body["content"] == ({"text": {"text": "SSN 123-45-6789"}},)
class _CustomApplyGuardrail(BedrockGuardrail):
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
return inputs
@pytest.mark.asyncio
async def test_attachment_scan_skipped_when_subclass_overrides_apply_guardrail():
guardrail = _CustomApplyGuardrail(
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
@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"}
@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

@ -41,7 +41,14 @@ from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrai
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.utils import CallTypes, Delta, GenericGuardrailAPIInputs, ModelResponseStream, StreamingChoices
from litellm.types.utils import (
CallTypes,
CallTypesLiteral,
Delta,
GenericGuardrailAPIInputs,
ModelResponseStream,
StreamingChoices,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -2395,3 +2402,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) -> None:
super().__init__()
self.attachment_scans: list[dict[str, object]] = []
async def async_scan_request_attachments(self, data: dict, call_type: CallTypesLiteral) -> None:
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

View file

@ -0,0 +1,515 @@
import base64
import pytest
from litellm.llms.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(
{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": PNG_B64}},
{"type": "input_image", "image_url": f"data:image/jpeg;base64,{JPEG_B64}"},
{"type": "input_file", "file_id": "file-1"},
),
CallTypes.acompletion.value,
[_png_item(), _jpeg_item()],
["input_file"],
id="chat-anthropic-and-responses-shaped-parts",
),
pytest.param(
{"input": [{"role": "user", "content": [{"type": "image_url", "image_url": {"url": "https://x/y.png"}}]}]},
CallTypes.aresponses.value,
[],
["image_url (remote URL or file id)"],
id="responses-chat-shaped-image-url",
),
pytest.param(
_converse({"audio": {"format": "mp3", "source": {"bytes": PDF_B64}}}),
CallTypes.allm_passthrough_route.value,
[],
["audio"],
id="converse-audio",
),
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(
{
"input": [
{
"type": "computer_call_output",
"call_id": "call_1",
"output": {"type": "computer_screenshot", "image_url": f"data:image/png;base64,{PNG_B64}"},
},
{
"type": "computer_call_output",
"call_id": "call_2",
"output": {"type": "computer_screenshot", "file_id": "file-1"},
},
]
},
CallTypes.responses.value,
[_png_item()],
["computer_screenshot (remote URL or file id)"],
id="responses-computer-screenshot",
),
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 == ()
def test_non_string_role_and_type_do_not_raise():
data = {
"messages": [
{"role": "user", "content": [{"type": "image_url", "image_url": f"data:image/png;base64,{PNG_B64}"}]},
{"role": ["tool"], "type": {"x": 1}, "content": [{"type": "file", "file": {"file_id": "f"}}]},
{"role": {"r": "tool"}, "type": ["function_call_output"], "content": [{"type": "file", "file": {}}]},
]
}
found = find_request_attachments(data, CallTypes.acompletion.value, True, False)
assert list(found.images) == [_png_item()]
assert found.unscannable == ("file", "file")
def test_converse_null_document_is_not_an_attachment():
data = _converse(_png_item(), {"document": None}, {"video": None, "audio": None}, {"audio": {"format": "mp3"}})
found = find_request_attachments(data, CallTypes.allm_passthrough_route.value, False, False)
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": [{"type": "text", "text": "hi"}]}},
id="content",
),
pytest.param({"type": "document", "source": {"type": "content", "content": "hi"}}, id="content-string"),
],
)
@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",)
assert found.document_texts == ("hi",)
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)",)
@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}}
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), call_type, False, False)
assert list(found.images) == [_png_item()]
assert found.unscannable == ("document",)
assert found.document_texts == ("hello",)
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]})
scanned = find_request_attachments(data, CallTypes.anthropic_messages.value, False, False)
skipped = find_request_attachments(data, CallTypes.anthropic_messages.value, True, False)
assert list(scanned.images) == [_png_item()]
assert skipped.images == ()
def test_document_title_and_context_are_scanned_as_text():
text_doc = {
"type": "document",
"title": "record",
"context": "SSN 123-45-6789",
"source": {"type": "text", "media_type": "text/plain", "data": "body"},
}
pdf = {
"type": "document",
"context": "note",
"source": {"type": "base64", "media_type": "application/pdf", "data": PDF_B64},
}
found = find_request_attachments(_chat(text_doc, pdf), CallTypes.anthropic_messages.value, False, False)
assert found.document_texts == ("record\nSSN 123-45-6789\nbody", "note")
assert found.unscannable == ("document",)
def _content_document(*inner: object) -> dict:
return {"type": "document", "source": {"type": "content", "content": list(inner)}}
def test_nested_documents_are_scanned_and_deep_nesting_is_refused():
inner_text = {"type": "document", "source": {"type": "text", "media_type": "text/plain", "data": "inner"}}
nested = _content_document(inner_text, _content_document({"type": "text", "text": "deeper"}))
too_deep = _content_document(_content_document(_content_document(_content_document(TEXT))))
found = find_request_attachments(_chat(nested, too_deep), CallTypes.anthropic_messages.value, False, False)
assert found.document_texts == ("inner", "deeper")
assert found.unscannable == ("document (nested too deep)",)

View file

@ -25781,7 +25781,7 @@ export interface components {
skip_tool_message_in_guardrail?: boolean | null;
/**
* Skip Unscannable Attachments
* @description Implemented by guardrail='model_armor'. When True, attachment references that carry no inline bytes (file_id, gs://, or http(s) URLs) pass through unscanned instead of blocking, while fail_on_error still governs real Model Armor API errors. Default False blocks them.
* @description Implemented by guardrail='model_armor' and guardrail='bedrock'. When True, attachments the guardrail cannot scan pass through unscanned instead of blocking. For Model Armor these are references with no inline bytes (file_id, gs://, or http(s) URLs), and fail_on_error still governs real Model Armor API errors. For Bedrock these are documents, files, audio, video, and images that are not inline PNG or JPEG up to 4 MB. Default False blocks them.
* @default false
*/
skip_unscannable_attachments: boolean | null;
@ -35007,7 +35007,7 @@ export interface components {
skip_tool_message_in_guardrail?: boolean | null;
/**
* Skip Unscannable Attachments
* @description Implemented by guardrail='model_armor'. When True, attachment references that carry no inline bytes (file_id, gs://, or http(s) URLs) pass through unscanned instead of blocking, while fail_on_error still governs real Model Armor API errors. Default False blocks them.
* @description Implemented by guardrail='model_armor' and guardrail='bedrock'. When True, attachments the guardrail cannot scan pass through unscanned instead of blocking. For Model Armor these are references with no inline bytes (file_id, gs://, or http(s) URLs), and fail_on_error still governs real Model Armor API errors. For Bedrock these are documents, files, audio, video, and images that are not inline PNG or JPEG up to 4 MB. Default False blocks them.
* @default false
*/
skip_unscannable_attachments: boolean | null;