mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(guardrails): classify every attachment shape the bedrock guardrail can receive
This commit is contained in:
parent
e12921ce23
commit
5a1ac6c39c
2 changed files with 38 additions and 19 deletions
|
|
@ -51,11 +51,12 @@ _IMAGE_FORMAT_BY_MIME: Final[Mapping[str, BedrockImageFormat]] = MappingProxyTyp
|
|||
)
|
||||
_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"})
|
||||
_OPENAI_IMAGE_TYPES: Final = frozenset({"image_url", "input_image", "computer_screenshot"})
|
||||
_OPENAI_UNSCANNABLE_TYPES: Final = frozenset(
|
||||
{"file", "input_file", "input_audio", "video_url", "audio_url", "document", "container_upload"}
|
||||
)
|
||||
_ANTHROPIC_UNSCANNABLE_TYPES: Final = frozenset({"document", "container_upload"})
|
||||
_RESPONSES_IMAGE_TYPES: Final = frozenset({"input_image", "computer_screenshot"})
|
||||
_RESPONSES_UNSCANNABLE_TYPES: Final = frozenset({"input_file", "input_audio"})
|
||||
_CONVERSE_UNSCANNABLE_KEYS: Final = ("document", "video")
|
||||
_CONVERSE_UNSCANNABLE_KEYS: Final = ("document", "video", "audio")
|
||||
_MAX_IMAGE_BYTES: Final = 4 * 1024 * 1024
|
||||
_CONVERSE_ACTIONS: Final = frozenset({"converse", "converse-stream"})
|
||||
_TOOL_ROLES: Final = frozenset({"tool", "function"})
|
||||
|
|
@ -90,11 +91,11 @@ 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
|
||||
return _mappings(data.get("messages")), _classify_openai_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
|
||||
return _mappings(data.get("input")), _classify_openai_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 _is_mapping(body) else None)
|
||||
|
|
@ -166,13 +167,15 @@ def _classify_nothing(block: Mapping[str, object]) -> _Classified:
|
|||
return None
|
||||
|
||||
|
||||
def _classify_chat_block(block: Mapping[str, object]) -> _Classified:
|
||||
def _classify_openai_block(block: Mapping[str, object]) -> _Classified:
|
||||
block_type: Final = block.get("type")
|
||||
if block_type == "image_url":
|
||||
if block_type == "image":
|
||||
return _classify_anthropic_block(block)
|
||||
if isinstance(block_type, str) and block_type in _OPENAI_IMAGE_TYPES:
|
||||
image_url: Final = block.get("image_url")
|
||||
url: Final = image_url.get("url") if _is_mapping(image_url) else image_url
|
||||
return _classify_data_uri(url, "image_url")
|
||||
if isinstance(block_type, str) and block_type in _CHAT_UNSCANNABLE_TYPES:
|
||||
return _classify_data_uri(url, block_type)
|
||||
if isinstance(block_type, str) and block_type in _OPENAI_UNSCANNABLE_TYPES:
|
||||
return _Unscannable(block_type)
|
||||
return None
|
||||
|
||||
|
|
@ -189,15 +192,6 @@ def _classify_anthropic_block(block: Mapping[str, object]) -> _Classified:
|
|||
return None
|
||||
|
||||
|
||||
def _classify_responses_block(block: Mapping[str, object]) -> _Classified:
|
||||
block_type: Final = block.get("type")
|
||||
if isinstance(block_type, str) and block_type in _RESPONSES_IMAGE_TYPES:
|
||||
return _classify_data_uri(block.get("image_url"), block_type)
|
||||
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 _is_mapping(image):
|
||||
|
|
|
|||
|
|
@ -105,6 +105,31 @@ TEXT = {"type": "text", "text": "hello"}
|
|||
["image_url (invalid base64)"],
|
||||
id="chat-empty-base64",
|
||||
),
|
||||
pytest.param(
|
||||
_chat(
|
||||
{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": PNG_B64}},
|
||||
{"type": "input_image", "image_url": f"data:image/jpeg;base64,{JPEG_B64}"},
|
||||
{"type": "input_file", "file_id": "file-1"},
|
||||
),
|
||||
CallTypes.acompletion.value,
|
||||
[_png_item(), _jpeg_item()],
|
||||
["input_file"],
|
||||
id="chat-anthropic-and-responses-shaped-parts",
|
||||
),
|
||||
pytest.param(
|
||||
{"input": [{"role": "user", "content": [{"type": "image_url", "image_url": {"url": "https://x/y.png"}}]}]},
|
||||
CallTypes.aresponses.value,
|
||||
[],
|
||||
["image_url (remote URL or file id)"],
|
||||
id="responses-chat-shaped-image-url",
|
||||
),
|
||||
pytest.param(
|
||||
_converse({"audio": {"format": "mp3", "source": {"bytes": PDF_B64}}}),
|
||||
CallTypes.allm_passthrough_route.value,
|
||||
[],
|
||||
["audio"],
|
||||
id="converse-audio",
|
||||
),
|
||||
pytest.param(
|
||||
_chat(TEXT, {"type": "file", "file": {"file_data": f"data:application/pdf;base64,{PDF_B64}"}}),
|
||||
CallTypes.acompletion.value,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue