mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
208484ce65
commit
506dc80b15
6 changed files with 1378 additions and 0 deletions
|
|
@ -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,
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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=())
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue