fix: Vertex AI image edit credential source

This commit is contained in:
Sameer Kankute 2025-12-17 15:56:10 +05:30
parent 1c5f16fb6e
commit 969ed1efb8
2 changed files with 54 additions and 10 deletions

View file

@ -8,7 +8,6 @@ import httpx
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.gemini.vertex_and_google_ai_studio_gemini import VertexLLM
@ -94,10 +93,22 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM):
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
headers = headers or {}
vertex_project = self._resolve_vertex_project()
vertex_credentials = self._resolve_vertex_credentials()
litellm_params = litellm_params or {}
# If a custom api_base is provided, skip credential validation
# This allows users to use proxies or mock endpoints without needing Vertex AI credentials
_api_base = litellm_params.get("api_base") or api_base
if _api_base is not None:
return headers
# First check litellm_params (where vertex_ai_project/vertex_ai_credentials are passed)
# then fall back to environment variables and other sources
vertex_project = self.safe_get_vertex_ai_project(litellm_params) or self._resolve_vertex_project()
vertex_credentials = self.safe_get_vertex_ai_credentials(litellm_params) or self._resolve_vertex_credentials()
access_token, _ = self._ensure_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
@ -114,19 +125,27 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM):
"""
Get the complete URL for Vertex AI Gemini generateContent API
"""
vertex_project = self._resolve_vertex_project()
vertex_location = self._resolve_vertex_location()
if not vertex_project or not vertex_location:
raise ValueError("vertex_project and vertex_location are required for Vertex AI")
# Use the model name as provided, handling vertex_ai prefix
model_name = model
if model.startswith("vertex_ai/"):
model_name = model.replace("vertex_ai/", "")
# If a custom api_base is provided, use it directly
# This allows users to use proxies or mock endpoints
if api_base:
base_url = api_base.rstrip("/")
return api_base.rstrip("/")
# First check litellm_params (where vertex_ai_project/vertex_ai_location are passed)
# then fall back to environment variables and other sources
vertex_project = self.safe_get_vertex_ai_project(litellm_params) or self._resolve_vertex_project()
vertex_location = self.safe_get_vertex_ai_location(litellm_params) or self._resolve_vertex_location()
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"

View file

@ -140,6 +140,31 @@ class TestVertexAIGeminiImageEditTransformation:
headers={},
)
def test_validate_environment_with_litellm_params(self) -> None:
"""Test validate_environment uses credentials from litellm_params"""
with patch.object(
self.config, "_ensure_access_token", return_value=("test-token", "test-expiry")
) as mock_token:
with patch.object(self.config, "set_headers", return_value={"Authorization": "Bearer test-token"}) as mock_headers:
litellm_params = {
"vertex_ai_project": "custom-project",
"vertex_ai_credentials": "/path/to/custom/credentials.json",
}
result = self.config.validate_environment(
headers={"X-Custom": "header"},
model=self.model,
litellm_params=litellm_params,
api_base=None,
)
# Verify that safe_get_vertex_ai_project and safe_get_vertex_ai_credentials were used
mock_token.assert_called_once()
call_kwargs = mock_token.call_args[1]
assert call_kwargs["credentials"] == "/path/to/custom/credentials.json"
assert call_kwargs["project_id"] == "custom-project"
assert result == {"Authorization": "Bearer test-token"}
class TestVertexAIImagenImageEditTransformation:
def setup_method(self) -> None: