diff --git a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py index 35cd54d65f6..f4e0aa5d81d 100644 --- a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py +++ b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py @@ -158,6 +158,7 @@ class VertexGemmaConfig(OpenAIGPTConfig): logging_obj=logging_obj, optional_params=optional_params, litellm_params=litellm_params, + client= client, timeout=timeout, encoding=encoding, ) @@ -172,6 +173,7 @@ class VertexGemmaConfig(OpenAIGPTConfig): logging_obj=logging_obj, optional_params=optional_params, litellm_params=litellm_params, + client= client, timeout=timeout, encoding=encoding, ) @@ -187,11 +189,12 @@ class VertexGemmaConfig(OpenAIGPTConfig): logging_obj: Any, optional_params: dict, litellm_params: dict, - timeout: Optional[Union[float, httpx.Timeout]], - encoding: Any, + client: Any = None, + timeout: Optional[Union[float, httpx.Timeout]]= None, + encoding: Any= None, ): """Synchronous completion request""" - from litellm.llms.custom_httpx.http_handler import HTTPHandler + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.utils import convert_to_model_response_object # Check if streaming is requested (will be faked) @@ -223,13 +226,17 @@ class VertexGemmaConfig(OpenAIGPTConfig): ) # Make the HTTP request - http_handler = HTTPHandler(concurrent_limit=1) - response = http_handler.post( - url=api_base, - headers=headers, - json=request_data, - timeout=timeout, - ) + if client is None or isinstance(client, (httpx.AsyncClient, AsyncHTTPHandler)): + http_handler = HTTPHandler(concurrent_limit=1) + elif isinstance(client, httpx.Client): + http_handler = HTTPHandler(client=client) + elif isinstance(client, HTTPHandler): + http_handler = client + else: + raise BaseLLMException( + status_code=400, + message=f"Invalid sync client type: {type(client)}. Expected HTTPHandler or httpx.Client.", + ) if response.status_code != 200: raise BaseLLMException( @@ -279,11 +286,16 @@ class VertexGemmaConfig(OpenAIGPTConfig): logging_obj: Any, optional_params: dict, litellm_params: dict, - timeout: Optional[Union[float, httpx.Timeout]], - encoding: Any, + client: Optional[Any]= None, + timeout: Optional[Union[float, httpx.Timeout]]= None, + encoding: Any= None, ): """Asynchronous completion request""" - from litellm.llms.custom_httpx.http_handler import get_async_httpx_client + from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + HTTPHandler, + get_async_httpx_client, + ) from litellm.types.utils import LlmProviders from litellm.utils import convert_to_model_response_object @@ -316,9 +328,18 @@ class VertexGemmaConfig(OpenAIGPTConfig): ) # Make the HTTP request - http_handler = get_async_httpx_client( - llm_provider=LlmProviders.VERTEX_AI, - ) + if client is None or isinstance(client, (httpx.Client, HTTPHandler)): + http_handler = get_async_httpx_client( + llm_provider=LlmProviders.VERTEX_AI, + ) + elif isinstance(client, (httpx.AsyncClient, AsyncHTTPHandler)): + http_handler = client + else: + raise BaseLLMException( + status_code=400, + message=f"Invalid async client type: {type(client)}. Expected AsyncHTTPHandler or httpx.AsyncClient.", + ) + response = await http_handler.post( url=api_base, headers=headers,