fix(guardrails): scan text document content with the Bedrock guardrail

This commit is contained in:
Yucheng He 2026-09-30 01:40:35 -07:00
parent 561214afff
commit 69fb98540f
4 changed files with 132 additions and 6 deletions

View file

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

View file

@ -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],

View file

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

View file

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