(UI + Backend) Fix Adding Azure, Azure AI Studio models on LiteLLM (#8856)

* fix Azure_AI_Studio

* fix flow for adding custom model names

* fix _should_use_api_key_header

* handle custom model name change

* test_azure_ai_request_format

* Azure AI Foundry (Studio)

* fix _should_use_api_key_header
This commit is contained in:
Ishaan Jaff 2025-02-26 18:45:02 -08:00 • committed by GitHub
parent 6c669734a6
commit c07dd16d88
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 179 additions and 19 deletions

View file

@ -1,4 +1,5 @@
from typing import Any, List, Optional, Tuple, cast
from urllib.parse import urlparse
import httpx
from httpx import Response
@ -28,13 +29,26 @@ class AzureAIStudioConfig(OpenAIConfig):
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
if api_base and "services.ai.azure.com" in api_base:
if api_base and self._should_use_api_key_header(api_base):
headers["api-key"] = api_key
else:
headers["Authorization"] = f"Bearer {api_key}"
return headers
def _should_use_api_key_header(self, api_base: str) -> bool:
"""
Returns True if the request should use `api-key` header for authentication.
"""
parsed_url = urlparse(api_base)
host = parsed_url.hostname
if host and (
host.endswith(".services.ai.azure.com")
or host.endswith(".openai.azure.com")
):
return True
return False
def get_complete_url(
self,
api_base: str,

View file

@ -0,0 +1,91 @@
import json
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
sys.path.insert(
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
import litellm
@pytest.mark.parametrize("is_async", [True, False])
@pytest.mark.asyncio
async def test_azure_ai_request_format(is_async):
"""
Test that Azure AI requests are formatted correctly with the proper endpoint and parameters
for both synchronous and asynchronous calls
"""
litellm._turn_on_debug()
# Set up the test parameters
api_key = "00xxx"
api_base = "https://my-endpoint-europe-berri-992.openai.azure.com/openai/deployments/gpt-4o-mini/chat/completions?api-version=2024-08-01-preview"
model = "azure_ai/gpt-4o-mini"
messages = [
{"role": "user", "content": "hi"},
{"role": "assistant", "content": "Hello! How can I assist you today?"},
{"role": "user", "content": "hi"},
]
if is_async:
# Mock AsyncHTTPHandler.post method for async test
with patch(
"litellm.llms.custom_httpx.llm_http_handler.AsyncHTTPHandler.post"
) as mock_post:
# Set up mock response
mock_post.return_value = AsyncMock()
# Call the acompletion function
try:
await litellm.acompletion(
custom_llm_provider="azure_ai",
api_key=api_key,
api_base=api_base,
model=model,
messages=messages,
)
except Exception as e:
# We expect an exception since we're mocking the response
pass
else:
# Mock HTTPHandler.post method for sync test
with patch(
"litellm.llms.custom_httpx.llm_http_handler.HTTPHandler.post"
) as mock_post:
# Set up mock response
mock_post.return_value = MagicMock()
# Call the completion function
try:
litellm.completion(
custom_llm_provider="azure_ai",
api_key=api_key,
api_base=api_base,
model=model,
messages=messages,
)
except Exception as e:
# We expect an exception since we're mocking the response
pass
# Verify the request was made with the correct parameters
mock_post.assert_called_once()
call_args = mock_post.call_args
print("sync request call=", json.dumps(call_args.kwargs, indent=4, default=str))
# Check URL
assert call_args.kwargs["url"] == api_base
# Check headers
assert call_args.kwargs["headers"]["api-key"] == api_key
# Check request body
request_body = json.loads(call_args.kwargs["data"])
assert (
request_body["model"] == "gpt-4o-mini"
) # Model name should be stripped of provider prefix
assert request_body["messages"] == messages

View file

@ -1,4 +1,4 @@
import React, { useEffect } from "react";
import React, { useEffect, useState } from "react";
import { Form, Table, Input } from "antd";
import { Text, TextInput } from "@tremor/react";
import { Row, Col } from "antd";
@ -6,21 +6,51 @@ import { Row, Col } from "antd";
const ConditionalPublicModelName: React.FC = () => {
// Access the form instance
const form = Form.useFormInstance();
const [tableKey, setTableKey] = useState(0); // Add a key to force table re-render
// Watch the 'model' field for changes
const selectedModels = Form.useWatch('model', form) || [];
// Watch the 'model' field for changes and ensure it's always an array
const modelValue = Form.useWatch('model', form) || [];
const selectedModels = Array.isArray(modelValue) ? modelValue : [modelValue];
const customModelName = Form.useWatch('custom_model_name', form);
const showPublicModelName = !selectedModels.includes('all-wildcard');
// Auto-populate model mappings when selected models change
// Force table to re-render when custom model name changes
useEffect(() => {
if (customModelName && selectedModels.includes('custom')) {
const currentMappings = form.getFieldValue('model_mappings') || [];
const updatedMappings = currentMappings.map((mapping: any) => {
if (mapping.public_name === 'custom' || mapping.litellm_model === 'custom') {
return {
public_name: customModelName,
litellm_model: customModelName
};
}
return mapping;
});
form.setFieldValue('model_mappings', updatedMappings);
setTableKey(prev => prev + 1); // Force table re-render
}
}, [customModelName, selectedModels, form]);
// Initial setup of model mappings when models are selected
useEffect(() => {
if (selectedModels.length > 0 && !selectedModels.includes('all-wildcard')) {
const mappings = selectedModels.map((model: string) => ({
public_name: model,
litellm_model: model
}));
const mappings = selectedModels.map((model: string) => {
if (model === 'custom' && customModelName) {
return {
public_name: customModelName,
litellm_model: customModelName
};
}
return {
public_name: model,
litellm_model: model
};
});
form.setFieldValue('model_mappings', mappings);
setTableKey(prev => prev + 1); // Force table re-render
}
}, [selectedModels, form]);
}, [selectedModels, customModelName, form]);
if (!showPublicModelName) return null;
@ -32,7 +62,7 @@ const ConditionalPublicModelName: React.FC = () => {
render: (text: string, record: any, index: number) => {
return (
<TextInput
defaultValue={text}
value={text}
onChange={(e) => {
const newMappings = [...form.getFieldValue('model_mappings')];
newMappings[index].public_name = e.target.value;
@ -61,6 +91,7 @@ const ConditionalPublicModelName: React.FC = () => {
required={true}
>
<Table
key={tableKey} // Add key to force re-render
dataSource={form.getFieldValue('model_mappings')}
columns={columns}
pagination={false}

View file

@ -35,6 +35,25 @@ const LiteLLMModelNameField: React.FC<LiteLLMModelNameFieldProps> = ({
}
};
// Handle custom model name changes
const handleCustomModelNameChange = (e: React.ChangeEvent<HTMLInputElement>) => {
const customName = e.target.value;
// Immediately update the model mappings
const currentMappings = form.getFieldValue('model_mappings') || [];
const updatedMappings = currentMappings.map((mapping: any) => {
if (mapping.public_name === 'custom' || mapping.litellm_model === 'custom') {
return {
public_name: customName,
litellm_model: customName
};
}
return mapping;
});
form.setFieldsValue({ model_mappings: updatedMappings });
};
return (
<>
<Form.Item
@ -63,6 +82,10 @@ const LiteLLMModelNameField: React.FC<LiteLLMModelNameFieldProps> = ({
(option?.label ?? '').toLowerCase().includes(input.toLowerCase())
}
options={[
{
label: 'Custom Model Name (Enter below)',
value: 'custom'
},
{
label: `All ${selectedProvider} Models (Wildcard)`,
value: 'all-wildcard'
@ -70,11 +93,7 @@ const LiteLLMModelNameField: React.FC<LiteLLMModelNameFieldProps> = ({
...providerModels.map(model => ({
label: model,
value: model
})),
{
label: 'Custom Model Name (Enter below)',
value: 'custom'
}
}))
]}
style={{ width: '100%' }}
/>
@ -92,13 +111,17 @@ const LiteLLMModelNameField: React.FC<LiteLLMModelNameFieldProps> = ({
>
{({ getFieldValue }) => {
const selectedModels = getFieldValue('model') || [];
return selectedModels.includes('custom') && (
const modelArray = Array.isArray(selectedModels) ? selectedModels : [selectedModels];
return modelArray.includes('custom') && (
<Form.Item
name="custom_model_name"
rules={[{ required: true, message: "Please enter a custom model name." }]}
className="mt-2"
>
<TextInput placeholder="Enter custom model name" />
<TextInput
placeholder="Enter custom model name"
onChange={handleCustomModelNameChange}
/>
</Form.Item>
);
}}

View file

@ -86,6 +86,7 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({
)}
{(selectedProviderEnum === Providers.Azure ||
selectedProviderEnum === Providers.Azure_AI_Studio ||
selectedProviderEnum === Providers.OpenAI_Compatible
) && (
<Form.Item

View file

@ -3,7 +3,7 @@ import React from "react";
export enum Providers {
OpenAI = "OpenAI",
Azure = "Azure",
Azure_AI_Studio = "Azure AI Studio",
Azure_AI_Studio = "Azure AI Foundry (Studio)",
Anthropic = "Anthropic",
Vertex_AI = "Vertex AI (Anthropic, Gemini, etc.)",
Google_AI_Studio = "Google AI Studio",