mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Fix: skip auth for custom api base in vertex ai
This commit is contained in:
parent
ea2e360cb5
commit
3f8e985b58
7 changed files with 258 additions and 34 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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 == ""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue