mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix: Vertex AI image edit credential source
This commit is contained in:
parent
1c5f16fb6e
commit
969ed1efb8
2 changed files with 54 additions and 10 deletions
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue