Merge pull request #18410 from BerriAI/litellm_refactor_location_code

feat(vertex-ai): add centralized get_vertex_base_url() helper for global location support
This commit is contained in:
Sameer Kankute 2025-12-24 21:42:18 +05:30 • committed by GitHub
commit fd7c928204
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 518 additions and 62 deletions

View file

@ -23,6 +23,7 @@ from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMExcepti
from litellm.llms.vertex_ai.agent_engine.sse_iterator import (
VertexAgentEngineResponseIterator,
)
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import Choices, Message, ModelResponse, Usage
@ -130,8 +131,7 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase):
)
resource_path = f"projects/{vertex_project}/locations/{vertex_location}/reasoningEngines/{engine_id}"
# Build the base URL
base_url = f"https://{vertex_location}-aiplatform.googleapis.com"
base_url = get_vertex_base_url(vertex_location)
# Always use :streamQuery endpoint for actual queries
# The :query endpoint only supports session management methods

View file

@ -8,6 +8,7 @@ from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
get_async_httpx_client,
)
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM
from litellm.types.llms.openai import CreateBatchRequest
from litellm.types.llms.vertex_ai import (
@ -128,7 +129,8 @@ class VertexAIBatchPrediction(VertexLLM):
) -> str:
"""Return the base url for the vertex garden models"""
# POST https://LOCATION-aiplatform.googleapis.com/v1/projects/PROJECT_ID/locations/LOCATION/batchPredictionJobs
return f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/batchPredictionJobs"
base_url = get_vertex_base_url(vertex_location)
return f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/batchPredictionJobs"
def retrieve_batch(
self,

View file

@ -193,6 +193,18 @@ def get_vertex_base_model_name(model: str) -> str:
return model
def get_vertex_base_url(
vertex_location: Optional[str],
) -> str:
"""
Get the base URL for Vertex AI API calls.
"""
if vertex_location == "global":
return "https://aiplatform.googleapis.com"
else:
return f"https://{vertex_location}-aiplatform.googleapis.com"
def _get_embedding_url(
model: str,
vertex_project: Optional[str],
@ -212,10 +224,18 @@ def _get_embedding_url(
# Strip routing prefixes (bge/, gemma/, etc.) for endpoint URL construction
model = get_vertex_base_model_name(model=model)
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}"
# Get base URL (handles global vs regional)
base_url = get_vertex_base_url(vertex_location)
if model.isdigit():
# https://us-central1-aiplatform.googleapis.com/v1/projects/$PROJECT_ID/locations/us-central1/endpoints/$ENDPOINT_ID:predict
url = f"https://{vertex_location}-aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}:{endpoint}"
# https://aiplatform.googleapis.com/v1/projects/$PROJECT_ID/locations/global/endpoints/$ENDPOINT_ID:predict
url = f"{base_url}/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}:{endpoint}"
else:
# Regular model -> publisher model
# https://us-central1-aiplatform.googleapis.com/v1/projects/$PROJECT_ID/locations/us-central1/publishers/google/models/{model}:predict
# https://aiplatform.googleapis.com/v1/projects/$PROJECT_ID/locations/global/publishers/google/models/{model}:predict
url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}"
return url, endpoint
@ -236,26 +256,23 @@ def _get_vertex_url(
if mode == "chat":
### SET RUNTIME ENDPOINT ###
endpoint = "generateContent"
base_url = get_vertex_base_url(vertex_location)
if stream is True:
endpoint = "streamGenerateContent"
if vertex_location == "global":
url = f"https://aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/global/publishers/google/models/{model}:{endpoint}?alt=sse"
else:
url = f"https://{vertex_location}-aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}?alt=sse"
else:
if vertex_location == "global":
url = f"https://aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/global/publishers/google/models/{model}:{endpoint}"
else:
url = f"https://{vertex_location}-aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}"
# if model is only numeric chars then it's a fine tuned gemini model
# model = 4965075652664360960
# send to this url: url = f"https://{vertex_location}-aiplatform.googleapis.com/{version}/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}:{endpoint}"
# send to this url: url = f"{base_url}/{version}/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}:{endpoint}"
if model.isdigit():
# It's a fine-tuned Gemini model
url = f"https://{vertex_location}-aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}:{endpoint}"
if stream is True:
url += "?alt=sse"
# It's a fine-tuned Gemini model - use endpoints/ path
url = f"{base_url}/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}:{endpoint}"
else:
# Regular model - use publishers/google/models/ path
url = f"{base_url}/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}"
if stream is True:
url += "?alt=sse"
elif mode == "embedding":
return _get_embedding_url(
model=model,
@ -265,15 +282,17 @@ def _get_vertex_url(
)
elif mode == "image_generation":
endpoint = "predict"
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}"
base_url = get_vertex_base_url(vertex_location)
if model.isdigit():
url = f"https://{vertex_location}-aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}:{endpoint}"
# Numeric model -> custom endpoint
url = f"{base_url}/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}:{endpoint}"
else:
# Regular model -> publisher model
url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}"
elif mode == "count_tokens":
endpoint = "countTokens"
if vertex_location == "global":
url = f"https://aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/global/publishers/google/models/{model}:{endpoint}"
else:
url = f"https://{vertex_location}-aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}"
base_url = get_vertex_base_url(vertex_location)
url = f"{base_url}/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}"
if not url or not endpoint:
raise ValueError(f"Unable to get vertex url/endpoint for mode: {mode}")
return url, endpoint

View file

@ -8,6 +8,7 @@ import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.llms.custom_httpx.http_handler import HTTPHandler, get_async_httpx_client
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM
from litellm.types.fine_tuning import OpenAIFineTuningHyperparameters
from litellm.types.llms.openai import FineTuningJobCreate
@ -261,7 +262,8 @@ class VertexFineTuningAPI(VertexLLM):
original_hyperparameters=original_hyperparameters or {},
)
fine_tuning_url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/tuningJobs"
base_url = get_vertex_base_url(vertex_location)
fine_tuning_url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/tuningJobs"
if _is_async is True:
return self.acreate_fine_tuning_job( # type: ignore
fine_tuning_url=fine_tuning_url,
@ -329,19 +331,21 @@ class VertexFineTuningAPI(VertexLLM):
"Content-Type": "application/json",
}
base_url = get_vertex_base_url(vertex_location)
url = None
if request_route == "/tuningJobs":
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/tuningJobs"
url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/tuningJobs"
elif "/tuningJobs/" in request_route and "cancel" in request_route:
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/tuningJobs{request_route}"
url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/tuningJobs{request_route}"
elif "generateContent" in request_route:
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}{request_route}"
url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}{request_route}"
elif "predict" in request_route:
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}{request_route}"
url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}{request_route}"
elif "/batchPredictionJobs" in request_route:
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}{request_route}"
url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}{request_route}"
elif "countTokens" in request_route:
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}{request_route}"
url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}{request_route}"
elif "cachedContents" in request_route:
_model = request_data.get("model")
if _model is not None and "/publishers/google/models/" not in _model:
@ -349,7 +353,7 @@ class VertexFineTuningAPI(VertexLLM):
f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{_model}"
)
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1beta1/projects/{vertex_project}/locations/{vertex_location}{request_route}"
url = f"{base_url}/v1beta1/projects/{vertex_project}/locations/{vertex_location}{request_route}"
else:
raise ValueError(f"Unsupported Vertex AI request route: {request_route}")
if self.async_handler is None:

