diff --git a/litellm/llms/google_code_assist/transformation.py b/litellm/llms/google_code_assist/transformation.py index 88e570cb591..f2908dc3a81 100644 --- a/litellm/llms/google_code_assist/transformation.py +++ b/litellm/llms/google_code_assist/transformation.py @@ -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 diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index ca127a4391e..fe8b60bcc29 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -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, diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py index 2371bc4865a..df77f7654b0 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py @@ -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, diff --git a/litellm/llms/vertex_ai/multimodal_embeddings/embedding_handler.py b/litellm/llms/vertex_ai/multimodal_embeddings/embedding_handler.py index d0ffc7be0a6..53ff92fb9f2 100644 --- a/litellm/llms/vertex_ai/multimodal_embeddings/embedding_handler.py +++ b/litellm/llms/vertex_ai/multimodal_embeddings/embedding_handler.py @@ -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, diff --git a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py index 5fffd983c24..e1a787f8147 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py @@ -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, diff --git a/litellm/main.py b/litellm/main.py index d1971c53c0c..a20ffb03bb6 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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,