Fix: skip auth for custom api base in vertex ai

This commit is contained in:
Sameer Kankute 2026-01-20 12:36:25 +05:30
parent ea2e360cb5
commit 3f8e985b58
7 changed files with 258 additions and 34 deletions

View file

@ -2329,15 +2329,17 @@ class VertexLLM(VertexBase):
optional_params=optional_params
)
# Extract use_psc_endpoint_format from optional_params
use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False)
_auth_header, vertex_project = await self._ensure_access_token_async(
credentials=vertex_credentials,
project_id=vertex_project,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
use_psc_endpoint_format=use_psc_endpoint_format,
)
# Extract use_psc_endpoint_format from optional_params
use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False)
auth_header, api_base = self._get_token_and_url(
model=model,
gemini_api_key=gemini_api_key,
@ -2427,15 +2429,17 @@ class VertexLLM(VertexBase):
optional_params=optional_params
)
# Extract use_psc_endpoint_format from optional_params
use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False)
_auth_header, vertex_project = await self._ensure_access_token_async(
credentials=vertex_credentials,
project_id=vertex_project,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
use_psc_endpoint_format=use_psc_endpoint_format,
)
# Extract use_psc_endpoint_format from optional_params
use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False)
auth_header, api_base = self._get_token_and_url(
model=model,
gemini_api_key=gemini_api_key,
@ -2615,15 +2619,17 @@ class VertexLLM(VertexBase):
optional_params=optional_params
)
# Extract use_psc_endpoint_format from optional_params
use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False)
_auth_header, vertex_project = self._ensure_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
use_psc_endpoint_format=use_psc_endpoint_format,
)
# Extract use_psc_endpoint_format from optional_params
use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False)
auth_header, url = self._get_token_and_url(
model=model,
gemini_api_key=gemini_api_key,

View file

@ -135,19 +135,24 @@ class VertexAIPartnerModels(VertexBase):
try:
vertex_httpx_logic = VertexLLM()
## CONSTRUCT API BASE
stream: bool = optional_params.get("stream", False) or False
# Extract use_psc_endpoint_format from optional_params
use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False)
access_token, project_id = vertex_httpx_logic._ensure_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
custom_llm_provider="vertex_ai",
api_base=api_base,
use_psc_endpoint_format=use_psc_endpoint_format,
)
openai_like_chat_completions = OpenAILikeChatHandler()
codestral_fim_completions = CodestralTextCompletion()
anthropic_chat_completions = AnthropicChatCompletion()
## CONSTRUCT API BASE
stream: bool = optional_params.get("stream", False) or False
optional_params["stream"] = stream
if self.should_use_openai_handler(model):

View file

@ -67,13 +67,16 @@ class VertexEmbedding(VertexBase):
optional_params=optional_params
)
# Extract use_psc_endpoint_format from optional_params
use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False)
_auth_header, vertex_project = self._ensure_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
use_psc_endpoint_format=use_psc_endpoint_format,
)
# Extract use_psc_endpoint_format from optional_params
use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False)
auth_header, api_base = self._get_token_and_url(
model=model,
@ -163,13 +166,16 @@ class VertexEmbedding(VertexBase):
should_use_v1beta1_features = self.is_using_v1beta1_features(
optional_params=optional_params
)
# Extract use_psc_endpoint_format from optional_params
use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False)
_auth_header, vertex_project = await self._ensure_access_token_async(
credentials=vertex_credentials,
project_id=vertex_project,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
use_psc_endpoint_format=use_psc_endpoint_format,
)
# Extract use_psc_endpoint_format from optional_params
use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False)
auth_header, api_base = self._get_token_and_url(
model=model,

View file

@ -86,16 +86,21 @@ class VertexAIGemmaModels(VertexBase):
model = get_vertex_base_model_name(model=model)
vertex_httpx_logic = VertexLLM()
## CONSTRUCT API BASE
stream: bool = optional_params.get("stream", False) or False
# Extract use_psc_endpoint_format from optional_params
use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False)
access_token, project_id = vertex_httpx_logic._ensure_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
custom_llm_provider="vertex_ai",
api_base=api_base,
use_psc_endpoint_format=use_psc_endpoint_format,
)
gemma_transformation = VertexGemmaConfig()
## CONSTRUCT API BASE
stream: bool = optional_params.get("stream", False) or False
optional_params["stream"] = stream
# If api_base is not provided, it should be set as an environment variable

