From b04b456bc27e8c0590d93f5fabc8d00de5647a5f Mon Sep 17 00:00:00 2001 From: Cole McIntosh <82463175+colesmcintosh@users.noreply.github.com> Date: Sat, 19 Jul 2025 16:20:40 -0600 Subject: [PATCH] fix(proxy): Fix Model Armor project_id initialization order (#12766) When using Model Armor guardrail with explicit project_id in config, the project_id was being overwritten to None due to incorrect initialization order between ModelArmorGuardrail and VertexBase parent class. This fix ensures that user-provided project_id is preserved by initializing parent classes before setting instance attributes. Fixes #12757 --- .../model_armor/model_armor.py | 8 +- .../guardrail_hooks/test_model_armor.py | 80 ++++++++++++++++++- 2 files changed, 84 insertions(+), 4 deletions(-) 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 101b0e76d46..889d02bb736 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -55,6 +55,11 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): api_endpoint: Optional[str] = None, **kwargs, ): + # Initialize parent classes first + super().__init__(**kwargs) + VertexBase.__init__(self) + + # Then set our attributes (this ensures project_id is not overwritten) self.async_handler = get_async_httpx_client( llm_provider=httpxSpecialProvider.GuardrailCallback ) @@ -67,9 +72,6 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): # Store optional params self.optional_params = kwargs - super().__init__(**kwargs) - VertexBase.__init__(self) - verbose_proxy_logger.debug( "Model Armor Guardrail initialized with template_id: %s, project_id: %s, location: %s", self.template_id, 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 297baf8860a..fc408efdf39 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 @@ -815,4 +815,82 @@ def mock_open(read_data=''): file_object.__exit__ = lambda self, *args: None mock_file = MagicMock(return_value=file_object) - return mock_file \ No newline at end of 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 + + +@pytest.mark.asyncio +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", + project_id="cloud-test-project", + location="eu", + 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 + mock_response.json = AsyncMock(return_value={ + "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 + guardrail.async_handler = AsyncMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) + + request_data = { + "model": "gpt-4", + "messages": [ + {"role": "user", "content": "Test content"} + ], + "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, + cache=mock_cache, + data=request_data, + call_type="completion" + ) + + # Verify the project_id was used correctly in the API call + guardrail.async_handler.post.assert_called_once() + call_args = guardrail.async_handler.post.call_args + assert "cloud-test-project" in call_args[1]["url"] \ No newline at end of file