From a2f8c87e9f540399d1eb784a861ef9504309b92d Mon Sep 17 00:00:00 2001 From: cursor Date: Fri, 24 Apr 2026 09:23:03 +0000 Subject: [PATCH] feat(ui): migrate add_model subtree to shadcn Migrate the add-model surface (AddModelForm, add_model_tab, add_auto_router_tab, advanced_settings, cache_control_settings, conditional_public_model_name, litellm_model_name, provider_specific_fields) from antd Form + tremor to shadcn primitives and react-hook-form. - Root form state moves from antd FormInstance (Form.useForm) to RHF UseFormReturn, created in ModelsAndEndpointsView.tsx and threaded through AddModelTab / AddModelForm. Children use useFormContext / Controller / useWatch / useFieldArray. - antd Modal -> shadcn Dialog; antd Tag -> Badge; antd Select -> shadcn Select; antd Radio.Group -> RadioGroup; antd Accordion (tremor) -> shadcn Accordion; antd Form.List -> RHF useFieldArray. - Submit helpers (handle_add_model_submit, handle_add_auto_router_submit) feature-detect antd.resetFields vs RHF.reset so both paths keep working through the mixed migration window. - UploadProps type import migrates from antd/es/upload to a local shim at add_model/add_model_upload_types.ts (structurally compatible with the antd shape still used by model_add/AddCredentialModal.tsx and model_add/EditCredentialModal.tsx). - cache_control_settings ships both an RHF path (new callers) and a legacy antd bridge used by model_info_view.tsx so the out-of-scope detail view keeps working without cascading its migration. - validateJsonValue added to utils/textUtils.ts to replace the antd Promise-style formItemValidateJSON for RHF rules. - Tests repaired: AddModelForm, add_model_tab, conditional_public_model_name, litellm_model_name, provider_specific_fields, advanced_settings now wrap with RHF FormProvider instead of antd Form. Co-authored-by: yuneng-jiang --- .../ModelsAndEndpointsView.tsx | 44 +- .../add_model/AddModelForm.test.tsx | 22 +- .../src/components/add_model/AddModelForm.tsx | 895 ++++++++++++------ .../add_model/add_auto_router_tab.tsx | 737 ++++++++------ .../add_model/add_model_tab.test.tsx | 32 +- .../components/add_model/add_model_tab.tsx | 43 +- .../add_model/add_model_upload_types.ts | 25 + .../add_model/advanced_settings.test.tsx | 55 +- .../add_model/advanced_settings.tsx | 704 +++++++++----- .../add_model/cache_control_settings.tsx | 408 +++++--- .../conditional_public_model_name.test.tsx | 35 +- .../conditional_public_model_name.tsx | 286 +++--- .../handle_add_auto_router_submit.tsx | 10 +- .../add_model/handle_add_model_submit.tsx | 8 +- .../add_model/litellm_model_name.test.tsx | 21 +- .../add_model/litellm_model_name.tsx | 372 +++++--- .../provider_specific_fields.test.tsx | 33 +- .../add_model/provider_specific_fields.tsx | 335 ++++--- ui/litellm-dashboard/src/utils/textUtils.ts | 16 + 19 files changed, 2660 insertions(+), 1421 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/add_model/add_model_upload_types.ts diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx index 6b2086fa9d8..42346567273 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx @@ -6,6 +6,10 @@ import AllModelsTab from "@/app/(dashboard)/models-and-endpoints/components/AllM import ModelRetrySettingsTab from "@/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab"; import PriceDataManagementTab from "@/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab"; import { handleAddModelSubmit } from "@/components/add_model/handle_add_model_submit"; +import type { + AddModelFormValues, +} from "@/components/add_model/AddModelForm"; +import type { UploadProps } from "@/components/add_model/add_model_upload_types"; import { Team } from "@/components/key_team_helpers/key_list"; import CredentialsPanel from "@/components/model_add/credentials"; import { getCallbacksCall, setCallbacksCall } from "@/components/networking"; @@ -17,8 +21,7 @@ import { RefreshCcw } from "lucide-react"; import { useQueryClient } from "@tanstack/react-query"; // eslint-disable-next-line litellm-ui/no-banned-ui-imports import { Tab, TabGroup, TabList, TabPanel, TabPanels } from "@tremor/react"; -import type { UploadProps } from "antd"; -import { Form } from "antd"; +import { useForm } from "react-hook-form"; import { PlusCircle as PlusCircleOutlined } from "lucide-react"; import React, { useEffect, useMemo, useState } from "react"; import AddModelTab from "../../../components/add_model/add_model_tab"; @@ -49,7 +52,11 @@ interface GlobalRetryPolicyObject { const ModelsAndEndpointsView: React.FC = ({ premiumUser, teams }) => { const { accessToken, token, userRole, userId: userID } = useAuthorized(); - const [addModelForm] = Form.useForm(); + const addModelForm = useForm({ + defaultValues: { + model_mappings: [], + }, + }); const [lastRefreshed, setLastRefreshed] = useState(""); const [providerModels, setProviderModels] = useState>([]); const [selectedProvider, setSelectedProvider] = useState(Providers.Anthropic); @@ -140,22 +147,6 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te }; const uploadProps: UploadProps = { - name: "file", - accept: ".json", - pastable: false, - beforeUpload: (file) => { - if (file.type === "application/json") { - const reader = new FileReader(); - reader.onload = (e) => { - if (e.target) { - const jsonStr = e.target.result as string; - addModelForm.setFieldsValue({ vertex_credentials: jsonStr }); - } - }; - reader.readAsText(file); - } - return false; - }, onChange(info) { if (info.file.status === "done") { NotificationsManager.success(`${info.file.name} file uploaded successfully`); @@ -242,17 +233,16 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te } const handleOk = async () => { + // Form validation is handled inside the shadcn `AddModelForm` via + // `form.handleSubmit` before calling this callback; at this point the + // form values are already valid and we can proceed with the submit. + const values = addModelForm.getValues(); try { - const values = await addModelForm.validateFields(); await handleAddModelSubmit(values, accessToken, addModelForm, handleRefreshClick); } catch (error: any) { - const errorMessages = - error.errorFields - ?.map((field: any) => { - return `${field.name.join(".")}: ${field.errors.join(", ")}`; - }) - .join(" | ") || "Unknown validation error"; - NotificationsManager.fromBackend(`Please fill in the following required fields: ${errorMessages}`); + NotificationsManager.fromBackend( + `Failed to add model: ${error?.message ?? "Unknown error"}`, + ); } }; diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx index 4a3dbaedf74..baa39493d8f 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx @@ -1,12 +1,12 @@ import { renderHook, screen, waitFor, renderWithProviders } from "../../../tests/test-utils"; import userEvent from "@testing-library/user-event"; -import { Form } from "antd"; -import type { UploadProps } from "antd/es/upload"; +import { useForm } from "react-hook-form"; import { describe, expect, it, vi } from "vitest"; import type { Team } from "../key_team_helpers/key_list"; import type { CredentialItem } from "../networking"; import { Providers } from "../provider_info_helpers"; -import AddModelForm from "./AddModelForm"; +import AddModelForm, { type AddModelFormValues } from "./AddModelForm"; +import type { UploadProps } from "./add_model_upload_types"; vi.mock("../molecules/models/ProviderLogo", () => ({ ProviderLogo: ({ provider, className }: { provider: string; className?: string }) => ( @@ -101,6 +101,8 @@ vi.mock("@/app/(dashboard)/hooks/tags/useTags", () => ({ })); const mockAuthorizedUser = (userRole: string, userId: string, premiumUser: boolean) => ({ + isLoading: false, + isAuthorized: true, token: "test-token", accessToken: "test-access-token", userId, @@ -123,11 +125,14 @@ const testTeam: Team = { created_at: "2024-01-01T00:00:00Z", keys: [], members_with_roles: [], -}; + spend: 0, +} as Team; const createTestProps = (userRole = "proxy_admin", userId = "user-1", isTeamAdmin = false) => { - const { result } = renderHook(() => Form.useForm()); - const [form] = result.current; + const { result } = renderHook(() => + useForm({ defaultValues: { model_mappings: [] } }), + ); + const form = result.current; const teams = [ { @@ -147,10 +152,7 @@ const createTestProps = (userRole = "proxy_admin", userId = "user-1", isTeamAdmi }, ]; - const uploadProps: UploadProps = { - beforeUpload: () => false, - showUploadList: false, - }; + const uploadProps: UploadProps = {}; return { form, diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx index 2f92aafcf82..039aec1eef5 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx @@ -1,16 +1,48 @@ import { useProviderFields } from "@/app/(dashboard)/hooks/providers/useProviderFields"; import { useGuardrails } from "@/app/(dashboard)/hooks/guardrails/useGuardrails"; import { useTags } from "@/app/(dashboard)/hooks/tags/useTags"; -// eslint-disable-next-line litellm-ui/no-banned-ui-imports import { all_admin_roles, isUserTeamAdminForAnyTeam } from "@/utils/roles"; -import { Switch, Text } from "@tremor/react"; -import type { FormInstance } from "antd"; -import { Select as AntdSelect, Button, Card, Col, Form, Modal, Row, Tooltip, Typography, Alert } from "antd"; -import type { UploadProps } from "antd/es/upload"; +import { Button } from "@/components/ui/button"; +import { Card } from "@/components/ui/card"; +import { Label } from "@/components/ui/label"; +import { Switch } from "@/components/ui/switch"; +import { Badge } from "@/components/ui/badge"; +import { Input } from "@/components/ui/input"; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from "@/components/ui/select"; +import { + Alert, + AlertDescription, + AlertTitle, +} from "@/components/ui/alert"; +import { + Dialog, + DialogContent, + DialogFooter, + DialogHeader, + DialogTitle, +} from "@/components/ui/dialog"; +import { Info, Loader2, X } from "lucide-react"; import React, { useEffect, useMemo, useState } from "react"; +import { + Controller, + FormProvider, + UseFormReturn, + useFormContext, + useWatch, +} from "react-hook-form"; import TeamDropdown from "../common_components/team_dropdown"; import type { Team } from "../key_team_helpers/key_list"; -import { type CredentialItem, type ProviderCreateInfo, modelAvailableCall } from "../networking"; +import { + type CredentialItem, + type ProviderCreateInfo, + modelAvailableCall, +} from "../networking"; import { Providers, providerLogoMap } from "../provider_info_helpers"; import { ProviderLogo } from "../molecules/models/ProviderLogo"; import AdvancedSettings from "./advanced_settings"; @@ -20,9 +52,42 @@ import ConnectionErrorDisplay from "./model_connection_test"; import ProviderSpecificFields from "./provider_specific_fields"; import { TEST_MODES } from "./add_model_modes"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import type { UploadProps } from "./add_model_upload_types"; + +export interface AddModelFormValues { + team_id?: string; + custom_llm_provider?: string; + model?: string | string[]; + custom_model_name?: string; + model_name?: string; + mode?: string; + litellm_credential_name?: string | null; + model_access_group?: string[]; + model_mappings?: { public_name: string; litellm_model: string }[]; + // Advanced settings + custom_pricing?: boolean; + pricing_model?: "per_token" | "per_second"; + input_cost_per_token?: string | number | null; + output_cost_per_token?: string | number | null; + input_cost_per_second?: string | number | null; + vector_store_ids?: string[]; + guardrails?: string[]; + tags?: string[]; + cache_control?: boolean; + cache_control_injection_points?: { + location: "message"; + role?: string; + index?: number | null; + }[]; + use_in_pass_through?: boolean; + litellm_extra_params?: string; + model_info_params?: string; + // Allow provider-specific credential fields to be stored under arbitrary keys. + [key: string]: unknown; +} interface AddModelFormProps { - form: FormInstance; // For the Add Model tab + form: UseFormReturn; handleOk: () => Promise; selectedProvider: Providers; setSelectedProvider: (provider: Providers) => void; @@ -36,8 +101,6 @@ interface AddModelFormProps { credentials: CredentialItem[]; } -const { Title, Link } = Typography; - const AddModelForm: React.FC = ({ form, handleOk, @@ -53,9 +116,10 @@ const AddModelForm: React.FC = ({ credentials, }) => { const [testMode, setTestMode] = useState("chat"); - const [isResultModalVisible, setIsResultModalVisible] = useState(false); - const [isTestingConnection, setIsTestingConnection] = useState(false); - // Using a unique ID to force the ConnectionErrorDisplay to remount and run a fresh test + const [isResultModalVisible, setIsResultModalVisible] = + useState(false); + const [isTestingConnection, setIsTestingConnection] = + useState(false); const [connectionTestId, setConnectionTestId] = useState(""); const { accessToken, userRole, premiumUser, userId } = useAuthorized(); @@ -65,8 +129,10 @@ const AddModelForm: React.FC = ({ error: providerMetadataError, } = useProviderFields(); const { data: guardrailsData } = useGuardrails(); - const guardrailsList = guardrailsData?.guardrails.map((g) => g.guardrail_name); - const { data: tagsList, isLoading: isTagsLoading, error: tagsError } = useTags(); + const guardrailsList = guardrailsData?.guardrails.map( + (g) => g.guardrail_name, + ); + const { data: tagsList } = useTags(); const handleTestConnection = async () => { setIsTestingConnection(true); @@ -76,13 +142,24 @@ const AddModelForm: React.FC = ({ const [isTeamOnly, setIsTeamOnly] = useState(false); const [modelAccessGroups, setModelAccessGroups] = useState([]); - // Team admin specific state - const [teamAdminSelectedTeam, setTeamAdminSelectedTeam] = useState(null); + const [teamAdminSelectedTeam, setTeamAdminSelectedTeam] = useState< + string | null + >(null); useEffect(() => { const fetchModelAccessGroups = async () => { - const response = await modelAvailableCall(accessToken, "", "", false, null, true, true); - setModelAccessGroups(response["data"].map((model: any) => model["id"])); + const response = await modelAvailableCall( + accessToken, + "", + "", + false, + null, + true, + true, + ); + setModelAccessGroups( + response["data"].map((model: { id: string }) => model.id), + ); }; fetchModelAccessGroups(); }, [accessToken]); @@ -91,7 +168,9 @@ const AddModelForm: React.FC = ({ if (!providerMetadata) { return []; } - return [...providerMetadata].sort((a, b) => a.provider_display_name.localeCompare(b.provider_display_name)); + return [...providerMetadata].sort((a, b) => + a.provider_display_name.localeCompare(b.provider_display_name), + ); }, [providerMetadata]); const providerMetadataErrorText = providerMetadataError @@ -103,321 +182,535 @@ const AddModelForm: React.FC = ({ const isAdmin = all_admin_roles.includes(userRole); const isTeamAdmin = isUserTeamAdminForAnyTeam(teams, userId); + const onSubmit = form.handleSubmit(async () => { + await handleOk().then(() => { + setTeamAdminSelectedTeam(null); + }); + }); + return ( - <> - Add Model + +

