mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
(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:
parent
6c669734a6
commit
c07dd16d88
6 changed files with 179 additions and 19 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
}}
|
||||
|
|
|
|||
|
|
@ -86,6 +86,7 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({
|
|||
)}
|
||||
|
||||
{(selectedProviderEnum === Providers.Azure ||
|
||||
selectedProviderEnum === Providers.Azure_AI_Studio ||
|
||||
selectedProviderEnum === Providers.OpenAI_Compatible
|
||||
) && (
|
||||
<Form.Item
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue