From 680f81599143add8eaf55bec1ba81f991d419718 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Wed, 26 Feb 2025 16:06:35 -0800 Subject: [PATCH] fix flow for adding custom model names --- .../conditional_public_model_name.tsx | 56 ++++++++++++++++--- .../add_model/litellm_model_name.tsx | 14 +++-- 2 files changed, 56 insertions(+), 14 deletions(-) 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..d87552106cc 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 @@ -7,17 +7,57 @@ const ConditionalPublicModelName: React.FC = () => { // Access the form instance const form = Form.useFormInstance(); - // 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 + // Update model mappings immediately when custom model name changes + const handleCustomModelNameChange = (value: string) => { + if (selectedModels.includes('custom') && value) { + const currentMappings = form.getFieldValue('model_mappings') || []; + const updatedMappings = currentMappings.map((mapping: any) => { + if (mapping.public_name === 'custom' || + (mapping.public_name !== value && mapping.litellm_model !== value && + mapping.public_name === mapping.litellm_model)) { + return { + public_name: value, + litellm_model: value + }; + } + return mapping; + }); + form.setFieldValue('model_mappings', updatedMappings); + } + }; + + // Listen for changes to the custom_model_name field + useEffect(() => { + const unsubscribe = form.getFieldInstance('custom_model_name')?.addEventListener('input', (e: any) => { + handleCustomModelNameChange(e.target.value); + }); + + return () => { + if (unsubscribe) unsubscribe(); + }; + }, [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); } }, [selectedModels, form]); @@ -32,7 +72,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; diff --git a/ui/litellm-dashboard/src/components/add_model/litellm_model_name.tsx b/ui/litellm-dashboard/src/components/add_model/litellm_model_name.tsx index 883923316b9..54fff4d3863 100644 --- a/ui/litellm-dashboard/src/components/add_model/litellm_model_name.tsx +++ b/ui/litellm-dashboard/src/components/add_model/litellm_model_name.tsx @@ -63,6 +63,10 @@ const LiteLLMModelNameField: React.FC = ({ (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 +74,7 @@ const LiteLLMModelNameField: React.FC = ({ ...providerModels.map(model => ({ label: model, value: model - })), - { - label: 'Custom Model Name (Enter below)', - value: 'custom' - } + })) ]} style={{ width: '100%' }} /> @@ -92,7 +92,9 @@ const LiteLLMModelNameField: React.FC = ({ > {({ getFieldValue }) => { const selectedModels = getFieldValue('model') || []; - return selectedModels.includes('custom') && ( + // Ensure selectedModels is always an array + const modelArray = Array.isArray(selectedModels) ? selectedModels : [selectedModels]; + return modelArray.includes('custom') && (