View file

@ -10,6 +10,7 @@ from httpx._types import RequestFiles
import litellm
from litellm.images.utils import ImageEditRequestUtils
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM
from litellm.secret_managers.main import get_secret_str
from litellm.types.images.main import ImageEditOptionalRequestParams
@ -143,11 +144,7 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM):
if not vertex_project or not vertex_location:
raise ValueError("vertex_project and vertex_location are required for Vertex AI")
# Handle global location differently (no region prefix in URL)
if vertex_location == "global":
base_url = "https://aiplatform.googleapis.com"
else:
base_url = f"https://{vertex_location}-aiplatform.googleapis.com"
base_url = get_vertex_base_url(vertex_location)
return f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model_name}:generateContent"

View file

@ -9,9 +9,9 @@ import httpx
from httpx._types import RequestFiles
import litellm
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM
from litellm.secret_managers.main import get_secret_str
from litellm.types.images.main import ImageEditOptionalRequestParams
@ -136,7 +136,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM):
if api_base:
base_url = api_base.rstrip("/")
else:
base_url = f"https://{vertex_location}-aiplatform.googleapis.com"
base_url = get_vertex_base_url(vertex_location)
return f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model_name}:predict"

View file

@ -7,13 +7,19 @@ import litellm
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
AllMessageValues,
OpenAIImageGenerationOptionalParams,
)
from litellm.types.utils import ImageObject, ImageResponse, ImageUsage, ImageUsageInputTokensDetails
from litellm.types.utils import (
ImageObject,
ImageResponse,
ImageUsage,
ImageUsageInputTokensDetails,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
@ -140,11 +146,7 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
if not vertex_project or not vertex_location:
raise ValueError("vertex_project and vertex_location are required for Vertex AI")
# Handle global location differently (no region prefix in URL)
if vertex_location == "global":
base_url = "https://aiplatform.googleapis.com"
else:
base_url = f"https://{vertex_location}-aiplatform.googleapis.com"
base_url = get_vertex_base_url(vertex_location)
return f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model_name}:generateContent"

