From 11b742a4b5a25fc1572d7d4b434bacfc266ac344 Mon Sep 17 00:00:00 2001 From: splendor023 Date: Wed, 19 Aug 2026 20:36:17 +0800 Subject: [PATCH] fix(guardrails): fix structured content scanning and violation content blocking --- .../aliyun/aliyun_ai_guardrail.py | 7 +++++ .../guardrails/guardrail_hooks/aliyun/base.py | 8 +++-- .../aliyun/test_aliyun_ai_guardrail.py | 31 +++++++++++++++++-- 3 files changed, 42 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/aliyun/aliyun_ai_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/aliyun/aliyun_ai_guardrail.py index 27ce33e4ca7..894ad3694ac 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aliyun/aliyun_ai_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aliyun/aliyun_ai_guardrail.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/aliyun/base.py b/litellm/proxy/guardrails/guardrail_hooks/aliyun/base.py index 452e8638b8a..1d1f6073511 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aliyun/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aliyun/base.py @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/aliyun/test_aliyun_ai_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/aliyun/test_aliyun_ai_guardrail.py index 3c28632cbd8..8b85b37d63a 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/aliyun/test_aliyun_ai_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/aliyun/test_aliyun_ai_guardrail.py @@ -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 # ---------------------------------------------------------------------------