mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
test(azure): drop mock-theater tests from azure split files per CI audit
Function-level deletions per the keep/drop audit (5): - test_azure_ai.py: test_azure_ai_services_handler, test_azure_ai_services_with_api_version (URL/api-key header asserted on patched post, pattern b) - test_azure_o_series.py: test_azure_o3_streaming, test_azure_o_series_routing, test_openai_o_series_max_retries_0, test_azure_o1_series_response_format_extra_params (kwarg reached the patched SDK client, pattern c) - test_azure_openai.py: test_azure_extra_headers, test_get_azure_ad_token_from_username_password, test_azure_openai_gpt_4o_naming, test_azure_gpt_4o_with_tool_call_and_response_format, test_azure_max_retries_0, test_async_azure_max_retries_0, test_azure_instruct, test_azure_embedding_max_retries_0, test_azure_openai_responses_bridge (patterns a/b/c) Unused imports removed via ruff F401.
This commit is contained in:
parent
00eb3dbbaf
commit
178700e351
3 changed files with 1 additions and 625 deletions
|
|
@ -1,7 +1,6 @@
|
|||
# What is this?
|
||||
## Unit tests for Azure AI integration
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
import traceback
|
||||
|
|
@ -10,29 +9,22 @@ from dotenv import load_dotenv
|
|||
|
||||
import litellm.types
|
||||
import litellm.types.utils
|
||||
from litellm.llms.anthropic.chat import ModelResponseIterator
|
||||
import httpx
|
||||
import json
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
# from base_rerank_unit_tests import BaseLLMRerankTest
|
||||
|
||||
load_dotenv()
|
||||
import io
|
||||
import os
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
from typing import Optional
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm import completion
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
|
||||
AZURE_AI_API_BASE = os.getenv("AZURE_AI_API_BASE")
|
||||
|
||||
|
|
@ -116,80 +108,6 @@ async def test_azure_ai_with_image_url():
|
|||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api_base, expected_url",
|
||||
[
|
||||
(
|
||||
"https://litellm8397336933.services.ai.azure.com/models/chat/completions?api-version=2024-05-01-preview",
|
||||
"https://litellm8397336933.services.ai.azure.com/models/chat/completions?api-version=2024-05-01-preview",
|
||||
),
|
||||
(
|
||||
"https://litellm8397336933.services.ai.azure.com/models/chat/completions",
|
||||
"https://litellm8397336933.services.ai.azure.com/models/chat/completions",
|
||||
),
|
||||
(
|
||||
"https://litellm8397336933.services.ai.azure.com/models",
|
||||
"https://litellm8397336933.services.ai.azure.com/models/chat/completions",
|
||||
),
|
||||
(
|
||||
"https://litellm8397336933.services.ai.azure.com",
|
||||
"https://litellm8397336933.services.ai.azure.com/models/chat/completions",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_azure_ai_services_handler(api_base, expected_url):
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
litellm.set_verbose = True
|
||||
|
||||
client = HTTPHandler()
|
||||
|
||||
with patch.object(client, "post") as mock_client:
|
||||
try:
|
||||
response = litellm.completion(
|
||||
model="azure_ai/Meta-Llama-3.1-70B-Instruct",
|
||||
messages=[{"role": "user", "content": "Hello, how are you?"}],
|
||||
api_key="my-fake-api-key",
|
||||
api_base=api_base,
|
||||
client=client,
|
||||
)
|
||||
|
||||
print(response)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error: {e}")
|
||||
|
||||
mock_client.assert_called_once()
|
||||
assert mock_client.call_args.kwargs["headers"]["api-key"] == "my-fake-api-key"
|
||||
assert mock_client.call_args.kwargs["url"] == expected_url
|
||||
|
||||
|
||||
def test_azure_ai_services_with_api_version():
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler
|
||||
|
||||
client = HTTPHandler()
|
||||
|
||||
with patch.object(client, "post") as mock_client:
|
||||
try:
|
||||
response = litellm.completion(
|
||||
model="azure_ai/Meta-Llama-3.1-70B-Instruct",
|
||||
messages=[{"role": "user", "content": "Hello, how are you?"}],
|
||||
api_key="my-fake-api-key",
|
||||
api_version="2024-05-01-preview",
|
||||
api_base="https://litellm8397336933.services.ai.azure.com/models",
|
||||
client=client,
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"Error: {e}")
|
||||
|
||||
mock_client.assert_called_once()
|
||||
assert mock_client.call_args.kwargs["headers"]["api-key"] == "my-fake-api-key"
|
||||
assert (
|
||||
mock_client.call_args.kwargs["url"]
|
||||
== "https://litellm8397336933.services.ai.azure.com/models/chat/completions?api-version=2024-05-01-preview"
|
||||
)
|
||||
|
||||
|
||||
def test_azure_deepseek_reasoning_content():
|
||||
import json
|
||||
|
||||
|
|
@ -251,7 +169,6 @@ async def test_azure_ai_request_format():
|
|||
Test that Azure AI requests are formatted correctly with the proper endpoint and parameters
|
||||
for both synchronous and asynchronous calls
|
||||
"""
|
||||
from openai import AsyncAzureOpenAI, AzureOpenAI
|
||||
|
||||
litellm._turn_on_debug()
|
||||
|
||||
|
|
|
|||
|
|
@ -1,19 +1,13 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock, patch, MagicMock
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm import Choices, Message, ModelResponse
|
||||
from base_llm_unit_tests import BaseLLMChatTest, BaseOSeriesModelsTest
|
||||
|
||||
|
||||
|
|
@ -97,143 +91,3 @@ class TestAzureOpenAIO3(BaseOSeriesModelsTest):
|
|||
)
|
||||
|
||||
|
||||
def test_azure_o3_streaming():
|
||||
"""
|
||||
Test that o3 models handles fake streaming correctly.
|
||||
"""
|
||||
from openai import AzureOpenAI
|
||||
from litellm import completion
|
||||
|
||||
client = AzureOpenAI(
|
||||
api_key="my-fake-o1-key",
|
||||
base_url="https://openai-gpt-4-test-v-1.openai.azure.com",
|
||||
api_version="2024-02-15-preview",
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
client.chat.completions.with_raw_response, "create"
|
||||
) as mock_create:
|
||||
try:
|
||||
completion(
|
||||
model="azure/o3-mini",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
stream=True,
|
||||
client=client,
|
||||
)
|
||||
except (
|
||||
Exception
|
||||
) as e: # expect output translation error as mock response doesn't return a json
|
||||
print(e)
|
||||
assert mock_create.call_count == 1
|
||||
assert "stream" in mock_create.call_args.kwargs
|
||||
|
||||
|
||||
def test_azure_o_series_routing():
|
||||
"""
|
||||
Allows user to pass model="azure/o_series/<any-deployment-name>" for explicit o_series model routing.
|
||||
"""
|
||||
from openai import AzureOpenAI
|
||||
from litellm import completion
|
||||
|
||||
client = AzureOpenAI(
|
||||
api_key="my-fake-o1-key",
|
||||
base_url="https://openai-gpt-4-test-v-1.openai.azure.com",
|
||||
api_version="2024-02-15-preview",
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
client.chat.completions.with_raw_response, "create"
|
||||
) as mock_create:
|
||||
try:
|
||||
completion(
|
||||
model="azure/o_series/my-random-deployment-name",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
stream=True,
|
||||
client=client,
|
||||
)
|
||||
except (
|
||||
Exception
|
||||
) as e: # expect output translation error as mock response doesn't return a json
|
||||
print(e)
|
||||
assert mock_create.call_count == 1
|
||||
assert "stream" not in mock_create.call_args.kwargs
|
||||
|
||||
|
||||
@patch("litellm.main.azure_o1_chat_completions._get_openai_client")
|
||||
def test_openai_o_series_max_retries_0(mock_get_openai_client):
|
||||
import litellm
|
||||
|
||||
litellm.set_verbose = True
|
||||
response = litellm.completion(
|
||||
model="azure/o1-preview",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
max_retries=0,
|
||||
)
|
||||
|
||||
mock_get_openai_client.assert_called_once()
|
||||
assert mock_get_openai_client.call_args.kwargs["max_retries"] == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_o1_series_response_format_extra_params():
|
||||
"""
|
||||
Tool calling should work for all azure o_series models.
|
||||
"""
|
||||
litellm._turn_on_debug()
|
||||
|
||||
from openai import AsyncAzureOpenAI
|
||||
|
||||
litellm.set_verbose = True
|
||||
|
||||
client = AsyncAzureOpenAI(
|
||||
api_key="fake-api-key",
|
||||
base_url="https://openai-prod-test.openai.azure.com/openai/deployments/o1/chat/completions?api-version=2025-01-01-preview",
|
||||
api_version="2025-01-01-preview",
|
||||
)
|
||||
|
||||
tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_current_time",
|
||||
"description": "Get the current time in a given location.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"location": {
|
||||
"type": "string",
|
||||
"description": "The city name, e.g. San Francisco",
|
||||
}
|
||||
},
|
||||
"required": ["location"],
|
||||
},
|
||||
},
|
||||
}
|
||||
]
|
||||
response_format = {"type": "json_object"}
|
||||
tool_choice = "auto"
|
||||
with patch.object(
|
||||
client.chat.completions.with_raw_response, "create"
|
||||
) as mock_client:
|
||||
try:
|
||||
await litellm.acompletion(
|
||||
client=client,
|
||||
model="azure/o_series/<my-deployment-name>",
|
||||
api_key="xxxxx",
|
||||
api_base="https://openai-prod-test.openai.azure.com/openai/deployments/o1/chat/completions?api-version=2025-01-01-preview",
|
||||
api_version="2024-12-01-preview",
|
||||
messages=[{"role": "user", "content": "Hello! return a json object"}],
|
||||
tools=tools,
|
||||
response_format=response_format,
|
||||
tool_choice=tool_choice,
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"Error: {e}")
|
||||
|
||||
mock_client.assert_called_once()
|
||||
request_body = mock_client.call_args.kwargs
|
||||
|
||||
print("request_body: ", json.dumps(request_body, indent=4))
|
||||
assert request_body["tools"] == tools
|
||||
assert request_body["response_format"] == response_format
|
||||
assert request_body["tool_choice"] == tool_choice
|
||||
|
|
|
|||
|
|
@ -5,7 +5,6 @@ sys.path.insert(
|
|||
0, os.path.abspath("../../")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from litellm.llms.azure.common_utils import process_azure_headers
|
||||
from httpx import Headers
|
||||
|
|
@ -98,79 +97,12 @@ def test_process_azure_headers_with_dict_input():
|
|||
assert result == expected_output, "Unexpected output for dict input"
|
||||
|
||||
|
||||
from httpx import Client
|
||||
from unittest.mock import MagicMock, patch
|
||||
from openai import AzureOpenAI
|
||||
from unittest.mock import MagicMock
|
||||
import litellm
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"input, call_type",
|
||||
[
|
||||
({"messages": [{"role": "user", "content": "Hello world"}]}, "completion"),
|
||||
({"input": "Hello world"}, "embedding"),
|
||||
({"prompt": "Hello world"}, "image_generation"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"header_value",
|
||||
[
|
||||
"headers",
|
||||
"extra_headers",
|
||||
],
|
||||
)
|
||||
def test_azure_extra_headers(input, call_type, header_value):
|
||||
from litellm import embedding, image_generation
|
||||
|
||||
# Clear the LLM clients cache to ensure the new http_client is used
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
http_client = Client()
|
||||
|
||||
messages = [{"role": "user", "content": "Hello world"}]
|
||||
with patch.object(http_client, "send", new=MagicMock()) as mock_client:
|
||||
litellm.client_session = http_client
|
||||
try:
|
||||
if call_type == "completion":
|
||||
func = completion
|
||||
elif call_type == "embedding":
|
||||
func = embedding
|
||||
elif call_type == "image_generation":
|
||||
func = image_generation
|
||||
|
||||
data = {
|
||||
"model": "azure/gpt-4.1-mini",
|
||||
"api_base": "https://openai-gpt-4-test-v-1.openai.azure.com",
|
||||
"api_version": "2023-07-01-preview",
|
||||
"api_key": "my-azure-api-key",
|
||||
header_value: {
|
||||
"Authorization": "my-bad-key",
|
||||
"Ocp-Apim-Subscription-Key": "hello-world-testing",
|
||||
},
|
||||
**input,
|
||||
}
|
||||
response = func(**data)
|
||||
print(response)
|
||||
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
mock_client.assert_called()
|
||||
|
||||
print(f"mock_client.call_args: {mock_client.call_args}")
|
||||
request = mock_client.call_args[0][0]
|
||||
print(request.method) # This will print 'POST'
|
||||
print(request.url) # This will print the full URL
|
||||
print(request.headers) # This will print the full URL
|
||||
auth_header = request.headers.get("Authorization")
|
||||
apim_key = request.headers.get("Ocp-Apim-Subscription-Key")
|
||||
print(auth_header)
|
||||
assert auth_header == "my-bad-key"
|
||||
assert apim_key == "hello-world-testing"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api_base, model, expected_endpoint",
|
||||
[
|
||||
|
|
@ -217,160 +149,6 @@ class TestAzureEmbedding(BaseLLMEmbeddingTest):
|
|||
return litellm.LlmProviders.AZURE
|
||||
|
||||
|
||||
@patch("azure.identity.UsernamePasswordCredential")
|
||||
@patch("azure.identity.get_bearer_token_provider")
|
||||
def test_get_azure_ad_token_from_username_password(
|
||||
mock_get_bearer_token_provider, mock_credential
|
||||
):
|
||||
from litellm.llms.azure.common_utils import (
|
||||
get_azure_ad_token_from_username_password,
|
||||
)
|
||||
|
||||
# Test inputs
|
||||
client_id = "test-client-id"
|
||||
username = "test-username"
|
||||
password = "test-password"
|
||||
|
||||
# Mock the token provider function
|
||||
mock_token_provider = lambda: "mock-token"
|
||||
mock_get_bearer_token_provider.return_value = mock_token_provider
|
||||
|
||||
# Call the function
|
||||
result = get_azure_ad_token_from_username_password(
|
||||
client_id=client_id, azure_username=username, azure_password=password
|
||||
)
|
||||
|
||||
# Verify UsernamePasswordCredential was called with correct arguments
|
||||
mock_credential.assert_called_once_with(
|
||||
client_id=client_id, username=username, password=password
|
||||
)
|
||||
|
||||
# Verify get_bearer_token_provider was called
|
||||
mock_get_bearer_token_provider.assert_called_once_with(
|
||||
mock_credential.return_value, "https://cognitiveservices.azure.com/.default"
|
||||
)
|
||||
|
||||
# Verify the result is the mock token provider
|
||||
assert result == mock_token_provider
|
||||
|
||||
|
||||
def test_azure_openai_gpt_4o_naming(monkeypatch):
|
||||
from openai import AzureOpenAI
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
monkeypatch.setenv("AZURE_API_VERSION", "2024-10-21")
|
||||
|
||||
client = AzureOpenAI(
|
||||
api_key="test-api-key",
|
||||
base_url="https://fake-azure-endpoint.invalid",
|
||||
api_version="2023-12-01-preview",
|
||||
)
|
||||
|
||||
class ResponseFormat(BaseModel):
|
||||
|
||||
number: str = Field(description="total number of days in a week")
|
||||
days: list[str] = Field(description="name of days in a week")
|
||||
|
||||
with patch.object(client.chat.completions.with_raw_response, "create") as mock_post:
|
||||
try:
|
||||
completion(
|
||||
model="azure/gpt4o",
|
||||
messages=[{"role": "user", "content": "Hello world"}],
|
||||
response_format=ResponseFormat,
|
||||
client=client,
|
||||
)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
|
||||
print(mock_post.call_args.kwargs)
|
||||
|
||||
assert "tool_calls" not in mock_post.call_args.kwargs
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api_version",
|
||||
[
|
||||
"2024-10-21",
|
||||
# "2024-02-15-preview",
|
||||
],
|
||||
)
|
||||
def test_azure_gpt_4o_with_tool_call_and_response_format(api_version):
|
||||
from litellm import completion
|
||||
from typing import Optional
|
||||
from pydantic import BaseModel
|
||||
import litellm
|
||||
|
||||
from openai import AzureOpenAI
|
||||
|
||||
client = AzureOpenAI(
|
||||
api_key="fake-key",
|
||||
base_url="https://fake-azure.openai.azure.com",
|
||||
api_version=api_version,
|
||||
)
|
||||
|
||||
class InvestigationOutput(BaseModel):
|
||||
alert_explanation: Optional[str] = None
|
||||
investigation: Optional[str] = None
|
||||
conclusions_and_possible_root_causes: Optional[str] = None
|
||||
next_steps: Optional[str] = None
|
||||
related_logs: Optional[str] = None
|
||||
app_or_infra: Optional[str] = None
|
||||
external_links: Optional[str] = None
|
||||
|
||||
tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_current_time",
|
||||
"description": "Returns the current date and time",
|
||||
"strict": True,
|
||||
"parameters": {
|
||||
"properties": {
|
||||
"timezone": {
|
||||
"type": "string",
|
||||
"description": "The timezone to get the current time for (e.g., 'UTC', 'America/New_York')",
|
||||
}
|
||||
},
|
||||
"required": ["timezone"],
|
||||
"type": "object",
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
with patch.object(client.chat.completions.with_raw_response, "create") as mock_post:
|
||||
response = litellm.completion(
|
||||
model="azure/gpt-4.1-mini",
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": "You are a tool-calling AI assist provided with common devops and IT tools that you can use to troubleshoot problems or answer questions.\nWhenever possible you MUST first use tools to investigate then answer the question.",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What is the current date and time in NYC?",
|
||||
},
|
||||
],
|
||||
drop_params=True,
|
||||
temperature=0.00000001,
|
||||
tools=tools,
|
||||
tool_choice="auto",
|
||||
response_format=InvestigationOutput, # commenting this line will cause the output to be correct
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
|
||||
if api_version == "2024-10-21":
|
||||
assert "response_format" in mock_post.call_args.kwargs
|
||||
else:
|
||||
assert "response_format" not in mock_post.call_args.kwargs
|
||||
|
||||
|
||||
def test_map_openai_params():
|
||||
"""
|
||||
Ensure response_format does not override tools
|
||||
|
|
@ -466,151 +244,6 @@ def test_map_openai_params():
|
|||
assert len(optional_params["tools"]) > 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("max_retries", [0, 4])
|
||||
@pytest.mark.parametrize("stream", [True, False])
|
||||
@patch(
|
||||
"litellm.main.azure_chat_completions.make_sync_azure_openai_chat_completion_request"
|
||||
)
|
||||
def test_azure_max_retries_0(
|
||||
mock_make_sync_azure_openai_chat_completion_request, max_retries, stream
|
||||
):
|
||||
import litellm
|
||||
from litellm import completion
|
||||
|
||||
# Clear the LLM clients cache to ensure max_retries is set correctly
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
try:
|
||||
completion(
|
||||
model="azure/gpt-4.1-mini",
|
||||
messages=[{"role": "user", "content": "Hello world"}],
|
||||
max_retries=max_retries,
|
||||
stream=stream,
|
||||
)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
mock_make_sync_azure_openai_chat_completion_request.assert_called_once()
|
||||
assert (
|
||||
mock_make_sync_azure_openai_chat_completion_request.call_args.kwargs[
|
||||
"azure_client"
|
||||
].max_retries
|
||||
== max_retries
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("max_retries", [0, 4])
|
||||
@pytest.mark.parametrize("stream", [True, False])
|
||||
@patch("litellm.main.azure_chat_completions.make_azure_openai_chat_completion_request")
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_azure_max_retries_0(
|
||||
make_azure_openai_chat_completion_request, max_retries, stream
|
||||
):
|
||||
import litellm
|
||||
from litellm import acompletion
|
||||
|
||||
# Clear the LLM clients cache to ensure max_retries is set correctly
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
try:
|
||||
await acompletion(
|
||||
model="azure/gpt-4.1-mini",
|
||||
messages=[{"role": "user", "content": "Hello world"}],
|
||||
max_retries=max_retries,
|
||||
stream=stream,
|
||||
)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
make_azure_openai_chat_completion_request.assert_called_once()
|
||||
assert (
|
||||
make_azure_openai_chat_completion_request.call_args.kwargs[
|
||||
"azure_client"
|
||||
].max_retries
|
||||
== max_retries
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("max_retries", [0, 4])
|
||||
@pytest.mark.parametrize("stream", [True, False])
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@patch("litellm.llms.azure.common_utils.select_azure_base_url_or_endpoint")
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_instruct(
|
||||
mock_select_azure_base_url_or_endpoint, max_retries, stream, sync_mode
|
||||
):
|
||||
import litellm
|
||||
from litellm import completion, acompletion
|
||||
|
||||
# Clear the LLM clients cache to ensure select_azure_base_url_or_endpoint is called
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
args = {
|
||||
"model": "azure_text/instruct-model",
|
||||
"messages": [
|
||||
{"role": "user", "content": "What is the weather like in Boston?"}
|
||||
],
|
||||
"max_tokens": 10,
|
||||
"max_retries": max_retries,
|
||||
}
|
||||
|
||||
try:
|
||||
if sync_mode:
|
||||
completion(**args)
|
||||
else:
|
||||
await acompletion(**args)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
mock_select_azure_base_url_or_endpoint.assert_called_once()
|
||||
assert (
|
||||
mock_select_azure_base_url_or_endpoint.call_args.kwargs["azure_client_params"][
|
||||
"max_retries"
|
||||
]
|
||||
== max_retries
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("max_retries", [0, 4])
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@patch("litellm.llms.azure.common_utils.select_azure_base_url_or_endpoint")
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_embedding_max_retries_0(
|
||||
mock_select_azure_base_url_or_endpoint, max_retries, sync_mode
|
||||
):
|
||||
import litellm
|
||||
from litellm import aembedding, embedding
|
||||
|
||||
# Clear the LLM clients cache to ensure select_azure_base_url_or_endpoint is called
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
args = {
|
||||
"model": "azure/text-embedding-ada-002",
|
||||
"input": "Hello world",
|
||||
"max_retries": max_retries,
|
||||
}
|
||||
|
||||
try:
|
||||
if sync_mode:
|
||||
embedding(**args)
|
||||
else:
|
||||
await aembedding(**args)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
mock_select_azure_base_url_or_endpoint.assert_called_once()
|
||||
print(
|
||||
"mock_select_azure_base_url_or_endpoint.call_args.kwargs",
|
||||
mock_select_azure_base_url_or_endpoint.call_args.kwargs,
|
||||
)
|
||||
assert (
|
||||
mock_select_azure_base_url_or_endpoint.call_args.kwargs["azure_client_params"][
|
||||
"max_retries"
|
||||
]
|
||||
== max_retries
|
||||
)
|
||||
|
||||
|
||||
def test_azure_safety_result():
|
||||
"""Bubble up safety result from Azure OpenAI"""
|
||||
from litellm import completion
|
||||
|
|
@ -629,32 +262,6 @@ def test_azure_safety_result():
|
|||
assert response.choices[0].provider_specific_fields is not None
|
||||
|
||||
|
||||
def test_azure_openai_responses_bridge():
|
||||
from litellm import completion
|
||||
import litellm
|
||||
|
||||
litellm._turn_on_debug()
|
||||
|
||||
with patch.object(litellm, "responses") as mock_responses:
|
||||
try:
|
||||
response = completion(
|
||||
model="azure/responses/test-azure-computer-use-preview",
|
||||
messages=[{"role": "user", "content": "Hello world"}],
|
||||
api_base=os.getenv("AZURE_COMPUTER_USE_API_BASE"),
|
||||
api_version="2025-04-01-preview",
|
||||
api_key=os.getenv("AZURE_COMPUTER_USE_API_KEY"),
|
||||
)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
mock_responses.assert_called_once()
|
||||
assert (
|
||||
mock_responses.call_args.kwargs["model"]
|
||||
== "test-azure-computer-use-preview"
|
||||
)
|
||||
assert mock_responses.call_args.kwargs["custom_llm_provider"] == "azure"
|
||||
|
||||
|
||||
def test_completion_azure_deployment_id():
|
||||
"""
|
||||
Ensure deployment_id takes precedence over model.
|
||||
|
|
@ -678,10 +285,8 @@ def test_azure_with_content_safety_error():
|
|||
"""
|
||||
Verify user can access innererror from the Azure OpenAI exception
|
||||
"""
|
||||
from litellm import completion
|
||||
from litellm.exceptions import ContentPolicyViolationError
|
||||
from litellm.litellm_core_utils.exception_mapping_utils import exception_type
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
mock_exception = Exception(
|
||||
"The response was filtered due to the prompt triggering Azure OpenAI's content management policy"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue