mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(guardrails): keep base behavior for Bedrock subclasses and malformed messages
- skip the attachment scan when a subclass overrides apply_guardrail - ignore non-string message role/type instead of raising - treat a null Converse document/video/audio block as absent - keep the base request-body debug log for text-only calls
This commit is contained in:
parent
eeda4cb61d
commit
73431b0340
4 changed files with 94 additions and 5 deletions
|
|
@ -131,11 +131,13 @@ def _latest_user_message(messages: Sequence[Mapping[str, object]]) -> tuple[Mapp
|
||||||
|
|
||||||
|
|
||||||
def _message_blocks(message: Mapping[str, object], nested_tool_blocks: _NestedToolBlocks) -> tuple[_Block, ...]:
|
def _message_blocks(message: Mapping[str, object], nested_tool_blocks: _NestedToolBlocks) -> tuple[_Block, ...]:
|
||||||
if message.get("type") in _TOOL_OUTPUT_ITEM_TYPES:
|
message_type: Final = message.get("type")
|
||||||
|
if isinstance(message_type, str) and message_type in _TOOL_OUTPUT_ITEM_TYPES:
|
||||||
output: Final = message.get("output")
|
output: Final = message.get("output")
|
||||||
blocks: Final = (output,) if _is_mapping(output) else _mappings(output)
|
blocks: Final = (output,) if _is_mapping(output) else _mappings(output)
|
||||||
return tuple(_Block(block, from_tool=True) for block in blocks)
|
return tuple(_Block(block, from_tool=True) for block in blocks)
|
||||||
from_tool_message: Final = message.get("role") in _TOOL_ROLES
|
role: Final = message.get("role")
|
||||||
|
from_tool_message: Final = isinstance(role, str) and role in _TOOL_ROLES
|
||||||
return tuple(
|
return tuple(
|
||||||
entry
|
entry
|
||||||
for block in _mappings(message.get("content"))
|
for block in _mappings(message.get("content"))
|
||||||
|
|
@ -203,7 +205,7 @@ def _classify_converse_block(block: Mapping[str, object]) -> _Classified:
|
||||||
mime: Final = f"image/{image_format}" if isinstance(image_format, str) else None
|
mime: Final = f"image/{image_format}" if isinstance(image_format, str) else None
|
||||||
return _classify_base64(mime, encoded, "image")
|
return _classify_base64(mime, encoded, "image")
|
||||||
for key in _CONVERSE_UNSCANNABLE_KEYS:
|
for key in _CONVERSE_UNSCANNABLE_KEYS:
|
||||||
if key in block:
|
if block.get(key) is not None:
|
||||||
return _Unscannable(key)
|
return _Unscannable(key)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1279,7 +1279,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
||||||
headers_dict: Final = dict(prepared_request.headers) # mutable-ok: the masking helper requires a dict
|
headers_dict: Final = dict(prepared_request.headers) # mutable-ok: the masking helper requires a dict
|
||||||
verbose_proxy_logger.debug(
|
verbose_proxy_logger.debug(
|
||||||
"Bedrock AI request body: %s, url %s, headers: %s",
|
"Bedrock AI request body: %s, url %s, headers: %s",
|
||||||
{**bedrock_request_data, "content": _without_image_bytes(content)},
|
{**bedrock_request_data, "content": _without_image_bytes(content)}
|
||||||
|
if any("image" in item for item in content)
|
||||||
|
else bedrock_request_data,
|
||||||
prepared_request.url,
|
prepared_request.url,
|
||||||
_get_masked_values(headers_dict),
|
_get_masked_values(headers_dict),
|
||||||
)
|
)
|
||||||
|
|
@ -2564,8 +2566,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
||||||
Documents, files, audio, video, and images sent by URL, by file id or in another format block the
|
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
|
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
|
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.
|
call raises the same error the text scan of the same hook raises. A subclass that overrides
|
||||||
|
``apply_guardrail`` skips this scan.
|
||||||
"""
|
"""
|
||||||
|
if type(self).apply_guardrail is not BedrockGuardrail.apply_guardrail:
|
||||||
|
return
|
||||||
attachments: Final = find_request_attachments(
|
attachments: Final = find_request_attachments(
|
||||||
data,
|
data,
|
||||||
call_type,
|
call_type,
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,7 @@ Unit tests for Bedrock Guardrails
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import logging
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
import sys
|
import sys
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
@ -14,6 +15,7 @@ from fastapi import HTTPException
|
||||||
|
|
||||||
|
|
||||||
import litellm
|
import litellm
|
||||||
|
from litellm._logging import verbose_proxy_logger
|
||||||
from litellm.caching.caching import DualCache
|
from litellm.caching.caching import DualCache
|
||||||
from litellm.exceptions import ModifyResponseException
|
from litellm.exceptions import ModifyResponseException
|
||||||
from litellm.proxy._types import UserAPIKeyAuth
|
from litellm.proxy._types import UserAPIKeyAuth
|
||||||
|
|
@ -6374,3 +6376,59 @@ async def test_attachment_scan_debug_log_omits_image_bytes():
|
||||||
logged = " ".join(str(call.args) for call in mock_debug.call_args_list)
|
logged = " ".join(str(call.args) for call in mock_debug.call_args_list)
|
||||||
assert _ATTACHMENT_PNG_B64 not in logged
|
assert _ATTACHMENT_PNG_B64 not in logged
|
||||||
assert f"<{len(_ATTACHMENT_PNG_B64)} base64 chars>" in logged
|
assert f"<{len(_ATTACHMENT_PNG_B64)} base64 chars>" in logged
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_text_only_debug_log_prints_the_signed_request_body():
|
||||||
|
guardrail = _attachment_guardrail()
|
||||||
|
post_patch, credentials_patch, prepare_patch = _patched_bedrock_post(
|
||||||
|
guardrail, _passing_bedrock_httpx_response("ok")
|
||||||
|
)
|
||||||
|
captured_records: list[logging.LogRecord] = []
|
||||||
|
|
||||||
|
class _RecordingHandler(logging.Handler):
|
||||||
|
def emit(self, record: logging.LogRecord) -> None:
|
||||||
|
captured_records.append(record)
|
||||||
|
|
||||||
|
handler = _RecordingHandler(level=logging.DEBUG)
|
||||||
|
previous_level = verbose_proxy_logger.level
|
||||||
|
verbose_proxy_logger.addHandler(handler)
|
||||||
|
verbose_proxy_logger.setLevel(logging.DEBUG)
|
||||||
|
try:
|
||||||
|
with post_patch, credentials_patch, prepare_patch:
|
||||||
|
await guardrail.make_bedrock_api_request(
|
||||||
|
source="INPUT", messages=[{"role": "user", "content": "hello"}], request_data={}
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
verbose_proxy_logger.removeHandler(handler)
|
||||||
|
verbose_proxy_logger.setLevel(previous_level)
|
||||||
|
|
||||||
|
body_lines = [
|
||||||
|
record.getMessage() for record in captured_records if record.getMessage().startswith("Bedrock AI request body")
|
||||||
|
]
|
||||||
|
assert len(body_lines) == 1
|
||||||
|
assert "'content': ({'text': {'text': 'hello'}},)" in body_lines[0]
|
||||||
|
|
||||||
|
|
||||||
|
class _CustomApplyGuardrail(BedrockGuardrail):
|
||||||
|
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
|
||||||
|
return inputs
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_attachment_scan_skipped_when_subclass_overrides_apply_guardrail():
|
||||||
|
guardrail = _CustomApplyGuardrail(
|
||||||
|
guardrail_name="bedrock-attachments", guardrailIdentifier="gid", guardrailVersion="DRAFT"
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
|
||||||
|
pdf_result = await guardrail.async_scan_request_attachments(
|
||||||
|
data=_pdf_chat_request(), call_type=CallTypes.acompletion.value
|
||||||
|
)
|
||||||
|
image_result = await guardrail.async_scan_request_attachments(
|
||||||
|
data=_image_only_chat_request(), call_type=CallTypes.acompletion.value
|
||||||
|
)
|
||||||
|
|
||||||
|
assert pdf_result is None
|
||||||
|
assert image_result is None
|
||||||
|
assert mock_post.await_count == 0
|
||||||
|
|
|
||||||
|
|
@ -373,3 +373,27 @@ def test_latest_user_message_only():
|
||||||
|
|
||||||
assert list(found.images) == [_png_item()]
|
assert list(found.images) == [_png_item()]
|
||||||
assert found.unscannable == ()
|
assert found.unscannable == ()
|
||||||
|
|
||||||
|
|
||||||
|
def test_non_string_role_and_type_do_not_raise():
|
||||||
|
data = {
|
||||||
|
"messages": [
|
||||||
|
{"role": "user", "content": [{"type": "image_url", "image_url": f"data:image/png;base64,{PNG_B64}"}]},
|
||||||
|
{"role": ["tool"], "type": {"x": 1}, "content": [{"type": "file", "file": {"file_id": "f"}}]},
|
||||||
|
{"role": {"r": "tool"}, "type": ["function_call_output"], "content": [{"type": "file", "file": {}}]},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
found = find_request_attachments(data, CallTypes.acompletion.value, True, False)
|
||||||
|
|
||||||
|
assert list(found.images) == [_png_item()]
|
||||||
|
assert found.unscannable == ("file", "file")
|
||||||
|
|
||||||
|
|
||||||
|
def test_converse_null_document_is_not_an_attachment():
|
||||||
|
data = _converse(_png_item(), {"document": None}, {"video": None, "audio": None}, {"audio": {"format": "mp3"}})
|
||||||
|
|
||||||
|
found = find_request_attachments(data, CallTypes.allm_passthrough_route.value, False, False)
|
||||||
|
|
||||||
|
assert list(found.images) == [_png_item()]
|
||||||
|
assert found.unscannable == ("audio",)
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue