Keep attachment scan within the type-discipline budget
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled

This commit is contained in:
Yucheng He 2026-09-30 14:24:16 -07:00
parent 63f6095c35
commit ced72d9621
2 changed files with 45 additions and 33 deletions

View file

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

View file

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