From a94fefe580adbc300dfaa5cb714caf00a0827f16 Mon Sep 17 00:00:00 2001 From: AlexsanderHamir Date: Fri, 10 Oct 2025 13:47:39 -0700 Subject: [PATCH] Fix: Add shared_session support for embedding calls with connection pooling The shared_session parameter was not being properly handled in embedding calls, causing it to be passed through to provider API requests where it's not needed. Changes: - Added shared_session to all_litellm_params to filter it from provider API request body - Extract shared_session in main embedding() function and pass it explicitly - Updated OpenAI embedding handlers (embedding() and aembedding()) to accept shared_session - Pass shared_session to _get_openai_client for HTTP client creation This enables proper connection pooling for embedding requests when shared_session is provided, improving performance for high-throughput scenarios. --- litellm/llms/openai/openai.py | 4 ++++ litellm/main.py | 2 ++ 2 files changed, 6 insertions(+) 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