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:
Hi120ki 2026-02-01 08:04:19 +09:00 • committed by Sameer Kankute
parent 6d86808eaf
commit c9757cd0d7
2 changed files with 224 additions and 128 deletions

View file

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

View file

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