mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
0e8e2a656e
commit
a94fefe580
2 changed files with 6 additions and 0 deletions
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue