From 6e351136d7588c21519c54846329da8785af45a0 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 08:56:08 -0700 Subject: [PATCH 01/33] handle _get_async_http_client for OpenAI --- litellm/llms/openai/openai.py | 26 ++++++++++++++++++++++++-- 1 file changed, 24 insertions(+), 2 deletions(-) diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 880a043d08a..a5e33795c82 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -370,7 +370,7 @@ class OpenAIChatCompletion(BaseLLM): _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,7 +379,7 @@ 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, @@ -401,6 +401,28 @@ class OpenAIChatCompletion(BaseLLM): ) return client + @staticmethod + def _get_async_http_client() -> Optional[httpx.AsyncClient]: + if litellm.ssl_verify: + return httpx.AsyncClient( + limits=httpx.Limits( + max_connections=1000, max_keepalive_connections=100 + ), + verify=litellm.ssl_verify, + ) + return litellm.aclient_session + + @staticmethod + def _get_sync_http_client() -> Optional[httpx.Client]: + if litellm.ssl_verify: + return httpx.Client( + limits=httpx.Limits( + max_connections=1000, max_keepalive_connections=100 + ), + verify=litellm.ssl_verify, + ) + return litellm.client_session + @track_llm_api_timing() async def make_openai_chat_completion_request( self, From c26cf69c5948ab84b873dd0ce2f2fa79cabf267b Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 09:10:59 -0700 Subject: [PATCH 02/33] docs release notes Azure OpenAI --- docs/my-website/release_notes/v1.63.11-stable/index.md | 3 +++ 1 file changed, 3 insertions(+) diff --git a/docs/my-website/release_notes/v1.63.11-stable/index.md b/docs/my-website/release_notes/v1.63.11-stable/index.md index 1a9583c8c35..30bc8297ac7 100644 --- a/docs/my-website/release_notes/v1.63.11-stable/index.md +++ b/docs/my-website/release_notes/v1.63.11-stable/index.md @@ -34,6 +34,9 @@ This release will be live on 03/16/2025 +## Known Issues +- 🚨 Known issue on Azure OpenAI - We don't recommend upgrading if you use Azure OpenAI. This version failed our Azure OpenAI load test + ## Docker Run LiteLLM Proxy ``` From 26be805ad371f5f546990a7ab5fad7df74f78caf Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 09:25:26 -0700 Subject: [PATCH 03/33] rename to _get_azure_openai_client --- litellm/llms/azure/azure.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 7fba70141c2..74a6ab58d3c 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -141,7 +141,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): return headers - def _get_sync_azure_client( + def _get_azure_openai_client( self, api_version: Optional[str], api_base: Optional[str], @@ -1240,7 +1240,7 @@ 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, @@ -1279,7 +1279,7 @@ 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, From b74f3cb76c7710fb59109de24f70718ff8e35f2c Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 09:38:27 -0700 Subject: [PATCH 04/33] _get_azure_openai_client --- litellm/llms/azure/azure.py | 24 +++++++++++++++++++----- 1 file changed, 19 insertions(+), 5 deletions(-) diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 74a6ab58d3c..545fc85e294 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -331,6 +331,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): client=client, max_retries=max_retries, azure_client_params=azure_client_params, + litellm_params=litellm_params, ) else: return self.acompletion( @@ -349,6 +350,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): 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( @@ -460,15 +462,26 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): 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, + azure_ad_token=azure_ad_token, + azure_ad_token_provider=azure_ad_token_provider, + model=model, + max_retries=max_retries, + timeout=timeout, + client=client, + client_type="async", + 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"], @@ -622,6 +635,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): 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: From f2026ef907c06d94440930917add71314b901413 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 09:51:28 -0700 Subject: [PATCH 05/33] fix - correctly re-use azure openai client --- litellm/llms/azure/azure.py | 129 +++++++++++++----- .../azure_client_usage_test.py | 108 +++++++++++++++ 2 files changed, 203 insertions(+), 34 deletions(-) create mode 100644 tests/code_coverage_tests/azure_client_usage_test.py diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 545fc85e294..c42657d87b4 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -149,8 +149,8 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): azure_ad_token: Optional[str], azure_ad_token_provider: Optional[Callable], model: str, - max_retries: int, - timeout: Union[float, httpx.Timeout], + max_retries: Optional[int], + timeout: Optional[Union[float, httpx.Timeout]], client: Optional[Any], client_type: Literal["sync", "async"], litellm_params: Optional[dict] = None, @@ -366,6 +366,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): timeout=timeout, client=client, max_retries=max_retries, + litellm_params=litellm_params, ) else: ## LOGGING @@ -387,21 +388,19 @@ 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, + azure_ad_token=azure_ad_token, + azure_ad_token_provider=azure_ad_token_provider, + model=model, + max_retries=max_retries, + timeout=timeout, + client=client, + client_type="sync", + litellm_params=litellm_params, + ) if not isinstance(azure_client, AzureOpenAI): raise AzureOpenAIError( status_code=500, @@ -567,6 +566,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 = { @@ -589,10 +589,24 @@ 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, + azure_ad_token=azure_ad_token, + azure_ad_token_provider=azure_ad_token_provider, + model=model, + max_retries=max_retries, + timeout=timeout, + client=client, + client_type="sync", + 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"], @@ -638,10 +652,22 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): 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, + azure_ad_token=azure_ad_token, + azure_ad_token_provider=azure_ad_token_provider, + model=model, + max_retries=max_retries, + timeout=timeout, + client=client, + client_type="async", + 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"], @@ -692,6 +718,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): async def aembedding( self, + model: str, data: dict, model_response: EmbeddingResponse, azure_client_params: dict, @@ -699,15 +726,33 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): logging_obj: LiteLLMLoggingObj, api_key: Optional[str] = None, client: Optional[AsyncAzureOpenAI] = None, - timeout=None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + max_retries: Optional[int] = None, + api_version: Optional[str] = None, + api_base: Optional[str] = None, + azure_ad_token: Optional[str] = None, + azure_ad_token_provider: Optional[Callable] = None, + litellm_params: Optional[dict] = {}, ): 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, + azure_ad_token=azure_ad_token, + azure_ad_token_provider=azure_ad_token_provider, + model=model, + max_retries=max_retries, + timeout=timeout, + client=client, + client_type="async", + 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 ) @@ -799,11 +844,27 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): azure_client_params=azure_client_params, timeout=timeout, client=client, + litellm_params=litellm_params, ) - 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, + azure_ad_token=azure_ad_token, + azure_ad_token_provider=azure_ad_token_provider, + model=model, + max_retries=max_retries, + timeout=timeout, + client=client, + client_type="sync", + 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) 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() From edfbf21c39e0d5dec8dffd4226e5776900331f1a Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 10:06:56 -0700 Subject: [PATCH 06/33] fix re-using azure openai client --- litellm/llms/azure/azure.py | 99 +++++------------------------- litellm/llms/azure/common_utils.py | 14 +++-- 2 files changed, 26 insertions(+), 87 deletions(-) diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index c42657d87b4..6613ce57baa 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -141,41 +141,6 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): return headers - def _get_azure_openai_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: Optional[int], - timeout: Optional[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, @@ -388,17 +353,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): status_code=422, message="max retries must be an int" ) # init AzureOpenAI Client - azure_client = self._get_azure_openai_client( + azure_client = self.get_azure_openai_client( api_version=api_version, api_base=api_base, 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, client=client, - client_type="sync", + _is_async=False, litellm_params=litellm_params, ) if not isinstance(azure_client, AzureOpenAI): @@ -466,17 +427,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): response = None try: # setting Azure client - azure_client = self._get_azure_openai_client( + azure_client = self.get_azure_openai_client( api_version=api_version, api_base=api_base, 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, client=client, - client_type="async", + _is_async=True, litellm_params=litellm_params, ) if not isinstance(azure_client, AsyncAzureOpenAI): @@ -589,17 +546,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): elif azure_ad_token_provider is not None: azure_client_params["azure_ad_token_provider"] = azure_ad_token_provider - azure_client = self._get_azure_openai_client( + azure_client = self.get_azure_openai_client( api_version=api_version, api_base=api_base, 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, client=client, - client_type="sync", + _is_async=False, litellm_params=litellm_params, ) if not isinstance(azure_client, AzureOpenAI): @@ -652,17 +605,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): litellm_params: Optional[dict] = {}, ): try: - azure_client = self._get_azure_openai_client( + azure_client = self.get_azure_openai_client( api_version=api_version, api_base=api_base, 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, client=client, - client_type="async", + _is_async=True, litellm_params=litellm_params, ) if not isinstance(azure_client, AsyncAzureOpenAI): @@ -737,17 +686,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): response = None try: - openai_aclient = self._get_azure_openai_client( + openai_aclient = self.get_azure_openai_client( api_version=api_version, api_base=api_base, 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, ) if not isinstance(openai_aclient, AsyncAzureOpenAI): @@ -846,17 +791,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): client=client, litellm_params=litellm_params, ) - azure_client = self._get_azure_openai_client( + azure_client = self.get_azure_openai_client( api_version=api_version, api_base=api_base, 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, ) if not isinstance(azure_client, AzureOpenAI): @@ -1315,17 +1256,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): litellm_params=litellm_params, ) # type: ignore - azure_client: AzureOpenAI = self._get_azure_openai_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 @@ -1354,17 +1291,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): litellm_params: Optional[dict] = None, ) -> HttpxBinaryResponseContent: - azure_client: AsyncAzureOpenAI = self._get_azure_openai_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..24eb758653a 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -247,20 +247,21 @@ def select_azure_base_url_or_endpoint(azure_client_params: dict): class BaseAzureLLM: 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 if client is None: 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, ) if _is_async is True: @@ -269,6 +270,11 @@ 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) return openai_client @@ -277,7 +283,7 @@ class BaseAzureLLM: litellm_params: dict, api_key: Optional[str], api_base: Optional[str], - model_name: str, + model_name: Optional[str], api_version: Optional[str], ) -> dict: From 34142a1b62f2e10dc0d25477ad9517f176db70b5 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 10:11:54 -0700 Subject: [PATCH 07/33] _init_azure_client_for_cloudflare_ai_gateway --- litellm/llms/azure/azure.py | 43 +++++++++--------------------- litellm/llms/azure/common_utils.py | 42 +++++++++++++++++++++++++++++ 2 files changed, 54 insertions(+), 31 deletions(-) diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 6613ce57baa..6a16b50c31b 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -238,37 +238,18 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): ### 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: diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index 24eb758653a..ac500fc2d60 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -357,3 +357,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 = { + "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) + return client From 0601768bb86980f0df33cd870642cd5a2cd7823c Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 10:14:51 -0700 Subject: [PATCH 08/33] use ssl on initialize_azure_sdk_client --- litellm/llms/azure/common_utils.py | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index ac500fc2d60..a7a1954e8fe 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -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.openai import OpenAIChatCompletion from litellm.secret_managers.get_azure_ad_token_provider import ( get_azure_ad_token_provider, ) @@ -244,7 +245,7 @@ def select_azure_base_url_or_endpoint(azure_client_params: dict): return azure_client_params -class BaseAzureLLM: +class BaseAzureLLM(OpenAIChatCompletion): def get_azure_openai_client( self, api_key: Optional[str], @@ -263,6 +264,7 @@ class BaseAzureLLM: api_base=api_base, model_name=model, api_version=api_version, + is_async=_is_async, ) if _is_async is True: openai_client = AsyncAzureOpenAI(**azure_client_params) @@ -285,6 +287,7 @@ class BaseAzureLLM: api_base: Optional[str], model_name: Optional[str], api_version: Optional[str], + is_async: bool, ) -> dict: azure_ad_token_provider: Optional[Callable[[], str]] = None @@ -340,8 +343,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: From a0c5fb81b8c449ddbab807aacac2040b51b6654c Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 10:23:30 -0700 Subject: [PATCH 09/33] fix logic for intializing openai clients --- litellm/llms/azure/azure.py | 3 +++ litellm/llms/azure/common_utils.py | 4 ++-- litellm/llms/openai/common_utils.py | 29 ++++++++++++++++++++++++++++ litellm/llms/openai/openai.py | 30 ++++++----------------------- 4 files changed, 40 insertions(+), 26 deletions(-) diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 6a16b50c31b..94e3d4732a5 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -234,6 +234,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): api_base=api_base, model_name=model, api_version=api_version, + is_async=False, ) ### CHECK IF CLOUDFLARE AI GATEWAY ### ### if so - set the model as part of the base url @@ -749,6 +750,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): model_name=model, api_version=api_version, api_base=api_base, + is_async=False, ) ## LOGGING logging_obj.pre_call( @@ -1152,6 +1154,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 diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index a7a1954e8fe..fa41b079734 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -9,7 +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.openai import OpenAIChatCompletion +from litellm.llms.openai.common_utils import BaseOpenAILLM from litellm.secret_managers.get_azure_ad_token_provider import ( get_azure_ad_token_provider, ) @@ -245,7 +245,7 @@ def select_azure_base_url_or_endpoint(azure_client_params: dict): return azure_client_params -class BaseAzureLLM(OpenAIChatCompletion): +class BaseAzureLLM(BaseOpenAILLM): def get_azure_openai_client( self, api_key: Optional[str], diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index a8412f867b5..99a6c8837a5 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -8,6 +8,7 @@ from typing import Any, Dict, List, Optional, Union import httpx import openai +import litellm from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -92,3 +93,31 @@ 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_async_http_client() -> Optional[httpx.AsyncClient]: + if litellm.ssl_verify: + return httpx.AsyncClient( + limits=httpx.Limits( + max_connections=1000, max_keepalive_connections=100 + ), + verify=litellm.ssl_verify, + ) + return litellm.aclient_session + + @staticmethod + def _get_sync_http_client() -> Optional[httpx.Client]: + if litellm.ssl_verify: + return httpx.Client( + limits=httpx.Limits( + max_connections=1000, max_keepalive_connections=100 + ), + verify=litellm.ssl_verify, + ) + return litellm.client_session diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index a5e33795c82..8045b5ef32d 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -50,7 +50,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 +321,7 @@ class OpenAIChatCompletionResponseIterator(BaseModelResponseIterator): raise e -class OpenAIChatCompletion(BaseLLM): +class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): def __init__(self) -> None: super().__init__() @@ -401,28 +405,6 @@ class OpenAIChatCompletion(BaseLLM): ) return client - @staticmethod - def _get_async_http_client() -> Optional[httpx.AsyncClient]: - if litellm.ssl_verify: - return httpx.AsyncClient( - limits=httpx.Limits( - max_connections=1000, max_keepalive_connections=100 - ), - verify=litellm.ssl_verify, - ) - return litellm.aclient_session - - @staticmethod - def _get_sync_http_client() -> Optional[httpx.Client]: - if litellm.ssl_verify: - return httpx.Client( - limits=httpx.Limits( - max_connections=1000, max_keepalive_connections=100 - ), - verify=litellm.ssl_verify, - ) - return litellm.client_session - @track_llm_api_timing() async def make_openai_chat_completion_request( self, From e34be5a3b6feb7009fb81c8058d774260e6abf1b Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 10:28:39 -0700 Subject: [PATCH 10/33] use get_azure_openai_client --- litellm/llms/azure/completion/handler.py | 132 +++++++++++++---------- 1 file changed, 75 insertions(+), 57 deletions(-) diff --git a/litellm/llms/azure/completion/handler.py b/litellm/llms/azure/completion/handler.py index 4ec5c435dac..91a00ebc2f8 100644 --- a/litellm/llms/azure/completion/handler.py +++ b/litellm/llms/azure/completion/handler.py @@ -76,33 +76,25 @@ class AzureTextCompletion(BaseAzureLLM): model_name=model, api_version=api_version, api_base=api_base, + is_async=False, ) ### 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: @@ -174,17 +166,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 @@ -234,20 +230,27 @@ class AzureTextCompletion(BaseAzureLLM): 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"], @@ -291,6 +294,7 @@ class AzureTextCompletion(BaseAzureLLM): 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 +302,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"], @@ -340,18 +352,24 @@ class AzureTextCompletion(BaseAzureLLM): 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"], From c1e0cb136e04c70c5baa4a622a2334f5b5a3ba77 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 10:47:29 -0700 Subject: [PATCH 11/33] fix using azure openai clients --- litellm/llms/azure/audio_transcriptions.py | 52 +++++++++++++--------- 1 file changed, 32 insertions(+), 20 deletions(-) diff --git a/litellm/llms/azure/audio_transcriptions.py b/litellm/llms/azure/audio_transcriptions.py index 52a3d780fbd..6daeff75e35 100644 --- a/litellm/llms/azure/audio_transcriptions.py +++ b/litellm/llms/azure/audio_transcriptions.py @@ -14,6 +14,7 @@ from litellm.utils import ( ) from .azure import AzureChatCompletion +from .common_utils import AzureOpenAIError class AzureAudioTranscription(AzureChatCompletion): @@ -36,15 +37,6 @@ class AzureAudioTranscription(AzureChatCompletion): ) -> 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 audio_file=audio_file, @@ -54,14 +46,24 @@ 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, ) - if client is None: - azure_client = AzureOpenAI(http_client=litellm.client_session, **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", + ) ## LOGGING logging_obj.pre_call( @@ -98,24 +100,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, ): 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( From 3458c69eb0b66dafc8435d8a8413243fd31eb110 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 11:04:02 -0700 Subject: [PATCH 12/33] fix common utils --- litellm/llms/azure/common_utils.py | 4 ++-- litellm/proxy/proxy_config.yaml | 7 +++++-- 2 files changed, 7 insertions(+), 4 deletions(-) diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index fa41b079734..9933ce50b6a 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -346,9 +346,9 @@ class BaseAzureLLM(BaseOpenAILLM): } # init http client + SSL Verification settings if is_async is True: - azure_client_params["http_client"] = self._get_async_http_client + azure_client_params["http_client"] = self._get_async_http_client() else: - azure_client_params["http_client"] = self._get_sync_http_client + azure_client_params["http_client"] = self._get_sync_http_client() if max_retries is not None: azure_client_params["max_retries"] = max_retries 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 From dfd7a7d547fef16cc7c681dc1139b25532ba9840 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 11:38:31 -0700 Subject: [PATCH 13/33] fix linting error --- litellm/llms/azure/assistants.py | 2 ++ litellm/llms/azure/common_utils.py | 10 +++++----- 2 files changed, 7 insertions(+), 5 deletions(-) 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/common_utils.py b/litellm/llms/azure/common_utils.py index 9933ce50b6a..34cca8fc8a2 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 @@ -385,7 +385,7 @@ class BaseAzureLLM(BaseOpenAILLM): api_base += "/" api_base += f"{model}" - azure_client_params = { + azure_client_params: Dict[str, Any] = { "api_version": api_version, "base_url": f"{api_base}", "http_client": litellm.client_session, @@ -399,11 +399,11 @@ class BaseAzureLLM(BaseOpenAILLM): 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: + 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) + client = AsyncAzureOpenAI(**azure_client_params) # type: ignore else: - client = AzureOpenAI(**azure_client_params) + client = AzureOpenAI(**azure_client_params) # type: ignore return client From 38e2dd00ccc508c05970742a0d306e62b66f6b40 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 11:42:11 -0700 Subject: [PATCH 14/33] fix amebedding issue on ssl azure --- litellm/llms/azure/azure.py | 1 + 1 file changed, 1 insertion(+) diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 94e3d4732a5..6e4b87c7321 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -766,6 +766,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): return self.aembedding( # type: ignore data=data, input=input, + model=model, logging_obj=logging_obj, api_key=api_key, model_response=model_response, From d4b3082ca24b343f4c3aeeece8c62d71932548da Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 12:19:12 -0700 Subject: [PATCH 15/33] fix azure embedding test --- litellm/llms/azure/azure.py | 28 ++++++++++++++-------------- 1 file changed, 14 insertions(+), 14 deletions(-) diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 6e4b87c7321..e9bb5733ffd 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 @@ -655,16 +655,16 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): 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: Optional[Union[float, httpx.Timeout]] = None, max_retries: Optional[int] = None, - api_version: Optional[str] = None, - api_base: Optional[str] = None, azure_ad_token: Optional[str] = None, azure_ad_token_provider: Optional[Callable] = None, litellm_params: Optional[dict] = {}, - ): + ) -> EmbeddingResponse: response = None try: @@ -693,13 +693,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( @@ -728,7 +734,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: @@ -737,13 +743,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, @@ -763,7 +762,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): ) if aembedding is True: - return self.aembedding( # type: ignore + return self.aembedding( data=data, input=input, model=model, @@ -774,6 +773,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): timeout=timeout, client=client, litellm_params=litellm_params, + api_base=api_base, ) azure_client = self.get_azure_openai_client( api_version=api_version, From 842625a6f092bfc0929a8289f2d97c5bed20c028 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 12:25:32 -0700 Subject: [PATCH 16/33] :test_completion_azure_ad_toke --- litellm/llms/openai/common_utils.py | 29 +++++++++++++---------------- 1 file changed, 13 insertions(+), 16 deletions(-) diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index 99a6c8837a5..649ce2e0f15 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -102,22 +102,19 @@ class BaseOpenAILLM: @staticmethod def _get_async_http_client() -> Optional[httpx.AsyncClient]: - if litellm.ssl_verify: - return httpx.AsyncClient( - limits=httpx.Limits( - max_connections=1000, max_keepalive_connections=100 - ), - verify=litellm.ssl_verify, - ) - return litellm.aclient_session + 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.ssl_verify: - return httpx.Client( - limits=httpx.Limits( - max_connections=1000, max_keepalive_connections=100 - ), - verify=litellm.ssl_verify, - ) - return litellm.client_session + 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, + ) From 6987a73e36df84fec6dff8a6a2a80523cf83be4b Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 12:27:17 -0700 Subject: [PATCH 17/33] initialize_azure_sdk_client --- tests/litellm/llms/azure/test_azure_common_utils.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/tests/litellm/llms/azure/test_azure_common_utils.py b/tests/litellm/llms/azure/test_azure_common_utils.py index 21fa3b37eee..f29cc795db6 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 From b316911120594c0801e1405947bed6fb059c2e22 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 12:31:44 -0700 Subject: [PATCH 18/33] fix typing errors --- litellm/main.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/litellm/main.py b/litellm/main.py index 64049c31d12..85aa0e96a96 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: From 80a5cfa01dabbadc4315e90f521ef98bb23bae88 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 12:35:34 -0700 Subject: [PATCH 19/33] test_azure_embedding_max_retries_0 --- litellm/llms/azure/azure.py | 10 ---------- litellm/llms/azure/completion/handler.py | 12 ------------ 2 files changed, 22 deletions(-) diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index e9bb5733ffd..172c963acba 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -652,7 +652,6 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): model: str, data: dict, model_response: EmbeddingResponse, - azure_client_params: dict, input: list, logging_obj: LiteLLMLoggingObj, api_base: str, @@ -743,14 +742,6 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): data = {"model": model, "input": input, **optional_params} if max_retries is None: max_retries = litellm.DEFAULT_MAX_RETRIES - 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, - is_async=False, - ) ## LOGGING logging_obj.pre_call( input=input, @@ -769,7 +760,6 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): 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, diff --git a/litellm/llms/azure/completion/handler.py b/litellm/llms/azure/completion/handler.py index 91a00ebc2f8..1bc9aaba9bf 100644 --- a/litellm/llms/azure/completion/handler.py +++ b/litellm/llms/azure/completion/handler.py @@ -12,18 +12,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__() From b60178f5344df89cf144ef6f119eaf56e93d294f Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 12:42:24 -0700 Subject: [PATCH 20/33] fix azure chat logic --- litellm/llms/azure/azure.py | 12 ------------ 1 file changed, 12 deletions(-) diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 172c963acba..03c5cc09ebe 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -228,14 +228,6 @@ 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, - is_async=False, - ) ### CHECK IF CLOUDFLARE AI GATEWAY ### ### if so - set the model as part of the base url if "gateway.ai.cloudflare.com" in api_base: @@ -277,7 +269,6 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): timeout=timeout, client=client, max_retries=max_retries, - azure_client_params=azure_client_params, litellm_params=litellm_params, ) else: @@ -296,7 +287,6 @@ 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: @@ -403,7 +393,6 @@ 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 @@ -583,7 +572,6 @@ 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: From 2cd49ef0969931db9dc6d8fb386b4a573cbc0d8e Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 12:46:55 -0700 Subject: [PATCH 21/33] fix test_ensure_initialize_azure_sdk_client_always_used --- litellm/llms/azure/completion/handler.py | 15 --------------- 1 file changed, 15 deletions(-) diff --git a/litellm/llms/azure/completion/handler.py b/litellm/llms/azure/completion/handler.py index 1bc9aaba9bf..9e4fcabb6b2 100644 --- a/litellm/llms/azure/completion/handler.py +++ b/litellm/llms/azure/completion/handler.py @@ -58,15 +58,6 @@ 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, - is_async=False, - ) - ### CHECK IF CLOUDFLARE AI GATEWAY ### ### if so - set the model as part of the base url if "gateway.ai.cloudflare.com" in api_base: @@ -104,7 +95,6 @@ class AzureTextCompletion(BaseAzureLLM): azure_ad_token=azure_ad_token, timeout=timeout, client=client, - azure_client_params=azure_client_params, ) else: return self.acompletion( @@ -119,7 +109,6 @@ class AzureTextCompletion(BaseAzureLLM): client=client, logging_obj=logging_obj, max_retries=max_retries, - azure_client_params=azure_client_params, ) elif "stream" in optional_params and optional_params["stream"] is True: return self.streaming( @@ -132,7 +121,6 @@ class AzureTextCompletion(BaseAzureLLM): azure_ad_token=azure_ad_token, timeout=timeout, client=client, - azure_client_params=azure_client_params, ) else: ## LOGGING @@ -217,7 +205,6 @@ 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 @@ -281,7 +268,6 @@ 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) @@ -339,7 +325,6 @@ class AzureTextCompletion(BaseAzureLLM): timeout: Any, azure_ad_token: Optional[str] = None, client=None, - azure_client_params: dict = {}, litellm_params: dict = {}, ): try: From dc3d7b3afc6ebdfc6a580cc8ee901bf3ecf57a0d Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 12:56:11 -0700 Subject: [PATCH 22/33] test_azure_instruct --- litellm/llms/azure/completion/handler.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/litellm/llms/azure/completion/handler.py b/litellm/llms/azure/completion/handler.py index 9e4fcabb6b2..67088a663af 100644 --- a/litellm/llms/azure/completion/handler.py +++ b/litellm/llms/azure/completion/handler.py @@ -95,6 +95,7 @@ class AzureTextCompletion(BaseAzureLLM): azure_ad_token=azure_ad_token, timeout=timeout, client=client, + litellm_params=litellm_params, ) else: return self.acompletion( @@ -109,6 +110,7 @@ class AzureTextCompletion(BaseAzureLLM): client=client, logging_obj=logging_obj, max_retries=max_retries, + litellm_params=litellm_params, ) elif "stream" in optional_params and optional_params["stream"] is True: return self.streaming( From b20a69f9fc4d355f7d2887863c80c5c8e53d362e Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 12:58:59 -0700 Subject: [PATCH 23/33] fix code quality --- litellm/llms/azure/audio_transcriptions.py | 1 - litellm/llms/azure/completion/handler.py | 1 - 2 files changed, 2 deletions(-) diff --git a/litellm/llms/azure/audio_transcriptions.py b/litellm/llms/azure/audio_transcriptions.py index 6daeff75e35..ffe43a40052 100644 --- a/litellm/llms/azure/audio_transcriptions.py +++ b/litellm/llms/azure/audio_transcriptions.py @@ -4,7 +4,6 @@ from typing import Any, Optional 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 ( diff --git a/litellm/llms/azure/completion/handler.py b/litellm/llms/azure/completion/handler.py index 67088a663af..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 From 7384d45ef08973a90381fe2a8242965c12f167ca Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 14:22:30 -0700 Subject: [PATCH 24/33] fix type errors on transcription azure --- litellm/main.py | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/litellm/main.py b/litellm/main.py index 85aa0e96a96..e75c23f0fc1 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -4951,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" @@ -4984,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. @@ -5053,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") From 55ea2370ba3d114c515c4a4107f9dbe13d48f245 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 14:23:14 -0700 Subject: [PATCH 25/33] Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]]: --- litellm/llms/azure/audio_transcriptions.py | 17 ++++++++++++----- .../llms/azure/test_azure_common_utils.py | 10 ++++------ 2 files changed, 16 insertions(+), 11 deletions(-) diff --git a/litellm/llms/azure/audio_transcriptions.py b/litellm/llms/azure/audio_transcriptions.py index ffe43a40052..be7d0fa30da 100644 --- a/litellm/llms/azure/audio_transcriptions.py +++ b/litellm/llms/azure/audio_transcriptions.py @@ -1,5 +1,5 @@ import uuid -from typing import Any, Optional +from typing import Any, Coroutine, Optional, Union from openai import AsyncAzureOpenAI, AzureOpenAI from pydantic import BaseModel @@ -33,11 +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} 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, @@ -47,6 +47,8 @@ class AzureAudioTranscription(AzureChatCompletion): client=client, max_retries=max_retries, logging_obj=logging_obj, + model=model, + litellm_params=litellm_params, ) azure_client = self.get_azure_openai_client( @@ -110,7 +112,7 @@ class AzureAudioTranscription(AzureChatCompletion): client=None, max_retries=None, litellm_params: Optional[dict] = None, - ): + ) -> TranscriptionResponse: response = None try: async_azure_client = self.get_azure_openai_client( @@ -179,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/tests/litellm/llms/azure/test_azure_common_utils.py b/tests/litellm/llms/azure/test_azure_common_utils.py index f29cc795db6..4d24009685b 100644 --- a/tests/litellm/llms/azure/test_azure_common_utils.py +++ b/tests/litellm/llms/azure/test_azure_common_utils.py @@ -308,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" ) From c010cdef5912a7538c04302eebb44dbe1886c08f Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 17:26:23 -0700 Subject: [PATCH 26/33] test_dynamic_azure_params --- tests/local_testing/test_completion.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) 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 From bdf77f6f4bbbc009b19c80b8fac5f033eccd1569 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 17:29:10 -0700 Subject: [PATCH 27/33] fix ensure async client test --- tests/code_coverage_tests/ensure_async_clients_test.py | 2 ++ 1 file changed, 2 insertions(+) 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" From f73e9047dc07e4eb9e05f7b420c8b9736c0b4424 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 17:56:32 -0700 Subject: [PATCH 28/33] use common logic for re-using openai clients --- litellm/llms/azure/common_utils.py | 16 ++++++ litellm/llms/openai/common_utils.py | 80 ++++++++++++++++++++++++++++- 2 files changed, 95 insertions(+), 1 deletion(-) diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index 34cca8fc8a2..4d9c35a5fb8 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -257,7 +257,17 @@ class BaseAzureLLM(BaseOpenAILLM): 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 and isinstance( + cached_client, (AzureOpenAI, AsyncAzureOpenAI) + ): + return cached_client + azure_client_params = self.initialize_azure_sdk_client( litellm_params=litellm_params or {}, api_key=api_key, @@ -278,6 +288,12 @@ class BaseAzureLLM(BaseOpenAILLM): # 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( diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index 649ce2e0f15..ac84fbacf1b 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -2,14 +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): @@ -100,6 +103,81 @@ 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')}", + ] + + for param in BaseOpenAILLM.get_openai_client_initialization_param_fields( + client_type=client_type + ): + 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: From a45830dac35c3813c593cfb291d0e7da8e215f59 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 17:57:03 -0700 Subject: [PATCH 29/33] use common caching logic for openai/azure clients --- litellm/llms/openai/openai.py | 30 +++++++++++------------------- 1 file changed, 11 insertions(+), 19 deletions(-) diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 8045b5ef32d..475c83be34b 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, @@ -348,6 +347,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): organization: Optional[str] = None, client: Optional[Union[OpenAI, AsyncOpenAI]] = None, ): + client_initialization_params: Dict = locals() if client is None: if not isinstance(max_retries, int): raise OpenAIError( @@ -356,20 +356,12 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): 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() - - _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 + cached_client = self.get_cached_openai_client( + client_initialization_params=client_initialization_params, + client_type="openai", + ) + if cached_client: + return cached_client if is_async: _new_client: Union[OpenAI, AsyncOpenAI] = AsyncOpenAI( api_key=api_key, @@ -390,10 +382,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): ) ## 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 From 3daef0d7402e0597c7ace0ea3fafee861095793a Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 17:59:46 -0700 Subject: [PATCH 30/33] fix common utils --- litellm/llms/openai/common_utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index ac84fbacf1b..f9ba366cb5d 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -163,7 +163,7 @@ class BaseOpenAILLM: @staticmethod def get_openai_client_initialization_param_fields( client_type: Literal["openai", "azure"] - ) -> list[str]: + ) -> List[str]: """Returns a list of fields that are used to initialize the OpenAI client""" import inspect From d5150e000dea40628e524783594da8907d45a51d Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 18:08:24 -0700 Subject: [PATCH 31/33] test openai common utils --- .../llms/openai/test_openai_common_utils.py | 226 ++++++++++++++++++ 1 file changed, 226 insertions(+) create mode 100644 tests/litellm/llms/openai/test_openai_common_utils.py 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..5dbab868034 --- /dev/null +++ b/tests/litellm/llms/openai/test_openai_common_utils.py @@ -0,0 +1,226 @@ +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 + + +def test_openai_client_reuse_sync(): + """ + Test that multiple synchronous completion calls reuse the same OpenAI client + """ + litellm.set_verbose = True + + # Mock the OpenAI client creation to track how many times it's called + with patch("litellm.llms.openai.openai.OpenAI") as mock_openai, 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 completion calls + for _ in range(10): + try: + litellm.completion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Hello"}], + max_tokens=10, + ) + except Exception: + # We expect exceptions since we're mocking the client + pass + + # Verify OpenAI client was created only once + assert mock_openai.call_count == 1, "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" + + +@pytest.mark.asyncio +async def test_openai_client_reuse_async(): + """ + Test that multiple asynchronous completion calls reuse the same OpenAI client + """ + litellm.set_verbose = True + + # Mock the AsyncOpenAI client creation to track how many times it's called + with patch( + "litellm.llms.openai.openai.AsyncOpenAI" + ) as mock_async_openai, 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 async completion calls + for _ in range(10): + try: + await litellm.acompletion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Hello"}], + max_tokens=10, + ) + except Exception: + # We expect exceptions since we're mocking the client + pass + + # Verify AsyncOpenAI client was created only once + assert ( + mock_async_openai.call_count == 1 + ), "AsyncOpenAI 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" + + +@pytest.mark.asyncio +async def test_openai_client_reuse_streaming(): + """ + Test that multiple streaming completion calls reuse the same OpenAI client + """ + litellm.set_verbose = True + + # Mock the AsyncOpenAI client creation to track how many times it's called + with patch( + "litellm.llms.openai.openai.AsyncOpenAI" + ) as mock_async_openai, 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 streaming completion calls + for _ in range(10): + try: + await litellm.acompletion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Hello"}], + max_tokens=10, + stream=True, + ) + except Exception: + # We expect exceptions since we're mocking the client + pass + + # Verify AsyncOpenAI client was created only once + assert ( + mock_async_openai.call_count == 1 + ), "AsyncOpenAI 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" + + +def test_openai_client_reuse_with_different_params(): + """ + Test that different client parameters create different cached clients + """ + litellm.set_verbose = True + + # Mock the OpenAI client creation + with patch("litellm.llms.openai.openai.OpenAI") as mock_openai, patch.object( + BaseOpenAILLM, "set_cached_openai_client" + ) as mock_set_cache, patch.object( + BaseOpenAILLM, "get_cached_openai_client", return_value=None + ) as mock_get_cache: + + # Make calls with different API keys + try: + litellm.completion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Hello"}], + api_key="test_key_1", + ) + except Exception: + pass + + try: + litellm.completion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Hello"}], + api_key="test_key_2", + ) + except Exception: + pass + + # Verify OpenAI client was created twice (different API keys) + assert ( + mock_openai.call_count == 2 + ), "Different API keys should create different clients" + + # Verify the clients were cached + assert mock_set_cache.call_count == 2, "Both clients should be cached" + + # Verify we tried to get from cache twice + assert mock_get_cache.call_count == 2, "Should check cache for each request" + + +def test_openai_client_reuse_with_custom_client(): + """ + Test that when a custom client is provided, it's used directly without caching + """ + litellm.set_verbose = True + + # Create a mock custom client + custom_client = MagicMock() + + # Mock the cache functions + with patch.object( + BaseOpenAILLM, "set_cached_openai_client" + ) as mock_set_cache, patch.object( + BaseOpenAILLM, "get_cached_openai_client" + ) as mock_get_cache: + + # Make multiple calls with the custom client + for _ in range(5): + try: + litellm.completion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Hello"}], + client=custom_client, + ) + except Exception: + pass + + # Verify we never tried to cache the client + assert mock_set_cache.call_count == 0, "Custom client should not be cached" + + # Verify we never tried to get from cache + assert ( + mock_get_cache.call_count == 0 + ), "Should not check cache when custom client is provided" From 40418c7bd89c55fc74676d28d28932a0a746c9d6 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 18:13:36 -0700 Subject: [PATCH 32/33] test_openai_client_reuse --- .../llms/openai/test_openai_common_utils.py | 268 ++++++------------ 1 file changed, 87 insertions(+), 181 deletions(-) diff --git a/tests/litellm/llms/openai/test_openai_common_utils.py b/tests/litellm/llms/openai/test_openai_common_utils.py index 5dbab868034..a343fcf25c5 100644 --- a/tests/litellm/llms/openai/test_openai_common_utils.py +++ b/tests/litellm/llms/openai/test_openai_common_utils.py @@ -11,59 +11,89 @@ sys.path.insert( import litellm from litellm.llms.openai.common_utils import BaseOpenAILLM - -def test_openai_client_reuse_sync(): - """ - Test that multiple synchronous completion calls reuse the same OpenAI client - """ - litellm.set_verbose = True - - # Mock the OpenAI client creation to track how many times it's called - with patch("litellm.llms.openai.openai.OpenAI") as mock_openai, 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 completion calls - for _ in range(10): - try: - litellm.completion( - model="gpt-4o-mini", - messages=[{"role": "user", "content": "Hello"}], - max_tokens=10, - ) - except Exception: - # We expect exceptions since we're mocking the client - pass - - # Verify OpenAI client was created only once - assert mock_openai.call_count == 1, "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" +# 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_async(): +async def test_openai_client_reuse(function_name, is_async, args): """ - Test that multiple asynchronous completion calls reuse the same OpenAI client + Test that multiple API calls reuse the same OpenAI client """ litellm.set_verbose = True - # Mock the AsyncOpenAI client creation to track how many times it's called - with patch( + # Determine which client class to mock based on whether the test is async + client_path = ( "litellm.llms.openai.openai.AsyncOpenAI" - ) as mock_async_openai, patch.object( + 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" @@ -75,152 +105,28 @@ async def test_openai_client_reuse_async(): mock_client ] * 9 # First call returns None, rest return the mock client - # Make 10 async completion calls + # Make 10 API calls for _ in range(10): try: - await litellm.acompletion( - model="gpt-4o-mini", - messages=[{"role": "user", "content": "Hello"}], - max_tokens=10, - ) + # 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 AsyncOpenAI client was created only once + # Verify client was created only once assert ( - mock_async_openai.call_count == 1 - ), "AsyncOpenAI client should be created only once" + 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" - - -@pytest.mark.asyncio -async def test_openai_client_reuse_streaming(): - """ - Test that multiple streaming completion calls reuse the same OpenAI client - """ - litellm.set_verbose = True - - # Mock the AsyncOpenAI client creation to track how many times it's called - with patch( - "litellm.llms.openai.openai.AsyncOpenAI" - ) as mock_async_openai, 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 streaming completion calls - for _ in range(10): - try: - await litellm.acompletion( - model="gpt-4o-mini", - messages=[{"role": "user", "content": "Hello"}], - max_tokens=10, - stream=True, - ) - except Exception: - # We expect exceptions since we're mocking the client - pass - - # Verify AsyncOpenAI client was created only once - assert ( - mock_async_openai.call_count == 1 - ), "AsyncOpenAI 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" - - -def test_openai_client_reuse_with_different_params(): - """ - Test that different client parameters create different cached clients - """ - litellm.set_verbose = True - - # Mock the OpenAI client creation - with patch("litellm.llms.openai.openai.OpenAI") as mock_openai, patch.object( - BaseOpenAILLM, "set_cached_openai_client" - ) as mock_set_cache, patch.object( - BaseOpenAILLM, "get_cached_openai_client", return_value=None - ) as mock_get_cache: - - # Make calls with different API keys - try: - litellm.completion( - model="gpt-4o-mini", - messages=[{"role": "user", "content": "Hello"}], - api_key="test_key_1", - ) - except Exception: - pass - - try: - litellm.completion( - model="gpt-4o-mini", - messages=[{"role": "user", "content": "Hello"}], - api_key="test_key_2", - ) - except Exception: - pass - - # Verify OpenAI client was created twice (different API keys) - assert ( - mock_openai.call_count == 2 - ), "Different API keys should create different clients" - - # Verify the clients were cached - assert mock_set_cache.call_count == 2, "Both clients should be cached" - - # Verify we tried to get from cache twice - assert mock_get_cache.call_count == 2, "Should check cache for each request" - - -def test_openai_client_reuse_with_custom_client(): - """ - Test that when a custom client is provided, it's used directly without caching - """ - litellm.set_verbose = True - - # Create a mock custom client - custom_client = MagicMock() - - # Mock the cache functions - with patch.object( - BaseOpenAILLM, "set_cached_openai_client" - ) as mock_set_cache, patch.object( - BaseOpenAILLM, "get_cached_openai_client" - ) as mock_get_cache: - - # Make multiple calls with the custom client - for _ in range(5): - try: - litellm.completion( - model="gpt-4o-mini", - messages=[{"role": "user", "content": "Hello"}], - client=custom_client, - ) - except Exception: - pass - - # Verify we never tried to cache the client - assert mock_set_cache.call_count == 0, "Custom client should not be cached" - - # Verify we never tried to get from cache - assert ( - mock_get_cache.call_count == 0 - ), "Should not check cache when custom client is provided" From 65083ca8da81c02c0c06bde6f9d0d2ef29a4a7c0 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 18 Mar 2025 18:35:50 -0700 Subject: [PATCH 33/33] get_openai_client_cache_key --- litellm/llms/azure/common_utils.py | 9 +- litellm/llms/openai/common_utils.py | 18 +- litellm/llms/openai/openai.py | 8 +- .../llms/azure/test_azure_common_utils.py | 174 ++++++++++++++++++ 4 files changed, 199 insertions(+), 10 deletions(-) diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index 4d9c35a5fb8..71092c8b993 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -263,10 +263,11 @@ class BaseAzureLLM(BaseOpenAILLM): client_initialization_params=client_initialization_params, client_type="azure", ) - if cached_client and isinstance( - cached_client, (AzureOpenAI, AsyncAzureOpenAI) - ): - return cached_client + 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 or {}, diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index f9ba366cb5d..55da16d6cd0 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -151,13 +151,23 @@ class BaseOpenAILLM: f"is_async={client_initialization_params.get('is_async')}", ] - for param in BaseOpenAILLM.get_openai_client_initialization_param_fields( - client_type=client_type - ): + 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 diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 475c83be34b..98ef95239e0 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -346,7 +346,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): 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): @@ -360,8 +360,12 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): client_initialization_params=client_initialization_params, client_type="openai", ) + if cached_client: - return 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, diff --git a/tests/litellm/llms/azure/test_azure_common_utils.py b/tests/litellm/llms/azure/test_azure_common_utils.py index 4d24009685b..a9e63f84f24 100644 --- a/tests/litellm/llms/azure/test_azure_common_utils.py +++ b/tests/litellm/llms/azure/test_azure_common_utils.py @@ -461,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"