mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-11 00:25:15 +00:00
fix: fix azure tests
This commit is contained in:
@@ -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}"
|
||||
|
||||
Reference in New Issue
Block a user