mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): populate applied_guardrails when Model Armor blocks content (#20034)
Previously, when Model Armor guardrail blocked a request/response, the `applied_guardrails` field was not populated in the logs because `add_guardrail_to_applied_guardrails_header()` was called after the HTTPException was raised. This fix moves the `add_guardrail_to_applied_guardrails_header()` call to before the blocking check in all hooks: - async_pre_call_hook (pre_call mode) - async_moderation_hook (during_call mode) - async_post_call_success_hook (post_call mode) - async_post_call_streaming_iterator_hook (streaming) This ensures that even when a guardrail blocks content, the guardrail name is properly recorded in the logs for observability. Added regression tests to verify applied_guardrails is populated when content is blocked. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
9dc55ea69a
commit
b524a46e35
2 changed files with 224 additions and 128 deletions
|
|
@ -421,6 +421,13 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
)
|
||||
else "success"
|
||||
)
|
||||
|
||||
# Add guardrail to applied_guardrails BEFORE potential blocking
|
||||
# This ensures guardrail is recorded even when it blocks the request
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
|
||||
# Check if content should be blocked
|
||||
if self._should_block_content(
|
||||
armor_response, allow_sanitization=self.mask_request_content
|
||||
|
|
@ -456,11 +463,6 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
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
|
||||
|
|
@ -517,6 +519,12 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
else "success"
|
||||
)
|
||||
|
||||
# Add guardrail to applied_guardrails BEFORE potential blocking
|
||||
# This ensures guardrail is recorded even when it blocks the request
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
|
||||
# Check if content should be blocked
|
||||
if self._should_block_content(
|
||||
armor_response, allow_sanitization=self.mask_request_content
|
||||
|
|
@ -550,11 +558,6 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
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
|
||||
|
|
@ -622,6 +625,12 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
guardrail_response=standard_logging_guardrail_information,
|
||||
)
|
||||
|
||||
# Add guardrail to applied_guardrails BEFORE potential blocking
|
||||
# This ensures guardrail is recorded even when it blocks the request
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
|
||||
# Check if content should be blocked
|
||||
if self._should_block_content(
|
||||
armor_response, allow_sanitization=self.mask_response_content
|
||||
|
|
@ -654,11 +663,6 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
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 response
|
||||
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
|
|
@ -703,6 +707,16 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
else "success"
|
||||
)
|
||||
|
||||
# Add guardrail to applied_guardrails BEFORE potential blocking
|
||||
# This ensures guardrail is recorded even when it blocks the request
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=request_data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
|
||||
# Check if blocked
|
||||
if self._should_block_content(armor_response):
|
||||
raise HTTPException(
|
||||
|
|
|
|||
|
|
@ -24,7 +24,7 @@ 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",
|
||||
|
|
@ -32,7 +32,7 @@ async def test_model_armor_pre_call_hook_sanitization():
|
|||
guardrail_name="model-armor-test",
|
||||
mask_request_content=True,
|
||||
)
|
||||
|
||||
|
||||
# Mock the Model Armor API response
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
|
|
@ -53,10 +53,10 @@ async def test_model_armor_pre_call_hook_sanitization():
|
|||
}
|
||||
}
|
||||
})
|
||||
|
||||
|
||||
# Mock the access token method
|
||||
guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project"))
|
||||
|
||||
|
||||
# Mock the async handler
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock_response)):
|
||||
request_data = {
|
||||
|
|
@ -66,17 +66,17 @@ async def test_model_armor_pre_call_hook_sanitization():
|
|||
],
|
||||
"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
|
||||
# Note: we need to use the captured mock from the patch if we want to assert on it
|
||||
# But for now, we'll just verify the behavior.
|
||||
|
|
@ -89,14 +89,14 @@ 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
|
||||
|
|
@ -118,10 +118,10 @@ async def test_model_armor_pre_call_hook_blocked():
|
|||
}
|
||||
}
|
||||
})
|
||||
|
||||
|
||||
# Mock the access token method
|
||||
guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project"))
|
||||
|
||||
|
||||
# Mock the async handler
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock_response)):
|
||||
request_data = {
|
||||
|
|
@ -131,7 +131,7 @@ async def test_model_armor_pre_call_hook_blocked():
|
|||
],
|
||||
"metadata": {"guardrails": ["model-armor-test"]}
|
||||
}
|
||||
|
||||
|
||||
# Should raise HTTPException for blocked content
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.async_pre_call_hook(
|
||||
|
|
@ -140,16 +140,21 @@ async def test_model_armor_pre_call_hook_blocked():
|
|||
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)
|
||||
|
||||
# IMPORTANT: Verify that applied_guardrails is populated even when blocked
|
||||
# This is a regression test for the issue where applied_guardrails was null when blocked
|
||||
assert "applied_guardrails" in request_data["metadata"]
|
||||
assert "model-armor-test" in request_data["metadata"]["applied_guardrails"]
|
||||
|
||||
|
||||
@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",
|
||||
|
|
@ -157,7 +162,7 @@ async def test_model_armor_post_call_hook_sanitization():
|
|||
guardrail_name="model-armor-test",
|
||||
mask_response_content=True,
|
||||
)
|
||||
|
||||
|
||||
# Mock the Model Armor API response
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
|
|
@ -178,10 +183,10 @@ async def test_model_armor_post_call_hook_sanitization():
|
|||
}
|
||||
}
|
||||
})
|
||||
|
||||
|
||||
# Mock the access token method
|
||||
guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project"))
|
||||
|
||||
|
||||
# Mock the async handler
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock_response)):
|
||||
# Create a mock response
|
||||
|
|
@ -193,36 +198,108 @@ async def test_model_armor_post_call_hook_sanitization():
|
|||
)
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
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]"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_armor_post_call_hook_blocked():
|
||||
"""Test Model Armor post-call hook when response is blocked and applied_guardrails is populated"""
|
||||
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 the Model Armor API response for blocked content
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json = AsyncMock(return_value={
|
||||
"sanitizationResult": {
|
||||
"filterMatchState": "MATCH_FOUND",
|
||||
"filterResults": {
|
||||
"rai": {
|
||||
"raiFilterResult": {
|
||||
"matchState": "MATCH_FOUND",
|
||||
"raiFilterTypeResults": {
|
||||
"dangerous": {
|
||||
"matchState": "MATCH_FOUND",
|
||||
"reason": "Harmful response detected"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
# Mock the access token method
|
||||
guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project"))
|
||||
|
||||
# Mock the async handler
|
||||
with patch.object(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 some harmful content..."
|
||||
)
|
||||
)
|
||||
]
|
||||
|
||||
request_data = {
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "Some prompt"}],
|
||||
"metadata": {"guardrails": ["model-armor-test"]}
|
||||
}
|
||||
|
||||
# Should raise HTTPException for blocked response
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.async_post_call_success_hook(
|
||||
data=request_data,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
response=mock_llm_response
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "Response blocked by Model Armor" in str(exc_info.value.detail)
|
||||
|
||||
# IMPORTANT: Verify that applied_guardrails is populated even when blocked
|
||||
# This is a regression test for the issue where applied_guardrails was null when blocked
|
||||
assert "applied_guardrails" in request_data["metadata"]
|
||||
assert "model-armor-test" in request_data["metadata"]["applied_guardrails"]
|
||||
|
||||
|
||||
@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
|
||||
|
|
@ -231,17 +308,17 @@ async def test_model_armor_with_list_content():
|
|||
"filterMatchState": "NO_MATCH_FOUND"
|
||||
}
|
||||
})
|
||||
|
||||
|
||||
# Mock the access token method
|
||||
guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project"))
|
||||
|
||||
|
||||
# Mock the async handler
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock_response)) as mock_post:
|
||||
request_data = {
|
||||
"model": "gpt-4",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Hello world"},
|
||||
{"type": "text", "text": "How are you?"}
|
||||
|
|
@ -250,14 +327,14 @@ async def test_model_armor_with_list_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"
|
||||
)
|
||||
|
||||
|
||||
# Verify the content was extracted correctly
|
||||
mock_post.assert_called_once()
|
||||
call_args = mock_post.call_args
|
||||
|
|
@ -269,7 +346,7 @@ 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",
|
||||
|
|
@ -277,15 +354,15 @@ async def test_model_armor_api_error_handling():
|
|||
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
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock_response)):
|
||||
request_data = {
|
||||
|
|
@ -293,7 +370,7 @@ async def test_model_armor_api_error_handling():
|
|||
"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(
|
||||
|
|
@ -302,7 +379,7 @@ async def test_model_armor_api_error_handling():
|
|||
data=request_data,
|
||||
call_type="completion"
|
||||
)
|
||||
|
||||
|
||||
assert exc_info.value.status_code == 500
|
||||
assert "Model Armor API error" in str(exc_info.value.detail)
|
||||
|
||||
|
|
@ -316,7 +393,7 @@ async def test_model_armor_credentials_handling():
|
|||
# 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"}')):
|
||||
|
|
@ -326,16 +403,16 @@ async def test_model_armor_credentials_handling():
|
|||
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"
|
||||
|
||||
|
|
@ -344,7 +421,7 @@ async def test_model_armor_credentials_handling():
|
|||
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",
|
||||
|
|
@ -352,7 +429,7 @@ async def test_model_armor_streaming_response():
|
|||
guardrail_name="model-armor-test",
|
||||
mask_response_content=True,
|
||||
)
|
||||
|
||||
|
||||
# Mock the Model Armor API response
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
|
|
@ -362,10 +439,10 @@ async def test_model_armor_streaming_response():
|
|||
"sanitizedText": "Sanitized response"
|
||||
}
|
||||
})
|
||||
|
||||
|
||||
# Mock the access token method
|
||||
guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project"))
|
||||
|
||||
|
||||
# Mock the async handler
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock_response)) as mock_post:
|
||||
# Create mock streaming chunks
|
||||
|
|
@ -388,13 +465,13 @@ async def test_model_armor_streaming_response():
|
|||
]
|
||||
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(
|
||||
|
|
@ -403,7 +480,7 @@ async def test_model_armor_streaming_response():
|
|||
request_data=request_data
|
||||
):
|
||||
result_chunks.append(chunk)
|
||||
|
||||
|
||||
# Should have processed the chunks through Model Armor
|
||||
assert len(result_chunks) > 0
|
||||
mock_post.assert_called()
|
||||
|
|
@ -423,19 +500,19 @@ 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,
|
||||
|
|
@ -443,7 +520,7 @@ async def test_model_armor_no_messages():
|
|||
data=request_data,
|
||||
call_type="completion"
|
||||
)
|
||||
|
||||
|
||||
assert result == request_data
|
||||
|
||||
|
||||
|
|
@ -452,14 +529,14 @@ 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": [
|
||||
|
|
@ -468,7 +545,7 @@ async def test_model_armor_empty_message_content():
|
|||
],
|
||||
"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,
|
||||
|
|
@ -476,7 +553,7 @@ async def test_model_armor_empty_message_content():
|
|||
data=request_data,
|
||||
call_type="completion"
|
||||
)
|
||||
|
||||
|
||||
assert result == request_data
|
||||
|
||||
|
||||
|
|
@ -485,14 +562,14 @@ 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": [
|
||||
|
|
@ -501,7 +578,7 @@ async def test_model_armor_system_assistant_messages():
|
|||
],
|
||||
"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,
|
||||
|
|
@ -509,7 +586,7 @@ async def test_model_armor_system_assistant_messages():
|
|||
data=request_data,
|
||||
call_type="completion"
|
||||
)
|
||||
|
||||
|
||||
assert result == request_data
|
||||
|
||||
|
||||
|
|
@ -518,7 +595,7 @@ 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",
|
||||
|
|
@ -526,7 +603,7 @@ async def test_model_armor_fail_on_error_false():
|
|||
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"))
|
||||
# Make it raise a non-HTTP exception to test the fail_on_error logic
|
||||
|
|
@ -536,7 +613,7 @@ async def test_model_armor_fail_on_error_false():
|
|||
"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,
|
||||
|
|
@ -544,7 +621,7 @@ async def test_model_armor_fail_on_error_false():
|
|||
data=request_data,
|
||||
call_type="completion"
|
||||
)
|
||||
|
||||
|
||||
# Should return original data
|
||||
assert result == request_data
|
||||
|
||||
|
|
@ -554,7 +631,7 @@ 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",
|
||||
|
|
@ -563,12 +640,12 @@ async def test_model_armor_custom_api_endpoint():
|
|||
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"))
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock_response)) as mock_post:
|
||||
request_data = {
|
||||
|
|
@ -576,14 +653,14 @@ async def test_model_armor_custom_api_endpoint():
|
|||
"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 = mock_post.call_args
|
||||
assert call_args[1]["url"].startswith(custom_endpoint)
|
||||
|
|
@ -597,13 +674,13 @@ async def test_model_armor_dict_credentials():
|
|||
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",
|
||||
|
|
@ -611,16 +688,16 @@ async def test_model_armor_dict_credentials():
|
|||
"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"
|
||||
|
||||
|
|
@ -630,7 +707,7 @@ 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",
|
||||
|
|
@ -638,7 +715,7 @@ async def test_model_armor_action_none():
|
|||
guardrail_name="model-armor-test",
|
||||
mask_request_content=True,
|
||||
)
|
||||
|
||||
|
||||
# Mock response with action=NO_MATCH_FOUND
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
|
|
@ -647,7 +724,7 @@ async def test_model_armor_action_none():
|
|||
"filterMatchState": "NO_MATCH_FOUND"
|
||||
}
|
||||
})
|
||||
|
||||
|
||||
guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project"))
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock_response)):
|
||||
original_content = "This content is fine"
|
||||
|
|
@ -656,14 +733,14 @@ async def test_model_armor_action_none():
|
|||
"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
|
||||
|
||||
|
|
@ -672,7 +749,7 @@ async def test_model_armor_action_none():
|
|||
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",
|
||||
|
|
@ -680,7 +757,7 @@ async def test_model_armor_missing_sanitized_text():
|
|||
guardrail_name="model-armor-test",
|
||||
mask_response_content=True,
|
||||
)
|
||||
|
||||
|
||||
# Mock response without sanitized_text
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
|
|
@ -689,7 +766,7 @@ async def test_model_armor_missing_sanitized_text():
|
|||
"filterMatchState": "NO_MATCH_FOUND"
|
||||
}
|
||||
})
|
||||
|
||||
|
||||
guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "test-project"))
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock_response)):
|
||||
# Create a mock response
|
||||
|
|
@ -699,19 +776,19 @@ async def test_model_armor_missing_sanitized_text():
|
|||
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 == "Original content"
|
||||
|
||||
|
|
@ -792,8 +869,8 @@ async def test_model_armor_no_circular_reference_in_logging():
|
|||
|
||||
# Verify the logging decorator properly added the guardrail information
|
||||
assert "standard_logging_guardrail_information" in request_data.get("metadata", {})
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_armor_bomb_content_blocked():
|
||||
"""Test Model Armor correctly blocks harmful content like bomb-making instructions"""
|
||||
|
|
@ -936,24 +1013,24 @@ async def test_model_armor_success_case_serializable():
|
|||
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,
|
||||
|
|
@ -967,26 +1044,26 @@ 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)
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock_response)):
|
||||
request_data = {
|
||||
|
|
@ -994,14 +1071,14 @@ async def test_model_armor_token_refresh():
|
|||
"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
|
||||
|
||||
|
|
@ -1011,25 +1088,25 @@ 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={
|
||||
|
|
@ -1040,7 +1117,7 @@ async def test_model_armor_non_model_response():
|
|||
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
|
||||
|
||||
|
|
@ -1049,36 +1126,36 @@ 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
|
||||
return mock_file
|
||||
|
||||
|
||||
def test_model_armor_initialization_preserves_project_id():
|
||||
"""Test that ModelArmorGuardrail initialization preserves the project_id correctly"""
|
||||
# This tests the fix for issue #12757 where project_id was being overwritten to None
|
||||
# due to incorrect initialization order with VertexBase parent class
|
||||
|
||||
|
||||
test_project_id = "cloud-xxxxx-yyyyy"
|
||||
test_template_id = "global-armor"
|
||||
test_location = "eu"
|
||||
|
||||
|
||||
guardrail = ModelArmorGuardrail(
|
||||
template_id=test_template_id,
|
||||
project_id=test_project_id,
|
||||
location=test_location,
|
||||
guardrail_name="model-armor-test",
|
||||
)
|
||||
|
||||
|
||||
# Assert that project_id is preserved after initialization
|
||||
assert guardrail.project_id == test_project_id
|
||||
assert guardrail.template_id == test_template_id
|
||||
assert guardrail.location == test_location
|
||||
|
||||
|
||||
# Also check that the VertexBase initialization didn't reset project_id to None
|
||||
assert hasattr(guardrail, 'project_id')
|
||||
assert guardrail.project_id is not None
|
||||
|
|
@ -1089,7 +1166,7 @@ async def test_model_armor_with_default_credentials():
|
|||
"""Test Model Armor with default credentials and explicit project_id"""
|
||||
mock_user_api_key_dict = UserAPIKeyAuth()
|
||||
mock_cache = MagicMock(spec=DualCache)
|
||||
|
||||
|
||||
# Initialize with explicit project_id but no credentials (simulating default auth)
|
||||
guardrail = ModelArmorGuardrail(
|
||||
template_id="test-template",
|
||||
|
|
@ -1098,7 +1175,7 @@ async def test_model_armor_with_default_credentials():
|
|||
guardrail_name="model-armor-test",
|
||||
credentials=None, # Explicitly set to None to test default auth
|
||||
)
|
||||
|
||||
|
||||
# Mock the Model Armor API response
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
|
|
@ -1106,10 +1183,10 @@ async def test_model_armor_with_default_credentials():
|
|||
"sanitized_text": "Test content",
|
||||
"action": "SANITIZE"
|
||||
})
|
||||
|
||||
|
||||
# Mock the access token method to simulate successful auth
|
||||
guardrail._ensure_access_token_async = AsyncMock(return_value=("test-token", "cloud-test-project"))
|
||||
|
||||
|
||||
# Mock the async handler
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock_response)) as mock_post:
|
||||
request_data = {
|
||||
|
|
@ -1119,7 +1196,7 @@ async def test_model_armor_with_default_credentials():
|
|||
],
|
||||
"metadata": {"guardrails": ["model-armor-test"]}
|
||||
}
|
||||
|
||||
|
||||
# This should not raise ValueError about project_id
|
||||
result = await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
|
|
@ -1127,7 +1204,7 @@ async def test_model_armor_with_default_credentials():
|
|||
data=request_data,
|
||||
call_type="completion"
|
||||
)
|
||||
|
||||
|
||||
# Verify the project_id was used correctly in the API call
|
||||
mock_post.assert_called_once()
|
||||
call_args = mock_post.call_args
|
||||
|
|
@ -1241,6 +1318,11 @@ async def test_async_moderation_hook_content_blocked():
|
|||
assert "_model_armor_response" in request_data["metadata"]
|
||||
assert request_data["metadata"]["_model_armor_status"] == "blocked"
|
||||
|
||||
# IMPORTANT: Verify that applied_guardrails is populated even when blocked
|
||||
# This is a regression test for the issue where applied_guardrails was null when blocked
|
||||
assert "applied_guardrails" in request_data["metadata"]
|
||||
assert "model-armor-test" in request_data["metadata"]["applied_guardrails"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_moderation_hook_with_sanitization():
|
||||
|
|
@ -1446,4 +1528,4 @@ async def test_async_moderation_hook_api_error_fail_on_error_false():
|
|||
call_type="completion"
|
||||
)
|
||||
|
||||
assert "API Error" in str(exc_info.value)
|
||||
assert "API Error" in str(exc_info.value)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue