mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
fix(guardrails): fix structured content scanning and violation content blocking
This commit is contained in:
parent
b7291c3ea4
commit
11b742a4b5
3 changed files with 42 additions and 4 deletions
|
|
@ -885,9 +885,14 @@ class AliyunAIGuardrail(AliyunGuardrailBase, CustomGuardrail):
|
|||
"""
|
||||
if target is None:
|
||||
return False
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
set_mcp_tool_result_structured_content,
|
||||
)
|
||||
|
||||
content: Final = getattr(target, "content", None)
|
||||
if isinstance(content, list):
|
||||
content[:] = blocked_content
|
||||
set_mcp_tool_result_structured_content(target, None)
|
||||
if hasattr(target, "isError"):
|
||||
try:
|
||||
# Kept duck-typed on purpose: any MCP result shape carrying `isError`
|
||||
|
|
@ -903,10 +908,12 @@ class AliyunAIGuardrail(AliyunGuardrailBase, CustomGuardrail):
|
|||
result: Final = target.get("result")
|
||||
if isinstance(result, dict) and isinstance(result.get("content"), list):
|
||||
result["content"] = list(blocked_content) # mutable-ok: the replaced field must stay a JSON list
|
||||
set_mcp_tool_result_structured_content(result, None)
|
||||
return True
|
||||
if isinstance(target.get("content"), list):
|
||||
blocked: Final = list(blocked_content) # mutable-ok: must stay a JSON list
|
||||
target["content"] = blocked # rebind-ok: overwrite the caller's result
|
||||
set_mcp_tool_result_structured_content(target, None)
|
||||
return True
|
||||
return False
|
||||
|
||||
|
|
|
|||
|
|
@ -26,6 +26,10 @@ class AliyunGuardrailBase:
|
|||
"""
|
||||
return (message for message in messages if message.get("role") == "user")
|
||||
|
||||
@staticmethod
|
||||
def _iter_audited_text_messages(messages: Sequence[AllMessageValues]) -> Iterator[AllMessageValues]:
|
||||
return (message for message in messages if message.get("role") in ("user", "tool"))
|
||||
|
||||
@staticmethod
|
||||
def _extract_image_url(part: object) -> str | None:
|
||||
"""
|
||||
|
|
@ -35,7 +39,7 @@ class AliyunGuardrailBase:
|
|||
Returns:
|
||||
The URL string, or None when the part carries no image URL
|
||||
"""
|
||||
if not isinstance(part, dict) or part.get("type") != "image_url":
|
||||
if not isinstance(part, dict) or part.get("type") not in ("image_url", "input_image"):
|
||||
return None
|
||||
image_url: Final = part.get("image_url")
|
||||
if isinstance(image_url, dict):
|
||||
|
|
@ -62,7 +66,7 @@ class AliyunGuardrailBase:
|
|||
)
|
||||
|
||||
user_prompt: Final = "\n".join(
|
||||
convert_content_list_to_str(message) for message in self._iter_user_messages(messages)
|
||||
convert_content_list_to_str(message) for message in self._iter_audited_text_messages(messages)
|
||||
).strip()
|
||||
return user_prompt or None
|
||||
|
||||
|
|
|
|||
|
|
@ -792,8 +792,8 @@ class TestPreCallHook:
|
|||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "结构化输入文本"},
|
||||
{"type": "image_url", "image_url": {"url": IMG_A}},
|
||||
{"type": "input_text", "text": "结构化输入文本"},
|
||||
{"type": "input_image", "image_url": IMG_A, "detail": "auto"},
|
||||
],
|
||||
}
|
||||
]
|
||||
|
|
@ -810,6 +810,32 @@ class TestPreCallHook:
|
|||
assert any("结构化输入文本" in (sp.get("content") or "") for sp in sent)
|
||||
assert any(IMG_A in (sp.get("imageUrls") or []) for sp in sent)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scans_responses_api_function_call_output(self):
|
||||
g = _make_guardrail(level="medium")
|
||||
clean = _make_aliyun_api_response(suggestion="pass", detail=[])
|
||||
data = {
|
||||
"input": [
|
||||
{
|
||||
"type": "function_call_output",
|
||||
"call_id": "call_1",
|
||||
"output": "工具输出里的违规内容",
|
||||
}
|
||||
]
|
||||
}
|
||||
with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=clean) as mock_post:
|
||||
await g.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=MagicMock(),
|
||||
data=data,
|
||||
call_type="responses",
|
||||
)
|
||||
|
||||
scanned = "".join(
|
||||
json.loads(call.kwargs["data"]["ServiceParameters"]).get("content", "") for call in mock_post.call_args_list
|
||||
)
|
||||
assert "工具输出里的违规内容" in scanned
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_blocks_violation(self):
|
||||
g = _make_guardrail(level="medium")
|
||||
|
|
@ -1350,6 +1376,7 @@ class TestExtractMcpToolText:
|
|||
remaining = " ".join(getattr(item, "text", "") for item in tool_result.content)
|
||||
assert CONTENT_MODERATION_TYPE in remaining
|
||||
assert tool_result.isError is True
|
||||
assert tool_result.structuredContent is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue