From 69fb98540fae60f8ff26b62323b60f4cc7edf311 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Wed, 30 Sep 2026 01:40:35 -0700 Subject: [PATCH] fix(guardrails): scan text document content with the Bedrock guardrail --- litellm/llms/bedrock/guardrail_attachments.py | 35 ++++++++-- .../guardrail_hooks/bedrock_guardrails.py | 28 +++++++- .../test_bedrock_guardrails.py | 67 +++++++++++++++++++ .../bedrock/test_guardrail_attachments.py | 8 ++- 4 files changed, 132 insertions(+), 6 deletions(-) diff --git a/litellm/llms/bedrock/guardrail_attachments.py b/litellm/llms/bedrock/guardrail_attachments.py index c63377d0ec0..a66f1d815c0 100644 --- a/litellm/llms/bedrock/guardrail_attachments.py +++ b/litellm/llms/bedrock/guardrail_attachments.py @@ -26,6 +26,7 @@ BedrockImageFormat = Literal["png", "jpeg"] class RequestAttachments(NamedTuple): images: tuple[BedrockContentItem, ...] unscannable: tuple[str, ...] + document_texts: tuple[str, ...] = () class _Image(NamedTuple): @@ -36,12 +37,17 @@ class _Unscannable(NamedTuple): label: str +class _DocumentText(NamedTuple): + text: str + + class _Block(NamedTuple): block: Mapping[str, object] from_tool: bool + in_document: bool = False -_Classified = _Image | _Unscannable | None +_Classified = _Image | _Unscannable | _DocumentText | None _BlockClassifier = Callable[[Mapping[str, object]], _Classified] _NestedToolBlocks = Callable[[Mapping[str, object]], tuple[Mapping[str, object], ...]] @@ -75,7 +81,7 @@ def find_request_attachments( 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.""" + """List the scannable images, the document text and the unscannable attachments in the message content and tool results.""" messages, classify, nested_tool_blocks = _messages_and_classifier(data, call_type) selected: Final = messages if not latest_user_message_only else _latest_user_message(messages) classified: Final = tuple( @@ -83,14 +89,32 @@ def find_request_attachments( 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 + and (result := _classify_entry(entry, classify)) 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)), + document_texts=tuple(result.text for result in classified if isinstance(result, _DocumentText)), ) +def _classify_entry(entry: _Block, classify: _BlockClassifier) -> _Classified: + text: Final = _document_text(entry) + return _DocumentText(text) if text else classify(entry.block) + + +def _document_text(entry: _Block) -> str | None: + block: Final = entry.block + if entry.in_document: + text: Final = block.get("text") if block.get("type") == "text" else None + return text if isinstance(text, str) else None + source: Final = block.get("source") if block.get("type") == "document" else None + if not _is_mapping(source): + return None + source_text: Final = source.get("data") if source.get("type") == "text" else source.get("content") + return source_text if source.get("type") in _TEXT_DOCUMENT_SOURCE_TYPES and isinstance(source_text, str) else None + + def _messages_and_classifier( data: Mapping[str, object], call_type: str ) -> tuple[Sequence[Mapping[str, object]], _BlockClassifier, _NestedToolBlocks]: @@ -161,7 +185,10 @@ def _with_document_content(entry: _Block) -> tuple[_Block, ...]: source: Final = entry.block.get("source") if entry.block.get("type") != "document" or not _is_mapping(source) or source.get("type") != "content": return (entry,) - return (entry, *(_Block(inner, from_tool=entry.from_tool) for inner in _mappings(source.get("content")))) + return ( + entry, + *(_Block(inner, from_tool=entry.from_tool, in_document=True) for inner in _mappings(source.get("content"))), + ) def _in_scope(entry: _Block, skip_tool_messages: bool, scan_only_tool_results: bool) -> bool: diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 30e2b60000d..8f2bfe58267 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -2561,7 +2561,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): 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. + """Scan the request's inline PNG/JPEG images and text documents, and block attachments it cannot scan. + + Images go to ApplyGuardrail 20 per call. Text documents go through ``make_bedrock_api_request`` + as one user turn, and any intervention on them blocks the request, since their text cannot be + rewritten with a mask. Documents, files, audio, video, and images sent by URL, by file id or in another format block the request unless ``skip_unscannable_attachments`` is set. ``checks`` mode calls the text-only @@ -2593,6 +2597,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): self.guardrail_name, len(unscannable), ) + if attachments.document_texts: + await self._scan_document_texts(attachments.document_texts, request_data=data, event_type=event_type) if not images: return bedrock_request_data, api_key = self._bedrock_request_body(BedrockRequest(source="INPUT"), data) @@ -2615,6 +2621,26 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): verbose_proxy_logger.error("Bedrock Guardrail: Failed to scan attachments: %s", str(e)) raise Exception(f"Bedrock guardrail failed: {e}") from e + async def _scan_document_texts( + self, + document_texts: Sequence[str], + request_data: dict, # mutable-ok: proxy request body dict, mutated by the logging helper + event_type: GuardrailEventHooks, + ) -> None: + """Scan text documents like a user turn and block when the guardrail intervenes on them.""" + response: Final = await self.make_bedrock_api_request( + source="INPUT", + messages=[{"role": "user", "content": text} for text in document_texts], + request_data=request_data, + logging_event_type=event_type, + ) + if response.get("action") == "GUARDRAIL_INTERVENED": + raise self._unscannable_attachments_exception( + tuple("document (text the guardrail would mask)" for _ in document_texts), + request_data=request_data, + event_type=event_type, + ) + def _unscannable_attachments_exception( self, unscannable: Sequence[str], 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 3f8e85dcaff..ea3a0195bbe 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 @@ -6410,6 +6410,73 @@ async def test_text_only_debug_log_prints_the_signed_request_body(): assert "'content': ({'text': {'text': 'hello'}},)" in body_lines[0] +def _text_document_messages_request(text: str) -> dict: + return { + "model": "claude", + "messages": [ + { + "role": "user", + "content": [ + {"type": "document", "source": {"type": "text", "media_type": "text/plain", "data": text}}, + {"type": "text", "text": "summarize"}, + ], + } + ], + } + + +def _anonymizing_bedrock_httpx_response() -> MagicMock: + response = MagicMock() + response.status_code = 200 + response.json.return_value = { + "action": "GUARDRAIL_INTERVENED", + "outputs": [{"text": "SSN {US_SOCIAL_SECURITY_NUMBER}"}], + "assessments": [ + { + "sensitiveInformationPolicy": { + "piiEntities": [{"type": "US_SOCIAL_SECURITY_NUMBER", "action": "ANONYMIZED"}] + } + } + ], + "usage": {"contentPolicyUnits": 1}, + } + return response + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "response, blocked", + [ + pytest.param(_passing_bedrock_httpx_response("ok"), False, id="pass"), + pytest.param(_blocking_bedrock_httpx_response("deny"), True, id="block"), + pytest.param(_anonymizing_bedrock_httpx_response(), True, id="anonymize"), + ], +) +async def test_attachment_scan_sends_text_document_as_text(response, blocked): + guardrail = _attachment_guardrail() + post_patch, credentials_patch, prepare_patch = _patched_bedrock_post(guardrail, response) + + with post_patch as mock_post, credentials_patch, prepare_patch as mock_prepare: + if blocked: + with pytest.raises(HTTPException) as exc_info: + await guardrail.async_scan_request_attachments( + data=_text_document_messages_request("SSN 123-45-6789"), + call_type=CallTypes.anthropic_messages.value, + ) + assert exc_info.value.status_code == 400 + else: + result = await guardrail.async_scan_request_attachments( + data=_text_document_messages_request("SSN 123-45-6789"), + call_type=CallTypes.anthropic_messages.value, + ) + assert result is None + + assert mock_post.await_count == 1 + sent_body = mock_prepare.call_args.kwargs["data"] + assert sent_body["source"] == "INPUT" + assert sent_body["content"] == ({"text": {"text": "SSN 123-45-6789"}},) + + class _CustomApplyGuardrail(BedrockGuardrail): async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): return inputs diff --git a/tests/unit/llms/bedrock/test_guardrail_attachments.py b/tests/unit/llms/bedrock/test_guardrail_attachments.py index 9d6245a9714..c2657ce7077 100644 --- a/tests/unit/llms/bedrock/test_guardrail_attachments.py +++ b/tests/unit/llms/bedrock/test_guardrail_attachments.py @@ -405,7 +405,11 @@ def test_converse_null_document_is_not_an_attachment(): pytest.param( {"type": "document", "source": {"type": "text", "media_type": "text/plain", "data": "hi"}}, id="text" ), - pytest.param({"type": "document", "source": {"type": "content", "content": [TEXT]}}, id="content"), + pytest.param( + {"type": "document", "source": {"type": "content", "content": [{"type": "text", "text": "hi"}]}}, + id="content", + ), + pytest.param({"type": "document", "source": {"type": "content", "content": "hi"}}, id="content-string"), ], ) @pytest.mark.parametrize( @@ -418,6 +422,7 @@ def test_text_source_document_is_not_an_attachment(block, call_type): assert found.images == () assert found.unscannable == ("document",) + assert found.document_texts == ("hi",) def test_unpadded_and_url_safe_base64_are_sent_as_standard_base64(): @@ -461,6 +466,7 @@ def test_images_inside_a_content_source_document_are_scanned(call_type): assert list(found.images) == [_png_item()] assert found.unscannable == ("document",) + assert found.document_texts == ("hello",) def test_document_images_inside_a_tool_result_follow_the_tool_scope():