From 03bd6cec34acdf36aa96f2e0c2dfa92162a5ea49 Mon Sep 17 00:00:00 2001 From: "feng.tsai" Date: Tue, 22 Sep 2026 14:37:06 +0800 Subject: [PATCH] fix(guardrails): send image parts to bedrock ApplyGuardrail Bedrock guardrails with the IMAGE modality enabled only ever saw the request's text: images went to the model unscanned while the guardrail still reported as applied, leaving contentPolicyImageUnits at zero. Inline png and jpeg images are now extracted and sent to ApplyGuardrail as image content blocks. Image forms the API cannot scan (other attachment types, remote urls, more than the 20 images one call accepts) are refused before any decode, fetch or AWS call rather than passing silently. The text-only shortcut and the incremental scan paths are image-aware, so an image-bearing request cannot skip the scan. convert_to_bedrock_format is now async, so its callers await it. --- .../guardrail_hooks/bedrock_guardrails.py | 377 ++++++++- .../guardrail_hooks/bedrock_guardrails.py | 17 +- .../test_bedrock_guardrails.py | 4 +- .../test_bedrock_guardrails.py | 766 +++++++++++++++++- .../guardrails/test_guardrail_endpoints.py | 4 +- .../proxy/guardrails/test_init_guardrails.py | 5 +- 6 files changed, 1128 insertions(+), 45 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 434c52ca6f3..4bfc5035779 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -10,6 +10,7 @@ import sys sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path import asyncio +import binascii import copy import json import re @@ -19,7 +20,7 @@ from collections.abc import AsyncGenerator, Mapping, Sequence from datetime import datetime, timezone from itertools import accumulate, groupby from types import MappingProxyType -from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, NamedTuple, Optional, cast +from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, NamedTuple, NoReturn, Optional, cast import httpx from fastapi import HTTPException @@ -40,6 +41,7 @@ from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import ( bedrock_guardrail_cost_by_unit, guardrail_cost_total, ) +from litellm.litellm_core_utils.prompt_templates.factory import BedrockImageProcessor 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, @@ -65,16 +67,24 @@ from litellm.types.guardrails import ( GuardrailEventHooks, LitellmParams, ) -from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage +from litellm.types.llms.openai import ( + AllMessageValues, + ChatCompletionImageObject, + ChatCompletionImageUrlObject, + ChatCompletionUserMessage, +) from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( BedrockChecksMessage, BedrockChecksViolation, BedrockContentItem, BedrockGuardrailChecksResponse, + BedrockGuardrailImageFormat, + BedrockGuardrailImageSource, BedrockGuardrailOutput, BedrockGuardrailQualifier, BedrockGuardrailResponse, BedrockGuardrailUsage, + BedrockImageContent, BedrockRequest, BedrockTextContent, ) @@ -146,6 +156,20 @@ _CONTENT_TYPE_TO_QUALIFIER: Final[dict[str, BedrockGuardrailQualifier]] = { # must not be graded against as if it were the application's own source material. _GROUNDING_SOURCE_TRUSTED_ROLES: Final = frozenset({"system", "developer"}) +# AWS image-filter limits, none of which litellm checked before. Anything other than +# png/jpeg has no representation in the payload and cannot be scanned at all. +# https://docs.aws.amazon.com/bedrock/latest/userguide/guardrails-mmfilter.html +_MAX_IMAGE_BYTES: Final = 4 * 1024 * 1024 +_MAX_IMAGES_PER_APPLY_GUARDRAIL_CALL: Final = 20 + +_APPLY_GUARDRAIL_IMAGE_FORMATS: Final[ + dict[str, BedrockGuardrailImageFormat] +] = { # mutable-ok: module-level lookup table, never mutated + "png": "png", + "jpeg": "jpeg", + "jpg": "jpeg", +} + class QualifiedTextBlock(NamedTuple): """A piece of message text paired with its Bedrock grounding qualifier (if any).""" @@ -374,28 +398,266 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) return cleaned or None - def _create_bedrock_input_content_request(self, messages: list[AllMessageValues] | None) -> BedrockRequest: - """ - Create a bedrock request for the input content - the LLM request. - """ + async def _create_bedrock_input_content_request(self, messages: list[AllMessageValues] | None) -> BedrockRequest: + """Create a bedrock request for the input content - the LLM request.""" bedrock_request: Final[BedrockRequest] = BedrockRequest(source="INPUT") - bedrock_request_content: Final[list[BedrockContentItem]] = [] if messages is None: return bedrock_request - for message in messages: - blocks = self.get_content_items_for_message(message=message) - if blocks is None: - continue - for block in blocks: - # INPUT scans send plain text only. Grounding qualifiers are attached - # exclusively when assembling the OUTPUT request, so a caller cannot use - # a grounding_source/query tag to change how input-safety policies treat - # their content (which would be an input-guardrail bypass). - bedrock_request_content.append(BedrockContentItem(text=BedrockTextContent(text=block.text))) - bedrock_request["content"] = bedrock_request_content + image_count: Final = self._image_count_in(messages) + if image_count > _MAX_IMAGES_PER_APPLY_GUARDRAIL_CALL: + raise HTTPException( + status_code=400, + detail={ # mutable-ok: HTTPException detail payload, serialized immediately + "error": "Violated guardrail policy", + "bedrock_guardrail_response": ( + f"Request contains {image_count} images; Bedrock ApplyGuardrail accepts at most " + f"{_MAX_IMAGES_PER_APPLY_GUARDRAIL_CALL} images per request" + ), + "guardrail_name": self.guardrail_name, + }, + ) + + per_message: Final = await asyncio.gather( + *(self._build_input_content_items(message=message) for message in messages) + ) + # mutable-ok: BedrockRequest["content"] is a list in the AWS wire format + bedrock_request["content"] = [item for items in per_message for item in items] return bedrock_request + async def _build_input_content_items(self, message: AllMessageValues) -> tuple[BedrockContentItem, ...]: + """Flatten one request message into ApplyGuardrail INPUT content items. + + Grounding qualifiers are attached only when assembling the OUTPUT request, so a + grounding_source/query tag cannot change how input-safety policies treat content + """ + content: Final = message.get("content") + if content is None: + return () + if isinstance(content, str): + return (BedrockContentItem(text=BedrockTextContent(text=content)),) + if not isinstance(content, list): + return () + parts: Final = cast( # cast-ok: AllMessageValues content is a union of part TypedDicts + tuple[object, ...], tuple(content) + ) + items: Final = await asyncio.gather(*(self._build_input_content_item(item=item) for item in parts)) + return tuple(item for item in items if item is not None) + + async def _build_input_content_item(self, item: object) -> BedrockContentItem | None: + if isinstance(item, str): + return BedrockContentItem(text=BedrockTextContent(text=item)) + if not isinstance(item, dict): + return None + part: Final = cast(Mapping[str, object], item) # cast-ok: narrowed to dict on the line above + # Provider transformations branch on `type`, so an image_url part reaches the + # model as an image even when it also carries `text`. Reading `text` first would + # scan that decoy and forward the image unscanned + if part.get("type") == "image_url": + image_url: Final = self._get_image_url(item=part) + if image_url is None: + return None + return await self._build_image_content_item(image_url=image_url) + text: Final = part.get("text") + if isinstance(text, str): + return BedrockContentItem(text=BedrockTextContent(text=text)) + return None + + @staticmethod + def _get_image_url(item: Mapping[str, object]) -> str | None: + """Pull the url out of an image_url part. The caller owns the type dispatch.""" + image_url: Final = item.get("image_url") + if isinstance(image_url, str): + return image_url + if isinstance(image_url, dict): + url: Final = cast(Mapping[str, object], image_url).get("url") # cast-ok: narrowed to dict on the line above + return url if isinstance(url, str) else None + return None + + @classmethod + def _image_urls_in(cls, messages: "Sequence[AllMessageValues] | None") -> frozenset[str]: + """Normalized image urls already carried by these messages.""" + found: Final[set[str]] = set() # mutable-ok: accumulator, frozen on return + for message in messages or (): + content = message.get("content") + if not isinstance(content, list): + continue + for part in content: + if not isinstance(part, dict) or part.get("type") != "image_url": + continue + url = cls._get_image_url(item=part) + if url is not None: + found.add(cls._normalize_image_input(url)) + return frozenset(found) + + @classmethod + def _image_count_in(cls, messages: "Sequence[AllMessageValues] | None") -> int: + """Count image occurrences in the exact messages sent to ApplyGuardrail.""" + count = 0 # rebind-ok: running count over request content + for message in messages or (): + content = message.get("content") + if not isinstance(content, list): + continue + count += sum( + 1 + for part in content + if isinstance(part, dict) and part.get("type") == "image_url" and cls._get_image_url(part) is not None + ) + return count + + def _handle_unscannable_image(self, reason: str) -> NoReturn: + """Block an image part ApplyGuardrail cannot scan. + + The image reaches the model either way, so skipping it silently would let a + caller defeat an IMAGE-modality guardrail by picking a format the API rejects + """ + raise HTTPException( + status_code=400, + detail={ # mutable-ok: HTTPException detail payload, serialized immediately + "error": "Violated guardrail policy", + "bedrock_guardrail_response": ( + f"Request contains an image the guardrail cannot scan ({reason}). " + "Bedrock ApplyGuardrail accepts inline png/jpeg images only" + ), + "guardrail_name": self.guardrail_name, + }, + ) + + #: base64 magic-byte prefixes for the formats ApplyGuardrail accepts. + _BASE64_IMAGE_PREFIXES: ClassVar[tuple[tuple[str, str], ...]] = ( + ("iVBORw0KGgo", "image/png"), + ("/9j/", "image/jpeg"), + ) + + @staticmethod + def _image_content_part(url: str) -> ChatCompletionImageObject: + """One OpenAI-format image content part, so the payload builder handles it.""" + return ChatCompletionImageObject(type="image_url", image_url=ChatCompletionImageUrlObject(url=url)) + + @classmethod + def _normalize_image_input(cls, value: str) -> str: + """Return a data URI or URL that `_build_image_content_item` can consume. + + `GenericGuardrailAPIInputs["images"]` is not a single shape. The OpenAI chat + translation appends the caller's `image_url` verbatim, so entries are already a + `data:` URI or an `https://` URL. The Anthropic translation's `_image_sources` + returns `source["data"]` only, which is bare base64 with the `media_type` + dropped. Sniff the format back from the base64 prefix so both shapes reach the + same decoder instead of the bare-base64 one failing as unreadable. + """ + if value.startswith(("data:", "http://", "https://")): + return value + for prefix, media_type in cls._BASE64_IMAGE_PREFIXES: + if value.startswith(prefix): + return f"data:{media_type};base64,{value}" + # Unrecognized inputs reach the decoder so the guardrail fails closed. + return value + + def _refuse_file_backed_images(self, request_data: Mapping[str, object], input_type: str) -> None: + """Reject a file-backed image rather than ignoring it. + + Reads the raw request, not inputs["structured_messages"]: the /v1/messages + handler fills that field by translating to OpenAI spec, which drops a file + source outright, so a check reading it never fires while the provider still + forwards the file. The cost is that the scope flags no longer narrow this + check, and over-refusing is the safer error here. + """ + if input_type != "request": + return + found: Final = self._file_backed_image_count(request_data.get("messages")) + if not found: + return + self._handle_unscannable_image( + reason=f"{found} image(s) reference a provider file id, whose bytes are not available here" + ) + + @classmethod + def _file_backed_image_count(cls, messages: object) -> int: + """Count Anthropic `{"type": "image", "source": {"type": "file"}}` parts. + + Matches that one shape rather than comparing counts against + `inputs["images"]`, whose length the skip and scope flags already narrow, so + a mismatch there is not by itself evidence of a dropped image. + """ + if not isinstance(messages, list): + return 0 + return sum(cls._file_backed_parts(message.get("content")) for message in messages if isinstance(message, dict)) + + @classmethod + def _file_backed_parts(cls, content: object) -> int: + """Count file-backed images in one content list, descending into tool_result. + + The extractor pulls scannable images out of a tool_result's nested blocks, so + a file source sitting there has to be refused for the same reason a top-level + one is: nothing else in the request will surface it. + """ + if not isinstance(content, list): + return 0 + found = 0 # rebind-ok: running count over the content list + for part in content: + if not isinstance(part, dict): + continue + if part.get("type") == "tool_result": + found += cls._file_backed_parts(part.get("content")) + continue + if part.get("type") != "image": + continue + source = part.get("source") + if isinstance(source, dict) and source.get("type") == "file": + found += 1 + return found + + async def _build_image_content_item(self, image_url: str) -> BedrockContentItem: + """Decode an inline image into an ApplyGuardrail image block. + + A remote url is named as its own rejection rather than left to the decoder: + fetching one is a separate piece of work (size cap, SSRF, and handing the + same bytes to the model), so the operator gets "not supported" instead of + the decoder's "could not be read". Anything else that is not a data URI is + an unrecognized payload and falls through to the decoder, which rejects it. + + The test is a substring, not a prefix, deliberately: it has to reject + everything `BedrockImageProcessor.process_image_async` would treat as + remote, and that check is `"http://" in image_url or "https://" in + image_url`. A prefix test reads more naturally but leaves a hole -- a url + carrying leading whitespace fails it, then matches downstream and is + fetched, so the fail-closed policy here would be bypassed into an + uncapped server-side download. Keep the two predicates identical. + """ + if "http://" in image_url or "https://" in image_url: + self._handle_unscannable_image(reason="remote image URLs are not supported") + + try: + block: Final = await BedrockImageProcessor.process_image_async(image_url=image_url, format=None) + except (ValueError, TypeError, KeyError, binascii.Error) as e: + self._handle_unscannable_image(reason=f"image content could not be read: {e}") + + image_block: Final = block.get("image") + image_format: Final = ( + _APPLY_GUARDRAIL_IMAGE_FORMATS.get(str(image_block.get("format"))) if image_block else None + ) + image_source: Final = image_block.get("source") if image_block else None + image_bytes: Final = image_source.get("bytes") if image_source else None + if image_format is None or not image_bytes: + self._handle_unscannable_image(reason="attachment is not a png/jpeg image") + + # base64 carries 3 bytes per 4 characters, less whatever the padding stands in + # for, so the size follows from the string without decoding the whole image a + # second time. Counting the padding matters at the boundary: without it a 4 MiB + # image measures 2 bytes over and gets refused. + decoded_size: Final = len(image_bytes) // 4 * 3 - image_bytes.count("=") + if decoded_size > _MAX_IMAGE_BYTES: + self._handle_unscannable_image( + reason=f"image is {decoded_size / 1024 / 1024:.1f} MB, over ApplyGuardrail's 4 MB limit" + ) + + return BedrockContentItem( + image=BedrockImageContent( + format=image_format, + source=BedrockGuardrailImageSource(bytes=image_bytes), + ) + ) + def _create_bedrock_output_content_request( self, response: object, @@ -442,7 +704,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): items.append(self._build_content_item(block)) return items - def convert_to_bedrock_format( + async def convert_to_bedrock_format( self, source: Literal["INPUT", "OUTPUT"], messages: list[AllMessageValues] | None = None, @@ -457,12 +719,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): Returns: BedrockRequest: The bedrock request object. """ - bedrock_request: BedrockRequest = BedrockRequest(source=source) if source == "INPUT": - bedrock_request = self._create_bedrock_input_content_request(messages=messages) - elif source == "OUTPUT": - bedrock_request = self._create_bedrock_output_content_request(response=response, messages=messages) - return bedrock_request + return await self._create_bedrock_input_content_request(messages=messages) + if source == "OUTPUT": + return self._create_bedrock_output_content_request(response=response, messages=messages) + return BedrockRequest(source=source) def get_content_items_for_message(self, message: AllMessageValues) -> list[QualifiedTextBlock] | None: """ @@ -925,7 +1186,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): """ start_time: Final = datetime.now(timezone.utc) bedrock_request_data: Final[dict] = dict( - self.convert_to_bedrock_format(source=source, messages=messages, response=response) + await self.convert_to_bedrock_format(source=source, messages=messages, response=response) ) api_key: str | None = None if request_data: @@ -3185,7 +3446,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): logging_obj: Optional logging object Returns: - GenericGuardrailAPIInputs - processed_texts may be masked, images unchanged + GenericGuardrailAPIInputs - processed_texts may be masked, images are + scanned but returned unchanged (ApplyGuardrail does not rewrite images) Raises: Exception: If content is blocked by Bedrock guardrail @@ -3193,10 +3455,28 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): # NOTE: Use `or []` to handle case where inputs["texts"] is explicitly None. # dict.get("texts", []) would return None if the key exists with a None value. texts: Final = inputs.get("texts") or [] + # Images only exist on the request side; ApplyGuardrail's OUTPUT source takes + # model-generated text. The endpoint translations already extract them: + # OpenAIChatCompletionsHandler from `image_url` parts, AnthropicMessagesHandler + # from `image`/`source` blocks. Five other guardrails already consume this + # field; Bedrock was the one that dropped it on the floor. + image_urls: Final = tuple(inputs.get("images") or ()) if input_type == "request" else () + # Before the shortcuts below, so a file-backed image cannot be skipped either. + # Both conditions live in the callee: apply_guardrail sits one branch under + # ruff-strict's complexity ceiling, and two more here would cross it. + self._refuse_file_backed_images(request_data=request_data, input_type=input_type) try: - verbose_proxy_logger.debug("Bedrock Guardrail: Applying guardrail to %s text(s)", len(texts)) + verbose_proxy_logger.debug( + "Bedrock Guardrail: Applying guardrail to %s text(s) and %s image(s)", len(texts), len(image_urls) + ) - if input_type == "request": + # Both of the optimizations below decide what to scan by looking at + # `texts` alone, so an image rides along unscanned and the proxy still + # reports the guardrail as run. Neither tracks images in its session + # cache either, so "already seen" cannot be established for them. + # Presence of an image therefore forces the full image-aware path: + # correctness over the optimization, and it fails closed. + if input_type == "request" and not image_urls: incremental_result: Final = await self._apply_incremental_request_scan( texts=texts, inputs=inputs, @@ -3213,14 +3493,17 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): request_data=request_data, input_type=input_type, ) - if selection.skip_scan: + # `experimental_use_latest_role_message_only` marks a latest message with + # no text as skip_scan; an image-only message is exactly that shape. + if selection.skip_scan and not image_urls: return inputs filtered_messages: Final = selection.filtered_messages scanned_slice: Final = selection.scanned_slice scanned_role_subset: Final = selection.scanned_role_subset - # Bedrock will throw an error if there is no text to process - if filtered_messages: + # Bedrock rejects an empty content list, so only skip when there is + # neither text nor an image to scan. + if filtered_messages or image_urls: _log_hook = GuardrailEventHooks.pre_call if input_type == "request" else GuardrailEventHooks.post_call # Map the abstract input_type to the Bedrock source parameter. # "request" -> INPUT (scan user-supplied content) @@ -3243,7 +3526,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ), finish_reason="stop", ) - for _idx, _msg in enumerate(filtered_messages) + # `or ()`: skip_scan is now bypassed when an image is present, + # and images exist on the request side only, so a response scan + # still always has messages here. Spelled out rather than left + # leaning on that indirection. + for _idx, _msg in enumerate(filtered_messages or ()) ] ) bedrock_response = await self.make_bedrock_api_request( @@ -3254,9 +3541,33 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): logging_event_type=_log_hook, ) else: + # Append the images as one extra user message. Reusing the normal + # message path means `_create_bedrock_input_content_request` does the + # decoding and fail-closed format checks, so the + # unified and native lifecycle paths cannot drift apart. + # `experimental_use_latest_role_message_only` puts the selected + # message itself into filtered_messages, image parts included, and + # those go through the same builder below. Appending them again + # would scan and bill each one twice. + already_scanned: Final = self._image_urls_in(filtered_messages) + image_parts: Final = [ # mutable-ok: OpenAI message content is a list in the wire format + self._image_content_part(normalized) + for normalized in (self._normalize_image_input(url) for url in image_urls) + if normalized not in already_scanned + ] + image_message: Final = ( + (ChatCompletionUserMessage(role="user", content=image_parts),) if image_parts else () + ) + # `filtered_messages` is Optional; the guard above only proves one of + # it and `image_urls` is truthy, so it can still be None here when the + # request carries an image and no text. + scan_messages: Final = [ # mutable-ok: make_bedrock_api_request takes a list of messages + *(filtered_messages or ()), + *image_message, + ] bedrock_response = await self.make_bedrock_api_request( source="INPUT", - messages=filtered_messages, + messages=scan_messages, request_data=request_data, logging_event_type=_log_hook, ) diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 8d66b624341..39db8d613b7 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -1,6 +1,6 @@ -from typing import Literal +from typing import Literal, TypeAlias -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,21 @@ class BedrockTextContent(TypedDict, total=False): qualifiers: list[BedrockGuardrailQualifier] +BedrockGuardrailImageFormat: TypeAlias = Literal["png", "jpeg"] + + +class BedrockGuardrailImageSource(TypedDict, total=False): + bytes: ReadOnly[str] + + +class BedrockImageContent(TypedDict, total=False): + format: ReadOnly[BedrockGuardrailImageFormat] + source: ReadOnly[BedrockGuardrailImageSource] + + class BedrockContentItem(TypedDict, total=False): text: BedrockTextContent + image: ReadOnly[BedrockImageContent] class BedrockRequest(TypedDict, total=False): diff --git a/tests/guardrails_tests/test_bedrock_guardrails.py b/tests/guardrails_tests/test_bedrock_guardrails.py index 43d088268eb..bb958bcc304 100644 --- a/tests/guardrails_tests/test_bedrock_guardrails.py +++ b/tests/guardrails_tests/test_bedrock_guardrails.py @@ -945,7 +945,7 @@ async def test_convert_to_bedrock_format_input_source(): ] # Call the method - result = guardrail.convert_to_bedrock_format(source="INPUT", messages=mock_messages) + result = await guardrail.convert_to_bedrock_format(source="INPUT", messages=mock_messages) # Verify the result structure assert isinstance(result, dict) @@ -1007,7 +1007,7 @@ async def test_convert_to_bedrock_format_output_source(): ) # Call the method - result = guardrail.convert_to_bedrock_format( + result = await guardrail.convert_to_bedrock_format( source="OUTPUT", response=mock_response ) 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 1d3d7a452b6..befb8cfb880 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 @@ -2,8 +2,9 @@ Unit tests for Bedrock Guardrails """ -import json import asyncio +import base64 +import json from datetime import datetime, timezone import sys from unittest.mock import AsyncMock, MagicMock, patch @@ -2419,13 +2420,15 @@ _GUARD_BLOCK = {"text": {"text": _GROUNDING_RESPONSE_TEXT, "qualifiers": ["guard def _input_request(messages: list) -> dict: """Arrange a guardrail and act: build the Bedrock INPUT payload.""" - return _grounding_guardrail().convert_to_bedrock_format(source="INPUT", messages=messages) + return asyncio.run(_grounding_guardrail().convert_to_bedrock_format(source="INPUT", messages=messages)) def _output_request(messages: list, response=None, from_messages: bool = False) -> dict: """Arrange a guardrail and act: build the Bedrock OUTPUT payload.""" - return _grounding_guardrail(from_messages).convert_to_bedrock_format( - source="OUTPUT", response=response, messages=messages + return asyncio.run( + _grounding_guardrail(from_messages).convert_to_bedrock_format( + source="OUTPUT", response=response, messages=messages + ) ) @@ -2598,8 +2601,8 @@ def test_grounding_input_ignores_plain_message_derivation(): "content": [{"text": {"text": _GROUNDING_SOURCE_TEXT}}, {"text": {"text": _GROUNDING_QUERY_TEXT}}], } - actual_request = _grounding_guardrail(from_messages=True).convert_to_bedrock_format( - source="INPUT", messages=messages + actual_request = asyncio.run( + _grounding_guardrail(from_messages=True).convert_to_bedrock_format(source="INPUT", messages=messages) ) assert actual_request == expected_request @@ -5496,6 +5499,757 @@ async def test_terminal_failure_logs_usage_and_cost_of_prior_passed_chunks(monke assert "error" in logged["guardrail_response"] +class TestBedrockGuardrailImageInput: + """Image parts must reach ApplyGuardrail, not just the text sitting next to them.""" + + _PNG_DATA_URI = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUg==" + _GIF_DATA_URI = "data:image/gif;base64,R0lGODlhAQABAAAAACw=" + def _guardrail(self, **kwargs) -> BedrockGuardrail: + return BedrockGuardrail( + guardrail_name="bedrock-image", + guardrailIdentifier="gr-image", + guardrailVersion="DRAFT", + **kwargs, + ) + + @pytest.mark.asyncio + async def test_inline_image_is_sent_for_scanning(self): + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "what does this say?"}, + {"type": "image_url", "image_url": {"url": self._PNG_DATA_URI}}, + ], + } + ] + + request = await self._guardrail().convert_to_bedrock_format(source="INPUT", messages=messages) + + assert request["content"] == [ + {"text": {"text": "what does this say?"}}, + {"image": {"format": "png", "source": {"bytes": self._PNG_DATA_URI.split(",")[1]}}}, + ] + + @pytest.mark.asyncio + async def test_image_part_carrying_a_text_field_is_still_scanned_as_an_image(self): + """A part tagged image_url reaches the model as an image, text field or not. + + Provider transformations branch on `type`, so reading `text` first would scan + the decoy and forward the image unscanned - the exact bypass this change closes. + """ + messages = [ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": {"url": self._PNG_DATA_URI}, + "text": "just a friendly note", + } + ], + } + ] + + request = await self._guardrail().convert_to_bedrock_format(source="INPUT", messages=messages) + + assert request["content"] == [ + {"image": {"format": "png", "source": {"bytes": self._PNG_DATA_URI.split(",")[1]}}} + ] + + @pytest.mark.asyncio + async def test_unscannable_image_part_carrying_a_text_field_still_blocks(self): + """The decoy text must not turn an unscannable image into a scanned request.""" + messages = [ + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": self._GIF_DATA_URI}, "text": "hello"}], + } + ] + + with pytest.raises(HTTPException) as exc_info: + await self._guardrail().convert_to_bedrock_format(source="INPUT", messages=messages) + + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_unscannable_image_blocks_the_request_by_default(self): + """ApplyGuardrail takes png/jpeg only, and the image reaches the model either way. + + Skipping it silently would let a caller defeat an IMAGE-modality guardrail by + sending a gif, so the default is to reject the request. + """ + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello"}, + {"type": "image_url", "image_url": {"url": self._GIF_DATA_URI}}, + ], + } + ] + + with pytest.raises(HTTPException) as exc_info: + await self._guardrail().convert_to_bedrock_format(source="INPUT", messages=messages) + + assert exc_info.value.status_code == 400 + assert exc_info.value.detail["error"] == "Violated guardrail policy" + + @pytest.mark.asyncio + async def test_undecodable_image_blocks_the_request_by_default(self): + messages = [ + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": "not-an-image"}}], + } + ] + + with pytest.raises(HTTPException) as exc_info: + await self._guardrail().convert_to_bedrock_format(source="INPUT", messages=messages) + + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_incremental_scan_does_not_skip_a_request_carrying_an_image(self): + """`only_scan_new_messages` decides what to scan from `texts` alone. + + With a session id and every text segment already seen, the incremental path + returns before the image-aware code runs, so benign text plus a violating + image is never scanned -- and the proxy still reports the guardrail as run. + Nothing caches images either, so "already seen" cannot cover them. + """ + g = self._guardrail(only_scan_new_messages=True) + sent: list = [] + + async def spy(**kwargs): + sent.append(await g.convert_to_bedrock_format(source="INPUT", messages=kwargs["messages"])) + return {"action": "NONE", "outputs": []} + + # Every text is already in the session cache: the exact state that made the + # incremental path return early. + with ( + patch.object(g, "filter_new_texts_for_session", new=AsyncMock(return_value=[])), + patch.object(g, "make_bedrock_api_request", new=spy), + ): + await g.apply_guardrail( + inputs={"texts": ["look at this"], "images": [self._PNG_DATA_URI]}, + request_data={"litellm_session_id": "sess-1"}, + input_type="request", + ) + + assert sent, "image-carrying request was skipped entirely by the incremental path" + kinds = [k for item in sent[0]["content"] for k in item] + assert "image" in kinds, f"image never reached the payload: {kinds}" + + @pytest.mark.asyncio + async def test_latest_message_only_does_not_skip_an_image_only_message(self): + """`experimental_use_latest_role_message_only` skips a latest message with no text. + + An image-only user message is exactly that shape, so the whole request was + dropped from the scan. + """ + g = self._guardrail(experimental_use_latest_role_message_only=True) + sent: list = [] + + async def spy(**kwargs): + sent.append(await g.convert_to_bedrock_format(source="INPUT", messages=kwargs["messages"])) + return {"action": "NONE", "outputs": []} + + with patch.object(g, "make_bedrock_api_request", new=spy): + await g.apply_guardrail( + inputs={ + "texts": [], + "images": [self._PNG_DATA_URI], + "structured_messages": [ + {"role": "user", "content": [{"type": "image_url", "image_url": {"url": self._PNG_DATA_URI}}]} + ], + }, + request_data={}, + input_type="request", + ) + + assert sent, "image-only latest message was skipped entirely" + kinds = [k for item in sent[0]["content"] for k in item] + assert "image" in kinds, f"image never reached the payload: {kinds}" + + @pytest.mark.asyncio + async def test_twenty_inline_images_are_accepted(self): + messages = [ + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": self._PNG_DATA_URI}} for _ in range(20)], + } + ] + + request = await self._guardrail().convert_to_bedrock_format(source="INPUT", messages=messages) + + assert len(request["content"]) == 20 + + @pytest.mark.asyncio + async def test_twenty_one_duplicate_images_are_rejected_before_decode(self): + """Repeated tiny data URIs still consume image slots and must not amplify AWS calls.""" + messages = [ + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": self._PNG_DATA_URI}} for _ in range(21)], + } + ] + + with patch( # test-quality-ok: the decoder is the thing under assertion -- the cap must reject before it is ever awaited + "litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails.BedrockImageProcessor.process_image_async", + new_callable=AsyncMock, + ) as decode: + with pytest.raises(HTTPException) as exc_info: + await self._guardrail().convert_to_bedrock_format(source="INPUT", messages=messages) + + decode.assert_not_awaited() + assert "at most 20 images" in str(exc_info.value.detail) + + @pytest.mark.parametrize( + "content", + [ + pytest.param(None, id="no content"), + pytest.param(123, id="content is not a list"), + pytest.param([123], id="part is not a mapping"), + pytest.param([{"type": "image_url"}], id="image part with no url"), + pytest.param([{"type": "image_url", "image_url": 123}], id="url is not a string or mapping"), + pytest.param([{"type": "image_url", "image_url": {"url": 123}}], id="url value is not a string"), + pytest.param([{"type": "input_audio"}], id="part carries neither image nor text"), + ], + ) + @pytest.mark.asyncio + async def test_a_malformed_part_never_becomes_scannable_content(self, content): + """Default-deny is the point of this PR, so pin it rather than trust it. + + Each shape below is one a caller can put on the wire. None of them may turn + into a content item: an unrecognised part that fell through to the text + branch would be reported to the operator as scanned when it was not. + """ + request = await self._guardrail().convert_to_bedrock_format( + source="INPUT", messages=[{"role": "user", "content": content}] + ) + + assert request["content"] == [] + + @pytest.mark.asyncio + async def test_a_bare_string_part_is_scanned_as_text(self): + """A content list may hold plain strings, not only typed parts.""" + request = await self._guardrail().convert_to_bedrock_format( + source="INPUT", messages=[{"role": "user", "content": ["just text"]}] + ) + + assert request["content"] == [{"text": {"text": "just text"}}] + + @pytest.mark.asyncio + async def test_image_url_given_as_a_plain_string_is_accepted(self): + """OpenAI accepts `image_url` as a bare string as well as `{"url": ...}`. + + Both reach the model as an image, so both have to reach the scan. + """ + request = await self._guardrail().convert_to_bedrock_format( + source="INPUT", + messages=[{"role": "user", "content": [{"type": "image_url", "image_url": self._PNG_DATA_URI}]}], + ) + + kinds = [k for item in request["content"] for k in item] + assert kinds == ["image"] + + @pytest.mark.asyncio + async def test_an_unrecognized_payload_is_rejected(self): + """_normalize_image_input sniffs png and jpeg out of bare base64. + + Anything else is handed to the decoder as-is rather than guessed at, and the + guardrail rejects it instead of forwarding the image unscanned. + """ + # Reached through apply_guardrail: bare base64 arrives in inputs["images"], + # which is the only caller that normalizes before decoding. + with pytest.raises(HTTPException) as exc_info: + await self._guardrail().apply_guardrail( + inputs={"texts": [], "images": ["R0lGODlhAQABAAAAACw="]}, + request_data={}, + input_type="request", + ) + + assert "could not be read" in str(exc_info.value.detail) or "not a png/jpeg" in str(exc_info.value.detail) + + @pytest.mark.parametrize( + "url", + [ + pytest.param("https://example.com/pic.png", id="plain https url"), + pytest.param("http://example.com/pic.png", id="plain http url"), + # The rejection used to be a prefix test, which these slip past while + # `process_image_async` still treats them as remote and fetches them -- + # an uncapped server-side download straight through the fail-closed path. + pytest.param(" https://example.com/pic.png", id="leading space"), + pytest.param("\thttps://example.com/pic.png", id="leading tab"), + pytest.param("\nhttps://example.com/pic.png", id="leading newline"), + ], + ) + @pytest.mark.asyncio + async def test_a_remote_image_url_is_rejected_without_being_fetched(self, url): + """A remote url is named as its own rejection rather than left to the decoder. + + Fetching one safely (size cap, SSRF/redirect validation) is separate work; this + PR only scans inline images, so a url has to fail closed rather than be ignored + or silently forwarded to the model unscanned. The refusal has to catch every + shape the decoder would fetch, not just the well-formed ones. + """ + messages = [{"role": "user", "content": [{"type": "image_url", "image_url": {"url": url}}]}] + + with patch( # test-quality-ok: the fetch is the thing under assertion -- it must never be reached + "litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails.BedrockImageProcessor.process_image_async", + new_callable=AsyncMock, + ) as decode: + with pytest.raises(HTTPException) as exc_info: + await self._guardrail().convert_to_bedrock_format(source="INPUT", messages=messages) + + decode.assert_not_awaited() + assert "remote image URLs are not supported" in str(exc_info.value.detail) + + def test_file_backed_part_counting_skips_non_mapping_entries(self): + """A content list may mix plain strings in with typed parts. + + `_file_backed_image_count` walks the raw request rather than a normalized + shape, so a non-dict entry must be skipped rather than raise or be miscounted. + """ + content = [ + "just a string", + {"type": "text", "text": "hi"}, + {"type": "image", "source": {"type": "file", "file_id": "file_abc"}}, + ] + + assert BedrockGuardrail._file_backed_parts(content) == 1 + + def test_the_url_helper_guards_its_own_inputs(self): + """Exercised directly so the guards are not dropped in a later refactor.""" + assert BedrockGuardrail._get_image_url(item={"type": "image_url"}) is None + assert BedrockGuardrail._get_image_url(item={"type": "image_url", "image_url": {"url": 7}}) is None + assert BedrockGuardrail._get_image_url(item={"type": "image_url", "image_url": 7}) is None + + @pytest.mark.asyncio + async def test_a_file_backed_image_is_refused_rather_than_ignored(self): + """`{"type": "file"}` carries no bytes, so nothing reaches inputs["images"]. + + The provider still forwards the file to the model, so ignoring it is exactly + the silent pass this path exists to remove. Documented is not the same as + safe; under the default policy the request is refused. + """ + inputs = {"texts": ["what does this say?"], "images": []} + request_data = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "what does this say?"}, + {"type": "image", "source": {"type": "file", "file_id": "file_abc"}}, + ], + } + ] + } + + g = self._guardrail() + sent: list = [] + + async def spy(**kwargs): + sent.append(kwargs["messages"]) + return {"action": "NONE", "outputs": []} + + # Stubbed so that without the refusal this request would simply succeed: + # the failure mode being pinned is a silent pass, not an AWS error. + with patch.object(g, "make_bedrock_api_request", new=spy): + with pytest.raises(HTTPException) as exc_info: + await g.apply_guardrail(inputs=inputs, request_data=request_data, input_type="request") + + assert "file id" in str(exc_info.value.detail) + assert sent == [], "refused before any scan was attempted" + + @pytest.mark.asyncio + async def test_a_file_backed_image_is_refused_through_the_real_translation(self): + """Drive the /v1/messages handler instead of hand-building its output. + + The handler translates to OpenAI spec before filling structured_messages, and + that translation drops a file source, so a check reading structured_messages + passes every hand-written fixture and never fires in production. + """ + from litellm.llms.anthropic.chat.guardrail_translation.handler import AnthropicMessagesHandler + + data = { + "model": "claude-sonnet-4-5", + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "what does this say?"}, + {"type": "image", "source": {"type": "file", "file_id": "file_abc"}}, + ], + } + ], + } + assert not self._file_parts_in(AnthropicMessagesHandler().get_structured_messages(data)), ( + "the translation is expected to drop the file source; that is why this test exists" + ) + + g = self._guardrail() + with patch.object(g, "make_bedrock_api_request", new=AsyncMock(return_value={"action": "NONE"})): + with pytest.raises(HTTPException) as exc_info: + await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=g) + + assert "file id" in str(exc_info.value.detail) + + @pytest.mark.asyncio + async def test_a_file_backed_image_inside_a_tool_result_is_refused(self): + """The extractor pulls scannable images out of a tool_result's nested blocks. + + A file source sitting there is invisible to both: it yields no bytes to scan + and, until the count descended into tool_result, no refusal either. + """ + g = self._guardrail() + sent: list = [] + + async def spy(**kwargs): + sent.append(kwargs["messages"]) + return {"action": "NONE", "outputs": []} + + with patch.object(g, "make_bedrock_api_request", new=spy): + with pytest.raises(HTTPException) as exc_info: + await g.apply_guardrail( + inputs={"texts": ["what does this say?"], "images": []}, + request_data={ + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "what does this say?"}, + { + "type": "tool_result", + "tool_use_id": "tu_1", + "content": [ + {"type": "image", "source": {"type": "file", "file_id": "file_abc"}} + ], + }, + ], + } + ] + }, + input_type="request", + ) + + assert "file id" in str(exc_info.value.detail) + assert sent == [] + + @pytest.mark.asyncio + async def test_a_scannable_image_inside_a_tool_result_is_not_refused(self): + """Descending into tool_result must not start refusing what can be scanned.""" + g = self._guardrail() + sent: list = [] + + async def spy(**kwargs): + sent.append(kwargs["messages"]) + return {"action": "NONE", "outputs": []} + + with patch.object(g, "make_bedrock_api_request", new=spy): + await g.apply_guardrail( + inputs={"texts": ["hello"], "images": [self._PNG_DATA_URI]}, + request_data={ + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello"}, + { + "type": "tool_result", + "tool_use_id": "tu_1", + "content": [ + {"type": "image", "source": {"type": "base64", "data": "AAAA"}} + ], + }, + ], + } + ] + }, + input_type="request", + ) + + assert sent, "a scannable nested image still has to be scanned, not refused" + + @staticmethod + def _file_parts_in(messages) -> int: + return sum( + 1 + for message in messages or () + if isinstance(message, dict) + for part in (message.get("content") if isinstance(message.get("content"), list) else ()) + if isinstance(part, dict) and part.get("type") == "image" + ) + + @pytest.mark.asyncio + async def test_latest_message_only_does_not_scan_the_same_image_twice(self): + """The selected message carries its own image parts into the scan payload. + + `inputs["images"]` holds that same url, so appending it again would fetch and + bill the image twice. + """ + url = self._PNG_DATA_URI + g = self._guardrail(experimental_use_latest_role_message_only=True) + sent: list = [] + + async def spy(**kwargs): + sent.append(await g.convert_to_bedrock_format(source="INPUT", messages=kwargs["messages"])) + return {"action": "NONE", "outputs": []} + + with patch.object(g, "make_bedrock_api_request", new=spy): + await g.apply_guardrail( + inputs={ + "texts": ["hello"], + "images": [url], + "structured_messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello"}, + {"type": "image_url", "image_url": {"url": url}}, + ], + } + ], + }, + request_data={}, + input_type="request", + ) + + images = [item for item in sent[0]["content"] if "image" in item] + assert len(images) == 1, f"the image was sent {len(images)} times" + + @pytest.mark.asyncio + async def test_the_scannable_source_shapes_are_not_refused(self): + """The refusal has to be specific to the shape that cannot be read. + + structured_messages is already narrowed by the skip and scope flags, so a + count mismatch against inputs["images"] is not evidence of a dropped image. + Blocking a legitimate request would be worse than the gap being closed. + """ + g = self._guardrail() + sent: list = [] + + async def spy(**kwargs): + sent.append(await g.convert_to_bedrock_format(source="INPUT", messages=kwargs["messages"])) + return {"action": "NONE", "outputs": []} + + with patch.object(g, "make_bedrock_api_request", new=spy): + await g.apply_guardrail( + inputs={"texts": ["hello"], "images": [self._PNG_DATA_URI]}, + request_data={ + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello"}, + {"type": "image", "source": {"type": "base64", "data": "AAAA"}}, + {"type": "image", "source": {"type": "url", "url": "https://example.com/a.png"}}, + ], + } + ] + }, + input_type="request", + ) + + kinds = [k for item in sent[0]["content"] for k in item] + assert "image" in kinds + + @pytest.mark.asyncio + async def test_a_malformed_structured_message_does_not_derail_the_file_check(self): + """structured_messages comes from the caller, so its shape is not guaranteed. + + The scan must keep walking past an entry it cannot read rather than throwing + or giving up, or a single junk element would hide a file image sitting after + it -- turning a defensive guard into the bypass it was meant to prevent. + """ + g = self._guardrail() + sent: list = [] + + async def spy(**kwargs): + sent.append(kwargs["messages"]) + return {"action": "NONE", "outputs": []} + + with patch.object(g, "make_bedrock_api_request", new=spy): + with pytest.raises(HTTPException) as exc_info: + await g.apply_guardrail( + inputs={"texts": ["hello"], "images": []}, + request_data={ + "messages": [ + "not a message", + 123, + {"role": "user", "content": "a plain string, not a list"}, + { + "role": "user", + "content": [{"type": "image", "source": {"type": "file", "file_id": "file_abc"}}], + }, + ] + }, + input_type="request", + ) + + assert "file id" in str(exc_info.value.detail) + assert sent == [] + + @pytest.mark.asyncio + async def test_a_file_backed_image_on_the_response_side_is_not_refused(self): + """Images are a request-side concern; an OUTPUT scan takes generated text.""" + g = self._guardrail() + + async def spy(**kwargs): + return {"action": "NONE", "outputs": []} + + with patch.object(g, "make_bedrock_api_request", new=spy): + result = await g.apply_guardrail( + inputs={ + "texts": ["the model said this"], + "structured_messages": [ + { + "role": "user", + "content": [{"type": "image", "source": {"type": "file", "file_id": "file_abc"}}], + } + ], + }, + request_data={}, + input_type="response", + ) + + assert result is not None + + @pytest.mark.asyncio + async def test_apply_guardrail_scans_images_from_inputs(self): + """The proxy reaches BedrockGuardrail through `apply_guardrail`, not the native hook. + + ProxyLogging._execute_guardrail_hook routes any guardrail that defines + `apply_guardrail` through unified_guardrail, so a fix that only touches + async_pre_call_hook never runs on a real request. The endpoint translations + already put images in inputs["images"]; this asserts they reach the payload. + """ + g = self._guardrail() + sent: list = [] + + async def spy(**kwargs): + sent.append(await g.convert_to_bedrock_format(source="INPUT", messages=kwargs["messages"])) + return {"action": "NONE", "outputs": []} + + with patch.object(g, "make_bedrock_api_request", new=spy): + await g.apply_guardrail( + inputs={"texts": ["what does this say?"], "images": [self._PNG_DATA_URI]}, + request_data={}, + input_type="request", + ) + + kinds = [k for item in sent[0]["content"] for k in item] + assert "image" in kinds, f"image never reached the payload: {kinds}" + + @pytest.mark.asyncio + async def test_apply_guardrail_scans_bare_base64_images(self): + """Anthropic's translation drops media_type and passes bare base64. + + `_image_sources` returns source["data"] only, so the entry is not a data URI. + Without sniffing the format back it would be rejected as unreadable and turn + a legitimate /v1/messages call into a 400. + """ + g = self._guardrail() + sent: list = [] + + async def spy(**kwargs): + sent.append(await g.convert_to_bedrock_format(source="INPUT", messages=kwargs["messages"])) + return {"action": "NONE", "outputs": []} + + bare_base64 = self._PNG_DATA_URI.split(",")[1] + with patch.object(g, "make_bedrock_api_request", new=spy): + await g.apply_guardrail( + inputs={"texts": [], "images": [bare_base64]}, + request_data={}, + input_type="request", + ) + + kinds = [k for item in sent[0]["content"] for k in item] + assert "image" in kinds, f"bare base64 image never reached the payload: {kinds}" + + @pytest.mark.asyncio + async def test_apply_guardrail_ignores_images_on_the_response_side(self): + """ApplyGuardrail's OUTPUT source scans model-generated text, not input images.""" + g = self._guardrail() + sent: list = [] + + async def spy(**kwargs): + sent.append(kwargs) + return {"action": "NONE", "outputs": []} + + with patch.object(g, "make_bedrock_api_request", new=spy): + await g.apply_guardrail( + inputs={"texts": ["some model output"], "images": [self._PNG_DATA_URI]}, + request_data={}, + input_type="response", + ) + + assert sent and sent[0]["source"] == "OUTPUT" + + @pytest.mark.asyncio + async def test_an_unknown_source_yields_an_empty_request(self): + """The dispatch went from seed-then-assign to early returns when INPUT became + async. The trailing fallback preserves what the seeded request used to return + for a source that is neither INPUT nor OUTPUT, so nothing starts scanning an + unrecognized source as if it were input. + """ + from typing import Literal, cast + + # cast-ok: pins the runtime fallback the Literal forbids at type-check time + source = cast(Literal["INPUT", "OUTPUT"], "SOMETHING_ELSE") + request = await self._guardrail().convert_to_bedrock_format( + source=source, + messages=[{"role": "user", "content": "hello"}], + ) + + assert request == {"source": "SOMETHING_ELSE"} + + @pytest.mark.asyncio + async def test_oversized_image_is_rejected_before_sending(self): + """ApplyGuardrail caps images at 4 MB; AWS's rejection is not worth a round trip. + + Built as a real data URI rather than a patched decoder so the size check runs + against what the decoder actually produces. + """ + oversized_png = base64.b64encode(b"\x89PNG\r\n\x1a\n" + b"\x00" * (5 * 1024 * 1024)).decode() + + with pytest.raises(HTTPException) as exc_info: + await self._guardrail().convert_to_bedrock_format( + source="INPUT", + messages=[ + { + "role": "user", + "content": [ + {"type": "image_url", "image_url": {"url": f"data:image/png;base64,{oversized_png}"}} + ], + } + ], + ) + + assert "4 MB limit" in str(exc_info.value.detail) + + @pytest.mark.asyncio + async def test_image_at_exactly_the_limit_is_scanned(self): + """4 MB is what ApplyGuardrail accepts, so an image of exactly that size belongs + in the scan. Sizing it from the base64 length alone reads 2 bytes over, because + padding rounds the string up to a multiple of four, and the image gets refused. + """ + exact_png = base64.b64encode(b"\x89PNG\r\n\x1a\n" + b"\x00" * (4 * 1024 * 1024 - 8)).decode() + + request = await self._guardrail().convert_to_bedrock_format( + source="INPUT", + messages=[ + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": f"data:image/png;base64,{exact_png}"}}], + } + ], + ) + + assert request["content"] == [{"image": {"format": "png", "source": {"bytes": exact_png}}}] + + def test_load_credentials_assumes_role_with_external_id(): """A trust policy requiring sts:ExternalId must be satisfied by the guardrail's aws_external_id.""" import datetime diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index bf641fd6cd0..6dfe9e30d22 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -944,7 +944,9 @@ async def test_bedrock_guardrail_make_api_request_passes_api_key(): guardrail_hook.async_handler, "post", AsyncMock(return_value=mock_response) ), patch.object(guardrail_hook, "_load_credentials") as mock_load_creds, - patch.object(guardrail_hook, "convert_to_bedrock_format") as mock_convert, + patch.object( + guardrail_hook, "convert_to_bedrock_format", new_callable=AsyncMock + ) as mock_convert, patch.object( guardrail_hook, "get_guardrail_dynamic_request_body_params" ) as mock_get_params, diff --git a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py index fc2fb949143..37a75c77dae 100644 --- a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py @@ -1,3 +1,4 @@ +import asyncio import json from unittest.mock import MagicMock, patch @@ -113,7 +114,9 @@ def test_initialize_bedrock_forwards_contextual_grounding_from_messages(): if isinstance(callback, BedrockGuardrail) and callback.guardrail_name == "test_bedrock_grounding_from_messages" ] assert initialized, "bedrock guardrail was not registered as a callback" - actual_request = initialized[-1].convert_to_bedrock_format(source="OUTPUT", response=response, messages=messages) + actual_request = asyncio.run( + initialized[-1].convert_to_bedrock_format(source="OUTPUT", response=response, messages=messages) + ) assert json.loads(json.dumps(actual_request)) == expected_request