diff --git a/litellm/__init__.py b/litellm/__init__.py index 0d6a788e368..f615766578d 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -260,6 +260,7 @@ azure_key: Optional[str] = None anthropic_key: Optional[str] = None replicate_key: Optional[str] = None bytez_key: Optional[str] = None +gdc_key: Optional[str] = None cohere_key: Optional[str] = None infinity_key: Optional[str] = None clarifai_key: Optional[str] = None diff --git a/litellm/constants.py b/litellm/constants.py index a3ea68c7949..a33cd940166 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -550,6 +550,7 @@ LITELLM_CHAT_PROVIDERS = [ "openai", "openai_like", "bytez", + "gdc", "xai", "custom_openai", "text-completion-openai", diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index 4941d52d7d6..ac45b58420b 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -506,6 +506,8 @@ def get_llm_provider( # bytez models elif model.startswith("bytez/"): custom_llm_provider = "bytez" + elif model.startswith("gdc/"): + custom_llm_provider = "gdc" elif model.startswith("lemonade/"): custom_llm_provider = "lemonade" elif model.startswith("heroku/"): diff --git a/litellm/llms/gdc/__init__.py b/litellm/llms/gdc/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/gdc/chat/__init__.py b/litellm/llms/gdc/chat/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/gdc/chat/transformation.py b/litellm/llms/gdc/chat/transformation.py new file mode 100644 index 00000000000..79b569dc8aa --- /dev/null +++ b/litellm/llms/gdc/chat/transformation.py @@ -0,0 +1,211 @@ +""" +GDC Gemini chat completion transformation +""" + +import litellm +from typing import Any, List, Optional +from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig + + +class GDCGeminiConfig(OpenAILikeChatConfig): + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + api_base = ( + api_base or litellm.api_base or getattr(litellm, "gdc_api_base", None) + ) + if not api_base: + raise litellm.utils.AuthenticationError( + message="api_base/host is required for GDC Gemini. Please set it or pass it.", + llm_provider="gdc", + model=model, + ) + + if not api_base.startswith("http"): + api_base = f"https://{api_base}" + + project = ( + optional_params.get("vertex_project") + or litellm_params.get("vertex_project") + or getattr(litellm, "vertex_project", None) + ) + if not project: + raise litellm.utils.AuthenticationError( + message="project is required for GDC Gemini. Please pass vertex_project.", + llm_provider="gdc", + model=model, + ) + + api_base = api_base.rstrip("/") + + # If the endpoint structure is already in the api_base, don't append it again + if "/v1/projects/" in api_base: + base_url = api_base + else: + base_url = f"{api_base}/v1/projects/{project}/locations/{project}" + + endpoint = "chat/completions" + if "chat/completions" in base_url: + return base_url + + return f"{base_url}/{endpoint}" + + def validate_environment( + self, + headers: dict, + model: str, + messages: List[Any], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + import json + import google.auth + import requests + from google.auth.transport import requests as auth_requests + + # Extract GDC API base + api_base = ( + api_base or litellm.api_base or getattr(litellm, "gdc_api_base", None) + ) + + if not api_base: + raise litellm.utils.AuthenticationError( + message="api_base/host is required for GDC Gemini. Please set it or pass it.", + llm_provider="gdc", + model=model, + ) + + if not api_key: + raise litellm.utils.AuthenticationError( + message="api_key is required for GDC Gemini. Please pass your service account string or token as the api_key.", + llm_provider="gdc", + model=model, + ) + + project = ( + optional_params.get("vertex_project") + or litellm_params.get("vertex_project") + or getattr(litellm, "vertex_project", None) + ) + if not project: + raise litellm.utils.AuthenticationError( + message="project is required for GDC Gemini. Please pass vertex_project.", + llm_provider="gdc", + model=model, + ) + + # Ensure we have the audience for token fetch + audience = api_base + if not audience.startswith("http"): + audience = f"https://{audience}" + + # Generate GDC token + is_service_account = False + try: + creds = None + if api_key: + import os + + try: + # Check if api_key is a file path + # Limit length to avoid OSError for 'File name too long' + if len(api_key) < 2000 and os.path.exists(api_key): + with open(api_key, "r") as f: + json_obj = json.load(f) + + is_service_account = True + else: + json_obj = json.loads(api_key) + is_service_account = True + + if is_service_account: + creds, _ = google.auth.load_credentials_from_dict(json_obj) + except (json.JSONDecodeError, OSError, UnicodeDecodeError): + # It's not a valid JSON string or file path, treat as raw token + is_service_account = False + except Exception as e: + raise litellm.utils.AuthenticationError( + message=f"Failed to load service account credentials from api_key: {str(e)}", + llm_provider="gdc", + model=model, + ) + + if creds is not None: + # if hasattr(creds, "with_gdch_audience"): + creds = creds.with_gdch_audience(audience) + auth_session = requests.Session() + + # Retrieve ssl_verify configuration + ssl_verify = litellm_params.get("ssl_verify") + if ssl_verify is None: + import os + + ssl_verify = os.getenv("SSL_VERIFY", True) + if isinstance(ssl_verify, str): + if ssl_verify.lower() == "false": + ssl_verify = False + elif ssl_verify.lower() == "true": + ssl_verify = True + + auth_session.verify = ssl_verify + auth_request = auth_requests.Request(session=auth_session) + + creds.refresh(auth_request) + token = creds.token + headers["Authorization"] = f"Bearer {token}" + except Exception as e: + raise e + + if "Authorization" not in headers and api_key and not is_service_account: + headers["Authorization"] = f"Bearer {api_key}" + + # Ensure Content-Type is set to application/json + if "content-type" not in headers and "Content-Type" not in headers: + headers["Content-Type"] = "application/json" + + # Ensure Vertex Project is included in headers + if ( + "x-goog-user-project" not in headers + and "X-Goog-User_project" not in headers + ): + headers["x-goog-user-project"] = f"projects/{project}" + + return headers + + def transform_request( + self, + model: str, + messages: List[Any], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + """ + Transforms the request to the GDC provider + """ + # Strip provider prefix for GDC + if model.startswith("gdc/"): + model = model.split("/", 1)[1] + + data = super().transform_request( + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=headers, + ) + + # Remove extra params used for routing/auth + data.pop("vertex_project", None) + data.pop("ssl_verify", None) + + return data diff --git a/litellm/main.py b/litellm/main.py index 80176cc8b16..38fa63ebf34 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -200,6 +200,7 @@ from .llms.bedrock.embed.embedding import BedrockEmbedding from .llms.bedrock.image_edit.handler import BedrockImageEdit from .llms.bedrock.image_generation.image_handler import BedrockImageGeneration from .llms.bytez.chat.transformation import BytezChatConfig +from .llms.gdc.chat.transformation import GDCGeminiConfig from .llms.clarifai.chat.transformation import ClarifaiConfig from .llms.codestral.completion.handler import CodestralTextCompletion from .llms.cohere.embed import handler as cohere_embed @@ -308,6 +309,7 @@ google_batch_embeddings = GoogleBatchEmbeddings() vertex_partner_models_chat_completion = VertexAIPartnerModels() vertex_gemma_chat_completion = VertexAIGemmaModels() vertex_model_garden_chat_completion = VertexAIModelGardenModels() +gdc_transformation = GDCGeminiConfig() # vertex_text_to_speech is now replaced by VertexAITextToSpeechConfig sagemaker_llm = SagemakerLLM() watsonx_chat_completion = WatsonXChatHandler() @@ -4345,6 +4347,33 @@ def completion( # type: ignore logging_obj=logging, ) + elif custom_llm_provider == "gdc": + api_key = ( + api_key + or litellm.gdc_key + or get_secret_str("GDC_API_KEY") + or litellm.api_key + ) + + response = base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + timeout=timeout, # type: ignore + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + provider_config=gdc_transformation, + ) + elif custom_llm_provider == "bytez": api_key = ( api_key diff --git a/litellm/utils.py b/litellm/utils.py index 30b5691a140..34efe260a6c 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -3878,6 +3878,7 @@ class PreProcessNonDefaultParams: k.startswith("vertex_") and custom_llm_provider != "vertex_ai" and custom_llm_provider != "vertex_ai_beta" + and custom_llm_provider != "gdc" ): # allow dynamically setting vertex ai init logic continue passed_params[k] = v diff --git a/tests/test_litellm/llms/gdc/chat/test_gdc_chat_transformation.py b/tests/test_litellm/llms/gdc/chat/test_gdc_chat_transformation.py new file mode 100644 index 00000000000..5e87de75123 --- /dev/null +++ b/tests/test_litellm/llms/gdc/chat/test_gdc_chat_transformation.py @@ -0,0 +1,104 @@ +import os +import sys +import pytest +from unittest.mock import MagicMock, patch + +# Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.gdc.chat.transformation import GDCGeminiConfig + +TEST_API_KEY = '{"type": "service_account", "project_id": "test-project"}' +TEST_MODEL = "gdc/gemini-2.5-flash" +TEST_API_BASE = "https://gdc-endpoint.com" +TEST_PROJECT = "test-project" + + +class TestGDCGeminiConfig: + def test_get_complete_url(self): + config = GDCGeminiConfig() + + # Test basic URL formatting + url = config.get_complete_url( + api_base=TEST_API_BASE, + api_key=None, + model=TEST_MODEL, + optional_params={"vertex_project": TEST_PROJECT}, + litellm_params={}, + ) + assert url == f"{TEST_API_BASE}/v1/projects/{TEST_PROJECT}/locations/{TEST_PROJECT}/chat/completions" + + def test_get_complete_url_missing_api_base(self): + config = GDCGeminiConfig() + with pytest.raises(Exception, match="api_base/host is required for GDC Gemini"): + config.get_complete_url( + api_base=None, + api_key=None, + model=TEST_MODEL, + optional_params={"vertex_project": TEST_PROJECT}, + litellm_params={}, + ) + + def test_get_complete_url_missing_project(self): + config = GDCGeminiConfig() + with pytest.raises(Exception, match="project is required for GDC Gemini"): + config.get_complete_url( + api_base=TEST_API_BASE, + api_key=None, + model=TEST_MODEL, + optional_params={}, + litellm_params={}, + ) + + @patch("google.auth.load_credentials_from_dict") + @patch("requests.Session") + def test_validate_environment(self, mock_session, mock_load_creds): + # Setup mocks for google.auth and requests.Session + mock_creds = MagicMock() + mock_creds.token = "mock-token" + mock_creds.with_gdch_audience.return_value = mock_creds + mock_load_creds.return_value = (mock_creds, None) + + mock_session_instance = MagicMock() + mock_session.return_value = mock_session_instance + + config = GDCGeminiConfig() + headers = {} + + result = config.validate_environment( + headers=headers, + model=TEST_MODEL, + messages=[], + optional_params={"vertex_project": TEST_PROJECT}, + litellm_params={}, + api_key=TEST_API_KEY, + api_base=TEST_API_BASE, + ) + + assert result["Authorization"] == "Bearer mock-token" + assert result["Content-Type"] == "application/json" + assert result["x-goog-user-project"] == f"projects/{TEST_PROJECT}" + + mock_creds.with_gdch_audience.assert_called_once_with(TEST_API_BASE) + mock_creds.refresh.assert_called_once() + assert mock_session_instance.verify is True + + def test_transform_request(self): + config = GDCGeminiConfig() + + messages = [{"role": "user", "content": "Hello"}] + headers = {} + + data = config.transform_request( + model=TEST_MODEL, + messages=messages, + optional_params={"vertex_project": TEST_PROJECT}, + litellm_params={"ssl_verify": True}, + headers=headers, + ) + + # Verify provider prefix is stripped + assert data["model"] == "gemini-2.5-flash" + # Verify routing/auth fields are popped + assert "vertex_project" not in data + assert "ssl_verify" not in data