diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 3347e533242..324205237dc 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -1125,6 +1125,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): api_base: Optional[str] = None, client: Optional[AsyncOpenAI] = None, max_retries=None, + shared_session: Optional["ClientSession"] = None, ): try: openai_aclient: AsyncOpenAI = self._get_openai_client( # type: ignore @@ -1134,6 +1135,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): timeout=timeout, max_retries=max_retries, client=client, + shared_session=shared_session, ) headers, response = await self.make_openai_embedding_request( openai_aclient=openai_aclient, @@ -1197,6 +1199,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): client=None, aembedding=None, max_retries: Optional[int] = None, + shared_session: Optional["ClientSession"] = None, ) -> EmbeddingResponse: super().embedding() try: @@ -1223,6 +1226,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): timeout=timeout, client=client, max_retries=max_retries, + shared_session=shared_session, ) openai_client: OpenAI = self._get_openai_client( # type: ignore diff --git a/litellm/main.py b/litellm/main.py index cfb0bef0797..18ad93237fe 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -3969,6 +3969,7 @@ def embedding( # noqa: PLR0915 """ azure = kwargs.get("azure", None) client = kwargs.pop("client", None) + shared_session = kwargs.get("shared_session", None) max_retries = kwargs.get("max_retries", None) litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore mock_response: Optional[List[float]] = kwargs.get("mock_response", None) # type: ignore @@ -4158,6 +4159,7 @@ def embedding( # noqa: PLR0915 client=client, aembedding=aembedding, max_retries=max_retries, + shared_session=shared_session, ) elif custom_llm_provider == "databricks": api_base = api_base or litellm.api_base or get_secret("DATABRICKS_API_BASE") # type: ignore