diff --git a/enterprise/enterprise_hooks/banned_keywords.py b/enterprise/enterprise_hooks/banned_keywords.py index 4df138939a..aafbb8cf40 100644 --- a/enterprise/enterprise_hooks/banned_keywords.py +++ b/enterprise/enterprise_hooks/banned_keywords.py @@ -11,6 +11,7 @@ from typing import Literal import litellm from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails._content_utils import iter_user_text from litellm.integrations.custom_logger import CustomLogger from litellm._logging import verbose_proxy_logger from fastapi import HTTPException @@ -73,10 +74,10 @@ class _ENTERPRISE_BannedKeywords(CustomLogger): - check if user id part of blocked list """ self.print_verbose("Inside Banned Keyword List Pre-Call Hook") - if call_type == "completion" and "messages" in data: - for m in data["messages"]: - if "content" in m and isinstance(m["content"], str): - self.test_violation(test_str=m["content"]) + if call_type == "completion": + # Covers multimodal list content + Responses-API input. + for text in iter_user_text(data): + self.test_violation(test_str=text) except HTTPException as e: raise e diff --git a/enterprise/enterprise_hooks/google_text_moderation.py b/enterprise/enterprise_hooks/google_text_moderation.py index 1f26d52adf..75829c41ca 100644 --- a/enterprise/enterprise_hooks/google_text_moderation.py +++ b/enterprise/enterprise_hooks/google_text_moderation.py @@ -12,6 +12,7 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails._content_utils import iter_user_text from litellm.types.utils import CallTypesLiteral @@ -94,11 +95,9 @@ class _ENTERPRISE_GoogleTextModeration(CustomLogger): - Calls Google's Text Moderation API - Rejects request if it fails safety check """ - if "messages" in data and isinstance(data["messages"], list): - text = "" - for m in data["messages"]: # assume messages is a list - if "content" in m and isinstance(m["content"], str): - text += m["content"] + # Covers multimodal list content + Responses-API input. + text = "".join(iter_user_text(data)) + if text: document = self.language_document(content=text, type_=self.document_type) request = self.moderate_text_request( diff --git a/enterprise/enterprise_hooks/openai_moderation.py b/enterprise/enterprise_hooks/openai_moderation.py index a1db9818e5..a98d5b30d7 100644 --- a/enterprise/enterprise_hooks/openai_moderation.py +++ b/enterprise/enterprise_hooks/openai_moderation.py @@ -19,6 +19,7 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails._content_utils import iter_user_text from litellm.types.utils import CallTypesLiteral @@ -37,11 +38,8 @@ class _ENTERPRISE_OpenAI_Moderation(CustomLogger): user_api_key_dict: UserAPIKeyAuth, call_type: CallTypesLiteral, ): - text = "" - if "messages" in data and isinstance(data["messages"], list): - for m in data["messages"]: # assume messages is a list - if "content" in m and isinstance(m["content"], str): - text += m["content"] + # Covers multimodal list content + Responses-API input. + text = "".join(iter_user_text(data)) from litellm.proxy.proxy_server import llm_router diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/secret_detection.py b/enterprise/litellm_enterprise/enterprise_callbacks/secret_detection.py index 8a7a82df68..f441ce71ab 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/secret_detection.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/secret_detection.py @@ -18,6 +18,7 @@ from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails._content_utils import walk_user_text GUARDRAIL_NAME = "hide_secrets" @@ -473,23 +474,19 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail): if await self.should_run_check(user_api_key_dict) is False: return - if "messages" in data and isinstance(data["messages"], list): - for message in data["messages"]: - if "content" in message and isinstance(message["content"], str): - detected_secrets = self.scan_message_for_secrets(message["content"]) + # Covers multimodal list content + Responses-API input. + def _redact_message_text(text: str) -> str: + detected_secrets = self.scan_message_for_secrets(text) + for secret in detected_secrets: + text = text.replace(secret["value"], "[REDACTED]") + if detected_secrets: + secret_types = [secret["type"] for secret in detected_secrets] + verbose_proxy_logger.warning( + f"Detected and redacted secrets in message: {secret_types}" + ) + return text - for secret in detected_secrets: - message["content"] = message["content"].replace( - secret["value"], "[REDACTED]" - ) - - if len(detected_secrets) > 0: - secret_types = [secret["type"] for secret in detected_secrets] - verbose_proxy_logger.warning( - f"Detected and redacted secrets in message: {secret_types}" - ) - else: - verbose_proxy_logger.debug("No secrets detected on input.") + walk_user_text(data, _redact_message_text) if "prompt" in data: if isinstance(data["prompt"], str): @@ -504,11 +501,15 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail): f"Detected and redacted secrets in prompt: {secret_types}" ) elif isinstance(data["prompt"], list): - for item in data["prompt"]: + # Index back into the list — assigning to ``item`` would only + # rebind the loop variable and leave ``data["prompt"]`` + # carrying the unredacted secret. + for idx, item in enumerate(data["prompt"]): if isinstance(item, str): detected_secrets = self.scan_message_for_secrets(item) for secret in detected_secrets: item = item.replace(secret["value"], "[REDACTED]") + data["prompt"][idx] = item if len(detected_secrets) > 0: secret_types = [ secret["type"] for secret in detected_secrets @@ -517,31 +518,6 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail): f"Detected and redacted secrets in prompt: {secret_types}" ) - if "input" in data: - if isinstance(data["input"], str): - detected_secrets = self.scan_message_for_secrets(data["input"]) - for secret in detected_secrets: - data["input"] = data["input"].replace(secret["value"], "[REDACTED]") - if len(detected_secrets) > 0: - secret_types = [secret["type"] for secret in detected_secrets] - verbose_proxy_logger.warning( - f"Detected and redacted secrets in input: {secret_types}" - ) - elif isinstance(data["input"], list): - _input_in_request = data["input"] - for idx, item in enumerate(_input_in_request): - if isinstance(item, str): - detected_secrets = self.scan_message_for_secrets(item) - for secret in detected_secrets: - _input_in_request[idx] = item.replace( - secret["value"], "[REDACTED]" - ) - if len(detected_secrets) > 0: - secret_types = [ - secret["type"] for secret in detected_secrets - ] - verbose_proxy_logger.warning( - f"Detected and redacted secrets in input: {secret_types}" - ) - verbose_proxy_logger.debug("Data after redacting input %s", data) + # ``data["input"]`` (Responses API and embeddings/moderation) is + # already covered by ``walk_user_text`` above. return diff --git a/litellm/proxy/guardrails/_content_utils.py b/litellm/proxy/guardrails/_content_utils.py new file mode 100644 index 0000000000..97a022768e --- /dev/null +++ b/litellm/proxy/guardrails/_content_utils.py @@ -0,0 +1,160 @@ +""" +Shared helpers for guardrail hooks: extract user-supplied text from a +request body regardless of whether it uses Chat Completions ``messages``, +Responses-API ``input``, or multimodal list-format ``content`` parts. + +Hooks that only check ``data["messages"]`` for string content silently +skip the other shapes — these helpers normalise that so every hook sees +every text fragment. +""" + +from typing import Any, Callable, Dict, Iterator, List + + +def _iter_text_parts_in_content(content: Any) -> Iterator[str]: + """Yield text fragments from a ``message.content`` value (string or + multimodal list). Non-text parts (images, audio, …) are skipped.""" + if isinstance(content, str): + if content: + yield content + elif isinstance(content, list): + for part in content: + if not isinstance(part, dict): + continue + if part.get("type") == "text": + text = part.get("text") + if isinstance(text, str) and text: + yield text + + +def _coerce_input_to_messages(input_value: Any) -> List[Dict[str, Any]]: + """Coerce a Responses-API ``data["input"]`` value into chat-style messages.""" + if isinstance(input_value, str): + return [{"role": "user", "content": input_value}] + if isinstance(input_value, list): + if input_value and all( + isinstance(item, dict) and "role" in item for item in input_value + ): + return list(input_value) + if input_value and all(isinstance(item, str) for item in input_value): + return [{"role": "user", "content": item} for item in input_value] + return [{"role": "user", "content": input_value}] + return [] + + +def _resolve_messages(data: Dict[str, Any]) -> List[Dict[str, Any]]: + """Return the messages list to inspect, falling back to ``input``.""" + messages = data.get("messages") + if isinstance(messages, list) and messages: + return messages + return _coerce_input_to_messages(data.get("input")) + + +def _iter_inspection_messages(data: Dict[str, Any]) -> Iterator[Dict[str, Any]]: + """Yield every message-like dict, walking ``messages`` AND ``input``.""" + messages = data.get("messages") + if isinstance(messages, list): + yield from messages + yield from _coerce_input_to_messages(data.get("input")) + + +def iter_user_text(data: Dict[str, Any]) -> Iterator[str]: + """Yield every user-supplied text fragment from ``messages`` and ``input``.""" + for message in _iter_inspection_messages(data): + if not isinstance(message, dict): + continue + yield from _iter_text_parts_in_content(message.get("content")) + + +def walk_user_text(data: Dict[str, Any], visit: Callable[[str], str]) -> int: + """Rewrite every user-supplied text fragment in place via ``visit``. + + Mutates ``data["messages"]`` and ``data["input"]``. Returns the number + of fragments visited so callers can short-circuit when nothing was + inspected. + """ + visited = 0 + + def _rewrite_content(content: Any) -> Any: + nonlocal visited + if isinstance(content, str): + if content: + visited += 1 + return visit(content) + return content + if isinstance(content, list): + new_parts = [] + for part in content: + if ( + isinstance(part, dict) + and part.get("type") == "text" + and isinstance(part.get("text"), str) + and part["text"] + ): + visited += 1 + new_parts.append({**part, "text": visit(part["text"])}) + else: + new_parts.append(part) + return new_parts + return content + + messages = data.get("messages") + if isinstance(messages, list): + for message in messages: + if isinstance(message, dict) and "content" in message: + message["content"] = _rewrite_content(message["content"]) + + input_value = data.get("input") + if isinstance(input_value, str): + if input_value: + visited += 1 + data["input"] = visit(input_value) + return visited + if isinstance(input_value, list): + # List of full messages: rewrite each message's content. + if input_value and all( + isinstance(item, dict) and "role" in item for item in input_value + ): + for item in input_value: + if "content" in item: + item["content"] = _rewrite_content(item["content"]) + return visited + # List of content parts or strings: rewrite in place. + for idx, item in enumerate(input_value): + if isinstance(item, str) and item: + visited += 1 + input_value[idx] = visit(item) + elif ( + isinstance(item, dict) + and item.get("type") == "text" + and isinstance(item.get("text"), str) + and item["text"] + ): + visited += 1 + input_value[idx] = {**item, "text": visit(item["text"])} + return visited + + return visited + + +def build_inspection_messages(data: Dict[str, Any]) -> List[Dict[str, str]]: + """Synthesize a chat-style messages list for posting to a guardrail API. + + Each returned message has a plain-string ``content`` — multimodal text + parts are joined with newlines and Responses-API ``input`` is lifted + into synthetic messages. Messages with no inspectable text are dropped. + + Hooks that POST ``{"messages": [...]}`` to an external service should + call this instead of ``data.get("messages", [])`` so the Responses API + and multimodal content are covered. + """ + flattened: List[Dict[str, str]] = [] + for message in _iter_inspection_messages(data): + if not isinstance(message, dict): + continue + text = "\n".join(_iter_text_parts_in_content(message.get("content"))) + if not text: + continue + role = message.get("role", "user") or "user" + flattened.append({"role": role, "content": text}) + return flattened diff --git a/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py b/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py index 1ae87e99c9..a386617428 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py @@ -22,6 +22,7 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails._content_utils import build_inspection_messages from litellm.types.utils import ( CallTypesLiteral, Choices, @@ -101,10 +102,11 @@ class AimGuardrail(CustomGuardrail): user_email=user_email, litellm_call_id=call_id, ) + # Covers multimodal list content + Responses-API input. response = await self.async_handler.post( f"{self.api_base}/fw/v1/analyze", headers=headers, - json={"messages": data.get("messages", [])}, + json={"messages": build_inspection_messages(data)}, ) response.raise_for_status() res = response.json() @@ -162,7 +164,7 @@ class AimGuardrail(CustomGuardrail): litellm_call_id=call_id, ), json={ - "messages": request_data.get("messages", []) + "messages": build_inspection_messages(request_data) + [{"role": "assistant", "content": output}] }, ) @@ -233,15 +235,27 @@ class AimGuardrail(CustomGuardrail): user_api_key_dict: UserAPIKeyAuth, response: Union[Any, ModelResponse, EmbeddingResponse, ImageResponse], ) -> Any: - if ( - isinstance(response, ModelResponse) - and response.choices - and isinstance(response.choices[0], Choices) - ): - content = response.choices[0].message.content or "" - aim_output_guardrail_result = await self.call_aim_guardrail_on_output( - data, content, hook="output", key_alias=user_api_key_dict.key_alias + if not (isinstance(response, ModelResponse) and response.choices): + return response + # Inspect every choice — when ``n>1`` the additional completions + # used to bypass Aim entirely because the hook only inspected + # ``choices[0]``. Run inspections concurrently so multi-completion + # responses don't pay an n× latency penalty. + choices_to_inspect = [c for c in response.choices if isinstance(c, Choices)] + if not choices_to_inspect: + return response + results = await asyncio.gather( + *( + self.call_aim_guardrail_on_output( + data, + choice.message.content or "", + hook="output", + key_alias=user_api_key_dict.key_alias, + ) + for choice in choices_to_inspect ) + ) + for choice, aim_output_guardrail_result in zip(choices_to_inspect, results): if aim_output_guardrail_result and aim_output_guardrail_result.get( "detection_message" ): @@ -252,7 +266,7 @@ class AimGuardrail(CustomGuardrail): if aim_output_guardrail_result and aim_output_guardrail_result.get( "redacted_output" ): - response.choices[0].message.content = aim_output_guardrail_result.get( + choice.message.content = aim_output_guardrail_result.get( "redacted_output" ) return response diff --git a/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/ibm_detector.py b/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/ibm_detector.py index 2fc0521364..435bb383df 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/ibm_detector.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/ibm_detector.py @@ -20,6 +20,7 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails._content_utils import iter_user_text from litellm.types.guardrails import GuardrailEventHooks from litellm.types.proxy.guardrails.guardrail_hooks.ibm import ( IBMDetectorDetection, @@ -463,65 +464,53 @@ class IBMGuardrailDetector(CustomGuardrail): if self.should_run_guardrail(data=data, event_type=event_type) is not True: return data - _messages = data.get("messages") - if _messages: - contents_to_check: List[str] = [] - for message in _messages: - _content = message.get("content") - if isinstance(_content, str): - contents_to_check.append(_content) + # Covers multimodal list content + Responses-API input. + contents_to_check: List[str] = list(iter_user_text(data)) + if contents_to_check: + if self.is_detector_server: + # Call detector server with all contents at once + result = await self._call_detector_server( + contents=contents_to_check, + request_data=data, + event_type=GuardrailEventHooks.pre_call, + ) - if contents_to_check: - if self.is_detector_server: - # Call detector server with all contents at once - result = await self._call_detector_server( - contents=contents_to_check, + verbose_proxy_logger.debug( + "IBM Detector Server async_pre_call_hook result: %s", result + ) + + # Check if any detections were found + has_violations = False + for message_detections in result: + filtered = self._filter_detections_by_threshold(message_detections) + if filtered: + has_violations = True + break + + if has_violations and self.block_on_detection: + error_message = self._create_error_message_detector_server(result) + raise ValueError(error_message) + + else: + # Call orchestrator for each content separately + for content in contents_to_check: + orchestrator_result = await self._call_orchestrator( + content=content, request_data=data, event_type=GuardrailEventHooks.pre_call, ) verbose_proxy_logger.debug( - "IBM Detector Server async_pre_call_hook result: %s", result + "IBM Orchestrator async_pre_call_hook result: %s", + orchestrator_result, ) - # Check if any detections were found - has_violations = False - for message_detections in result: - filtered = self._filter_detections_by_threshold( - message_detections - ) - if filtered: - has_violations = True - break - - if has_violations and self.block_on_detection: - error_message = self._create_error_message_detector_server( - result - ) - raise ValueError(error_message) - - else: - # Call orchestrator for each content separately - for content in contents_to_check: - orchestrator_result = await self._call_orchestrator( - content=content, - request_data=data, - event_type=GuardrailEventHooks.pre_call, - ) - - verbose_proxy_logger.debug( - "IBM Orchestrator async_pre_call_hook result: %s", - orchestrator_result, - ) - - filtered = self._filter_detections_by_threshold( + filtered = self._filter_detections_by_threshold(orchestrator_result) + if filtered and self.block_on_detection: + error_message = self._create_error_message_orchestrator( orchestrator_result ) - if filtered and self.block_on_detection: - error_message = self._create_error_message_orchestrator( - orchestrator_result - ) - raise ValueError(error_message) + raise ValueError(error_message) # Add guardrail to applied guardrails header add_guardrail_to_applied_guardrails_header( @@ -550,65 +539,53 @@ class IBMGuardrailDetector(CustomGuardrail): if self.should_run_guardrail(data=data, event_type=event_type) is not True: return - _messages = data.get("messages") - if _messages: - contents_to_check: List[str] = [] - for message in _messages: - _content = message.get("content") - if isinstance(_content, str): - contents_to_check.append(_content) + # Covers multimodal list content + Responses-API input. + contents_to_check: List[str] = list(iter_user_text(data)) + if contents_to_check: + if self.is_detector_server: + # Call detector server with all contents at once + result = await self._call_detector_server( + contents=contents_to_check, + request_data=data, + event_type=GuardrailEventHooks.during_call, + ) - if contents_to_check: - if self.is_detector_server: - # Call detector server with all contents at once - result = await self._call_detector_server( - contents=contents_to_check, + verbose_proxy_logger.debug( + "IBM Detector Server async_moderation_hook result: %s", result + ) + + # Check if any detections were found + has_violations = False + for message_detections in result: + filtered = self._filter_detections_by_threshold(message_detections) + if filtered: + has_violations = True + break + + if has_violations and self.block_on_detection: + error_message = self._create_error_message_detector_server(result) + raise ValueError(error_message) + + else: + # Call orchestrator for each content separately + for content in contents_to_check: + orchestrator_result = await self._call_orchestrator( + content=content, request_data=data, event_type=GuardrailEventHooks.during_call, ) verbose_proxy_logger.debug( - "IBM Detector Server async_moderation_hook result: %s", result + "IBM Orchestrator async_moderation_hook result: %s", + orchestrator_result, ) - # Check if any detections were found - has_violations = False - for message_detections in result: - filtered = self._filter_detections_by_threshold( - message_detections - ) - if filtered: - has_violations = True - break - - if has_violations and self.block_on_detection: - error_message = self._create_error_message_detector_server( - result - ) - raise ValueError(error_message) - - else: - # Call orchestrator for each content separately - for content in contents_to_check: - orchestrator_result = await self._call_orchestrator( - content=content, - request_data=data, - event_type=GuardrailEventHooks.during_call, - ) - - verbose_proxy_logger.debug( - "IBM Orchestrator async_moderation_hook result: %s", - orchestrator_result, - ) - - filtered = self._filter_detections_by_threshold( + filtered = self._filter_detections_by_threshold(orchestrator_result) + if filtered and self.block_on_detection: + error_message = self._create_error_message_orchestrator( orchestrator_result ) - if filtered and self.block_on_detection: - error_message = self._create_error_message_orchestrator( - orchestrator_result - ) - raise ValueError(error_message) + raise ValueError(error_message) # Add guardrail to applied guardrails header add_guardrail_to_applied_guardrails_header( diff --git a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py index 7aa26435ba..167ad4f790 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py @@ -13,6 +13,7 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails._content_utils import build_inspection_messages from litellm.secret_managers.main import get_secret_str from litellm.types.guardrails import GuardrailEventHooks from litellm.types.llms.openai import AllMessageValues @@ -214,10 +215,11 @@ class LakeraAIGuardrail(CustomGuardrail): ) return data - new_messages: Optional[List[AllMessageValues]] = data.get("messages") - if new_messages is None: + # Covers multimodal list content + Responses-API input. + new_messages = build_inspection_messages(data) + if not new_messages: verbose_proxy_logger.warning( - "Lakera AI: not running guardrail. No messages in data" + "Lakera AI: not running guardrail. No inspectable text in data" ) return data @@ -225,7 +227,7 @@ class LakeraAIGuardrail(CustomGuardrail): ########## 1. Make the Lakera AI v2 guard API request ########## ######################################################### lakera_guardrail_response, masked_entity_count = await self.call_v2_guard( - messages=new_messages, + messages=new_messages, # type: ignore[arg-type] request_data=data, event_type=GuardrailEventHooks.pre_call, ) @@ -280,10 +282,10 @@ class LakeraAIGuardrail(CustomGuardrail): if self.should_run_guardrail(data=data, event_type=event_type) is not True: return - new_messages: Optional[List[AllMessageValues]] = data.get("messages") - if new_messages is None: + new_messages = build_inspection_messages(data) + if not new_messages: verbose_proxy_logger.warning( - "Lakera AI: not running guardrail. No messages in data" + "Lakera AI: not running guardrail. No inspectable text in data" ) return @@ -291,7 +293,7 @@ class LakeraAIGuardrail(CustomGuardrail): ########## 1. Make the Lakera AI v2 guard API request ########## ######################################################### lakera_guardrail_response, masked_entity_count = await self.call_v2_guard( - messages=new_messages, + messages=new_messages, # type: ignore[arg-type] request_data=data, event_type=GuardrailEventHooks.during_call, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py index a72f3e4c3f..a7719124b2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py @@ -50,6 +50,7 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails._content_utils import build_inspection_messages from litellm.types.guardrails import GuardrailEventHooks import litellm @@ -366,7 +367,8 @@ class LassoGuardrail(CustomGuardrail): LassoGuardrailAPIError: If the Lasso API call fails HTTPException: If blocking violations are detected """ - messages: List[Dict[str, str]] = data.get("messages", []) + # Covers multimodal list content + Responses-API input. + messages: List[Dict[str, str]] = build_inspection_messages(data) if not messages: return data diff --git a/litellm/proxy/hooks/azure_content_safety.py b/litellm/proxy/hooks/azure_content_safety.py index b35d671117..e66ada4208 100644 --- a/litellm/proxy/hooks/azure_content_safety.py +++ b/litellm/proxy/hooks/azure_content_safety.py @@ -8,6 +8,7 @@ from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails._content_utils import iter_user_text class _PROXY_AzureContentSafety( @@ -118,10 +119,10 @@ class _PROXY_AzureContentSafety( ): verbose_proxy_logger.debug("Inside Azure Content-Safety Pre-Call Hook") try: - if call_type == "completion" and "messages" in data: - for m in data["messages"]: - if "content" in m and isinstance(m["content"], str): - await self.test_violation(content=m["content"], source="input") + if call_type == "completion": + # Covers multimodal list content + Responses-API input. + for text in iter_user_text(data): + await self.test_violation(content=text, source="input") except HTTPException as e: raise e diff --git a/tests/test_litellm/proxy/guardrails/test_content_utils.py b/tests/test_litellm/proxy/guardrails/test_content_utils.py new file mode 100644 index 0000000000..2491363330 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/test_content_utils.py @@ -0,0 +1,201 @@ +"""Tests for the shared guardrail content extraction helpers.""" + +from litellm.proxy.guardrails._content_utils import ( + build_inspection_messages, + iter_user_text, + walk_user_text, +) + + +# ── iter_user_text ──────────────────────────────────────────────────────────── + + +def test_iter_user_text_string_messages(): + data = { + "messages": [ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": "hi"}, + ] + } + assert list(iter_user_text(data)) == ["hello", "hi"] + + +def test_iter_user_text_multimodal_list_content(): + """VERIA-11: list-format content must be inspected, not silently skipped.""" + data = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "AWS_KEY=AKIA..."}, + {"type": "image_url", "image_url": {"url": "..."}}, + {"type": "text", "text": "more secrets"}, + ], + } + ] + } + assert list(iter_user_text(data)) == ["AWS_KEY=AKIA...", "more secrets"] + + +def test_iter_user_text_responses_api_string_input(): + """fniVO9-F: Responses-API ``input`` must be inspectable when ``messages`` absent.""" + data = {"input": "tell me a secret"} + assert list(iter_user_text(data)) == ["tell me a secret"] + + +def test_iter_user_text_responses_api_list_input_messages(): + data = { + "input": [ + {"role": "user", "content": "first"}, + {"role": "user", "content": "second"}, + ] + } + assert list(iter_user_text(data)) == ["first", "second"] + + +def test_iter_user_text_responses_api_list_input_content_parts(): + data = { + "input": [ + {"type": "text", "text": "alpha"}, + {"type": "image_url", "image_url": {"url": "..."}}, + {"type": "text", "text": "beta"}, + ] + } + assert list(iter_user_text(data)) == ["alpha", "beta"] + + +def test_iter_user_text_walks_messages_and_input_independently(): + """When both are present (rare), every fragment from either field is + inspected — a stricter guarantee than "first one wins".""" + data = { + "messages": [{"role": "user", "content": "msg-content"}], + "input": "input-content", + } + assert list(iter_user_text(data)) == ["msg-content", "input-content"] + + +def test_iter_user_text_empty_data(): + assert list(iter_user_text({})) == [] + assert list(iter_user_text({"messages": []})) == [] + assert list(iter_user_text({"input": ""})) == [] + + +# ── walk_user_text ──────────────────────────────────────────────────────────── + + +def test_walk_user_text_redacts_string_messages_in_place(): + data = { + "messages": [ + {"role": "user", "content": "leak: AKIAEXAMPLE"}, + {"role": "assistant", "content": "ok"}, + ] + } + visited = walk_user_text(data, lambda s: s.replace("AKIAEXAMPLE", "[REDACTED]")) + assert visited == 2 + assert data["messages"][0]["content"] == "leak: [REDACTED]" + assert data["messages"][1]["content"] == "ok" + + +def test_walk_user_text_redacts_multimodal_text_parts(): + """VERIA-11: list-content text parts must be mutable for in-place redaction.""" + data = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "AKIAEXAMPLE here"}, + {"type": "image_url", "image_url": {"url": "..."}}, + {"type": "text", "text": "no secret"}, + ], + } + ] + } + visited = walk_user_text(data, lambda s: s.replace("AKIAEXAMPLE", "[REDACTED]")) + assert visited == 2 + parts = data["messages"][0]["content"] + assert parts[0] == {"type": "text", "text": "[REDACTED] here"} + # Non-text part must be left untouched. + assert parts[1] == {"type": "image_url", "image_url": {"url": "..."}} + assert parts[2] == {"type": "text", "text": "no secret"} + + +def test_walk_user_text_redacts_responses_api_string_input(): + data = {"input": "leak AKIAEXAMPLE"} + visited = walk_user_text(data, lambda s: s.replace("AKIAEXAMPLE", "[REDACTED]")) + assert visited == 1 + assert data["input"] == "leak [REDACTED]" + + +def test_walk_user_text_redacts_responses_api_list_input(): + data = { + "input": [ + {"type": "text", "text": "AKIAEXAMPLE"}, + {"type": "image_url", "image_url": {"url": "..."}}, + ] + } + visited = walk_user_text(data, lambda s: f"[redacted]{s}[/]") + assert visited == 1 + assert data["input"][0] == {"type": "text", "text": "[redacted]AKIAEXAMPLE[/]"} + assert data["input"][1] == {"type": "image_url", "image_url": {"url": "..."}} + + +# ── build_inspection_messages ───────────────────────────────────────────────── + + +def test_build_inspection_messages_chat_completion_passthrough(): + data = { + "messages": [ + {"role": "system", "content": "be helpful"}, + {"role": "user", "content": "hi"}, + ] + } + assert build_inspection_messages(data) == [ + {"role": "system", "content": "be helpful"}, + {"role": "user", "content": "hi"}, + ] + + +def test_build_inspection_messages_joins_multimodal_text_parts(): + data = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "first part"}, + {"type": "image_url", "image_url": {"url": "..."}}, + {"type": "text", "text": "second part"}, + ], + } + ] + } + assert build_inspection_messages(data) == [ + {"role": "user", "content": "first part\nsecond part"} + ] + + +def test_build_inspection_messages_lifts_responses_api_input(): + """fniVO9-F: ``input`` must be visible to hooks that POST messages to a remote API.""" + data = {"input": "responses-api content"} + assert build_inspection_messages(data) == [ + {"role": "user", "content": "responses-api content"} + ] + + +def test_build_inspection_messages_drops_messages_with_no_text(): + data = { + "messages": [ + {"role": "user", "content": ""}, + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": "..."}}], + }, + {"role": "user", "content": "kept"}, + ] + } + assert build_inspection_messages(data) == [{"role": "user", "content": "kept"}] + + +def test_build_inspection_messages_empty_data(): + assert build_inspection_messages({}) == [] + assert build_inspection_messages({"messages": []}) == [] + assert build_inspection_messages({"input": ""}) == [] diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_coverage.py b/tests/test_litellm/proxy/guardrails/test_guardrail_coverage.py new file mode 100644 index 0000000000..4cfcbf3c6d --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_coverage.py @@ -0,0 +1,437 @@ +""" +Regression tests for guardrail-coverage gaps. + +Each test confirms that a previously-bypassable input shape now triggers +inspection by the relevant guardrail hook: + +- VERIA-11: multimodal list-format ``content`` is inspected (no longer + silently skipped because of an ``isinstance(content, str)`` check). +- fniVO9-F: Responses-API ``data["input"]`` is inspected (no longer + silently skipped because the hook only looked at ``data["messages"]``). +- yVS0wMDO: Aim's post-call hook inspects every choice when ``n>1``, + not just ``choices[0]``. +""" + +from typing import Any, Dict +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from httpx import Request, Response + +from litellm import DualCache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.utils import Choices, Message, ModelResponse + + +@pytest.fixture +def user_api_key(): + return UserAPIKeyAuth(api_key="hashed", user_id="u", key_alias=None) + + +# ── Aim ─────────────────────────────────────────────────────────────────────── + + +def _aim_no_action_response() -> Response: + return Response( + status_code=200, + json={"required_action": None}, + request=Request("POST", "https://api.aim.security/fw/v1/analyze"), + ) + + +@pytest.mark.asyncio +async def test_aim_inspects_multimodal_list_content(user_api_key, monkeypatch): + monkeypatch.setenv("AIM_API_KEY", "hs-aim-key") + from litellm.proxy.guardrails.guardrail_hooks.aim.aim import AimGuardrail + + guard = AimGuardrail() + sent_payload: Dict[str, Any] = {} + + async def capture(url, headers, json): + sent_payload.update(json) + return _aim_no_action_response() + + with patch.object(guard.async_handler, "post", side_effect=capture): + await guard.async_pre_call_hook( + user_api_key_dict=user_api_key, + cache=DualCache(), + data={ + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "secret payload"}, + {"type": "image_url", "image_url": {"url": "..."}}, + ], + } + ] + }, + call_type="acompletion", + ) + + # The multimodal text part must be visible to Aim. + assert sent_payload["messages"] == [{"role": "user", "content": "secret payload"}] + + +@pytest.mark.asyncio +async def test_aim_inspects_responses_api_input(user_api_key, monkeypatch): + monkeypatch.setenv("AIM_API_KEY", "hs-aim-key") + from litellm.proxy.guardrails.guardrail_hooks.aim.aim import AimGuardrail + + guard = AimGuardrail() + sent_payload: Dict[str, Any] = {} + + async def capture(url, headers, json): + sent_payload.update(json) + return _aim_no_action_response() + + with patch.object(guard.async_handler, "post", side_effect=capture): + await guard.async_pre_call_hook( + user_api_key_dict=user_api_key, + cache=DualCache(), + data={"input": "responses-api content"}, + call_type="acompletion", + ) + + assert sent_payload["messages"] == [ + {"role": "user", "content": "responses-api content"} + ] + + +@pytest.mark.asyncio +async def test_aim_post_call_inspects_all_choices(user_api_key, monkeypatch): + """yVS0wMDO: ``n>1`` no longer bypasses Aim by hiding violations in + ``choices[1+]``.""" + monkeypatch.setenv("AIM_API_KEY", "hs-aim-key") + from litellm.proxy.guardrails.guardrail_hooks.aim.aim import AimGuardrail + + guard = AimGuardrail() + inspected_outputs = [] + + async def capture(request_data, output, hook, key_alias): + inspected_outputs.append(output) + return {"redacted_output": output} + + response = ModelResponse( + choices=[ + Choices(index=0, message=Message(role="assistant", content="first")), + Choices(index=1, message=Message(role="assistant", content="second")), + Choices(index=2, message=Message(role="assistant", content="third")), + ] + ) + + with patch.object(guard, "call_aim_guardrail_on_output", side_effect=capture): + await guard.async_post_call_success_hook( + data={"messages": [{"role": "user", "content": "hi"}]}, + user_api_key_dict=user_api_key, + response=response, + ) + + # ``asyncio.gather`` is used for parallelism, so order of inspection is + # not guaranteed. + assert sorted(inspected_outputs) == ["first", "second", "third"] + + +# ── Lakera v2 ───────────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_lakera_v2_inspects_responses_api_input(user_api_key, monkeypatch): + monkeypatch.setenv("LAKERA_API_KEY", "lk-test") + from litellm.proxy.guardrails.guardrail_hooks.lakera_ai_v2 import ( + LakeraAIGuardrail, + ) + + guard = LakeraAIGuardrail(api_key="lk-test", on_flagged="monitor") + + seen_messages = [] + + async def fake_call_v2_guard(messages, request_data, event_type): + seen_messages.append(messages) + return {"flagged": False}, {} + + with patch.object(guard, "call_v2_guard", side_effect=fake_call_v2_guard): + await guard.async_pre_call_hook( + user_api_key_dict=user_api_key, + cache=DualCache(), + data={"input": "responses-api content"}, + call_type="responses", + ) + + assert seen_messages == [[{"role": "user", "content": "responses-api content"}]] + + +@pytest.mark.asyncio +async def test_lakera_v2_inspects_multimodal_list_content(user_api_key, monkeypatch): + monkeypatch.setenv("LAKERA_API_KEY", "lk-test") + from litellm.proxy.guardrails.guardrail_hooks.lakera_ai_v2 import ( + LakeraAIGuardrail, + ) + + guard = LakeraAIGuardrail(api_key="lk-test", on_flagged="monitor") + seen_messages = [] + + async def fake_call_v2_guard(messages, request_data, event_type): + seen_messages.append(messages) + return {"flagged": False}, {} + + with patch.object(guard, "call_v2_guard", side_effect=fake_call_v2_guard): + await guard.async_pre_call_hook( + user_api_key_dict=user_api_key, + cache=DualCache(), + data={ + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "AKIAEXAMPLE"}, + {"type": "image_url", "image_url": {"url": "..."}}, + ], + } + ] + }, + call_type="acompletion", + ) + + assert seen_messages == [[{"role": "user", "content": "AKIAEXAMPLE"}]] + + +# ── Lasso ───────────────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_lasso_inspects_responses_api_input(user_api_key, monkeypatch): + monkeypatch.setenv("LASSO_API_KEY", "ls-test") + from litellm.proxy.guardrails.guardrail_hooks.lasso.lasso import LassoGuardrail + + guard = LassoGuardrail(lasso_api_key="ls-test") + + seen_messages = [] + + async def fake_handle_classification(data, cache, message_type, messages): + seen_messages.append(messages) + return data + + with patch.object( + guard, "_handle_classification", side_effect=fake_handle_classification + ): + await guard._run_lasso_guardrail( + data={"input": "responses-api content"}, + cache=DualCache(), + message_type="PROMPT", + ) + + assert seen_messages == [[{"role": "user", "content": "responses-api content"}]] + + +# ── Banned Keywords ─────────────────────────────────────────────────────────── + + +def test_banned_keywords_blocks_multimodal_content(monkeypatch): + """VERIA-11: a banned word hidden in a multimodal text part is now caught.""" + monkeypatch.setattr("litellm.banned_keywords_list", ["forbidden"], raising=False) + from enterprise.enterprise_hooks.banned_keywords import _ENTERPRISE_BannedKeywords + from fastapi import HTTPException + + guard = _ENTERPRISE_BannedKeywords() + + async def _run(): + await guard.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="hashed", user_id="u"), + cache=DualCache(), + data={ + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "forbidden word here"}, + {"type": "image_url", "image_url": {"url": "..."}}, + ], + } + ] + }, + call_type="completion", + ) + + import asyncio + + with pytest.raises(HTTPException) as exc: + asyncio.run(_run()) + assert "forbidden" in str(exc.value.detail).lower() + + +def test_banned_keywords_blocks_responses_api_input(monkeypatch): + monkeypatch.setattr("litellm.banned_keywords_list", ["forbidden"], raising=False) + from enterprise.enterprise_hooks.banned_keywords import _ENTERPRISE_BannedKeywords + from fastapi import HTTPException + + guard = _ENTERPRISE_BannedKeywords() + + async def _run(): + await guard.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="hashed", user_id="u"), + cache=DualCache(), + data={"input": "this contains forbidden content"}, + call_type="completion", + ) + + import asyncio + + with pytest.raises(HTTPException): + asyncio.run(_run()) + + +# ── Secret Detection ────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_secret_detection_redacts_multimodal_text_parts(user_api_key): + from enterprise.litellm_enterprise.enterprise_callbacks.secret_detection import ( + _ENTERPRISE_SecretDetection, + ) + + guard = _ENTERPRISE_SecretDetection() + data = { + "messages": [ + { + "role": "user", + "content": [ + { + "type": "text", + "text": "AKIAIOSFODNN7EXAMPLE is the key", + }, + {"type": "image_url", "image_url": {"url": "..."}}, + ], + } + ] + } + + await guard.async_pre_call_hook( + user_api_key_dict=user_api_key, + cache=DualCache(), + data=data, + call_type="completion", + ) + + parts = data["messages"][0]["content"] + assert "AKIAIOSFODNN7EXAMPLE" not in parts[0]["text"] + assert "[REDACTED]" in parts[0]["text"] + # Non-text part is preserved untouched. + assert parts[1] == {"type": "image_url", "image_url": {"url": "..."}} + + +@pytest.mark.asyncio +async def test_secret_detection_redacts_responses_api_input(user_api_key): + from enterprise.litellm_enterprise.enterprise_callbacks.secret_detection import ( + _ENTERPRISE_SecretDetection, + ) + + guard = _ENTERPRISE_SecretDetection() + data = {"input": "leak: AKIAIOSFODNN7EXAMPLE"} + + await guard.async_pre_call_hook( + user_api_key_dict=user_api_key, + cache=DualCache(), + data=data, + call_type="moderation", + ) + + assert "AKIAIOSFODNN7EXAMPLE" not in data["input"] + assert "[REDACTED]" in data["input"] + + +# ── OpenAI Moderation ───────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_openai_moderation_inspects_multimodal_content(monkeypatch, user_api_key): + """The aggregated text passed to ``llm_router.amoderation`` must include + list-format text parts and Responses-API input — without this, multimodal + content silently passed moderation.""" + from enterprise.enterprise_hooks.openai_moderation import ( + _ENTERPRISE_OpenAI_Moderation, + ) + + guard = _ENTERPRISE_OpenAI_Moderation() + + seen_inputs = [] + + class FakeModeration: + results = [type("R", (), {"flagged": False})()] + + async def fake_amoderation(model, input): + seen_inputs.append(input) + return FakeModeration() + + fake_router = MagicMock() + fake_router.amoderation = AsyncMock(side_effect=fake_amoderation) + + monkeypatch.setattr( + "litellm.proxy.proxy_server.llm_router", fake_router, raising=False + ) + + await guard.async_moderation_hook( + data={ + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "alpha "}, + {"type": "image_url", "image_url": {"url": "..."}}, + {"type": "text", "text": "beta"}, + ], + } + ] + }, + user_api_key_dict=user_api_key, + call_type="acompletion", + ) + + assert seen_inputs == ["alpha beta"] + + +# ── Google Text Moderation ──────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_google_text_moderation_inspects_multimodal_content(user_api_key): + """The text passed to Google's moderation client must include list-format + text parts.""" + from enterprise.enterprise_hooks.google_text_moderation import ( + _ENTERPRISE_GoogleTextModeration, + ) + + guard = _ENTERPRISE_GoogleTextModeration.__new__(_ENTERPRISE_GoogleTextModeration) + seen_documents = [] + + def fake_language_document(content, type_): + seen_documents.append(content) + return MagicMock() + + fake_response = MagicMock() + fake_response.moderation_categories = [] + + guard.language_document = fake_language_document + guard.moderate_text_request = MagicMock(return_value=MagicMock()) + guard.document_type = MagicMock() + guard.client = MagicMock() + guard.client.moderate_text = MagicMock(return_value=fake_response) + + await guard.async_moderation_hook( + data={ + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello "}, + {"type": "image_url", "image_url": {"url": "..."}}, + {"type": "text", "text": "world"}, + ], + } + ] + }, + user_api_key_dict=user_api_key, + call_type="acompletion", + ) + + assert seen_documents == ["hello world"]