diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index 7436a0e1b00..c26c1df821d 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -491,6 +491,10 @@ class BaseAzureLLM(BaseOpenAILLM): client_initialization_params: Final[dict] = locals() client_initialization_params["is_async"] = _is_async _lp: Final = litellm_params or {} + configured_max_retries: Final = _lp.get("max_retries") + client_initialization_params["max_retries"] = ( + DEFAULT_MAX_RETRIES if configured_max_retries is None else configured_max_retries + ) _ad_provider: Final = _lp.get("azure_ad_token_provider") _ad_token: Final = _lp.get("azure_ad_token") _client_secret: Final = _lp.get("client_secret") diff --git a/tests/unit/llms/azure/test_azure_common_utils.py b/tests/unit/llms/azure/test_azure_common_utils.py index caf941ebd19..42aced3cd57 100644 --- a/tests/unit/llms/azure/test_azure_common_utils.py +++ b/tests/unit/llms/azure/test_azure_common_utils.py @@ -1,7 +1,7 @@ import json import os import traceback -from typing import Callable, Optional +from typing import Callable, Final, Optional from unittest.mock import MagicMock, patch import pytest @@ -440,6 +440,66 @@ def test_default_max_retries_env_var_reaches_azure_sdk_client(): assert completed.stdout.strip() == "0" +@pytest.mark.asyncio +@pytest.mark.parametrize("is_async", [False, True]) +@pytest.mark.parametrize("api_version", ["2024-02-01", "v1"]) +@pytest.mark.parametrize( + "retry_params, expected_retries, reuse_first", + [ + pytest.param( + ({"max_retries": 3}, {"max_retries": 0}, {"max_retries": 3}), + (3, 0, 3), + False, + id="disable-retries", + ), + pytest.param( + ({"max_retries": 0}, {"max_retries": 3}, {"max_retries": 0}), + (0, 3, 0), + False, + id="enable-retries", + ), + pytest.param( + ({}, {"max_retries": None}, {"max_retries": litellm.constants.DEFAULT_MAX_RETRIES}), + (litellm.constants.DEFAULT_MAX_RETRIES,) * 3, + True, + id="equivalent-defaults", + ), + ], +) +async def test_azure_client_cache_respects_effective_max_retries( + monkeypatch: pytest.MonkeyPatch, + is_async: bool, + api_version: str, + retry_params: tuple[dict[str, int | None], ...], + expected_retries: tuple[int, ...], + reuse_first: bool, +) -> None: + import httpx + + from litellm.caching.llm_caching_handler import LLMClientCache + + monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache()) + async with httpx.AsyncClient() as async_session: + with httpx.Client() as sync_session: + monkeypatch.setattr(litellm, "aclient_session", async_session) + monkeypatch.setattr(litellm, "client_session", sync_session) + clients: Final = tuple( + BaseAzureLLM().get_azure_openai_client( + api_key="test-api-key", + api_base="https://test.openai.azure.com", + api_version=api_version, + litellm_params=params, + _is_async=is_async, + ) + for params in retry_params + ) + + assert all(client is not None for client in clients) + assert tuple(client.max_retries for client in clients if client is not None) == expected_retries + assert clients[0] is clients[2] + assert (clients[0] is clients[1]) is reuse_first + + @pytest.mark.parametrize( "call_type", [