Add Model

- -
{ - console.log("๐Ÿ”ฅ Form onFinish triggered with values:", values); - await handleOk().then(() => { - setTeamAdminSelectedTeam(null); - }); - }} - onFinishFailed={(errorInfo) => { - console.log("๐Ÿ’ฅ Form onFinishFailed triggered:", errorInfo); - }} - labelCol={{ span: 10 }} - wrapperCol={{ span: 16 }} - labelAlign="left" - > - <> - {isTeamAdmin && !isAdmin && ( - <> - + + {isTeamAdmin && !isAdmin && ( + <> + + {!teamAdminSelectedTeam && ( + + + Team Selection Required + + As a team admin, you need to select your team first before + adding models. + + + )} + + )} + + {(isAdmin || (isTeamAdmin && teamAdminSelectedTeam)) && ( + <> +
+ +
+ ( + )} - {sortedProviderMetadata.map((providerInfo) => { - const displayName = providerInfo.provider_display_name; - const providerKey = providerInfo.provider; - const logoSrc = providerLogoMap[displayName] ?? ""; - - return ( - -
- - {displayName} -
-
- ); - })} - - - - - {/* Conditionally Render "Public Model Name" */} - - - {/* Select Mode */} - - setTestMode(value)} - options={TEST_MODES} /> - - - - - - Optional - LiteLLM endpoint to use when health checking this model{" "} - - Learn more - - - - - - {/* Credentials */} -
- - Either select existing credentials OR enter new provider credentials below - + {form.formState.errors.custom_llm_provider?.message && ( +

+ {String( + form.formState.errors.custom_llm_provider.message, + )} +

+ )}
+
- - (option?.label ?? "").toLowerCase().includes(input.toLowerCase())} - options={[ - { value: null, label: "None" }, - ...credentials.map((credential) => ({ - value: credential.credential_name, - label: credential.credential_name, - })), - ]} - allowClear + + + + +
+ +
+ ( + + )} /> - - - - prevValues.litellm_credential_name !== currentValues.litellm_credential_name || - prevValues.provider !== currentValues.provider - } - > - {({ getFieldValue }) => { - const credentialName = getFieldValue("litellm_credential_name"); - console.log("๐Ÿ”‘ Credential Name Changed:", credentialName); - // Only show provider specific fields if no credentials selected - if (!credentialName) { - return ( - <> -
-
- OR -
-
- - - ); - } - return null; - }} -
-
-
- Additional Model Info Settings -
- {/* Team-only Model Switch - Only show for proxy admins, not team admins */} - {(isAdmin || !isTeamAdmin) && ( - + +
+
+

+ Optional - LiteLLM endpoint to use when + health checking this model{" "} + - +

+
+ +
+

+ Either select existing credentials OR enter new provider + credentials below +

+
+ +
+ +
+ ( + + )} + /> +
+
+ + + +
+
+ + Additional Model Info Settings + +
+
+ + {(isAdmin || !isTeamAdmin) && ( +
+ +
+ { - setIsTeamOnly(checked); + onCheckedChange={(checked) => { + setIsTeamOnly(!!checked); if (!checked) { - form.setFieldValue("team_id", undefined); + form.setValue("team_id", undefined); } }} disabled={!premiumUser} /> - - - )} + +
+
+ )} - {/* Conditional Team Selection */} - {isTeamOnly && (isAdmin || !isTeamAdmin) && ( - - - - )} - {isAdmin && ( - <> - - ({ - value: group, - label: group, - }))} - maxTagCount="responsive" - allowClear - /> - - - )} - {}} + required={isTeamOnly && !isAdmin} + disabled={!premiumUser} + tooltip="Only keys for this team will be able to call this model." /> - - )} -
- - Need Help? - -
- - -
+ )} + + {isAdmin && ( +
+ +
+ ( + + )} + /> +
+
+ )} + + + + )} + +
+ + Need Help? + +
+ +
- - +
+ {/* Test Connection Results Modal */} - { - setIsResultModalVisible(false); - setIsTestingConnection(false); + onOpenChange={(open) => { + if (!open) { + setIsResultModalVisible(false); + setIsTestingConnection(false); + } }} - footer={[ - , - ]} - width={700} > - {/* Only render the ConnectionErrorDisplay when modal is visible and we have a test ID */} - {isResultModalVisible && ( - { - setIsResultModalVisible(false); - setIsTestingConnection(false); - }} - onTestComplete={() => setIsTestingConnection(false)} - /> - )} - - + + + Connection Test Results + + {isResultModalVisible && ( + { + setIsResultModalVisible(false); + setIsTestingConnection(false); + }} + onTestComplete={() => setIsTestingConnection(false)} + /> + )} + + + + + + ); }; +function TeamSelectField({ + onTeamSelected, + required, + disabled, + tooltip, +}: { + onTeamSelected: (teamId: string | null) => void; + required?: boolean; + disabled?: boolean; + tooltip?: string; +}) { + const { control, formState } = useFormContext(); + const error = (formState.errors as Record).team_id; + + return ( +
+ +
+ ( + { + field.onChange(value); + onTeamSelected(value || null); + }} + disabled={disabled} + /> + )} + /> + {error?.message && ( +

{String(error.message)}

+ )} +
+
+ ); +} + +function CredentialsGate({ + selectedProvider, + uploadProps, +}: { + selectedProvider: Providers; + uploadProps: UploadProps; +}) { + const { control } = useFormContext(); + const credentialName = useWatch({ control, name: "litellm_credential_name" }); + if (credentialName) return null; + return ( + <> +
+
+ OR +
+
+ + + ); +} + +function AccessGroupTagInput({ + value, + onChange, + options, +}: { + value: string[]; + onChange: (next: string[]) => void; + options: string[]; +}) { + const [input, setInput] = React.useState(""); + const remaining = options.filter((o) => !value.includes(o)); + + const addValue = (next: string) => { + const trimmed = next.trim(); + if (!trimmed || value.includes(trimmed)) return; + onChange([...value, trimmed]); + }; + + return ( +
+
+ setInput(e.target.value)} + placeholder="Type a group name and press Enter" + onKeyDown={(e) => { + if (e.key === "Enter" || e.key === ",") { + e.preventDefault(); + addValue(input); + setInput(""); + } + }} + /> + +
+ {value.length > 0 && ( +
+ {value.map((v) => ( + + {v} + + + ))} +
+ )} +
+ ); +} + export default AddModelForm; diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx index eaaf701bf43..37006b9eb1c 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx @@ -1,20 +1,67 @@ import React, { useEffect, useState } from "react"; -import { Card, Form, Button, Tooltip, Typography, Select as AntdSelect, Modal, Radio, Badge, Space } from "antd"; -// eslint-disable-next-line litellm-ui/no-banned-ui-imports -import type { FormInstance } from "antd"; -import { Text, TextInput } from "@tremor/react"; +import { + Controller, + FormProvider, + UseFormReturn, + useFormContext, +} from "react-hook-form"; +import { Button } from "@/components/ui/button"; +import { Card } from "@/components/ui/card"; +import { Label } from "@/components/ui/label"; +import { Badge } from "@/components/ui/badge"; +import { Input } from "@/components/ui/input"; +import { X, Loader2 } from "lucide-react"; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from "@/components/ui/select"; +import { + RadioGroup, + RadioGroupItem, +} from "@/components/ui/radio-group"; +import { + Dialog, + DialogContent, + DialogFooter, + DialogHeader, + DialogTitle, +} from "@/components/ui/dialog"; import { modelAvailableCall } from "../networking"; import ConnectionErrorDisplay from "./model_connection_test"; import { all_admin_roles } from "@/utils/roles"; import { handleAddAutoRouterSubmit } from "./handle_add_auto_router_submit"; -import { fetchAvailableModels, ModelGroup } from "../playground/llm_calls/fetch_models"; +import { + fetchAvailableModels, + ModelGroup, +} from "../playground/llm_calls/fetch_models"; import RouterConfigBuilder from "./RouterConfigBuilder"; import ComplexityRouterConfig from "./ComplexityRouterConfig"; import NotificationManager from "../molecules/notifications_manager"; import { Zap as ThunderboltOutlined, GitBranch as BranchesOutlined } from "lucide-react"; + +export interface AutoRouterFormValues { + auto_router_name: string; + auto_router_default_model?: string; + auto_router_embedding_model?: string; + model_access_group?: string[]; + // populated at submit time only + // eslint-disable-next-line @typescript-eslint/no-explicit-any + auto_router_config?: any; + // eslint-disable-next-line @typescript-eslint/no-explicit-any + complexity_router_config?: any; + model_type?: "semantic_router" | "complexity_router"; + custom_llm_provider?: string; + model?: string; + api_key?: string; + team_id?: string; +} + interface AddAutoRouterTabProps { - form: FormInstance; - handleOk: () => void; + form: UseFormReturn; + handleOk: (values: AutoRouterFormValues) => void | Promise; accessToken: string; userRole: string; } @@ -28,26 +75,24 @@ interface ComplexityTiers { REASONING: string; } -const { Title, Link } = Typography; - -const AddAutoRouterTab: React.FC = ({ form, handleOk, accessToken, userRole }) => { - // State for connection testing - const [isResultModalVisible, setIsResultModalVisible] = useState(false); - const [isTestingConnection, setIsTestingConnection] = useState(false); +const AddAutoRouterTab: React.FC = ({ + form, + handleOk, + accessToken, + userRole, +}) => { + const [isResultModalVisible, setIsResultModalVisible] = + useState(false); + const [isTestingConnection, setIsTestingConnection] = + useState(false); const [connectionTestId, setConnectionTestId] = useState(""); const [modelAccessGroups, setModelAccessGroups] = useState([]); const [modelInfo, setModelInfo] = useState([]); - const [showCustomDefaultModel, setShowCustomDefaultModel] = useState(false); - const [showCustomEmbeddingModel, setShowCustomEmbeddingModel] = useState(false); - - // Router type state - default to complexity router const [routerType, setRouterType] = useState("complexity"); - - // Semantic router config (existing) + + // eslint-disable-next-line @typescript-eslint/no-explicit-any const [routerConfig, setRouterConfig] = useState(null); - - // Complexity router config (new) const [complexityTiers, setComplexityTiers] = useState({ SIMPLE: "", MEDIUM: "", @@ -57,8 +102,18 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc useEffect(() => { const fetchModelAccessGroups = async () => { - const response = await modelAvailableCall(accessToken, "", "", false, null, true, true); - setModelAccessGroups(response["data"].map((model: any) => model["id"])); + const response = await modelAvailableCall( + accessToken, + "", + "", + false, + null, + true, + true, + ); + setModelAccessGroups( + response["data"].map((model: { id: string }) => model.id), + ); }; fetchModelAccessGroups(); }, [accessToken]); @@ -67,7 +122,6 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc const loadModels = async () => { try { const uniqueModels = await fetchAvailableModels(accessToken); - console.log("Fetched models for auto router:", uniqueModels); setModelInfo(uniqueModels); } catch (error) { console.error("Error fetching model info for auto router:", error); @@ -78,95 +132,64 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc const isAdmin = all_admin_roles.includes(userRole); - // Test connection when button is clicked const handleTestConnection = async () => { setIsTestingConnection(true); setConnectionTestId(`test-${Date.now()}`); setIsResultModalVisible(true); }; - // Auto router specific form submit handler - const handleAutoRouterSubmit = () => { - console.log("Auto router submit triggered!"); - console.log("Router type:", routerType); - - const currentFormValues = form.getFieldsValue(); - console.log("Form values:", currentFormValues); - - // Check basic required fields first + const handleAutoRouterSubmit = form.handleSubmit(async (currentFormValues) => { if (!currentFormValues.auto_router_name) { NotificationManager.fromBackend("Please enter an Auto Router Name"); return; } - // Validation differs based on router type if (routerType === "complexity") { - // Complexity Router validation const filledTiers = Object.values(complexityTiers).filter(Boolean); if (filledTiers.length === 0) { - NotificationManager.fromBackend("Please select at least one model for a complexity tier"); + NotificationManager.fromBackend( + "Please select at least one model for a complexity tier", + ); return; } - // For complexity router, use the first non-empty tier as default - const defaultModel = complexityTiers.MEDIUM || complexityTiers.SIMPLE || complexityTiers.COMPLEX || complexityTiers.REASONING; - - // Set form values for complexity router - form.setFieldsValue({ - custom_llm_provider: "auto_router", - model: currentFormValues.auto_router_name, - api_key: "not_required_for_auto_router", - auto_router_default_model: defaultModel, - }); + const defaultModel = + complexityTiers.MEDIUM || + complexityTiers.SIMPLE || + complexityTiers.COMPLEX || + complexityTiers.REASONING; - form - .validateFields(["auto_router_name"]) - .then((values) => { - console.log("Complexity router validation passed"); - - // Build the complexity router config - const submitValues = { - ...values, - auto_router_name: currentFormValues.auto_router_name, - auto_router_default_model: defaultModel, - // Use special model prefix for complexity router - model_type: "complexity_router", - complexity_router_config: { - tiers: complexityTiers, - }, - model_access_group: currentFormValues.model_access_group, - }; - - console.log("Final submit values:", submitValues); - handleAddAutoRouterSubmit(submitValues, accessToken, form, handleOk); - }) - .catch((error) => { - console.error("Validation failed:", error); - NotificationManager.fromBackend("Please fill in all required fields"); - }); - + const submitValues: AutoRouterFormValues = { + ...currentFormValues, + auto_router_default_model: defaultModel, + model_type: "complexity_router", + complexity_router_config: { tiers: complexityTiers }, + }; + + await handleAddAutoRouterSubmit(submitValues, accessToken, form, () => + handleOk(submitValues), + ); } else { - // Semantic Router validation (existing logic) if (!currentFormValues.auto_router_default_model) { NotificationManager.fromBackend("Please select a Default Model"); return; } - form.setFieldsValue({ - custom_llm_provider: "auto_router", - model: currentFormValues.auto_router_name, - api_key: "not_required_for_auto_router", - }); - - // Custom validation for router config - if (!routerConfig || !routerConfig.routes || routerConfig.routes.length === 0) { - NotificationManager.fromBackend("Please configure at least one route for the auto router"); + if ( + !routerConfig || + !routerConfig.routes || + routerConfig.routes.length === 0 + ) { + NotificationManager.fromBackend( + "Please configure at least one route for the auto router", + ); return; } - // Check if all routes have required fields const invalidRoutes = routerConfig.routes.filter( - (route: any) => !route.name || !route.description || route.utterances.length === 0, + // eslint-disable-next-line @typescript-eslint/no-explicit-any + (route: any) => + !route.name || !route.description || route.utterances.length === 0, ); if (invalidRoutes.length > 0) { @@ -176,282 +199,398 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc return; } - form - .validateFields() - .then((values) => { - console.log("Form validation passed, submitting with values:", values); - const submitValues = { - ...values, - auto_router_config: routerConfig, - model_type: "semantic_router", - }; - console.log("Final submit values:", submitValues); - handleAddAutoRouterSubmit(submitValues, accessToken, form, handleOk); - }) - .catch((error) => { - console.error("Validation failed:", error); - const fieldErrors = error.errorFields || []; - if (fieldErrors.length > 0) { - const missingFields = fieldErrors.map((field: any) => { - const fieldName = field.name[0]; - const friendlyNames: { [key: string]: string } = { - auto_router_name: "Auto Router Name", - auto_router_default_model: "Default Model", - auto_router_embedding_model: "Embedding Model", - }; - return friendlyNames[fieldName] || fieldName; - }); - NotificationManager.fromBackend(`Please fill in the following required fields: ${missingFields.join(", ")}`); - } else { - NotificationManager.fromBackend("Please fill in all required fields"); - } - }); + const submitValues: AutoRouterFormValues = { + ...currentFormValues, + auto_router_config: routerConfig, + model_type: "semantic_router", + }; + + await handleAddAutoRouterSubmit(submitValues, accessToken, form, () => + handleOk(submitValues), + ); } - }; + }); return ( - <> - Add Auto Router - - Create an auto router that automatically selects the best model based on request complexity or semantic matching. - + +

Add Auto Router

+

+ Create an auto router that automatically selects the best model based on + request complexity or semantic matching. +

- +
- Router Type - setRouterType(e.target.value)} - className="w-full" + + setRouterType(v as RouterType)} + className="w-full flex flex-col gap-4" > - - +