mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(azure): include effective max_retries in client cache key
This commit is contained in:
parent
615ed7900f
commit
d61a23f2de
2 changed files with 65 additions and 1 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue