mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(gemini): finish review follow-ups for fallback params and async auth
This commit is contained in:
parent
05ceff8391
commit
90a7177db8
6 changed files with 214 additions and 81 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue