From 196a7269f68e7b2788a44516d7eca7eb16dab460 Mon Sep 17 00:00:00 2001 From: Alvaro Sevilla Date: Fri, 13 Jun 2025 18:51:04 +0200 Subject: [PATCH] fix(vertex_ai): support global location in vertex ai passthrough (#11661) * Support global location in vertex ai passthrough * refactor --- .../llm_passthrough_endpoints.py | 15 +++- .../test_llm_pass_through_endpoints.py | 90 ++++++++++++++++++- 2 files changed, 99 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 5d140eb511..f9c7230494 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -1,5 +1,5 @@ """ -What is this? +What is this? Provider-specific Pass-Through Endpoints @@ -697,14 +697,21 @@ class VertexAIDiscoveryPassThroughHandler(BaseVertexAIPassThroughHandler): class VertexAIPassThroughHandler(BaseVertexAIPassThroughHandler): @staticmethod def get_default_base_target_url(vertex_location: Optional[str]) -> str: - return f"https://{vertex_location}-aiplatform.googleapis.com/" + return get_vertex_base_url(vertex_location) @staticmethod def update_base_target_url_with_credential_location( base_target_url: str, vertex_location: Optional[str] ) -> str: - return f"https://{vertex_location}-aiplatform.googleapis.com/" + return get_vertex_base_url(vertex_location) +def get_vertex_base_url(vertex_location: Optional[str]) -> str: + """ + Returns the base URL for Vertex AI based on the provided location. + """ + if vertex_location == "global": + return "https://aiplatform.googleapis.com/" + return f"https://{vertex_location}-aiplatform.googleapis.com/" def get_vertex_pass_through_handler( call_type: Literal["discovery", "aiplatform"] @@ -805,7 +812,7 @@ async def _base_vertex_proxy_route( ) if base_target_url is None: - base_target_url = f"https://{vertex_location}-aiplatform.googleapis.com/" + base_target_url = get_vertex_base_url(vertex_location) request_route = encoded_endpoint verbose_proxy_logger.debug("request_route %s", request_route) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index cf88dfb14f..afcc8f8ec7 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -253,7 +253,7 @@ class TestVertexAIPassThroughHandler: mock_get_token.return_value = (test_token, "") mock_get_virtual_key.return_value = "Bearer test-key" mock_user_auth.return_value = {"api_key": "test-key"} - + # Mock create_pass_through_route to return a function that returns a mock response mock_endpoint_func = AsyncMock(return_value={"status": "success"}) mock_create_route.return_value = mock_endpoint_func @@ -276,6 +276,92 @@ class TestVertexAIPassThroughHandler: custom_headers={"Authorization": f"Bearer {test_token}"}, ) + @pytest.mark.asyncio + async def test_vertex_passthrough_with_global_location(self, monkeypatch): + """ + Test that when global location is used, it is correctly handled in the request + """ + from litellm.proxy.pass_through_endpoints.passthrough_endpoint_router import ( + PassthroughEndpointRouter, + ) + + vertex_project = "test-project" + vertex_location = "global" + vertex_credentials = "test-creds" + + pass_through_router = PassthroughEndpointRouter() + + pass_through_router.add_vertex_credentials( + project_id=vertex_project, + location=vertex_location, + vertex_credentials=vertex_credentials, + ) + + monkeypatch.setattr( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router", + pass_through_router, + ) + + endpoint = f"/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/gemini-1.5-flash:generateContent" + + # Mock request + mock_request = Mock() + mock_request.method = "POST" + mock_request.headers = { + "Authorization": "Bearer test-creds", + "Content-Type": "application/json", + } + mock_request.url = Mock() + mock_request.url.path = endpoint + + # Mock response + mock_response = Response() + + # Mock vertex credentials + test_project = vertex_project + test_location = vertex_location + test_token = vertex_credentials + + with mock.patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.vertex_llm_base._ensure_access_token_async" + ) as mock_ensure_token, mock.patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.vertex_llm_base._get_token_and_url" + ) as mock_get_token, mock.patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" + ) as mock_create_route, mock.patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_litellm_virtual_key" + ) as mock_get_virtual_key, mock.patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth" + ) as mock_user_auth: + # Setup mocks + mock_ensure_token.return_value = ("test-auth-header", test_project) + mock_get_token.return_value = (test_token, "") + mock_get_virtual_key.return_value = "Bearer test-key" + mock_user_auth.return_value = {"api_key": "test-key"} + + # Mock create_pass_through_route to return a function that returns a mock response + mock_endpoint_func = AsyncMock(return_value={"status": "success"}) + mock_create_route.return_value = mock_endpoint_func + + # Call the route + try: + result = await vertex_proxy_route( + endpoint=endpoint, + request=mock_request, + fastapi_response=mock_response, + user_api_key_dict={"api_key": "test-key"}, + ) + except Exception as e: + print(f"Error: {e}") + + # Verify create_pass_through_route was called with correct arguments + mock_create_route.assert_called_once_with( + endpoint=endpoint, + target=f"https://aiplatform.googleapis.com/v1/projects/{test_project}/locations/{test_location}/publishers/google/models/gemini-1.5-flash:generateContent", + custom_headers={"Authorization": f"Bearer {test_token}"}, + ) + + @pytest.mark.parametrize( "initial_endpoint", [ @@ -534,7 +620,7 @@ class TestVertexAIDiscoveryPassThroughHandler: mock_get_token.return_value = (test_token, "") mock_get_virtual_key.return_value = "Bearer test-key" mock_user_auth.return_value = {"api_key": "test-key"} - + # Mock create_pass_through_route to return a function that returns a mock response mock_endpoint_func = AsyncMock(return_value={"status": "success"}) mock_create_route.return_value = mock_endpoint_func