Merge pull request #9341 from BerriAI/litellm_fix_ssl_verify

[Bug Fix] - Azure OpenAI - ensure SSL verification runs
This commit is contained in:
Ishaan Jaff 2025-03-19 21:03:24 -07:00 • committed by GitHub
commit 23a09f1359
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
14 changed files with 914 additions and 306 deletions

View file

@ -43,6 +43,7 @@ class AzureAssistantsAPI(BaseAzureLLM):
api_base=api_base,
model_name="",
api_version=api_version,
is_async=False,
)
azure_openai_client = AzureOpenAI(**azure_client_params) # type: ignore
else:
@ -68,6 +69,7 @@ class AzureAssistantsAPI(BaseAzureLLM):
api_base=api_base,
model_name="",
api_version=api_version,
is_async=True,
)
azure_openai_client = AsyncAzureOpenAI(**azure_client_params)

View file

@ -1,10 +1,9 @@
import uuid
from typing import Any, Optional
from typing import Any, Coroutine, Optional, Union
from openai import AsyncAzureOpenAI, AzureOpenAI
from pydantic import BaseModel
import litellm
from litellm.litellm_core_utils.audio_utils.utils import get_audio_file_name
from litellm.types.utils import FileTypes
from litellm.utils import (
@ -14,6 +13,7 @@ from litellm.utils import (
)
from .azure import AzureChatCompletion
from .common_utils import AzureOpenAIError
class AzureAudioTranscription(AzureChatCompletion):
@ -33,20 +33,11 @@ class AzureAudioTranscription(AzureChatCompletion):
azure_ad_token: Optional[str] = None,
atranscription: bool = False,
litellm_params: Optional[dict] = None,
) -> TranscriptionResponse:
) -> Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]]:
data = {"model": model, "file": audio_file, **optional_params}
# init AzureOpenAI Client
azure_client_params = self.initialize_azure_sdk_client(
litellm_params=litellm_params or {},
api_key=api_key,
model_name=model,
api_version=api_version,
api_base=api_base,
)
if atranscription is True:
return self.async_audio_transcriptions( # type: ignore
return self.async_audio_transcriptions(
audio_file=audio_file,
data=data,
model_response=model_response,
@ -54,14 +45,26 @@ class AzureAudioTranscription(AzureChatCompletion):
api_key=api_key,
api_base=api_base,
client=client,
azure_client_params=azure_client_params,
max_retries=max_retries,
logging_obj=logging_obj,
model=model,
litellm_params=litellm_params,
)
azure_client = self.get_azure_openai_client(
api_version=api_version,
api_base=api_base,
api_key=api_key,
model=model,
_is_async=False,
client=client,
litellm_params=litellm_params,
)
if not isinstance(azure_client, AzureOpenAI):
raise AzureOpenAIError(
status_code=500,
message="azure_client is not an instance of AzureOpenAI",
)
if client is None:
azure_client = AzureOpenAI(http_client=litellm.client_session, **azure_client_params) # type: ignore
else:
azure_client = client
## LOGGING
logging_obj.pre_call(
@ -98,24 +101,34 @@ class AzureAudioTranscription(AzureChatCompletion):
async def async_audio_transcriptions(
self,
audio_file: FileTypes,
model: str,
data: dict,
model_response: TranscriptionResponse,
timeout: float,
azure_client_params: dict,
logging_obj: Any,
api_version: Optional[str] = None,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
client=None,
max_retries=None,
):
litellm_params: Optional[dict] = None,
) -> TranscriptionResponse:
response = None
try:
if client is None:
async_azure_client = AsyncAzureOpenAI(
**azure_client_params,
async_azure_client = self.get_azure_openai_client(
api_version=api_version,
api_base=api_base,
api_key=api_key,
model=model,
_is_async=True,
client=client,
litellm_params=litellm_params,
)
if not isinstance(async_azure_client, AsyncAzureOpenAI):
raise AzureOpenAIError(
status_code=500,
message="async_azure_client is not an instance of AsyncAzureOpenAI",
)
else:
async_azure_client = client
## LOGGING
logging_obj.pre_call(
@ -168,7 +181,12 @@ class AzureAudioTranscription(AzureChatCompletion):
model_response_object=model_response,
hidden_params=hidden_params,
response_type="audio_transcription",
) # type: ignore
)
if not isinstance(response, TranscriptionResponse):
raise AzureOpenAIError(
status_code=500,
message="response is not an instance of TranscriptionResponse",
)
return response
except Exception as e:
## LOGGING

View file

@ -1,7 +1,7 @@
import asyncio
import json
import time
from typing import Any, Callable, Dict, List, Literal, Optional, Union
from typing import Any, Callable, Coroutine, Dict, List, Optional, Union
import httpx # type: ignore
from openai import APITimeoutError, AsyncAzureOpenAI, AzureOpenAI
@ -141,41 +141,6 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
return headers
def _get_sync_azure_client(
self,
api_version: Optional[str],
api_base: Optional[str],
api_key: Optional[str],
azure_ad_token: Optional[str],
azure_ad_token_provider: Optional[Callable],
model: str,
max_retries: int,
timeout: Union[float, httpx.Timeout],
client: Optional[Any],
client_type: Literal["sync", "async"],
litellm_params: Optional[dict] = None,
):
# init AzureOpenAI Client
azure_client_params: Dict[str, Any] = self.initialize_azure_sdk_client(
litellm_params=litellm_params or {},
api_key=api_key,
model_name=model,
api_version=api_version,
api_base=api_base,
)
if client is None:
if client_type == "sync":
azure_client = AzureOpenAI(**azure_client_params) # type: ignore
elif client_type == "async":
azure_client = AsyncAzureOpenAI(**azure_client_params) # type: ignore
else:
azure_client = client
if api_version is not None and isinstance(azure_client._custom_query, dict):
# set api_version to version passed by user
azure_client._custom_query.setdefault("api-version", api_version)
return azure_client
def make_sync_azure_openai_chat_completion_request(
self,
azure_client: AzureOpenAI,
@ -263,47 +228,21 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
max_retries = DEFAULT_MAX_RETRIES
json_mode: Optional[bool] = optional_params.pop("json_mode", False)
azure_client_params = self.initialize_azure_sdk_client(
litellm_params=litellm_params or {},
api_key=api_key,
api_base=api_base,
model_name=model,
api_version=api_version,
)
### CHECK IF CLOUDFLARE AI GATEWAY ###
### if so - set the model as part of the base url
if "gateway.ai.cloudflare.com" in api_base:
## build base url - assume api base includes resource name
if client is None:
if not api_base.endswith("/"):
api_base += "/"
api_base += f"{model}"
azure_client_params = {
"api_version": api_version,
"base_url": f"{api_base}",
"http_client": litellm.client_session,
"max_retries": max_retries,
"timeout": timeout,
}
if api_key is not None:
azure_client_params["api_key"] = api_key
elif azure_ad_token is not None:
if azure_ad_token.startswith("oidc/"):
azure_ad_token = get_azure_ad_token_from_oidc(
azure_ad_token
)
azure_client_params["azure_ad_token"] = azure_ad_token
elif azure_ad_token_provider is not None:
azure_client_params["azure_ad_token_provider"] = (
azure_ad_token_provider
)
if acompletion is True:
client = AsyncAzureOpenAI(**azure_client_params)
else:
client = AzureOpenAI(**azure_client_params)
client = self._init_azure_client_for_cloudflare_ai_gateway(
api_base=api_base,
model=model,
api_version=api_version,
max_retries=max_retries,
timeout=timeout,
api_key=api_key,
azure_ad_token=azure_ad_token,
azure_ad_token_provider=azure_ad_token_provider,
acompletion=acompletion,
client=client,
)
data = {"model": None, "messages": messages, **optional_params}
else:
@ -330,7 +269,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
timeout=timeout,
client=client,
max_retries=max_retries,
azure_client_params=azure_client_params,
litellm_params=litellm_params,
)
else:
return self.acompletion(
@ -348,7 +287,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
logging_obj=logging_obj,
max_retries=max_retries,
convert_tool_call_to_json_mode=json_mode,
azure_client_params=azure_client_params,
litellm_params=litellm_params,
)
elif "stream" in optional_params and optional_params["stream"] is True:
return self.streaming(
@ -364,6 +303,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
timeout=timeout,
client=client,
max_retries=max_retries,
litellm_params=litellm_params,
)
else:
## LOGGING
@ -385,21 +325,15 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
status_code=422, message="max retries must be an int"
)
# init AzureOpenAI Client
if (
client is None
or not isinstance(client, AzureOpenAI)
or dynamic_params
):
azure_client = AzureOpenAI(**azure_client_params)
else:
azure_client = client
if api_version is not None and isinstance(
azure_client._custom_query, dict
):
# set api_version to version passed by user
azure_client._custom_query.setdefault(
"api-version", api_version
)
azure_client = self.get_azure_openai_client(
api_version=api_version,
api_base=api_base,
api_key=api_key,
model=model,
client=client,
_is_async=False,
litellm_params=litellm_params,
)
if not isinstance(azure_client, AzureOpenAI):
raise AzureOpenAIError(
status_code=500,
@ -459,16 +393,22 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
azure_ad_token_provider: Optional[Callable] = None,
convert_tool_call_to_json_mode: Optional[bool] = None,
client=None, # this is the AsyncAzureOpenAI
azure_client_params: dict = {},
litellm_params: Optional[dict] = {},
):
response = None
try:
# setting Azure client
if client is None or dynamic_params:
azure_client = AsyncAzureOpenAI(**azure_client_params)
else:
azure_client = client
azure_client = self.get_azure_openai_client(
api_version=api_version,
api_base=api_base,
api_key=api_key,
model=model,
client=client,
_is_async=True,
litellm_params=litellm_params,
)
if not isinstance(azure_client, AsyncAzureOpenAI):
raise ValueError("Azure client is not an instance of AsyncAzureOpenAI")
## LOGGING
logging_obj.pre_call(
input=data["messages"],
@ -554,6 +494,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
azure_ad_token: Optional[str] = None,
azure_ad_token_provider: Optional[Callable] = None,
client=None,
litellm_params: Optional[dict] = {},
):
# init AzureOpenAI Client
azure_client_params = {
@ -576,10 +517,20 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
elif azure_ad_token_provider is not None:
azure_client_params["azure_ad_token_provider"] = azure_ad_token_provider
if client is None or dynamic_params:
azure_client = AzureOpenAI(**azure_client_params)
else:
azure_client = client
azure_client = self.get_azure_openai_client(
api_version=api_version,
api_base=api_base,
api_key=api_key,
model=model,
client=client,
_is_async=False,
litellm_params=litellm_params,
)
if not isinstance(azure_client, AzureOpenAI):
raise AzureOpenAIError(
status_code=500,
message="azure_client is not an instance of AzureOpenAI",
)
## LOGGING
logging_obj.pre_call(
input=data["messages"],
@ -621,13 +572,21 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
azure_ad_token: Optional[str] = None,
azure_ad_token_provider: Optional[Callable] = None,
client=None,
azure_client_params: dict = {},
litellm_params: Optional[dict] = {},
):
try:
if client is None or dynamic_params:
azure_client = AsyncAzureOpenAI(**azure_client_params)
else:
azure_client = client
azure_client = self.get_azure_openai_client(
api_version=api_version,
api_base=api_base,
api_key=api_key,
model=model,
client=client,
_is_async=True,
litellm_params=litellm_params,
)
if not isinstance(azure_client, AsyncAzureOpenAI):
raise ValueError("Azure client is not an instance of AsyncAzureOpenAI")
## LOGGING
logging_obj.pre_call(
input=data["messages"],
@ -678,22 +637,36 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
async def aembedding(
self,
model: str,
data: dict,
model_response: EmbeddingResponse,
azure_client_params: dict,
input: list,
logging_obj: LiteLLMLoggingObj,
api_base: str,
api_key: Optional[str] = None,
api_version: Optional[str] = None,
client: Optional[AsyncAzureOpenAI] = None,
timeout=None,
):
timeout: Optional[Union[float, httpx.Timeout]] = None,
max_retries: Optional[int] = None,
azure_ad_token: Optional[str] = None,
azure_ad_token_provider: Optional[Callable] = None,
litellm_params: Optional[dict] = {},
) -> EmbeddingResponse:
response = None
try:
if client is None:
openai_aclient = AsyncAzureOpenAI(**azure_client_params)
else:
openai_aclient = client
openai_aclient = self.get_azure_openai_client(
api_version=api_version,
api_base=api_base,
api_key=api_key,
model=model,
_is_async=True,
client=client,
litellm_params=litellm_params,
)
if not isinstance(openai_aclient, AsyncAzureOpenAI):
raise ValueError("Azure client is not an instance of AsyncAzureOpenAI")
raw_response = await openai_aclient.embeddings.with_raw_response.create(
**data, timeout=timeout
)
@ -707,13 +680,19 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
additional_args={"complete_input_dict": data},
original_response=stringified_response,
)
return convert_to_model_response_object(
embedding_response = convert_to_model_response_object(
response_object=stringified_response,
model_response_object=model_response,
hidden_params={"headers": headers},
_response_headers=process_azure_headers(headers),
response_type="embedding",
)
if not isinstance(embedding_response, EmbeddingResponse):
raise AzureOpenAIError(
status_code=500,
message="embedding_response is not an instance of EmbeddingResponse",
)
return embedding_response
except Exception as e:
## LOGGING
logging_obj.post_call(
@ -742,7 +721,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
aembedding=None,
headers: Optional[dict] = None,
litellm_params: Optional[dict] = None,
) -> EmbeddingResponse:
) -> Union[EmbeddingResponse, Coroutine[Any, Any, EmbeddingResponse]]:
if headers:
optional_params["extra_headers"] = headers
if self._client_session is None:
@ -751,20 +730,6 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
data = {"model": model, "input": input, **optional_params}
if max_retries is None:
max_retries = litellm.DEFAULT_MAX_RETRIES
if not isinstance(max_retries, int):
raise AzureOpenAIError(
status_code=422, message="max retries must be an int"
)
# init AzureOpenAI Client
azure_client_params = self.initialize_azure_sdk_client(
litellm_params=litellm_params or {},
api_key=api_key,
model_name=model,
api_version=api_version,
api_base=api_base,
)
## LOGGING
logging_obj.pre_call(
input=input,
@ -776,20 +741,33 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
)
if aembedding is True:
return self.aembedding( # type: ignore
return self.aembedding(
data=data,
input=input,
model=model,
logging_obj=logging_obj,
api_key=api_key,
model_response=model_response,
azure_client_params=azure_client_params,
timeout=timeout,
client=client,
litellm_params=litellm_params,
api_base=api_base,
)
if client is None:
azure_client = AzureOpenAI(**azure_client_params) # type: ignore
else:
azure_client = client
azure_client = self.get_azure_openai_client(
api_version=api_version,
api_base=api_base,
api_key=api_key,
model=model,
_is_async=False,
client=client,
litellm_params=litellm_params,
)
if not isinstance(azure_client, AzureOpenAI):
raise AzureOpenAIError(
status_code=500,
message="azure_client is not an instance of AzureOpenAI",
)
## COMPLETION CALL
raw_response = azure_client.embeddings.with_raw_response.create(**data, timeout=timeout) # type: ignore
headers = dict(raw_response.headers)
@ -1155,6 +1133,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
model_name=model or "",
api_version=api_version,
api_base=api_base,
is_async=False,
)
if aimg_generation is True:
return self.aimage_generation(data=data, input=input, logging_obj=logging_obj, model_response=model_response, api_key=api_key, client=client, azure_client_params=azure_client_params, timeout=timeout, headers=headers) # type: ignore
@ -1240,17 +1219,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
litellm_params=litellm_params,
) # type: ignore
azure_client: AzureOpenAI = self._get_sync_azure_client(
azure_client: AzureOpenAI = self.get_azure_openai_client(
api_base=api_base,
api_version=api_version,
api_key=api_key,
azure_ad_token=azure_ad_token,
azure_ad_token_provider=azure_ad_token_provider,
model=model,
max_retries=max_retries,
timeout=timeout,
_is_async=False,
client=client,
client_type="sync",
litellm_params=litellm_params,
) # type: ignore
@ -1279,17 +1254,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
litellm_params: Optional[dict] = None,
) -> HttpxBinaryResponseContent:
azure_client: AsyncAzureOpenAI = self._get_sync_azure_client(
azure_client: AsyncAzureOpenAI = self.get_azure_openai_client(
api_base=api_base,
api_version=api_version,
api_key=api_key,
azure_ad_token=azure_ad_token,
azure_ad_token_provider=azure_ad_token_provider,
model=model,
max_retries=max_retries,
timeout=timeout,
_is_async=True,
client=client,
client_type="async",
litellm_params=litellm_params,
) # type: ignore

View file

@ -1,6 +1,6 @@
import json
import os
from typing import Callable, Optional, Union
from typing import Any, Callable, Dict, Optional, Union
import httpx
from openai import AsyncAzureOpenAI, AzureOpenAI
@ -9,6 +9,7 @@ import litellm
from litellm._logging import verbose_logger
from litellm.caching.caching import DualCache
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.openai.common_utils import BaseOpenAILLM
from litellm.secret_managers.get_azure_ad_token_provider import (
get_azure_ad_token_provider,
)
@ -244,24 +245,37 @@ def select_azure_base_url_or_endpoint(azure_client_params: dict):
return azure_client_params
class BaseAzureLLM:
class BaseAzureLLM(BaseOpenAILLM):
def get_azure_openai_client(
self,
litellm_params: dict,
api_key: Optional[str],
api_base: Optional[str],
api_version: Optional[str] = None,
client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI]] = None,
litellm_params: Optional[dict] = None,
_is_async: bool = False,
model: Optional[str] = None,
) -> Optional[Union[AzureOpenAI, AsyncAzureOpenAI]]:
openai_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI]] = None
client_initialization_params: dict = locals()
if client is None:
cached_client = self.get_cached_openai_client(
client_initialization_params=client_initialization_params,
client_type="azure",
)
if cached_client:
if isinstance(cached_client, AzureOpenAI) or isinstance(
cached_client, AsyncAzureOpenAI
):
return cached_client
azure_client_params = self.initialize_azure_sdk_client(
litellm_params=litellm_params,
litellm_params=litellm_params or {},
api_key=api_key,
api_base=api_base,
model_name="",
model_name=model,
api_version=api_version,
is_async=_is_async,
)
if _is_async is True:
openai_client = AsyncAzureOpenAI(**azure_client_params)
@ -269,7 +283,18 @@ class BaseAzureLLM:
openai_client = AzureOpenAI(**azure_client_params) # type: ignore
else:
openai_client = client
if api_version is not None and isinstance(
openai_client._custom_query, dict
):
# set api_version to version passed by user
openai_client._custom_query.setdefault("api-version", api_version)
# save client in-memory cache
self.set_cached_openai_client(
openai_client=openai_client,
client_initialization_params=client_initialization_params,
client_type="azure",
)
return openai_client
def initialize_azure_sdk_client(
@ -277,8 +302,9 @@ class BaseAzureLLM:
litellm_params: dict,
api_key: Optional[str],
api_base: Optional[str],
model_name: str,
model_name: Optional[str],
api_version: Optional[str],
is_async: bool,
) -> dict:
azure_ad_token_provider: Optional[Callable[[], str]] = None
@ -334,8 +360,13 @@ class BaseAzureLLM:
"api_version": api_version,
"azure_ad_token": azure_ad_token,
"azure_ad_token_provider": azure_ad_token_provider,
"http_client": litellm.client_session,
}
# init http client + SSL Verification settings
if is_async is True:
azure_client_params["http_client"] = self._get_async_http_client()
else:
azure_client_params["http_client"] = self._get_sync_http_client()
if max_retries is not None:
azure_client_params["max_retries"] = max_retries
if timeout is not None:
@ -351,3 +382,45 @@ class BaseAzureLLM:
)
return azure_client_params
def _init_azure_client_for_cloudflare_ai_gateway(
self,
api_base: str,
model: str,
api_version: str,
max_retries: int,
timeout: Union[float, httpx.Timeout],
api_key: Optional[str],
azure_ad_token: Optional[str],
azure_ad_token_provider: Optional[Callable[[], str]],
acompletion: bool,
client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI]] = None,
) -> Union[AzureOpenAI, AsyncAzureOpenAI]:
## build base url - assume api base includes resource name
if client is None:
if not api_base.endswith("/"):
api_base += "/"
api_base += f"{model}"
azure_client_params: Dict[str, Any] = {
"api_version": api_version,
"base_url": f"{api_base}",
"http_client": litellm.client_session,
"max_retries": max_retries,
"timeout": timeout,
}
if api_key is not None:
azure_client_params["api_key"] = api_key
elif azure_ad_token is not None:
if azure_ad_token.startswith("oidc/"):
azure_ad_token = get_azure_ad_token_from_oidc(azure_ad_token)
azure_client_params["azure_ad_token"] = azure_ad_token
if azure_ad_token_provider is not None:
azure_client_params["azure_ad_token_provider"] = azure_ad_token_provider
if acompletion is True:
client = AsyncAzureOpenAI(**azure_client_params) # type: ignore
else:
client = AzureOpenAI(**azure_client_params) # type: ignore
return client

