fix: fix azure tests

This commit is contained in:
Krrish Dholakia
2026-03-28 18:19:41 -07:00
parent c75fb3650a
commit 44b03e6138
2 changed files with 70 additions and 35 deletions
+3 -3
View File
@@ -287,8 +287,8 @@ def test_audio_speech_cost_calc():
from litellm.integrations.custom_logger import CustomLogger
model = "azure/azure-tts"
api_base = os.getenv("AZURE_SWEDEN_API_BASE")
api_key = os.getenv("AZURE_SWEDEN_API_KEY")
api_base = os.getenv("AZURE_TTS_API_BASE")
api_key = os.getenv("AZURE_TTS_API_KEY")
custom_logger = CustomLogger()
litellm.set_verbose = True
@@ -301,7 +301,7 @@ def test_audio_speech_cost_calc():
input="the quick brown fox jumped over the lazy dogs",
api_base=api_base,
api_key=api_key,
base_model="azure/tts-1",
base_model="azure/tts",
)
time.sleep(1)
@@ -539,7 +539,11 @@ async def test_ensure_initialize_azure_sdk_client_always_used(call_type):
patch_target = (
"litellm.rerank_api.main.azure_rerank.initialize_azure_sdk_client"
)
elif call_type == CallTypes.acreate_batch or call_type == CallTypes.aretrieve_batch or call_type == CallTypes.acancel_batch:
elif (
call_type == CallTypes.acreate_batch
or call_type == CallTypes.aretrieve_batch
or call_type == CallTypes.acancel_batch
):
patch_target = (
"litellm.batches.main.azure_batches_instance.initialize_azure_sdk_client"
)
@@ -570,7 +574,9 @@ async def test_ensure_initialize_azure_sdk_client_always_used(call_type):
or call_type == CallTypes.avideo_extension
):
# Skip video call types as they don't use Azure SDK client initialization
pytest.skip(f"Skipping {call_type.value} because Azure video calls don't use initialize_azure_sdk_client")
pytest.skip(
f"Skipping {call_type.value} because Azure video calls don't use initialize_azure_sdk_client"
)
elif (
call_type == CallTypes.alist_containers
or call_type == CallTypes.aretrieve_container
@@ -580,13 +586,26 @@ async def test_ensure_initialize_azure_sdk_client_always_used(call_type):
or call_type == CallTypes.aupload_container_file
):
# Skip container call types as they're not supported for Azure (only OpenAI)
pytest.skip(f"Skipping {call_type.value} because Azure doesn't support container operations")
elif call_type == CallTypes.avector_store_file_create or call_type == CallTypes.avector_store_file_list or call_type == CallTypes.avector_store_file_retrieve or call_type == CallTypes.avector_store_file_content or call_type == CallTypes.avector_store_file_update or call_type == CallTypes.avector_store_file_delete:
pytest.skip(
f"Skipping {call_type.value} because Azure doesn't support container operations"
)
elif (
call_type == CallTypes.avector_store_file_create
or call_type == CallTypes.avector_store_file_list
or call_type == CallTypes.avector_store_file_retrieve
or call_type == CallTypes.avector_store_file_content
or call_type == CallTypes.avector_store_file_update
or call_type == CallTypes.avector_store_file_delete
):
# Skip vector store file call types as they're not supported for Azure (only OpenAI)
pytest.skip(f"Skipping {call_type.value} because Azure doesn't support vector store file operations")
pytest.skip(
f"Skipping {call_type.value} because Azure doesn't support vector store file operations"
)
elif call_type == CallTypes.aocr or call_type == CallTypes.ocr:
# Skip OCR call types as they don't use Azure SDK client initialization
pytest.skip(f"Skipping {call_type.value} because OCR calls don't use initialize_azure_sdk_client")
pytest.skip(
f"Skipping {call_type.value} because OCR calls don't use initialize_azure_sdk_client"
)
# Mock the initialize_azure_sdk_client function
with patch(patch_target) as mock_init_azure:
# Also mock async_function_with_fallbacks to prevent actual API calls
@@ -767,7 +786,7 @@ AZURE_API_FUNCTION_PARAMS = [
"speech",
False,
{
"model": "azure/tts-1",
"model": "azure/tts",
"input": "Hello, this is a test of text to speech",
"voice": "alloy",
"api_key": "test-api-key",
@@ -1434,43 +1453,44 @@ def test_token_provider_raises_exception(setup_mocks):
def test_get_azure_ad_token_provider_with_default_azure_credential():
"""
Test that get_azure_ad_token_provider correctly uses DefaultAzureCredential
Test that get_azure_ad_token_provider correctly uses DefaultAzureCredential
when explicitly specified as the credential type. This verifies that the function
can dynamically instantiate DefaultAzureCredential and return a working token provider.
"""
# Mock Azure identity classes
with patch('azure.identity.DefaultAzureCredential') as mock_default_cred, \
patch('azure.identity.get_bearer_token_provider') as mock_token_provider:
with patch("azure.identity.DefaultAzureCredential") as mock_default_cred, patch(
"azure.identity.get_bearer_token_provider"
) as mock_token_provider:
# Configure mocks
mock_credential_instance = MagicMock()
mock_default_cred.return_value = mock_credential_instance
mock_token_provider.return_value = lambda: "test-default-azure-token"
# Test with DefaultAzureCredential specified explicitly
token_provider = get_azure_ad_token_provider(
azure_scope="https://cognitiveservices.azure.com/.default",
azure_credential=AzureCredentialType.DefaultAzureCredential
azure_credential=AzureCredentialType.DefaultAzureCredential,
)
# Verify DefaultAzureCredential was instantiated
mock_default_cred.assert_called_once_with()
# Verify get_bearer_token_provider was called with the right parameters
mock_token_provider.assert_called_once_with(
mock_credential_instance,
"https://cognitiveservices.azure.com/.default"
mock_credential_instance, "https://cognitiveservices.azure.com/.default"
)
# Verify the returned token provider works
token = token_provider()
assert token == "test-default-azure-token"
def test_get_azure_ad_token_fallback_to_default_azure_credential(setup_mocks, monkeypatch):
def test_get_azure_ad_token_fallback_to_default_azure_credential(
setup_mocks, monkeypatch
):
"""
Test that get_azure_ad_token falls back to DefaultAzureCredential when the
service principal method fails but token refresh is enabled. This tests the
Test that get_azure_ad_token falls back to DefaultAzureCredential when the
service principal method fails but token refresh is enabled. This tests the
complete fallback flow from service principal to DefaultAzureCredential.
"""
# Clear environment variables that might interfere
@@ -1486,7 +1506,7 @@ def test_get_azure_ad_token_fallback_to_default_azure_credential(setup_mocks, mo
# Enable token refresh
setup_mocks["litellm"].enable_azure_ad_token_refresh = True
# Configure get_azure_ad_token_provider to fail first (service principal)
# Configure get_azure_ad_token_provider to fail first (service principal)
# but succeed on second call (DefaultAzureCredential)
def mock_token_provider_side_effect(*args, **kwargs):
# If called with azure_credential=DefaultAzureCredential, return a working provider
@@ -1512,19 +1532,22 @@ def test_get_azure_ad_token_fallback_to_default_azure_credential(setup_mocks, mo
# 1. First with just azure_scope (service principal attempt)
# 2. Second with azure_credential=DefaultAzureCredential (fallback)
assert setup_mocks["token_provider"].call_count == 2
# Verify the calls were made with expected parameters
calls = setup_mocks["token_provider"].call_args_list
# First call should be service principal attempt (no azure_credential)
first_call_kwargs = calls[0][1]
assert "azure_scope" in first_call_kwargs
assert first_call_kwargs.get("azure_credential") is None
# Second call should be DefaultAzureCredential attempt
second_call_kwargs = calls[1][1]
assert "azure_scope" in second_call_kwargs
assert second_call_kwargs.get("azure_credential") == AzureCredentialType.DefaultAzureCredential
assert (
second_call_kwargs.get("azure_credential")
== AzureCredentialType.DefaultAzureCredential
)
# Verify the token is what we expect from our DefaultAzureCredential mock
assert token == "mock-default-azure-credential-token"
@@ -1584,9 +1607,13 @@ def test_azure_v1_api_uses_openai_client(api_version):
)
# Should be OpenAI client, not AzureOpenAI
assert isinstance(client, OpenAI), f"Expected OpenAI client for api_version={api_version}"
assert isinstance(
client, OpenAI
), f"Expected OpenAI client for api_version={api_version}"
# base_url should be /openai/v1/ (not /deployments/)
assert "/openai/v1/" in str(client.base_url), f"base_url should contain /openai/v1/, got {client.base_url}"
assert "/openai/v1/" in str(
client.base_url
), f"base_url should contain /openai/v1/, got {client.base_url}"
# Test async client
with patch.object(base_llm, "initialize_azure_sdk_client") as mock_init:
@@ -1606,9 +1633,13 @@ def test_azure_v1_api_uses_openai_client(api_version):
)
# Should be AsyncOpenAI client, not AsyncAzureOpenAI
assert isinstance(async_client, AsyncOpenAI), f"Expected AsyncOpenAI client for api_version={api_version}"
assert isinstance(
async_client, AsyncOpenAI
), f"Expected AsyncOpenAI client for api_version={api_version}"
# base_url should be /openai/v1/
assert "/openai/v1/" in str(async_client.base_url), f"base_url should contain /openai/v1/, got {async_client.base_url}"
assert "/openai/v1/" in str(
async_client.base_url
), f"base_url should contain /openai/v1/, got {async_client.base_url}"
def test_azure_traditional_api_uses_azure_openai_client():
@@ -1643,7 +1674,9 @@ def test_azure_traditional_api_uses_azure_openai_client():
)
# Should be AzureOpenAI client
assert isinstance(client, AzureOpenAI), f"Expected AzureOpenAI client for api_version={api_version}"
assert isinstance(
client, AzureOpenAI
), f"Expected AzureOpenAI client for api_version={api_version}"
# Test async client
with patch.object(base_llm, "initialize_azure_sdk_client") as mock_init:
@@ -1663,4 +1696,6 @@ def test_azure_traditional_api_uses_azure_openai_client():
)
# Should be AsyncAzureOpenAI client
assert isinstance(async_client, AsyncAzureOpenAI), f"Expected AsyncAzureOpenAI client for api_version={api_version}"
assert isinstance(
async_client, AsyncAzureOpenAI
), f"Expected AsyncAzureOpenAI client for api_version={api_version}"