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 <lostpaige@google.com>
Co-authored-by: Tim Laubach <tlaubach@google.com>
This commit is contained in:
Mateo Wang 2026-07-01 17:31:07 -07:00 • committed by GitHub
parent 0b0fd6a4d1
commit fde4c7c97a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
14 changed files with 1126 additions and 6 deletions

View file

@ -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

View file

@ -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",

View file

@ -460,6 +460,7 @@ LITELLM_CHAT_PROVIDERS = [
"openai",
"openai_like",
"bytez",
"gdc",
"xai",
"custom_openai",
"text-completion-openai",

View file

@ -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/"):

View file

View file

View file

@ -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

View file

@ -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":

View file

@ -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

View file

@ -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

View file

@ -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",

View file

@ -180,7 +180,7 @@
"limit": 34
},
"PLR1714": {
"limit": 267
"limit": 265
},
"PLR1730": {
"limit": 10

View file

@ -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"

View file

@ -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