diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py new file mode 100644 index 00000000000..4be98a19db4 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py @@ -0,0 +1,42 @@ +from typing import TYPE_CHECKING + +from litellm.types.guardrails import SupportedGuardrailIntegrations + +from .model_armor import ModelArmorGuardrail + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"): + import litellm + from litellm.proxy.guardrails.guardrail_hooks.model_armor import ( + ModelArmorGuardrail, + ) + + _model_armor_callback = ModelArmorGuardrail( + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + template_id=litellm_params.template_id, + project_id=litellm_params.project_id, + location=litellm_params.location, + credentials=litellm_params.credentials, + api_endpoint=litellm_params.api_endpoint, + default_on=litellm_params.default_on, + mask_request_content=litellm_params.mask_request_content, + mask_response_content=litellm_params.mask_response_content, + fail_on_error=litellm_params.fail_on_error, + ) + litellm.logging_callback_manager.add_litellm_callback(_model_armor_callback) + + return _model_armor_callback + + +guardrail_initializer_registry = { + SupportedGuardrailIntegrations.MODEL_ARMOR.value: initialize_guardrail, +} + + +guardrail_class_registry = { + SupportedGuardrailIntegrations.MODEL_ARMOR.value: ModelArmorGuardrail, +} \ No newline at end of file diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py new file mode 100644 index 00000000000..101b0e76d46 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -0,0 +1,460 @@ +from typing import ( + TYPE_CHECKING, + Any, + AsyncGenerator, + List, + Literal, + Optional, + Type, + Union, +) + +from fastapi import HTTPException + +if TYPE_CHECKING: + from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.caching import DualCache +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) +from litellm.llms.vertex_ai.vertex_llm_base import VertexBase +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.utils import ( + Choices, + ModelResponse, + ModelResponseStream, +) + +GUARDRAIL_NAME = "model_armor" + + +class ModelArmorGuardrail(CustomGuardrail, VertexBase): + """ + Google Cloud Model Armor Guardrail integration for LiteLLM. + + Supports: + - Pre-call sanitization (sanitizeUserPrompt) + - Post-call sanitization (sanitizeModelResponse) + """ + + def __init__( + self, + template_id: Optional[str] = None, + project_id: Optional[str] = None, + location: Optional[str] = None, + credentials: Optional[Any] = None, + api_endpoint: Optional[str] = None, + **kwargs, + ): + self.async_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback + ) + self.template_id = template_id + self.project_id = project_id + self.location = location or "us-central1" + self.credentials = credentials + self.api_endpoint = api_endpoint + + # Store optional params + self.optional_params = kwargs + + super().__init__(**kwargs) + VertexBase.__init__(self) + + verbose_proxy_logger.debug( + "Model Armor Guardrail initialized with template_id: %s, project_id: %s, location: %s", + self.template_id, + self.project_id, + self.location, + ) + + def _get_api_endpoint(self) -> str: + """Get the API endpoint for Model Armor.""" + if self.api_endpoint: + return self.api_endpoint + return f"https://modelarmor.{self.location}.rep.googleapis.com" + + def _create_sanitize_request( + self, content: str, source: Literal["user_prompt", "model_response"] + ) -> dict: + """Create request body for Model Armor API.""" + if source == "user_prompt": + return {"user_prompt_data": {"text": content}} + else: + return {"model_response_data": {"text": content}} + + + + def _extract_content_from_response( + self, response: Union[Any, ModelResponse] + ) -> str: + """ + Extract text content from model response. + + Returns empty string for non-text responses (TTS, images, etc.) to skip guardrail processing. + """ + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + get_content_from_model_response, + ) + + # Handle ModelResponse objects + if isinstance(response, litellm.ModelResponse): + return get_content_from_model_response(response) + + # For non-ModelResponse types (e.g., TTS, images), return empty string + # These response types are not text-based and shouldn't be processed by text guardrails + verbose_proxy_logger.debug( + "Model Armor: Skipping non-ModelResponse type: %s", type(response).__name__ + ) + return "" + + async def make_model_armor_request( + self, + content: str, + source: Literal["user_prompt", "model_response"], + request_data: Optional[dict] = None, + ) -> dict: + """Make request to Model Armor API.""" + # Get access token using VertexBase auth + access_token, resolved_project_id = await self._ensure_access_token_async( + credentials=self.credentials, + project_id=self.project_id, + custom_llm_provider="vertex_ai", + ) + + # Use resolved project ID if not explicitly set + if not self.project_id and resolved_project_id: + self.project_id = resolved_project_id + + # Construct URL + endpoint = self._get_api_endpoint() + if source == "user_prompt": + url = f"{endpoint}/v1/projects/{self.project_id}/locations/{self.location}/templates/{self.template_id}:sanitizeUserPrompt" + else: + url = f"{endpoint}/v1/projects/{self.project_id}/locations/{self.location}/templates/{self.template_id}:sanitizeModelResponse" + + # Create request body + body = self._create_sanitize_request(content, source) + + # Set headers + headers = { + "Content-Type": "application/json", + "Authorization": f"Bearer {access_token}", + } + + verbose_proxy_logger.debug( + "Model Armor request - URL: %s, Body: %s", + url, + body, + ) + + # Make request + if self.async_handler is None: + raise ValueError("Async handler not initialized") + + response = await self.async_handler.post( + url=url, + json=body, + headers=headers, + ) + + verbose_proxy_logger.debug( + "Model Armor response - Status: %s, Body: %s", + response.status_code, + response.text, + ) + + if response.status_code != 200: + verbose_proxy_logger.error( + "Model Armor API error - Status: %s, Response: %s", + response.status_code, + response.text, + ) + raise HTTPException( + status_code=response.status_code, + detail=f"Model Armor API error: {response.text}", + ) + + json_response = response.json() + if hasattr(json_response, "__await__"): + return await json_response + return json_response + + def _should_block_content(self, armor_response: dict) -> bool: + """Check if Model Armor response indicates content should be blocked.""" + # Model Armor may return different response structures + # This is a basic implementation - adjust based on actual API response + if armor_response.get("blocked", False): + return True + + # Check for sanitization actions + if armor_response.get("action") == "BLOCK": + return True + + return False + + def _get_sanitized_content(self, armor_response: dict) -> Optional[str]: + """Extract sanitized content from Model Armor response.""" + # This depends on the actual Model Armor API response structure + # Adjust based on documentation + return armor_response.get("sanitized_text") or armor_response.get("text") + + @log_guardrail_information + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, + call_type: Literal[ + "completion", + "text_completion", + "embeddings", + "image_generation", + "moderation", + "audio_transcription", + "pass_through_endpoint", + "rerank", + ], + ) -> Union[Exception, str, dict, None]: + """Pre-call hook to sanitize user prompts.""" + verbose_proxy_logger.debug("Inside Model Armor Pre-Call Hook") + + from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_to_applied_guardrails_header, + ) + + event_type = GuardrailEventHooks.pre_call + if self.should_run_guardrail(data=data, event_type=event_type) is not True: + return data + + messages = data.get("messages") + if not messages: + verbose_proxy_logger.warning( + "Model Armor: not running guardrail. No messages in data" + ) + return data + + # Extract content from messages using helper from common_utils + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + get_last_user_message, + ) + + content = get_last_user_message(messages) + if not content: + return data + + # Make Model Armor request + try: + armor_response = await self.make_model_armor_request( + content=content, + source="user_prompt", + request_data=data, + ) + + # Check if content should be blocked + if self._should_block_content(armor_response): + raise HTTPException( + status_code=400, + detail={ + "error": "Content blocked by Model Armor", + "model_armor_response": armor_response, + }, + ) + + # If mask_request_content is enabled, update messages with sanitized content + if self.mask_request_content: + sanitized_content = self._get_sanitized_content(armor_response) + if sanitized_content and sanitized_content != content: + # Use the helper to set the last user message with sanitized content + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + set_last_user_message, + ) + + data["messages"] = set_last_user_message( + messages, sanitized_content + ) + + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.error( + "Model Armor pre-call error: %s", str(e), exc_info=True + ) + # Depending on configuration, either fail or continue + if self.optional_params.get("fail_on_error", True): + raise + + # Add guardrail to headers + add_guardrail_to_applied_guardrails_header( + request_data=data, guardrail_name=self.guardrail_name + ) + + return data + + @log_guardrail_information + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + response, + ): + """Post-call hook to sanitize model responses.""" + from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_to_applied_guardrails_header, + ) + + if ( + self.should_run_guardrail( + data=data, event_type=GuardrailEventHooks.post_call + ) + is not True + ): + return + + # Extract content from response + content = self._extract_content_from_response(response) + if not content: + verbose_proxy_logger.debug( + "Model Armor: No text content to process in response, skipping guardrail" + ) + return + + # Make Model Armor request + try: + armor_response = await self.make_model_armor_request( + content=content, + source="model_response", + request_data=data, + ) + + # Check if content should be blocked + if self._should_block_content(armor_response): + raise HTTPException( + status_code=400, + detail={ + "error": "Response blocked by Model Armor", + "model_armor_response": armor_response, + }, + ) + + # If mask_response_content is enabled, update response with sanitized content + if self.mask_response_content: + sanitized_content = self._get_sanitized_content(armor_response) + if sanitized_content and sanitized_content != content: + # Update response content + if isinstance(response, litellm.ModelResponse): + for choice in response.choices: + if isinstance(choice, Choices): + if choice.message.content: + choice.message.content = sanitized_content + + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.error( + "Model Armor post-call error: %s", str(e), exc_info=True + ) + if self.optional_params.get("fail_on_error", True): + raise + + # Add guardrail to headers + add_guardrail_to_applied_guardrails_header( + request_data=data, guardrail_name=self.guardrail_name + ) + + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + request_data: dict, + ) -> AsyncGenerator[ModelResponseStream, None]: + """Process streaming response chunks.""" + + from litellm.llms.base_llm.base_model_iterator import MockResponseIterator + from litellm.main import stream_chunk_builder + + # Collect all chunks + all_chunks: List[ModelResponseStream] = [] + async for chunk in response: + all_chunks.append(chunk) + + # Build complete response + assembled_response = stream_chunk_builder(chunks=all_chunks) + + if isinstance(assembled_response, ModelResponse): + # Extract content + content = self._extract_content_from_response(assembled_response) + + if content: + try: + # Check with Model Armor + armor_response = await self.make_model_armor_request( + content=content, + source="model_response", + request_data=request_data, + ) + + # Check if blocked + if self._should_block_content(armor_response): + raise HTTPException( + status_code=400, + detail={ + "error": "Streaming response blocked by Model Armor", + "model_armor_response": armor_response, + }, + ) + + # Apply sanitization if enabled + if self.mask_response_content: + sanitized_content = self._get_sanitized_content(armor_response) + if sanitized_content and sanitized_content != content: + # Update assembled response + for choice in assembled_response.choices: + if isinstance(choice, Choices): + if choice.message.content: + choice.message.content = sanitized_content + + # Return sanitized stream + mock_response = MockResponseIterator( + model_response=assembled_response + ) + async for chunk in mock_response: + yield chunk + return + + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.error( + "Model Armor streaming error: %s", str(e), exc_info=True + ) + if self.optional_params.get("fail_on_error", True): + raise + else: + verbose_proxy_logger.debug( + "Model Armor: No text content in streaming response, skipping guardrail" + ) + + # Return original chunks if no sanitization needed + for chunk in all_chunks: + yield chunk + + @staticmethod + def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + """ + Get the config model for the Model Armor guardrail. + """ + from litellm.types.proxy.guardrails.guardrail_hooks.model_armor import ( + ModelArmorGuardrailConfigModel, + ) + + return ModelArmorGuardrailConfigModel diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index 2e3ec44e4a5..1cf0b15ab5c 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -121,3 +121,5 @@ def initialize_hide_secrets(litellm_params: LitellmParams, guardrail: Guardrail) ) litellm.logging_callback_manager.add_litellm_callback(_secret_detection_object) return _secret_detection_object + + diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index cd93a1884af..18d1156be74 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -37,6 +37,7 @@ class SupportedGuardrailIntegrations(Enum): PANW_PRISMA_AIRS = "panw_prisma_airs" AZURE_PROMPT_SHIELD = "azure/prompt_shield" AZURE_TEXT_MODERATIONS = "azure/text_moderations" + MODEL_ARMOR = "model_armor" OPENAI_MODERATION = "openai_moderation" class Role(Enum): @@ -309,6 +310,7 @@ class BedrockGuardrailConfigModel(BaseModel): ) + class LakeraV2GuardrailConfigModel(BaseModel): """Configuration parameters for the Lakera AI v2 guardrail""" @@ -398,6 +400,25 @@ class BaseLitellmParams(BaseModel): # works for new and patch update guardrails default=None, description="Optional field if guardrail requires a 'model' parameter" ) + # Model Armor params + template_id: Optional[str] = Field( + default=None, description="The ID of your Model Armor template" + ) + location: Optional[str] = Field( + default=None, description="Google Cloud location/region (e.g., us-central1)" + ) + credentials: Optional[str] = Field( + default=None, + description="Path to Google Cloud credentials JSON file or JSON string", + ) + api_endpoint: Optional[str] = Field( + default=None, description="Optional custom API endpoint for Model Armor" + ) + fail_on_error: Optional[bool] = Field( + default=True, + description="Whether to fail the request if Model Armor encounters an error", + ) + model_config = ConfigDict(extra="allow", protected_namespaces=()) diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/model_armor.py b/litellm/types/proxy/guardrails/guardrail_hooks/model_armor.py new file mode 100644 index 00000000000..2d8fa0606b8 --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/model_armor.py @@ -0,0 +1,35 @@ +from typing import Optional + +from pydantic import Field + +from .base import GuardrailConfigModel + + +class ModelArmorGuardrailConfigModel(GuardrailConfigModel): + """Configuration parameters for Google Cloud Model Armor guardrail""" + + template_id: Optional[str] = Field( + default=None, description="The ID of your Model Armor template" + ) + project_id: Optional[str] = Field( + default=None, description="Google Cloud project ID" + ) + location: Optional[str] = Field( + default=None, description="Google Cloud location/region (e.g., us-central1)" + ) + credentials: Optional[str] = Field( + default=None, + description="Path to Google Cloud credentials JSON file or JSON string", + ) + api_endpoint: Optional[str] = Field( + default=None, description="Optional custom API endpoint for Model Armor" + ) + fail_on_error: Optional[bool] = Field( + default=True, + description="Whether to fail the request if Model Armor encounters an error", + ) + + @staticmethod + def ui_friendly_name() -> str: + """Return the UI-friendly name for Model Armor guardrail""" + return "Google Cloud Model Armor" \ No newline at end of file diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py new file mode 100644 index 00000000000..297baf8860a --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py @@ -0,0 +1,818 @@ +import sys +import os +import io, asyncio +import pytest +import json +from unittest.mock import MagicMock, AsyncMock, patch, Mock + +sys.path.insert(0, os.path.abspath("../../../../..")) + +import litellm +import litellm.types.utils +from litellm.proxy.guardrails.guardrail_hooks.model_armor import ModelArmorGuardrail +from litellm.proxy._types import UserAPIKeyAuth +from litellm.caching import DualCache +from litellm.types.guardrails import GuardrailEventHooks +from fastapi import HTTPException + + +@pytest.mark.asyncio +async def test_model_armor_pre_call_hook_sanitization(): + """Test Model Armor pre-call hook with content sanitization""" + mock_user_api_key_dict = UserAPIKeyAuth() + mock_cache = MagicMock(spec=DualCache) + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + mask_request_content=True, + ) + + # Mock the Model Armor API response + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.json = AsyncMock(return_value={ + "sanitized_text": "Hello, my phone number is [REDACTED]", + "action": "SANITIZE" + }) + + # Mock the access token method + guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project")) + + # Mock the async handler + guardrail.async_handler = AsyncMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) + + request_data = { + "model": "gpt-4", + "messages": [ + {"role": "user", "content": "Hello, my phone number is +1 412 555 1212"} + ], + "metadata": {"guardrails": ["model-armor-test"]} + } + + result = await guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key_dict, + cache=mock_cache, + data=request_data, + call_type="completion" + ) + + # Assert the message was sanitized + assert result["messages"][0]["content"] == "Hello, my phone number is [REDACTED]" + + # Verify API was called correctly + guardrail.async_handler.post.assert_called_once() + call_args = guardrail.async_handler.post.call_args + assert "sanitizeUserPrompt" in call_args[1]["url"] + assert call_args[1]["json"]["user_prompt_data"]["text"] == "Hello, my phone number is +1 412 555 1212" + + +@pytest.mark.asyncio +async def test_model_armor_pre_call_hook_blocked(): + """Test Model Armor pre-call hook when content is blocked""" + mock_user_api_key_dict = UserAPIKeyAuth() + mock_cache = MagicMock(spec=DualCache) + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + ) + + # Mock the Model Armor API response for blocked content + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.json = AsyncMock(return_value={ + "action": "BLOCK", + "blocked": True, + "reason": "Prohibited content detected" + }) + + # Mock the access token method + guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project")) + + # Mock the async handler + guardrail.async_handler = AsyncMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) + + request_data = { + "model": "gpt-4", + "messages": [ + {"role": "user", "content": "Some harmful content"} + ], + "metadata": {"guardrails": ["model-armor-test"]} + } + + # Should raise HTTPException for blocked content + with pytest.raises(HTTPException) as exc_info: + await guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key_dict, + cache=mock_cache, + data=request_data, + call_type="completion" + ) + + assert exc_info.value.status_code == 400 + assert "Content blocked by Model Armor" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_model_armor_post_call_hook_sanitization(): + """Test Model Armor post-call hook with response sanitization""" + mock_user_api_key_dict = UserAPIKeyAuth() + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + mask_response_content=True, + ) + + # Mock the Model Armor API response + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.json = AsyncMock(return_value={ + "sanitized_text": "Here is the information: [REDACTED]", + "action": "SANITIZE" + }) + + # Mock the access token method + guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project")) + + # Mock the async handler + guardrail.async_handler = AsyncMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) + + # Create a mock response + mock_llm_response = litellm.ModelResponse() + mock_llm_response.choices = [ + litellm.Choices( + message=litellm.Message( + content="Here is the information: Credit card 1234-5678-9012-3456" + ) + ) + ] + + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "What's my credit card?"}], + "metadata": {"guardrails": ["model-armor-test"]} + } + + await guardrail.async_post_call_success_hook( + data=request_data, + user_api_key_dict=mock_user_api_key_dict, + response=mock_llm_response + ) + + # Assert the response was sanitized + assert mock_llm_response.choices[0].message.content == "Here is the information: [REDACTED]" + + # Verify API was called correctly + guardrail.async_handler.post.assert_called_once() + call_args = guardrail.async_handler.post.call_args + assert "sanitizeModelResponse" in call_args[1]["url"] + + +@pytest.mark.asyncio +async def test_model_armor_with_list_content(): + """Test Model Armor with messages containing list content""" + mock_user_api_key_dict = UserAPIKeyAuth() + mock_cache = MagicMock(spec=DualCache) + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + ) + + # Mock the Model Armor API response + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.json = AsyncMock(return_value={ + "action": "NONE" + }) + + # Mock the access token method + guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project")) + + # Mock the async handler + guardrail.async_handler = AsyncMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) + + request_data = { + "model": "gpt-4", + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Hello world"}, + {"type": "text", "text": "How are you?"} + ] + } + ], + "metadata": {"guardrails": ["model-armor-test"]} + } + + result = await guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key_dict, + cache=mock_cache, + data=request_data, + call_type="completion" + ) + + # Verify the content was extracted correctly + guardrail.async_handler.post.assert_called_once() + call_args = guardrail.async_handler.post.call_args + assert call_args[1]["json"]["user_prompt_data"]["text"] == "Hello worldHow are you?" + + +@pytest.mark.asyncio +async def test_model_armor_api_error_handling(): + """Test Model Armor error handling when API returns error""" + mock_user_api_key_dict = UserAPIKeyAuth() + mock_cache = MagicMock(spec=DualCache) + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + fail_on_error=True, + ) + + # Mock the Model Armor API error response + mock_response = AsyncMock() + mock_response.status_code = 500 + mock_response.text = "Internal Server Error" + + # Mock the access token method + guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project")) + + # Mock the async handler + guardrail.async_handler = AsyncMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) + + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": {"guardrails": ["model-armor-test"]} + } + + # Should raise HTTPException for API error + with pytest.raises(HTTPException) as exc_info: + await guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key_dict, + cache=mock_cache, + data=request_data, + call_type="completion" + ) + + assert exc_info.value.status_code == 500 + assert "Model Armor API error" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_model_armor_credentials_handling(): + """Test Model Armor handling of different credential types""" + try: + from google.auth.credentials import Credentials + except ImportError: + # If google.auth is not installed, skip this test + pytest.skip("google.auth not installed") + return + + # Test with string credentials (file path) + with patch('os.path.exists', return_value=True): + with patch('builtins.open', mock_open(read_data='{"type": "service_account", "project_id": "test-project"}')): + with patch.object(ModelArmorGuardrail, '_credentials_from_service_account') as mock_creds: + mock_creds_obj = Mock() + mock_creds_obj.token = "test-token" + mock_creds_obj.expired = False + mock_creds_obj.project_id = "test-project" # Add project_id + mock_creds.return_value = mock_creds_obj + + guardrail = ModelArmorGuardrail( + template_id="test-template", + credentials="/path/to/creds.json", + project_id="test-project", # Provide project_id + ) + + # Force credential loading + creds, project_id = guardrail.load_auth(credentials="/path/to/creds.json", project_id="test-project") + + assert mock_creds.called + assert project_id == "test-project" + + +@pytest.mark.asyncio +async def test_model_armor_streaming_response(): + """Test Model Armor with streaming responses""" + mock_user_api_key_dict = UserAPIKeyAuth() + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + mask_response_content=True, + ) + + # Mock the Model Armor API response + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.json = AsyncMock(return_value={ + "sanitized_text": "Sanitized response", + "action": "SANITIZE" + }) + + # Mock the access token method + guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project")) + + # Mock the async handler + guardrail.async_handler = AsyncMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) + + # Create mock streaming chunks + async def mock_stream(): + chunks = [ + litellm.ModelResponseStream( + choices=[ + litellm.types.utils.StreamingChoices( + delta=litellm.types.utils.Delta(content="Sensitive ") + ) + ] + ), + litellm.ModelResponseStream( + choices=[ + litellm.types.utils.StreamingChoices( + delta=litellm.types.utils.Delta(content="information") + ) + ] + ), + ] + for chunk in chunks: + yield chunk + + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Tell me secrets"}], + "metadata": {"guardrails": ["model-armor-test"]} + } + + # Process streaming response + result_chunks = [] + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=mock_user_api_key_dict, + response=mock_stream(), + request_data=request_data + ): + result_chunks.append(chunk) + + # Should have processed the chunks through Model Armor + assert len(result_chunks) > 0 + guardrail.async_handler.post.assert_called() + +def test_model_armor_ui_friendly_name(): + """Test the UI-friendly name of the Model Armor guardrail""" + from litellm.types.proxy.guardrails.guardrail_hooks.model_armor import ( + ModelArmorGuardrailConfigModel, + ) + + assert ( + ModelArmorGuardrailConfigModel.ui_friendly_name() == "Google Cloud Model Armor" + ) + +@pytest.mark.asyncio +async def test_model_armor_no_messages(): + """Test Model Armor when request has no messages""" + mock_user_api_key_dict = UserAPIKeyAuth() + mock_cache = MagicMock(spec=DualCache) + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + ) + + request_data = { + "model": "gpt-4", + "metadata": {"guardrails": ["model-armor-test"]} + } + + # Should return data unchanged when no messages + result = await guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key_dict, + cache=mock_cache, + data=request_data, + call_type="completion" + ) + + assert result == request_data + + +@pytest.mark.asyncio +async def test_model_armor_empty_message_content(): + """Test Model Armor when message content is empty""" + mock_user_api_key_dict = UserAPIKeyAuth() + mock_cache = MagicMock(spec=DualCache) + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + ) + + request_data = { + "model": "gpt-4", + "messages": [ + {"role": "user", "content": ""}, + {"role": "assistant", "content": "Previous response"} + ], + "metadata": {"guardrails": ["model-armor-test"]} + } + + # Should return data unchanged when no content + result = await guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key_dict, + cache=mock_cache, + data=request_data, + call_type="completion" + ) + + assert result == request_data + + +@pytest.mark.asyncio +async def test_model_armor_system_assistant_messages(): + """Test Model Armor with only system/assistant messages (no user messages)""" + mock_user_api_key_dict = UserAPIKeyAuth() + mock_cache = MagicMock(spec=DualCache) + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + ) + + request_data = { + "model": "gpt-4", + "messages": [ + {"role": "system", "content": "You are a helpful assistant"}, + {"role": "assistant", "content": "How can I help you?"} + ], + "metadata": {"guardrails": ["model-armor-test"]} + } + + # Should return data unchanged when no user messages + result = await guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key_dict, + cache=mock_cache, + data=request_data, + call_type="completion" + ) + + assert result == request_data + + +@pytest.mark.asyncio +async def test_model_armor_fail_on_error_false(): + """Test Model Armor with fail_on_error=False when API fails""" + mock_user_api_key_dict = UserAPIKeyAuth() + mock_cache = MagicMock(spec=DualCache) + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + fail_on_error=False, + ) + + # Mock the async handler to raise an exception + guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project")) + guardrail.async_handler = AsyncMock() + # Make it raise a non-HTTP exception to test the fail_on_error logic + guardrail.async_handler.post = AsyncMock(side_effect=Exception("Connection error")) + + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": {"guardrails": ["model-armor-test"]} + } + + # Should not raise exception when fail_on_error=False + result = await guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key_dict, + cache=mock_cache, + data=request_data, + call_type="completion" + ) + + # Should return original data + assert result == request_data + + +@pytest.mark.asyncio +async def test_model_armor_custom_api_endpoint(): + """Test Model Armor with custom API endpoint""" + mock_user_api_key_dict = UserAPIKeyAuth() + mock_cache = MagicMock(spec=DualCache) + + custom_endpoint = "https://custom-modelarmor.example.com" + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + api_endpoint=custom_endpoint, + ) + + # Mock successful response + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.json = AsyncMock(return_value={"action": "NONE"}) + + guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project")) + guardrail.async_handler = AsyncMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) + + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Test message"}], + "metadata": {"guardrails": ["model-armor-test"]} + } + + await guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key_dict, + cache=mock_cache, + data=request_data, + call_type="completion" + ) + + # Verify custom endpoint was used + call_args = guardrail.async_handler.post.call_args + assert call_args[1]["url"].startswith(custom_endpoint) + + +@pytest.mark.asyncio +async def test_model_armor_dict_credentials(): + """Test Model Armor with dictionary credentials instead of file path""" + try: + from google.auth import default + except ImportError: + pytest.skip("google.auth not installed") + return + + # Use patch context manager properly + mock_creds_obj = Mock() + mock_creds_obj.token = "test-token" + mock_creds_obj.expired = False + mock_creds_obj.project_id = "test-project" + + with patch.object(ModelArmorGuardrail, '_credentials_from_service_account', return_value=mock_creds_obj) as mock_creds: + creds_dict = { + "type": "service_account", + "project_id": "test-project", + "private_key": "test-key", + "client_email": "test@example.com" + } + + guardrail = ModelArmorGuardrail( + template_id="test-template", + credentials=creds_dict, + location="us-central1", + ) + + # Force credential loading + creds, project_id = guardrail.load_auth(credentials=creds_dict, project_id=None) + + assert mock_creds.called + assert project_id == "test-project" + + +@pytest.mark.asyncio +async def test_model_armor_action_none(): + """Test Model Armor when action is NONE (no sanitization needed)""" + mock_user_api_key_dict = UserAPIKeyAuth() + mock_cache = MagicMock(spec=DualCache) + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + mask_request_content=True, + ) + + # Mock response with action=NONE + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.json = AsyncMock(return_value={"action": "NONE"}) + + guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project")) + guardrail.async_handler = AsyncMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) + + original_content = "This content is fine" + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": original_content}], + "metadata": {"guardrails": ["model-armor-test"]} + } + + result = await guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key_dict, + cache=mock_cache, + data=request_data, + call_type="completion" + ) + + # Content should remain unchanged + assert result["messages"][0]["content"] == original_content + + +@pytest.mark.asyncio +async def test_model_armor_missing_sanitized_text(): + """Test Model Armor when response has no sanitized_text field""" + mock_user_api_key_dict = UserAPIKeyAuth() + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + mask_response_content=True, + ) + + # Mock response without sanitized_text + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.json = AsyncMock(return_value={ + "action": "SANITIZE", + "text": "Fallback sanitized content" + }) + + guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project")) + guardrail.async_handler = AsyncMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) + + # Create a mock response + mock_llm_response = litellm.ModelResponse() + mock_llm_response.choices = [ + litellm.Choices( + message=litellm.Message(content="Original content") + ) + ] + + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Test"}], + "metadata": {"guardrails": ["model-armor-test"]} + } + + await guardrail.async_post_call_success_hook( + data=request_data, + user_api_key_dict=mock_user_api_key_dict, + response=mock_llm_response + ) + + # Should use 'text' field as fallback + assert mock_llm_response.choices[0].message.content == "Fallback sanitized content" + + +@pytest.mark.asyncio +async def test_model_armor_non_text_response(): + """Test Model Armor with non-text response types (TTS, image generation)""" + mock_user_api_key_dict = UserAPIKeyAuth() + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + ) + + # Mock a non-ModelResponse object (like TTS or image response) + mock_tts_response = Mock() + mock_tts_response.audio = b"audio_data" + + request_data = { + "model": "tts-1", + "input": "Text to speak", + "metadata": {"guardrails": ["model-armor-test"]} + } + + # Should not raise an error for non-text responses + await guardrail.async_post_call_success_hook( + data=request_data, + user_api_key_dict=mock_user_api_key_dict, + response=mock_tts_response + ) + + +@pytest.mark.asyncio +async def test_model_armor_token_refresh(): + """Test Model Armor handling expired auth tokens""" + mock_user_api_key_dict = UserAPIKeyAuth() + mock_cache = MagicMock(spec=DualCache) + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + ) + + # Mock successful response + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.json = AsyncMock(return_value={"action": "NONE"}) + + # Mock token refresh - first call returns expired token, second returns fresh + call_count = 0 + async def mock_token_method(*args, **kwargs): + nonlocal call_count + call_count += 1 + return (f"token-{call_count}", "test-project") + + guardrail._ensure_access_token_async = AsyncMock(side_effect=mock_token_method) + guardrail.async_handler = AsyncMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) + + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Test"}], + "metadata": {"guardrails": ["model-armor-test"]} + } + + await guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key_dict, + cache=mock_cache, + data=request_data, + call_type="completion" + ) + + # Verify token method was called + assert guardrail._ensure_access_token_async.called + + +@pytest.mark.asyncio +async def test_model_armor_non_model_response(): + """Test Model Armor handles non-ModelResponse types (e.g., TTS) correctly""" + mock_user_api_key_dict = UserAPIKeyAuth() + mock_cache = MagicMock(spec=DualCache) + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + ) + + # Mock a TTS response (not a ModelResponse) + class TTSResponse: + def __init__(self): + self.audio_data = b"fake audio data" + + tts_response = TTSResponse() + + # Mock the access token + guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project")) + guardrail.async_handler = AsyncMock() + + # Call post-call hook with non-ModelResponse + await guardrail.async_post_call_success_hook( + data={ + "model": "tts-1", + "input": "Hello world", + "metadata": {"guardrails": ["model-armor-test"]} + }, + user_api_key_dict=mock_user_api_key_dict, + response=tts_response + ) + + # Verify that Model Armor API was NOT called since there's no text content + assert not guardrail.async_handler.post.called + + +def mock_open(read_data=''): + """Helper to create a mock file object""" + import io + from unittest.mock import MagicMock + + file_object = io.StringIO(read_data) + file_object.__enter__ = lambda self: self + file_object.__exit__ = lambda self, *args: None + + mock_file = MagicMock(return_value=file_object) + return mock_file \ No newline at end of file