mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): scan text document content with the Bedrock guardrail
This commit is contained in:
parent
561214afff
commit
69fb98540f
4 changed files with 132 additions and 6 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue