fix(guardrails): fix structured content scanning and violation content blocking

This commit is contained in:
splendor023 2026-08-19 20:36:17 +08:00
parent b7291c3ea4
commit 11b742a4b5
3 changed files with 42 additions and 4 deletions

View file

@ -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

View file

@ -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

View file

@ -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
# ---------------------------------------------------------------------------