(UI Fixes for add new model flow) (#8216)

* ui fix add model flow

* fix provider info + add flow

* cleanup add model setup

* use 1 file for ProviderSpecificFields

* use 1 file for ProviderSpecificFields

* use antd select for model / providers

* fix selectedProviderEnum

* fix upload vertex models
This commit is contained in:
Ishaan Jaff 2025-02-03 08:21:00 -08:00 • committed by GitHub
parent 6f11137d6f
commit b3154be6f5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 251 additions and 222 deletions

View file

@ -1,5 +1,5 @@
import { message } from "antd";
import { provider_map } from "../provider_info_helpers";
import { provider_map, Providers } from "../provider_info_helpers";
import { modelCreateCall, Model } from "../networking";
@ -12,7 +12,9 @@ export const handleAddModelSubmit = async (
console.log("handling submit for formValues:", formValues);
// If model_name is not provided, use provider.toLowerCase() + "/*"
if (formValues["model"] && formValues["model"].includes("all-wildcard")) {
const wildcardModel = formValues["custom_llm_provider"].toLowerCase() + "/*";
const customProvider: Providers = formValues["custom_llm_provider"];
const litellm_custom_provider = provider_map[customProvider as keyof typeof Providers];
const wildcardModel = litellm_custom_provider + "/*";
formValues["model_name"] = wildcardModel;
formValues["model"] = wildcardModel;
}

View file

@ -0,0 +1,171 @@
import React from "react";
import { Form } from "antd";
import { TextInput, Text } from "@tremor/react";
import { Row, Col, Typography, Button as Button2, Upload, UploadProps } from "antd";
import { UploadOutlined } from "@ant-design/icons";
import { Providers } from "../provider_info_helpers";
const { Link } = Typography;
interface ProviderSpecificFieldsProps {
selectedProvider: Providers;
uploadProps?: UploadProps;
}
const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({
selectedProvider,
uploadProps
}) => {
console.log(`Selected provider: ${selectedProvider}`);
console.log(`type of selectedProvider: ${typeof selectedProvider}`);
// cast selectedProvider to Providers
const selectedProviderEnum = Providers[selectedProvider as keyof typeof Providers] as Providers;
console.log(`selectedProviderEnum: ${selectedProviderEnum}`);
console.log(`type of selectedProviderEnum: ${typeof selectedProviderEnum}`);
return (
<>
{selectedProviderEnum === Providers.OpenAI && (
<Form.Item label="Organization ID" name="organization">
<TextInput placeholder="[OPTIONAL] my-unique-org" />
</Form.Item>
)}
{selectedProviderEnum === Providers.Vertex_AI && (
<>
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="Vertex Project"
name="vertex_project"
>
<TextInput placeholder="adroit-cadet-1234.." />
</Form.Item>
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="Vertex Location"
name="vertex_location"
>
<TextInput placeholder="us-east-1" />
</Form.Item>
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="Vertex Credentials"
name="vertex_credentials"
className="mb-0"
>
<Upload {...uploadProps}>
<Button2 icon={<UploadOutlined />}>
Click to Upload
</Button2>
</Upload>
</Form.Item>
<Row>
<Col span={10}></Col>
<Col span={10}>
<Text className="mb-3 mt-1">
Give litellm a gcp service account(.json file), so it
can make the relevant calls
</Text>
</Col>
</Row>
</>
)}
{(selectedProviderEnum === Providers.Azure ||
selectedProviderEnum === Providers.OpenAI_Compatible) && (
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="API Base"
name="api_base"
>
<TextInput placeholder="https://..." />
</Form.Item>
)}
{selectedProviderEnum === Providers.Azure && (
<>
<Form.Item
label="API Version"
name="api_version"
tooltip="By default litellm will use the latest version. If you want to use a different version, you can specify it here"
>
<TextInput placeholder="2023-07-01-preview" />
</Form.Item>
<div>
<Form.Item
label="Base Model"
name="base_model"
className="mb-0"
>
<TextInput placeholder="azure/gpt-3.5-turbo" />
</Form.Item>
<Row>
<Col span={10}></Col>
<Col span={10}>
<Text className="mb-2">
The actual model your azure deployment uses. Used
for accurate cost tracking. Select name from{" "}
<Link
href="https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json"
target="_blank"
>
here
</Link>
</Text>
</Col>
</Row>
</div>
</>
)}
{selectedProviderEnum === Providers.Bedrock && (
<>
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="AWS Access Key ID"
name="aws_access_key_id"
tooltip="You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`)."
>
<TextInput placeholder="" />
</Form.Item>
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="AWS Secret Access Key"
name="aws_secret_access_key"
tooltip="You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`)."
>
<TextInput placeholder="" />
</Form.Item>
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="AWS Region Name"
name="aws_region_name"
tooltip="You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`)."
>
<TextInput placeholder="us-east-1" />
</Form.Item>
</>
)}
{selectedProviderEnum != Providers.Bedrock &&
selectedProviderEnum != Providers.Vertex_AI &&
selectedProviderEnum != Providers.Ollama &&
(
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="API Key"
name="api_key"
tooltip="LLM API Credentials"
>
<TextInput placeholder="sk-" type="password" />
</Form.Item>
)}
</>
);
};
export default ProviderSpecificFields;

View file

@ -19,6 +19,7 @@ import {
import ConditionalPublicModelName from "./add_model/conditional_public_model_name";
import LiteLLMModelNameField from "./add_model/litellm_model_name";
import AdvancedSettings from "./add_model/advanced_settings";
import ProviderSpecificFields from "./add_model/provider_specific_fields";
import { handleAddModelSubmit } from "./add_model/handle_add_model_submit";
import EditModelModal from "./edit_model/edit_model_modal";
import {
@ -65,7 +66,7 @@ import {
Popover,
Form,
Input,
Select as Select2,
Select as AntdSelect,
InputNumber,
message,
Descriptions,
@ -99,7 +100,7 @@ import { Upload } from "antd";
import TimeToFirstToken from "./model_metrics/time_to_first_token";
import DynamicFields from "./model_add/dynamic_form";
import { Prism as SyntaxHighlighter } from "react-syntax-highlighter";
import { Providers, provider_map, providerLogoMap, getProviderLogoAndName, getPlaceholder } from "./provider_info_helpers";
import { Providers, provider_map, providerLogoMap, getProviderLogoAndName, getPlaceholder, getProviderModels } from "./provider_info_helpers";
interface ModelDashboardProps {
accessToken: string | null;
@ -178,7 +179,7 @@ const ModelDashboard: React.FC<ModelDashboardProps> = ({
const [providerSettings, setProviderSettings] = useState<ProviderSettings[]>(
[]
);
const [selectedProvider, setSelectedProvider] = useState<String>("OpenAI");
const [selectedProvider, setSelectedProvider] = useState<Providers>(Providers.OpenAI);
const [healthCheckResponse, setHealthCheckResponse] = useState<string>("");
const [editModalVisible, setEditModalVisible] = useState<boolean>(false);
const [infoModalVisible, setInfoModalVisible] = useState<boolean>(false);
@ -226,6 +227,12 @@ const ModelDashboard: React.FC<ModelDashboardProps> = ({
// Add state for advanced settings visibility
const [showAdvancedSettings, setShowAdvancedSettings] = useState<boolean>(false);
const setProviderModelsFn = (provider: Providers) => {
const _providerModels = getProviderModels(provider, modelMap);
setProviderModels(_providerModels);
console.log(`providerModels: ${_providerModels}`);
};
const updateModelMetrics = async (
modelGroup: string | null,
startTime: Date | undefined,
@ -439,7 +446,7 @@ const ModelDashboard: React.FC<ModelDashboardProps> = ({
}
};
const props: UploadProps = {
const uploadProps: UploadProps = {
name: "file",
accept: ".json",
beforeUpload: (file) => {
@ -786,8 +793,6 @@ const ModelDashboard: React.FC<ModelDashboardProps> = ({
}
});
}
if (userRole && userRole == "Admin Viewer") {
const { Title, Paragraph } = Typography;
return (
@ -800,48 +805,6 @@ const ModelDashboard: React.FC<ModelDashboardProps> = ({
);
}
const setProviderModelsFn = (provider: string) => {
console.log(`received provider string: ${provider}`);
let providerKey = provider;
if (providerKey) {
let _providerModels: Array<string> = [];
if (typeof modelMap === "object") {
Object.entries(modelMap).forEach(([key, value]) => {
if (
value !== null &&
typeof value === "object" &&
"litellm_provider" in (value as object) &&
((value as any)["litellm_provider"] === providerKey ||
(value as any)["litellm_provider"].includes(providerKey))
) {
_providerModels.push(key);
}
});
// Special case for cohere_chat
// we need both cohere_chat and cohere models to show on dropdown
if (providerKey == Providers.Cohere) {
console.log("adding cohere chat model")
Object.entries(modelMap).forEach(([key, value]) => {
if (
value !== null &&
typeof value === "object" &&
"litellm_provider" in (value as object) &&
((value as any)["litellm_provider"] === "cohere")
) {
_providerModels.push(key);
}
});
}
}
setProviderModels(_providerModels);
console.log(`providerModels: ${providerModels}`);
}
};
const runHealthCheck = async () => {
try {
message.info("Running health check...");
@ -1536,32 +1499,27 @@ const ModelDashboard: React.FC<ModelDashboardProps> = ({
labelCol={{ span: 10 }}
labelAlign="left"
>
<Select
value={selectedProvider as string}
<AntdSelect
showSearch={true}
value={selectedProvider}
onChange={(value) => {
// Set the selected provider
setSelectedProvider(value as unknown as string);
// Update provider-specific models
setProviderModelsFn(provider_map[value as unknown as string]);
// Reset the 'model' field
form.setFieldsValue({ model: [] });
// Reset the 'model_name' field
form.setFieldsValue({ model_name: undefined });
setSelectedProvider(value);
setProviderModelsFn(value);
form.setFieldsValue({
model: [],
model_name: undefined
});
}}
>
{Object.keys(Providers).map((providerKey) => (
<SelectItem
key={providerKey}
value={Providers[providerKey as keyof typeof Providers]}
onClick={() => {
setProviderModelsFn(provider_map[providerKey as keyof typeof Providers]);
setSelectedProvider(Providers[providerKey as keyof typeof Providers]);
}}
{Object.entries(Providers).map(([providerEnum, providerDisplayName]) => (
<AntdSelect.Option
key={providerEnum}
value={providerEnum}
>
<div className="flex items-center space-x-2">
<img
src={providerLogoMap[Providers[providerKey as keyof typeof Providers]]}
alt={`${Providers[providerKey as keyof typeof Providers]} logo`}
src={providerLogoMap[providerDisplayName]}
alt={`${providerEnum} logo`}
className="w-5 h-5"
onError={(e) => {
// Create a div with provider initial as fallback
@ -1570,19 +1528,19 @@ const ModelDashboard: React.FC<ModelDashboardProps> = ({
if (parent) {
const fallbackDiv = document.createElement('div');
fallbackDiv.className = 'w-5 h-5 rounded-full bg-gray-200 flex items-center justify-center text-xs';
fallbackDiv.textContent = Providers[providerKey as keyof typeof Providers].charAt(0);
fallbackDiv.textContent = providerDisplayName.charAt(0);
parent.replaceChild(fallbackDiv, target);
}
}}
/>
<span>{Providers[providerKey as keyof typeof Providers]}</span>
<span>{providerDisplayName}</span>
</div>
</SelectItem>
</AntdSelect.Option>
))}
</Select>
</AntdSelect>
</Form.Item>
<LiteLLMModelNameField
selectedProvider={selectedProvider as string}
selectedProvider={selectedProvider}
providerModels={providerModels}
getPlaceholder={getPlaceholder}
/>
@ -1590,153 +1548,10 @@ const ModelDashboard: React.FC<ModelDashboardProps> = ({
{/* Conditionally Render "Public Model Name" */}
<ConditionalPublicModelName />
{/* Provider-specific fields */}
{dynamicProviderForm !== undefined &&
dynamicProviderForm.fields.length > 0 && (
<DynamicFields
fields={dynamicProviderForm.fields}
selectedProvider={dynamicProviderForm.name}
/>
)}
{selectedProvider != Providers.Bedrock &&
selectedProvider != Providers.Vertex_AI &&
selectedProvider != Providers.Ollama &&
(dynamicProviderForm === undefined ||
dynamicProviderForm.fields.length == 0) && (
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="API Key"
name="api_key"
tooltip="LLM API Credentials"
>
<TextInput placeholder="sk-" type="password" />
</Form.Item>
)}
{selectedProvider == Providers.OpenAI && (
<Form.Item label="Organization ID" name="organization">
<TextInput placeholder="[OPTIONAL] my-unique-org" />
</Form.Item>
)}
{selectedProvider == Providers.Vertex_AI && (
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="Vertex Project"
name="vertex_project"
>
<TextInput placeholder="adroit-cadet-1234.." />
</Form.Item>
)}
{selectedProvider == Providers.Vertex_AI && (
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="Vertex Location"
name="vertex_location"
>
<TextInput placeholder="us-east-1" />
</Form.Item>
)}
{selectedProvider == Providers.Vertex_AI && (
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="Vertex Credentials"
name="vertex_credentials"
className="mb-0"
>
<Upload {...props}>
<Button2 icon={<UploadOutlined />}>
Click to Upload
</Button2>
</Upload>
</Form.Item>
)}
{selectedProvider == Providers.Vertex_AI && (
<Row>
<Col span={10}></Col>
<Col span={10}>
<Text className="mb-3 mt-1">
Give litellm a gcp service account(.json file), so it
can make the relevant calls
</Text>
</Col>
</Row>
)}
{(selectedProvider == Providers.Azure ||
selectedProvider == Providers.OpenAI_Compatible) && (
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="API Base"
name="api_base"
>
<TextInput placeholder="https://..." />
</Form.Item>
)}
{selectedProvider == Providers.Azure && (
<Form.Item
label="API Version"
name="api_version"
tooltip="By default litellm will use the latest version. If you want to use a different version, you can specify it here"
>
<TextInput placeholder="2023-07-01-preview" />
</Form.Item>
)}
{selectedProvider == Providers.Azure && (
<div>
<Form.Item
label="Base Model"
name="base_model"
className="mb-0"
>
<TextInput placeholder="azure/gpt-3.5-turbo" />
</Form.Item>
<Row>
<Col span={10}></Col>
<Col span={10}>
<Text className="mb-2">
The actual model your azure deployment uses. Used
for accurate cost tracking. Select name from{" "}
<Link
href="https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json"
target="_blank"
>
here
</Link>
</Text>
</Col>
</Row>
</div>
)}
{selectedProvider == Providers.Bedrock && (
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="AWS Access Key ID"
name="aws_access_key_id"
tooltip="You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`)."
>
<TextInput placeholder="" />
</Form.Item>
)}
{selectedProvider == Providers.Bedrock && (
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="AWS Secret Access Key"
name="aws_secret_access_key"
tooltip="You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`)."
>
<TextInput placeholder="" />
</Form.Item>
)}
{selectedProvider == Providers.Bedrock && (
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="AWS Region Name"
name="aws_region_name"
tooltip="You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`)."
>
<TextInput placeholder="us-east-1" />
</Form.Item>
)}
<ProviderSpecificFields
selectedProvider={selectedProvider}
uploadProps={uploadProps}
/>
<AdvancedSettings
showAdvancedSettings={showAdvancedSettings}
setShowAdvancedSettings={setShowAdvancedSettings}

View file

@ -92,3 +92,44 @@ export const getPlaceholder = (selectedProvider: string): string => {
return "gpt-3.5-turbo";
}
};
export const getProviderModels = (provider: Providers, modelMap: any): Array<string> => {
let providerKey = provider;
console.log(`Provider key: ${providerKey}`);
let custom_llm_provider = provider_map[providerKey];
console.log(`Provider mapped to: ${custom_llm_provider}`);
let providerModels: Array<string> = [];
if (providerKey && typeof modelMap === "object") {
Object.entries(modelMap).forEach(([key, value]) => {
if (
value !== null &&
typeof value === "object" &&
"litellm_provider" in (value as object) &&
((value as any)["litellm_provider"] === custom_llm_provider ||
(value as any)["litellm_provider"].includes(custom_llm_provider))
) {
providerModels.push(key);
}
});
// Special case for cohere_chat
// we need both cohere_chat and cohere models to show on dropdown
if (providerKey == Providers.Cohere) {
console.log("Adding cohere chat models");
Object.entries(modelMap).forEach(([key, value]) => {
if (
value !== null &&
typeof value === "object" &&
"litellm_provider" in (value as object) &&
((value as any)["litellm_provider"] === "cohere")
) {
providerModels.push(key);
}
});
}
}
return providerModels;
};