From c07dd16d88621cbf082a6ce1d7d920585ddec442 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Wed, 26 Feb 2025 18:45:02 -0800 Subject: [PATCH] (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 --- litellm/llms/azure_ai/chat/transformation.py | 16 +++- .../chat/test_azure_ai_transformation.py | 91 +++++++++++++++++++ .../conditional_public_model_name.tsx | 51 +++++++++-- .../add_model/litellm_model_name.tsx | 37 ++++++-- .../add_model/provider_specific_fields.tsx | 1 + .../src/components/provider_info_helpers.tsx | 2 +- 6 files changed, 179 insertions(+), 19 deletions(-) create mode 100644 tests/litellm/llms/azure_ai/chat/test_azure_ai_transformation.py diff --git a/litellm/llms/azure_ai/chat/transformation.py b/litellm/llms/azure_ai/chat/transformation.py index afedc950019..2815eaa14c9 100644 --- a/litellm/llms/azure_ai/chat/transformation.py +++ b/litellm/llms/azure_ai/chat/transformation.py @@ -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, diff --git a/tests/litellm/llms/azure_ai/chat/test_azure_ai_transformation.py b/tests/litellm/llms/azure_ai/chat/test_azure_ai_transformation.py new file mode 100644 index 00000000000..239bea950cd --- /dev/null +++ b/tests/litellm/llms/azure_ai/chat/test_azure_ai_transformation.py @@ -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 diff --git a/ui/litellm-dashboard/src/components/add_model/conditional_public_model_name.tsx b/ui/litellm-dashboard/src/components/add_model/conditional_public_model_name.tsx index a9229c8006c..8e0f6a288b3 100644 --- a/ui/litellm-dashboard/src/components/add_model/conditional_public_model_name.tsx +++ b/ui/litellm-dashboard/src/components/add_model/conditional_public_model_name.tsx @@ -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 ( { const newMappings = [...form.getFieldValue('model_mappings')]; newMappings[index].public_name = e.target.value; @@ -61,6 +91,7 @@ const ConditionalPublicModelName: React.FC = () => { required={true} > = ({ } }; + // Handle custom model name changes + const handleCustomModelNameChange = (e: React.ChangeEvent) => { + 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 ( <> = ({ (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 = ({ ...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 = ({ > {({ getFieldValue }) => { const selectedModels = getFieldValue('model') || []; - return selectedModels.includes('custom') && ( + const modelArray = Array.isArray(selectedModels) ? selectedModels : [selectedModels]; + return modelArray.includes('custom') && ( - + ); }} diff --git a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx index f23bf1e0156..d50c08b6328 100644 --- a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx +++ b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx @@ -86,6 +86,7 @@ const ProviderSpecificFields: React.FC = ({ )} {(selectedProviderEnum === Providers.Azure || + selectedProviderEnum === Providers.Azure_AI_Studio || selectedProviderEnum === Providers.OpenAI_Compatible ) && (