View file

@ -7,6 +7,7 @@ import litellm
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
@ -140,7 +141,7 @@ class VertexAIImagenImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
if not vertex_project or not vertex_location:
raise ValueError("vertex_project and vertex_location are required for Vertex AI")
base_url = f"https://{vertex_location}-aiplatform.googleapis.com"
base_url = get_vertex_base_url(vertex_location)
return f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model_name}:predict"

View file

@ -10,6 +10,7 @@ from litellm.litellm_core_utils.prompt_templates.image_handling import (
)
from litellm.llms.base_llm.ocr.transformation import DocumentType, OCRRequestData
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
@ -104,7 +105,7 @@ class VertexAIOCRConfig(MistralOCRConfig):
# Get API base URL
if api_base is None:
api_base = f"https://{vertex_location}-aiplatform.googleapis.com"
api_base = get_vertex_base_url(vertex_location)
# Ensure no trailing slash
api_base = api_base.rstrip("/")

View file

@ -8,6 +8,7 @@ from typing import Any, Dict, Optional
from litellm._logging import verbose_logger
from litellm.constants import DEFAULT_CHUNK_OVERLAP, DEFAULT_CHUNK_SIZE
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.types.rag import RAGChunkingStrategy
@ -37,8 +38,8 @@ class VertexAIRAGTransformation(VertexBase):
Note: The REST endpoint for importRagFiles may not be publicly available.
Vertex AI RAG Engine primarily uses gRPC-based SDK.
"""
base_url = f"https://{vertex_location}-aiplatform.googleapis.com/v1"
return f"{base_url}/projects/{vertex_project}/locations/{vertex_location}/ragCorpora/{corpus_id}:importRagFiles"
base_url = get_vertex_base_url(vertex_location)
return f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/ragCorpora/{corpus_id}:importRagFiles"
def get_retrieve_contexts_url(
self,
@ -46,8 +47,8 @@ class VertexAIRAGTransformation(VertexBase):
vertex_location: str,
) -> str:
"""Get the URL for retrieving contexts (search)."""
base_url = f"https://{vertex_location}-aiplatform.googleapis.com/v1"
return f"{base_url}/projects/{vertex_project}/locations/{vertex_location}:retrieveContexts"
base_url = get_vertex_base_url(vertex_location)
return f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}:retrieveContexts"
def transform_chunking_strategy_to_vertex_format(
self,

View file

@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
import httpx
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.types.router import GenericLiteLLMParams
from litellm.types.vector_stores import (
@ -88,7 +89,8 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
return api_base.rstrip("/")
# Vertex AI RAG API endpoint for retrieveContexts
return f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}"
base_url = get_vertex_base_url(vertex_location)
return f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}"
def transform_search_vector_store_request(
self,

View file

@ -8,6 +8,7 @@ their respective publisher-specific count-tokens endpoints.
from typing import Any, Dict, Optional
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
@ -65,10 +66,8 @@ class VertexAIPartnerModelsTokenCounter(VertexBase):
# Use custom api_base if provided, otherwise construct default
if api_base:
base_url = api_base
elif vertex_location == "global":
base_url = "https://aiplatform.googleapis.com"
else:
base_url = f"https://{vertex_location}-aiplatform.googleapis.com"
base_url = get_vertex_base_url(vertex_location)
# Construct the count-tokens endpoint
# Format: /v1/projects/{project}/locations/{location}/publishers/{publisher}/models/count-tokens:rawPredict

View file

@ -20,6 +20,7 @@ from typing import Callable, Optional, Union
import httpx # type: ignore
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.utils import ModelResponse
from ..common_utils import VertexAIError, get_vertex_base_model_name
@ -34,8 +35,8 @@ def create_vertex_url(
api_base: Optional[str] = None,
) -> str:
"""Return the base url for the vertex garden models"""
# f"https://{self.endpoint.location}-aiplatform.googleapis.com/v1beta1/projects/{PROJECT_ID}/locations/{self.endpoint.location}"
return f"https://{vertex_location}-aiplatform.googleapis.com/v1beta1/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}"
base_url = get_vertex_base_url(vertex_location)
return f"{base_url}/v1beta1/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}"
class VertexAIModelGardenModels(VertexBase):

View file

@ -17,6 +17,7 @@ from litellm.images.utils import ImageEditRequestUtils
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
from litellm.llms.vertex_ai.common_utils import (
_convert_vertex_datetime_to_openai_datetime,
get_vertex_base_url,
)
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.types.router import GenericLiteLLMParams
@ -222,10 +223,8 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
# Construct the URL
if api_base:
base_url = api_base.rstrip("/")
elif vertex_location == "global":
base_url = "https://aiplatform.googleapis.com"
else:
base_url = f"https://{vertex_location}-aiplatform.googleapis.com"
base_url = get_vertex_base_url(vertex_location)
url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model_name}"

View file

@ -0,0 +1,428 @@
"""
Comprehensive tests for Vertex AI global URL support across all endpoints.
This test suite ensures that all Vertex AI endpoints properly handle the 'global' location,
which uses a different URL format than regional endpoints.
Regional: https://{region}-aiplatform.googleapis.com/...
Global: https://aiplatform.googleapis.com/...
"""
from unittest.mock import patch
import pytest
from litellm.llms.vertex_ai.common_utils import (
_get_embedding_url,
_get_vertex_url,
get_vertex_base_url,
)
class TestVertexBaseURL:
"""Test the centralized get_vertex_base_url helper function."""
@pytest.mark.parametrize(
"vertex_location, expected_base_url",
[
("us-central1", "https://us-central1-aiplatform.googleapis.com"),
("us-east1", "https://us-east1-aiplatform.googleapis.com"),
("europe-west1", "https://europe-west1-aiplatform.googleapis.com"),
("asia-northeast1", "https://asia-northeast1-aiplatform.googleapis.com"),
("global", "https://aiplatform.googleapis.com"),
],
)
def test_get_vertex_base_url(self, vertex_location, expected_base_url):
"""Test that get_vertex_base_url returns correct URL for all location types."""
result = get_vertex_base_url(vertex_location)
assert result == expected_base_url
assert not result.endswith("/") # No trailing slash
class TestChatCompletionURLs:
"""Test chat/completion endpoint URL construction with global location."""
@pytest.mark.parametrize(
"vertex_location, stream, expected_url_pattern",
[
# Regional, non-streaming
(
"us-central1",
False,
"https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-1.5-pro:generateContent",
),
# Regional, streaming
(
"us-central1",
True,
"https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-1.5-pro:streamGenerateContent?alt=sse",
),
# Global, non-streaming
(
"global",
False,
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/google/models/gemini-1.5-pro:generateContent",
),
# Global, streaming
(
"global",
True,
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/google/models/gemini-1.5-pro:streamGenerateContent?alt=sse",
),
],
)
def test_chat_url_construction(
self, vertex_location, stream, expected_url_pattern
):
"""Test that chat URLs are correctly constructed for regional and global locations."""
with patch(
"litellm.VertexGeminiConfig.get_model_for_vertex_ai_url",
side_effect=lambda model: model,
):
url, endpoint = _get_vertex_url(
mode="chat",
model="gemini-1.5-pro",
stream=stream,
vertex_project="test-project",
vertex_location=vertex_location,
vertex_api_version="v1",
)
assert url == expected_url_pattern
if stream:
assert endpoint == "streamGenerateContent"
assert "?alt=sse" in url
else:
assert endpoint == "generateContent"
assert "?alt=sse" not in url
@pytest.mark.parametrize(
"vertex_location, stream",
[
("us-central1", False),
("us-central1", True),
("global", False),
("global", True),
],
)
def test_finetuned_model_url_construction(self, vertex_location, stream):
"""Test that fine-tuned models (numeric IDs) use endpoints/ path correctly."""
with patch(
"litellm.VertexGeminiConfig.get_model_for_vertex_ai_url",
side_effect=lambda model: model,
):
url, endpoint = _get_vertex_url(
mode="chat",
model="1234567890", # Numeric model ID
stream=stream,
vertex_project="test-project",
vertex_location=vertex_location,
vertex_api_version="v1",
)
# Should use endpoints/ path instead of publishers/google/models/
assert "/endpoints/1234567890:" in url
assert "/publishers/google/models/" not in url
# Check base URL is correct
if vertex_location == "global":
assert url.startswith("https://aiplatform.googleapis.com")
else:
assert url.startswith(f"https://{vertex_location}-aiplatform.googleapis.com")
class TestEmbeddingURLs:
"""Test embedding endpoint URL construction with global location."""
@pytest.mark.parametrize(
"vertex_location, model, expected_url_pattern",
[
# Regional, regular model
(
"us-central1",
"text-embedding-004",
"https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/text-embedding-004:predict",
),
# Global, regular model
(
"global",
"text-embedding-004",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/google/models/text-embedding-004:predict",
),
# Regional, numeric endpoint
(
"us-central1",
"1234567890",
"https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/endpoints/1234567890:predict",
),
# Global, numeric endpoint
(
"global",
"1234567890",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/endpoints/1234567890:predict",
),
],
)
def test_embedding_url_construction(
self, vertex_location, model, expected_url_pattern
):
"""Test that embedding URLs are correctly constructed for regional and global locations."""
url, endpoint = _get_embedding_url(
model=model,
vertex_project="test-project",
vertex_location=vertex_location,
vertex_api_version="v1",
)
assert url == expected_url_pattern
assert endpoint == "predict"
# Verify base URL format
if vertex_location == "global":
assert url.startswith("https://aiplatform.googleapis.com")
assert "-aiplatform.googleapis.com" not in url
else:
assert url.startswith(f"https://{vertex_location}-aiplatform.googleapis.com")
@pytest.mark.parametrize(
"vertex_location",
["us-central1", "europe-west1", "global"],
)
def test_embedding_url_with_routing_prefix(self, vertex_location):
"""Test that routing prefixes (bge/, gemma/, etc.) are stripped from URLs."""
url, endpoint = _get_embedding_url(
model="bge/1234567890", # Model with routing prefix
vertex_project="test-project",
vertex_location=vertex_location,
vertex_api_version="v1",
)
# Routing prefix should be stripped
assert "bge/" not in url
assert "/endpoints/1234567890:" in url
class TestCountTokensURLs:
"""Test count_tokens endpoint URL construction with global location."""
@pytest.mark.parametrize(
"vertex_location, expected_url_pattern",
[
(
"us-central1",
"https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-1.5-pro:countTokens",
),
(
"global",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/google/models/gemini-1.5-pro:countTokens",
),
],
)
def test_count_tokens_url_construction(self, vertex_location, expected_url_pattern):
"""Test that count_tokens URLs are correctly constructed for regional and global locations."""
with patch(
"litellm.VertexGeminiConfig.get_model_for_vertex_ai_url",
side_effect=lambda model: model,
):
url, endpoint = _get_vertex_url(
mode="count_tokens",
model="gemini-1.5-pro",
stream=None,
vertex_project="test-project",
vertex_location=vertex_location,
vertex_api_version="v1",
)
assert url == expected_url_pattern
assert endpoint == "countTokens"
class TestImageGenerationURLs:
"""Test image_generation endpoint URL construction with global location."""
@pytest.mark.parametrize(
"vertex_location, model, expected_url_pattern",
[
# Regional, regular model
(
"us-central1",
"imagen-3.0-generate-001",
"https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/imagen-3.0-generate-001:predict",
),
# Global, regular model
(
"global",
"imagen-3.0-generate-001",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/google/models/imagen-3.0-generate-001:predict",
),
# Regional, numeric endpoint
(
"us-central1",
"9876543210",
"https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/endpoints/9876543210:predict",
),
# Global, numeric endpoint
(
"global",
"9876543210",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/endpoints/9876543210:predict",
),
],
)
def test_image_generation_url_construction(
self, vertex_location, model, expected_url_pattern
):
"""Test that image_generation URLs are correctly constructed for regional and global locations."""
with patch(
"litellm.VertexGeminiConfig.get_model_for_vertex_ai_url",
side_effect=lambda model: model,
):
url, endpoint = _get_vertex_url(
mode="image_generation",
model=model,
stream=None,
vertex_project="test-project",
vertex_location=vertex_location,
vertex_api_version="v1",
)
assert url == expected_url_pattern
assert endpoint == "predict"
class TestAPIVersions:
"""Test that both v1 and v1beta1 API versions work with global location."""
@pytest.mark.parametrize(
"api_version, vertex_location",
[
("v1", "us-central1"),
("v1", "global"),
("v1beta1", "us-central1"),
("v1beta1", "global"),
],
)
def test_api_versions_in_urls(self, api_version, vertex_location):
"""Test that API version is correctly included in URLs for all locations."""
with patch(
"litellm.VertexGeminiConfig.get_model_for_vertex_ai_url",
side_effect=lambda model: model,
):
url, _ = _get_vertex_url(
mode="chat",
model="gemini-1.5-pro",
stream=False,
vertex_project="test-project",
vertex_location=vertex_location,
vertex_api_version=api_version,
)
# API version should be in the URL
assert f"/{api_version}/" in url
class TestEdgeCases:
"""Test edge cases and special scenarios."""
def test_global_location_no_region_prefix(self):
"""Ensure global URLs never have a region prefix."""
base_url = get_vertex_base_url("global")
assert base_url == "https://aiplatform.googleapis.com"
assert "global-aiplatform" not in base_url
assert "-aiplatform.googleapis.com" not in base_url
@pytest.mark.parametrize(
"mode",
["chat", "embedding", "count_tokens", "image_generation"],
)
def test_all_modes_support_global(self, mode):
"""Test that all URL modes support global location."""
with patch(
"litellm.VertexGeminiConfig.get_model_for_vertex_ai_url",
side_effect=lambda model: model,
):
if mode == "embedding":
url, _ = _get_embedding_url(
model="text-embedding-004",
vertex_project="test-project",
vertex_location="global",
vertex_api_version="v1",
)
else:
url, _ = _get_vertex_url(
mode=mode,
model="gemini-1.5-pro",
stream=False,
vertex_project="test-project",
vertex_location="global",
vertex_api_version="v1",
)
# All URLs should use global format
assert url.startswith("https://aiplatform.googleapis.com")
assert "/locations/global/" in url
def test_location_in_path_matches_parameter(self):
"""Ensure the location in the URL path matches the vertex_location parameter."""
test_locations = ["us-central1", "europe-west1", "global"]
for location in test_locations:
with patch(
"litellm.VertexGeminiConfig.get_model_for_vertex_ai_url",
side_effect=lambda model: model,
):
url, _ = _get_vertex_url(
mode="chat",
model="gemini-1.5-pro",
stream=False,
vertex_project="test-project",
vertex_location=location,
vertex_api_version="v1",
)
# Location should appear in the path
assert f"/locations/{location}/" in url
class TestBackwardCompatibility:
"""Ensure changes don't break existing functionality."""
def test_regional_urls_unchanged(self):
"""Test that regional URL construction hasn't changed."""
with patch(
"litellm.VertexGeminiConfig.get_model_for_vertex_ai_url",
side_effect=lambda model: model,
):
url, _ = _get_vertex_url(
mode="chat",
model="gemini-1.5-pro",
stream=False,
vertex_project="my-project",
vertex_location="us-central1",
vertex_api_version="v1",
)
# Should match the traditional regional format
assert (
url
== "https://us-central1-aiplatform.googleapis.com/v1/projects/my-project/locations/us-central1/publishers/google/models/gemini-1.5-pro:generateContent"
)
def test_streaming_urls_unchanged(self):
"""Test that streaming URL construction hasn't changed."""
with patch(
"litellm.VertexGeminiConfig.get_model_for_vertex_ai_url",
side_effect=lambda model: model,
):
url, _ = _get_vertex_url(
mode="chat",
model="gemini-1.5-pro",
stream=True,
vertex_project="my-project",
vertex_location="us-central1",
vertex_api_version="v1",
)
# Should include streaming endpoint and alt=sse
assert ":streamGenerateContent?alt=sse" in url