View file

@ -2,7 +2,6 @@ from typing import Any, Callable, Optional
from openai import AsyncAzureOpenAI, AzureOpenAI
import litellm
from litellm.litellm_core_utils.prompt_templates.factory import prompt_factory
from litellm.utils import CustomStreamWrapper, ModelResponse, TextCompletionResponse
@ -12,18 +11,6 @@ from ..common_utils import AzureOpenAIError, BaseAzureLLM
openai_text_completion_config = OpenAITextCompletionConfig()
def select_azure_base_url_or_endpoint(azure_client_params: dict):
azure_endpoint = azure_client_params.get("azure_endpoint", None)
if azure_endpoint is not None:
# see : https://github.com/openai/openai-python/blob/3d61ed42aba652b547029095a7eb269ad4e1e957/src/openai/lib/azure.py#L192
if "/openai/deployments" in azure_endpoint:
# this is base_url, not an azure_endpoint
azure_client_params["base_url"] = azure_endpoint
azure_client_params.pop("azure_endpoint")
return azure_client_params
class AzureTextCompletion(BaseAzureLLM):
def __init__(self) -> None:
super().__init__()
@ -70,39 +57,22 @@ class AzureTextCompletion(BaseAzureLLM):
messages=messages, model=model, custom_llm_provider="azure_text"
)
azure_client_params = self.initialize_azure_sdk_client(
litellm_params=litellm_params or {},
api_key=api_key,
model_name=model,
api_version=api_version,
api_base=api_base,
)
### CHECK IF CLOUDFLARE AI GATEWAY ###
### if so - set the model as part of the base url
if "gateway.ai.cloudflare.com" in api_base:
## build base url - assume api base includes resource name
if client is None:
if not api_base.endswith("/"):
api_base += "/"
api_base += f"{model}"
azure_client_params = {
"api_version": api_version,
"base_url": f"{api_base}",
"http_client": litellm.client_session,
"max_retries": max_retries,
"timeout": timeout,
}
if api_key is not None:
azure_client_params["api_key"] = api_key
elif azure_ad_token is not None:
azure_client_params["azure_ad_token"] = azure_ad_token
if acompletion is True:
client = AsyncAzureOpenAI(**azure_client_params)
else:
client = AzureOpenAI(**azure_client_params)
client = self._init_azure_client_for_cloudflare_ai_gateway(
api_key=api_key,
api_version=api_version,
api_base=api_base,
model=model,
client=client,
max_retries=max_retries,
timeout=timeout,
azure_ad_token=azure_ad_token,
azure_ad_token_provider=azure_ad_token_provider,
acompletion=acompletion,
)
data = {"model": None, "prompt": prompt, **optional_params}
else:
@ -124,7 +94,7 @@ class AzureTextCompletion(BaseAzureLLM):
azure_ad_token=azure_ad_token,
timeout=timeout,
client=client,
azure_client_params=azure_client_params,
litellm_params=litellm_params,
)
else:
return self.acompletion(
@ -139,7 +109,7 @@ class AzureTextCompletion(BaseAzureLLM):
client=client,
logging_obj=logging_obj,
max_retries=max_retries,
azure_client_params=azure_client_params,
litellm_params=litellm_params,
)
elif "stream" in optional_params and optional_params["stream"] is True:
return self.streaming(
@ -152,7 +122,6 @@ class AzureTextCompletion(BaseAzureLLM):
azure_ad_token=azure_ad_token,
timeout=timeout,
client=client,
azure_client_params=azure_client_params,
)
else:
## LOGGING
@ -174,17 +143,21 @@ class AzureTextCompletion(BaseAzureLLM):
status_code=422, message="max retries must be an int"
)
# init AzureOpenAI Client
if client is None:
azure_client = AzureOpenAI(**azure_client_params)
else:
azure_client = client
if api_version is not None and isinstance(
azure_client._custom_query, dict
):
# set api_version to version passed by user
azure_client._custom_query.setdefault(
"api-version", api_version
)
azure_client = self.get_azure_openai_client(
api_key=api_key,
api_base=api_base,
api_version=api_version,
client=client,
litellm_params=litellm_params,
_is_async=False,
model=model,
)
if not isinstance(azure_client, AzureOpenAI):
raise AzureOpenAIError(
status_code=500,
message="azure_client is not an instance of AzureOpenAI",
)
raw_response = azure_client.completions.with_raw_response.create(
**data, timeout=timeout
@ -233,21 +206,27 @@ class AzureTextCompletion(BaseAzureLLM):
max_retries: int,
azure_ad_token: Optional[str] = None,
client=None, # this is the AsyncAzureOpenAI
azure_client_params: dict = {},
litellm_params: dict = {},
):
response = None
try:
# init AzureOpenAI Client
# setting Azure client
if client is None:
azure_client = AsyncAzureOpenAI(**azure_client_params)
else:
azure_client = client
if api_version is not None and isinstance(
azure_client._custom_query, dict
):
# set api_version to version passed by user
azure_client._custom_query.setdefault("api-version", api_version)
azure_client = self.get_azure_openai_client(
api_version=api_version,
api_base=api_base,
api_key=api_key,
model=model,
_is_async=True,
client=client,
litellm_params=litellm_params,
)
if not isinstance(azure_client, AsyncAzureOpenAI):
raise AzureOpenAIError(
status_code=500,
message="azure_client is not an instance of AsyncAzureOpenAI",
)
## LOGGING
logging_obj.pre_call(
input=data["prompt"],
@ -290,7 +269,7 @@ class AzureTextCompletion(BaseAzureLLM):
timeout: Any,
azure_ad_token: Optional[str] = None,
client=None,
azure_client_params: dict = {},
litellm_params: dict = {},
):
max_retries = data.pop("max_retries", 2)
if not isinstance(max_retries, int):
@ -298,13 +277,21 @@ class AzureTextCompletion(BaseAzureLLM):
status_code=422, message="max retries must be an int"
)
# init AzureOpenAI Client
if client is None:
azure_client = AzureOpenAI(**azure_client_params)
else:
azure_client = client
if api_version is not None and isinstance(azure_client._custom_query, dict):
# set api_version to version passed by user
azure_client._custom_query.setdefault("api-version", api_version)
azure_client = self.get_azure_openai_client(
api_version=api_version,
api_base=api_base,
api_key=api_key,
model=model,
_is_async=False,
client=client,
litellm_params=litellm_params,
)
if not isinstance(azure_client, AzureOpenAI):
raise AzureOpenAIError(
status_code=500,
message="azure_client is not an instance of AzureOpenAI",
)
## LOGGING
logging_obj.pre_call(
input=data["prompt"],
@ -339,19 +326,24 @@ class AzureTextCompletion(BaseAzureLLM):
timeout: Any,
azure_ad_token: Optional[str] = None,
client=None,
azure_client_params: dict = {},
litellm_params: dict = {},
):
try:
# init AzureOpenAI Client
if client is None:
azure_client = AsyncAzureOpenAI(**azure_client_params)
else:
azure_client = client
if api_version is not None and isinstance(
azure_client._custom_query, dict
):
# set api_version to version passed by user
azure_client._custom_query.setdefault("api-version", api_version)
azure_client = self.get_azure_openai_client(
api_version=api_version,
api_base=api_base,
api_key=api_key,
model=model,
_is_async=True,
client=client,
litellm_params=litellm_params,
)
if not isinstance(azure_client, AsyncAzureOpenAI):
raise AzureOpenAIError(
status_code=500,
message="azure_client is not an instance of AsyncAzureOpenAI",
)
## LOGGING
logging_obj.pre_call(
input=data["prompt"],

View file

@ -2,13 +2,17 @@
Common helpers / utils across al OpenAI endpoints
"""
import hashlib
import json
from typing import Any, Dict, List, Optional, Union
from typing import Any, Dict, List, Literal, Optional, Union
import httpx
import openai
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
import litellm
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.custom_httpx.http_handler import _DEFAULT_TTL_FOR_HTTPX_CLIENTS
class OpenAIError(BaseLLMException):
@ -92,3 +96,113 @@ def drop_params_from_unprocessable_entity_error(
new_data = {k: v for k, v in data.items() if k not in invalid_params}
return new_data
class BaseOpenAILLM:
"""
Base class for OpenAI LLMs for getting their httpx clients and SSL verification settings
"""
@staticmethod
def get_cached_openai_client(
client_initialization_params: dict, client_type: Literal["openai", "azure"]
) -> Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]]:
"""Retrieves the OpenAI client from the in-memory cache based on the client initialization parameters"""
_cache_key = BaseOpenAILLM.get_openai_client_cache_key(
client_initialization_params=client_initialization_params,
client_type=client_type,
)
_cached_client = litellm.in_memory_llm_clients_cache.get_cache(_cache_key)
return _cached_client
@staticmethod
def set_cached_openai_client(
openai_client: Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI],
client_type: Literal["openai", "azure"],
client_initialization_params: dict,
):
"""Stores the OpenAI client in the in-memory cache for _DEFAULT_TTL_FOR_HTTPX_CLIENTS SECONDS"""
_cache_key = BaseOpenAILLM.get_openai_client_cache_key(
client_initialization_params=client_initialization_params,
client_type=client_type,
)
litellm.in_memory_llm_clients_cache.set_cache(
key=_cache_key,
value=openai_client,
ttl=_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
)
@staticmethod
def get_openai_client_cache_key(
client_initialization_params: dict, client_type: Literal["openai", "azure"]
) -> str:
"""Creates a cache key for the OpenAI client based on the client initialization parameters"""
hashed_api_key = None
if client_initialization_params.get("api_key") is not None:
hash_object = hashlib.sha256(
client_initialization_params.get("api_key", "").encode()
)
# Hexadecimal representation of the hash
hashed_api_key = hash_object.hexdigest()
# Create a more readable cache key using a list of key-value pairs
key_parts = [
f"hashed_api_key={hashed_api_key}",
f"is_async={client_initialization_params.get('is_async')}",
]
LITELLM_CLIENT_SPECIFIC_PARAMS = [
"timeout",
"max_retries",
"organization",
"api_base",
]
openai_client_fields = (
BaseOpenAILLM.get_openai_client_initialization_param_fields(
client_type=client_type
)
+ LITELLM_CLIENT_SPECIFIC_PARAMS
)
for param in openai_client_fields:
key_parts.append(f"{param}={client_initialization_params.get(param)}")
_cache_key = ",".join(key_parts)
return _cache_key
@staticmethod
def get_openai_client_initialization_param_fields(
client_type: Literal["openai", "azure"]
) -> List[str]:
"""Returns a list of fields that are used to initialize the OpenAI client"""
import inspect
from openai import AzureOpenAI, OpenAI
if client_type == "openai":
signature = inspect.signature(OpenAI.__init__)
else:
signature = inspect.signature(AzureOpenAI.__init__)
# Extract parameter names, excluding 'self'
param_names = [param for param in signature.parameters if param != "self"]
return param_names
@staticmethod
def _get_async_http_client() -> Optional[httpx.AsyncClient]:
if litellm.aclient_session is not None:
return litellm.aclient_session
return httpx.AsyncClient(
limits=httpx.Limits(max_connections=1000, max_keepalive_connections=100),
verify=litellm.ssl_verify,
)
@staticmethod
def _get_sync_http_client() -> Optional[httpx.Client]:
if litellm.client_session is not None:
return litellm.client_session
return httpx.Client(
limits=httpx.Limits(max_connections=1000, max_keepalive_connections=100),
verify=litellm.ssl_verify,
)

View file

@ -33,7 +33,6 @@ from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.llms.bedrock.chat.invoke_handler import MockResponseIterator
from litellm.llms.custom_httpx.http_handler import _DEFAULT_TTL_FOR_HTTPX_CLIENTS
from litellm.types.utils import (
EmbeddingResponse,
ImageResponse,
@ -50,7 +49,11 @@ from litellm.utils import (
from ...types.llms.openai import *
from ..base import BaseLLM
from .chat.o_series_transformation import OpenAIOSeriesConfig
from .common_utils import OpenAIError, drop_params_from_unprocessable_entity_error
from .common_utils import (
BaseOpenAILLM,
OpenAIError,
drop_params_from_unprocessable_entity_error,
)
openaiOSeriesConfig = OpenAIOSeriesConfig()
@ -317,7 +320,7 @@ class OpenAIChatCompletionResponseIterator(BaseModelResponseIterator):
raise e
class OpenAIChatCompletion(BaseLLM):
class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
def __init__(self) -> None:
super().__init__()
@ -343,7 +346,8 @@ class OpenAIChatCompletion(BaseLLM):
max_retries: Optional[int] = DEFAULT_MAX_RETRIES,
organization: Optional[str] = None,
client: Optional[Union[OpenAI, AsyncOpenAI]] = None,
):
) -> Optional[Union[OpenAI, AsyncOpenAI]]:
client_initialization_params: Dict = locals()
if client is None:
if not isinstance(max_retries, int):
raise OpenAIError(
@ -352,25 +356,21 @@ class OpenAIChatCompletion(BaseLLM):
max_retries
),
)
# Creating a new OpenAI Client
# check in memory cache before creating a new one
# Convert the API key to bytes
hashed_api_key = None
if api_key is not None:
hash_object = hashlib.sha256(api_key.encode())
# Hexadecimal representation of the hash
hashed_api_key = hash_object.hexdigest()
cached_client = self.get_cached_openai_client(
client_initialization_params=client_initialization_params,
client_type="openai",
)
_cache_key = f"hashed_api_key={hashed_api_key},api_base={api_base},timeout={timeout},max_retries={max_retries},organization={organization},is_async={is_async}"
_cached_client = litellm.in_memory_llm_clients_cache.get_cache(_cache_key)
if _cached_client:
return _cached_client
if cached_client:
if isinstance(cached_client, OpenAI) or isinstance(
cached_client, AsyncOpenAI
):
return cached_client
if is_async:
_new_client: Union[OpenAI, AsyncOpenAI] = AsyncOpenAI(
api_key=api_key,
base_url=api_base,
http_client=litellm.aclient_session,
http_client=OpenAIChatCompletion._get_async_http_client(),
timeout=timeout,
max_retries=max_retries,
organization=organization,
@ -379,17 +379,17 @@ class OpenAIChatCompletion(BaseLLM):
_new_client = OpenAI(
api_key=api_key,
base_url=api_base,
http_client=litellm.client_session,
http_client=OpenAIChatCompletion._get_sync_http_client(),
timeout=timeout,
max_retries=max_retries,
organization=organization,
)
## SAVE CACHE KEY
litellm.in_memory_llm_clients_cache.set_cache(
key=_cache_key,
value=_new_client,
ttl=_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
self.set_cached_openai_client(
openai_client=_new_client,
client_initialization_params=client_initialization_params,
client_type="openai",
)
return _new_client

View file

@ -25,6 +25,7 @@ from functools import partial
from typing import (
Any,
Callable,
Coroutine,
Dict,
List,
Literal,
@ -3288,7 +3289,7 @@ def embedding( # noqa: PLR0915
litellm_call_id=None,
logger_fn=None,
**kwargs,
) -> EmbeddingResponse:
) -> Union[EmbeddingResponse, Coroutine[Any, Any, EmbeddingResponse]]:
"""
Embedding function that calls an API to generate embeddings for the given input.
@ -3409,7 +3410,9 @@ def embedding( # noqa: PLR0915
if mock_response is not None:
return mock_embedding(model=model, mock_response=mock_response)
try:
response: Optional[EmbeddingResponse] = None
response: Optional[
Union[EmbeddingResponse, Coroutine[Any, Any, EmbeddingResponse]]
] = None
if azure is True or custom_llm_provider == "azure":
# azure configs
@ -3901,7 +3904,11 @@ def embedding( # noqa: PLR0915
raise LiteLLMUnknownProvider(
model=model, custom_llm_provider=custom_llm_provider
)
if response is not None and hasattr(response, "_hidden_params"):
if (
response is not None
and hasattr(response, "_hidden_params")
and isinstance(response, EmbeddingResponse)
):
response._hidden_params["custom_llm_provider"] = custom_llm_provider
if response is None:
@ -4944,6 +4951,10 @@ async def atranscription(*args, **kwargs) -> TranscriptionResponse:
else:
# Call the synchronous function using run_in_executor
response = await loop.run_in_executor(None, func_with_context)
if not isinstance(response, TranscriptionResponse):
raise ValueError(
f"Invalid response from transcription provider, expected TranscriptionResponse, but got {type(response)}"
)
return response
except Exception as e:
custom_llm_provider = custom_llm_provider or "openai"
@ -4977,7 +4988,7 @@ def transcription(
max_retries: Optional[int] = None,
custom_llm_provider=None,
**kwargs,
) -> TranscriptionResponse:
) -> Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]]:
"""
Calls openai + azure whisper endpoints.
@ -5046,7 +5057,9 @@ def transcription(
custom_llm_provider=custom_llm_provider,
)
response: Optional[TranscriptionResponse] = None
response: Optional[
Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]]
] = None
if custom_llm_provider == "azure":
# azure configs
api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")

View file

@ -1,6 +1,9 @@
model_list:
- model_name: gpt-4o
- model_name: gpt-3.5-turbo-end-user-test
litellm_params:
model: gpt-4o
model: azure/chatgpt-v-2
api_base: https://openai-gpt-4-test-v-1.openai.azure.com/
api_version: "2023-05-15"
api_key: os.environ/AZURE_API_KEY

View file

@ -0,0 +1,108 @@
import ast
import os
import re
def find_azure_files(base_dir):
"""
Find all Python files in the Azure directory.
"""
azure_files = []
for root, _, files in os.walk(base_dir):
for file in files:
if file.endswith(".py"):
azure_files.append(os.path.join(root, file))
return azure_files
def check_direct_instantiation(file_path):
"""
Check if a file directly instantiates AzureOpenAI or AsyncAzureOpenAI
outside of the BaseAzureLLM class methods.
"""
with open(file_path, "r") as file:
content = file.read()
# Parse the file
tree = ast.parse(content)
# Track issues found
issues = []
# Find all class definitions
for node in ast.walk(tree):
if isinstance(node, ast.ClassDef):
class_name = node.name
# Skip BaseAzureLLM class since it's allowed to define the client creation methods
if class_name == "BaseAzureLLM":
continue
# Check method bodies for direct instantiation
for method in node.body:
if isinstance(method, ast.FunctionDef) or isinstance(
method, ast.AsyncFunctionDef
):
method_name = method.name
# Skip methods that are specifically for client creation
if method_name in [
"get_azure_openai_client",
"initialize_azure_sdk_client",
]:
continue
# Look for direct instantiation in the method body
for subnode in ast.walk(method):
if isinstance(subnode, ast.Call):
if hasattr(subnode, "func") and hasattr(subnode.func, "id"):
if subnode.func.id in [
"AzureOpenAI",
"AsyncAzureOpenAI",
]:
issues.append(
f"Direct instantiation of {subnode.func.id} in {class_name}.{method_name}"
)
elif hasattr(subnode, "func") and hasattr(
subnode.func, "attr"
):
if subnode.func.attr in [
"AzureOpenAI",
"AsyncAzureOpenAI",
]:
issues.append(
f"Direct instantiation of {subnode.func.attr} in {class_name}.{method_name}"
)
return issues
def main():
"""
Main function to run the test.
"""
# local
base_dir = "../../litellm/llms/azure"
azure_files = find_azure_files(base_dir)
print(f"Found {len(azure_files)} Azure Python files to check")
all_issues = []
for file_path in azure_files:
issues = check_direct_instantiation(file_path)
if issues:
all_issues.extend([f"{file_path}: {issue}" for issue in issues])
if all_issues:
print("Found direct instantiations of AzureOpenAI or AsyncAzureOpenAI:")
for issue in all_issues:
print(f" - {issue}")
raise Exception(
f"Found {len(all_issues)} direct instantiations of AzureOpenAI or AsyncAzureOpenAI classes. Use get_azure_openai_client instead."
)
else:
print("All Azure modules are correctly using get_azure_openai_client!")
if __name__ == "__main__":
main()

View file

@ -10,6 +10,7 @@ ALLOWED_FILES = [
"../../litellm/llms/huggingface_restapi.py",
"../../litellm/llms/base.py",
"../../litellm/llms/custom_httpx/httpx_handler.py",
"../../litellm/llms/openai/common_utils.py",
# when running on ci/cd
"./litellm/__init__.py",
"./litellm/llms/custom_httpx/http_handler.py",
@ -18,6 +19,7 @@ ALLOWED_FILES = [
"./litellm/llms/huggingface_restapi.py",
"./litellm/llms/base.py",
"./litellm/llms/custom_httpx/httpx_handler.py",
"./litellm/llms/openai/common_utils.py",
]
warning_msg = "this is a serious violation that can impact latency. Creating Async clients per request can add +500ms per request"

View file

@ -68,6 +68,7 @@ def test_initialize_with_api_key(setup_mocks):
api_base="https://test.openai.azure.com",
model_name="gpt-4",
api_version="2023-06-01",
is_async=False,
)
# Verify expected result
@ -90,6 +91,7 @@ def test_initialize_with_tenant_credentials(setup_mocks):
api_base="https://test.openai.azure.com",
model_name="gpt-4",
api_version=None,
is_async=False,
)
# Verify that get_azure_ad_token_from_entrata_id was called
@ -117,6 +119,7 @@ def test_initialize_with_username_password(setup_mocks):
api_base="https://test.openai.azure.com",
model_name="gpt-4",
api_version=None,
is_async=False,
)
# Verify that get_azure_ad_token_from_username_password was called
@ -138,6 +141,7 @@ def test_initialize_with_oidc_token(setup_mocks):
api_base="https://test.openai.azure.com",
model_name="gpt-4",
api_version=None,
is_async=False,
)
# Verify that get_azure_ad_token_from_oidc was called
@ -158,6 +162,7 @@ def test_initialize_with_enable_token_refresh(setup_mocks):
api_base="https://test.openai.azure.com",
model_name="gpt-4",
api_version=None,
is_async=False,
)
# Verify that get_azure_ad_token_provider was called
@ -179,6 +184,7 @@ def test_initialize_with_token_refresh_error(setup_mocks):
api_base="https://test.openai.azure.com",
model_name="gpt-4",
api_version=None,
is_async=False,
)
# Verify error was logged
@ -196,6 +202,7 @@ def test_api_version_from_env_var(setup_mocks):
api_base="https://test.openai.azure.com",
model_name="gpt-4",
api_version=None,
is_async=False,
)
# Verify expected result
@ -210,6 +217,7 @@ def test_select_azure_base_url_called(setup_mocks):
api_base="https://test.openai.azure.com",
model_name="gpt-4",
api_version="2023-06-01",
is_async=False,
)
# Verify that select_azure_base_url_or_endpoint was called
@ -300,12 +308,10 @@ async def test_ensure_initialize_azure_sdk_client_always_used(call_type):
# Get appropriate input for this call type
input_kwarg = test_inputs.get(call_type.value, {})
patch_target = "litellm.main.azure_chat_completions.initialize_azure_sdk_client"
if call_type == CallTypes.atranscription:
patch_target = (
"litellm.main.azure_audio_transcriptions.initialize_azure_sdk_client"
)
elif call_type == CallTypes.arerank:
patch_target = (
"litellm.llms.azure.common_utils.BaseAzureLLM.initialize_azure_sdk_client"
)
if call_type == CallTypes.arerank:
patch_target = (
"litellm.rerank_api.main.azure_rerank.initialize_azure_sdk_client"
)
@ -455,3 +461,177 @@ async def test_ensure_initialize_azure_sdk_client_always_used_azure_text(call_ty
for call in azure_calls:
assert "api_key" in call.kwargs, "api_key not found in parameters"
assert "api_base" in call.kwargs, "api_base not found in parameters"
# Test parameters for different API functions with Azure models
AZURE_API_FUNCTION_PARAMS = [
# (function_name, is_async, args)
(
"completion",
False,
{
"model": "azure/gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"max_tokens": 10,
"api_key": "test-api-key",
"api_base": "https://test.openai.azure.com",
"api_version": "2023-05-15",
},
),
(
"completion",
True,
{
"model": "azure/gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"max_tokens": 10,
"stream": True,
"api_key": "test-api-key",
"api_base": "https://test.openai.azure.com",
"api_version": "2023-05-15",
},
),
(
"embedding",
False,
{
"model": "azure/text-embedding-ada-002",
"input": "Hello world",
"api_key": "test-api-key",
"api_base": "https://test.openai.azure.com",
"api_version": "2023-05-15",
},
),
(
"embedding",
True,
{
"model": "azure/text-embedding-ada-002",
"input": "Hello world",
"api_key": "test-api-key",
"api_base": "https://test.openai.azure.com",
"api_version": "2023-05-15",
},
),
(
"speech",
False,
{
"model": "azure/tts-1",
"input": "Hello, this is a test of text to speech",
"voice": "alloy",
"api_key": "test-api-key",
"api_base": "https://test.openai.azure.com",
"api_version": "2023-05-15",
},
),
(
"speech",
True,
{
"model": "azure/tts-1",
"input": "Hello, this is a test of text to speech",
"voice": "alloy",
"api_key": "test-api-key",
"api_base": "https://test.openai.azure.com",
"api_version": "2023-05-15",
},
),
(
"transcription",
False,
{
"model": "azure/whisper-1",
"file": MagicMock(),
"api_key": "test-api-key",
"api_base": "https://test.openai.azure.com",
"api_version": "2023-05-15",
},
),
(
"transcription",
True,
{
"model": "azure/whisper-1",
"file": MagicMock(),
"api_key": "test-api-key",
"api_base": "https://test.openai.azure.com",
"api_version": "2023-05-15",
},
),
]
@pytest.mark.parametrize("function_name,is_async,args", AZURE_API_FUNCTION_PARAMS)
@pytest.mark.asyncio
async def test_azure_client_reuse(function_name, is_async, args):
"""
Test that multiple Azure API calls reuse the same Azure OpenAI client
"""
litellm.set_verbose = True
# Determine which client class to mock based on whether the test is async
client_path = (
"litellm.llms.azure.common_utils.AsyncAzureOpenAI"
if is_async
else "litellm.llms.azure.common_utils.AzureOpenAI"
)
# Create a proper mock class that can pass isinstance checks
mock_client = MagicMock()
# Create the appropriate patches
with patch(client_path) as mock_client_class, patch.object(
BaseAzureLLM, "set_cached_openai_client"
) as mock_set_cache, patch.object(
BaseAzureLLM, "get_cached_openai_client"
) as mock_get_cache, patch.object(
BaseAzureLLM, "initialize_azure_sdk_client"
) as mock_init_azure:
# Configure the mock client class to return our mock instance
mock_client_class.return_value = mock_client
# Setup the mock to return None first time (cache miss) then a client for subsequent calls
mock_get_cache.side_effect = [None] + [
mock_client
] * 9 # First call returns None, rest return the mock client
# Mock the initialize_azure_sdk_client to return a dict with the necessary params
mock_init_azure.return_value = {
"api_key": args.get("api_key"),
"azure_endpoint": args.get("api_base"),
"api_version": args.get("api_version"),
"azure_ad_token": None,
"azure_ad_token_provider": None,
}
# Make 10 API calls
for _ in range(10):
try:
# Call the appropriate function based on parameters
if is_async:
# Add 'a' prefix for async functions
func = getattr(litellm, f"a{function_name}")
await func(**args)
else:
func = getattr(litellm, function_name)
func(**args)
except Exception:
# We expect exceptions since we're mocking the client
pass
# Verify client was created only once
assert (
mock_client_class.call_count == 1
), f"{'Async' if is_async else ''}AzureOpenAI client should be created only once"
# Verify initialize_azure_sdk_client was called once
assert (
mock_init_azure.call_count == 1
), "initialize_azure_sdk_client should be called once"
# Verify the client was cached
assert mock_set_cache.call_count == 1, "Client should be cached once"
# Verify we tried to get from cache 10 times (once per request)
assert mock_get_cache.call_count == 10, "Should check cache for each request"

View file

@ -0,0 +1,132 @@
import os
import sys
from unittest.mock import MagicMock, call, patch
import pytest
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
import litellm
from litellm.llms.openai.common_utils import BaseOpenAILLM
# Test parameters for different API functions
API_FUNCTION_PARAMS = [
# (function_name, is_async, args)
(
"completion",
False,
{
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "Hello"}],
"max_tokens": 10,
},
),
(
"completion",
True,
{
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "Hello"}],
"max_tokens": 10,
},
),
(
"completion",
True,
{
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "Hello"}],
"max_tokens": 10,
"stream": True,
},
),
("embedding", False, {"model": "text-embedding-ada-002", "input": "Hello world"}),
("embedding", True, {"model": "text-embedding-ada-002", "input": "Hello world"}),
(
"image_generation",
False,
{"model": "dall-e-3", "prompt": "A beautiful sunset over mountains"},
),
(
"image_generation",
True,
{"model": "dall-e-3", "prompt": "A beautiful sunset over mountains"},
),
(
"speech",
False,
{
"model": "tts-1",
"input": "Hello, this is a test of text to speech",
"voice": "alloy",
},
),
(
"speech",
True,
{
"model": "tts-1",
"input": "Hello, this is a test of text to speech",
"voice": "alloy",
},
),
("transcription", False, {"model": "whisper-1", "file": MagicMock()}),
("transcription", True, {"model": "whisper-1", "file": MagicMock()}),
]
@pytest.mark.parametrize("function_name,is_async,args", API_FUNCTION_PARAMS)
@pytest.mark.asyncio
async def test_openai_client_reuse(function_name, is_async, args):
"""
Test that multiple API calls reuse the same OpenAI client
"""
litellm.set_verbose = True
# Determine which client class to mock based on whether the test is async
client_path = (
"litellm.llms.openai.openai.AsyncOpenAI"
if is_async
else "litellm.llms.openai.openai.OpenAI"
)
# Create the appropriate patches
with patch(client_path) as mock_client_class, patch.object(
BaseOpenAILLM, "set_cached_openai_client"
) as mock_set_cache, patch.object(
BaseOpenAILLM, "get_cached_openai_client"
) as mock_get_cache:
# Setup the mock to return None first time (cache miss) then a client for subsequent calls
mock_client = MagicMock()
mock_get_cache.side_effect = [None] + [
mock_client
] * 9 # First call returns None, rest return the mock client
# Make 10 API calls
for _ in range(10):
try:
# Call the appropriate function based on parameters
if is_async:
# Add 'a' prefix for async functions
func = getattr(litellm, f"a{function_name}")
await func(**args)
else:
func = getattr(litellm, function_name)
func(**args)
except Exception:
# We expect exceptions since we're mocking the client
pass
# Verify client was created only once
assert (
mock_client_class.call_count == 1
), f"{'Async' if is_async else ''}OpenAI client should be created only once"
# Verify the client was cached
assert mock_set_cache.call_count == 1, "Client should be cached once"
# Verify we tried to get from cache 10 times (once per request)
assert mock_get_cache.call_count == 10, "Should check cache for each request"

View file

@ -4364,14 +4364,14 @@ async def test_dynamic_azure_params(stream, sync_mode):
## recreate mock client
if sync_mode:
mock_client = MagicMock(return_value="Hello world!")
new_mock_client = MagicMock(return_value="Hello world!")
else:
mock_client = AsyncMock(return_value="Hello world!")
new_mock_client = AsyncMock(return_value="Hello world!")
## CHECK IF NEW CLIENT IS USED (PARAM CHANGE)
with patch.object(
client.chat.completions.with_raw_response, "create", new=mock_client
) as mock_client:
client.chat.completions.with_raw_response, "create", new=new_mock_client
) as new_mock_client:
try:
if sync_mode:
_ = completion(
@ -4393,7 +4393,7 @@ async def test_dynamic_azure_params(stream, sync_mode):
pass
try:
mock_client.assert_not_called()
new_mock_client.assert_called()
except Exception as e:
raise e