diff --git a/tests/litellm/llms/vertex_ai/test_vertex_ai_common_utils.py b/tests/litellm/llms/vertex_ai/test_vertex_ai_common_utils.py new file mode 100644 index 0000000000..e89355443f --- /dev/null +++ b/tests/litellm/llms/vertex_ai/test_vertex_ai_common_utils.py @@ -0,0 +1,43 @@ +import os +import sys +from unittest.mock import MagicMock, call, patch + +import pytest + +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system path + +import litellm +from litellm.llms.vertex_ai.common_utils import ( + get_vertex_location_from_url, + get_vertex_project_id_from_url, +) + + +@pytest.mark.asyncio +async def test_get_vertex_project_id_from_url(): + """Test _get_vertex_project_id_from_url with various URLs""" + # Test with valid URL + url = "https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:streamGenerateContent" + project_id = get_vertex_project_id_from_url(url) + assert project_id == "test-project" + + # Test with invalid URL + url = "https://invalid-url.com" + project_id = get_vertex_project_id_from_url(url) + assert project_id is None + + +@pytest.mark.asyncio +async def test_get_vertex_location_from_url(): + """Test _get_vertex_location_from_url with various URLs""" + # Test with valid URL + url = "https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:streamGenerateContent" + location = get_vertex_location_from_url(url) + assert location == "us-central1" + + # Test with invalid URL + url = "https://invalid-url.com" + location = get_vertex_location_from_url(url) + assert location is None diff --git a/tests/litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 74a3dd45c8..da08dea605 100644 --- a/tests/litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -3,7 +3,7 @@ import os import sys import traceback from unittest import mock -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, Mock, patch import httpx import pytest @@ -419,3 +419,40 @@ class TestVertexAIPassThroughHandler: target=f"https://{test_location}-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.asyncio + async def test_async_vertex_proxy_route_api_key_auth(self): + """ + Critical + + This is how Vertex AI JS SDK will Auth to Litellm Proxy + """ + # Mock dependencies + mock_request = Mock() + mock_request.headers = {"x-litellm-api-key": "test-key-123"} + mock_request.method = "POST" + mock_response = Mock() + + with patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth" + ) as mock_auth: + mock_auth.return_value = {"api_key": "test-key-123"} + + with patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" + ) as mock_pass_through: + mock_pass_through.return_value = AsyncMock( + return_value={"status": "success"} + ) + + # Call the function + result = await vertex_proxy_route( + endpoint="v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-1.5-pro:generateContent", + request=mock_request, + fastapi_response=mock_response, + ) + + # Verify user_api_key_auth was called with the correct Bearer token + mock_auth.assert_called_once() + call_args = mock_auth.call_args[1] + assert call_args["api_key"] == "Bearer test-key-123" diff --git a/tests/litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py b/tests/litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py new file mode 100644 index 0000000000..bd8c5f5a99 --- /dev/null +++ b/tests/litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py @@ -0,0 +1,44 @@ +import json +import os +import sys +import traceback +from unittest import mock +from unittest.mock import MagicMock, patch + +import httpx +import pytest +from fastapi import Request, Response +from fastapi.testclient import TestClient + +sys.path.insert( + 0, os.path.abspath("../../../..") +) # Adds the parent directory to the system path + +from unittest.mock import Mock + +from litellm.proxy.pass_through_endpoints.common_utils import get_litellm_virtual_key + + +@pytest.mark.asyncio +async def test_get_litellm_virtual_key(): + """ + Test that the get_litellm_virtual_key function correctly handles the API key authentication + """ + # Test with x-litellm-api-key + mock_request = Mock() + mock_request.headers = {"x-litellm-api-key": "test-key-123"} + result = get_litellm_virtual_key(mock_request) + assert result == "Bearer test-key-123" + + # Test with Authorization header + mock_request.headers = {"Authorization": "Bearer auth-key-456"} + result = get_litellm_virtual_key(mock_request) + assert result == "Bearer auth-key-456" + + # Test with both headers (x-litellm-api-key should take precedence) + mock_request.headers = { + "x-litellm-api-key": "test-key-123", + "Authorization": "Bearer auth-key-456", + } + result = get_litellm_virtual_key(mock_request) + assert result == "Bearer test-key-123" diff --git a/tests/pass_through_unit_tests/test_unit_test_passthrough_router.py b/tests/pass_through_unit_tests/test_unit_test_passthrough_router.py index 6e8296876a..8e016b68d0 100644 --- a/tests/pass_through_unit_tests/test_unit_test_passthrough_router.py +++ b/tests/pass_through_unit_tests/test_unit_test_passthrough_router.py @@ -11,6 +11,7 @@ from unittest.mock import patch from litellm.proxy.pass_through_endpoints.passthrough_endpoint_router import ( PassthroughEndpointRouter, ) +from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials passthrough_endpoint_router = PassthroughEndpointRouter() @@ -132,3 +133,185 @@ class TestPassthroughEndpointRouter(unittest.TestCase): ), "COHERE_API_KEY", ) + + def test_get_deployment_key(self): + """Test _get_deployment_key with various inputs""" + router = PassthroughEndpointRouter() + + # Test with valid inputs + key = router._get_deployment_key("test-project", "us-central1") + assert key == "test-project-us-central1" + + # Test with None values + key = router._get_deployment_key(None, "us-central1") + assert key is None + + key = router._get_deployment_key("test-project", None) + assert key is None + + key = router._get_deployment_key(None, None) + assert key is None + + def test_add_vertex_credentials(self): + """Test add_vertex_credentials functionality""" + router = PassthroughEndpointRouter() + + # Test adding valid credentials + router.add_vertex_credentials( + project_id="test-project", + location="us-central1", + vertex_credentials='{"credentials": "test-creds"}', + ) + + assert "test-project-us-central1" in router.deployment_key_to_vertex_credentials + creds = router.deployment_key_to_vertex_credentials["test-project-us-central1"] + assert creds.vertex_project == "test-project" + assert creds.vertex_location == "us-central1" + assert creds.vertex_credentials == '{"credentials": "test-creds"}' + + # Test adding with None values + router.add_vertex_credentials( + project_id=None, + location=None, + vertex_credentials='{"credentials": "test-creds"}', + ) + # Should not add None values + assert len(router.deployment_key_to_vertex_credentials) == 1 + + def test_default_credentials(self): + """ + Test get_vertex_credentials with stored credentials. + + Tests if default credentials are used if set. + + Tests if no default credentials are used, if no default set + """ + router = PassthroughEndpointRouter() + router.add_vertex_credentials( + project_id="test-project", + location="us-central1", + vertex_credentials='{"credentials": "test-creds"}', + ) + + creds = router.get_vertex_credentials( + project_id="test-project", location="us-central2" + ) + + assert creds is None + + def test_get_vertex_env_vars(self): + """Test that _get_vertex_env_vars correctly reads environment variables""" + # Set environment variables for the test + os.environ["DEFAULT_VERTEXAI_PROJECT"] = "test-project-123" + os.environ["DEFAULT_VERTEXAI_LOCATION"] = "us-central1" + os.environ["DEFAULT_GOOGLE_APPLICATION_CREDENTIALS"] = "/path/to/creds" + + try: + result = self.router._get_vertex_env_vars() + print(result) + + # Verify the result + assert isinstance(result, VertexPassThroughCredentials) + assert result.vertex_project == "test-project-123" + assert result.vertex_location == "us-central1" + assert result.vertex_credentials == "/path/to/creds" + + finally: + # Clean up environment variables + del os.environ["DEFAULT_VERTEXAI_PROJECT"] + del os.environ["DEFAULT_VERTEXAI_LOCATION"] + del os.environ["DEFAULT_GOOGLE_APPLICATION_CREDENTIALS"] + + def test_set_default_vertex_config(self): + """Test set_default_vertex_config with various inputs""" + # Test with None config - set environment variables first + os.environ["DEFAULT_VERTEXAI_PROJECT"] = "env-project" + os.environ["DEFAULT_VERTEXAI_LOCATION"] = "env-location" + os.environ["DEFAULT_GOOGLE_APPLICATION_CREDENTIALS"] = "env-creds" + os.environ["GOOGLE_CREDS"] = "secret-creds" + + try: + # Test with None config + self.router.set_default_vertex_config() + + assert self.router.default_vertex_config.vertex_project == "env-project" + assert self.router.default_vertex_config.vertex_location == "env-location" + assert self.router.default_vertex_config.vertex_credentials == "env-creds" + + # Test with valid config.yaml settings on vertex_config + test_config = { + "vertex_project": "my-project-123", + "vertex_location": "us-central1", + "vertex_credentials": "path/to/creds", + } + self.router.set_default_vertex_config(test_config) + + assert self.router.default_vertex_config.vertex_project == "my-project-123" + assert self.router.default_vertex_config.vertex_location == "us-central1" + assert ( + self.router.default_vertex_config.vertex_credentials == "path/to/creds" + ) + + # Test with environment variable reference + test_config = { + "vertex_project": "my-project-123", + "vertex_location": "us-central1", + "vertex_credentials": "os.environ/GOOGLE_CREDS", + } + self.router.set_default_vertex_config(test_config) + + assert ( + self.router.default_vertex_config.vertex_credentials == "secret-creds" + ) + + finally: + # Clean up environment variables + del os.environ["DEFAULT_VERTEXAI_PROJECT"] + del os.environ["DEFAULT_VERTEXAI_LOCATION"] + del os.environ["DEFAULT_GOOGLE_APPLICATION_CREDENTIALS"] + del os.environ["GOOGLE_CREDS"] + + def test_vertex_passthrough_router_init(self): + """Test VertexPassThroughRouter initialization""" + router = PassthroughEndpointRouter() + assert isinstance(router.deployment_key_to_vertex_credentials, dict) + assert len(router.deployment_key_to_vertex_credentials) == 0 + + def test_get_vertex_credentials_none(self): + """Test get_vertex_credentials with various inputs""" + router = PassthroughEndpointRouter() + + router.set_default_vertex_config( + config={ + "vertex_project": None, + "vertex_location": None, + "vertex_credentials": None, + } + ) + + # Test with None project_id and location - should return default config + creds = router.get_vertex_credentials(None, None) + assert isinstance(creds, VertexPassThroughCredentials) + + # Test with valid project_id and location but no stored credentials + creds = router.get_vertex_credentials("test-project", "us-central1") + assert isinstance(creds, VertexPassThroughCredentials) + assert creds.vertex_project is None + assert creds.vertex_location is None + assert creds.vertex_credentials is None + + def test_get_vertex_credentials_stored(self): + """Test get_vertex_credentials with stored credentials""" + router = PassthroughEndpointRouter() + router.add_vertex_credentials( + project_id="test-project", + location="us-central1", + vertex_credentials='{"credentials": "test-creds"}', + ) + + creds = router.get_vertex_credentials( + project_id="test-project", location="us-central1" + ) + assert creds.vertex_project == "test-project" + assert creds.vertex_location == "us-central1" + assert creds.vertex_credentials == '{"credentials": "test-creds"}' diff --git a/tests/pass_through_unit_tests/test_unit_test_vertex_pass_through.py b/tests/pass_through_unit_tests/test_unit_test_vertex_pass_through.py index 9b354a84c9..066e434c26 100644 --- a/tests/pass_through_unit_tests/test_unit_test_vertex_pass_through.py +++ b/tests/pass_through_unit_tests/test_unit_test_vertex_pass_through.py @@ -26,292 +26,3 @@ from litellm.proxy.vertex_ai_endpoints.vertex_endpoints import ( from litellm.proxy.vertex_ai_endpoints.vertex_passthrough_router import ( VertexPassThroughRouter, ) - - -@pytest.mark.asyncio -async def test_get_litellm_virtual_key(): - """ - Test that the get_litellm_virtual_key function correctly handles the API key authentication - """ - # Test with x-litellm-api-key - mock_request = Mock() - mock_request.headers = {"x-litellm-api-key": "test-key-123"} - result = get_litellm_virtual_key(mock_request) - assert result == "Bearer test-key-123" - - # Test with Authorization header - mock_request.headers = {"Authorization": "Bearer auth-key-456"} - result = get_litellm_virtual_key(mock_request) - assert result == "Bearer auth-key-456" - - # Test with both headers (x-litellm-api-key should take precedence) - mock_request.headers = { - "x-litellm-api-key": "test-key-123", - "Authorization": "Bearer auth-key-456", - } - result = get_litellm_virtual_key(mock_request) - assert result == "Bearer test-key-123" - - -@pytest.mark.asyncio -async def test_async_vertex_proxy_route_api_key_auth(): - """ - Critical - - This is how Vertex AI JS SDK will Auth to Litellm Proxy - """ - # Mock dependencies - mock_request = Mock() - mock_request.headers = {"x-litellm-api-key": "test-key-123"} - mock_request.method = "POST" - mock_response = Mock() - - with patch( - "litellm.proxy.vertex_ai_endpoints.vertex_endpoints.user_api_key_auth" - ) as mock_auth: - mock_auth.return_value = {"api_key": "test-key-123"} - - with patch( - "litellm.proxy.vertex_ai_endpoints.vertex_endpoints.create_pass_through_route" - ) as mock_pass_through: - mock_pass_through.return_value = AsyncMock( - return_value={"status": "success"} - ) - - # Call the function - result = await vertex_proxy_route( - endpoint="v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-1.5-pro:generateContent", - request=mock_request, - fastapi_response=mock_response, - ) - - # Verify user_api_key_auth was called with the correct Bearer token - mock_auth.assert_called_once() - call_args = mock_auth.call_args[1] - assert call_args["api_key"] == "Bearer test-key-123" - - -@pytest.mark.asyncio -async def test_get_vertex_env_vars(): - """Test that _get_vertex_env_vars correctly reads environment variables""" - # Set environment variables for the test - os.environ["DEFAULT_VERTEXAI_PROJECT"] = "test-project-123" - os.environ["DEFAULT_VERTEXAI_LOCATION"] = "us-central1" - os.environ["DEFAULT_GOOGLE_APPLICATION_CREDENTIALS"] = "/path/to/creds" - - try: - result = _get_vertex_env_vars() - print(result) - - # Verify the result - assert isinstance(result, VertexPassThroughCredentials) - assert result.vertex_project == "test-project-123" - assert result.vertex_location == "us-central1" - assert result.vertex_credentials == "/path/to/creds" - - finally: - # Clean up environment variables - del os.environ["DEFAULT_VERTEXAI_PROJECT"] - del os.environ["DEFAULT_VERTEXAI_LOCATION"] - del os.environ["DEFAULT_GOOGLE_APPLICATION_CREDENTIALS"] - - -@pytest.mark.asyncio -async def test_set_default_vertex_config(): - """Test set_default_vertex_config with various inputs""" - # Test with None config - set environment variables first - os.environ["DEFAULT_VERTEXAI_PROJECT"] = "env-project" - os.environ["DEFAULT_VERTEXAI_LOCATION"] = "env-location" - os.environ["DEFAULT_GOOGLE_APPLICATION_CREDENTIALS"] = "env-creds" - os.environ["GOOGLE_CREDS"] = "secret-creds" - - try: - # Test with None config - set_default_vertex_config() - from litellm.proxy.vertex_ai_endpoints.vertex_endpoints import ( - default_vertex_config, - ) - - assert default_vertex_config.vertex_project == "env-project" - assert default_vertex_config.vertex_location == "env-location" - assert default_vertex_config.vertex_credentials == "env-creds" - - # Test with valid config.yaml settings on vertex_config - test_config = { - "vertex_project": "my-project-123", - "vertex_location": "us-central1", - "vertex_credentials": "path/to/creds", - } - set_default_vertex_config(test_config) - from litellm.proxy.vertex_ai_endpoints.vertex_endpoints import ( - default_vertex_config, - ) - - assert default_vertex_config.vertex_project == "my-project-123" - assert default_vertex_config.vertex_location == "us-central1" - assert default_vertex_config.vertex_credentials == "path/to/creds" - - # Test with environment variable reference - test_config = { - "vertex_project": "my-project-123", - "vertex_location": "us-central1", - "vertex_credentials": "os.environ/GOOGLE_CREDS", - } - set_default_vertex_config(test_config) - from litellm.proxy.vertex_ai_endpoints.vertex_endpoints import ( - default_vertex_config, - ) - - assert default_vertex_config.vertex_credentials == "secret-creds" - - finally: - # Clean up environment variables - del os.environ["DEFAULT_VERTEXAI_PROJECT"] - del os.environ["DEFAULT_VERTEXAI_LOCATION"] - del os.environ["DEFAULT_GOOGLE_APPLICATION_CREDENTIALS"] - del os.environ["GOOGLE_CREDS"] - - -@pytest.mark.asyncio -async def test_vertex_passthrough_router_init(): - """Test VertexPassThroughRouter initialization""" - router = VertexPassThroughRouter() - assert isinstance(router.deployment_key_to_vertex_credentials, dict) - assert len(router.deployment_key_to_vertex_credentials) == 0 - - -@pytest.mark.asyncio -async def test_get_vertex_credentials_none(): - """Test get_vertex_credentials with various inputs""" - from litellm.proxy.vertex_ai_endpoints import vertex_endpoints - - setattr(vertex_endpoints, "default_vertex_config", VertexPassThroughCredentials()) - router = VertexPassThroughRouter() - - # Test with None project_id and location - should return default config - creds = router.get_vertex_credentials(None, None) - assert isinstance(creds, VertexPassThroughCredentials) - - # Test with valid project_id and location but no stored credentials - creds = router.get_vertex_credentials("test-project", "us-central1") - assert isinstance(creds, VertexPassThroughCredentials) - assert creds.vertex_project is None - assert creds.vertex_location is None - assert creds.vertex_credentials is None - - -@pytest.mark.asyncio -async def test_get_vertex_credentials_stored(): - """Test get_vertex_credentials with stored credentials""" - router = VertexPassThroughRouter() - router.add_vertex_credentials( - project_id="test-project", - location="us-central1", - vertex_credentials='{"credentials": "test-creds"}', - ) - - creds = router.get_vertex_credentials( - project_id="test-project", location="us-central1" - ) - assert creds.vertex_project == "test-project" - assert creds.vertex_location == "us-central1" - assert creds.vertex_credentials == '{"credentials": "test-creds"}' - - -@pytest.mark.asyncio -async def test_default_credentials(): - """ - Test get_vertex_credentials with stored credentials. - - Tests if default credentials are used if set. - - Tests if no default credentials are used, if no default set - """ - router = VertexPassThroughRouter() - router.add_vertex_credentials( - project_id="test-project", - location="us-central1", - vertex_credentials='{"credentials": "test-creds"}', - ) - - creds = router.get_vertex_credentials( - project_id="test-project", location="us-central2" - ) - - assert creds is None - - -@pytest.mark.asyncio -async def test_add_vertex_credentials(): - """Test add_vertex_credentials functionality""" - router = VertexPassThroughRouter() - - # Test adding valid credentials - router.add_vertex_credentials( - project_id="test-project", - location="us-central1", - vertex_credentials='{"credentials": "test-creds"}', - ) - - assert "test-project-us-central1" in router.deployment_key_to_vertex_credentials - creds = router.deployment_key_to_vertex_credentials["test-project-us-central1"] - assert creds.vertex_project == "test-project" - assert creds.vertex_location == "us-central1" - assert creds.vertex_credentials == '{"credentials": "test-creds"}' - - # Test adding with None values - router.add_vertex_credentials( - project_id=None, - location=None, - vertex_credentials='{"credentials": "test-creds"}', - ) - # Should not add None values - assert len(router.deployment_key_to_vertex_credentials) == 1 - - -@pytest.mark.asyncio -async def test_get_deployment_key(): - """Test _get_deployment_key with various inputs""" - router = VertexPassThroughRouter() - - # Test with valid inputs - key = router._get_deployment_key("test-project", "us-central1") - assert key == "test-project-us-central1" - - # Test with None values - key = router._get_deployment_key(None, "us-central1") - assert key is None - - key = router._get_deployment_key("test-project", None) - assert key is None - - key = router._get_deployment_key(None, None) - assert key is None - - -@pytest.mark.asyncio -async def test_get_vertex_project_id_from_url(): - """Test _get_vertex_project_id_from_url with various URLs""" - # Test with valid URL - url = "https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:streamGenerateContent" - project_id = VertexPassThroughRouter._get_vertex_project_id_from_url(url) - assert project_id == "test-project" - - # Test with invalid URL - url = "https://invalid-url.com" - project_id = VertexPassThroughRouter._get_vertex_project_id_from_url(url) - assert project_id is None - - -@pytest.mark.asyncio -async def test_get_vertex_location_from_url(): - """Test _get_vertex_location_from_url with various URLs""" - # Test with valid URL - url = "https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:streamGenerateContent" - location = VertexPassThroughRouter._get_vertex_location_from_url(url) - assert location == "us-central1" - - # Test with invalid URL - url = "https://invalid-url.com" - location = VertexPassThroughRouter._get_vertex_location_from_url(url) - assert location is None