From fde4c7c97ae49370dc7986326aace4c5fb7bd91b Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 1 Jul 2026 17:31:07 -0700 Subject: [PATCH] feat(gdc): implement Google Distributed Cloud (GDC) Gemini provider (#31895) * feat(gdc): add Google Distributed Cloud Gemini provider support Introduce support for the Google Distributed Cloud (GDC) Gemini provider by adding "gdc" to the list of chat providers and enabling the gdc/ model prefix. The implementation defines a new GDCGeminiConfig class which handles authentication via Google Distributed Cloud service account credentials, manages token generation, formats GDC Gemini request URLs, and transforms request structures accordingly The PreProcessNonDefaultParams class is also updated to exclude vertex parameters from filtering when the custom LLM provider is GDC, allowing vertex parameters to be passed properly during GDC initialization * fix: resolve issues identified in PR #30702 * fix(gdc): harden credentials, fix vertex param filtering, add tests The supports_vertex_params branch regressed vertex_ai and vertex_ai_beta: the `if custom_llm_provider in [...]: pass` was a no-op, so those providers fell through to the config lookup, found no supports_vertex_params, and had their vertex_ params stripped. The check is now a single _provider_supports_vertex_params helper that keeps vertex_ params for the vertex family and for any config that opts in, and only swallows the expected ValueError from an unknown provider string instead of a blanket except GDC project and location now resolve from the deployment's litellm_params and the litellm.vertex_project / litellm.vertex_location globals before falling back to request optional_params, matching how vertex_ai resolves them, so a proxy caller can no longer route a request to a project the deployment did not expose A request api_key is no longer treated as a filesystem path, so a caller can't make the host open a local service-account file; api_key must be a literal service-account JSON string or a bearer token The opt-in token cache is hardened: the lock and cache dict are created in __init__ instead of via a racy hasattr lazy-init, the token is read inside the lock, and the audience is stripped of a trailing slash once so the cached and non-cached paths agree Also declares gdc_api_base, switches the lazy-import entry to the relative path every other entry uses, adds the missing trailing comma in the provider config map, and drops the api_base fallback that only ran when api_key was None Adds unit tests covering the vertex-param filter, deployment-over-request precedence, the api_key file-path rejection, URL construction branches, environment validation, token caching, and the gdc completion dispatch; transformation.py is fully covered * fix(gdc): prefer GDC-specific config, honor vertex_ai aliases, harden URL and bool parsing * fix(gdc): mint the GDCH token audience from the host, not the full base When api_base embedded /v1/projects/... and the deployment set project/location, get_complete_url rebuilt the request URL from the host while validate_environment still derived the token audience from the full original api_base, so the bearer token could target a different audience than the URL actually called. The audience is now the scheme://host of api_base in every case, matching the host get_complete_url builds against * fix(gdc): restrict JSON api_key to GDCH service accounts Only accept a credential whose type is gdch_service_account before calling google.auth.load_credentials_from_dict, so a caller-supplied external_account/identity_pool/pluggable credential carrying arbitrary token or credential_source endpoints is rejected before any token refresh runs. GDC only ever uses GDCH service accounts, and non-GDCH credentials could not have completed auth anyway (with_gdch_audience is GDCH-only), so this narrows the credential-refresh surface without changing valid GDC behavior. * fix(gdc): validate project and location as plain identifiers vertex_project and vertex_location can come from request params and were interpolated as raw path text into the GDC request URL and the x-goog-user-project header. A caller-supplied value containing / ? # or .. could reshape the path and make the proxy send its GDC-authorized request to a different endpoint under the configured host. Validate both against a strict identifier pattern before building the URL or header and raise an auth error otherwise; GCP project ids and locations are plain identifiers so valid deployments are unaffected. * fix(gdc): bind x-goog-user-project quota header to the deployment The quota project header was resolved with request-level vertex_project taking effect, so with a preformed deployment api_base a caller could set vertex_project to a different project and have it sent under the proxy's GDC credential, misattributing quota or billing. Resolve the header project the same way the URL is resolved: a preformed api_base without a deployment override binds to the project embedded in the URL, otherwise deployment and global config win over request params. This keeps the URL and the quota header consistent. * fix(gdc): always rebind x-goog-user-project, stripping caller-forwarded values The quota project header was only set when absent, so with client header forwarding an authenticated caller could send their own x-goog-user-project (any casing) and have it ride on the proxy's GDC credential, bypassing the deployment-derived binding. Strip every casing of the header and always set it from _effective_project before the request is signed. * fix(gdc): make a preformed api_base authoritative for project routing get_litellm_params copies caller-supplied vertex_project and vertex_location into litellm_params via OPTIONAL_KWARGS_KEYS, so litellm_params cannot be treated as a deployment-only source. The previous _deployment_overrides_path inference let an authenticated caller flip a pinned preformed api_base such as /v1/projects/pinned/... to /v1/projects/attacker/..., driving requests to a caller-chosen project with the proxy's configured GDC credentials and quota header A preformed /v1/projects/ api_base is now authoritative; get_complete_url returns it unchanged and _effective_project binds the x-goog-user-project quota header to the project embedded in that URL, so a caller can no longer redirect a pinned deployment or move the quota header off it. The two tests that asserted the override behavior are now regression tests that fail if the rewrite is reintroduced * fix(gdc): make a preformed api_base self-sufficient in get_complete_url get_complete_url resolved and required a params-derived vertex_project before returning a preformed /v1/projects/ api_base, so a deployment that pins its project in the api_base path was forced to also pass vertex_project or hit 'project is required'. validate_environment already extracts the project from a preformed URL and needs no such param, so the two paths disagreed The preformed-URL early return now runs before project/location resolution, matching validate_environment: a preformed api_base is returned as-is with no redundant param, and non-preformed bases still require vertex_project and vertex_location as before. Adds a regression test that a preformed base with no project/location params returns the URL unchanged --------- Co-authored-by: Paige O'Connor Co-authored-by: Tim Laubach --- litellm/__init__.py | 3 + litellm/_lazy_imports_registry.py | 5 + litellm/constants.py | 1 + .../get_llm_provider_logic.py | 2 + litellm/llms/gdc/__init__.py | 0 litellm/llms/gdc/chat/__init__.py | 0 litellm/llms/gdc/chat/transformation.py | 285 +++++++ litellm/main.py | 43 ++ litellm/types/utils.py | 1 + litellm/utils.py | 21 +- provider_endpoints_support.json | 10 + ruff-strict-budget.json | 2 +- .../gdc/chat/test_gdc_chat_transformation.py | 717 ++++++++++++++++++ tests/test_litellm/test_utils.py | 42 + 14 files changed, 1126 insertions(+), 6 deletions(-) create mode 100644 litellm/llms/gdc/__init__.py create mode 100644 litellm/llms/gdc/chat/__init__.py create mode 100644 litellm/llms/gdc/chat/transformation.py create mode 100644 tests/test_litellm/llms/gdc/chat/test_gdc_chat_transformation.py diff --git a/litellm/__init__.py b/litellm/__init__.py index 15e95ded906..9327e121b1d 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -263,6 +263,8 @@ 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 +gdc_api_base: Optional[str] = None cohere_key: Optional[str] = None infinity_key: Optional[str] = None clarifai_key: Optional[str] = None @@ -1787,6 +1789,7 @@ if TYPE_CHECKING: from .llms.nvidia_nim.embed import ( NvidiaNimEmbeddingConfig as NvidiaNimEmbeddingConfig, ) + from .llms.gdc.chat.transformation import GDCGeminiConfig as GDCGeminiConfig # Type stubs for lazy-loaded config instances openaiOSeriesConfig: OpenAIOSeriesConfig diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 4f131354d2e..0f9d3a560d1 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -323,6 +323,7 @@ LLM_CONFIG_NAMES = ( "SnowflakeEmbeddingConfig", "AmazonNovaChatConfig", "SonioxAudioTranscriptionConfig", + "GDCGeminiConfig", ) # Types that support lazy loading via _lazy_import_types @@ -1157,6 +1158,10 @@ _LLM_CONFIGS_IMPORT_MAP = { ".llms.dashscope.chat.transformation", "DashScopeChatConfig", ), + "GDCGeminiConfig": ( + ".llms.gdc.chat.transformation", + "GDCGeminiConfig", + ), "ModelScopeChatConfig": ( ".llms.modelscope.chat.transformation", "ModelScopeChatConfig", diff --git a/litellm/constants.py b/litellm/constants.py index 93aed83852b..6eb2779dcae 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -460,6 +460,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 122d09c855b..a7a576ff167 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -446,6 +446,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..61631920a64 --- /dev/null +++ b/litellm/llms/gdc/chat/transformation.py @@ -0,0 +1,285 @@ +""" +GDC Gemini chat completion transformation +""" + +import json +import os +import re +import threading +from typing import Any, Final +from urllib.parse import urlsplit + +import litellm +from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig + + +class GDCGeminiConfig(OpenAILikeChatConfig): + supports_vertex_params: bool = True # Tell LiteLLM utilities not to strip vertex_ params + _GDCH_CREDENTIAL_TYPE: Final[str] = "gdch_service_account" + _PATH_ID_PATTERN: Final[re.Pattern[str]] = re.compile(r"^[a-zA-Z0-9_-]+$") + + def __init__(self, **kwargs: Any) -> None: + super().__init__(**kwargs) + self._creds_lock = threading.Lock() + self._gdch_creds_cache: dict = {} + + def get_supported_openai_params(self, model: str) -> list: + return [ + "vertex_project", + "vertex_location", + ] + super().get_supported_openai_params(model) + + def _resolve_project(self, optional_params: dict, litellm_params: dict) -> str | None: + return ( + litellm_params.get("vertex_project") + or litellm_params.get("vertex_ai_project") + or getattr(litellm, "vertex_project", None) + or optional_params.get("vertex_project") + or optional_params.get("vertex_ai_project") + ) + + def _resolve_location(self, optional_params: dict, litellm_params: dict) -> str | None: + return ( + litellm_params.get("vertex_location") + or litellm_params.get("vertex_ai_location") + or getattr(litellm, "vertex_location", None) + or optional_params.get("vertex_location") + or optional_params.get("vertex_ai_location") + ) + + def _effective_project(self, api_base: str, optional_params: dict, litellm_params: dict) -> str | None: + match = re.search(r"/v1/projects/([^/]+)", api_base) + if match: + return match.group(1) + return self._resolve_project(optional_params, litellm_params) + + def _validate_path_id(self, value: str, field: str, model: str) -> str: + if not self._PATH_ID_PATTERN.match(value): + raise litellm.utils.AuthenticationError( + message=f"{field} must be a plain identifier of letters, digits, hyphens or underscores.", + llm_provider="gdc", + model=model, + ) + return value + + def get_complete_url( + self, + api_base: str | None, + api_key: str | None, + model: str, + optional_params: dict, + litellm_params: dict, + stream: bool | None = None, + ) -> str: + api_base = api_base or litellm.gdc_api_base or litellm.api_base + 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}" + + api_base = api_base.rstrip("/") + + if "/v1/projects/" in api_base: + return api_base + + project = self._resolve_project(optional_params, litellm_params) + + if not project: + raise litellm.utils.AuthenticationError( + message="project is required for GDC Gemini. Please pass vertex_project.", + llm_provider="gdc", + model=model, + ) + + location = self._resolve_location(optional_params, litellm_params) + + if not location: + raise litellm.utils.AuthenticationError( + message="location is required for GDC Gemini. Please pass vertex_location.", + llm_provider="gdc", + model=model, + ) + + project = self._validate_path_id(project, "vertex_project", model) + location = self._validate_path_id(location, "vertex_location", model) + + return f"{api_base}/v1/projects/{project}/locations/{location}/chat/completions" + + def _read_env_bool(self, val: Any, env_var: str, default: bool = True) -> bool | str: + def _parse(s: str) -> bool | str: + cleaned = s.strip().lower() + if cleaned in ("false", "0", "no", "off"): + return False + if cleaned in ("true", "1", "yes", "on"): + return True + return s + + if val is not None: + if isinstance(val, str): + return _parse(val) + return val + + _env_val = os.getenv(env_var) + if _env_val is None: + return default + return _parse(_env_val) + + def _fetch_auth(self, gdch_creds: Any, ssl_verify: bool | str) -> None: + import requests + from google.auth.transport import requests as auth_requests + + auth_session = requests.Session() + auth_session.verify = ssl_verify + auth_request = auth_requests.Request(session=auth_session) + gdch_creds.refresh(auth_request) + + def _cached_fetch_token(self, creds: Any, audience: str, ssl_verify: bool | str, api_key: str | None = None) -> str: + # Key cache by both audience and credential identity to prevent cross-caller contamination + cache_key = (audience.rstrip("/"), api_key or str(id(creds))) + + with self._creds_lock: + if cache_key not in self._gdch_creds_cache: + self._gdch_creds_cache[cache_key] = creds.with_gdch_audience(audience.rstrip("/")) + + gdch_creds = self._gdch_creds_cache[cache_key] + + if not getattr(gdch_creds, "valid", False) or not getattr(gdch_creds, "token", None): + self._fetch_auth(gdch_creds, ssl_verify) + + token = gdch_creds.token + + return token + + def _load_creds_from_key(self, api_key: str) -> tuple[Any, bool]: + import google.auth + + try: + json_obj = json.loads(api_key) + except json.JSONDecodeError: + return None, False + if not isinstance(json_obj, dict) or json_obj.get("type") != self._GDCH_CREDENTIAL_TYPE: + raise ValueError( + "GDC only accepts a GDCH service account credential as a JSON api_key " + '(expected "type": "gdch_service_account"). Other Google credential types are ' + "rejected so their token or external-account endpoints cannot drive server-side requests." + ) + creds, _ = google.auth.load_credentials_from_dict(json_obj) + return creds, True + + def validate_environment( + self, + headers: dict, + model: str, + messages: list[Any], + optional_params: dict, + litellm_params: dict, + api_key: str | None = None, + api_base: str | None = None, + ) -> dict: + import google.auth.exceptions + + api_base = api_base or litellm.gdc_api_base or litellm.api_base + 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 = self._effective_project(api_base, optional_params, litellm_params) + if not project: + raise litellm.utils.AuthenticationError( + message="project is required for GDC Gemini. Please pass vertex_project.", + llm_provider="gdc", + model=model, + ) + project = self._validate_path_id(project, "vertex_project", model) + + _audience_parts = urlsplit(api_base if api_base.startswith("http") else f"https://{api_base}") + audience = f"{_audience_parts.scheme}://{_audience_parts.netloc}" + + try: + creds, is_service_account = self._load_creds_from_key(api_key) + except ( + google.auth.exceptions.GoogleAuthError, + ValueError, + TypeError, + KeyError, + AttributeError, + ) as e: + raise litellm.utils.AuthenticationError( + message=f"Failed to load service account credentials from api_key: {str(e)}", + llm_provider="gdc", + model=model, + ) from e + + if creds is not None: + ssl_verify = self._read_env_bool(litellm_params.get("ssl_verify"), "SSL_VERIFY", default=True) + if self._read_env_bool(litellm_params.get("gdc_token_caching"), "GDC_TOKEN_CACHING", default=False): + token = self._cached_fetch_token(creds, audience, ssl_verify, api_key) + else: + gdch_creds = creds.with_gdch_audience(audience) + self._fetch_auth(gdch_creds, ssl_verify) + token = gdch_creds.token + headers["Authorization"] = f"Bearer {token}" + + if "Authorization" not in headers and not is_service_account: + headers["Authorization"] = f"Bearer {api_key}" + + # Standardize necessary metadata headers + if "content-type" not in headers and "Content-Type" not in headers: + headers["Content-Type"] = "application/json" + + stale_quota_headers = tuple(h for h in headers if h.lower() == "x-goog-user-project") + for stale in stale_quota_headers: + headers.pop(stale, None) + 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 + """ + 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 + for param in [ + "vertex_project", + "vertex_ai_project", + "vertex_location", + "vertex_ai_location", + "ssl_verify", + "gdc_token_caching", + ]: + data.pop(param, None) + + return data diff --git a/litellm/main.py b/litellm/main.py index 18d2c367f8d..567930a4999 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -210,6 +210,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 @@ -318,6 +319,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() @@ -4336,6 +4338,45 @@ def _complete_gradient_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatc ) +def _complete_gdc(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + stream = ctx.stream + timeout = ctx.timeout + + api_key = api_key or litellm.gdc_key or get_secret_str("GDC_API_KEY") or litellm.api_key + api_base = api_base or litellm.gdc_api_base or get_secret_str("GDC_API_BASE") or litellm.api_base + + return 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, + ) + + def _complete_bytez(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: acompletion = ctx.acompletion api_base = ctx.api_base @@ -5533,6 +5574,8 @@ def completion( # type: ignore elif custom_llm_provider == "gradient_ai": response = _complete_gradient_ai(_dispatch_ctx) + elif custom_llm_provider == "gdc": + response = _complete_gdc(_dispatch_ctx) elif custom_llm_provider == "bytez": response = _complete_bytez(_dispatch_ctx) elif custom_llm_provider == "lemonade": diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 4f0c0c21bce..0e2a6b24806 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3369,6 +3369,7 @@ class LlmProviders(str, Enum): LITELLM_AGENT = "litellm_agent" CURSOR = "cursor" BEDROCK_MANTLE = "bedrock_mantle" + GDC = "gdc" # Create a set of all provider values for quick lookup diff --git a/litellm/utils.py b/litellm/utils.py index 26d3ae32739..226b94913b5 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -3497,6 +3497,17 @@ def filter_out_litellm_params(kwargs: dict) -> dict: return {key: value for key, value in kwargs.items() if key not in all_litellm_params} +def _provider_supports_vertex_params(custom_llm_provider: str) -> bool: + if custom_llm_provider in ("vertex_ai", "vertex_ai_beta"): + return True + try: + provider = LlmProviders(custom_llm_provider) + except ValueError: + return False + provider_config = ProviderConfigManager.get_provider_chat_config(model="", provider=provider) + return bool(getattr(provider_config, "supports_vertex_params", False)) + + class PreProcessNonDefaultParams: @staticmethod def base_pre_process_non_default_params( @@ -3518,11 +3529,7 @@ class PreProcessNonDefaultParams: continue elif k == "hf_model_name" and custom_llm_provider != "sagemaker": continue - elif ( - k.startswith("vertex_") - and custom_llm_provider != "vertex_ai" - and custom_llm_provider != "vertex_ai_beta" - ): # allow dynamically setting vertex ai init logic + elif k.startswith("vertex_") and not _provider_supports_vertex_params(custom_llm_provider): continue passed_params[k] = v @@ -7674,6 +7681,10 @@ class ProviderConfigManager: lambda: ProviderConfigManager._get_langflow_config(), False, ), + LlmProviders.GDC: ( + lambda: litellm.GDCGeminiConfig(), + False, + ), } @staticmethod diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index b137ec59a1f..edac8949f28 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -1059,6 +1059,16 @@ "interactions": true } }, + "gdc": { + "display_name": "Google Distributed Cloud (GDC)", + "url": "https://docs.litellm.ai/docs/providers/gdc", + "endpoints": { + "chat_completions": true, + "messages": false, + "responses": false, + "embeddings": false + } + }, "github_copilot": { "display_name": "GitHub Copilot (`github_copilot`)", "url": "https://docs.litellm.ai/docs/providers/github_copilot", diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index be62f8a9d67..f5aa600ab84 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -180,7 +180,7 @@ "limit": 34 }, "PLR1714": { - "limit": 267 + "limit": 265 }, "PLR1730": { "limit": 10 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..d106cf7ea21 --- /dev/null +++ b/tests/test_litellm/llms/gdc/chat/test_gdc_chat_transformation.py @@ -0,0 +1,717 @@ +import os +import sys +from unittest.mock import MagicMock, patch + +import pytest + +# Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../../../..")) + +import litellm +from litellm.llms.gdc.chat.transformation import GDCGeminiConfig + +TEST_API_KEY = '{"type": "gdch_service_account", "project_id": "test-project"}' +TEST_MODEL = "gdc/gemini-2.5-flash" +TEST_API_BASE = "https://gdc-endpoint.com" +TEST_PROJECT = "test-project" +TEST_LOCATION = "test-location" + + +class TestGDCGeminiConfig: + def test_get_complete_url(self): + config = GDCGeminiConfig() + url = config.get_complete_url( + api_base=TEST_API_BASE, + api_key=None, + model=TEST_MODEL, + optional_params={ + "vertex_project": TEST_PROJECT, + "vertex_location": TEST_LOCATION, + }, + litellm_params={}, + ) + assert ( + url + == f"{TEST_API_BASE}/v1/projects/{TEST_PROJECT}/locations/{TEST_LOCATION}/chat/completions" + ) + + def test_get_complete_url_adds_https_scheme(self): + config = GDCGeminiConfig() + url = config.get_complete_url( + api_base="gdc-endpoint.com", + api_key=None, + model=TEST_MODEL, + optional_params={}, + litellm_params={ + "vertex_project": TEST_PROJECT, + "vertex_location": TEST_LOCATION, + }, + ) + assert url.startswith("https://gdc-endpoint.com/v1/projects/") + + def test_get_complete_url_preformed_base_returned_as_is(self): + config = GDCGeminiConfig() + preformed = f"{TEST_API_BASE}/v1/projects/{TEST_PROJECT}/locations/{TEST_LOCATION}/chat/completions" + url = config.get_complete_url( + api_base=preformed, + api_key=None, + model=TEST_MODEL, + optional_params={"vertex_project": TEST_PROJECT}, + litellm_params={}, + ) + assert url == preformed + + 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, + "vertex_location": TEST_LOCATION, + }, + 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={}, + ) + + def test_get_complete_url_missing_location(self): + config = GDCGeminiConfig() + with pytest.raises(Exception, match="location is required for GDC Gemini"): + config.get_complete_url( + api_base=TEST_API_BASE, + api_key=None, + model=TEST_MODEL, + optional_params={"vertex_project": TEST_PROJECT}, + litellm_params={}, + ) + + def test_get_complete_url_accepts_vertex_ai_aliases(self): + config = GDCGeminiConfig() + url = config.get_complete_url( + api_base=TEST_API_BASE, + api_key=None, + model=TEST_MODEL, + optional_params={}, + litellm_params={ + "vertex_ai_project": TEST_PROJECT, + "vertex_ai_location": TEST_LOCATION, + }, + ) + assert ( + url + == f"{TEST_API_BASE}/v1/projects/{TEST_PROJECT}/locations/{TEST_LOCATION}/chat/completions" + ) + + def test_get_complete_url_preformed_base_is_authoritative_over_litellm_params(self): + config = GDCGeminiConfig() + preformed = f"{TEST_API_BASE}/v1/projects/pinned-project/locations/pinned-loc/chat/completions" + url = config.get_complete_url( + api_base=preformed, + api_key=None, + model=TEST_MODEL, + optional_params={"vertex_project": "attacker-optional", "vertex_location": "attacker-loc"}, + litellm_params={ + "vertex_project": "attacker-project", + "vertex_location": "attacker-loc", + }, + ) + assert url == preformed + + def test_get_complete_url_preformed_base_needs_no_project_param(self): + config = GDCGeminiConfig() + preformed = f"{TEST_API_BASE}/v1/projects/pinned-project/locations/pinned-loc/chat/completions" + url = config.get_complete_url( + api_base=preformed, + api_key=None, + model=TEST_MODEL, + optional_params={}, + litellm_params={}, + ) + assert url == preformed + + def test_deployment_project_takes_precedence_over_request(self): + config = GDCGeminiConfig() + url = config.get_complete_url( + api_base=TEST_API_BASE, + api_key=None, + model=TEST_MODEL, + optional_params={ + "vertex_project": "caller-project", + "vertex_location": "caller-location", + }, + litellm_params={ + "vertex_project": "deployment-project", + "vertex_location": "deployment-location", + }, + ) + assert url == ( + f"{TEST_API_BASE}/v1/projects/deployment-project" + "/locations/deployment-location/chat/completions" + ) + + @patch("google.auth.load_credentials_from_dict") + @patch("requests.Session") + def test_validate_environment(self, mock_session, mock_load_creds): + 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() + result = config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={ + "vertex_project": TEST_PROJECT, + "vertex_location": TEST_LOCATION, + }, + 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_validate_environment_strips_audience_trailing_slash(self): + config = GDCGeminiConfig() + mock_creds = MagicMock() + mock_creds.token = "mock-token" + mock_creds.with_gdch_audience.return_value = mock_creds + + with patch( + "google.auth.load_credentials_from_dict", return_value=(mock_creds, None) + ), patch("requests.Session"): + config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={}, + litellm_params={"vertex_project": TEST_PROJECT}, + api_key=TEST_API_KEY, + api_base="https://gdc-endpoint.com/", + ) + + mock_creds.with_gdch_audience.assert_called_once_with("https://gdc-endpoint.com") + + def test_validate_environment_audience_is_host_for_preformed_base(self): + config = GDCGeminiConfig() + mock_creds = MagicMock() + mock_creds.token = "mock-token" + mock_creds.with_gdch_audience.return_value = mock_creds + + with patch( + "google.auth.load_credentials_from_dict", return_value=(mock_creds, None) + ), patch("requests.Session"): + config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={}, + litellm_params={ + "vertex_project": "deployment-project", + "vertex_location": "deployment-loc", + }, + api_key=TEST_API_KEY, + api_base=f"{TEST_API_BASE}/v1/projects/embedded/locations/embedded/chat/completions", + ) + + mock_creds.with_gdch_audience.assert_called_once_with(TEST_API_BASE) + + def test_validate_environment_missing_api_base(self, monkeypatch): + monkeypatch.setattr(litellm, "api_base", None, raising=False) + monkeypatch.setattr(litellm, "gdc_api_base", None, raising=False) + config = GDCGeminiConfig() + with pytest.raises(Exception, match="api_base/host is required for GDC Gemini"): + config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={}, + litellm_params={"vertex_project": TEST_PROJECT}, + api_key=TEST_API_KEY, + api_base=None, + ) + + def test_validate_environment_missing_api_key(self): + config = GDCGeminiConfig() + with pytest.raises(Exception, match="api_key is required for GDC Gemini"): + config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={}, + litellm_params={"vertex_project": TEST_PROJECT}, + api_key=None, + api_base=TEST_API_BASE, + ) + + def test_validate_environment_missing_project(self): + config = GDCGeminiConfig() + with pytest.raises(Exception, match="project is required for GDC Gemini"): + config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={}, + litellm_params={}, + api_key=TEST_API_KEY, + api_base=TEST_API_BASE, + ) + + def test_validate_environment_raw_token_used_as_bearer(self): + config = GDCGeminiConfig() + headers = config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={}, + litellm_params={"vertex_project": TEST_PROJECT}, + api_key="ya29.raw-access-token", + api_base=TEST_API_BASE, + ) + assert headers["Authorization"] == "Bearer ya29.raw-access-token" + assert headers["x-goog-user-project"] == f"projects/{TEST_PROJECT}" + + def test_validate_environment_bad_credentials_raise_auth_error(self): + config = GDCGeminiConfig() + with patch( + "google.auth.load_credentials_from_dict", + side_effect=ValueError("bad creds"), + ): + with pytest.raises( + Exception, match="Failed to load service account credentials" + ): + config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={}, + litellm_params={"vertex_project": TEST_PROJECT}, + api_key=TEST_API_KEY, + api_base=TEST_API_BASE, + ) + + def test_validate_environment_string_false_disables_token_caching(self): + config = GDCGeminiConfig() + mock_creds = MagicMock() + mock_creds.token = "mock-token" + mock_creds.with_gdch_audience.return_value = mock_creds + + with patch( + "google.auth.load_credentials_from_dict", return_value=(mock_creds, None) + ), patch("requests.Session"), patch.object( + config, "_cached_fetch_token" + ) as mock_cached: + config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={}, + litellm_params={ + "vertex_project": TEST_PROJECT, + "gdc_token_caching": "false", + }, + api_key=TEST_API_KEY, + api_base=TEST_API_BASE, + ) + + mock_cached.assert_not_called() + + def test_validate_environment_token_caching_path(self): + config = GDCGeminiConfig() + mock_creds = MagicMock() + mock_creds.token = "cached-token" + mock_creds.valid = True + mock_creds.with_gdch_audience.return_value = mock_creds + + with patch( + "google.auth.load_credentials_from_dict", return_value=(mock_creds, None) + ): + headers = config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={}, + litellm_params={ + "vertex_project": TEST_PROJECT, + "gdc_token_caching": True, + }, + api_key=TEST_API_KEY, + api_base=TEST_API_BASE, + ) + + assert headers["Authorization"] == "Bearer cached-token" + mock_creds.refresh.assert_not_called() + + def test_validate_environment_preserves_content_type_but_rebinds_quota_project(self): + config = GDCGeminiConfig() + headers = config.validate_environment( + headers={ + "Content-Type": "text/plain", + "x-goog-user-project": "projects/attacker", + }, + model=TEST_MODEL, + messages=[], + optional_params={}, + litellm_params={"vertex_project": TEST_PROJECT}, + api_key="raw-token", + api_base=TEST_API_BASE, + ) + assert headers["Content-Type"] == "text/plain" + assert headers["x-goog-user-project"] == f"projects/{TEST_PROJECT}" + + @pytest.mark.parametrize( + "header_name", ["x-goog-user-project", "X-Goog-User-Project", "X-GOOG-USER-PROJECT"] + ) + def test_validate_environment_strips_caller_forwarded_quota_header(self, header_name): + config = GDCGeminiConfig() + mock_creds = MagicMock() + mock_creds.token = "tok" + mock_creds.with_gdch_audience.return_value = mock_creds + preformed = f"{TEST_API_BASE}/v1/projects/deployment-proj/locations/us-central1/chat/completions" + with patch( + "google.auth.load_credentials_from_dict", return_value=(mock_creds, None) + ): + headers = config.validate_environment( + headers={header_name: "projects/attacker"}, + model=TEST_MODEL, + messages=[], + optional_params={"vertex_project": "attacker-proj"}, + litellm_params={}, + api_key=TEST_API_KEY, + api_base=preformed, + ) + quota_values = [v for k, v in headers.items() if k.lower() == "x-goog-user-project"] + assert quota_values == ["projects/deployment-proj"] + + @pytest.mark.parametrize( + "bad", ["p/locations/l/chat/completions?", "a/b", "a?b", "a#b", "..", "a b", "a:b", "a%2Fb"] + ) + def test_get_complete_url_rejects_project_path_injection(self, bad): + config = GDCGeminiConfig() + with pytest.raises(Exception, match="vertex_project must be a plain identifier"): + config.get_complete_url( + api_base=TEST_API_BASE, + api_key=None, + model=TEST_MODEL, + optional_params={"vertex_project": bad, "vertex_location": TEST_LOCATION}, + litellm_params={}, + ) + + @pytest.mark.parametrize("bad", ["../../evil", "l/chat/completions", "l?x", ".."]) + def test_get_complete_url_rejects_location_path_injection(self, bad): + config = GDCGeminiConfig() + with pytest.raises(Exception, match="vertex_location must be a plain identifier"): + config.get_complete_url( + api_base=TEST_API_BASE, + api_key=None, + model=TEST_MODEL, + optional_params={"vertex_project": TEST_PROJECT, "vertex_location": bad}, + litellm_params={}, + ) + + @pytest.mark.parametrize("good", ["test-project", "us-central1", "123456", "proj_1", "MyProj-2"]) + def test_get_complete_url_accepts_valid_ids(self, good): + config = GDCGeminiConfig() + url = config.get_complete_url( + api_base=TEST_API_BASE, + api_key=None, + model=TEST_MODEL, + optional_params={"vertex_project": good, "vertex_location": good}, + litellm_params={}, + ) + assert url == f"{TEST_API_BASE}/v1/projects/{good}/locations/{good}/chat/completions" + + def test_validate_environment_rejects_project_path_injection(self): + config = GDCGeminiConfig() + with pytest.raises(Exception, match="vertex_project must be a plain identifier"): + config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={"vertex_project": "p/../admin"}, + litellm_params={}, + api_key="raw-token", + api_base=TEST_API_BASE, + ) + + def test_validate_environment_quota_header_bound_to_deployment_url(self): + config = GDCGeminiConfig() + mock_creds = MagicMock() + mock_creds.token = "tok" + mock_creds.with_gdch_audience.return_value = mock_creds + preformed = f"{TEST_API_BASE}/v1/projects/deployment-proj/locations/us-central1/chat/completions" + with patch( + "google.auth.load_credentials_from_dict", return_value=(mock_creds, None) + ): + headers = config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={"vertex_project": "attacker-proj"}, + litellm_params={}, + api_key=TEST_API_KEY, + api_base=preformed, + ) + assert headers["x-goog-user-project"] == "projects/deployment-proj" + + def test_validate_environment_quota_header_pinned_to_preformed_url(self): + config = GDCGeminiConfig() + mock_creds = MagicMock() + mock_creds.token = "tok" + mock_creds.with_gdch_audience.return_value = mock_creds + preformed = f"{TEST_API_BASE}/v1/projects/url-proj/locations/us-central1/chat/completions" + with patch( + "google.auth.load_credentials_from_dict", return_value=(mock_creds, None) + ): + headers = config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={"vertex_project": "attacker-proj"}, + litellm_params={"vertex_project": "override-proj"}, + api_key=TEST_API_KEY, + api_base=preformed, + ) + assert headers["x-goog-user-project"] == "projects/url-proj" + + def test_transform_request(self): + config = GDCGeminiConfig() + data = config.transform_request( + model=TEST_MODEL, + messages=[{"role": "user", "content": "Hello"}], + optional_params={ + "vertex_project": TEST_PROJECT, + "vertex_location": TEST_LOCATION, + }, + litellm_params={"ssl_verify": True}, + headers={}, + ) + assert data["model"] == "gemini-2.5-flash" + assert "vertex_project" not in data + assert "vertex_location" not in data + assert "ssl_verify" not in data + + def test_load_creds_from_key_ignores_file_paths(self, tmp_path): + config = GDCGeminiConfig() + creds_file = tmp_path / "service_account.json" + creds_file.write_text( + '{"type": "gdch_service_account", "project_id": "host-only-project"}' + ) + + creds, is_service_account = config._load_creds_from_key(str(creds_file)) + + assert creds is None + assert is_service_account is False + + def test_load_creds_from_key_rejects_non_gdch_credential_types(self): + config = GDCGeminiConfig() + external_account = ( + '{"type": "external_account", ' + '"token_url": "http://169.254.169.254/latest/api/token", ' + '"credential_source": {"url": "http://169.254.169.254/"}}' + ) + with patch( + "google.auth.load_credentials_from_dict", + return_value=(MagicMock(), None), + ) as mock_load: + with pytest.raises(ValueError, match="GDCH service account"): + config._load_creds_from_key(external_account) + mock_load.assert_not_called() + + def test_validate_environment_rejects_non_gdch_credential_without_refresh(self): + config = GDCGeminiConfig() + mock_creds = MagicMock() + mock_creds.token = "leaked-token" + mock_creds.with_gdch_audience.return_value = mock_creds + malicious = ( + '{"type": "external_account", ' + '"token_url": "http://169.254.169.254/latest/api/token"}' + ) + + with patch( + "google.auth.load_credentials_from_dict", + return_value=(mock_creds, None), + ) as mock_load, patch("requests.Session") as mock_session: + with pytest.raises( + Exception, match="Failed to load service account credentials" + ): + config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={}, + litellm_params={"vertex_project": TEST_PROJECT}, + api_key=malicious, + api_base=TEST_API_BASE, + ) + + mock_load.assert_not_called() + mock_session.assert_not_called() + mock_creds.refresh.assert_not_called() + + def test_validate_environment_does_not_read_api_key_file_path(self, tmp_path): + config = GDCGeminiConfig() + creds_file = tmp_path / "service_account.json" + creds_file.write_text( + '{"type": "service_account", "project_id": "host-only-project"}' + ) + + headers = config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params={}, + litellm_params={ + "vertex_project": TEST_PROJECT, + "vertex_location": TEST_LOCATION, + }, + api_key=str(creds_file), + api_base=TEST_API_BASE, + ) + + assert headers["Authorization"] == f"Bearer {creds_file}" + assert headers["x-goog-user-project"] == f"projects/{TEST_PROJECT}" + + @pytest.mark.parametrize( + "val, env_value, default, expected", + [ + (True, None, True, True), + (False, "true", True, False), + ("literal", None, True, "literal"), + (None, None, True, True), + (None, None, False, False), + (None, "true", False, True), + (None, "1", False, True), + (None, "on", False, True), + (None, "false", True, False), + (None, "0", True, False), + (None, "off", True, False), + (None, "verbose", True, "verbose"), + ], + ) + def test_read_env_bool(self, monkeypatch, val, env_value, default, expected): + config = GDCGeminiConfig() + env_var = "GDC_TEST_FLAG" + if env_value is None: + monkeypatch.delenv(env_var, raising=False) + else: + monkeypatch.setenv(env_var, env_value) + assert config._read_env_bool(val, env_var, default=default) == expected + + def test_cached_fetch_token_keys_by_credential(self): + config = GDCGeminiConfig() + + def make_creds(token): + creds = MagicMock() + creds.with_gdch_audience.return_value = creds + creds.valid = True + creds.token = token + return creds + + creds_a = make_creds("token-a") + creds_b = make_creds("token-b") + + assert ( + config._cached_fetch_token(creds_a, TEST_API_BASE, True, api_key="key-a") + == "token-a" + ) + assert ( + config._cached_fetch_token(creds_b, TEST_API_BASE, True, api_key="key-b") + == "token-b" + ) + # same credential identity reuses the cached entry + config._cached_fetch_token(creds_a, TEST_API_BASE, True, api_key="key-a") + creds_a.with_gdch_audience.assert_called_once() + + def test_cached_fetch_token_refreshes_when_invalid(self): + config = GDCGeminiConfig() + creds = MagicMock() + creds.with_gdch_audience.return_value = creds + creds.valid = False + creds.token = "refreshed" + + with patch.object(config, "_fetch_auth") as mock_fetch: + token = config._cached_fetch_token( + creds, TEST_API_BASE, True, api_key="key" + ) + + assert token == "refreshed" + mock_fetch.assert_called_once() + + def test_init_sets_up_lock_and_cache(self): + config = GDCGeminiConfig() + assert config._gdch_creds_cache == {} + assert config._creds_lock is not None + + +class TestCompleteGDC: + @patch("litellm.main.base_llm_http_handler.completion") + def test_complete_gdc_resolves_key_and_base(self, mock_completion, monkeypatch): + from litellm.main import gdc_transformation + + mock_completion.return_value = MagicMock() + monkeypatch.setattr(litellm, "gdc_key", "resolved-key", raising=False) + monkeypatch.setattr( + litellm, "gdc_api_base", "https://resolved-base.com", raising=False + ) + monkeypatch.setattr(litellm, "api_base", None, raising=False) + + litellm.completion( + model="gdc/gemini-2.5-flash", + messages=[{"role": "user", "content": "hi"}], + vertex_project=TEST_PROJECT, + vertex_location=TEST_LOCATION, + ) + + assert mock_completion.called + _, kwargs = mock_completion.call_args + assert kwargs["custom_llm_provider"] == "gdc" + assert kwargs["api_key"] == "resolved-key" + assert kwargs["api_base"] == "https://resolved-base.com" + assert kwargs["provider_config"] is gdc_transformation + + @patch("litellm.main.base_llm_http_handler.completion") + def test_complete_gdc_prefers_gdc_api_base_over_global( + self, mock_completion, monkeypatch + ): + mock_completion.return_value = MagicMock() + monkeypatch.setattr(litellm, "gdc_key", "resolved-key", raising=False) + monkeypatch.setattr( + litellm, "gdc_api_base", "https://gdc-specific.com", raising=False + ) + monkeypatch.setattr( + litellm, "api_base", "https://other-provider.com", raising=False + ) + + litellm.completion( + model="gdc/gemini-2.5-flash", + messages=[{"role": "user", "content": "hi"}], + vertex_project=TEST_PROJECT, + vertex_location=TEST_LOCATION, + ) + + _, kwargs = mock_completion.call_args + assert kwargs["api_base"] == "https://gdc-specific.com" diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 6f9f26bf6dd..28b82cf035b 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -1345,6 +1345,48 @@ def test_pre_process_non_default_params(model, custom_llm_provider): } +@pytest.mark.parametrize( + "custom_llm_provider, expected", + [ + ("vertex_ai", True), + ("vertex_ai_beta", True), + ("gdc", True), + ("openai", False), + ("bedrock", False), + ("not_a_real_provider", False), + ], +) +def test_provider_supports_vertex_params(custom_llm_provider, expected): + from litellm.utils import _provider_supports_vertex_params + + assert _provider_supports_vertex_params(custom_llm_provider) is expected + + +@pytest.mark.parametrize( + "model, custom_llm_provider, should_keep", + [ + ("gemini-2.5-pro", "vertex_ai", True), + ("gemini-2.5-pro", "vertex_ai_beta", True), + ("gdc/gemini-2.5-flash", "gdc", True), + ("gpt-4o", "openai", False), + ], +) +def test_vertex_params_not_stripped_for_vertex_family( + model, custom_llm_provider, should_keep +): + optional_params = litellm.utils.get_optional_params( + model=model, + custom_llm_provider=custom_llm_provider, + vertex_project="my-project", + vertex_location="us-central1", + ) + assert ("vertex_project" in optional_params) is should_keep + assert ("vertex_location" in optional_params) is should_keep + if should_keep: + assert optional_params["vertex_project"] == "my-project" + assert optional_params["vertex_location"] == "us-central1" + + from litellm.utils import supports_function_calling