mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(guardrails): scan or refuse attachments in the bedrock guardrail
This commit is contained in:
parent
69d2a3c24f
commit
567a1f78e6
8 changed files with 1004 additions and 17 deletions
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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 == ()
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue