Revert "Fix: Preserve full path structure for Gemini custom api_base (#12215)" (#12227)

This reverts commit f47254ecab.
This commit is contained in:
Ishaan Jaff 2025-07-01 20:41:39 -07:00 • committed by GitHub
parent 4b4e2dfde4
commit 7471a30dcd
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 8 additions and 151 deletions

View file

@ -83,11 +83,7 @@ class VertexBase:
if "type" in json_obj and json_obj["type"] == "external_account":
# If environment_id key contains "aws" value it corresponds to an AWS config file
credential_source = json_obj.get("credential_source", {})
environment_id = (
credential_source.get("environment_id", "")
if isinstance(credential_source, dict)
else ""
)
environment_id = credential_source.get("environment_id", "") if isinstance(credential_source, dict) else ""
if isinstance(environment_id, str) and "aws" in environment_id:
creds = self._credentials_from_identity_pool_with_aws(json_obj)
else:
@ -134,7 +130,7 @@ class VertexBase:
from google.auth import identity_pool
return identity_pool.Credentials.from_info(json_obj)
def _credentials_from_identity_pool_with_aws(self, json_obj):
from google.auth import aws
@ -301,21 +297,7 @@ class VertexBase:
"""
if api_base:
if custom_llm_provider == "gemini":
# Extract the path from the original URL to preserve the model structure
# Original URL format: https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash:generateContent?key=...
# We need to extract: /v1beta/models/gemini-2.5-flash:generateContent
from urllib.parse import urlparse
parsed_original = urlparse(url)
path_and_query = parsed_original.path
# Remove any query parameters from the path (they'll be re-added if needed)
if "?" in path_and_query:
path_and_query = path_and_query.split("?")[0]
# Construct the new URL with the custom base and the original path
url = "{}{}".format(api_base.rstrip("/"), path_and_query)
url = "{}:{}".format(api_base, endpoint)
if gemini_api_key is None:
raise ValueError(
"Missing gemini_api_key, please set `GEMINI_API_KEY`"
@ -324,7 +306,7 @@ class VertexBase:
gemini_api_key # cloudflare expects api key as bearer token
)
else:
url = "{}{}".format(api_base, endpoint)
url = "{}:{}".format(api_base, endpoint)
if stream is True:
url = url + "?alt=sse"
@ -359,7 +341,7 @@ class VertexBase:
stream=stream,
gemini_api_key=gemini_api_key,
)
auth_header = None # this field is not used for gemini
auth_header = None # this field is not used for gemin
else:
vertex_location = self.get_vertex_region(
vertex_region=vertex_location,
@ -508,7 +490,7 @@ class VertexBase:
headers.update(extra_headers)
return headers
@staticmethod
def get_vertex_ai_project(litellm_params: dict) -> Optional[str]:
return (
@ -517,7 +499,7 @@ class VertexBase:
or litellm.vertex_project
or get_secret_str("VERTEXAI_PROJECT")
)
@staticmethod
def get_vertex_ai_credentials(litellm_params: dict) -> Optional[str]:
return (
@ -525,7 +507,7 @@ class VertexBase:
or litellm_params.pop("vertex_ai_credentials", None)
or get_secret_str("VERTEXAI_CREDENTIALS")
)
@staticmethod
def get_vertex_ai_location(litellm_params: dict) -> Optional[str]:
return (

View file

@ -1,125 +0,0 @@
import os
import sys
from unittest.mock import MagicMock, patch
import pytest
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
import litellm
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
class TestCustomApiBaseFix:
"""Test that custom api_base works correctly for both Gemini and Vertex AI models"""
def test_gemini_custom_api_base_preserves_path(self):
"""Test that Gemini custom api_base preserves the full path structure"""
vertex_base = VertexBase()
# Test case from the issue: Cloudflare AI Gateway
cloudflare_base = "https://gateway.ai.cloudflare.com/v1/my-id/my-gateway/google-ai-studio"
original_url = "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash:generateContent?key=test-key"
token, final_url = vertex_base._check_custom_proxy(
api_base=cloudflare_base,
auth_header=None,
custom_llm_provider="gemini",
gemini_api_key="test-api-key",
endpoint=":generateContent",
stream=False,
url=original_url,
)
# Should preserve the full path structure
expected_url = "https://gateway.ai.cloudflare.com/v1/my-id/my-gateway/google-ai-studio/v1beta/models/gemini-2.5-flash:generateContent"
assert final_url == expected_url
assert token == "test-api-key"
def test_gemini_custom_api_base_with_streaming(self):
"""Test that streaming adds the alt=sse parameter correctly"""
vertex_base = VertexBase()
cloudflare_base = "https://gateway.ai.cloudflare.com/v1/my-id/my-gateway/google-ai-studio"
original_url = "https://generativelanguage.googleapis.com/v1beta/models/gemini-pro:streamGenerateContent?key=test-key&alt=sse"
token, final_url = vertex_base._check_custom_proxy(
api_base=cloudflare_base,
auth_header=None,
custom_llm_provider="gemini",
gemini_api_key="test-api-key",
endpoint=":streamGenerateContent",
stream=True,
url=original_url,
)
# Should preserve path and add streaming parameter
expected_url = "https://gateway.ai.cloudflare.com/v1/my-id/my-gateway/google-ai-studio/v1beta/models/gemini-pro:streamGenerateContent?alt=sse"
assert final_url == expected_url
assert token == "test-api-key"
def test_vertex_ai_custom_api_base_unchanged(self):
"""Test that Vertex AI behavior remains unchanged"""
vertex_base = VertexBase()
# Vertex AI uses a different URL structure
cloudflare_base = "https://gateway.ai.cloudflare.com/v1/my-id/my-gateway/google-vertex-ai/v1/projects/my-project/locations/us-central1/publishers/google/models/gemini-2.0-flash"
original_url = "https://us-central1-aiplatform.googleapis.com/v1/projects/my-project/locations/us-central1/publishers/google/models/gemini-2.0-flash:generateContent"
token, final_url = vertex_base._check_custom_proxy(
api_base=cloudflare_base,
auth_header="Bearer mock-token",
custom_llm_provider="vertex_ai",
gemini_api_key=None,
endpoint=":generateContent",
stream=False,
url=original_url,
)
# Vertex AI should still use the old behavior (appending endpoint)
expected_url = "https://gateway.ai.cloudflare.com/v1/my-id/my-gateway/google-vertex-ai/v1/projects/my-project/locations/us-central1/publishers/google/models/gemini-2.0-flash:generateContent"
assert final_url == expected_url
assert token == "Bearer mock-token" # Auth header unchanged for Vertex AI
def test_full_flow_with_get_token_and_url(self):
"""Test the full flow through _get_token_and_url"""
vertex_base = VertexBase()
cloudflare_base = "https://gateway.ai.cloudflare.com/v1/my-id/my-gateway/google-ai-studio"
token, url = vertex_base._get_token_and_url(
model="gemini/gemini-2.5-flash",
auth_header=None,
gemini_api_key="test-api-key",
vertex_project=None,
vertex_location=None,
vertex_credentials=None,
stream=False,
custom_llm_provider="gemini",
api_base=cloudflare_base,
)
# Should produce the correct URL with custom base
assert cloudflare_base in url
assert "/v1beta/models/gemini/gemini-2.5-flash:generateContent" in url
assert token == "test-api-key"
def test_missing_gemini_api_key_raises_error(self):
"""Test that missing Gemini API key raises an error"""
vertex_base = VertexBase()
cloudflare_base = "https://gateway.ai.cloudflare.com/v1/my-id/my-gateway/google-ai-studio"
original_url = "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash:generateContent"
with pytest.raises(ValueError, match="Missing gemini_api_key"):
vertex_base._check_custom_proxy(
api_base=cloudflare_base,
auth_header=None,
custom_llm_provider="gemini",
gemini_api_key=None, # Missing API key
endpoint=":generateContent",
stream=False,
url=original_url,
)