View file

@ -274,17 +274,44 @@ class VertexBase:
custom_llm_provider: Literal[
"vertex_ai", "vertex_ai_beta", "gemini"
], # if it's vertex_ai or gemini (google ai studio)
api_base: Optional[str] = None,
use_psc_endpoint_format: bool = False,
) -> Tuple[str, str]:
"""
Returns auth token and project id
Args:
credentials: Vertex AI credentials
project_id: Google Cloud project ID
custom_llm_provider: Provider type (vertex_ai, vertex_ai_beta, or gemini)
api_base: Custom API base URL (e.g., for proxies)
use_psc_endpoint_format: Whether using PSC endpoint format
Returns:
Tuple of (access_token, project_id)
Note:
Authentication is skipped when using a custom api_base that is not a PSC endpoint.
PSC endpoints still require Google authentication even with custom api_base.
"""
if custom_llm_provider == "gemini":
return "", ""
else:
return self.get_access_token(
credentials=credentials,
project_id=project_id,
# Skip authentication if custom api_base is provided and it's not a PSC endpoint
if api_base is not None and not use_psc_endpoint_format:
verbose_logger.debug(
"Skipping Vertex AI authentication - custom api_base provided without PSC endpoint format"
)
# Return empty token and use provided project_id or empty string
return "", project_id or ""
# Perform authentication for:
# 1. No custom api_base (standard Vertex AI)
# 2. Custom api_base with PSC endpoint format (PSC endpoints need auth)
return self.get_access_token(
credentials=credentials,
project_id=project_id,
)
def is_using_v1beta1_features(self, optional_params: dict) -> bool:
"""
@ -626,20 +653,47 @@ class VertexBase:
custom_llm_provider: Literal[
"vertex_ai", "vertex_ai_beta", "gemini"
], # if it's vertex_ai or gemini (google ai studio)
api_base: Optional[str] = None,
use_psc_endpoint_format: bool = False,
) -> Tuple[str, str]:
"""
Async version of _ensure_access_token
Args:
credentials: Vertex AI credentials
project_id: Google Cloud project ID
custom_llm_provider: Provider type (vertex_ai, vertex_ai_beta, or gemini)
api_base: Custom API base URL (e.g., for proxies)
use_psc_endpoint_format: Whether using PSC endpoint format
Returns:
Tuple of (access_token, project_id)
Note:
Authentication is skipped when using a custom api_base that is not a PSC endpoint.
PSC endpoints still require Google authentication even with custom api_base.
"""
if custom_llm_provider == "gemini":
return "", ""
else:
try:
return await asyncify(self.get_access_token)(
credentials=credentials,
project_id=project_id,
)
except Exception as e:
raise e
# Skip authentication if custom api_base is provided and it's not a PSC endpoint
if api_base is not None and not use_psc_endpoint_format:
verbose_logger.debug(
"Skipping Vertex AI authentication - custom api_base provided without PSC endpoint format"
)
# Return empty token and use provided project_id or empty string
return "", project_id or ""
# Perform authentication for:
# 1. No custom api_base (standard Vertex AI)
# 2. Custom api_base with PSC endpoint format (PSC endpoints need auth)
try:
return await asyncify(self.get_access_token)(
credentials=credentials,
project_id=project_id,
)
except Exception as e:
raise e
def set_headers(
self, auth_header: Optional[str], extra_headers: Optional[dict]

View file

@ -93,16 +93,21 @@ class VertexAIModelGardenModels(VertexBase):
model = get_vertex_base_model_name(model=model)
vertex_httpx_logic = VertexLLM()
## CONSTRUCT API BASE
stream: bool = optional_params.get("stream", False) or False
# Extract use_psc_endpoint_format from optional_params
use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False)
access_token, project_id = vertex_httpx_logic._ensure_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
custom_llm_provider="vertex_ai",
api_base=api_base,
use_psc_endpoint_format=use_psc_endpoint_format,
)
openai_like_chat_completions = OpenAILikeChatHandler()
## CONSTRUCT API BASE
stream: bool = optional_params.get("stream", False) or False
optional_params["stream"] = stream
default_api_base = create_vertex_url(
vertex_location=vertex_location or "us-central1",

View file

@ -1048,3 +1048,146 @@ class TestVertexBase:
MockCredentials.from_info.assert_called_once_with(json_obj)
mock_creds.with_scopes.assert_called_once_with(scopes)
assert result == "scoped_creds"
@pytest.mark.parametrize("is_async", [True, False], ids=["async", "sync"])
@pytest.mark.asyncio
async def test_skip_auth_with_custom_api_base(self, is_async):
"""
Test that authentication is skipped when using a custom api_base
without PSC endpoint format (e.g., custom proxy that doesn't require Google credentials)
"""
vertex_base = VertexBase()
# Test case 1: Custom api_base without PSC format should skip authentication
if is_async:
token, project = await vertex_base._ensure_access_token_async(
credentials=None,
project_id="test-project",
custom_llm_provider="vertex_ai",
api_base="https://custom-proxy.example.com",
use_psc_endpoint_format=False,
)
else:
token, project = vertex_base._ensure_access_token(
credentials=None,
project_id="test-project",
custom_llm_provider="vertex_ai",
api_base="https://custom-proxy.example.com",
use_psc_endpoint_format=False,
)
# Should return empty token and the provided project_id
assert token == ""
assert project == "test-project"
@pytest.mark.parametrize("is_async", [True, False], ids=["async", "sync"])
@pytest.mark.asyncio
async def test_require_auth_with_psc_endpoint(self, is_async):
"""
Test that authentication is still required when using PSC endpoint format,
even with custom api_base
"""
vertex_base = VertexBase()
# Mock credentials for PSC endpoint
mock_creds = MagicMock()
mock_creds.token = "psc-token"
mock_creds.expired = False
mock_creds.project_id = "psc-project"
mock_creds.quota_project_id = "psc-project"
# Test case 2: Custom api_base WITH PSC format should require authentication
with patch.object(
vertex_base, "load_auth", return_value=(mock_creds, "psc-project")
):
if is_async:
token, project = await vertex_base._ensure_access_token_async(
credentials={"type": "service_account"},
project_id="psc-project",
custom_llm_provider="vertex_ai",
api_base="https://10.0.0.1",
use_psc_endpoint_format=True,
)
else:
token, project = vertex_base._ensure_access_token(
credentials={"type": "service_account"},
project_id="psc-project",
custom_llm_provider="vertex_ai",
api_base="https://10.0.0.1",
use_psc_endpoint_format=True,
)
# Should return actual token from authentication
assert token == "psc-token"
assert project == "psc-project"
@pytest.mark.parametrize("is_async", [True, False], ids=["async", "sync"])
@pytest.mark.asyncio
async def test_require_auth_without_custom_api_base(self, is_async):
"""
Test that authentication is required when no custom api_base is provided
(standard Vertex AI usage)
"""
vertex_base = VertexBase()
# Mock credentials for standard Vertex AI
mock_creds = MagicMock()
mock_creds.token = "standard-token"
mock_creds.expired = False
mock_creds.project_id = "standard-project"
mock_creds.quota_project_id = "standard-project"
# Test case 3: No custom api_base should require authentication
with patch.object(
vertex_base, "load_auth", return_value=(mock_creds, "standard-project")
):
if is_async:
token, project = await vertex_base._ensure_access_token_async(
credentials={"type": "service_account"},
project_id="standard-project",
custom_llm_provider="vertex_ai",
api_base=None,
use_psc_endpoint_format=False,
)
else:
token, project = vertex_base._ensure_access_token(
credentials={"type": "service_account"},
project_id="standard-project",
custom_llm_provider="vertex_ai",
api_base=None,
use_psc_endpoint_format=False,
)
# Should return actual token from authentication
assert token == "standard-token"
assert project == "standard-project"
@pytest.mark.parametrize("is_async", [True, False], ids=["async", "sync"])
@pytest.mark.asyncio
async def test_skip_auth_with_custom_api_base_no_project(self, is_async):
"""
Test that authentication is skipped with custom api_base even when project_id is None
"""
vertex_base = VertexBase()
# Test case 4: Custom api_base without project_id should still skip auth
if is_async:
token, project = await vertex_base._ensure_access_token_async(
credentials=None,
project_id=None,
custom_llm_provider="vertex_ai",
api_base="https://custom-proxy.example.com",
use_psc_endpoint_format=False,
)
else:
token, project = vertex_base._ensure_access_token(
credentials=None,
project_id=None,
custom_llm_provider="vertex_ai",
api_base="https://custom-proxy.example.com",
use_psc_endpoint_format=False,
)
# Should return empty token and empty project
assert token == ""
assert project == ""