diff --git a/litellm/llms/azure/assistants.py b/litellm/llms/azure/assistants.py index 1328eb1fea8..2e8c78b259e 100644 --- a/litellm/llms/azure/assistants.py +++ b/litellm/llms/azure/assistants.py @@ -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) diff --git a/litellm/llms/azure/audio_transcriptions.py b/litellm/llms/azure/audio_transcriptions.py index 52a3d780fbd..be7d0fa30da 100644 --- a/litellm/llms/azure/audio_transcriptions.py +++ b/litellm/llms/azure/audio_transcriptions.py @@ -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 diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 7fba70141c2..03c5cc09ebe 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -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 diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index 909fcd88a5c..71092c8b993 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -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 diff --git a/litellm/llms/azure/completion/handler.py b/litellm/llms/azure/completion/handler.py index 4ec5c435dac..8301c4d617d 100644 --- a/litellm/llms/azure/completion/handler.py +++ b/litellm/llms/azure/completion/handler.py @@ -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"], diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index a8412f867b5..55da16d6cd0 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -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, + ) diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 880a043d08a..98ef95239e0 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -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 diff --git a/litellm/main.py b/litellm/main.py index 64049c31d12..e75c23f0fc1 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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") diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index c5add9ee090..6f37f0e1404 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -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 diff --git a/tests/code_coverage_tests/azure_client_usage_test.py b/tests/code_coverage_tests/azure_client_usage_test.py new file mode 100644 index 00000000000..e216f6902a8 --- /dev/null +++ b/tests/code_coverage_tests/azure_client_usage_test.py @@ -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() diff --git a/tests/code_coverage_tests/ensure_async_clients_test.py b/tests/code_coverage_tests/ensure_async_clients_test.py index 0565de9b383..db47973b692 100644 --- a/tests/code_coverage_tests/ensure_async_clients_test.py +++ b/tests/code_coverage_tests/ensure_async_clients_test.py @@ -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" diff --git a/tests/litellm/llms/azure/test_azure_common_utils.py b/tests/litellm/llms/azure/test_azure_common_utils.py index 21fa3b37eee..a9e63f84f24 100644 --- a/tests/litellm/llms/azure/test_azure_common_utils.py +++ b/tests/litellm/llms/azure/test_azure_common_utils.py @@ -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" diff --git a/tests/litellm/llms/openai/test_openai_common_utils.py b/tests/litellm/llms/openai/test_openai_common_utils.py new file mode 100644 index 00000000000..a343fcf25c5 --- /dev/null +++ b/tests/litellm/llms/openai/test_openai_common_utils.py @@ -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" diff --git a/tests/local_testing/test_completion.py b/tests/local_testing/test_completion.py index 5fe4984c17e..59f5a38f08f 100644 --- a/tests/local_testing/test_completion.py +++ b/tests/local_testing/test_completion.py @@ -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