diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrail_attachments.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrail_attachments.py new file mode 100644 index 00000000000..a5b90393c07 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrail_attachments.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 228b31604a3..5a58a8abbad 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index d68a55f9a88..e6e5e127bc7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index c422902d30d..51e2ad4d80e 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -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, diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 8d66b624341..4d46d3594f0 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -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): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrail_attachments.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrail_attachments.py new file mode 100644 index 00000000000..26e8e66f430 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrail_attachments.py @@ -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 == () diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index 45ad336368c..42bcdf678c1 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -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() diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index c90f88ec110..c544ee2013c 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -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