feat: integrate Google Cloud Model Armor guardrails (#12492)

* feat: integrate Google Cloud Model Armor guardrails (LIT-298)

- Add ModelArmorGuardrail class that extends CustomGuardrail and VertexBase
- Support for both pre-call (sanitizeUserPrompt) and post-call (sanitizeModelResponse) sanitization
- Integrate with existing Vertex AI authentication using VertexBase
- Add configuration model for Model Armor in guardrail types
- Register Model Armor in guardrail initializers and registry
- Include comprehensive test suite for Model Armor functionality
- Support content masking for both requests and responses
- Handle streaming responses with content sanitization

This integration allows LiteLLM to use Google Cloud Model Armor API for
content moderation and sanitization, providing similar functionality to
Bedrock Guards but using Google Cloud's security infrastructure.

* fix: remove unused imports flagged by ruff linter

- Remove unused asyncio import at top level (moved to local import where needed)
- Remove unused TextCompletionResponse import

* fix: remove additional unused imports

- Remove unused TextChoices import from line 34
- Remove duplicate asyncio import from line 384
- Replace asyncio.iscoroutine() with hasattr check for __await__

* fix: remove final unused imports

- Remove unused 'import sys' from line 9
- Remove unused 'StreamingChoices' from imports

* fix: remove unused import from model_armor.py

- Remove unused 'import os' from the top of the file

* fix: remove commented-out header from model_armor.py

- Eliminate unnecessary comments at the top of the file to improve code clarity.

* feat(guardrails): Add Model Armor UI support

- Convert model_armor.py to a directory structure with __init__.py for dynamic discovery
- Add get_config_model() method to ModelArmorGuardrail class for UI integration
- Add ui_friendly_name() to ModelArmorConfigModel returning "Google Cloud Model Armor"
- Remove manual registration from guardrail_registry.py to use dynamic discovery
- Model Armor now appears in the UI guardrails dropdown with proper configuration fields

This enables users to configure Model Armor guardrails through the LiteLLM UI interface.

* fix(guardrails): Fix undefined name 'GuardrailConfigModel' in Model Armor

- Import TYPE_CHECKING and GuardrailConfigModel from base module
- Fixes F821 linting error for undefined name in type annotation
- Follows same pattern as other guardrail implementations

* fix(guardrails): Fix Model Armor type errors and config model inheritance

- Create ModelArmorGuardrailConfigModel that properly inherits from GuardrailConfigModel base class
- Move config model to litellm/types/proxy/guardrails/guardrail_hooks/model_armor.py following convention
- Update get_config_model() to return the properly typed config model
- Remove ModelArmorConfigModel from LitellmParams inheritance chain
- Add template_id field to BaseLitellmParams instead

This fixes the mypy type errors and follows the same pattern as other guardrail implementations.

* fix(guardrails): Add missing Model Armor fields to BaseLitellmParams

- Add location, credentials, api_endpoint, and fail_on_error fields
- Fixes mypy errors about missing attributes in LitellmParams
- All Model Armor configuration parameters are now properly defined

* fix: Apply PR review feedback for Model Armor guardrail

- Move test file to tests/test_litellm/ for GitHub Actions
- Extract only last consecutive user messages to avoid context limits
- Use get_content_from_model_response helper for response extraction
- Handle non-ModelResponse types (e.g., TTS) gracefully
- Maintain newline separation for multi-part content

* refactor: Simplify message content extraction in ModelArmorGuardrail

- Removed the custom _extract_content_from_messages method.
- Integrated get_last_user_message helper for improved content extraction.
- Updated test to reflect changes in content formatting.

* refactor: Remove unused import in model_armor.py

- Deleted the unused import of AllMessageValues to clean up the codebase.

* Add unit tests for Model Armor guardrail functionality

- Implement tests for pre-call and post-call hooks, including content sanitization and blocking behavior.
- Validate error handling for API responses and credential management.
- Test streaming responses and handling of list content in user messages.
- Ensure proper assertions for API interactions and response sanitization.

* Add comprehensive test coverage for Model Armor guardrail

- Add test for requests with no messages field
- Add test for empty message content handling
- Add test for system/assistant-only messages
- Add test for fail_on_error=False behavior
- Add test for custom API endpoint configuration
- Add test for dictionary credentials (non-file path)
- Add test for action=NONE response handling
- Add test for missing sanitized_text field fallback
- Add test for non-text response types (TTS/image)
- Add test for auth token refresh behavior

All tests ensure robust edge case handling and proper error management.

* feat: improve Model Armor handling of non-ModelResponse types

- Add debug logging when skipping non-text responses (TTS, images, etc.)
- Improve docstring to clarify behavior for non-text responses
- Add test coverage for non-ModelResponse handling
- Ensure guardrail gracefully skips processing for response types it cannot handle
This commit is contained in:
Cole McIntosh 2025-07-18 13:51:09 -06:00 • committed by GitHub
parent 208484ce65
commit 506dc80b15
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 1378 additions and 0 deletions

View file

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

View file

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

View file

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

View file

@ -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=())

View file

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

View file

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