This commit is contained in:
huhupy 2026-10-05 01:58:35 +08:00 • committed by GitHub
commit a29c03699d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 65 additions and 1 deletions

View file

@ -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")

View file

@ -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",
[