From ced72d962132397b07c71590d7fdee4213acacda Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Wed, 30 Sep 2026 14:24:16 -0700 Subject: [PATCH] Keep attachment scan within the type-discipline budget --- litellm/llms/bedrock/guardrail_attachments.py | 55 +++++++++++-------- .../guardrail_hooks/bedrock_guardrails.py | 23 ++++---- 2 files changed, 45 insertions(+), 33 deletions(-) diff --git a/litellm/llms/bedrock/guardrail_attachments.py b/litellm/llms/bedrock/guardrail_attachments.py index fe05c4f6757..86be9a4b460 100644 --- a/litellm/llms/bedrock/guardrail_attachments.py +++ b/litellm/llms/bedrock/guardrail_attachments.py @@ -10,6 +10,7 @@ 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 @@ -17,6 +18,7 @@ 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 @@ -48,8 +50,11 @@ class _Block(NamedTuple): _Classified = _Image | _Unscannable | _DocumentText | None -_BlockClassifier = Callable[[Mapping[str, object]], _Classified] -_NestedToolBlocks = Callable[[Mapping[str, object]], tuple[Mapping[str, object], ...]] +_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=()) @@ -85,11 +90,11 @@ def find_request_attachments( """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 message in selected - for entry in _message_blocks(message, nested_tool_blocks) + for entry in entries if _in_scope(entry, skip_tool_messages, scan_only_tool_results) ) ) @@ -191,26 +196,28 @@ def _message_blocks(message: Mapping[str, object], nested_tool_blocks: _NestedTo def _with_document_contents(entries: Iterable[_Block]) -> tuple[_Block, ...]: - return tuple(chain.from_iterable(_with_document_content(entry) for entry in entries)) + return reduce(_with_document_level_expanded, range(_MAX_DOCUMENT_DEPTH), tuple(entries)) -def _with_document_content(entry: _Block) -> tuple[_Block, ...]: - expanded: Final[list[_Block]] = [] - pending: Final[list[_Block]] = [entry] - while pending: - current = pending.pop() - expanded.append(current) - source = _content_document_source(current.block) - if source is not None and current.document_depth < _MAX_DOCUMENT_DEPTH: - pending.extend( - reversed( - [ - _Block(inner, from_tool=current.from_tool, document_depth=current.document_depth + 1) - for inner in _mappings(source.get("content")) - ] - ) - ) - return tuple(expanded) +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: @@ -312,7 +319,9 @@ def _classify_base64(mime: object, encoded: object, label: str) -> _Classified: 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": standard}))) + return _Image( + BedrockContentItem(image=BedrockImageContent(format=image_format, source=BedrockImageSource(bytes=standard))) + ) def _standard_base64(encoded: object) -> str: diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 8f2bfe58267..024beefb9b3 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -227,8 +227,8 @@ 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]) -> list[BedrockContentItem]: - return [ +def _without_image_bytes(content: Sequence[BedrockContentItem]) -> tuple[BedrockContentItem, ...]: + return tuple( BedrockContentItem( image=BedrockImageContent( format=item["image"]["format"], @@ -238,7 +238,7 @@ def _without_image_bytes(content: Sequence[BedrockContentItem]) -> list[BedrockC if "image" in item else item for item in content - ] + ) def _is_responses_api_route(request_route: str | None) -> bool: @@ -980,12 +980,12 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): 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) + 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 @@ -1279,7 +1279,7 @@ 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, "content": _without_image_bytes(content)} + {**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, @@ -2630,7 +2630,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): """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], + messages=[{"role": "user", "content": text} for text in document_texts], # mutable-ok: API message payload request_data=request_data, logging_event_type=event_type, ) @@ -2650,12 +2650,12 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): """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" + 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)}, + guardrail_json_response={"unscannable_attachments": list(unscannable)}, # mutable-ok: JSON wire format request_data=request_data, guardrail_status="guardrail_intervened", start_time=now, @@ -2670,7 +2670,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): request_data=request_data, guardrail_name=self.guardrail_name, ) - detail: Final[dict[str, object]] = {"error": "Violated guardrail policy", "reason": reason} + 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: