fix(gemini): finish review follow-ups for fallback params and async auth

This commit is contained in:
balazss 2026-03-18 09:45:03 -07:00
parent 05ceff8391
commit 90a7177db8
6 changed files with 214 additions and 81 deletions

View file

@ -2,7 +2,7 @@ import json
import uuid
import copy
import httpx
from typing import Any, List, Optional
from typing import Any
from litellm._logging import verbose_logger
from litellm.llms.base_llm.chat.transformation import BaseLLMException
@ -49,36 +49,6 @@ class GoogleCodeAssistConfig(VertexGeminiConfig):
- `stop_sequences` (List[str]): The set of character sequences that will stop output generation.
"""
def __init__(
self,
temperature: Optional[float] = None,
max_output_tokens: Optional[int] = None,
top_p: Optional[float] = None,
top_k: Optional[int] = None,
stop_sequences: Optional[list] = None,
) -> None:
super().__init__(
temperature=temperature,
max_output_tokens=max_output_tokens,
top_p=top_p,
top_k=top_k,
stop_sequences=stop_sequences,
)
def get_supported_openai_params(self, model: str) -> List[str]:
return super().get_supported_openai_params(model)
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
messages: list,
) -> dict:
return super().map_openai_params(
non_default_params, optional_params, model, messages
)
def transform_request(
self,
model: str,
@ -128,12 +98,6 @@ class GoogleCodeAssistConfig(VertexGeminiConfig):
generation_config["thinkingConfig"] = {
"includeThoughts": base_params.pop("include_thoughts")
}
elif "thinkingConfig" in optional_params:
generation_config["thinkingConfig"] = optional_params["thinkingConfig"]
elif "include_thoughts" in optional_params:
generation_config["thinkingConfig"] = {
"includeThoughts": optional_params["include_thoughts"]
}
if (
"thinkingConfig" in optional_params

View file

@ -2814,7 +2814,6 @@ class VertexLLM(VertexBase):
auth_header, url = self._get_token_and_url(
model=model,
gemini_api_key=gemini_api_key,
gemini_auth_data=None,
auth_header=_auth_header,
vertex_project=vertex_project,
vertex_location=vertex_location,

View file

@ -2,6 +2,7 @@
Google AI Studio /batchEmbedContents Embeddings Endpoint
"""
import asyncio
import json
from typing import Any, Dict, Literal, Optional, Union
@ -130,6 +131,36 @@ class GoogleBatchEmbeddings(VertexLLM):
client=None,
extra_headers: Optional[dict] = None,
) -> EmbeddingResponse:
optional_params = optional_params or {}
is_multimodal = _is_multimodal_input(input)
use_embed_content = is_multimodal or (custom_llm_provider == "vertex_ai")
mode: Literal["embedding", "batch_embedding"]
if use_embed_content:
mode = "embedding"
else:
mode = "batch_embedding"
if aembedding is True:
return self._async_batch_embeddings_with_auth_resolution( # type: ignore
model=model,
input=input,
model_response=model_response,
custom_llm_provider=custom_llm_provider,
optional_params=optional_params,
logging_obj=logging_obj,
api_key=api_key,
api_base=api_base,
use_embed_content=use_embed_content,
mode=mode,
timeout=timeout,
client=client,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_credentials=vertex_credentials,
extra_headers=extra_headers,
)
_auth_header, vertex_project = self._ensure_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
@ -149,16 +180,6 @@ class GoogleBatchEmbeddings(VertexLLM):
else:
sync_handler = client # type: ignore
optional_params = optional_params or {}
is_multimodal = _is_multimodal_input(input)
use_embed_content = is_multimodal or (custom_llm_provider == "vertex_ai")
mode: Literal["embedding", "batch_embedding"]
if use_embed_content:
mode = "embedding"
else:
mode = "batch_embedding"
auth_header, url = self._get_token_and_url(
model=model,
auth_header=_auth_header,
@ -184,22 +205,6 @@ class GoogleBatchEmbeddings(VertexLLM):
if extra_headers is not None:
headers.update(extra_headers)
if aembedding is True:
return self.async_batch_embeddings( # type: ignore
model=model,
api_base=api_base,
url=url,
data=None,
model_response=model_response,
timeout=timeout,
headers=headers,
input=input,
use_embed_content=use_embed_content,
api_key=api_key,
optional_params=optional_params,
logging_obj=logging_obj,
)
### TRANSFORMATION (sync path) ###
request_data: Any
if use_embed_content:
@ -257,6 +262,78 @@ class GoogleBatchEmbeddings(VertexLLM):
input=input,
)
async def _async_batch_embeddings_with_auth_resolution(
self,
model: str,
input: EmbeddingInput,
model_response: EmbeddingResponse,
custom_llm_provider: Literal["gemini", "vertex_ai"],
optional_params: dict,
logging_obj: Any,
api_key: Optional[str],
api_base: Optional[str],
use_embed_content: bool,
mode: Literal["embedding", "batch_embedding"],
timeout: Optional[Union[float, httpx.Timeout]],
client: Optional[AsyncHTTPHandler],
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_credentials: Optional[Any],
extra_headers: Optional[dict],
) -> EmbeddingResponse:
_auth_header, vertex_project = await self._ensure_access_token_async(
credentials=vertex_credentials,
project_id=vertex_project,
custom_llm_provider=custom_llm_provider,
)
gemini_auth_data = None
if custom_llm_provider == "gemini" and api_key is None:
from litellm.llms.gemini.common_utils import get_gemini_oauth_token
gemini_auth_data = await asyncio.to_thread(get_gemini_oauth_token)
auth_header, url = self._get_token_and_url(
model=model,
auth_header=_auth_header,
gemini_api_key=api_key,
gemini_auth_data=gemini_auth_data,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_credentials=vertex_credentials,
stream=None,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
should_use_v1beta1_features=False,
mode=mode,
)
headers = {
"Content-Type": "application/json; charset=utf-8",
}
if auth_header is not None:
if isinstance(auth_header, dict):
headers.update(auth_header)
else:
headers["Authorization"] = f"Bearer {auth_header}"
if extra_headers is not None:
headers.update(extra_headers)
return await self.async_batch_embeddings(
model=model,
api_base=api_base,
url=url,
data=None,
model_response=model_response,
timeout=timeout,
headers=headers,
client=client,
input=input,
use_embed_content=use_embed_content,
api_key=api_key,
optional_params=optional_params,
logging_obj=logging_obj,
)
async def async_batch_embeddings(
self,
model: str,

View file

@ -1,3 +1,4 @@
import asyncio
import json
from typing import Literal, Optional, Union
@ -50,6 +51,25 @@ class VertexMultimodalEmbedding(VertexLLM):
timeout=300,
client=None,
) -> EmbeddingResponse:
if aembedding is True:
return self._async_multimodal_embedding_with_auth_resolution( # type: ignore
model=model,
input=input,
model_response=model_response,
custom_llm_provider=custom_llm_provider,
optional_params=optional_params,
litellm_params=litellm_params,
logging_obj=logging_obj,
api_key=api_key,
api_base=api_base,
headers=headers,
timeout=timeout,
client=client,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_credentials=vertex_credentials,
)
_auth_header, vertex_project = self._ensure_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
@ -108,21 +128,6 @@ class VertexMultimodalEmbedding(VertexLLM):
},
)
if aembedding is True:
return self.async_multimodal_embedding( # type: ignore
model=model,
api_base=url,
data=request_data,
timeout=timeout,
headers=headers,
client=client,
model_response=model_response,
optional_params=optional_params,
litellm_params=litellm_params,
logging_obj=logging_obj,
api_key=api_key,
)
response = sync_handler.post(
url=url,
headers=headers,
@ -140,6 +145,87 @@ class VertexMultimodalEmbedding(VertexLLM):
litellm_params=litellm_params,
)
async def _async_multimodal_embedding_with_auth_resolution(
self,
model: str,
input: Union[list, str],
model_response: EmbeddingResponse,
custom_llm_provider: Literal["gemini", "vertex_ai"],
optional_params: dict,
litellm_params: dict,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str],
api_base: Optional[str],
headers: dict,
timeout: Optional[Union[float, httpx.Timeout]],
client: Optional[AsyncHTTPHandler],
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES],
) -> EmbeddingResponse:
_auth_header, vertex_project = await self._ensure_access_token_async(
credentials=vertex_credentials,
project_id=vertex_project,
custom_llm_provider=custom_llm_provider,
)
gemini_auth_data = None
if custom_llm_provider == "gemini" and api_key is None:
from litellm.llms.gemini.common_utils import get_gemini_oauth_token
gemini_auth_data = await asyncio.to_thread(get_gemini_oauth_token)
auth_header, url = self._get_token_and_url(
model=model,
auth_header=_auth_header,
gemini_api_key=api_key,
gemini_auth_data=gemini_auth_data,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_credentials=vertex_credentials,
stream=None,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
should_use_v1beta1_features=False,
mode="embedding",
)
request_data = vertex_multimodal_embedding_handler.transform_embedding_request(
model, input, optional_params, headers
)
request_headers = vertex_multimodal_embedding_handler.validate_environment(
headers=headers,
model=model,
messages=[],
optional_params=optional_params,
api_key=auth_header,
api_base=api_base,
litellm_params=litellm_params,
)
logging_obj.pre_call(
input=input,
api_key="",
additional_args={
"complete_input_dict": request_data,
"api_base": url,
"headers": request_headers,
},
)
return await self.async_multimodal_embedding(
model=model,
api_base=url,
data=request_data,
timeout=timeout,
headers=request_headers,
client=client,
model_response=model_response,
optional_params=optional_params,
litellm_params=litellm_params,
logging_obj=logging_obj,
api_key=api_key,
)
async def async_multimodal_embedding(
self,
model: str,

View file

@ -1,3 +1,4 @@
import asyncio
from typing import Literal, Optional, Union
import httpx
@ -166,12 +167,18 @@ class VertexEmbedding(VertexBase):
project_id=vertex_project,
custom_llm_provider=custom_llm_provider,
)
gemini_auth_data = None
if custom_llm_provider == "gemini" and gemini_api_key is None:
from litellm.llms.gemini.common_utils import get_gemini_oauth_token
gemini_auth_data = await asyncio.to_thread(get_gemini_oauth_token)
# 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,
gemini_auth_data=gemini_auth_data,
auth_header=_auth_header,
vertex_project=vertex_project,
vertex_location=vertex_location,

View file

@ -3473,7 +3473,7 @@ def completion( # type: ignore # noqa: PLR0915
"messages": messages,
"model_response": model_response,
"print_verbose": print_verbose,
"optional_params": optional_params,
"optional_params": optional_params or {},
"litellm_params": litellm_params, # type: ignore
"logging_obj": logging,
"logger_fn": logger_fn,
@ -3508,7 +3508,7 @@ def completion( # type: ignore # noqa: PLR0915
"messages": messages,
"model_response": model_response,
"print_verbose": print_verbose,
"optional_params": optional_params,
"optional_params": optional_params or {},
"litellm_params": litellm_params, # type: ignore
"logging_obj": logging,
"logger_fn": logger_fn,