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.
This commit is contained in:
AlexsanderHamir 2025-10-10 13:47:39 -07:00
parent 0e8e2a656e
commit a94fefe580
2 changed files with 6 additions and 0 deletions

View file

@ -1125,6 +1125,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
api_base: Optional[str] = None, api_base: Optional[str] = None,
client: Optional[AsyncOpenAI] = None, client: Optional[AsyncOpenAI] = None,
max_retries=None, max_retries=None,
shared_session: Optional["ClientSession"] = None,
): ):
try: try:
openai_aclient: AsyncOpenAI = self._get_openai_client( # type: ignore openai_aclient: AsyncOpenAI = self._get_openai_client( # type: ignore
@ -1134,6 +1135,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
timeout=timeout, timeout=timeout,
max_retries=max_retries, max_retries=max_retries,
client=client, client=client,
shared_session=shared_session,
) )
headers, response = await self.make_openai_embedding_request( headers, response = await self.make_openai_embedding_request(
openai_aclient=openai_aclient, openai_aclient=openai_aclient,
@ -1197,6 +1199,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
client=None, client=None,
aembedding=None, aembedding=None,
max_retries: Optional[int] = None, max_retries: Optional[int] = None,
shared_session: Optional["ClientSession"] = None,
) -> EmbeddingResponse: ) -> EmbeddingResponse:
super().embedding() super().embedding()
try: try:
@ -1223,6 +1226,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
timeout=timeout, timeout=timeout,
client=client, client=client,
max_retries=max_retries, max_retries=max_retries,
shared_session=shared_session,
) )
openai_client: OpenAI = self._get_openai_client( # type: ignore openai_client: OpenAI = self._get_openai_client( # type: ignore

View file

@ -3969,6 +3969,7 @@ def embedding( # noqa: PLR0915
""" """
azure = kwargs.get("azure", None) azure = kwargs.get("azure", None)
client = kwargs.pop("client", None) client = kwargs.pop("client", None)
shared_session = kwargs.get("shared_session", None)
max_retries = kwargs.get("max_retries", None) max_retries = kwargs.get("max_retries", None)
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
mock_response: Optional[List[float]] = kwargs.get("mock_response", None) # type: ignore mock_response: Optional[List[float]] = kwargs.get("mock_response", None) # type: ignore
@ -4158,6 +4159,7 @@ def embedding( # noqa: PLR0915
client=client, client=client,
aembedding=aembedding, aembedding=aembedding,
max_retries=max_retries, max_retries=max_retries,
shared_session=shared_session,
) )
elif custom_llm_provider == "databricks": elif custom_llm_provider == "databricks":
api_base = api_base or litellm.api_base or get_secret("DATABRICKS_API_BASE") # type: ignore api_base = api_base or litellm.api_base or get_secret("DATABRICKS_API_BASE") # type: ignore