diff --git a/litellm/proxy/guardrails/guardrail_hooks/aliyun/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/aliyun/__init__.py new file mode 100644 index 00000000000..1fa0e880584 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/aliyun/__init__.py @@ -0,0 +1,103 @@ +""" +Aliyun AI Security Guardrail Integration for LiteLLM +阿里云AI安全护栏集成 +This module provides integration with Aliyun's AI Security Guardrail service for: +- ContentModeration 内容合规检测 +- PromptAttack 提示词攻击检测 +- SensitiveData 敏感内容检测 +- ModelHallucination 模型幻觉 +- MaliciousUrl 恶意URL检测 +... +Documentation: https://help.aliyun.com/document_detail/2873209.html +""" + +from typing import TYPE_CHECKING, Final + +from litellm.types.guardrails import SupportedGuardrailIntegrations + +from .aliyun_ai_guardrail import AliyunAIGuardrail + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def _resolve_os_environ_reference(value: str | None) -> str | None: + """Resolve an ``os.environ/`` reference. + + guardrail_registry.py only auto-resolves api_key/api_base, so the Aliyun + credential fields have to be resolved here. + """ + from litellm.secret_managers.main import get_secret_str + + if isinstance(value, str) and value.startswith("os.environ/"): + return get_secret_str(value) + return value + + +def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail") -> AliyunAIGuardrail: + """ + Initialize an Aliyun AI Guardrail instance. + Credentials are configured in config.yaml (litellm_params) and support + os.environ/ references: + - access_key_id: Aliyun Access Key ID + - access_key_secret: Aliyun Access Key Secret + Args: + litellm_params: The LiteLLM parameters for the guardrail + guardrail: The guardrail configuration + Returns: + AliyunAIGuardrail instance + """ + import litellm + + guardrail_name: Final = guardrail.get("guardrail_name") + if not guardrail_name: + raise ValueError("Aliyun AI Guardrail: guardrail_name is required") + + level: Final = getattr(litellm_params, "level", None) + max_text_length: Final = getattr(litellm_params, "max_text_length", None) + stream_window_size: Final = getattr(litellm_params, "stream_window_size", None) + stream_slide_step: Final = getattr(litellm_params, "stream_slide_step", None) + stream_first_check_step: Final = getattr(litellm_params, "stream_first_check_step", None) + region_id: Final = getattr(litellm_params, "region_id", None) + service_input: Final = getattr(litellm_params, "service_input", None) + service_output: Final = getattr(litellm_params, "service_output", None) + service_mcp: Final = getattr(litellm_params, "service_mcp", None) + + # These custom credential fields are not auto-resolved by guardrail_registry.py + # (only api_key/api_base are), so os.environ/ references are resolved here. + access_key_id: Final = _resolve_os_environ_reference(getattr(litellm_params, "access_key_id", None)) + access_key_secret: Final = _resolve_os_environ_reference(getattr(litellm_params, "access_key_secret", None)) + + aliyun_guardrail: Final = AliyunAIGuardrail( + guardrail_name=guardrail_name, + access_key_id=access_key_id, + access_key_secret=access_key_secret, + level=level, + max_text_length=max_text_length, + stream_window_size=stream_window_size, + stream_slide_step=stream_slide_step, + stream_first_check_step=stream_first_check_step, + region_id=region_id, + service_input=service_input, + service_output=service_output, + service_mcp=service_mcp, + default_on=litellm_params.default_on, + event_hook=litellm_params.mode, + ) + + litellm.logging_callback_manager.add_litellm_callback(aliyun_guardrail) + + return aliyun_guardrail + + +# Registry for guardrail initializers. +# Plain dicts: guardrail_registry.py gates discovery on `isinstance(registry, dict)`, +# which a MappingProxyType would fail, silently skipping this guardrail's registration. +guardrail_initializer_registry: Final = { # mutable-ok: loader requires a real dict + SupportedGuardrailIntegrations.ALIYUN_AI_GUARDRAIL.value: initialize_guardrail, +} + +# Registry for guardrail classes +guardrail_class_registry: Final = { # mutable-ok: loader requires a real dict + SupportedGuardrailIntegrations.ALIYUN_AI_GUARDRAIL.value: AliyunAIGuardrail, +} diff --git a/litellm/proxy/guardrails/guardrail_hooks/aliyun/aliyun_ai_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/aliyun/aliyun_ai_guardrail.py new file mode 100644 index 00000000000..7a3742d50e7 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/aliyun/aliyun_ai_guardrail.py @@ -0,0 +1,1169 @@ +""" +Aliyun AI Security Guardrail Integration for LiteLLM +阿里云AI安全护栏集成 +This guardrail scans prompts and responses using the Aliyun AI Security Guardrail API to detect: +- Content moderation violations +- Sensitive data (PII) +- Prompt injection attacks +- Malicious URLs +Documentation: https://help.aliyun.com/document_detail/2875413.html +Credentials: +Configured in config.yaml (litellm_params), support os.environ/ references: +- access_key_id: Aliyun Access Key ID +- access_key_secret: Aliyun Access Key Secret +""" + +from __future__ import annotations + +import asyncio +import base64 +import hashlib +import hmac +import json +import re +import uuid +from collections.abc import AsyncGenerator, AsyncIterable, Iterable, Iterator, Mapping, Sequence +from datetime import datetime, timezone +from types import MappingProxyType + +# `Any` is needed for the heterogeneous LLM payloads this guardrail walks (request bodies, +# streaming chunks, MCP tool results); every field it actually reads is probed defensively. +# `cast` re-labels the upstream message iterator, whose element type is untyped upstream. +from typing import TYPE_CHECKING, Any, Final, Literal, cast # noqa: TID251 # see comment above +from urllib.parse import quote + +from fastapi import HTTPException + +from litellm._logging import verbose_proxy_logger +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) +from litellm.proxy._types import UserAPIKeyAuth + +# Only the private iterator walks `messages` *and* the Responses API's `input` while keeping +# multimodal parts intact; the public helpers flatten to text, which would drop image URLs. +from litellm.proxy.guardrails._content_utils import ( + _iter_inspection_messages, # pyright: ignore[reportPrivateUsage] # see comment above +) +from litellm.types.utils import CallTypes + +from .base import AliyunGuardrailBase + +if TYPE_CHECKING: + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.llms.openai import AllMessageValues + from litellm.types.mcp import MCPPostCallResponseObject + from litellm.types.proxy.guardrails.guardrail_hooks.aliyun.aliyun_ai_guardrail import ( + AliyunAIGuardrailResponse, + ) + from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel + from litellm.types.utils import CallTypesLiteral, LLMResponseTypes + +# Constants +ENCODING: Final = "UTF-8" +ISO8601_DATE_FORMAT: Final = "%Y-%m-%dT%H:%M:%SZ" +ALGORITHM: Final = "HmacSHA1" + +# Region to endpoint mapping +REGION_ENDPOINTS: Final = MappingProxyType( + { + "cn-shanghai": "green-cip.cn-shanghai.aliyuncs.com", + "cn-beijing": "green-cip.cn-beijing.aliyuncs.com", + "cn-hangzhou": "green-cip.cn-hangzhou.aliyuncs.com", + "cn-shenzhen": "green-cip.cn-shenzhen.aliyuncs.com", + "cn-chengdu": "green-cip.cn-chengdu.aliyuncs.com", + "ap-southeast-1": "green-cip.ap-southeast-1.aliyuncs.com", + "eu-central-1": "green-cip.eu-central-1.aliyuncs.com", + } +) + +# Service codes for domestic (China) regions +SERVICE_INPUT_DOMESTIC: Final = "query_security_check_pro" +SERVICE_OUTPUT_DOMESTIC: Final = "response_security_check_pro" + +# Service codes for international regions +SERVICE_INPUT_INTERNATIONAL: Final = "query_security_check_cb" +SERVICE_OUTPUT_INTERNATIONAL: Final = "response_security_check_cb" + + +# Detection types +CONTENT_MODERATION_TYPE: Final = "contentModeration" +PROMPT_ATTACK_TYPE: Final = "promptAttack" +SENSITIVE_DATA_TYPE: Final = "sensitiveData" +MALICIOUS_URL_TYPE: Final = "maliciousUrl" +MODEL_HALLUCINATION_TYPE: Final = "modelHallucination" +CUSTOM_LABEL_TYPE: Final = "customLabel" + +# Suggestion returned by Aliyun when it has decided the content must be rejected +BLOCK_SUGGESTION: Final = "block" + +# An explicit upstream block that carries no parseable severity is treated as the most +# severe level, so it is still weighed against the configured threshold rather than +# being silently downgraded to "none". +UNRESOLVED_BLOCK_LEVEL: Final = "high" + +# Risk level to integer, for threshold comparison. Covers both the standard levels +# (none/low/medium/high) and the sensitiveData levels (S0-S4). +RISK_LEVEL_TO_INT: Final = MappingProxyType( + { + # Standard risk levels + "none": 0, + "low": 1, + "medium": 2, + "high": 3, + # Sensitive data levels (mapped to standard levels) + "s0": 0, # No risk + "s1": 1, # Low risk + "s2": 2, # Medium risk + "s3": 3, # High risk + "s4": 3, # High risk (highest sensitive level) + } +) + + +def level_to_int(risk_level: str) -> int: + """ + Convert risk level string to integer for comparison. + Higher value = higher risk. + Supports both standard risk levels (none/low/medium/high) + and sensitive data levels (S0/S1/S2/S3/S4). + Mapping: + - none/S0 = 0 (no risk) + - low/S1 = 1 (low risk) + - medium/S2 = 2 (medium risk) + - high/S3/S4 = 3 (high risk) + """ + level_lower: Final = risk_level.lower() if risk_level else "none" + return RISK_LEVEL_TO_INT.get(level_lower, 0) + + +# Sentence-ending punctuation, used to cut long text on a boundary instead of mid-sentence +SENTENCE_BOUNDARY_PATTERN: Final = r"[。!?;:\.?!]+" + +# Parsed result of a response Aliyun did not flag at all +PASS_RESULT: Final = MappingProxyType( + {"flagged": False, "suggestion": "pass", "details": MappingProxyType({}), "message": ""} +) + +# Protection level thresholds +# If detected_level >= threshold, then block +PROTECTION_LEVEL_THRESHOLD: Final = MappingProxyType( + { + "low": 1, # High protection: block low, medium, high (threshold=1, block if >=1) + "medium": 2, # Medium protection: block medium, high (threshold=2, block if >=2) + "high": 3, # Low protection: block high only (threshold=3, block if >=3) + "max": 99, # Observation mode: never block (threshold very high) + } +) + + +class AliyunAIGuardrail(AliyunGuardrailBase, CustomGuardrail): + """ + LiteLLM Built-in Guardrail for Aliyun AI Security Guardrail. + This guardrail scans prompts and responses using the Aliyun AI Security Guardrail API to detect + malicious content, injection attempts, sensitive data, and policy violations. + Configuration: + guardrail_name: Name of the guardrail instance + access_key_id: Aliyun Access Key ID + access_key_secret: Aliyun Access Key Secret + region_id: Aliyun region ID (default: cn-shanghai) + default_on: Whether to enable by default + """ + + def __init__( + self, + guardrail_name: str, + access_key_id: str | None = None, + access_key_secret: str | None = None, + region_id: str | None = None, + level: str | None = None, + max_text_length: int | None = None, + stream_window_size: int | None = None, + stream_slide_step: int | None = None, + stream_first_check_step: int | None = None, + service_input: str | None = None, + service_output: str | None = None, + service_mcp: str | None = None, + **kwargs: object, # kwargs-ok: forwarded verbatim to CustomGuardrail's own **kwargs contract + ) -> None: + """ + Initialize Aliyun AI Guardrail handler. + Credentials (access_key_id / access_key_secret) are passed in from config.yaml + via the guardrail loader. + Args: + region_id: Aliyun region ID (default: cn-shanghai) + level: Protection level for risk filtering + - "low": High protection, block all risks (low, medium, high, S1+) + - "medium": Medium protection, block medium and high risks (medium, high, S2+) + - "high": Low protection, block only high risks (high, S3+) + - "max": Observation mode, no blocking + service_input: Service code for input detection (default: query_security_check_pro) + service_output: Service code for output detection (default: response_security_check_pro) + service_mcp: Service code for MCP tool call detection, used by both + pre_mcp_call and post_mcp_call (default: query_security_check_pro) + """ + super().__init__( + guardrail_name=guardrail_name, + **kwargs, + ) + self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) + self.access_key_id = access_key_id or "" + self.access_key_secret = access_key_secret or "" + self.region_id = region_id or "cn-shanghai" + if not self.access_key_id: + raise ValueError( + "Aliyun AI Guardrail: ak is required. Set access_key_id in config.yaml (supports os.environ/ reference)" + ) + if not self.access_key_secret: + raise ValueError( + "Aliyun AI Guardrail: sk is required. " + "Set access_key_secret in config.yaml (supports os.environ/ reference)" + ) + self.level = level or "medium" + if self.level not in PROTECTION_LEVEL_THRESHOLD: + raise ValueError( + f"Aliyun AI Guardrail: Invalid level '{self.level}'. " + f"Valid values are: {tuple(PROTECTION_LEVEL_THRESHOLD)}" + ) + self.max_text_length = max_text_length or 2000 + self.endpoint = REGION_ENDPOINTS.get(self.region_id, REGION_ENDPOINTS["cn-shanghai"]) + self.service_url = f"https://{self.endpoint}" + self.service_input = service_input or SERVICE_INPUT_DOMESTIC + self.service_output = service_output or SERVICE_OUTPUT_DOMESTIC + self.service_mcp = service_mcp or SERVICE_INPUT_DOMESTIC + self.stream_window_size = stream_window_size or 500 + self.stream_slide_step = stream_slide_step or 300 + self.stream_first_check_step = stream_first_check_step or 50 + verbose_proxy_logger.info( + "Initialized Aliyun AI Security Guardrail: %s, " + "region: %s, level: %s, " + "service_input: %s, service_output: %s, " + "service_mcp: %s", + guardrail_name, + self.region_id, + self.level, + self.service_input, + self.service_output, + self.service_mcp, + ) + + @staticmethod + def get_config_model() -> type[GuardrailConfigModel] | None: + from litellm.types.proxy.guardrails.guardrail_hooks.aliyun.aliyun_ai_guardrail import ( + AliyunAIGuardrailConfigModel, + ) + + return AliyunAIGuardrailConfigModel + + @staticmethod + def _format_iso8601_date() -> str: + """Format current timestamp in ISO8601 format.""" + return datetime.now(timezone.utc).strftime(ISO8601_DATE_FORMAT) + + @staticmethod + def _percent_encode(value: str | None) -> str: + """URL encode a value according to Aliyun signature requirements.""" + if value is None: + return "" + return quote(value.encode(ENCODING), safe="~").replace("+", "%20").replace("*", "%2A") + + def _create_signature(self, string_to_sign: str) -> str: + """Create HMAC-SHA1 signature for API request.""" + secret: Final = self.access_key_secret + "&" + signature: Final = hmac.new( + secret.encode(ENCODING), + string_to_sign.encode(ENCODING), + hashlib.sha1, + ).digest() + return base64.b64encode(signature).decode(ENCODING) + + def _create_string_to_sign(self, http_method: str, parameters: Mapping[str, str]) -> str: + """Create the string to sign for API request.""" + canonicalized_query_string: Final = "".join( + f"&{self._percent_encode(key)}={self._percent_encode(parameters[key])}" for key in sorted(parameters) + ) + return ( + http_method + "&" + self._percent_encode("/") + "&" + self._percent_encode(canonicalized_query_string[1:]) + ) + + @staticmethod + def _iter_text_segments(text: str, max_length: int) -> Iterator[str]: + """ + Yield ``text`` in segments of at most ``max_length`` characters. + Each cut lands on the last sentence boundary inside the window, so a + segment does not split a sentence in half when it can be avoided. + """ + start = 0 # rebind-ok: cursor advances by one emitted segment per iteration + while start < len(text): + if len(text) - start <= max_length: + yield text[start:] + return + window = text[start : start + max_length] + boundaries = tuple(re.finditer(SENTENCE_BOUNDARY_PATTERN, window)) + cut_point = boundaries[-1].end() if boundaries else max_length + yield window[:cut_point] + start += cut_point + + def _split_text(self, text: str, max_length: int = 2000) -> tuple[str, ...]: + """ + Split text into segments of maximum length, trying to preserve sentence boundaries. + Args: + text: Text to split + max_length: Maximum length of each segment + Returns: + The text segments, in order + """ + return tuple(self._iter_text_segments(text, max_length)) + + @staticmethod + def _build_service_parameters(text: str | None, image_urls: Sequence[str] | None) -> Mapping[str, Any]: + """ + Build the ServiceParameters payload of one guardrail request. + Args: + text: The text to audit, when there is any + image_urls: The image URLs to audit, when there are any + Returns: + The payload, as a real dict because json.dumps rejects a MappingProxyType + """ + return dict( # mutable-ok: json.dumps rejects a MappingProxyType + ( + ("requestFrom", "LiteLLM"), + *((("content", text),) if text else ()), + *((("imageUrls", tuple(image_urls)),) if image_urls else ()), + ) + ) + + async def async_make_request( + self, + text: str | None = None, + service_type: Literal["input", "output", "mcp"] = "input", + image_urls: Sequence[str] | None = None, + ) -> AliyunAIGuardrailResponse: + """ + Make a request to the Aliyun AI Security Guardrail API. + Args: + text: Text to check (optional when only images are checked) + service_type: "input" for query_security_check, "output" for response_security_check, + "mcp" for MCP tool call check (uses service_mcp config) + image_urls: Public image URLs to check (optional) + Returns: + AliyunAIGuardrailResponse + """ + from litellm.types.proxy.guardrails.guardrail_hooks.aliyun.aliyun_ai_guardrail import ( + AliyunAIGuardrailResponse, + ) + + service_code: Final = ( + self.service_mcp + if service_type == "mcp" + else self.service_input + if service_type == "input" + else self.service_output + ) + service_parameters: Final = self._build_service_parameters(text, image_urls) + # httpx form-encodes `data=` from a dict, so this payload stays a real dict too + parameters: Final[dict[str, str]] = { # mutable-ok: httpx form-encodes `data=` from a dict + "Action": "MultiModalGuard", + "Version": "2022-03-02", + "AccessKeyId": self.access_key_id, + "Timestamp": self._format_iso8601_date(), + "SignatureMethod": "HMAC-SHA1", + "SignatureVersion": "1.0", + "SignatureNonce": str(uuid.uuid4()), + "Format": "JSON", + "Service": service_code, + "ServiceParameters": json.dumps(service_parameters, ensure_ascii=False), + } + string_to_sign: Final = self._create_string_to_sign("POST", parameters) + parameters["Signature"] = self._create_signature(string_to_sign) + verbose_proxy_logger.debug( + "Aliyun AI Guardrail request: service=%s, text_length=%d, image_count=%d", + service_code, + len(text) if text else 0, + len(image_urls) if image_urls else 0, + ) + response: Final = await self.async_handler.post( + url=self.service_url, + data=parameters, + headers={"Content-Type": "application/x-www-form-urlencoded"}, # mutable-ok: httpx wants a dict + timeout=30.0, + ) + body: Final = response.json() + verbose_proxy_logger.debug("Aliyun AI Guardrail response: %s", body) + if response.status_code != 200: + raise HTTPException( + status_code=response.status_code, + detail={ # mutable-ok: HTTPException detail payload + "error": f"Aliyun AI Guardrail request failed. Status: {response.status_code}, Body: {body}" + }, + ) + if body.get("Code") != 200: + raise HTTPException( + status_code=400, + detail={ # mutable-ok: HTTPException detail payload + "error": f"Aliyun AI Guardrail API error. Code: {body.get('Code')}, Message: {body.get('Message')}" + }, + ) + return AliyunAIGuardrailResponse( + RequestId=body.get("RequestId", ""), + Code=body.get("Code", 0), + Message=body.get("Message"), + Data=body.get("Data"), + ) + + def _should_block_by_level(self, detected_level: str) -> bool: + """ + Check if the detected risk level should trigger blocking based on protection level. + Logic: If detected_level_int >= threshold_int, then should block. + Args: + detected_level: Risk level from API response (none/low/medium/high or S0/S1/S2/S3/S4) + Returns: + True if should block, False otherwise + """ + threshold: Final = PROTECTION_LEVEL_THRESHOLD.get(self.level, 99) + detected_int: Final = level_to_int(detected_level) + return detected_int >= threshold + + def _resolve_detail_level(self, detail: Mapping[str, Any]) -> str: + """ + Resolve the risk level of a single Detail entry. + + MultiModalGuard reports severity in two shapes: the ``_pro`` service codes + return ``Detail[].Level``, while the documented response carries it as + ``Detail[].Result[].RiskLevel``. Honouring only the former downgrades the + latter to "none", which would let content Aliyun rejected pass through. + Args: + detail: A single entry of the response ``Detail`` list + Returns: + The risk level string to weigh against the configured threshold + """ + level: Final = detail.get("Level") + if isinstance(level, str) and level.strip(): + return level + # Fall back to the highest RiskLevel reported across the individual results + risk_levels: Final = tuple( + risk_level + for result in detail.get("Result") or () + if isinstance(result, dict) + for risk_level in (result.get("RiskLevel"),) + if isinstance(risk_level, str) and risk_level.strip() + ) + if risk_levels: + # max() keeps the first of equal-severity levels, as the previous scan did + return max(risk_levels, key=level_to_int) + # No parseable severity: never downgrade an explicit block to "none" + if detail.get("Suggestion") == BLOCK_SUGGESTION: + return UNRESOLVED_BLOCK_LEVEL + return "none" + + @staticmethod + def _resolve_desensitization(detail_list: Sequence[Any]) -> str: + """ + Return the desensitized text reported for sensitive data. + The first ``Desensitization`` within one detail's results wins, and a later + sensitiveData detail overrides an earlier one. + Args: + detail_list: The response ``Detail`` entries + Returns: + The desensitized text, empty when none was reported + """ + resolved = "" # rebind-ok: a later sensitiveData detail overrides an earlier one + for detail in detail_list: + if detail.get("Type", "") != SENSITIVE_DATA_TYPE: + continue + for ext in (result.get("Ext") for result in detail.get("Result", ()) or ()): + if ext and ext.get("Desensitization"): + resolved = ext.get("Desensitization", "") + break + return resolved + + def _resolve_block_decision(self, detail_list: Sequence[Any], final_suggestion: str) -> tuple[str, str, str]: + """ + Decide whether the audited content must be blocked. + Args: + detail_list: The response ``Detail`` entries + final_suggestion: The overall ``Suggestion`` of the response + Returns: + (blocked_type, blocked_level, block_message); an empty message means pass + """ + for detail in detail_list: + detected_level = self._resolve_detail_level(detail) + if self._should_block_by_level(detected_level): + detection_type = detail.get("Type", "") + return detection_type, detected_level, f"检测到{detection_type} (风险等级: {detected_level})" + # Aliyun rejected the content overall but no single detection could be attributed + # (e.g. an empty Detail list, or every entry reporting pass). Treat it as the most + # severe level so the decision is not lost, still subject to the threshold. + if final_suggestion == BLOCK_SUGGESTION and self._should_block_by_level(UNRESOLVED_BLOCK_LEVEL): + return "", UNRESOLVED_BLOCK_LEVEL, f"阿里云返回阻断建议 (Suggestion: {final_suggestion})" + return "", "", "" + + def _summarise_detail(self, detail: Mapping[str, Any]) -> Mapping[str, Any]: + """Summarise one ``Detail`` entry for the parsed result payload.""" + return { # mutable-ok: returned payload + "level": self._resolve_detail_level(detail), + "suggestion": detail.get("Suggestion", "pass"), + "results": tuple(detail.get("Result") or ()), + } + + def _parse_response_and_check( + self, + response: AliyunAIGuardrailResponse, + check_type: Literal["input", "output"], + ) -> Mapping[str, Any]: + """ + Parse the guardrail response and check if content should be blocked. + Blocking logic: + Check if detected level >= threshold based on protection level setting + Args: + response: The API response + check_type: "input" or "output" + Returns: + Mapping with parsed results + Raises: + HTTPException if content should be blocked + """ + data: Final = response.get("Data") + if not data: + return PASS_RESULT + final_suggestion: Final = data.get("Suggestion", "pass") + detail_list: Final = tuple(data.get("Detail") or ()) + details: Final[dict[str, Mapping[str, Any]]] = { # mutable-ok: returned payload + detail.get("Type", ""): self._summarise_detail(detail) for detail in detail_list + } + desensitization: Final = self._resolve_desensitization(detail_list) + blocked_type, blocked_level, block_message = self._resolve_block_decision(detail_list, final_suggestion) + verbose_proxy_logger.debug( + "Aliyun AI Guardrail: level=%s, check_type=%s, should_block=%s, blocked_type=%s, blocked_level=%s", + self.level, + check_type, + bool(block_message), + blocked_type, + blocked_level, + ) + if block_message: + raise HTTPException( + status_code=400, + detail={ # mutable-ok: HTTPException detail payload + "error": f"Aliyun AI Guardrail: {block_message}", + "type": check_type, + "details": details, + }, + ) + return { # mutable-ok: returned payload + "flagged": final_suggestion != "pass", + "suggestion": final_suggestion, + "desensitization": desensitization, + "details": details, + "message": block_message, + } + + @staticmethod + def _build_input_payloads( + segments: Sequence[str], image_urls: Sequence[str] + ) -> tuple[tuple[str | None, Sequence[str] | None], ...]: + """ + Pair each text segment with the images to audit alongside it. + Every image URL rides with the first text segment so content and images are + checked together; later segments carry text only. With no text at all, a + single image-only request is produced. + Args: + segments: The text segments to audit + image_urls: The image URLs to audit + Returns: + (text, images) pairs, one per request to send + """ + if segments: + return tuple( + (segment, image_urls if idx == 0 and image_urls else None) for idx, segment in enumerate(segments) + ) + if image_urls: + return ((None, image_urls),) + return () + + @staticmethod + def _iter_responses_function_call_arguments(data: Mapping[str, object]) -> Iterator[str]: + input_value: Final = data.get("input") + if not isinstance(input_value, list): + return + for item in input_value: + if not isinstance(item, dict) or item.get("type") != "function_call": + continue + arguments = item.get("arguments") + if isinstance(arguments, str) and arguments: + yield arguments + + @log_guardrail_information + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict[str, Any], # mutable-ok: CustomLogger's hook contract + call_type: CallTypesLiteral, + ) -> dict[str, Any] | None: # mutable-ok: CustomLogger's hook contract + """ + Pre-call hook to scan user prompts before sending to LLM. + Raises HTTPException if content should be blocked. + """ + verbose_proxy_logger.info( + "Aliyun AI Guardrail: Running pre-call prompt scan, on call_type: %s", + call_type, + ) + if call_type == CallTypes.call_mcp_tool.value: + await self._mcp_pre_call_check(data) + return None + # Walks messages AND the Responses API's input, so /v1/responses is covered + new_messages: Final = cast( # cast-ok: the upstream iterator is typed as plain dicts + "Sequence[AllMessageValues]", tuple(_iter_inspection_messages(data)) + ) + function_call_arguments: Final = tuple(self._iter_responses_function_call_arguments(data)) + if not new_messages and not function_call_arguments: + verbose_proxy_logger.warning("Aliyun AI Guardrail: not running guardrail. No messages in data") + return data + message_prompt: Final = self.get_user_prompt(new_messages) + user_prompt: Final = ( + "\n".join(text for text in (message_prompt, *function_call_arguments) if text).strip() or None + ) + image_urls: Final = self.get_image_urls(new_messages) + if not user_prompt and not image_urls: + verbose_proxy_logger.warning("Aliyun AI Guardrail: No user prompt or image found") + return None + verbose_proxy_logger.info( + "Aliyun AI Guardrail: Pre-call scan started, prompt length: %d, image count: %d", + len(user_prompt) if user_prompt else 0, + len(image_urls), + ) + # _split_text already returns a single segment when the text fits + segments: Final = self._split_text(user_prompt, self.max_text_length) if user_prompt else () + payloads: Final = self._build_input_payloads(segments, image_urls) + semaphore: Final = asyncio.Semaphore(5) # Max 5 concurrent requests, MultiModalGuard API limit is 20 + + async def check_with_semaphore( + segment_text: str | None, segment_images: Sequence[str] | None + ) -> AliyunAIGuardrailResponse: + async with semaphore: + return await self.async_make_request(text=segment_text, service_type="input", image_urls=segment_images) + + responses: Final = await asyncio.gather(*(check_with_semaphore(t, imgs) for t, imgs in payloads)) + for response in responses: + self._parse_response_and_check(response, check_type="input") + verbose_proxy_logger.info("Aliyun AI Guardrail: Pre-call scan passed") + return None + + # ================================================================ + # MCP-specific guardrail methods + # ================================================================ + + async def _mcp_pre_call_check(self, data: Mapping[str, Any]) -> None: + """MCP pre-call: audit tool name + arguments before execution.""" + messages: Final = data.get("messages", ()) + content: Final = messages[0].get("content", "") if messages else "" + verbose_proxy_logger.info( + "Aliyun AI Guardrail: ★ MCP pre-call check started, content length: %d", + len(content), + ) + if not content: + return + # _split_text already returns a single segment when the text fits + segments: Final = self._split_text(content, self.max_text_length) + semaphore: Final = asyncio.Semaphore(5) + + async def check(text: str) -> AliyunAIGuardrailResponse: + async with semaphore: + return await self.async_make_request(text=text, service_type="mcp") + + responses: Final = await asyncio.gather(*(check(s) for s in segments)) + for resp in responses: + self._parse_response_and_check(resp, check_type="input") + verbose_proxy_logger.info("Aliyun AI Guardrail: ★ MCP pre-call check passed") + + def _should_run_post_mcp_call(self) -> bool: + """Check if post_mcp_call is configured in event_hook. + + Since GuardrailEventHooks enum does not include post_mcp_call, + we cannot use should_run_guardrail(). This method manually checks + whether the user configured 'post_mcp_call' in the guardrail mode. + + Returns True if post_mcp_call should run: + - event_hook is None → run for all events + - event_hook is a list containing 'post_mcp_call' + - event_hook is a string equal to 'post_mcp_call' + - event_hook is a Mode with 'post_mcp_call' in tags or default + """ + from litellm.types.guardrails import Mode + + if self.event_hook is None: + return True + if isinstance(self.event_hook, list): + return "post_mcp_call" in self.event_hook + if isinstance(self.event_hook, Mode): + for tag_value in self.event_hook.tags.values(): + if isinstance(tag_value, list): + if "post_mcp_call" in tag_value: + return True + elif tag_value == "post_mcp_call": + return True + if self.event_hook.default: + default_value: Final = self.event_hook.default + return "post_mcp_call" in (default_value if isinstance(default_value, list) else (default_value,)) + return False + return self.event_hook == "post_mcp_call" + + async def async_post_mcp_tool_call_hook( + self, + kwargs: Mapping[str, Any], + response_obj: MCPPostCallResponseObject, + start_time: datetime, + end_time: datetime, + ) -> MCPPostCallResponseObject | None: + """MCP post-call: audit tool output after execution. + + Raising from here does not block: the dispatcher treats every callback exception + as a non-blocking logging error and hands the untouched tool result back. Both + call sites also discard this hook's return value, so a violation has to be + written into the live tool result carried by ``kwargs["original_response"]``. + The replacement object is returned as well, to honour the hook's contract. + """ + # Since GuardrailEventHooks enum has no post_mcp_call, the framework + # always invokes this hook if implemented. We check config manually. + if not self._should_run_post_mcp_call(): + verbose_proxy_logger.info("Aliyun AI Guardrail: Skipping post_mcp_call — not configured in event_hook") + return None + verbose_proxy_logger.info("Aliyun AI Guardrail: ★ MCP post-call check started") + original_response: Final = kwargs.get("original_response") + # The live tool result is preferred; the wrapped copy is the fallback. + combined_text: Final = next( + ( + text + for text in ( + self._extract_mcp_tool_text(candidate) + for candidate in (original_response, getattr(response_obj, "mcp_tool_call_response", None)) + if candidate is not None + ) + if text + ), + "", + ) + if not combined_text: + return None + verbose_proxy_logger.info( + "Aliyun AI Guardrail: ★ MCP post-call check, response length: %d", + len(combined_text), + ) + # _split_text already returns a single segment when the text fits + segments: Final = self._split_text(combined_text, self.max_text_length) + semaphore: Final = asyncio.Semaphore(5) + + async def check(text: str) -> AliyunAIGuardrailResponse: + async with semaphore: + return await self.async_make_request(text=text, service_type="mcp") + + try: + responses: Final = await asyncio.gather(*(check(s) for s in segments)) + for resp in responses: + self._parse_response_and_check(resp, check_type="output") + except HTTPException as e: + verbose_proxy_logger.warning( + "Aliyun AI Guardrail: ★ MCP post-call blocked — tool output replaced with the violation detail" + ) + return self._block_mcp_tool_output( + detail=e.detail, + response_obj=response_obj, + original_response=original_response, + ) + except Exception as e: # noqa: BLE001 - fail closed on any guardrail failure + # Raising here is swallowed by the dispatcher as a non-blocking logging + # error, which would hand the unaudited tool output straight to the + # client. Fail closed, matching every other path of this integration. + verbose_proxy_logger.error( + "Aliyun AI Guardrail: ★ MCP post-call check failed, blocking unaudited tool output: %s", + e, + exc_info=True, + ) + return self._block_mcp_tool_output( + detail={ # mutable-ok: HTTPException detail payload + "error": f"Aliyun AI Guardrail 调用失败,工具输出未经审核,已拦截: {e!s}" + }, # mutable-ok: detail payload + response_obj=response_obj, + original_response=original_response, + ) + verbose_proxy_logger.info("Aliyun AI Guardrail: ★ MCP post-call check passed") + return None + + @staticmethod + def _find_coerced_field(payload: Sequence[object], field: str) -> object: + """ + Return one field of a CallToolResult that was coerced into (field, value) pairs. + Args: + payload: The coerced pairs + field: The field name to recover + Returns: + The field's value, or None when it is absent + """ + return next( + (entry[1] for entry in payload if isinstance(entry, tuple) and len(entry) == 2 and entry[0] == field), + None, + ) + + @classmethod + def _iter_mcp_content_items(cls, payload: object) -> tuple[object, ...]: + """ + Normalise an MCP tool result into its content items. + Args: + payload: A CallToolResult, its content list, or a raw response body + Returns: + The content items to audit, empty when none can be located + """ + if isinstance(payload, str): + return (payload,) + content: Final = getattr(payload, "content", None) + if isinstance(content, list): + return tuple(content) + if isinstance(payload, dict): + inner: Final = payload.get("content") + return tuple(inner) if isinstance(inner, list) else (payload,) + if isinstance(payload, list): + # MCPPostCallResponseObject declares mcp_tool_call_response as a list, so a + # CallToolResult handed to it is coerced by iterating the model into + # (field, value) pairs. Recover the real content instead of auditing reprs. + coerced: Final = cls._find_coerced_field(payload, "content") + return tuple(coerced) if isinstance(coerced, list) else tuple(payload) + return () + + @classmethod + def _extract_mcp_item_text(cls, item: object) -> str: + """ + Collect the text of a single MCP content item. + Args: + item: A content item, as a string, dict or MCP content model + Returns: + The item's text, empty when it carries none + """ + if isinstance(item, str): + return item + direct_text: Final = item.get("text") if isinstance(item, dict) else getattr(item, "text", None) + if direct_text: + return str(direct_text) + # An EmbeddedResource keeps its payload one level down, on the resource. + resource: Final = item.get("resource") if isinstance(item, dict) else getattr(item, "resource", None) + resource_text: Final = resource.get("text") if isinstance(resource, dict) else getattr(resource, "text", None) + return str(resource_text or "") + + @classmethod + def _extract_mcp_structured_content(cls, payload: object) -> str: + """ + Serialise a tool result's ``structuredContent``, when it carries one. + Args: + payload: A CallToolResult, a raw response body, or its coerced form + Returns: + The serialised structured payload, empty when there is none + """ + structured: Final = ( + payload.get("structuredContent") + if isinstance(payload, dict) + # MCPPostCallResponseObject coerces a CallToolResult into (field, value) pairs. + else cls._find_coerced_field(payload, "structuredContent") + if isinstance(payload, list) + else getattr(payload, "structuredContent", None) + ) + if structured is None: + return "" + if isinstance(structured, str): + return structured + return json.dumps(structured, ensure_ascii=False, default=str) + + def _extract_mcp_tool_text(self, payload: object) -> str: + """ + Collect the textual output of an MCP tool result. + A tool result carries text beyond ``content[].text``: ``structuredContent`` + holds an arbitrary JSON payload and an EmbeddedResource keeps its text on + ``resource.text``. Auditing only the plain text items would let a tool + return prohibited content in either field unchecked. + """ + item_texts: Final = tuple( + text for text in map(self._extract_mcp_item_text, self._iter_mcp_content_items(payload)) if text + ) + structured_text: Final = self._extract_mcp_structured_content(payload) + return "\n".join((*item_texts, structured_text) if structured_text else item_texts) + + @staticmethod + def _replace_tool_output_in_place(target: object, blocked_content: Sequence[Any]) -> bool: + """ + Overwrite an MCP tool result's content with ``blocked_content``, in place. + The in-place stores are the point: both dispatch sites hand the caller's own + tool result to the client, so the violation has to be written into it. + """ + 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` + # must be flagged, not just the SDK's CallToolResult. + target.isError = True # pyright: ignore[reportAttributeAccessIssue] # rebind-ok: flag the caller's result + except (AttributeError, TypeError, ValueError): + pass + return True + if isinstance(target, list): + target[:] = blocked_content # rebind-ok: the caller's result must be overwritten + return True + if isinstance(target, dict): + 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 + + def _block_mcp_tool_output( + self, + detail: object, + response_obj: MCPPostCallResponseObject, + original_response: object, + ) -> MCPPostCallResponseObject: + """Replace blocked MCP tool output, both in place and as the returned object.""" + from mcp.types import TextContent + + from litellm.types.llms.base import HiddenParams + from litellm.types.mcp import MCPPostCallResponseObject as _MCPPostCallResponseObject + + payload: Final = detail if isinstance(detail, dict) else {"error": str(detail)} # mutable-ok: JSON payload + blocked_content: Final = (TextContent(type="text", text=json.dumps(payload, ensure_ascii=False)),) + for target in (original_response, getattr(response_obj, "mcp_tool_call_response", None)): + self._replace_tool_output_in_place(target, blocked_content) + hidden_params: Final = getattr(response_obj, "hidden_params", None) + return _MCPPostCallResponseObject( + mcp_tool_call_response=blocked_content, + hidden_params=hidden_params if isinstance(hidden_params, HiddenParams) else HiddenParams(), + ) + + @staticmethod + def _iter_function_call_text(function: object) -> Iterator[str]: + """Yield the name and arguments of a tool/function call.""" + for field in ("name", "arguments"): + value = getattr(function, field, None) + if value: + yield str(value) + + @staticmethod + def _iter_output_item_text(output_items: Iterable[object] | None) -> Iterator[str]: + """Yield the text carried by Responses API output items.""" + for output_item in output_items or (): + content_parts = getattr(output_item, "content", None) + if content_parts: + for content_part in content_parts: + text = getattr(content_part, "text", None) + if text: + yield str(text) + else: + text = getattr(output_item, "text", None) + if text: + yield str(text) + + @classmethod + def _iter_completion_text(cls, message: object) -> Iterator[str]: + """ + Yield every text field a chat completion message or streaming delta carries. + Args: + message: A choice's ``message`` (non-streaming) or ``delta`` (streaming) + """ + for field in ("content", "reasoning_content"): + value = getattr(message, field, None) + if value: + yield str(value) + for call in getattr(message, "tool_calls", None) or (): + yield from cls._iter_function_call_text(getattr(call, "function", None)) + yield from cls._iter_function_call_text(getattr(message, "function_call", None)) + + @classmethod + def _extract_response_text(cls, response: object) -> str: + """ + Collect every text field a non-streaming response can carry. + Mirrors ``_extract_stream_chunk_text``: auditing only ``message.content`` + would release tool call arguments, reasoning text and every + /v1/responses body to the client unchecked whenever ``stream=False``. + Args: + response: A non-streaming response object + Returns: + The concatenated text to audit, empty when the response carries none + """ + return "\n".join( + ( + # Responses API bodies keep their text in output items, not in choices. + *cls._iter_output_item_text(getattr(response, "output", None)), + *( + text + for choice in getattr(response, "choices", None) or () + for text in cls._iter_completion_text(getattr(choice, "message", None)) + ), + ) + ) + + @log_guardrail_information + async def async_post_call_success_hook( + self, + data: dict, # mutable-ok: CustomLogger's hook contract + user_api_key_dict: UserAPIKeyAuth, + response: LLMResponseTypes, + ) -> LLMResponseTypes: + """ + Post-call hook to scan LLM responses. + Raises HTTPException if content should be blocked. + """ + # Every choice is scanned: n>1 responses would otherwise return unchecked content + content: Final = self._extract_response_text(response) + if content: + verbose_proxy_logger.info( + "Aliyun AI Guardrail: Post-call scan started, response length: %d", + len(content), + ) + # _split_text already returns a single segment when the text fits + segments: Final = self._split_text(content, self.max_text_length) + semaphore: Final = asyncio.Semaphore(5) # Max 5 concurrent requests + + async def check_with_semaphore(segment: str) -> AliyunAIGuardrailResponse: + async with semaphore: + return await self.async_make_request(text=segment, service_type="output") + + responses: Final = await asyncio.gather(*(check_with_semaphore(segment) for segment in segments)) + for guardrail_response in responses: + self._parse_response_and_check(guardrail_response, check_type="output") + verbose_proxy_logger.info("Aliyun AI Guardrail: Post-call scan passed") + return response + + @classmethod + def _extract_stream_chunk_text(cls, chunk: object) -> str: + """ + Collect every text field a streaming chunk can carry. + Covers both chat completion chunks and Responses API streaming events. + Auditing only ``delta.content`` would release tool call arguments, + reasoning text and every /v1/responses chunk to the client unchecked. + Args: + chunk: A streaming chunk + Returns: + The concatenated text to audit, empty when the chunk carries none + """ + # Responses API events carry text outside of choices: either a plain + # string delta, or the assembled output of a terminal event. + event_delta: Final = getattr(chunk, "delta", None) + return "".join( + ( + *((event_delta,) if isinstance(event_delta, str) and event_delta else ()), + *cls._iter_output_item_text(getattr(getattr(chunk, "response", None), "output", None)), + *( + text + for choice in getattr(chunk, "choices", None) or () + for text in cls._iter_completion_text(getattr(choice, "delta", None)) + ), + ) + ) + + @staticmethod + def _as_error_payload(detail: object) -> Mapping[str, Any]: + """Normalise an HTTPException detail into a JSON object.""" + return detail if isinstance(detail, dict) else {"message": str(detail)} # mutable-ok: JSON payload + + def _stream_check_windows(self, text: str, last_check_position: int) -> tuple[str, ...]: + current_length: Final = len(text) + window_size: Final = self.stream_window_size + latest_window_start: Final = max(0, current_length - window_size) + if current_length - last_check_position <= window_size: + return (text[latest_window_start:],) + step: Final = max(1, min(self.stream_slide_step, window_size)) + overlap: Final = window_size - step + first_window_start: Final = max(0, last_check_position - overlap) + window_starts: Final = (*range(first_window_start, latest_window_start, step), latest_window_start) + return tuple(text[start : min(start + window_size, current_length)] for start in window_starts) + + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncIterable[Any], + request_data: Mapping[str, Any], + ) -> AsyncGenerator[Any, None]: + """ + Process streaming response with sliding window guardrail checks. + This method implements sliding window guardrail checks based on + `stream_window_size` and `stream_slide_step`. + It triggers a guardrail API call when: + 1. Every `stream_slide_step` new chars accumulate since the last check. + 2. Stream ends and there's remaining unchecked content. + For example, if stream_window_size=2000, stream_slide_step=300: + - At 300 chars: check chars 0-300 (window: last 2000) + - At 600 chars: check chars 0-600 (window: last 2000) + - At 2100 chars: check chars 100-2100 (window slides forward) + - At 2400 chars: check chars 400-2400 (window slides forward) + - When stream ends with 2500 chars: check chars 500-2500 (final window) + """ + # Sliding-window state: every name below advances as the stream is consumed. + accumulated_text = "" # rebind-ok: grows with each chunk + last_check_position = 0 # rebind-ok: position (total length) of the last check + pending_chunks: Final[list[Any]] = [] # mutable-ok: buffer held back until a check passes + is_first_check = True # rebind-ok: the first check uses a smaller threshold + verbose_proxy_logger.info( + "Aliyun AI Guardrail: Streaming scan started, window=%d, step=%d, first_check_step=%d", + self.stream_window_size, + self.stream_slide_step, + self.stream_first_check_step, + ) + try: + async for chunk in response: + chunk_text = self._extract_stream_chunk_text(chunk) + accumulated_text += chunk_text + # Buffer the chunk, don't yield until guardrail check passes + pending_chunks.append(chunk) + current_length = len(accumulated_text) + new_chars_since_last_check = current_length - last_check_position + check_threshold = self.stream_first_check_step if is_first_check else self.stream_slide_step + if new_chars_since_last_check >= check_threshold: + check_windows = self._stream_check_windows(accumulated_text, last_check_position) + for text_to_check in check_windows: + guardrail_response = await self.async_make_request(text=text_to_check, service_type="output") + self._parse_response_and_check(guardrail_response, check_type="output") + verbose_proxy_logger.info( + "Aliyun AI Guardrail: Streaming check passed at position %d", current_length + ) + for pending_chunk in pending_chunks: + yield pending_chunk + pending_chunks.clear() + last_check_position = current_length + is_first_check = False + # Stream ended - check any remaining unchecked content with a final window + if len(accumulated_text) > last_check_position: + final_windows: Final = self._stream_check_windows(accumulated_text, last_check_position) + for remaining_text in final_windows: + final_response = await self.async_make_request(text=remaining_text, service_type="output") + self._parse_response_and_check(final_response, check_type="output") + verbose_proxy_logger.info( + "Aliyun AI Guardrail: Streaming scan completed, total length: %d", len(accumulated_text) + ) + for pending_chunk in pending_chunks: + yield pending_chunk + pending_chunks.clear() + except HTTPException as e: + detail_payload: Final = self._as_error_payload(e.detail) + verbose_proxy_logger.info("Aliyun AI Guardrail: Streaming blocked at position %d", len(accumulated_text)) + payload: Final = json.dumps({"error": detail_payload}, ensure_ascii=False) # mutable-ok: SSE payload + yield f"data: {payload}\n\n" + return + except Exception as e: + verbose_proxy_logger.error("Aliyun AI Guardrail streaming error: %s", e, exc_info=True) + raise diff --git a/litellm/proxy/guardrails/guardrail_hooks/aliyun/base.py b/litellm/proxy/guardrails/guardrail_hooks/aliyun/base.py new file mode 100644 index 00000000000..b50884b3de2 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/aliyun/base.py @@ -0,0 +1,100 @@ +""" +Base class for Aliyun guardrails +阿里云护栏基类 +""" + +from __future__ import annotations + +from collections.abc import Iterator, Sequence +from typing import TYPE_CHECKING, Final + +if TYPE_CHECKING: + from litellm.types.llms.openai import AllMessageValues + + +class AliyunGuardrailBase: + """ + Base class for Aliyun guardrails. + """ + + @staticmethod + def _iter_user_messages(messages: Sequence[AllMessageValues]) -> Iterator[AllMessageValues]: + """ + Yield every user message of the request, in order. + Restricting this to the trailing user block would let a caller hide a + prohibited turn behind an attacker-supplied assistant message. + """ + 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: + """ + Return the URL of an ``image_url`` content part. + Args: + part: A single content part of a message + Returns: + The URL string, or None when the part carries no 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): + url: Final = image_url.get("url") + return url if isinstance(url, str) else None + return image_url if isinstance(image_url, str) else None + + def get_user_prompt(self, messages: Sequence[AllMessageValues]) -> str | None: + """ + Collect the text of every user message in the request. + Scanning only the trailing user block would let a caller hide a + prohibited prompt behind an attacker-supplied assistant message, so all + user turns of the submitted request are audited. + Example: + messages = [ + {"role": "user", "content": "Hello, how are you?"}, + {"role": "assistant", "content": "I'm good, thank you!"}, + {"role": "user", "content": "What is the weather in Tokyo?"}, + ] + get_user_prompt(messages) -> "Hello, how are you?\nWhat is the weather in Tokyo?" + """ + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + convert_content_list_to_str, + ) + + user_prompt: Final = "\n".join( + convert_content_list_to_str(message) for message in self._iter_audited_text_messages(messages) + ).strip() + return user_prompt or None + + def _iter_public_image_urls(self, messages: Sequence[AllMessageValues]) -> Iterator[str]: + """Yield publicly reachable image URLs from user and tool messages.""" + for content in (message.get("content") for message in self._iter_audited_text_messages(messages)): + if not isinstance(content, list): + continue + for url in (self._extract_image_url(part) for part in content): + # Only public http(s) URLs are reachable by the Aliyun API, so + # data: URIs and other inline payloads are skipped. + if url is not None and url.startswith(("http://", "https://")): + yield url + + def get_image_urls(self, messages: Sequence[AllMessageValues]) -> tuple[str, ...]: + """ + Extract image URLs from every user and tool message in the request. + Only publicly accessible http(s) URLs are collected (in order, + de-duplicated). Uses the same message range as ``get_user_prompt``. + Example: + messages = [ + {"role": "user", "content": [ + {"type": "text", "text": "what is in this image?"}, + {"type": "image_url", "image_url": {"url": "https://a.com/x.png"}}, + ]}, + ] + get_image_urls(messages) -> ("https://a.com/x.png",) + """ + # dict.fromkeys is the order-preserving dedup; it is transient and the + # result is frozen into a tuple before it leaves this method. + return tuple(dict.fromkeys(self._iter_public_image_urls(messages))) # mutable-ok: transient dedup, frozen here diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index e48c36de8ba..6b769059b62 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -134,6 +134,7 @@ class SupportedGuardrailIntegrations(Enum): HEADROOM = "headroom" COMPRESR = "compresr" STRAIKER = "straiker" + ALIYUN_AI_GUARDRAIL = "aliyun_ai_guardrail" class Role(Enum): @@ -667,6 +668,55 @@ class JavelinGuardrailConfigModel(BaseModel): config: dict | None = Field(default=None, description="Additional configuration for the guardrail") +class AliyunAIGuardrailConfigModel(BaseModel): + """Configuration parameters for the Aliyun AI Security guardrail.""" + + level: str | None = Field( + default=None, + description="Protection level. 'low': block all risks (high protection), 'medium': block medium+high risks, 'high': block high only, 'max': observe mode. Default: medium", + ) + max_text_length: int | None = Field( + default=None, + description="Maximum text length for a single API call. Text longer than this will be split.", + ) + stream_window_size: int | None = Field( + default=None, + description="Sliding window size (in chars) for streaming output guardrail checks. Default: 500", + ) + stream_slide_step: int | None = Field( + default=None, + description="Sliding step (in chars) for streaming output guardrail checks. Default: 300", + ) + stream_first_check_step: int | None = Field( + default=None, + description="First check threshold (in chars) to reduce first-token latency. Default: 50", + ) + region_id: str | None = Field( + default=None, + description="Aliyun region ID. Default: cn-shanghai", + ) + service_input: str | None = Field( + default=None, + description="Service code for input (pre-call) detection. Default: query_security_check", + ) + service_output: str | None = Field( + default=None, + description="Service code for output (post-call) detection. Default: response_security_check", + ) + service_mcp: str | None = Field( + default=None, + description="Service code for MCP tool call detection (pre/post MCP call). Default: query_security_check", + ) + access_key_id: str | None = Field( + default=None, + description="Aliyun Access Key ID for the guardrail. Configure in config.yaml, supports os.environ/ reference", + ) + access_key_secret: str | None = Field( + default=None, + description="Aliyun Access Key Secret for the guardrail. Configure in config.yaml, supports os.environ/ reference", + ) + + class ContentFilterAction(str, Enum): """Action to take when content filter detects a match""" @@ -1012,6 +1062,7 @@ class LitellmParams( QostodianNexusConfigModel, VigilGuardGuardrailConfigModel, SingulrGuardrailConfigModel, + AliyunAIGuardrailConfigModel, ): guardrail: str = Field(description="The type of guardrail integration to use") mode: str | list[str] | Mode = Field( diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/aliyun/aliyun_ai_guardrail.py b/litellm/types/proxy/guardrails/guardrail_hooks/aliyun/aliyun_ai_guardrail.py new file mode 100644 index 00000000000..9d763602352 --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/aliyun/aliyun_ai_guardrail.py @@ -0,0 +1,168 @@ +""" +Type definitions for Aliyun AI Security Guardrail +阿里云AI安全护栏类型定义 +Aliyun AI Guardrail supports the following detection types: +- contentModeration: Content safety moderation +- sensitiveData: Sensitive data detection (PII, etc.) +- promptAttack: Prompt injection attack detection +- maliciousUrl: Malicious URL detection +""" + +from collections.abc import Sequence +from typing import Literal, TypeAlias + +from pydantic import BaseModel, Field +from typing_extensions import NotRequired, ReadOnly, TypedDict + +from ..base import GuardrailConfigModel + + +# Response types +class AliyunAIGuardrailResponseDetailResultExt(TypedDict, total=False): + """Extended information in result""" + + Desensitization: ReadOnly[str | None] # Desensitized text when action is mask + + +class AliyunAIGuardrailResponseDetailResult(TypedDict, total=False): + """Result item in detail""" + + Confidence: ReadOnly[float | None] + Label: ReadOnly[str | None] + Ext: ReadOnly[AliyunAIGuardrailResponseDetailResultExt | None] + # Per-result risk level. This is the shape documented for MultiModalGuard; the + # ``_pro`` service codes report the severity on the parent Detail as ``Level`` + # instead, so both have to be honoured when deciding whether to block. + RiskLevel: ReadOnly[str | None] + + +class AliyunAIGuardrailResponseDetail(TypedDict): + """Detail item in response data""" + + Type: ReadOnly[str] # contentModeration, sensitiveData, promptAttack, maliciousUrl + Suggestion: ReadOnly[str] # pass, block, mask + Result: ReadOnly[Sequence[AliyunAIGuardrailResponseDetailResult]] + # Risk level as returned by the ``_pro`` service codes (none/low/medium/high, or + # S0-S4 for sensitiveData). Absent in the documented response shape, which carries + # the severity as Result[].RiskLevel. + Level: ReadOnly[NotRequired[str]] + + +class AliyunAIGuardrailResponseData(TypedDict, total=False): + """Response data from Aliyun AI Guardrail API""" + + Suggestion: ReadOnly[str] # Overall suggestion: pass, block, mask + Detail: ReadOnly[Sequence[AliyunAIGuardrailResponseDetail] | None] + + +class AliyunAIGuardrailResponse(TypedDict): + """Response from Aliyun AI Guardrail API""" + + RequestId: ReadOnly[str] + Code: ReadOnly[int] + Message: ReadOnly[str | None] + Data: ReadOnly[AliyunAIGuardrailResponseData | None] + + +# Suggestion type +AliyunAIGuardrailSuggestion: TypeAlias = Literal["pass", "block", "watch"] + +# Detection type +AliyunAIGuardrailDetectionType: TypeAlias = Literal[ + "contentModeration", "sensitiveData", "promptAttack", "maliciousUrl" +] + + +class AliyunAIGuardrailRequestParams(TypedDict, total=False): + """Request parameters for Aliyun AI Guardrail API""" + + Action: ReadOnly[str] + Version: ReadOnly[str] + AccessKeyId: ReadOnly[str] + Timestamp: ReadOnly[str] + SignatureMethod: ReadOnly[str] + SignatureVersion: ReadOnly[str] + SignatureNonce: ReadOnly[str] + Format: ReadOnly[str] + Service: ReadOnly[str] + ServiceParameters: ReadOnly[str] + Signature: ReadOnly[str] + + +# Risk level literals +AliyunRiskLevel: TypeAlias = Literal["none", "low", "medium", "high"] + +# Protection level literals +AliyunProtectionLevel: TypeAlias = Literal["low", "medium", "high", "max"] + + +# Configuration models +class AliyunAIGuardrailOptionalParams(BaseModel): + """ + Optional parameters for Aliyun AI Guardrail. + Credentials (access_key_id / access_key_secret) are configured + in config.yaml on the AliyunAIGuardrailConfigModel and support os.environ/ references. + """ + + level: AliyunProtectionLevel | None = Field( + default="medium", + description="Protection level for risk filtering. 'low': block all risks (high protection), 'medium': block medium and high risks, 'high': block only high risks (low protection), 'max': observation mode (no blocking). Default: medium", + ) + max_text_length: int | None = Field( + default=2000, + description="Maximum text length for a single API call. Text longer than this will be split.", + ) + stream_window_size: int | None = Field( + default=500, + description="Sliding window size (in chars) for streaming output guardrail checks. Each check sends the most recent N chars to the API.", + ) + stream_slide_step: int | None = Field( + default=300, + description="Sliding step (in chars) for streaming output guardrail checks. A check is triggered every time N new chars accumulate since the last check.", + ) + stream_first_check_step: int | None = Field( + default=50, + description="First check threshold (in chars) for streaming output. The first guardrail check triggers earlier (at N chars) to reduce first-token latency, subsequent checks use stream_slide_step.", + ) + region_id: str | None = Field( + default="cn-shanghai", + description="Aliyun region ID. Default: cn-shanghai", + ) + service_input: str | None = Field( + default="query_security_check_pro", + description="Service code for input (pre-call) detection. Default: query_security_check_pro", + ) + service_output: str | None = Field( + default="response_security_check_pro", + description="Service code for output (post-call) detection. Default: response_security_check_pro", + ) + service_mcp: str | None = Field( + default="query_security_check_pro", + description="Service code for MCP tool call detection (pre_mcp_call and post_mcp_call). Default: query_security_check_pro", + ) + + +class AliyunAIGuardrailConfigModel(GuardrailConfigModel[AliyunAIGuardrailOptionalParams]): + """ + Configuration model for Aliyun AI Guardrail. + Credentials are configured in config.yaml and support os.environ/ references: + - access_key_id: Aliyun Access Key ID + - access_key_secret: Aliyun Access Key Secret + """ + + access_key_id: str | None = Field( + default=None, + description="Aliyun Access Key ID. Configure in config.yaml, supports os.environ/ reference", + ) + access_key_secret: str | None = Field( + default=None, + description="Aliyun Access Key Secret. Configure in config.yaml, supports os.environ/ reference", + ) + optional_params: AliyunAIGuardrailOptionalParams | None = Field( + default_factory=AliyunAIGuardrailOptionalParams, + description="Optional parameters for the Aliyun AI Guardrail", + ) + + @staticmethod + def ui_friendly_name() -> str: + return "Aliyun AI Security Guardrail" 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 new file mode 100644 index 00000000000..cc27f4cbf08 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/aliyun/test_aliyun_ai_guardrail.py @@ -0,0 +1,2070 @@ +""" +Unit tests for Aliyun AI Security Guardrail. + +Tests cover: +- Registration in guardrail system (enum, initializer, registry) +- Constructor validation (credentials, level, configurable service codes) +- Helper functions (level_to_int, _split_text) +- Blocking logic (_should_block_by_level, _parse_response_and_check) +- Pre-call hook (text + multimodal image URL detection) +- Post-call hook (blocks violations in response, passes clean response) +- Image URL extraction (get_image_urls) +- ServiceParameters construction (text / image / mixed combos) +- Config model (get_config_model, ui_friendly_name) +""" + +import json +from datetime import datetime +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import HTTPException + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.aliyun.aliyun_ai_guardrail import ( + CONTENT_MODERATION_TYPE, + PROMPT_ATTACK_TYPE, + SENSITIVE_DATA_TYPE, + AliyunAIGuardrail, + level_to_int, +) + +FAKE_AK = "test-access-key-id" +FAKE_SK = "test-access-key-secret" + +IMG_A = "https://example.com/a.png" +IMG_B = "http://example.com/b.jpg" +DATA_URI = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAAC0lEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" + + +def _make_guardrail(**kwargs) -> AliyunAIGuardrail: + defaults = dict( + guardrail_name="test-aliyun", + access_key_id=FAKE_AK, + access_key_secret=FAKE_SK, + level="medium", + ) + defaults.update(kwargs) + return AliyunAIGuardrail(**defaults) + + +def _make_aliyun_api_response( + suggestion: str = "pass", + detail: list = None, + code: int = 200, +) -> MagicMock: + """Build a mock httpx.Response mimicking Aliyun API output.""" + mock = MagicMock() + mock.status_code = 200 + mock.json.return_value = { + "Code": code, + "RequestId": "test-req-id", + "Message": None, + "Data": { + "Suggestion": suggestion, + "Detail": detail or [], + }, + } + return mock + + +def _make_detail( + detection_type: str = CONTENT_MODERATION_TYPE, + level: str = "high", + suggestion: str = "block", + results: list = None, +) -> dict: + """Build a Detail item. Pass level=None to omit the Level field entirely, + mimicking the response shape documented for MultiModalGuard, which carries + the severity as Result[].RiskLevel instead.""" + detail = { + "Type": detection_type, + "Suggestion": suggestion, + "Result": results or [], + } + if level is not None: + detail["Level"] = level + return detail + + +def _captured_service_parameters(mock_post: AsyncMock): + """Return (ServiceParameters dict, Service code) from a mocked post call.""" + _, kwargs = mock_post.call_args + params = kwargs["data"] + return json.loads(params["ServiceParameters"]), params["Service"] + + +# --------------------------------------------------------------------------- +# Registration tests +# --------------------------------------------------------------------------- + + +class TestAliyunGuardrailRegistration: + def test_supported_guardrail_enum_entry(self): + from litellm.types.guardrails import SupportedGuardrailIntegrations + + assert hasattr(SupportedGuardrailIntegrations, "ALIYUN_AI_GUARDRAIL") + assert SupportedGuardrailIntegrations.ALIYUN_AI_GUARDRAIL.value == "aliyun_ai_guardrail" + + def test_initialize_guardrail_function_exists(self): + from litellm.proxy.guardrails.guardrail_hooks.aliyun import ( + guardrail_initializer_registry, + initialize_guardrail, + ) + + assert initialize_guardrail is not None + assert "aliyun_ai_guardrail" in guardrail_initializer_registry + + def test_guardrail_class_registry_exists(self): + from litellm.proxy.guardrails.guardrail_hooks.aliyun import ( + guardrail_class_registry, + ) + + assert "aliyun_ai_guardrail" in guardrail_class_registry + assert guardrail_class_registry["aliyun_ai_guardrail"] is AliyunAIGuardrail + + def test_aliyun_in_global_registry(self): + from litellm.proxy.guardrails.guardrail_registry import ( + guardrail_initializer_registry, + ) + + assert "aliyun_ai_guardrail" in guardrail_initializer_registry + + def test_initialize_guardrail_creates_instance(self): + from litellm.proxy.guardrails.guardrail_hooks.aliyun import ( + initialize_guardrail, + ) + from litellm.types.guardrails import LitellmParams + + litellm_params = LitellmParams( + guardrail="aliyun_ai_guardrail", + mode="pre_call", + level="medium", + access_key_id=FAKE_AK, + access_key_secret=FAKE_SK, + ) + guardrail_config = {"guardrail_name": "test-aliyun-guard"} + + with patch("litellm.logging_callback_manager.add_litellm_callback") as mock_add: + result = initialize_guardrail(litellm_params, guardrail_config) + + assert isinstance(result, AliyunAIGuardrail) + assert result.guardrail_name == "test-aliyun-guard" + assert result.level == "medium" + assert result.access_key_id == FAKE_AK + assert result.access_key_secret == FAKE_SK + mock_add.assert_called_once_with(result) + + def test_initialize_guardrail_resolves_os_environ_reference(self): + from litellm.proxy.guardrails.guardrail_hooks.aliyun import ( + initialize_guardrail, + ) + from litellm.types.guardrails import LitellmParams + + litellm_params = LitellmParams( + guardrail="aliyun_ai_guardrail", + mode="pre_call", + access_key_id="os.environ/GUARD_ACCESS_KEY_ID", + access_key_secret="os.environ/GUARD_ACCESS_KEY_SECRET", + ) + guardrail_config = {"guardrail_name": "test-aliyun-guard"} + + with ( + patch.dict( + "os.environ", + { + "GUARD_ACCESS_KEY_ID": FAKE_AK, + "GUARD_ACCESS_KEY_SECRET": FAKE_SK, + }, + ), + patch("litellm.logging_callback_manager.add_litellm_callback"), + ): + result = initialize_guardrail(litellm_params, guardrail_config) + + assert result.access_key_id == FAKE_AK + assert result.access_key_secret == FAKE_SK + + def test_initialize_guardrail_forwards_service_mcp(self): + from litellm.proxy.guardrails.guardrail_hooks.aliyun import ( + initialize_guardrail, + ) + from litellm.types.guardrails import LitellmParams + + litellm_params = LitellmParams( + guardrail="aliyun_ai_guardrail", + mode="pre_mcp_call", + access_key_id=FAKE_AK, + access_key_secret=FAKE_SK, + service_mcp="text_img_mix_guard", + ) + guardrail_config = {"guardrail_name": "test-aliyun-guard"} + + with patch("litellm.logging_callback_manager.add_litellm_callback"): + result = initialize_guardrail(litellm_params, guardrail_config) + + assert result.service_mcp == "text_img_mix_guard" + + +# --------------------------------------------------------------------------- +# Constructor tests +# --------------------------------------------------------------------------- + + +class TestAliyunGuardrailConstructor: + def test_init_with_explicit_credentials(self): + g = _make_guardrail() + assert g.access_key_id == FAKE_AK + assert g.access_key_secret == FAKE_SK + assert g.level == "medium" + assert g.region_id == "cn-shanghai" + + def test_init_credentials_from_config(self): + g = AliyunAIGuardrail( + guardrail_name="config-test", + access_key_id="cfg-ak", + access_key_secret="cfg-sk", + level="low", + ) + assert g.access_key_id == "cfg-ak" + assert g.access_key_secret == "cfg-sk" + # region defaults to cn-shanghai when not provided via config + assert g.region_id == "cn-shanghai" + + def test_init_raises_without_api_key(self): + with pytest.raises(ValueError, match="ak is required"): + AliyunAIGuardrail(guardrail_name="test") + + def test_init_raises_without_api_secret(self): + with pytest.raises(ValueError, match="sk is required"): + AliyunAIGuardrail(guardrail_name="test", access_key_id=FAKE_AK) + + def test_init_invalid_level_raises(self): + with pytest.raises(ValueError, match="Invalid level"): + _make_guardrail(level="invalid") + + def test_init_default_service_codes_domestic(self): + g = _make_guardrail() + assert g.service_input == "query_security_check_pro" + assert g.service_output == "response_security_check_pro" + + def test_init_service_codes_are_configurable(self): + g = _make_guardrail( + service_input="text_img_mix_guard", + service_output="response_security_check_cb", + ) + assert g.service_input == "text_img_mix_guard" + assert g.service_output == "response_security_check_cb" + + def test_init_region_is_configurable(self): + g = _make_guardrail(region_id="eu-central-1") + assert g.region_id == "eu-central-1" + + +# --------------------------------------------------------------------------- +# Helper function tests +# --------------------------------------------------------------------------- + + +class TestLevelToInt: + def test_standard_levels(self): + assert level_to_int("none") == 0 + assert level_to_int("low") == 1 + assert level_to_int("medium") == 2 + assert level_to_int("high") == 3 + + def test_sensitive_data_levels(self): + assert level_to_int("S0") == 0 + assert level_to_int("S1") == 1 + assert level_to_int("S2") == 2 + assert level_to_int("S3") == 3 + assert level_to_int("S4") == 3 + + def test_case_insensitive(self): + assert level_to_int("HIGH") == 3 + assert level_to_int("Low") == 1 + + def test_empty_string_defaults_to_zero(self): + assert level_to_int("") == 0 + assert level_to_int(None) == 0 + + def test_unknown_level_defaults_to_zero(self): + assert level_to_int("unknown") == 0 + + +class TestSplitText: + def test_short_text_returns_single_segment(self): + g = _make_guardrail() + result = g._split_text("short text", max_length=100) + assert result == ("short text",) + + def test_long_text_splits_at_sentence_boundary(self): + g = _make_guardrail() + text = "Hello world. This is a test. Final sentence." + result = g._split_text(text, max_length=20) + assert len(result) >= 2 + assert "".join(result) == text + + def test_long_text_without_boundary_splits_at_max_length(self): + g = _make_guardrail() + text = "a" * 100 + result = g._split_text(text, max_length=30) + assert len(result) >= 3 + assert "".join(result) == text + + def test_empty_text_returns_no_segments(self): + g = _make_guardrail() + result = g._split_text("", max_length=100) + assert result == () + + +# --------------------------------------------------------------------------- +# Blocking logic tests +# --------------------------------------------------------------------------- + + +class TestShouldBlockByLevel: + def test_low_protection_blocks_all_risks(self): + g = _make_guardrail(level="low") + assert g._should_block_by_level("low") is True + assert g._should_block_by_level("medium") is True + assert g._should_block_by_level("high") is True + assert g._should_block_by_level("none") is False + + def test_medium_protection_blocks_medium_and_high(self): + g = _make_guardrail(level="medium") + assert g._should_block_by_level("low") is False + assert g._should_block_by_level("medium") is True + assert g._should_block_by_level("high") is True + + def test_high_protection_blocks_high_only(self): + g = _make_guardrail(level="high") + assert g._should_block_by_level("medium") is False + assert g._should_block_by_level("high") is True + + def test_max_observation_never_blocks(self): + g = _make_guardrail(level="max") + assert g._should_block_by_level("high") is False + assert g._should_block_by_level("S4") is False + + +class TestParseResponseAndCheck: + def test_blocks_when_level_meets_threshold(self): + g = _make_guardrail(level="medium") + response = { + "Data": { + "Suggestion": "block", + "Detail": [_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + }, + } + with pytest.raises(HTTPException) as exc_info: + g._parse_response_and_check(response, check_type="input") + assert exc_info.value.status_code == 400 + # Message uses the raw detection type returned by Aliyun + assert CONTENT_MODERATION_TYPE in str(exc_info.value.detail) + + def test_passes_when_level_below_threshold(self): + g = _make_guardrail(level="high") + response = { + "Data": { + "Suggestion": "pass", + "Detail": [_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="low")], + }, + } + result = g._parse_response_and_check(response, check_type="input") + assert result["flagged"] is False + + def test_empty_data_returns_pass(self): + g = _make_guardrail() + response = {"Data": {}} + result = g._parse_response_and_check(response, check_type="input") + assert result["flagged"] is False + assert result["suggestion"] == "pass" + + def test_extracts_desensitization_for_sensitive_data(self): + # level=max never blocks, so the parsed result (with desensitization) is returned + g = _make_guardrail(level="max") + response = { + "Data": { + "Suggestion": "mask", + "Detail": [ + { + "Type": SENSITIVE_DATA_TYPE, + "Level": "S2", + "Suggestion": "mask", + "Result": [{"Ext": {"Desensitization": "masked_text"}}], + }, + ], + }, + } + result = g._parse_response_and_check(response, check_type="input") + assert result["desensitization"] == "masked_text" + + def test_prompt_attack_block_message(self): + g = _make_guardrail(level="medium") + response = { + "Data": { + "Suggestion": "block", + "Detail": [_make_detail(detection_type=PROMPT_ATTACK_TYPE, level="medium")], + }, + } + with pytest.raises(HTTPException) as exc_info: + g._parse_response_and_check(response, check_type="input") + assert PROMPT_ATTACK_TYPE in str(exc_info.value.detail) + + +class TestRiskLevelResolution: + """MultiModalGuard reports severity in two shapes: Detail[].Level (returned by + the _pro service codes) and Detail[].Result[].RiskLevel (the documented shape). + Reading only the former silently downgrades the latter to "none" and lets + Aliyun's own block decision through. + """ + + def test_falls_back_to_result_risk_level_when_level_absent(self): + g = _make_guardrail(level="medium") + response = { + "Data": { + "Suggestion": "block", + "Detail": [ + _make_detail( + detection_type=CONTENT_MODERATION_TYPE, + level=None, + results=[{"Label": "violence", "Confidence": 99.5, "RiskLevel": "high"}], + ) + ], + }, + } + with pytest.raises(HTTPException) as exc_info: + g._parse_response_and_check(response, check_type="input") + assert CONTENT_MODERATION_TYPE in str(exc_info.value.detail) + + def test_result_risk_level_still_respects_threshold(self): + g = _make_guardrail(level="high") + response = { + "Data": { + "Suggestion": "pass", + "Detail": [ + _make_detail( + detection_type=CONTENT_MODERATION_TYPE, + level=None, + suggestion="pass", + results=[{"Label": "spam", "RiskLevel": "low"}], + ) + ], + }, + } + result = g._parse_response_and_check(response, check_type="input") + assert result["flagged"] is False + + def test_highest_result_risk_level_wins(self): + g = _make_guardrail(level="high") + response = { + "Data": { + "Suggestion": "block", + "Detail": [ + _make_detail( + detection_type=CONTENT_MODERATION_TYPE, + level=None, + results=[ + {"Label": "spam", "RiskLevel": "low"}, + {"Label": "violence", "RiskLevel": "high"}, + ], + ) + ], + }, + } + with pytest.raises(HTTPException): + g._parse_response_and_check(response, check_type="input") + + def test_detail_level_takes_precedence_over_result_risk_level(self): + g = _make_guardrail(level="medium") + response = { + "Data": { + "Suggestion": "pass", + "Detail": [ + _make_detail( + detection_type=CONTENT_MODERATION_TYPE, + level="none", + suggestion="pass", + results=[{"Label": "spam", "RiskLevel": "low"}], + ) + ], + }, + } + result = g._parse_response_and_check(response, check_type="input") + assert result["flagged"] is False + + def test_blocks_on_detail_block_suggestion_without_any_risk_level(self): + g = _make_guardrail(level="medium") + response = { + "Data": { + "Suggestion": "pass", + "Detail": [_make_detail(detection_type=PROMPT_ATTACK_TYPE, level=None, suggestion="block")], + }, + } + with pytest.raises(HTTPException) as exc_info: + g._parse_response_and_check(response, check_type="input") + assert PROMPT_ATTACK_TYPE in str(exc_info.value.detail) + + def test_blocks_on_overall_block_suggestion_without_any_risk_level(self): + g = _make_guardrail(level="medium") + response = { + "Data": { + "Suggestion": "block", + "Detail": [_make_detail(detection_type=CONTENT_MODERATION_TYPE, level=None, suggestion="pass")], + }, + } + with pytest.raises(HTTPException): + g._parse_response_and_check(response, check_type="input") + + def test_blocks_on_overall_block_suggestion_with_empty_detail(self): + g = _make_guardrail(level="medium") + response = {"Data": {"Suggestion": "block", "Detail": []}} + with pytest.raises(HTTPException): + g._parse_response_and_check(response, check_type="input") + + def test_observation_mode_never_blocks_on_block_suggestion(self): + """level=max is an explicit opt-in to logging only; fail-closed handling of an + unparseable severity must still honour it rather than bypass the threshold.""" + g = _make_guardrail(level="max") + response = { + "Data": { + "Suggestion": "block", + "Detail": [_make_detail(detection_type=CONTENT_MODERATION_TYPE, level=None, suggestion="block")], + }, + } + result = g._parse_response_and_check(response, check_type="input") + assert result["flagged"] is True + + def test_documented_response_shape_blocks(self): + """Regression guard for the exact payload in the integration README, which + carries no Detail[].Level at all.""" + g = _make_guardrail(level="medium") + response = { + "RequestId": "xxx", + "Code": 200, + "Data": { + "Suggestion": "block", + "Detail": [ + { + "Type": "contentModeration", + "Suggestion": "block", + "Result": [{"Label": "violence", "Confidence": 99.5, "RiskLevel": "high"}], + } + ], + }, + } + with pytest.raises(HTTPException) as exc_info: + g._parse_response_and_check(response, check_type="input") + assert exc_info.value.status_code == 400 + assert "contentModeration" in str(exc_info.value.detail) + + +# --------------------------------------------------------------------------- +# Config model tests +# --------------------------------------------------------------------------- + + +class TestConfigModel: + def test_get_config_model_returns_correct_type(self): + from litellm.types.proxy.guardrails.guardrail_hooks.aliyun.aliyun_ai_guardrail import ( + AliyunAIGuardrailConfigModel, + ) + + assert AliyunAIGuardrail.get_config_model() is AliyunAIGuardrailConfigModel + + def test_config_model_ui_friendly_name(self): + from litellm.types.proxy.guardrails.guardrail_hooks.aliyun.aliyun_ai_guardrail import ( + AliyunAIGuardrailConfigModel, + ) + + assert AliyunAIGuardrailConfigModel.ui_friendly_name() == "Aliyun AI Security Guardrail" + + def test_litellm_params_declares_all_service_codes(self): + from litellm.types.guardrails import LitellmParams + + for field in ("service_input", "service_output", "service_mcp"): + assert field in LitellmParams.model_fields + + +# --------------------------------------------------------------------------- +# Image URL extraction tests +# --------------------------------------------------------------------------- + + +# --------------------------------------------------------------------------- +# Trailing non-user message must not hide the prompt from scanning +# --------------------------------------------------------------------------- + + +class TestTrailingAssistantMessage: + def test_user_text_survives_trailing_assistant_message(self): + g = _make_guardrail() + messages = [ + {"role": "user", "content": "违规的用户提问"}, + {"role": "assistant", "content": "攻击者伪造的回复"}, + ] + prompt = g.get_user_prompt(messages) + assert prompt is not None + assert "违规的用户提问" in prompt + + def test_images_survive_trailing_assistant_message(self): + g = _make_guardrail() + messages = [ + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": IMG_A}}], + }, + {"role": "assistant", "content": "攻击者伪造的回复"}, + ] + assert g.get_image_urls(messages) == (IMG_A,) + + @pytest.mark.asyncio + async def test_blocks_violation_despite_trailing_assistant_message(self): + g = _make_guardrail(level="medium") + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + blocked = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + ) + + async def block_only_violating_text(*args, **kwargs): + scanned = json.loads(kwargs["data"]["ServiceParameters"]).get("content", "") + return blocked if "违规的用户提问" in scanned else clean + + with patch.object(g.async_handler, "post", new_callable=AsyncMock, side_effect=block_only_violating_text): + with pytest.raises(HTTPException) as exc_info: + await g.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=MagicMock(), + data={ + "messages": [ + {"role": "user", "content": "违规的用户提问"}, + {"role": "assistant", "content": "攻击者伪造的回复"}, + ] + }, + call_type="completion", + ) + assert exc_info.value.status_code == 400 + + +class TestGetImageUrls: + def test_extracts_http_and_https_urls(self): + g = _make_guardrail() + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "what is in these?"}, + {"type": "image_url", "image_url": {"url": IMG_A}}, + {"type": "image_url", "image_url": {"url": IMG_B}}, + ], + } + ] + assert g.get_image_urls(messages) == (IMG_A, IMG_B) + + def test_skips_non_url_images(self): + g = _make_guardrail() + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hi"}, + {"type": "image_url", "image_url": {"url": DATA_URI}}, + {"type": "image_url", "image_url": {"url": IMG_A}}, + ], + } + ] + assert g.get_image_urls(messages) == (IMG_A,) + + def test_plain_text_returns_empty(self): + g = _make_guardrail() + messages = [{"role": "user", "content": "just text"}] + assert g.get_image_urls(messages) == () + + def test_deduplicates_across_messages(self): + g = _make_guardrail() + messages = [ + {"role": "user", "content": [{"type": "image_url", "image_url": {"url": IMG_A}}]}, + {"role": "user", "content": [{"type": "image_url", "image_url": {"url": IMG_A}}]}, + ] + assert g.get_image_urls(messages) == (IMG_A,) + + def test_collects_images_from_every_user_message(self): + g = _make_guardrail() + messages = [ + {"role": "user", "content": [{"type": "image_url", "image_url": {"url": IMG_B}}]}, + {"role": "assistant", "content": "ok"}, + {"role": "user", "content": [{"type": "image_url", "image_url": {"url": IMG_A}}]}, + ] + assert g.get_image_urls(messages) == (IMG_B, IMG_A) + + def test_empty_messages_returns_empty(self): + g = _make_guardrail() + assert g.get_image_urls([]) == () + + +# --------------------------------------------------------------------------- +# ServiceParameters construction tests +# --------------------------------------------------------------------------- + + +class TestServiceParametersConstruction: + @pytest.mark.asyncio + async def test_text_only(self): + g = _make_guardrail(service_input="query_security_check") + with patch.object( + g.async_handler, "post", new_callable=AsyncMock, return_value=_make_aliyun_api_response() + ) as mock_post: + await g.async_make_request(text="hello", service_type="input") + sp, service = _captured_service_parameters(mock_post) + assert sp == {"requestFrom": "LiteLLM", "content": "hello"} + assert service == "query_security_check" + + @pytest.mark.asyncio + async def test_text_and_images(self): + g = _make_guardrail(service_input="text_img_mix_guard") + with patch.object( + g.async_handler, "post", new_callable=AsyncMock, return_value=_make_aliyun_api_response() + ) as mock_post: + await g.async_make_request(text="hello", service_type="input", image_urls=[IMG_A, IMG_B]) + sp, service = _captured_service_parameters(mock_post) + assert sp == {"requestFrom": "LiteLLM", "content": "hello", "imageUrls": [IMG_A, IMG_B]} + assert service == "text_img_mix_guard" + + @pytest.mark.asyncio + async def test_images_only(self): + g = _make_guardrail(service_input="img_query_security_check") + with patch.object( + g.async_handler, "post", new_callable=AsyncMock, return_value=_make_aliyun_api_response() + ) as mock_post: + await g.async_make_request(service_type="input", image_urls=[IMG_A]) + sp, service = _captured_service_parameters(mock_post) + assert sp == {"requestFrom": "LiteLLM", "imageUrls": [IMG_A]} + assert service == "img_query_security_check" + + +# --------------------------------------------------------------------------- +# Pre-call hook tests (text + multimodal) +# --------------------------------------------------------------------------- + + +class TestPreCallHook: + @pytest.mark.asyncio + async def test_scans_responses_api_string_input(self): + g = _make_guardrail(level="medium") + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + 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={"input": "违规的 responses 输入"}, + call_type="responses", + ) + + scanned = "".join( + json.loads(call.kwargs["data"]["ServiceParameters"]).get("content", "") for call in mock_post.call_args_list + ) + assert "违规的 responses 输入" in scanned + + @pytest.mark.asyncio + async def test_blocks_violating_responses_api_input(self): + g = _make_guardrail(level="medium") + blocked = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + ) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=blocked): + with pytest.raises(HTTPException) as exc_info: + await g.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=MagicMock(), + data={"input": "违规的 responses 输入"}, + call_type="responses", + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_scans_responses_api_structured_input_with_image(self): + g = _make_guardrail(level="medium") + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + data = { + "input": [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "结构化输入文本"}, + {"type": "input_image", "image_url": IMG_A, "detail": "auto"}, + ], + } + ] + } + 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", + ) + + sent = [json.loads(call.kwargs["data"]["ServiceParameters"]) for call in mock_post.call_args_list] + 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_scans_responses_api_function_call_output_image(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": [ + {"type": "input_image", "image_url": IMG_A, "detail": "auto"}, + ], + } + ] + } + 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", + ) + + sent = [json.loads(call.kwargs["data"]["ServiceParameters"]) for call in mock_post.call_args_list] + assert any(IMG_A in (service_parameters.get("imageUrls") or []) for service_parameters in sent) + + @pytest.mark.asyncio + async def test_scans_responses_api_function_call_arguments(self): + g = _make_guardrail(level="medium") + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + data = { + "input": [ + { + "type": "function_call", + "call_id": "call_1", + "name": "send_message", + "arguments": '{"text":"函数调用里的违规参数"}', + } + ] + } + 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") + mock_api_response = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + ) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=mock_api_response): + with pytest.raises(HTTPException) as exc_info: + await g.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=MagicMock(), + data={"messages": [{"role": "user", "content": "违规内容"}]}, + call_type="completion", + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_passes_clean_content(self): + g = _make_guardrail(level="medium") + mock_api_response = _make_aliyun_api_response(suggestion="pass", detail=[]) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=mock_api_response): + result = await g.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=MagicMock(), + data={"messages": [{"role": "user", "content": "你好"}]}, + call_type="completion", + ) + assert result is None + + @pytest.mark.asyncio + async def test_no_messages_returns_data(self): + g = _make_guardrail() + result = await g.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=MagicMock(), + data={}, + call_type="completion", + ) + assert result == {} + + @pytest.mark.asyncio + async def test_no_user_prompt_returns_none(self): + g = _make_guardrail() + result = await g.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=MagicMock(), + data={"messages": [{"role": "system", "content": "system prompt only"}]}, + call_type="completion", + ) + assert result is None + + @pytest.mark.asyncio + async def test_long_text_splits_into_multiple_requests(self): + g = _make_guardrail(max_text_length=10, level="medium") + mock_api_response = _make_aliyun_api_response(suggestion="pass", detail=[]) + long_content = "This is a very long text that exceeds the max_text_length limit." + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=mock_api_response) as mock_post: + await g.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=MagicMock(), + data={"messages": [{"role": "user", "content": long_content}]}, + call_type="completion", + ) + assert mock_post.call_count >= 2 + + @pytest.mark.asyncio + async def test_blocks_image_violation(self): + g = _make_guardrail(level="medium", service_input="text_img_mix_guard") + mock_resp = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + ) + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "look"}, + {"type": "image_url", "image_url": {"url": IMG_A}}, + ], + } + ] + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=mock_resp): + with pytest.raises(HTTPException) as exc_info: + await g.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=MagicMock(), + data={"messages": messages}, + call_type="completion", + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_image_only_request_sends_imageurls(self): + g = _make_guardrail(level="medium", service_input="img_query_security_check") + mock_resp = _make_aliyun_api_response(suggestion="pass", detail=[]) + messages = [{"role": "user", "content": [{"type": "image_url", "image_url": {"url": IMG_A}}]}] + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=mock_resp) as mock_post: + await g.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=MagicMock(), + data={"messages": messages}, + call_type="completion", + ) + assert mock_post.call_count == 1 + sp, _ = _captured_service_parameters(mock_post) + assert sp == {"requestFrom": "LiteLLM", "imageUrls": [IMG_A]} + + @pytest.mark.asyncio + async def test_first_segment_carries_images(self): + g = _make_guardrail(max_text_length=10, level="medium", service_input="text_img_mix_guard") + mock_resp = _make_aliyun_api_response(suggestion="pass", detail=[]) + long_text = "This is a very long text exceeding the limit for sure." + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": long_text}, + {"type": "image_url", "image_url": {"url": IMG_A}}, + ], + } + ] + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=mock_resp) as mock_post: + await g.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=MagicMock(), + data={"messages": messages}, + call_type="completion", + ) + assert mock_post.call_count >= 2 + image_carrying = 0 + for call in mock_post.call_args_list: + sp = json.loads(call.kwargs["data"]["ServiceParameters"]) + if "imageUrls" in sp: + image_carrying += 1 + assert image_carrying == 1 + + +# --------------------------------------------------------------------------- +# Post-call hook tests +# --------------------------------------------------------------------------- + + +class TestPostCallHook: + @pytest.mark.asyncio + async def test_blocks_violation_in_response(self): + import litellm + + g = _make_guardrail(level="medium") + mock_api_response = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + ) + response = litellm.ModelResponse( + id="test-id", + choices=[ + litellm.Choices( + index=0, + message=litellm.Message(role="assistant", content="违规响应内容"), + ) + ], + ) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=mock_api_response): + with pytest.raises(HTTPException) as exc_info: + await g.async_post_call_success_hook( + data={"messages": [{"role": "user", "content": "hi"}]}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_passes_clean_response(self): + import litellm + + g = _make_guardrail(level="medium") + mock_api_response = _make_aliyun_api_response(suggestion="pass", detail=[]) + response = litellm.ModelResponse( + id="test-id", + choices=[ + litellm.Choices( + index=0, + message=litellm.Message(role="assistant", content="正常的回复内容"), + ) + ], + ) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=mock_api_response): + result = await g.async_post_call_success_hook( + data={"messages": [{"role": "user", "content": "hi"}]}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + assert result is response + + @pytest.mark.asyncio + async def test_blocks_violation_in_later_choice(self): + import litellm + + g = _make_guardrail(level="medium") + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + blocked = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + ) + + async def block_only_violating_text(*args, **kwargs): + scanned = json.loads(kwargs["data"]["ServiceParameters"])["content"] + return blocked if "违规响应内容" in scanned else clean + + response = litellm.ModelResponse( + id="test-id", + choices=[ + litellm.Choices( + index=0, + message=litellm.Message(role="assistant", content="正常的回复内容"), + ), + litellm.Choices( + index=1, + message=litellm.Message(role="assistant", content="违规响应内容"), + ), + ], + ) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, side_effect=block_only_violating_text): + with pytest.raises(HTTPException) as exc_info: + await g.async_post_call_success_hook( + data={"messages": [{"role": "user", "content": "hi"}]}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_scans_content_from_every_choice(self): + import litellm + + g = _make_guardrail(level="medium", max_text_length=10000) + mock_api_response = _make_aliyun_api_response(suggestion="pass", detail=[]) + response = litellm.ModelResponse( + id="test-id", + choices=[ + litellm.Choices(index=0, message=litellm.Message(role="assistant", content="第一个回复")), + litellm.Choices(index=1, message=litellm.Message(role="assistant", content="第二个回复")), + ], + ) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=mock_api_response) as mock_post: + await g.async_post_call_success_hook( + data={"messages": [{"role": "user", "content": "hi"}]}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + + scanned = "".join( + json.loads(call.kwargs["data"]["ServiceParameters"])["content"] for call in mock_post.call_args_list + ) + assert "第一个回复" in scanned + assert "第二个回复" in scanned + + @pytest.mark.asyncio + async def test_non_model_response_passthrough(self): + g = _make_guardrail() + result = await g.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response="not a model response", + ) + assert result == "not a model response" + + @pytest.mark.asyncio + async def test_empty_content_response_passthrough(self): + import litellm + + g = _make_guardrail() + response = litellm.ModelResponse( + id="test-id", + choices=[ + litellm.Choices( + index=0, + message=litellm.Message(role="assistant", content=""), + ) + ], + ) + result = await g.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + assert result is response + + +def _make_tool_call_response(arguments: str, name: str = "send_email"): + """Build a non-streaming response whose only output is a tool call.""" + import litellm + from litellm.types.utils import ChatCompletionMessageToolCall, Function + + return litellm.ModelResponse( + id="test-id", + choices=[ + litellm.Choices( + index=0, + message=litellm.Message( + role="assistant", + content=None, + tool_calls=[ + ChatCompletionMessageToolCall( + id="call_1", + type="function", + function=Function(name=name, arguments=arguments), + ) + ], + ), + ) + ], + ) + + +def _make_responses_api_response(text: str): + """Build a non-streaming /v1/responses body carrying assistant text.""" + from litellm.types.llms.openai import ResponsesAPIResponse + from litellm.types.responses.main import GenericResponseOutputItem, OutputText + + return ResponsesAPIResponse( + id="resp_1", + created_at=0, + model="gpt-4o", + object="response", + output=[ + GenericResponseOutputItem( + type="message", + id="m1", + status="completed", + role="assistant", + content=[OutputText(type="output_text", text=text, annotations=[])], + ) + ], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + temperature=1.0, + top_p=1.0, + ) + + +class TestPostCallStructuredFields: + """The streaming path already audits tool calls, reasoning text and + /v1/responses output. Auditing only ``message.content`` here would let the + very same content reach the client unchecked whenever stream=False.""" + + @pytest.mark.asyncio + async def test_scans_tool_call_arguments(self): + g = _make_guardrail(level="medium") + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + response = _make_tool_call_response('{"body": "违规的工具参数"}') + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=clean) as mock_post: + await g.async_post_call_success_hook( + data={"messages": [{"role": "user", "content": "hi"}]}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + scanned = "".join( + json.loads(call.kwargs["data"]["ServiceParameters"])["content"] for call in mock_post.call_args_list + ) + assert "违规的工具参数" in scanned + + @pytest.mark.asyncio + async def test_blocks_violation_in_tool_call_arguments(self): + g = _make_guardrail(level="medium") + blocked = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + ) + response = _make_tool_call_response('{"body": "违规的工具参数"}') + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=blocked): + with pytest.raises(HTTPException) as exc_info: + await g.async_post_call_success_hook( + data={"messages": [{"role": "user", "content": "hi"}]}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_scans_reasoning_content(self): + import litellm + + g = _make_guardrail(level="medium") + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + response = litellm.ModelResponse( + id="test-id", + choices=[ + litellm.Choices( + index=0, + message=litellm.Message( + role="assistant", + content="正常的回复内容", + reasoning_content="推理过程里的违规内容", + ), + ) + ], + ) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=clean) as mock_post: + await g.async_post_call_success_hook( + data={"messages": [{"role": "user", "content": "hi"}]}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + scanned = "".join( + json.loads(call.kwargs["data"]["ServiceParameters"])["content"] for call in mock_post.call_args_list + ) + assert "推理过程里的违规内容" in scanned + + @pytest.mark.asyncio + async def test_scans_responses_api_output(self): + g = _make_guardrail(level="medium") + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + response = _make_responses_api_response("响应体里的违规内容") + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=clean) as mock_post: + await g.async_post_call_success_hook( + data={"input": "hi"}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + scanned = "".join( + json.loads(call.kwargs["data"]["ServiceParameters"])["content"] for call in mock_post.call_args_list + ) + assert "响应体里的违规内容" in scanned + + @pytest.mark.asyncio + async def test_blocks_violation_in_responses_api_output(self): + g = _make_guardrail(level="medium") + blocked = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + ) + response = _make_responses_api_response("响应体里的违规内容") + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=blocked): + with pytest.raises(HTTPException) as exc_info: + await g.async_post_call_success_hook( + data={"input": "hi"}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + assert exc_info.value.status_code == 400 + + +# --------------------------------------------------------------------------- +# Post-MCP hook tests +# --------------------------------------------------------------------------- + + +def _make_call_tool_result(text: str = "tool output"): + from mcp.types import CallToolResult, TextContent + + return CallToolResult(content=[TextContent(type="text", text=text)], isError=False) + + +class TestExtractMcpToolText: + """A tool result carries text outside of ``content[].text``. Auditing only that + field would release structured payloads and embedded resources unchecked.""" + + def test_collects_structured_content(self): + from mcp.types import CallToolResult + + g = _make_guardrail() + result = CallToolResult(content=[], structuredContent={"note": "结构化字段里的违规内容"}, isError=False) + assert "结构化字段里的违规内容" in g._extract_mcp_tool_text(result) + + def test_collects_structured_content_alongside_text(self): + from mcp.types import CallToolResult, TextContent + + g = _make_guardrail() + result = CallToolResult( + content=[TextContent(type="text", text="正常的工具输出")], + structuredContent={"note": "结构化字段里的违规内容"}, + isError=False, + ) + extracted = g._extract_mcp_tool_text(result) + assert "正常的工具输出" in extracted + assert "结构化字段里的违规内容" in extracted + + def test_collects_embedded_resource_text(self): + from mcp.types import CallToolResult, EmbeddedResource, TextResourceContents + + g = _make_guardrail() + result = CallToolResult( + content=[ + EmbeddedResource( + type="resource", + resource=TextResourceContents( + uri="file:///tmp/note.txt", + mimeType="text/plain", + text="内嵌资源里的违规内容", + ), + ) + ], + isError=False, + ) + assert "内嵌资源里的违规内容" in g._extract_mcp_tool_text(result) + + def test_collects_structured_content_from_dict_payload(self): + g = _make_guardrail() + payload = {"content": [], "structuredContent": {"note": "结构化字段里的违规内容"}} + assert "结构化字段里的违规内容" in g._extract_mcp_tool_text(payload) + + def test_collects_structured_content_from_coerced_tuple_list(self): + """MCPPostCallResponseObject coerces a CallToolResult into (field, value) pairs.""" + g = _make_guardrail() + payload = [("content", []), ("structuredContent", {"note": "结构化字段里的违规内容"}), ("isError", False)] + assert "结构化字段里的违规内容" in g._extract_mcp_tool_text(payload) + + @pytest.mark.asyncio + async def test_blocks_violation_in_structured_content(self): + from mcp.types import CallToolResult, TextContent + + g = _make_guardrail(level="medium") + tool_result = CallToolResult( + content=[TextContent(type="text", text="正常的工具输出")], + structuredContent={"note": "结构化字段里的违规内容"}, + isError=False, + ) + kwargs, response_obj = _make_post_mcp_hook_args(tool_result) + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + blocked = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + ) + + async def block_only_structured_content(*args, **kwargs): + scanned = json.loads(kwargs["data"]["ServiceParameters"]).get("content", "") + return blocked if "结构化字段里的违规内容" in scanned else clean + + with patch.object(g.async_handler, "post", new_callable=AsyncMock, side_effect=block_only_structured_content): + await g.async_post_mcp_tool_call_hook( + kwargs=kwargs, + response_obj=response_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + 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 + + +# --------------------------------------------------------------------------- +# Streaming hook tests +# --------------------------------------------------------------------------- + + +def _make_stream_chunk(content=None, tool_call_arguments=None): + """Build a ModelResponseStream chunk carrying content and/or tool call arguments.""" + from litellm.types.utils import ( + ChatCompletionDeltaToolCall, + Delta, + Function, + ModelResponseStream, + StreamingChoices, + ) + + tool_calls = None + if tool_call_arguments is not None: + tool_calls = [ + ChatCompletionDeltaToolCall( + id="call_1", + type="function", + index=0, + function=Function(name="send_message", arguments=tool_call_arguments), + ) + ] + return ModelResponseStream( + id="test-id", + choices=[StreamingChoices(index=0, delta=Delta(content=content, tool_calls=tool_calls))], + ) + + +async def _aiter(chunks): + for chunk in chunks: + yield chunk + + +# --------------------------------------------------------------------------- +# MCP pre-call hook tests +# --------------------------------------------------------------------------- + + +class TestMcpPreCallCheck: + @pytest.mark.asyncio + async def test_blocks_violating_tool_arguments(self): + g = _make_guardrail(level="medium") + blocked = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + ) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=blocked): + with pytest.raises(HTTPException) as exc_info: + await g.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=MagicMock(), + data={"messages": [{"role": "user", "content": "delete_all_files 违规参数"}]}, + call_type="call_mcp_tool", + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_uses_mcp_service_code(self): + g = _make_guardrail(level="medium", service_mcp="text_img_mix_guard") + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=clean) as mock_post: + result = await g.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=MagicMock(), + data={"messages": [{"role": "user", "content": "send_message hello"}]}, + call_type="call_mcp_tool", + ) + + assert result is None + service_parameters, service = _captured_service_parameters(mock_post) + assert service == "text_img_mix_guard" + assert service_parameters["content"] == "send_message hello" + + @pytest.mark.asyncio + async def test_logs_mcp_content_length_without_content(self): + g = _make_guardrail(level="medium") + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + sensitive_content = "send_message token=secret-value" + logger_path = "litellm.proxy.guardrails.guardrail_hooks.aliyun.aliyun_ai_guardrail.verbose_proxy_logger.info" + with ( + patch(logger_path) as mock_info, + patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=clean), + ): + await g.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=MagicMock(), + data={"messages": [{"role": "user", "content": sensitive_content}]}, + call_type="call_mcp_tool", + ) + + mock_info.assert_any_call( + "Aliyun AI Guardrail: ★ MCP pre-call check started, content length: %d", + len(sensitive_content), + ) + assert sensitive_content not in str(mock_info.call_args_list) + + @pytest.mark.asyncio + async def test_empty_content_skips_check(self): + g = _make_guardrail() + with patch.object(g.async_handler, "post", new_callable=AsyncMock) as mock_post: + result = await g.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=MagicMock(), + data={"messages": [{"role": "user", "content": ""}]}, + call_type="call_mcp_tool", + ) + + assert result is None + mock_post.assert_not_called() + + @pytest.mark.asyncio + async def test_long_content_is_split_into_multiple_requests(self): + g = _make_guardrail(max_text_length=10) + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + 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={"messages": [{"role": "user", "content": "一二三四五六七八九十。壹贰叁肆伍陆柒捌玖抾。"}]}, + call_type="call_mcp_tool", + ) + + assert mock_post.call_count > 1 + + +# --------------------------------------------------------------------------- +# post_mcp_call event hook gating tests +# --------------------------------------------------------------------------- + + +class TestShouldRunPostMcpCall: + def test_no_event_hook_runs(self): + g = _make_guardrail() + g.event_hook = None + assert g._should_run_post_mcp_call() is True + + def test_list_containing_post_mcp_call_runs(self): + g = _make_guardrail() + g.event_hook = ["pre_call", "post_mcp_call"] + assert g._should_run_post_mcp_call() is True + + def test_list_without_post_mcp_call_skips(self): + g = _make_guardrail() + g.event_hook = ["pre_call", "post_call"] + assert g._should_run_post_mcp_call() is False + + def test_plain_string_is_matched(self): + g = _make_guardrail() + g.event_hook = "post_mcp_call" + assert g._should_run_post_mcp_call() is True + g.event_hook = "pre_call" + assert g._should_run_post_mcp_call() is False + + def test_mode_tags_are_matched(self): + from litellm.types.guardrails import Mode + + g = _make_guardrail() + g.event_hook = Mode(tags={"team-a": ["pre_call", "post_mcp_call"]}) + assert g._should_run_post_mcp_call() is True + + g.event_hook = Mode(tags={"team-a": "post_mcp_call"}) + assert g._should_run_post_mcp_call() is True + + def test_mode_falls_back_to_default(self): + from litellm.types.guardrails import Mode + + g = _make_guardrail() + g.event_hook = Mode(tags={"team-a": "pre_call"}, default=["post_mcp_call"]) + assert g._should_run_post_mcp_call() is True + + g.event_hook = Mode(tags={"team-a": "pre_call"}, default="pre_call") + assert g._should_run_post_mcp_call() is False + + def test_mode_without_match_or_default_skips(self): + from litellm.types.guardrails import Mode + + g = _make_guardrail() + g.event_hook = Mode(tags={"team-a": "pre_call"}) + assert g._should_run_post_mcp_call() is False + + +# --------------------------------------------------------------------------- +# MCP payload shape handling tests +# --------------------------------------------------------------------------- + + +class TestIterMcpContentItems: + def test_plain_string_is_wrapped(self): + g = _make_guardrail() + assert g._iter_mcp_content_items("hello") == ("hello",) + + def test_object_with_content_list(self): + g = _make_guardrail() + payload = MagicMock() + payload.content = ["a", "b"] + assert g._iter_mcp_content_items(payload) == ("a", "b") + + def test_dict_with_content_list(self): + g = _make_guardrail() + assert g._iter_mcp_content_items({"content": ["a"]}) == ("a",) + + def test_dict_without_content_returns_itself(self): + g = _make_guardrail() + payload = {"text": "no content key"} + assert g._iter_mcp_content_items(payload) == (payload,) + + def test_coerced_tuple_pairs_recover_real_content(self): + g = _make_guardrail() + payload = [("meta", None), ("content", ["real"]), ("isError", False)] + assert g._iter_mcp_content_items(payload) == ("real",) + + def test_plain_list_passes_through(self): + g = _make_guardrail() + assert g._iter_mcp_content_items(["a", "b"]) == ("a", "b") + + def test_unsupported_payload_returns_empty(self): + g = _make_guardrail() + assert g._iter_mcp_content_items(123) == () + + +class TestReplaceToolOutputInPlace: + def test_none_target_is_rejected(self): + g = _make_guardrail() + assert g._replace_tool_output_in_place(None, ["blocked"]) is False + + def test_object_content_is_overwritten_and_flagged(self): + g = _make_guardrail() + target = MagicMock() + target.content = ["original"] + target.isError = False + assert g._replace_tool_output_in_place(target, ["blocked"]) is True + assert target.content == ["blocked"] + assert target.isError is True + + def test_list_target_is_overwritten(self): + g = _make_guardrail() + target = ["original"] + assert g._replace_tool_output_in_place(target, ["blocked"]) is True + assert target == ["blocked"] + + def test_nested_result_content_is_overwritten(self): + g = _make_guardrail() + target = {"result": {"content": ["original"]}} + assert g._replace_tool_output_in_place(target, ["blocked"]) is True + assert target["result"]["content"] == ["blocked"] + + def test_dict_content_is_overwritten(self): + g = _make_guardrail() + target = {"content": ["original"]} + assert g._replace_tool_output_in_place(target, ["blocked"]) is True + assert target["content"] == ["blocked"] + + def test_unsupported_shape_is_rejected(self): + g = _make_guardrail() + assert g._replace_tool_output_in_place({"unrelated": 1}, ["blocked"]) is False + + +class TestStreamingHook: + @pytest.mark.asyncio + async def test_scans_tool_call_arguments(self): + g = _make_guardrail(level="medium", stream_first_check_step=1, stream_slide_step=1) + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + chunk = _make_stream_chunk(tool_call_arguments='{"text": "违规工具参数"}') + + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=clean) as mock_post: + async for _ in g.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=_aiter([chunk]), + request_data={}, + ): + pass + + scanned = "".join( + json.loads(call.kwargs["data"]["ServiceParameters"])["content"] for call in mock_post.call_args_list + ) + assert "违规工具参数" in scanned + + @pytest.mark.asyncio + async def test_blocks_violating_tool_call_arguments(self): + g = _make_guardrail(level="medium", stream_first_check_step=1, stream_slide_step=1) + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + blocked = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + ) + + async def block_only_violating_text(*args, **kwargs): + scanned = json.loads(kwargs["data"]["ServiceParameters"])["content"] + return blocked if "违规工具参数" in scanned else clean + + chunk = _make_stream_chunk(tool_call_arguments='{"text": "违规工具参数"}') + + with patch.object(g.async_handler, "post", new_callable=AsyncMock, side_effect=block_only_violating_text): + emitted = [ + chunk + async for chunk in g.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=_aiter([chunk]), + request_data={}, + ) + ] + + # Streaming blocks by emitting an SSE error event instead of raising + assert not any(getattr(item, "choices", None) for item in emitted if hasattr(item, "choices")) + + @pytest.mark.asyncio + async def test_scans_plain_content(self): + g = _make_guardrail(level="medium", stream_first_check_step=1, stream_slide_step=1) + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + chunk = _make_stream_chunk(content="普通流式内容") + + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=clean) as mock_post: + emitted = [ + item + async for item in g.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=_aiter([chunk]), + request_data={}, + ) + ] + + scanned = "".join( + json.loads(call.kwargs["data"]["ServiceParameters"])["content"] for call in mock_post.call_args_list + ) + assert "普通流式内容" in scanned + assert emitted == [chunk] + + @pytest.mark.asyncio + async def test_blocks_violation_in_oversized_delta_prefix(self): + g = _make_guardrail( + level="medium", + stream_window_size=10, + stream_first_check_step=1, + stream_slide_step=6, + ) + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + blocked = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + ) + chunk = _make_stream_chunk(content="违规前缀" + "正常内容" * 8) + + async def block_prefix(*args, **kwargs): + scanned = json.loads(kwargs["data"]["ServiceParameters"])["content"] + return blocked if "违规前缀" in scanned else clean + + with patch.object(g.async_handler, "post", new_callable=AsyncMock, side_effect=block_prefix): + emitted = [ + item + async for item in g.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=_aiter([chunk]), + request_data={}, + ) + ] + + assert chunk not in emitted + + @pytest.mark.asyncio + async def test_scans_all_windows_before_emitting_oversized_delta_once(self): + g = _make_guardrail( + level="medium", + stream_window_size=10, + stream_first_check_step=1, + stream_slide_step=6, + ) + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + chunk = _make_stream_chunk(content="开头内容" + "中间内容" * 6 + "结尾内容") + + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=clean) as mock_post: + emitted = [ + item + async for item in g.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=_aiter([chunk]), + request_data={}, + ) + ] + + scanned_windows = tuple( + json.loads(call.kwargs["data"]["ServiceParameters"])["content"] for call in mock_post.call_args_list + ) + assert all(len(window) <= g.stream_window_size for window in scanned_windows) + assert any("开头内容" in window for window in scanned_windows) + assert any("结尾内容" in window for window in scanned_windows) + assert emitted == [chunk] + + @pytest.mark.asyncio + async def test_scans_responses_api_text_delta(self): + from litellm.types.llms.openai import OutputTextDeltaEvent + + g = _make_guardrail(level="medium", stream_first_check_step=1, stream_slide_step=1) + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + chunk = OutputTextDeltaEvent( + type="response.output_text.delta", + item_id="item_1", + output_index=0, + content_index=0, + delta="违规的 responses 输出", + ) + + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=clean) as mock_post: + async for _ in g.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=_aiter([chunk]), + request_data={}, + ): + pass + + scanned = "".join( + json.loads(call.kwargs["data"]["ServiceParameters"])["content"] for call in mock_post.call_args_list + ) + assert "违规的 responses 输出" in scanned + + @pytest.mark.asyncio + async def test_blocks_violating_responses_api_text_delta(self): + from litellm.types.llms.openai import OutputTextDeltaEvent + + g = _make_guardrail(level="medium", stream_first_check_step=1, stream_slide_step=1) + blocked = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + ) + chunk = OutputTextDeltaEvent( + type="response.output_text.delta", + item_id="item_1", + output_index=0, + content_index=0, + delta="违规的 responses 输出", + ) + + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=blocked): + emitted = [ + item + async for item in g.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=_aiter([chunk]), + request_data={}, + ) + ] + + assert chunk not in emitted + + @pytest.mark.asyncio + async def test_scans_responses_api_completed_event(self): + from litellm.types.llms.openai import ResponseCompletedEvent + + g = _make_guardrail(level="medium", stream_first_check_step=1, stream_slide_step=1) + clean = _make_aliyun_api_response(suggestion="pass", detail=[]) + content_part = MagicMock() + content_part.text = "完成事件里的违规内容" + output_item = MagicMock() + output_item.content = [content_part] + chunk = MagicMock(spec=ResponseCompletedEvent) + chunk.type = "response.completed" + chunk.response = MagicMock() + chunk.response.output = [output_item] + + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=clean) as mock_post: + async for _ in g.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=_aiter([chunk]), + request_data={}, + ): + pass + + scanned = "".join( + json.loads(call.kwargs["data"]["ServiceParameters"])["content"] for call in mock_post.call_args_list + ) + assert "完成事件里的违规内容" in scanned + + +def _make_post_mcp_hook_args(tool_result): + """Mimic how litellm_logging dispatches the post-MCP hook. + + The live CallToolResult is stored in model_call_details["original_response"] and + the same object is wrapped into MCPPostCallResponseObject, whose + mcp_tool_call_response field is declared as a list - so pydantic coerces the + CallToolResult by iterating it into (key, value) tuples. + """ + from litellm.types.llms.base import HiddenParams + from litellm.types.mcp import MCPPostCallResponseObject + + kwargs = {"original_response": tool_result} + response_obj = MCPPostCallResponseObject(mcp_tool_call_response=tool_result, hidden_params=HiddenParams()) + return kwargs, response_obj + + +class TestPostMcpToolCallHook: + @pytest.mark.asyncio + async def test_audits_tool_text_not_pydantic_tuple_repr(self): + """The wrapped response degrades into (key, value) tuples, so auditing it + verbatim would send Python reprs - and MCP envelope fields - to Aliyun + instead of the tool's own output.""" + g = _make_guardrail(level="medium") + tool_result = _make_call_tool_result("工具返回的内容") + kwargs, response_obj = _make_post_mcp_hook_args(tool_result) + mock_api_response = _make_aliyun_api_response(suggestion="pass", detail=[]) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=mock_api_response) as mock_post: + await g.async_post_mcp_tool_call_hook( + kwargs=kwargs, + response_obj=response_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + sp, service = _captured_service_parameters(mock_post) + assert sp["content"] == "工具返回的内容" + for envelope_field in ("isError", "structuredContent", "annotations", "TextContent"): + assert envelope_field not in sp["content"] + + @pytest.mark.asyncio + async def test_violation_replaces_tool_output_in_place(self): + """Both dispatch sites discard this hook's return value and hand the original + CallToolResult to the client, so the violation must be written into it.""" + g = _make_guardrail(level="medium") + tool_result = _make_call_tool_result("违规的工具输出") + kwargs, response_obj = _make_post_mcp_hook_args(tool_result) + mock_api_response = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=CONTENT_MODERATION_TYPE, level="high")], + ) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=mock_api_response): + await g.async_post_mcp_tool_call_hook( + kwargs=kwargs, + response_obj=response_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + remaining = " ".join(getattr(item, "text", "") for item in tool_result.content) + assert "违规的工具输出" not in remaining + assert CONTENT_MODERATION_TYPE in remaining + assert tool_result.isError is True + + @pytest.mark.asyncio + async def test_violation_returns_replacement_object_without_raising(self): + """Raising is a no-op here: the dispatcher catches every callback exception as a + non-blocking logging error and returns the untouched tool result.""" + from litellm.types.mcp import MCPPostCallResponseObject + + g = _make_guardrail(level="medium") + tool_result = _make_call_tool_result("违规的工具输出") + kwargs, response_obj = _make_post_mcp_hook_args(tool_result) + mock_api_response = _make_aliyun_api_response( + suggestion="block", + detail=[_make_detail(detection_type=PROMPT_ATTACK_TYPE, level="high")], + ) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=mock_api_response): + result = await g.async_post_mcp_tool_call_hook( + kwargs=kwargs, + response_obj=response_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + assert isinstance(result, MCPPostCallResponseObject) + replaced = " ".join(getattr(item, "text", "") for item in result.mcp_tool_call_response) + assert PROMPT_ATTACK_TYPE in replaced + + @pytest.mark.asyncio + async def test_clean_output_leaves_tool_result_untouched(self): + g = _make_guardrail(level="medium") + tool_result = _make_call_tool_result("正常的工具输出") + kwargs, response_obj = _make_post_mcp_hook_args(tool_result) + mock_api_response = _make_aliyun_api_response(suggestion="pass", detail=[]) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, return_value=mock_api_response): + result = await g.async_post_mcp_tool_call_hook( + kwargs=kwargs, + response_obj=response_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + assert result is None + assert tool_result.content[0].text == "正常的工具输出" + assert tool_result.isError is False + + @pytest.mark.asyncio + async def test_api_network_failure_fails_closed_in_place(self): + """A network error is swallowed by the dispatcher as a non-blocking logging + error, so unaudited tool output would reach the client untouched. The hook + must fail closed by replacing the live tool result instead of raising.""" + import httpx + + g = _make_guardrail(level="medium") + tool_result = _make_call_tool_result("未拦截将泄漏的敏感内容") + kwargs, response_obj = _make_post_mcp_hook_args(tool_result) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, side_effect=httpx.ConnectError("boom")): + await g.async_post_mcp_tool_call_hook( + kwargs=kwargs, + response_obj=response_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + remaining = " ".join(getattr(item, "text", "") for item in tool_result.content) + assert "未拦截将泄漏的敏感内容" not in remaining + assert "未经审核" in remaining + assert tool_result.isError is True + + @pytest.mark.asyncio + async def test_api_network_failure_returns_replacement_object(self): + """The failure notice must also be returned as the documented replacement + object, mirroring the violation path.""" + import httpx + + from litellm.types.mcp import MCPPostCallResponseObject + + g = _make_guardrail(level="medium") + tool_result = _make_call_tool_result("未拦截将泄漏的敏感内容") + kwargs, response_obj = _make_post_mcp_hook_args(tool_result) + with patch.object(g.async_handler, "post", new_callable=AsyncMock, side_effect=httpx.ConnectError("boom")): + result = await g.async_post_mcp_tool_call_hook( + kwargs=kwargs, + response_obj=response_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + assert isinstance(result, MCPPostCallResponseObject) + replaced = " ".join(getattr(item, "text", "") for item in result.mcp_tool_call_response) + assert "未经审核" in replaced + + @pytest.mark.asyncio + async def test_skipped_when_post_mcp_call_not_configured(self): + g = _make_guardrail(level="medium", event_hook=["pre_call"]) + tool_result = _make_call_tool_result("违规的工具输出") + kwargs, response_obj = _make_post_mcp_hook_args(tool_result) + with patch.object(g.async_handler, "post", new_callable=AsyncMock) as mock_post: + result = await g.async_post_mcp_tool_call_hook( + kwargs=kwargs, + response_obj=response_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + assert result is None + mock_post.assert_not_called() + assert tool_result.content[0].text == "违规的工具输出"