diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index a12eb2486d2..38462094b11 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -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( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py index 6d0a1b46559..987388a80c7 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py @@ -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) \ No newline at end of file + assert "API Error" in str(exc_info.value)