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 <yuneng-berri@users.noreply.github.com>
This commit is contained in:
cursor 2026-04-24 09:23:03 +00:00
parent ce1c604be2
commit a2f8c87e9f
No known key found for this signature in database
19 changed files with 2660 additions and 1421 deletions

View file

@ -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<ModelDashboardProps> = ({ premiumUser, teams }) => {
const { accessToken, token, userRole, userId: userID } = useAuthorized();
const [addModelForm] = Form.useForm();
const addModelForm = useForm<AddModelFormValues>({
defaultValues: {
model_mappings: [],
},
});
const [lastRefreshed, setLastRefreshed] = useState("");
const [providerModels, setProviderModels] = useState<Array<string>>([]);
const [selectedProvider, setSelectedProvider] = useState<Providers>(Providers.Anthropic);
@ -140,22 +147,6 @@ const ModelsAndEndpointsView: React.FC<ModelDashboardProps> = ({ 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<ModelDashboardProps> = ({ 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"}`,
);
}
};

View file

@ -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<AddModelFormValues>({ 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,

View file

@ -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<AddModelFormValues>;
handleOk: () => Promise<void>;
selectedProvider: Providers;
setSelectedProvider: (provider: Providers) => void;
@ -36,8 +101,6 @@ interface AddModelFormProps {
credentials: CredentialItem[];
}
const { Title, Link } = Typography;
const AddModelForm: React.FC<AddModelFormProps> = ({
form,
handleOk,
@ -53,9 +116,10 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
credentials,
}) => {
const [testMode, setTestMode] = useState<string>("chat");
const [isResultModalVisible, setIsResultModalVisible] = useState<boolean>(false);
const [isTestingConnection, setIsTestingConnection] = useState<boolean>(false);
// Using a unique ID to force the ConnectionErrorDisplay to remount and run a fresh test
const [isResultModalVisible, setIsResultModalVisible] =
useState<boolean>(false);
const [isTestingConnection, setIsTestingConnection] =
useState<boolean>(false);
const [connectionTestId, setConnectionTestId] = useState<string>("");
const { accessToken, userRole, premiumUser, userId } = useAuthorized();
@ -65,8 +129,10 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
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<AddModelFormProps> = ({
const [isTeamOnly, setIsTeamOnly] = useState<boolean>(false);
const [modelAccessGroups, setModelAccessGroups] = useState<string[]>([]);
// Team admin specific state
const [teamAdminSelectedTeam, setTeamAdminSelectedTeam] = useState<string | null>(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<AddModelFormProps> = ({
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<AddModelFormProps> = ({
const isAdmin = all_admin_roles.includes(userRole);
const isTeamAdmin = isUserTeamAdminForAnyTeam(teams, userId);
const onSubmit = form.handleSubmit(async () => {
await handleOk().then(() => {
setTeamAdminSelectedTeam(null);
});
});
return (
<>
<Title level={2}>Add Model</Title>
<FormProvider {...form}>
<h2 className="text-2xl font-semibold mb-4">Add Model</h2>
<Card>
<Form
form={form}
onFinish={async (values) => {
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 && (
<>
<Form.Item
label="Select Team"
name="team_id"
rules={[{ required: true, message: "Please select a team to continue" }]}
tooltip="Select the team for which you want to add this model"
<Card className="p-6">
<form onSubmit={onSubmit}>
{isTeamAdmin && !isAdmin && (
<>
<TeamSelectField
onTeamSelected={setTeamAdminSelectedTeam}
required
/>
{!teamAdminSelectedTeam && (
<Alert className="mb-4">
<Info className="h-4 w-4" />
<AlertTitle>Team Selection Required</AlertTitle>
<AlertDescription>
As a team admin, you need to select your team first before
adding models.
</AlertDescription>
</Alert>
)}
</>
)}
{(isAdmin || (isTeamAdmin && teamAdminSelectedTeam)) && (
<>
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label
className="col-span-10 pt-2"
title="E.g. OpenAI, Azure OpenAI, Anthropic, Bedrock, etc."
>
<TeamDropdown
onChange={(value) => {
setTeamAdminSelectedTeam(value);
}}
/>
</Form.Item>
{!teamAdminSelectedTeam && (
<Alert
message="Team Selection Required"
description="As a team admin, you need to select your team first before adding models."
type="info"
showIcon
className="mb-4"
/>
)}
</>
)}
{(isAdmin || (isTeamAdmin && teamAdminSelectedTeam)) && (
<>
<Form.Item
rules={[{ required: true, message: "Required" }]}
label="Provider:"
name="custom_llm_provider"
tooltip="E.g. OpenAI, Azure OpenAI, Anthropic, Bedrock, etc."
labelCol={{ span: 10 }}
labelAlign="left"
>
<AntdSelect
virtual={false}
showSearch
loading={isProviderMetadataLoading}
placeholder={isProviderMetadataLoading ? "Loading providers..." : "Select a provider"}
optionFilterProp="data-label"
onChange={(value) => {
setSelectedProvider(value as Providers);
setProviderModelsFn(value as Providers);
form.setFieldsValue({
custom_llm_provider: value,
});
form.setFieldsValue({
model: [],
model_name: undefined,
});
}}
>
{providerMetadataErrorText && sortedProviderMetadata.length === 0 && (
<AntdSelect.Option key="__error" value="">
{providerMetadataErrorText}
</AntdSelect.Option>
Provider <span className="text-destructive">*</span>
</Label>
<div className="col-span-14 space-y-1">
<Controller
control={form.control}
name="custom_llm_provider"
rules={{ required: "Required" }}
render={({ field }) => (
<Select
value={(field.value as string) || ""}
onValueChange={(value) => {
field.onChange(value);
setSelectedProvider(value as Providers);
setProviderModelsFn(value as Providers);
form.setValue("model", []);
form.setValue("model_name", undefined);
}}
>
<SelectTrigger data-testid="provider-select">
<SelectValue
placeholder={
isProviderMetadataLoading
? "Loading providers..."
: "Select a provider"
}
/>
</SelectTrigger>
<SelectContent>
{providerMetadataErrorText &&
sortedProviderMetadata.length === 0 && (
<SelectItem key="__error" value="__error">
{providerMetadataErrorText}
</SelectItem>
)}
{sortedProviderMetadata.map((providerInfo) => {
const displayName =
providerInfo.provider_display_name;
const providerKey = providerInfo.provider;
// referenced via data-label (search hint) only
void providerLogoMap[displayName];
return (
<SelectItem
key={providerKey}
value={providerKey}
>
<div className="flex items-center space-x-2">
<ProviderLogo
provider={providerKey}
className="w-5 h-5"
/>
<span>{displayName}</span>
</div>
</SelectItem>
);
})}
</SelectContent>
</Select>
)}
{sortedProviderMetadata.map((providerInfo) => {
const displayName = providerInfo.provider_display_name;
const providerKey = providerInfo.provider;
const logoSrc = providerLogoMap[displayName] ?? "";
return (
<AntdSelect.Option key={providerKey} value={providerKey} data-label={displayName}>
<div className="flex items-center space-x-2">
<ProviderLogo provider={providerKey} className="w-5 h-5" />
<span>{displayName}</span>
</div>
</AntdSelect.Option>
);
})}
</AntdSelect>
</Form.Item>
<LiteLLMModelNameField
selectedProvider={selectedProvider}
providerModels={providerModels}
getPlaceholder={getPlaceholder}
/>
{/* Conditionally Render "Public Model Name" */}
<ConditionalPublicModelName />
{/* Select Mode */}
<Form.Item label="Mode" name="mode" className="mb-1">
<AntdSelect
style={{ width: "100%" }}
value={testMode}
onChange={(value) => setTestMode(value)}
options={TEST_MODES}
/>
</Form.Item>
<Row>
<Col span={10}></Col>
<Col span={10}>
<Text className="mb-5 mt-1">
<strong>Optional</strong> - LiteLLM endpoint to use when health checking this model{" "}
<Link href="https://docs.litellm.ai/docs/proxy/health#health" target="_blank">
Learn more
</Link>
</Text>
</Col>
</Row>
{/* Credentials */}
<div className="mb-4">
<Typography.Text className="text-sm text-gray-500 mb-2">
Either select existing credentials OR enter new provider credentials below
</Typography.Text>
{form.formState.errors.custom_llm_provider?.message && (
<p className="text-sm text-destructive">
{String(
form.formState.errors.custom_llm_provider.message,
)}
</p>
)}
</div>
</div>
<Form.Item label="Existing Credentials" name="litellm_credential_name" initialValue={null}>
<AntdSelect
showSearch
placeholder="Select or search for existing credentials"
optionFilterProp="children"
filterOption={(input, option) => (option?.label ?? "").toLowerCase().includes(input.toLowerCase())}
options={[
{ value: null, label: "None" },
...credentials.map((credential) => ({
value: credential.credential_name,
label: credential.credential_name,
})),
]}
allowClear
<LiteLLMModelNameField
selectedProvider={selectedProvider}
providerModels={providerModels}
getPlaceholder={getPlaceholder}
/>
<ConditionalPublicModelName />
<div className="grid grid-cols-24 gap-2 mb-1 items-start">
<Label className="col-span-10 pt-2">Mode</Label>
<div className="col-span-14">
<Controller
control={form.control}
name="mode"
render={({ field }) => (
<Select
value={(field.value as string) || testMode}
onValueChange={(value) => {
field.onChange(value);
setTestMode(value);
}}
>
<SelectTrigger>
<SelectValue placeholder="Select a mode" />
</SelectTrigger>
<SelectContent>
{TEST_MODES.map((opt) => (
<SelectItem
key={opt.value}
value={opt.value as string}
>
{opt.label}
</SelectItem>
))}
</SelectContent>
</Select>
)}
/>
</Form.Item>
<Form.Item
noStyle
shouldUpdate={(prevValues, currentValues) =>
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 (
<>
<div className="flex items-center my-4">
<div className="flex-grow border-t border-gray-200"></div>
<span className="px-4 text-gray-500 text-sm">OR</span>
<div className="flex-grow border-t border-gray-200"></div>
</div>
<ProviderSpecificFields selectedProvider={selectedProvider} uploadProps={uploadProps} />
</>
);
}
return null;
}}
</Form.Item>
<div className="flex items-center my-4">
<div className="flex-grow border-t border-gray-200"></div>
<span className="px-4 text-gray-500 text-sm">Additional Model Info Settings</span>
<div className="flex-grow border-t border-gray-200"></div>
</div>
{/* Team-only Model Switch - Only show for proxy admins, not team admins */}
{(isAdmin || !isTeamAdmin) && (
<Form.Item
label="Team-BYOK Model"
tooltip="Only use this model + credential combination for this team. Useful when teams want to onboard their own OpenAI keys."
className="mb-4"
</div>
<div className="grid grid-cols-24 gap-2 mb-5">
<div className="col-span-10" />
<p className="col-span-14 text-sm mt-1">
<strong>Optional</strong> - LiteLLM endpoint to use when
health checking this model{" "}
<a
href="https://docs.litellm.ai/docs/proxy/health#health"
target="_blank"
rel="noopener noreferrer"
className="text-primary hover:text-primary/80 underline"
>
<Tooltip
Learn more
</a>
</p>
</div>
<div className="mb-4">
<p className="text-sm text-muted-foreground mb-2">
Either select existing credentials OR enter new provider
credentials below
</p>
</div>
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label className="col-span-10 pt-2">Existing Credentials</Label>
<div className="col-span-14">
<Controller
control={form.control}
name="litellm_credential_name"
defaultValue={null}
render={({ field }) => (
<Select
value={(field.value as string) || "__none"}
onValueChange={(v) =>
field.onChange(v === "__none" ? null : v)
}
>
<SelectTrigger>
<SelectValue placeholder="Select or search for existing credentials" />
</SelectTrigger>
<SelectContent>
<SelectItem value="__none">None</SelectItem>
{credentials.map((credential) => (
<SelectItem
key={credential.credential_name}
value={credential.credential_name}
>
{credential.credential_name}
</SelectItem>
))}
</SelectContent>
</Select>
)}
/>
</div>
</div>
<CredentialsGate
selectedProvider={selectedProvider}
uploadProps={uploadProps}
/>
<div className="flex items-center my-4">
<div className="flex-grow border-t border-border"></div>
<span className="px-4 text-muted-foreground text-sm">
Additional Model Info Settings
</span>
<div className="flex-grow border-t border-border"></div>
</div>
{(isAdmin || !isTeamAdmin) && (
<div className="grid grid-cols-24 gap-2 mb-4 items-center">
<Label
className="col-span-10"
title="Only use this model + credential combination for this team. Useful when teams want to onboard their own OpenAI keys."
>
Team-BYOK Model
</Label>
<div className="col-span-14">
<span
title={
!premiumUser
? "This is an enterprise-only feature. Upgrade to premium to restrict model+credential combinations to a specific team."
: ""
}
placement="top"
>
<Switch
checked={isTeamOnly}
onChange={(checked) => {
setIsTeamOnly(checked);
onCheckedChange={(checked) => {
setIsTeamOnly(!!checked);
if (!checked) {
form.setFieldValue("team_id", undefined);
form.setValue("team_id", undefined);
}
}}
disabled={!premiumUser}
/>
</Tooltip>
</Form.Item>
)}
</span>
</div>
</div>
)}
{/* Conditional Team Selection */}
{isTeamOnly && (isAdmin || !isTeamAdmin) && (
<Form.Item
label="Select Team"
name="team_id"
className="mb-4"
tooltip="Only keys for this team will be able to call this model."
rules={[
{
required: isTeamOnly && !isAdmin,
message: "Please select a team.",
},
]}
>
<TeamDropdown disabled={!premiumUser} />
</Form.Item>
)}
{isAdmin && (
<>
<Form.Item
label="Model Access Group"
name="model_access_group"
className="mb-4"
tooltip="Use model access groups to give users access to select models, and add new ones to the group over time."
>
<AntdSelect
mode="tags"
showSearch
placeholder="Select existing groups or type to create new ones"
optionFilterProp="children"
tokenSeparators={[","]}
options={modelAccessGroups.map((group) => ({
value: group,
label: group,
}))}
maxTagCount="responsive"
allowClear
/>
</Form.Item>
</>
)}
<AdvancedSettings
showAdvancedSettings={showAdvancedSettings}
setShowAdvancedSettings={setShowAdvancedSettings}
teams={teams}
guardrailsList={guardrailsList || []}
tagsList={tagsList || {}}
accessToken={accessToken || ""}
{isTeamOnly && (isAdmin || !isTeamAdmin) && (
<TeamSelectField
onTeamSelected={() => {}}
required={isTeamOnly && !isAdmin}
disabled={!premiumUser}
tooltip="Only keys for this team will be able to call this model."
/>
</>
)}
<div className="flex justify-between items-center mb-4">
<Tooltip title="Get help on our github">
<Typography.Link href="https://github.com/BerriAI/litellm/issues">Need Help?</Typography.Link>
</Tooltip>
<div className="space-x-2">
<Button data-testid="test-connect-btn" onClick={handleTestConnection} loading={isTestingConnection}>
Test Connect
</Button>
<Button data-testid="add-model-btn" htmlType="submit">Add Model</Button>
</div>
)}
{isAdmin && (
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label
className="col-span-10 pt-2"
title="Use model access groups to give users access to select models, and add new ones to the group over time."
>
Model Access Group
</Label>
<div className="col-span-14">
<Controller
control={form.control}
name="model_access_group"
defaultValue={[]}
render={({ field }) => (
<AccessGroupTagInput
value={(field.value as string[]) ?? []}
onChange={field.onChange}
options={modelAccessGroups}
/>
)}
/>
</div>
</div>
)}
<AdvancedSettings
showAdvancedSettings={showAdvancedSettings}
setShowAdvancedSettings={setShowAdvancedSettings}
teams={teams}
guardrailsList={guardrailsList || []}
tagsList={tagsList || {}}
accessToken={accessToken || ""}
/>
</>
)}
<div className="flex justify-between items-center mb-4">
<a
href="https://github.com/BerriAI/litellm/issues"
title="Get help on our github"
target="_blank"
rel="noopener noreferrer"
className="text-primary hover:text-primary/80 underline"
>
Need Help?
</a>
<div className="space-x-2">
<Button
type="button"
variant="outline"
data-testid="test-connect-btn"
onClick={handleTestConnection}
disabled={isTestingConnection}
>
{isTestingConnection && (
<Loader2 className="h-4 w-4 animate-spin mr-2" />
)}
Test Connect
</Button>
<Button data-testid="add-model-btn" type="submit">
Add Model
</Button>
</div>
</>
</Form>
</div>
</form>
</Card>
{/* Test Connection Results Modal */}
<Modal
title="Connection Test Results"
<Dialog
open={isResultModalVisible}
onCancel={() => {
setIsResultModalVisible(false);
setIsTestingConnection(false);
onOpenChange={(open) => {
if (!open) {
setIsResultModalVisible(false);
setIsTestingConnection(false);
}
}}
footer={[
<Button
key="close"
onClick={() => {
setIsResultModalVisible(false);
setIsTestingConnection(false);
}}
>
Close
</Button>,
]}
width={700}
>
{/* Only render the ConnectionErrorDisplay when modal is visible and we have a test ID */}
{isResultModalVisible && (
<ConnectionErrorDisplay
// The key prop tells React to create a fresh component instance when it changes
key={connectionTestId}
formValues={form.getFieldsValue()}
accessToken={accessToken}
testMode={testMode}
modelName={form.getFieldValue("model_name") || form.getFieldValue("model")}
onClose={() => {
setIsResultModalVisible(false);
setIsTestingConnection(false);
}}
onTestComplete={() => setIsTestingConnection(false)}
/>
)}
</Modal>
</>
<DialogContent className="sm:max-w-[700px]">
<DialogHeader>
<DialogTitle>Connection Test Results</DialogTitle>
</DialogHeader>
{isResultModalVisible && (
<ConnectionErrorDisplay
key={connectionTestId}
formValues={form.getValues()}
accessToken={accessToken}
testMode={testMode}
modelName={
(form.getValues("model_name") as string) ||
((Array.isArray(form.getValues("model"))
? (form.getValues("model") as string[])[0]
: (form.getValues("model") as string)) as string)
}
onClose={() => {
setIsResultModalVisible(false);
setIsTestingConnection(false);
}}
onTestComplete={() => setIsTestingConnection(false)}
/>
)}
<DialogFooter>
<Button
variant="outline"
onClick={() => {
setIsResultModalVisible(false);
setIsTestingConnection(false);
}}
>
<X className="h-4 w-4 mr-1" />
Close
</Button>
</DialogFooter>
</DialogContent>
</Dialog>
</FormProvider>
);
};
function TeamSelectField({
onTeamSelected,
required,
disabled,
tooltip,
}: {
onTeamSelected: (teamId: string | null) => void;
required?: boolean;
disabled?: boolean;
tooltip?: string;
}) {
const { control, formState } = useFormContext<AddModelFormValues>();
const error = (formState.errors as Record<string, { message?: string }>).team_id;
return (
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label
className="col-span-10 pt-2"
title={tooltip ?? "Select the team for which you want to add this model"}
>
Select Team
{required && <span className="text-destructive ml-1">*</span>}
</Label>
<div className="col-span-14 space-y-1">
<Controller
control={control}
name="team_id"
rules={required ? { required: "Please select a team to continue" } : {}}
render={({ field }) => (
<TeamDropdown
value={field.value as string | undefined}
onChange={(value) => {
field.onChange(value);
onTeamSelected(value || null);
}}
disabled={disabled}
/>
)}
/>
{error?.message && (
<p className="text-sm text-destructive">{String(error.message)}</p>
)}
</div>
</div>
);
}
function CredentialsGate({
selectedProvider,
uploadProps,
}: {
selectedProvider: Providers;
uploadProps: UploadProps;
}) {
const { control } = useFormContext<AddModelFormValues>();
const credentialName = useWatch({ control, name: "litellm_credential_name" });
if (credentialName) return null;
return (
<>
<div className="flex items-center my-4">
<div className="flex-grow border-t border-border"></div>
<span className="px-4 text-muted-foreground text-sm">OR</span>
<div className="flex-grow border-t border-border"></div>
</div>
<ProviderSpecificFields
selectedProvider={selectedProvider}
uploadProps={uploadProps}
/>
</>
);
}
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 (
<div className="space-y-2">
<div className="flex gap-2">
<Input
value={input}
onChange={(e) => 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("");
}
}}
/>
<Select
value=""
onValueChange={(v) => {
if (v) addValue(v);
}}
>
<SelectTrigger className="w-48">
<SelectValue placeholder="Pick existing" />
</SelectTrigger>
<SelectContent>
{remaining.length === 0 ? (
<div className="py-2 px-3 text-sm text-muted-foreground">
No options available
</div>
) : (
remaining.map((opt) => (
<SelectItem key={opt} value={opt}>
{opt}
</SelectItem>
))
)}
</SelectContent>
</Select>
</div>
{value.length > 0 && (
<div className="flex flex-wrap gap-1">
{value.map((v) => (
<Badge
key={v}
variant="secondary"
className="flex items-center gap-1"
>
{v}
<button
type="button"
onClick={() => onChange(value.filter((s) => s !== v))}
className="inline-flex items-center justify-center rounded-full hover:bg-muted-foreground/20"
aria-label={`Remove ${v}`}
>
<X size={12} />
</button>
</Badge>
))}
</div>
)}
</div>
);
}
export default AddModelForm;

View file

@ -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<AutoRouterFormValues>;
handleOk: (values: AutoRouterFormValues) => void | Promise<void>;
accessToken: string;
userRole: string;
}
@ -28,26 +75,24 @@ interface ComplexityTiers {
REASONING: string;
}
const { Title, Link } = Typography;
const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, accessToken, userRole }) => {
// State for connection testing
const [isResultModalVisible, setIsResultModalVisible] = useState<boolean>(false);
const [isTestingConnection, setIsTestingConnection] = useState<boolean>(false);
const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({
form,
handleOk,
accessToken,
userRole,
}) => {
const [isResultModalVisible, setIsResultModalVisible] =
useState<boolean>(false);
const [isTestingConnection, setIsTestingConnection] =
useState<boolean>(false);
const [connectionTestId, setConnectionTestId] = useState<string>("");
const [modelAccessGroups, setModelAccessGroups] = useState<string[]>([]);
const [modelInfo, setModelInfo] = useState<ModelGroup[]>([]);
const [showCustomDefaultModel, setShowCustomDefaultModel] = useState<boolean>(false);
const [showCustomEmbeddingModel, setShowCustomEmbeddingModel] = useState<boolean>(false);
// Router type state - default to complexity router
const [routerType, setRouterType] = useState<RouterType>("complexity");
// Semantic router config (existing)
// eslint-disable-next-line @typescript-eslint/no-explicit-any
const [routerConfig, setRouterConfig] = useState<any>(null);
// Complexity router config (new)
const [complexityTiers, setComplexityTiers] = useState<ComplexityTiers>({
SIMPLE: "",
MEDIUM: "",
@ -57,8 +102,18 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ 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<AddAutoRouterTabProps> = ({ 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<AddAutoRouterTabProps> = ({ 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<AddAutoRouterTabProps> = ({ 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 (
<>
<Title level={2}>Add Auto Router</Title>
<Text className="text-gray-600 mb-6">
Create an auto router that automatically selects the best model based on request complexity or semantic matching.
</Text>
<FormProvider {...form}>
<h2 className="text-2xl font-semibold mb-2">Add Auto Router</h2>
<p className="text-muted-foreground mb-6">
Create an auto router that automatically selects the best model based on
request complexity or semantic matching.
</p>
<Card className="mb-4">
<Card className="p-6 mb-4">
<div className="mb-4">
<Text className="text-sm font-medium mb-2 block">Router Type</Text>
<Radio.Group
value={routerType}
onChange={(e) => setRouterType(e.target.value)}
className="w-full"
<Label className="text-sm font-medium mb-2 block">Router Type</Label>
<RadioGroup
value={routerType}
onValueChange={(v) => setRouterType(v as RouterType)}
className="w-full flex flex-col gap-4"
>
<Space direction="vertical" className="w-full">
<Radio value="complexity" className="w-full">
<label className="flex items-start gap-3 cursor-pointer">
<RadioGroupItem value="complexity" className="mt-1" />
<div>
<div className="flex items-center gap-2">
{/* eslint-disable-next-line litellm-ui/no-raw-tailwind-colors */}
<ThunderboltOutlined className="text-yellow-500" />
<span className="font-medium">Complexity Router</span>
<Badge
count="Recommended"
style={{
backgroundColor: '#52c41a',
fontSize: '10px',
padding: '0 6px',
}}
/>
<Badge variant="secondary">Recommended</Badge>
</div>
<div className="text-xs text-gray-500 ml-6 mt-1">
Automatically routes based on request complexity. No training data needed — just pick 4 models and go.
<br />
<span className="text-green-600">✓ Zero API calls</span> · <span className="text-green-600">✓ &lt;1ms latency</span> · <span className="text-green-600">✓ No cost</span>
<div className="text-xs text-muted-foreground ml-0 mt-1">
Automatically routes based on request complexity. No
training data needed — just pick 4 models and go.
</div>
</Radio>
<Radio value="semantic" className="w-full mt-2">
</div>
</label>
<label className="flex items-start gap-3 cursor-pointer">
<RadioGroupItem value="semantic" className="mt-1" />
<div>
<div className="flex items-center gap-2">
{/* eslint-disable-next-line litellm-ui/no-raw-tailwind-colors */}
<BranchesOutlined className="text-blue-500" />
<span className="font-medium">Semantic Router</span>
</div>
<div className="text-xs text-gray-500 ml-6 mt-1">
Routes based on semantic similarity to example utterances. Requires embedding model and training examples.
<div className="text-xs text-muted-foreground ml-0 mt-1">
Routes based on semantic similarity to example utterances.
Requires embedding model and training examples.
</div>
</Radio>
</Space>
</Radio.Group>
</div>
</label>
</RadioGroup>
</div>
</Card>
<Card>
<Form
form={form}
onFinish={handleAutoRouterSubmit}
labelCol={{ span: 10 }}
wrapperCol={{ span: 16 }}
labelAlign="left"
>
{/* Auto Router Name */}
<Form.Item
rules={[{ required: true, message: "Auto router name is required" }]}
label="Auto Router Name"
name="auto_router_name"
tooltip="Unique name for this auto router configuration"
labelCol={{ span: 10 }}
labelAlign="left"
>
<TextInput placeholder="e.g., smart_router, auto_router_1" />
</Form.Item>
<Card className="p-6">
<form onSubmit={handleAutoRouterSubmit}>
<AutoRouterNameField />
{/* Conditional rendering based on router type */}
{routerType === "complexity" ? (
/* Complexity Router Configuration */
<div className="w-full mb-4">
<ComplexityRouterConfig
modelInfo={modelInfo}
value={complexityTiers}
onChange={(tiers) => {
setComplexityTiers(tiers);
}}
onChange={(tiers) => setComplexityTiers(tiers)}
/>
</div>
) : (
/* Semantic Router Configuration (existing) */
<>
{/* Router Configuration Builder */}
<div className="w-full mb-4">
<RouterConfigBuilder
modelInfo={modelInfo}
value={routerConfig}
onChange={(config) => {
setRouterConfig(config);
form.setFieldValue("auto_router_config", config);
}}
/>
</div>
{/* Auto Router Default Model */}
<Form.Item
rules={[{ required: routerType === "semantic", message: "Default model is required" }]}
label="Default Model"
name="auto_router_default_model"
tooltip="Fallback model to use when auto routing logic cannot determine the best model"
labelCol={{ span: 10 }}
labelAlign="left"
>
<AntdSelect
placeholder="Select a default model"
onChange={(value) => {
setShowCustomDefaultModel(value === "custom");
}}
options={[
...Array.from(new Set(modelInfo.map((option) => option.model_group))).map((model_group) => ({
value: model_group,
label: model_group,
})),
{ value: "custom", label: "Enter custom model name" },
]}
style={{ width: "100%" }}
showSearch={true}
/>
</Form.Item>
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label
className="col-span-10 pt-2"
title="Fallback model to use when auto routing logic cannot determine the best model"
>
Default Model{" "}
<span className="text-destructive">*</span>
</Label>
<div className="col-span-14">
<Controller
control={form.control}
name="auto_router_default_model"
rules={{
required:
routerType === "semantic"
? "Default model is required"
: false,
}}
render={({ field }) => (
<Select
value={(field.value as string) || ""}
onValueChange={field.onChange}
>
<SelectTrigger>
<SelectValue placeholder="Select a default model" />
</SelectTrigger>
<SelectContent>
{Array.from(
new Set(
modelInfo.map((option) => option.model_group),
),
).map((model_group) => (
<SelectItem key={model_group} value={model_group}>
{model_group}
</SelectItem>
))}
<SelectItem value="custom">
Enter custom model name
</SelectItem>
</SelectContent>
</Select>
)}
/>
</div>
</div>
{/* Auto Router Embedding Model */}
<Form.Item
label="Embedding Model"
name="auto_router_embedding_model"
tooltip="Optional: Embedding model to use for semantic routing decisions"
labelCol={{ span: 10 }}
labelAlign="left"
>
<AntdSelect
value={form.getFieldValue("auto_router_embedding_model")}
placeholder="Select an embedding model (optional)"
onChange={(value) => {
setShowCustomEmbeddingModel(value === "custom");
form.setFieldValue("auto_router_embedding_model", value);
}}
options={[
...Array.from(new Set(modelInfo.map((option) => option.model_group))).map((model_group) => ({
value: model_group,
label: model_group,
})),
{ value: "custom", label: "Enter custom model name" },
]}
style={{ width: "100%" }}
showSearch={true}
allowClear
/>
</Form.Item>
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label
className="col-span-10 pt-2"
title="Optional: Embedding model to use for semantic routing decisions"
>
Embedding Model
</Label>
<div className="col-span-14">
<Controller
control={form.control}
name="auto_router_embedding_model"
render={({ field }) => (
<Select
value={(field.value as string) || ""}
onValueChange={field.onChange}
>
<SelectTrigger>
<SelectValue placeholder="Select an embedding model (optional)" />
</SelectTrigger>
<SelectContent>
{Array.from(
new Set(
modelInfo.map((option) => option.model_group),
),
).map((model_group) => (
<SelectItem key={model_group} value={model_group}>
{model_group}
</SelectItem>
))}
<SelectItem value="custom">
Enter custom model name
</SelectItem>
</SelectContent>
</Select>
)}
/>
</div>
</div>
</>
)}
<div className="flex items-center my-4">
<div className="flex-grow border-t border-gray-200"></div>
<span className="px-4 text-gray-500 text-sm">Additional Settings</span>
<div className="flex-grow border-t border-gray-200"></div>
<div className="flex-grow border-t border-border"></div>
<span className="px-4 text-muted-foreground text-sm">
Additional Settings
</span>
<div className="flex-grow border-t border-border"></div>
</div>
{/* Model Access Groups - Admin only */}
{isAdmin && (
<Form.Item
label="Model Access Group"
name="model_access_group"
className="mb-4"
tooltip="Use model access groups to control who can access this auto router"
>
<AntdSelect
mode="tags"
showSearch
placeholder="Select existing groups or type to create new ones"
optionFilterProp="children"
tokenSeparators={[","]}
options={modelAccessGroups.map((group) => ({
value: group,
label: group,
}))}
maxTagCount="responsive"
allowClear
/>
</Form.Item>
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label
className="col-span-10 pt-2"
title="Use model access groups to control who can access this auto router"
>
Model Access Group
</Label>
<div className="col-span-14">
<Controller
control={form.control}
name="model_access_group"
defaultValue={[]}
render={({ field }) => (
<GroupTagInput
value={(field.value as string[]) ?? []}
onChange={field.onChange}
options={modelAccessGroups}
/>
)}
/>
</div>
</div>
)}
<div className="flex justify-between items-center mb-4">
<Tooltip title="Get help on our github">
<Typography.Link href="https://github.com/BerriAI/litellm/issues">Need Help?</Typography.Link>
</Tooltip>
<a
href="https://github.com/BerriAI/litellm/issues"
title="Get help on our github"
target="_blank"
rel="noopener noreferrer"
className="text-primary hover:text-primary/80 underline"
>
Need Help?
</a>
<div className="space-x-2">
<Button onClick={handleTestConnection} loading={isTestingConnection}>
<Button
type="button"
variant="outline"
onClick={handleTestConnection}
disabled={isTestingConnection}
>
{isTestingConnection && (
<Loader2 className="h-4 w-4 animate-spin mr-2" />
)}
Test Connection
</Button>
<Button
type="primary"
onClick={() => {
console.log("Add Auto Router button clicked!");
handleAutoRouterSubmit();
}}
>
Add Auto Router
</Button>
<Button type="submit">Add Auto Router</Button>
</div>
</div>
</Form>
</form>
</Card>
{/* Test Connection Results Modal */}
<Modal
title="Connection Test Results"
<Dialog
open={isResultModalVisible}
onCancel={() => {
setIsResultModalVisible(false);
setIsTestingConnection(false);
onOpenChange={(open) => {
if (!open) {
setIsResultModalVisible(false);
setIsTestingConnection(false);
}
}}
footer={[
<Button
key="close"
onClick={() => {
setIsResultModalVisible(false);
setIsTestingConnection(false);
}}
>
Close
</Button>,
]}
width={700}
>
{/* Only render the ConnectionErrorDisplay when modal is visible and we have a test ID */}
{isResultModalVisible && (
<ConnectionErrorDisplay
key={connectionTestId}
formValues={form.getFieldsValue()}
accessToken={accessToken}
testMode="chat"
modelName={form.getFieldValue("auto_router_name")}
onClose={() => {
setIsResultModalVisible(false);
setIsTestingConnection(false);
}}
onTestComplete={() => setIsTestingConnection(false)}
/>
)}
</Modal>
</>
<DialogContent className="sm:max-w-[700px]">
<DialogHeader>
<DialogTitle>Connection Test Results</DialogTitle>
</DialogHeader>
{isResultModalVisible && (
<ConnectionErrorDisplay
key={connectionTestId}
formValues={form.getValues()}
accessToken={accessToken}
testMode="chat"
modelName={form.getValues("auto_router_name")}
onClose={() => {
setIsResultModalVisible(false);
setIsTestingConnection(false);
}}
onTestComplete={() => setIsTestingConnection(false)}
/>
)}
<DialogFooter>
<Button
variant="outline"
onClick={() => {
setIsResultModalVisible(false);
setIsTestingConnection(false);
}}
>
<X className="h-4 w-4 mr-1" />
Close
</Button>
</DialogFooter>
</DialogContent>
</Dialog>
</FormProvider>
);
};
function AutoRouterNameField() {
const { control, formState } = useFormContext<AutoRouterFormValues>();
const error = (formState.errors as Record<string, { message?: string }>)
.auto_router_name;
return (
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label
className="col-span-10 pt-2"
title="Unique name for this auto router configuration"
>
Auto Router Name <span className="text-destructive">*</span>
</Label>
<div className="col-span-14 space-y-1">
<Controller
control={control}
name="auto_router_name"
rules={{ required: "Auto router name is required" }}
render={({ field }) => (
<Input
placeholder="e.g., smart_router, auto_router_1"
value={(field.value as string) ?? ""}
onChange={(e) => field.onChange(e.target.value)}
/>
)}
/>
{error?.message && (
<p className="text-sm text-destructive">{String(error.message)}</p>
)}
</div>
</div>
);
}
function GroupTagInput({
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 (
<div className="space-y-2">
<div className="flex gap-2">
<Input
value={input}
onChange={(e) => 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("");
}
}}
/>
<Select
value=""
onValueChange={(v) => {
if (v) addValue(v);
}}
>
<SelectTrigger className="w-48">
<SelectValue placeholder="Pick existing" />
</SelectTrigger>
<SelectContent>
{remaining.length === 0 ? (
<div className="py-2 px-3 text-sm text-muted-foreground">
No options available
</div>
) : (
remaining.map((opt) => (
<SelectItem key={opt} value={opt}>
{opt}
</SelectItem>
))
)}
</SelectContent>
</Select>
</div>
{value.length > 0 && (
<div className="flex flex-wrap gap-1">
{value.map((v) => (
<Badge
key={v}
variant="secondary"
className="flex items-center gap-1"
>
{v}
<button
type="button"
onClick={() => onChange(value.filter((s) => s !== v))}
className="inline-flex items-center justify-center rounded-full hover:bg-muted-foreground/20"
aria-label={`Remove ${v}`}
>
<X size={12} />
</button>
</Badge>
))}
</div>
)}
</div>
);
}
export default AddAutoRouterTab;

View file

@ -1,13 +1,14 @@
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { render, renderHook, screen, waitFor } from "@testing-library/react";
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 AddModelTab from "./add_model_tab";
import type { AddModelFormValues } from "./AddModelForm";
import type { UploadProps } from "./add_model_upload_types";
vi.mock("../molecules/models/ProviderLogo", () => ({
ProviderLogo: ({ provider, className }: { provider: string; className?: string }) => (
@ -64,9 +65,16 @@ vi.mock("@/app/(dashboard)/hooks/providers/useProviderFields", () => ({
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
default: vi.fn().mockReturnValue({
isLoading: false,
isAuthorized: true,
token: "test-token",
accessToken: "test-access-token",
userId: "user-1",
userEmail: "test@example.com",
userRole: "Admin",
premiumUser: true,
disabledPersonalKeyCreation: false,
showSSOBanner: false,
}),
}));
@ -85,8 +93,10 @@ const createQueryClient = () =>
});
const createTestProps = () => {
const { result } = renderHook(() => Form.useForm());
const [form] = result.current;
const { result } = renderHook(() =>
useForm<AddModelFormValues>({ defaultValues: { model_mappings: [] } }),
);
const form = result.current;
const handleOk = vi.fn();
const setSelectedProvider = vi.fn();
@ -111,7 +121,8 @@ const createTestProps = () => {
created_at: "2024-01-01T00:00:00Z",
keys: [],
members_with_roles: [],
},
spend: 0,
} as Team,
];
const credentials: CredentialItem[] = [
@ -125,10 +136,7 @@ const createTestProps = () => {
},
];
const uploadProps: UploadProps = {
beforeUpload: () => false,
showUploadList: false,
};
const uploadProps: UploadProps = {};
return {
form,
@ -259,7 +267,6 @@ describe("Add Model Tab", () => {
</QueryClientProvider>,
);
// Wait for async operations to complete and buttons to appear
await waitFor(
async () => {
const testConnectButtons = await screen.findAllByRole("button", { name: "Test Connect" });
@ -296,20 +303,15 @@ describe("Add Model Tab", () => {
</QueryClientProvider>,
);
// Wait for component to load
await screen.findByText("Provider");
// Find the team-BYOK switch by its role
const teamSwitch = screen.getByRole("switch");
expect(teamSwitch).toBeInTheDocument();
// Initially, team selection should not be visible
expect(screen.queryByText("Select Team")).not.toBeInTheDocument();
// Click the switch to enable team-only mode
await userEvent.click(teamSwitch!);
// Now team selection should be visible
expect(await screen.findByText("Select Team")).toBeInTheDocument();
});
});

View file

@ -4,19 +4,20 @@ import {
TabsList,
TabsTrigger,
} from "@/components/ui/tabs";
import type { FormInstance } from "antd";
import { Form } from "antd";
import type { UploadProps } from "antd/es/upload";
import React from "react";
import { useForm, UseFormReturn } from "react-hook-form";
import type { Team } from "../key_team_helpers/key_list";
import { type CredentialItem } from "../networking";
import { Providers } from "../provider_info_helpers";
import AddAutoRouterTab from "./add_auto_router_tab";
import AddModelForm from "./AddModelForm";
import AddAutoRouterTab, {
type AutoRouterFormValues,
} from "./add_auto_router_tab";
import AddModelForm, { type AddModelFormValues } from "./AddModelForm";
import { handleAddAutoRouterSubmit } from "./handle_add_auto_router_submit";
import type { UploadProps } from "./add_model_upload_types";
interface AddModelTabProps {
form: FormInstance;
form: UseFormReturn<AddModelFormValues>;
// eslint-disable-next-line @typescript-eslint/no-explicit-any
handleOk: (values?: any) => Promise<void>;
selectedProvider: Providers;
@ -49,22 +50,22 @@ const AddModelTab: React.FC<AddModelTabProps> = ({
accessToken,
userRole,
}) => {
const [autoRouterForm] = Form.useForm();
const autoRouterForm = useForm<AutoRouterFormValues>({
defaultValues: {
auto_router_name: "",
auto_router_default_model: "",
auto_router_embedding_model: "",
model_access_group: [],
},
});
const handleAutoRouterOk = () => {
autoRouterForm
.validateFields()
.then((values) => {
handleAddAutoRouterSubmit(
values,
accessToken,
autoRouterForm,
handleOk,
);
})
.catch((error) => {
console.error("Validation failed:", error);
});
const handleAutoRouterOk = async (values: AutoRouterFormValues) => {
await handleAddAutoRouterSubmit(
values,
accessToken,
autoRouterForm,
() => handleOk(),
);
};
return (

View file

@ -0,0 +1,25 @@
/**
* Minimal upload type compatible with the antd `UploadProps` shape used
* by the out-of-scope parent `ModelsAndEndpointsView.tsx` and by
* `model_add/AddCredentialModal.tsx` / `model_add/EditCredentialModal.tsx`,
* which still import the antd `UploadProps` type directly. We only depend
* on the fields that `provider_specific_fields.tsx` actually reads
* (`onChange` with a `file` payload exposing `name` / `status` / `type`),
* so the `onChange` signature is typed with an `any` info payload to
* accept both the antd `UploadChangeParam<UploadFile<any>>` shape and our
* minimal shape — this can be tightened once all callers migrate off antd.
*/
export interface MinimalUploadFileInfo {
name: string;
status?: string;
type?: string;
}
export interface MinimalUploadChangeParam {
file: MinimalUploadFileInfo;
}
export interface UploadProps {
// eslint-disable-next-line @typescript-eslint/no-explicit-any
onChange?: (info: any) => void;
}

View file

@ -1,32 +1,43 @@
import { act, fireEvent, render, waitFor } from "@testing-library/react";
import { FormProvider, useForm } from "react-hook-form";
import { beforeEach, describe, expect, it, vi } from "vitest";
import AdvancedSettings from "./advanced_settings";
function Wrapper({ children }: { children: React.ReactNode }) {
const form = useForm({ defaultValues: {} });
return <FormProvider {...form}>{children}</FormProvider>;
}
describe("AdvancedSettings", () => {
beforeEach(() => {
vi.clearAllMocks();
});
it("should render", () => {
render(
<AdvancedSettings
showAdvancedSettings={true}
setShowAdvancedSettings={() => {}}
guardrailsList={[]}
tagsList={{}}
accessToken="test-token"
/>,
<Wrapper>
<AdvancedSettings
showAdvancedSettings={true}
setShowAdvancedSettings={() => {}}
guardrailsList={[]}
tagsList={{}}
accessToken="test-token"
/>
</Wrapper>,
);
});
it("should render tags list", async () => {
const { getByText } = render(
<AdvancedSettings
showAdvancedSettings={true}
setShowAdvancedSettings={() => {}}
guardrailsList={[]}
tagsList={{}}
accessToken="test-token"
/>,
<Wrapper>
<AdvancedSettings
showAdvancedSettings={true}
setShowAdvancedSettings={() => {}}
guardrailsList={[]}
tagsList={{}}
accessToken="test-token"
/>
</Wrapper>,
);
fireEvent.click(getByText("Advanced Settings"));
await waitFor(() => {
@ -36,13 +47,15 @@ describe("AdvancedSettings", () => {
it("should render the litellm params", async () => {
const { getByText } = render(
<AdvancedSettings
showAdvancedSettings={true}
setShowAdvancedSettings={() => {}}
guardrailsList={[]}
tagsList={{}}
accessToken="test-token"
/>,
<Wrapper>
<AdvancedSettings
showAdvancedSettings={true}
setShowAdvancedSettings={() => {}}
guardrailsList={[]}
tagsList={{}}
accessToken="test-token"
/>
</Wrapper>,
);
act(() => {
fireEvent.click(getByText("Advanced Settings"));

View file

@ -1,22 +1,29 @@
import React from "react";
import { Form, Switch, Select, Tooltip } from "antd";
// eslint-disable-next-line litellm-ui/no-banned-ui-imports
import { Controller, useFormContext, useWatch } from "react-hook-form";
import {
Text,
Accordion,
AccordionHeader,
AccordionBody,
TextInput,
} from "@tremor/react";
import { Row, Col, Typography } from "antd";
import TextArea from "antd/es/input/TextArea";
import { Info as InfoCircleOutlined } from "lucide-react";
AccordionContent,
AccordionItem,
AccordionTrigger,
} from "@/components/ui/accordion";
import { Label } from "@/components/ui/label";
import { Switch } from "@/components/ui/switch";
import { Input } from "@/components/ui/input";
import { Textarea } from "@/components/ui/textarea";
import { Badge } from "@/components/ui/badge";
import {
Select,
SelectContent,
SelectItem,
SelectTrigger,
SelectValue,
} from "@/components/ui/select";
import { Info as InfoCircleOutlined, X } from "lucide-react";
import { Team } from "../key_team_helpers/key_list";
import CacheControlSettings from "./cache_control_settings";
import VectorStoreSelector from "../vector_store_management/VectorStoreSelector";
import { Tag } from "../tag_management/types";
import { formItemValidateJSON } from "../../utils/textUtils";
const { Link } = Typography;
import { validateJsonValue } from "../../utils/textUtils";
interface AdvancedSettingsProps {
showAdvancedSettings: boolean;
@ -27,278 +34,547 @@ interface AdvancedSettingsProps {
accessToken: string;
}
function TagMultiSelect({
value,
onChange,
options,
placeholder,
}: {
value: string[];
onChange: (next: string[]) => void;
options: { value: string; label: string }[];
placeholder: string;
}) {
const selected = value ?? [];
const remaining = options.filter(
(o) => !selected.includes(o.value as string),
);
const [input, setInput] = React.useState("");
const addValue = (next: string) => {
const trimmed = next.trim();
if (!trimmed) return;
if (selected.includes(trimmed)) return;
onChange([...selected, trimmed]);
};
return (
<div className="space-y-2">
<div className="flex gap-2">
<Input
value={input}
onChange={(e) => setInput(e.target.value)}
placeholder={placeholder}
onKeyDown={(e) => {
if (e.key === "Enter" || e.key === ",") {
e.preventDefault();
addValue(input);
setInput("");
}
}}
/>
<Select
value=""
onValueChange={(v) => {
if (v) {
addValue(v);
}
}}
>
<SelectTrigger className="w-48">
<SelectValue placeholder="Pick existing" />
</SelectTrigger>
<SelectContent>
{remaining.length === 0 ? (
<div className="py-2 px-3 text-sm text-muted-foreground">
No options available
</div>
) : (
remaining.map((opt) => (
<SelectItem key={opt.value} value={opt.value}>
{opt.label}
</SelectItem>
))
)}
</SelectContent>
</Select>
</div>
{selected.length > 0 && (
<div className="flex flex-wrap gap-1">
{selected.map((v) => (
<Badge
key={v}
variant="secondary"
className="flex items-center gap-1"
>
{v}
<button
type="button"
onClick={() => onChange(selected.filter((s) => s !== v))}
className="inline-flex items-center justify-center rounded-full hover:bg-muted-foreground/20"
aria-label={`Remove ${v}`}
>
<X size={12} />
</button>
</Badge>
))}
</div>
)}
</div>
);
}
const AdvancedSettings: React.FC<AdvancedSettingsProps> = ({
showAdvancedSettings,
setShowAdvancedSettings,
teams,
guardrailsList,
tagsList,
accessToken,
}) => {
const [form] = Form.useForm();
const [customPricing, setCustomPricing] = React.useState(false);
const [pricingModel, setPricingModel] = React.useState<"per_token" | "per_second">("per_token");
const [showCacheControl, setShowCacheControl] = React.useState(false);
const { control, getValues, setValue, formState } = useFormContext();
const customPricing = !!useWatch({ control, name: "custom_pricing" });
const pricingModel =
(useWatch({ control, name: "pricing_model" }) as
| "per_token"
| "per_second"
| undefined) ?? "per_token";
const showCacheControl = !!useWatch({ control, name: "cache_control" });
// Add validation function for numbers
const validateNumber = (_: any, value: string) => {
if (!value) {
return Promise.resolve();
}
if (isNaN(Number(value)) || Number(value) < 0) {
return Promise.reject("Please enter a valid positive number");
}
return Promise.resolve();
};
// Handle custom pricing changes
const handleCustomPricingChange = (checked: boolean) => {
setCustomPricing(checked);
setValue("custom_pricing", checked);
if (!checked) {
// Clear pricing fields when disabled
form.setFieldsValue({
input_cost_per_token: undefined,
output_cost_per_token: undefined,
input_cost_per_second: undefined,
});
setValue("input_cost_per_token", undefined);
setValue("output_cost_per_token", undefined);
setValue("input_cost_per_second", undefined);
}
};
const handlePassThroughChange = (checked: boolean) => {
const currentParams = form.getFieldValue("litellm_extra_params");
setValue("use_in_pass_through", checked);
const currentParams = getValues("litellm_extra_params") as
| string
| undefined;
try {
let paramsObj = currentParams ? JSON.parse(currentParams) : {};
const paramsObj = currentParams ? JSON.parse(currentParams) : {};
if (checked) {
paramsObj.use_in_pass_through = true;
} else {
delete paramsObj.use_in_pass_through;
}
// Only set the field value if there are remaining parameters
if (Object.keys(paramsObj).length > 0) {
form.setFieldValue("litellm_extra_params", JSON.stringify(paramsObj, null, 2));
setValue(
"litellm_extra_params",
JSON.stringify(paramsObj, null, 2),
);
} else {
form.setFieldValue("litellm_extra_params", "");
setValue("litellm_extra_params", "");
}
} catch (error) {
// If JSON parsing fails, only create new object if checked is true
} catch {
if (checked) {
form.setFieldValue("litellm_extra_params", JSON.stringify({ use_in_pass_through: true }, null, 2));
setValue(
"litellm_extra_params",
JSON.stringify({ use_in_pass_through: true }, null, 2),
);
} else {
form.setFieldValue("litellm_extra_params", "");
setValue("litellm_extra_params", "");
}
}
};
const handleCacheControlChange = (checked: boolean) => {
setShowCacheControl(checked);
setValue("cache_control", checked);
if (!checked) {
const currentParams = form.getFieldValue("litellm_extra_params");
const currentParams = getValues("litellm_extra_params") as
| string
| undefined;
try {
let paramsObj = currentParams ? JSON.parse(currentParams) : {};
const paramsObj = currentParams ? JSON.parse(currentParams) : {};
delete paramsObj.cache_control_injection_points;
if (Object.keys(paramsObj).length > 0) {
form.setFieldValue("litellm_extra_params", JSON.stringify(paramsObj, null, 2));
setValue(
"litellm_extra_params",
JSON.stringify(paramsObj, null, 2),
);
} else {
form.setFieldValue("litellm_extra_params", "");
setValue("litellm_extra_params", "");
}
} catch (error) {
form.setFieldValue("litellm_extra_params", "");
} catch {
setValue("litellm_extra_params", "");
}
}
};
const validateNumber = (value: unknown): true | string => {
if (value === undefined || value === null || value === "") {
return true;
}
const num = Number(value);
if (Number.isNaN(num) || num < 0) {
return "Please enter a valid positive number";
}
return true;
};
const errors = formState.errors as Record<
string,
{ message?: string } | undefined
>;
return (
<>
<Accordion className="mt-2 mb-4">
<AccordionHeader>
<Accordion type="single" collapsible className="mt-2 mb-4">
<AccordionItem value="advanced-settings">
<AccordionTrigger>
<b>Advanced Settings</b>
</AccordionHeader>
<AccordionBody>
<div className="bg-white rounded-lg">
<Form.Item label="Custom Pricing" name="custom_pricing" valuePropName="checked" className="mb-4">
<Switch onChange={handleCustomPricingChange} className="bg-gray-600" />
</Form.Item>
</AccordionTrigger>
<AccordionContent>
<div className="bg-background rounded-lg">
<div className="grid grid-cols-24 gap-2 mb-4 items-center">
<Label className="col-span-10">Custom Pricing</Label>
<div className="col-span-14">
<Controller
control={control}
name="custom_pricing"
render={({ field }) => (
<Switch
checked={!!field.value}
onCheckedChange={(checked) => {
field.onChange(checked);
handleCustomPricingChange(checked);
}}
/>
)}
/>
</div>
</div>
<Form.Item
label={
<span>
Attached Knowledge Bases (RAG){" "}
<Tooltip title="Vector stores to use for RAG. Every request to this model will automatically retrieve context from these knowledge bases.">
<a
href="https://docs.litellm.ai/docs/completion/knowledgebase"
target="_blank"
rel="noopener noreferrer"
onClick={(e) => e.stopPropagation()}
>
<InfoCircleOutlined style={{ marginLeft: "4px" }} />
</a>
</Tooltip>
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label className="col-span-10 pt-2">
Attached Knowledge Bases (RAG){" "}
<span
className="ml-1"
title="Vector stores to use for RAG. Every request to this model will automatically retrieve context from these knowledge bases."
>
<a
href="https://docs.litellm.ai/docs/completion/knowledgebase"
target="_blank"
rel="noopener noreferrer"
onClick={(e) => e.stopPropagation()}
>
<InfoCircleOutlined
className="h-3.5 w-3.5 inline-block"
style={{ marginLeft: "4px" }}
/>
</a>
</span>
}
name="vector_store_ids"
className="mt-4"
help="Select vector stores to attach. Requests to this model will automatically use these for RAG. Set up vector stores in Tools > Vector Stores."
>
<VectorStoreSelector
onChange={() => {}}
accessToken={accessToken}
placeholder="Select knowledge bases (optional)"
/>
</Form.Item>
</Label>
<div className="col-span-14">
<Controller
control={control}
name="vector_store_ids"
defaultValue={[]}
render={({ field }) => (
<VectorStoreSelector
value={(field.value as string[]) ?? []}
onChange={field.onChange}
accessToken={accessToken}
placeholder="Select knowledge bases (optional)"
/>
)}
/>
<p className="text-sm text-muted-foreground mt-1">
Select vector stores to attach. Requests to this model will
automatically use these for RAG. Set up vector stores in
Tools &gt; Vector Stores.
</p>
</div>
</div>
<Form.Item
label={
<span>
Guardrails{" "}
<Tooltip title="Apply safety guardrails to this key to filter content or enforce policies">
<a
href="https://docs.litellm.ai/docs/proxy/guardrails/quick_start"
target="_blank"
rel="noopener noreferrer"
onClick={(e) => e.stopPropagation()} // Prevent accordion from collapsing when clicking link
>
<InfoCircleOutlined style={{ marginLeft: "4px" }} />
</a>
</Tooltip>
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label className="col-span-10 pt-2">
Guardrails{" "}
<span
className="ml-1"
title="Apply safety guardrails to this key to filter content or enforce policies"
>
<a
href="https://docs.litellm.ai/docs/proxy/guardrails/quick_start"
target="_blank"
rel="noopener noreferrer"
onClick={(e) => e.stopPropagation()}
>
<InfoCircleOutlined
className="h-3.5 w-3.5 inline-block"
style={{ marginLeft: "4px" }}
/>
</a>
</span>
}
name="guardrails"
className="mt-4"
help="Select existing guardrails. Go to 'Guardrails' tab to create new guardrails."
>
<Select
mode="tags"
style={{ width: "100%" }}
placeholder="Select or enter guardrails"
options={guardrailsList.map((name) => ({ value: name, label: name }))}
/>
</Form.Item>
</Label>
<div className="col-span-14">
<Controller
control={control}
name="guardrails"
defaultValue={[]}
render={({ field }) => (
<TagMultiSelect
value={(field.value as string[]) ?? []}
onChange={field.onChange}
placeholder="Select or enter guardrails"
options={guardrailsList.map((name) => ({
value: name,
label: name,
}))}
/>
)}
/>
<p className="text-sm text-muted-foreground mt-1">
Select existing guardrails. Go to &apos;Guardrails&apos; tab
to create new guardrails.
</p>
</div>
</div>
<Form.Item label="Tags" name="tags" className="mb-4">
<Select
mode="tags"
style={{ width: "100%" }}
placeholder="Select or enter tags"
options={Object.values(tagsList).map((tag) => ({
value: tag.name,
label: tag.name,
title: tag.description || tag.name,
}))}
/>
</Form.Item>
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label className="col-span-10 pt-2">Tags</Label>
<div className="col-span-14">
<Controller
control={control}
name="tags"
defaultValue={[]}
render={({ field }) => (
<TagMultiSelect
value={(field.value as string[]) ?? []}
onChange={field.onChange}
placeholder="Select or enter tags"
options={Object.values(tagsList).map((tag) => ({
value: tag.name,
label: tag.name,
}))}
/>
)}
/>
</div>
</div>
{customPricing && (
<div className="ml-6 pl-4 border-l-2 border-gray-200">
<Form.Item label="Pricing Model" name="pricing_model" className="mb-4">
<Select
defaultValue="per_token"
onChange={(value: "per_token" | "per_second") => setPricingModel(value)}
options={[
{ value: "per_token", label: "Per Million Tokens" },
{ value: "per_second", label: "Per Second" },
]}
/>
</Form.Item>
<div className="ml-6 pl-4 border-l-2 border-border">
<div className="grid grid-cols-24 gap-2 mb-4 items-center">
<Label className="col-span-10">Pricing Model</Label>
<div className="col-span-14">
<Controller
control={control}
name="pricing_model"
defaultValue="per_token"
render={({ field }) => (
<Select
value={(field.value as string) || "per_token"}
onValueChange={field.onChange}
>
<SelectTrigger>
<SelectValue placeholder="Pricing model" />
</SelectTrigger>
<SelectContent>
<SelectItem value="per_token">
Per Million Tokens
</SelectItem>
<SelectItem value="per_second">
Per Second
</SelectItem>
</SelectContent>
</Select>
)}
/>
</div>
</div>
{pricingModel === "per_token" ? (
<>
<Form.Item
label="Input Cost (per 1M tokens)"
name="input_cost_per_token"
rules={[{ validator: validateNumber }]}
className="mb-4"
>
<TextInput />
</Form.Item>
<Form.Item
label="Output Cost (per 1M tokens)"
name="output_cost_per_token"
rules={[{ validator: validateNumber }]}
className="mb-4"
>
<TextInput />
</Form.Item>
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label className="col-span-10 pt-2">
Input Cost (per 1M tokens)
</Label>
<div className="col-span-14">
<Controller
control={control}
name="input_cost_per_token"
rules={{ validate: validateNumber }}
render={({ field }) => (
<Input
type="text"
value={(field.value as string) ?? ""}
onChange={(e) => field.onChange(e.target.value)}
/>
)}
/>
{errors.input_cost_per_token?.message && (
<p className="text-sm text-destructive mt-1">
{String(errors.input_cost_per_token.message)}
</p>
)}
</div>
</div>
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label className="col-span-10 pt-2">
Output Cost (per 1M tokens)
</Label>
<div className="col-span-14">
<Controller
control={control}
name="output_cost_per_token"
rules={{ validate: validateNumber }}
render={({ field }) => (
<Input
type="text"
value={(field.value as string) ?? ""}
onChange={(e) => field.onChange(e.target.value)}
/>
)}
/>
{errors.output_cost_per_token?.message && (
<p className="text-sm text-destructive mt-1">
{String(errors.output_cost_per_token.message)}
</p>
)}
</div>
</div>
</>
) : (
<Form.Item
label="Cost Per Second"
name="input_cost_per_second"
rules={[{ validator: validateNumber }]}
className="mb-4"
>
<TextInput />
</Form.Item>
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label className="col-span-10 pt-2">Cost Per Second</Label>
<div className="col-span-14">
<Controller
control={control}
name="input_cost_per_second"
rules={{ validate: validateNumber }}
render={({ field }) => (
<Input
type="text"
value={(field.value as string) ?? ""}
onChange={(e) => field.onChange(e.target.value)}
/>
)}
/>
{errors.input_cost_per_second?.message && (
<p className="text-sm text-destructive mt-1">
{String(errors.input_cost_per_second.message)}
</p>
)}
</div>
</div>
)}
</div>
)}
<Form.Item
label="Use in pass through routes"
name="use_in_pass_through"
valuePropName="checked"
className="mb-4 mt-4"
tooltip={
<span>
Allow using these credentials in pass through routes.{" "}
<Link href="https://docs.litellm.ai/docs/pass_through/vertex_ai" target="_blank">
Learn more
</Link>
</span>
}
>
<Switch onChange={handlePassThroughChange} className="bg-gray-600" />
</Form.Item>
<div className="grid grid-cols-24 gap-2 mb-4 mt-4 items-center">
<Label
className="col-span-10"
title="Allow using these credentials in pass through routes."
>
Use in pass through routes
</Label>
<div className="col-span-14">
<Controller
control={control}
name="use_in_pass_through"
render={({ field }) => (
<Switch
checked={!!field.value}
onCheckedChange={(checked) => {
field.onChange(checked);
handlePassThroughChange(checked);
}}
/>
)}
/>
</div>
</div>
<CacheControlSettings
form={form}
showCacheControl={showCacheControl}
onCacheControlChange={handleCacheControlChange}
/>
<Form.Item
label="LiteLLM Params"
name="litellm_extra_params"
tooltip="Optional litellm params used for making a litellm.completion() call."
className="mb-4 mt-4"
rules={[{ validator: formItemValidateJSON }]}
>
<TextArea
rows={4}
placeholder='{
<div className="grid grid-cols-24 gap-2 mb-4 mt-4 items-start">
<Label
className="col-span-10 pt-2"
title="Optional litellm params used for making a litellm.completion() call."
>
LiteLLM Params
</Label>
<div className="col-span-14">
<Controller
control={control}
name="litellm_extra_params"
rules={{ validate: validateJsonValue }}
render={({ field }) => (
<Textarea
rows={4}
placeholder='{
"rpm": 100,
"timeout": 0,
"stream_timeout": 0
}'
/>
</Form.Item>
<Row className="mb-4">
<Col span={10}></Col>
<Col span={10}>
<Text className="text-gray-600 text-sm">
value={(field.value as string) ?? ""}
onChange={field.onChange}
/>
)}
/>
{errors.litellm_extra_params?.message && (
<p className="text-sm text-destructive mt-1">
{String(errors.litellm_extra_params.message)}
</p>
)}
</div>
</div>
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<div className="col-span-10" />
<div className="col-span-14">
<p className="text-muted-foreground text-sm">
Pass JSON of litellm supported params{" "}
<Link href="https://docs.litellm.ai/docs/completion/input" target="_blank">
<a
href="https://docs.litellm.ai/docs/completion/input"
target="_blank"
rel="noopener noreferrer"
className="text-primary hover:text-primary/80 underline"
>
litellm.completion() call
</Link>
</Text>
</Col>
</Row>
<Form.Item
label="Model Info"
name="model_info_params"
tooltip="Optional model info params. Returned when calling `/model/info` endpoint."
className="mb-0"
rules={[{ validator: formItemValidateJSON }]}
>
<TextArea
rows={4}
placeholder='{
</a>
</p>
</div>
</div>
<div className="grid grid-cols-24 gap-2 mb-0 items-start">
<Label
className="col-span-10 pt-2"
title="Optional model info params. Returned when calling `/model/info` endpoint."
>
Model Info
</Label>
<div className="col-span-14">
<Controller
control={control}
name="model_info_params"
rules={{ validate: validateJsonValue }}
render={({ field }) => (
<Textarea
rows={4}
placeholder='{
"mode": "chat"
}'
/>
</Form.Item>
value={(field.value as string) ?? ""}
onChange={field.onChange}
/>
)}
/>
{errors.model_info_params?.message && (
<p className="text-sm text-destructive mt-1">
{String(errors.model_info_params.message)}
</p>
)}
</div>
</div>
</div>
</AccordionBody>
</Accordion>
</>
</AccordionContent>
</AccordionItem>
</Accordion>
);
};

View file

@ -1,18 +1,42 @@
import React from "react";
import { Form, Switch, Select } from "antd";
import {
useFieldArray,
useFormContext,
Controller,
FormProvider,
useForm,
} from "react-hook-form";
import { MinusCircle, Plus } from "lucide-react";
import NumericalInput from "../shared/numerical_input";
import { Label } from "@/components/ui/label";
import { Switch } from "@/components/ui/switch";
import {
Select,
SelectContent,
SelectItem,
SelectTrigger,
SelectValue,
} from "@/components/ui/select";
import { Input } from "@/components/ui/input";
interface CacheControlInjectionPoint {
location: "message";
role?: "user" | "system" | "assistant";
index?: number;
role?: "user" | "system" | "assistant" | "";
index?: number | null;
}
// eslint-disable-next-line @typescript-eslint/no-explicit-any
type AntdLikeFormInstance = any;
interface CacheControlSettingsProps {
// Form instance from parent (antd Form)
// eslint-disable-next-line @typescript-eslint/no-explicit-any
form: any;
/**
* Optional antd-style `FormInstance`. When provided (legacy consumers
* such as `model_info_view.tsx`), the component self-wraps in its own
* react-hook-form `FormProvider` so the controls keep working; the
* submit wiring stays on the antd side via `getFieldValue` /
* `setFieldValue`. When omitted (phase-1 shadcn callers), we use the
* ambient RHF context via `useFormContext`.
*/
form?: AntdLikeFormInstance;
showCacheControl: boolean;
onCacheControlChange: (checked: boolean) => void;
}
@ -22,19 +46,65 @@ const CacheControlSettings: React.FC<CacheControlSettingsProps> = ({
showCacheControl,
onCacheControlChange,
}) => {
const updateCacheControlPoints = (injectionPoints: CacheControlInjectionPoint[]) => {
const currentParams = form.getFieldValue("litellm_extra_params");
if (form && typeof form.getFieldValue === "function") {
return (
<LegacyAntdCacheControlSettings
form={form}
showCacheControl={showCacheControl}
onCacheControlChange={onCacheControlChange}
/>
);
}
return (
<RHFCacheControlSettings
showCacheControl={showCacheControl}
onCacheControlChange={onCacheControlChange}
/>
);
};
interface RHFCacheControlSettingsProps {
showCacheControl: boolean;
onCacheControlChange: (checked: boolean) => void;
}
const RHFCacheControlSettings: React.FC<RHFCacheControlSettingsProps> = ({
showCacheControl,
onCacheControlChange,
}) => {
const { control, getValues, setValue } = useFormContext();
const { fields, append, remove } = useFieldArray({
control,
name: "cache_control_injection_points",
});
const syncExtraParams = () => {
const injectionPoints = (getValues("cache_control_injection_points") ||
[]) as CacheControlInjectionPoint[];
const cleaned = injectionPoints
.filter((p) => p && (p.role || p.index !== undefined))
.map((p) => {
const next: Record<string, unknown> = {
location: p.location ?? "message",
};
if (p.role) next.role = p.role;
if (p.index !== undefined && p.index !== null) next.index = p.index;
return next;
});
const currentParams = getValues("litellm_extra_params");
try {
let paramsObj = currentParams ? JSON.parse(currentParams) : {};
if (injectionPoints.length > 0) {
paramsObj.cache_control_injection_points = injectionPoints;
const paramsObj = currentParams ? JSON.parse(currentParams) : {};
if (cleaned.length > 0) {
paramsObj.cache_control_injection_points = cleaned;
} else {
delete paramsObj.cache_control_injection_points;
}
if (Object.keys(paramsObj).length > 0) {
form.setFieldValue("litellm_extra_params", JSON.stringify(paramsObj, null, 2));
setValue("litellm_extra_params", JSON.stringify(paramsObj, null, 2));
} else {
form.setFieldValue("litellm_extra_params", "");
setValue("litellm_extra_params", "");
}
} catch (error) {
console.error("Error updating cache control points:", error);
@ -43,15 +113,22 @@ const CacheControlSettings: React.FC<CacheControlSettingsProps> = ({
return (
<>
<Form.Item
label="Cache Control Injection Points"
name="cache_control"
valuePropName="checked"
className="mb-4"
tooltip="Tell litellm where to inject cache control checkpoints. You can specify either by role (to apply to all messages of that role) or by specific message index."
>
<Switch onChange={onCacheControlChange} />
</Form.Item>
<div className="grid grid-cols-24 gap-2 mb-4 items-center">
<Label
className="col-span-10"
title="Tell litellm where to inject cache control checkpoints. You can specify either by role (to apply to all messages of that role) or by specific message index."
>
Cache Control Injection Points
</Label>
<div className="col-span-14">
<Switch
checked={showCacheControl}
onCheckedChange={(checked) => {
onCacheControlChange(!!checked);
}}
/>
</div>
</div>
{showCacheControl && (
<div className="ml-6 pl-4 border-l-2 border-border">
@ -61,102 +138,215 @@ const CacheControlSettings: React.FC<CacheControlSettingsProps> = ({
automatically add them for you as a cost saving feature.
</p>
<Form.List name="cache_control_injection_points" initialValue={[{ location: "message" }]}>
{(fields, { add, remove }) => (
<>
{fields.map((field, index) => (
<div key={field.key} className="flex items-center mb-4 gap-4">
<Form.Item
{...field}
label="Type"
name={[field.name, "location"]}
initialValue="message"
className="mb-0"
style={{ width: "180px" }}
{fields.map((field, index) => (
<div
key={field.id}
className="flex items-center mb-4 gap-4 flex-wrap"
>
<div className="flex flex-col gap-1" style={{ width: "180px" }}>
<Label>Type</Label>
<Controller
control={control}
name={`cache_control_injection_points.${index}.location`}
defaultValue="message"
render={({ field: locationField }) => (
<Select
value={(locationField.value as string) || "message"}
onValueChange={locationField.onChange}
disabled
>
<Select disabled options={[{ value: "message", label: "Message" }]} />
</Form.Item>
<SelectTrigger>
<SelectValue placeholder="Message" />
</SelectTrigger>
<SelectContent>
<SelectItem value="message">Message</SelectItem>
</SelectContent>
</Select>
)}
/>
</div>
<Form.Item
{...field}
label="Role"
name={[field.name, "role"]}
className="mb-0"
style={{ width: "180px" }}
tooltip="LiteLLM will mark all messages of this role as cacheable"
<div className="flex flex-col gap-1" style={{ width: "180px" }}>
<Label title="LiteLLM will mark all messages of this role as cacheable">
Role
</Label>
<Controller
control={control}
name={`cache_control_injection_points.${index}.role`}
render={({ field: roleField }) => (
<Select
value={(roleField.value as string) || ""}
onValueChange={(v) => {
roleField.onChange(v);
syncExtraParams();
}}
>
<Select
placeholder="Select a role"
allowClear
options={[
{ value: "user", label: "User" },
{ value: "system", label: "System" },
{ value: "assistant", label: "Assistant" },
]}
onChange={() => {
const values = form.getFieldValue("cache_control_points");
updateCacheControlPoints(values);
}}
/>
</Form.Item>
<SelectTrigger>
<SelectValue placeholder="Select a role" />
</SelectTrigger>
<SelectContent>
<SelectItem value="user">User</SelectItem>
<SelectItem value="system">System</SelectItem>
<SelectItem value="assistant">Assistant</SelectItem>
</SelectContent>
</Select>
)}
/>
</div>
<Form.Item
{...field}
label="Index"
name={[field.name, "index"]}
className="mb-0"
style={{ width: "180px" }}
tooltip="(Optional) If set litellm will mark the message at this index as cacheable"
>
<NumericalInput
type="number"
placeholder="Optional"
step={1}
onChange={() => {
const values = form.getFieldValue("cache_control_points");
updateCacheControlPoints(values);
}}
/>
</Form.Item>
<div className="flex flex-col gap-1" style={{ width: "180px" }}>
<Label title="(Optional) If set litellm will mark the message at this index as cacheable">
Index
</Label>
<Controller
control={control}
name={`cache_control_injection_points.${index}.index`}
render={({ field: indexField }) => (
<Input
type="number"
step={1}
placeholder="Optional"
value={
indexField.value === null ||
indexField.value === undefined
? ""
: (indexField.value as number)
}
onChange={(e) => {
const raw = e.target.value;
indexField.onChange(raw === "" ? null : Number(raw));
syncExtraParams();
}}
onWheel={(event) =>
(event.currentTarget as HTMLInputElement).blur()
}
/>
)}
/>
</div>
{fields.length > 1 && (
<button
type="button"
className="text-destructive cursor-pointer ml-12"
aria-label="Remove injection point"
onClick={() => {
remove(field.name);
setTimeout(() => {
const values = form.getFieldValue(
"cache_control_points",
);
updateCacheControlPoints(values);
}, 0);
}}
>
<MinusCircle className="h-5 w-5" />
</button>
)}
</div>
))}
{fields.length > 1 && (
<button
type="button"
className="text-destructive cursor-pointer ml-12"
aria-label="Remove injection point"
onClick={() => {
remove(index);
setTimeout(() => {
syncExtraParams();
}, 0);
}}
>
<MinusCircle className="h-5 w-5" />
</button>
)}
</div>
))}
<Form.Item>
<button
type="button"
className="flex items-center justify-center w-full border border-dashed border-border py-2 px-4 text-muted-foreground hover:text-primary hover:border-primary/50 transition-all rounded"
onClick={() => add()}
>
<Plus className="mr-2 h-4 w-4" />
Add Injection Point
</button>
</Form.Item>
</>
)}
</Form.List>
<button
type="button"
className="flex items-center justify-center w-full border border-dashed border-border py-2 px-4 text-muted-foreground hover:text-primary hover:border-primary/50 transition-all rounded"
onClick={() =>
append({ location: "message", role: "", index: null })
}
>
<Plus className="mr-2 h-4 w-4" />
Add Injection Point
</button>
</div>
)}
</>
);
};
/**
* Legacy wrapper used by consumers that still drive the form with an antd
* `FormInstance` (e.g. `model_info_view.tsx`). Bridges the antd form's
* `getFieldValue` / `setFieldValue` helpers to the RHF-based UI.
*/
interface LegacyAntdCacheControlSettingsProps {
form: AntdLikeFormInstance;
showCacheControl: boolean;
onCacheControlChange: (checked: boolean) => void;
}
const LegacyAntdCacheControlSettings: React.FC<
LegacyAntdCacheControlSettingsProps
> = ({ form, showCacheControl, onCacheControlChange }) => {
const rhf = useForm<{
cache_control_injection_points: CacheControlInjectionPoint[];
}>({
defaultValues: {
cache_control_injection_points:
(form?.getFieldValue?.("cache_control_injection_points") as
| CacheControlInjectionPoint[]
| undefined) ?? [{ location: "message" }],
},
});
// Mirror RHF state back into the antd form as it changes so submit wiring
// keeps working without rewriting the parent.
React.useEffect(() => {
const subscription = rhf.watch((values) => {
if (form?.setFieldValue) {
form.setFieldValue(
"cache_control_injection_points",
values.cache_control_injection_points,
);
}
// Also sync to litellm_extra_params the same way the non-legacy flow
// does, if the caller wires through that field. This matches the
// prior antd implementation's behavior.
if (
form?.getFieldValue &&
form?.setFieldValue &&
typeof form.getFieldValue === "function"
) {
const currentParams = form.getFieldValue("litellm_extra_params");
try {
const paramsObj = currentParams ? JSON.parse(currentParams) : {};
const points = (values.cache_control_injection_points ||
[]) as CacheControlInjectionPoint[];
const cleaned = points
.filter((p) => p && (p.role || p.index !== undefined))
.map((p) => {
const next: Record<string, unknown> = {
location: p.location ?? "message",
};
if (p.role) next.role = p.role;
if (p.index !== undefined && p.index !== null)
next.index = p.index;
return next;
});
if (cleaned.length > 0) {
paramsObj.cache_control_injection_points = cleaned;
} else {
delete paramsObj.cache_control_injection_points;
}
if (Object.keys(paramsObj).length > 0) {
form.setFieldValue(
"litellm_extra_params",
JSON.stringify(paramsObj, null, 2),
);
} else {
form.setFieldValue("litellm_extra_params", "");
}
} catch {
/* best-effort — caller may not track litellm_extra_params */
}
}
});
return () => subscription.unsubscribe();
}, [rhf, form]);
return (
<FormProvider {...rhf}>
<RHFCacheControlSettings
showCacheControl={showCacheControl}
onCacheControlChange={onCacheControlChange}
/>
</FormProvider>
);
};
export default CacheControlSettings;

View file

@ -1,24 +1,33 @@
import { render, screen } from "@testing-library/react";
import { Form } from "antd";
import { FormProvider, useForm } from "react-hook-form";
import { describe, expect, it } from "vitest";
import ConditionalPublicModelName from "./conditional_public_model_name";
function Wrapper({
children,
}: {
children: React.ReactNode;
}) {
const form = useForm({
defaultValues: {
model: ["gpt-4"],
model_mappings: [
{
public_name: "gpt-4",
litellm_model: "gpt-4",
},
],
},
});
return <FormProvider {...form}>{children}</FormProvider>;
}
describe("ConditionalPublicModelName", () => {
it("should render", () => {
render(
<Form
initialValues={{
model: ["gpt-4"],
model_mappings: [
{
public_name: "gpt-4",
litellm_model: "gpt-4",
},
],
}}
>
<Wrapper>
<ConditionalPublicModelName />
</Form>,
</Wrapper>,
);
expect(screen.getByText("Model Mappings")).toBeInTheDocument();

View file

@ -1,57 +1,83 @@
import React, { useEffect, useState } from "react";
import { Form, Table } from "antd";
import React, { useEffect } from "react";
import { Controller, useFormContext, useWatch } from "react-hook-form";
import { Input } from "@/components/ui/input";
import {
Table,
TableBody,
TableCell,
TableHead,
TableHeader,
TableRow,
} from "@/components/ui/table";
import { Label } from "@/components/ui/label";
import { Tooltip } from "../atoms/index";
import { Providers } from "../provider_info_helpers";
const ConditionalPublicModelName: React.FC = () => {
const form = Form.useFormInstance();
const [tableKey, setTableKey] = useState(0); // Add a key to force table re-render
interface ModelMapping {
public_name: string;
litellm_model: string;
}
// 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 ConditionalPublicModelName: React.FC = () => {
const { control, setValue, getValues, formState } = useFormContext();
const modelValue = useWatch({ control, name: "model" });
const customModelName = useWatch({ control, name: "custom_model_name" });
const selectedProvider = useWatch({ control, name: "custom_llm_provider" });
const modelMappings = (useWatch({ control, name: "model_mappings" }) ??
[]) as ModelMapping[];
const selectedModels = Array.isArray(modelValue)
? (modelValue as string[])
: modelValue
? [modelValue as string]
: [];
const showPublicModelName = !selectedModels.includes("all-wildcard");
const selectedProvider = Form.useWatch("custom_llm_provider", form);
// Force table to re-render when custom model name changes
useEffect(() => {
if (customModelName && selectedModels.includes("custom")) {
const currentMappings = form.getFieldValue("model_mappings") || [];
// eslint-disable-next-line @typescript-eslint/no-explicit-any
const updatedMappings = currentMappings.map((mapping: any) => {
if (mapping.public_name === "custom" || mapping.litellm_model === "custom") {
const currentMappings = (getValues("model_mappings") ||
[]) as ModelMapping[];
const updatedMappings = currentMappings.map((mapping) => {
if (
mapping.public_name === "custom" ||
mapping.litellm_model === "custom"
) {
if (selectedProvider === Providers.Azure) {
return {
public_name: customModelName,
litellm_model: `azure/${customModelName}`,
public_name: customModelName as string,
litellm_model: `azure/${customModelName as string}`,
};
}
return {
public_name: customModelName,
litellm_model: customModelName,
public_name: customModelName as string,
litellm_model: customModelName as string,
};
}
return mapping;
});
form.setFieldValue("model_mappings", updatedMappings);
setTableKey((prev) => prev + 1); // Force table re-render
setValue("model_mappings", updatedMappings);
}
}, [customModelName, selectedModels, selectedProvider, form]);
// eslint-disable-next-line react-hooks/exhaustive-deps
}, [customModelName, JSON.stringify(selectedModels), selectedProvider]);
// Initial setup of model mappings when models are selected
useEffect(() => {
if (selectedModels.length > 0 && !selectedModels.includes("all-wildcard")) {
// Check if we already have mappings that match the selected models
const currentMappings = form.getFieldValue("model_mappings") || [];
if (
selectedModels.length > 0 &&
!selectedModels.includes("all-wildcard")
) {
const currentMappings = (getValues("model_mappings") ||
[]) as ModelMapping[];
// Only update if the mappings don't exist or don't match the selected models
const shouldUpdateMappings =
currentMappings.length !== selectedModels.length ||
!selectedModels.every((model) =>
currentMappings.some((mapping: { public_name: string; litellm_model: string }) => {
currentMappings.some((mapping) => {
if (model === "custom") {
return mapping.litellm_model === "custom" || mapping.litellm_model === customModelName;
return (
mapping.litellm_model === "custom" ||
mapping.litellm_model === customModelName
);
}
if (selectedProvider === Providers.Azure) {
return mapping.litellm_model === `azure/${model}`;
@ -65,13 +91,13 @@ const ConditionalPublicModelName: React.FC = () => {
if (model === "custom" && customModelName) {
if (selectedProvider === Providers.Azure) {
return {
public_name: customModelName,
litellm_model: `azure/${customModelName}`,
public_name: customModelName as string,
litellm_model: `azure/${customModelName as string}`,
};
}
return {
public_name: customModelName,
litellm_model: customModelName,
public_name: customModelName as string,
litellm_model: customModelName as string,
};
}
if (selectedProvider === Providers.Azure) {
@ -86,11 +112,11 @@ const ConditionalPublicModelName: React.FC = () => {
};
});
form.setFieldValue("model_mappings", mappings);
setTableKey((prev) => prev + 1); // Force table re-render
setValue("model_mappings", mappings);
}
}
}, [selectedModels, customModelName, selectedProvider, form]);
// eslint-disable-next-line react-hooks/exhaustive-deps
}, [JSON.stringify(selectedModels), customModelName, selectedProvider]);
if (!showPublicModelName) return null;
@ -126,104 +152,128 @@ const ConditionalPublicModelName: React.FC = () => {
</>
);
const liteLLMModelTooltipContent = <div>The model name LiteLLM will send to the LLM API</div>;
const liteLLMModelTooltipContent = (
<div>The model name LiteLLM will send to the LLM API</div>
);
const columns = [
{
title: (
<span className="flex items-center">
Public Model Name
<Tooltip content={publicNameTooltipContent} width="500px" />
</span>
),
dataIndex: "public_name",
key: "public_name",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
render: (text: string, record: any, index: number) => {
return (
<Input
value={text}
onChange={(e) => {
const newValue = e.target.value;
const newMappings = [...form.getFieldValue("model_mappings")];
const handlePublicNameChange = (index: number, newValue: string) => {
const newMappings = [
...((getValues("model_mappings") || []) as ModelMapping[]),
];
// Check conditions for Anthropic -1m suffix handling
const isAnthropic = selectedProvider === Providers.Anthropic;
const endsWith1m = newValue.endsWith("-1m");
const litellmParams = form.getFieldValue("litellm_extra_params");
const isLitellmParamsEmpty = !litellmParams || litellmParams.trim() === "";
const isAnthropic = selectedProvider === Providers.Anthropic;
const endsWith1m = newValue.endsWith("-1m");
const litellmParams = getValues("litellm_extra_params") as
| string
| undefined;
const isLitellmParamsEmpty = !litellmParams || litellmParams.trim() === "";
let finalPublicName = newValue;
let finalPublicName = newValue;
if (isAnthropic && endsWith1m && isLitellmParamsEmpty) {
// Set litellm params with extra_headers
const litellmParamsValue = JSON.stringify(
{ extra_headers: { "anthropic-beta": "context-1m-2025-08-07" } },
null,
2,
);
form.setFieldValue("litellm_extra_params", litellmParamsValue);
if (isAnthropic && endsWith1m && isLitellmParamsEmpty) {
const litellmParamsValue = JSON.stringify(
{ extra_headers: { "anthropic-beta": "context-1m-2025-08-07" } },
null,
2,
);
setValue("litellm_extra_params", litellmParamsValue);
finalPublicName = newValue.slice(0, -3);
}
// Remove -1m suffix from public_name
finalPublicName = newValue.slice(0, -3); // Remove "-1m" (3 characters)
}
newMappings[index] = {
...newMappings[index],
public_name: finalPublicName,
};
setValue("model_mappings", newMappings);
};
newMappings[index].public_name = finalPublicName;
form.setFieldValue("model_mappings", newMappings);
}}
/>
);
},
},
{
title: (
<span className="flex items-center">
LiteLLM Model Name
<Tooltip content={liteLLMModelTooltipContent} width="360px" />
</span>
),
dataIndex: "litellm_model",
key: "litellm_model",
},
];
const mappingsError = (formState.errors as Record<string, { message?: string }>)
.model_mappings;
return (
<>
<Form.Item
label="Model Mappings"
name="model_mappings"
tooltip="Map public model names to LiteLLM model names for load balancing"
labelCol={{ span: 10 }}
wrapperCol={{ span: 16 }}
labelAlign="left"
rules={[
{
required: true,
validator: async (_, value) => {
<div className="grid grid-cols-24 gap-2 mb-4 items-start">
<Label
className="col-span-10 pt-2"
title="Map public model names to LiteLLM model names for load balancing"
>
Model Mappings
</Label>
<div className="col-span-14 space-y-2">
<Controller
control={control}
name="model_mappings"
rules={{
validate: (value) => {
if (!value || value.length === 0) {
throw new Error("At least one model mapping is required");
return "At least one model mapping is required";
}
const invalidMappings = value.filter(
// eslint-disable-next-line @typescript-eslint/no-explicit-any
(mapping: any) =>
const invalidMappings = (value as ModelMapping[]).filter(
(mapping) =>
!mapping.public_name || mapping.public_name.trim() === "",
);
if (invalidMappings.length > 0) {
throw new Error("All model mappings must have valid public names");
return "All model mappings must have valid public names";
}
return true;
},
},
]}
>
<Table
key={tableKey} // Add key to force re-render
dataSource={form.getFieldValue("model_mappings")}
columns={columns}
pagination={false}
size="small"
}}
render={() => (
<Table>
<TableHeader>
<TableRow>
<TableHead>
<span className="flex items-center">
Public Model Name
<Tooltip content={publicNameTooltipContent} width="500px" />
</span>
</TableHead>
<TableHead>
<span className="flex items-center">
LiteLLM Model Name
<Tooltip
content={liteLLMModelTooltipContent}
width="360px"
/>
</span>
</TableHead>
</TableRow>
</TableHeader>
<TableBody>
{modelMappings.length === 0 ? (
<TableRow>
<TableCell
colSpan={2}
className="text-center text-muted-foreground py-4"
>
Select at least one model
</TableCell>
</TableRow>
) : (
modelMappings.map((record, index) => (
<TableRow key={`${record.litellm_model}-${index}`}>
<TableCell>
<Input
value={record.public_name}
onChange={(e) =>
handlePublicNameChange(index, e.target.value)
}
/>
</TableCell>
<TableCell>{record.litellm_model}</TableCell>
</TableRow>
))
)}
</TableBody>
</Table>
)}
/>
</Form.Item>
</>
{mappingsError?.message && (
<p className="text-sm text-destructive">
{String(mappingsError.message)}
</p>
)}
</div>
</div>
);
};

View file

@ -73,8 +73,14 @@ export const handleAddAutoRouterSubmit = async (values: any, accessToken: string
const routerTypeName = values.model_type === "complexity_router" ? "Complexity Router" : "Semantic Router";
NotificationManager.success(`Successfully created ${routerTypeName}: ${values.auto_router_name}`);
// Reset the form
form.resetFields();
// Reset the form. The `form` argument may be an antd FormInstance
// (legacy) or a RHF `UseFormReturn` (phase-1 shadcn migration); we
// support both by feature-detecting the reset method.
if (typeof form?.resetFields === "function") {
form.resetFields();
} else if (typeof form?.reset === "function") {
form.reset();
}
// Call the callback if provided (e.g., to close modal)
if (callback) {

View file

@ -165,7 +165,13 @@ export const handleAddModelSubmit = async (values: any, accessToken: string, for
}
callback && callback();
form.resetFields();
// `form` may be an antd FormInstance (legacy) or a RHF `UseFormReturn`
// (phase-1 shadcn migration); we support both.
if (typeof form?.resetFields === "function") {
form.resetFields();
} else if (typeof form?.reset === "function") {
form.reset();
}
} catch (error) {
NotificationManager.fromBackend("Failed to add model: " + error);
}

View file

@ -1,28 +1,37 @@
import { render } from "@testing-library/react";
import { Form } from "antd";
import { FormProvider, useForm } from "react-hook-form";
import { describe, expect, it } from "vitest";
import { getPlaceholder, Providers } from "../provider_info_helpers";
import LiteLLMModelNameField from "./litellm_model_name";
function Wrapper({ children }: { children: React.ReactNode }) {
const form = useForm({ defaultValues: {} });
return <FormProvider {...form}>{children}</FormProvider>;
}
describe("LitellmModelNameField", () => {
it("should render", () => {
const { getByText } = render(
<Form>
<Wrapper>
<LiteLLMModelNameField
selectedProvider={Providers.OpenAI}
providerModels={[]}
getPlaceholder={getPlaceholder}
/>
</Form>,
</Wrapper>,
);
expect(getByText("LiteLLM Model Name(s)")).toBeInTheDocument();
});
it("should show Azure placeholder as 'my-deployment'", () => {
const { getByPlaceholderText, queryByPlaceholderText } = render(
<Form>
<LiteLLMModelNameField selectedProvider={Providers.Azure} providerModels={[]} getPlaceholder={getPlaceholder} />
</Form>,
<Wrapper>
<LiteLLMModelNameField
selectedProvider={Providers.Azure}
providerModels={[]}
getPlaceholder={getPlaceholder}
/>
</Wrapper>,
);
expect(getByPlaceholderText("my-deployment")).toBeInTheDocument();
expect(queryByPlaceholderText("gpt-3.5-turbo")).toBeNull();

View file

@ -1,6 +1,16 @@
import React from "react";
import { Form, Select as AntSelect } from "antd";
import { Controller, useFormContext, useWatch } from "react-hook-form";
import { Input } from "@/components/ui/input";
import { Label } from "@/components/ui/label";
import {
Select,
SelectContent,
SelectItem,
SelectTrigger,
SelectValue,
} from "@/components/ui/select";
import { Badge } from "@/components/ui/badge";
import { X } from "lucide-react";
import { Providers } from "../provider_info_helpers";
interface LiteLLMModelNameFieldProps {
@ -9,53 +19,128 @@ interface LiteLLMModelNameFieldProps {
getPlaceholder: (provider: Providers) => string;
}
/**
* Multi-select rendered with shadcn Select + chip list below. Accepts any
* `Array<{ label: string; value: string }>` list of options. Mirrors the
* pattern established in `AccessGroupBaseForm.tsx`.
*/
function MultiSelect({
value,
onChange,
options,
placeholder,
testId,
}: {
value: string[];
onChange: (next: string[]) => void;
options: { label: string; value: string }[];
placeholder: string;
testId?: string;
}) {
const selected = value ?? [];
const remaining = options.filter((o) => !selected.includes(o.value));
return (
<div className="space-y-2" data-testid={testId}>
<Select
value=""
onValueChange={(v) => {
if (v) onChange([...selected, v]);
}}
>
<SelectTrigger>
<SelectValue placeholder={placeholder} />
</SelectTrigger>
<SelectContent>
{remaining.length === 0 ? (
<div className="py-2 px-3 text-sm text-muted-foreground">
No options available
</div>
) : (
remaining.map((opt) => (
<SelectItem key={opt.value} value={opt.value}>
{opt.label}
</SelectItem>
))
)}
</SelectContent>
</Select>
{selected.length > 0 && (
<div className="flex flex-wrap gap-1">
{selected.map((v) => {
const opt = options.find((o) => o.value === v);
return (
<Badge
key={v}
variant="secondary"
className="flex items-center gap-1"
>
{opt?.label ?? v}
<button
type="button"
onClick={() =>
onChange(selected.filter((s) => s !== v))
}
className="inline-flex items-center justify-center rounded-full hover:bg-muted-foreground/20"
aria-label={`Remove ${opt?.label ?? v}`}
>
<X size={12} />
</button>
</Badge>
);
})}
</div>
)}
</div>
);
}
const LiteLLMModelNameField: React.FC<LiteLLMModelNameFieldProps> = ({
selectedProvider,
providerModels,
getPlaceholder,
}) => {
const form = Form.useFormInstance();
const { control, setValue, getValues, formState } = useFormContext();
const selectedModels = (useWatch({ control, name: "model" }) || []) as
| string[]
| string;
const modelArray = Array.isArray(selectedModels)
? selectedModels
: selectedModels
? [selectedModels]
: [];
const handleModelChange = (value: string | string[]) => {
// Ensure value is always treated as an array
const values = Array.isArray(value) ? value : [value];
// If "all-wildcard" is selected, clear the model_name field
const handleModelChange = (values: string[]) => {
if (values.includes("all-wildcard")) {
form.setFieldsValue({ model_name: undefined, model_mappings: [] });
} else {
// Get current model value to check if we need to update
const currentModel = form.getFieldValue("model");
setValue("model_name", undefined);
setValue("model_mappings", []);
setValue("model", values);
return;
}
// Only update if the value has actually changed
if (JSON.stringify(currentModel) !== JSON.stringify(values)) {
// Create mappings first
const mappings = values.map((model) => {
if (selectedProvider === Providers.Azure) {
return {
public_name: model,
litellm_model: `azure/${model}`,
};
}
const currentModel = getValues("model");
if (JSON.stringify(currentModel) !== JSON.stringify(values)) {
const mappings = values.map((model) => {
if (selectedProvider === Providers.Azure) {
return {
public_name: model,
litellm_model: model,
litellm_model: `azure/${model}`,
};
});
// Update both fields in one call to reduce re-renders
form.setFieldsValue({
model: values,
model_mappings: mappings,
});
}
}
return {
public_name: model,
litellm_model: model,
};
});
setValue("model", values);
setValue("model_mappings", mappings);
}
};
const handleAzureDeploymentNameChange = (e: React.ChangeEvent<HTMLInputElement>) => {
const handleAzureDeploymentNameChange = (
e: React.ChangeEvent<HTMLInputElement>,
) => {
const deploymentName = e.target.value;
// Create mapping with Azure-specific format
const mappings = deploymentName
? [
{
@ -64,23 +149,23 @@ const LiteLLMModelNameField: React.FC<LiteLLMModelNameFieldProps> = ({
},
]
: [];
// Update both fields
form.setFieldsValue({
model: deploymentName,
model_mappings: mappings,
});
setValue("model", deploymentName);
setValue("model_mappings", mappings);
};
// Handle custom model name changes
const handleCustomModelNameChange = (e: React.ChangeEvent<HTMLInputElement>) => {
const handleCustomModelNameChange = (
e: React.ChangeEvent<HTMLInputElement>,
) => {
const customName = e.target.value;
setValue("custom_model_name", customName);
// Immediately update the model mappings
const currentMappings = form.getFieldValue("model_mappings") || [];
const currentMappings = getValues("model_mappings") || [];
// eslint-disable-next-line @typescript-eslint/no-explicit-any
const updatedMappings = currentMappings.map((mapping: any) => {
if (mapping.public_name === "custom" || mapping.litellm_model === "custom") {
if (
mapping.public_name === "custom" ||
mapping.litellm_model === "custom"
) {
if (selectedProvider === Providers.Azure) {
return {
public_name: customName,
@ -94,100 +179,149 @@ const LiteLLMModelNameField: React.FC<LiteLLMModelNameFieldProps> = ({
}
return mapping;
});
form.setFieldsValue({ model_mappings: updatedMappings });
setValue("model_mappings", updatedMappings);
};
const showSelectVariant =
selectedProvider !== Providers.Azure &&
selectedProvider !== Providers.OpenAI_Compatible &&
selectedProvider !== Providers.Ollama &&
providerModels.length > 0;
const modelError = (formState.errors as Record<string, { message?: string }>)
.model;
const customNameError = (
formState.errors as Record<string, { message?: string }>
).custom_model_name;
const requiredMessage = `Please enter ${
selectedProvider === Providers.Azure
? "a deployment name"
: "at least one model"
}.`;
return (
<>
<Form.Item
label="LiteLLM Model Name(s)"
tooltip="The model name LiteLLM will send to the LLM API"
className="mb-0"
>
<Form.Item
name="model"
rules={[
{
required: true,
message: `Please enter ${selectedProvider === Providers.Azure ? "a deployment name" : "at least one model"}.`,
},
]}
noStyle
>
<div className="grid grid-cols-24 gap-2 mb-0 items-start">
<Label className="col-span-10 pt-2" title="The model name LiteLLM will send to the LLM API">
LiteLLM Model Name(s)
</Label>
<div className="col-span-14 space-y-2">
{selectedProvider === Providers.Azure ||
selectedProvider === Providers.OpenAI_Compatible ||
selectedProvider === Providers.Ollama ? (
<Input
placeholder={getPlaceholder(selectedProvider)}
onChange={
selectedProvider === Providers.Azure
? handleAzureDeploymentNameChange
: undefined
}
<Controller
control={control}
name="model"
rules={{
validate: (value) =>
(typeof value === "string" && value.length > 0) ||
(Array.isArray(value) && value.length > 0) ||
requiredMessage,
}}
render={({ field }) => (
<Input
placeholder={getPlaceholder(selectedProvider)}
value={(field.value as string) ?? ""}
onChange={(e) => {
field.onChange(e.target.value);
if (selectedProvider === Providers.Azure) {
handleAzureDeploymentNameChange(e);
}
}}
/>
)}
/>
) : providerModels.length > 0 ? (
<AntSelect
data-testid="model-name-select"
mode="multiple"
allowClear
showSearch
placeholder="Select models"
onChange={handleModelChange}
optionFilterProp="children"
filterOption={(input, option) => (option?.label ?? "").toLowerCase().includes(input.toLowerCase())}
options={[
{
label: "Custom Model Name (Enter below)",
value: "custom",
},
{
label: `All ${selectedProvider} Models (Wildcard)`,
value: "all-wildcard",
},
...providerModels.map((model) => ({
label: model,
value: model,
})),
]}
style={{ width: "100%" }}
) : showSelectVariant ? (
<Controller
control={control}
name="model"
rules={{
validate: (value) =>
(Array.isArray(value) && value.length > 0) ||
requiredMessage,
}}
render={({ field }) => (
<MultiSelect
testId="model-name-select"
value={(field.value as string[]) ?? []}
onChange={(next) => {
field.onChange(next);
handleModelChange(next);
}}
placeholder="Select models"
options={[
{
label: "Custom Model Name (Enter below)",
value: "custom",
},
{
label: `All ${selectedProvider} Models (Wildcard)`,
value: "all-wildcard",
},
...providerModels.map((model) => ({
label: model,
value: model,
})),
]}
/>
)}
/>
) : (
<Input placeholder={getPlaceholder(selectedProvider)} />
<Controller
control={control}
name="model"
rules={{
validate: (value) =>
(typeof value === "string" && value.length > 0) ||
(Array.isArray(value) && value.length > 0) ||
requiredMessage,
}}
render={({ field }) => (
<Input
placeholder={getPlaceholder(selectedProvider)}
value={(field.value as string) ?? ""}
onChange={(e) => field.onChange(e.target.value)}
/>
)}
/>
)}
{modelError?.message && (
<p className="text-sm text-destructive">
{String(modelError.message)}
</p>
)}
</Form.Item>
{/* Custom Model Name field */}
<Form.Item noStyle shouldUpdate={(prevValues, currentValues) => prevValues.model !== currentValues.model}>
{({ getFieldValue }) => {
const selectedModels = getFieldValue("model") || [];
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"
>
{modelArray.includes("custom") && (
<div className="mt-2 space-y-1">
<Controller
control={control}
name="custom_model_name"
rules={{ required: "Please enter a custom model name." }}
render={({ field }) => (
<Input
placeholder={
selectedProvider === Providers.Azure
? "Enter Azure deployment name"
: "Enter custom model name"
}
onChange={handleCustomModelNameChange}
value={(field.value as string) ?? ""}
onChange={(e) => {
field.onChange(e.target.value);
handleCustomModelNameChange(e);
}}
/>
</Form.Item>
)
);
}}
</Form.Item>
</Form.Item>
)}
/>
{customNameError?.message && (
<p className="text-sm text-destructive">
{String(customNameError.message)}
</p>
)}
</div>
)}
</div>
</div>
<div className="grid grid-cols-24 gap-2">
<div className="col-span-10" />
<p className="col-span-14 mb-3 mt-1 text-sm text-muted-foreground">

View file

@ -1,6 +1,6 @@
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { render, screen, waitFor } from "@testing-library/react";
import { Form } from "antd";
import { FormProvider, useForm } from "react-hook-form";
import { beforeAll, describe, expect, it, vi } from "vitest";
import { Providers } from "../provider_info_helpers";
import ProviderSpecificFields from "./provider_specific_fields";
@ -124,19 +124,24 @@ const createQueryClient = () =>
},
});
function Wrapper({ children }: { children: React.ReactNode }) {
const form = useForm({ defaultValues: {} });
return <FormProvider {...form}>{children}</FormProvider>;
}
describe("ProviderSpecificFields", () => {
it("should render", async () => {
const queryClient = createQueryClient();
render(
<QueryClientProvider client={queryClient}>
<Form>
<Wrapper>
<ProviderSpecificFields selectedProvider={Providers.OpenAI} />
</Form>
</Wrapper>
</QueryClientProvider>,
);
await waitFor(() => {
expect(screen.getByLabelText("OpenAI API Key")).toBeInTheDocument();
expect(screen.getByLabelText(/OpenAI API Key/i)).toBeInTheDocument();
});
});
@ -144,14 +149,14 @@ describe("ProviderSpecificFields", () => {
const queryClient = createQueryClient();
render(
<QueryClientProvider client={queryClient}>
<Form>
<Wrapper>
<ProviderSpecificFields selectedProvider={Providers.OpenAI} />
</Form>
</Wrapper>
</QueryClientProvider>,
);
await waitFor(() => {
const apiKeyLabel = screen.getByLabelText("OpenAI API Key");
const apiKeyLabel = screen.getByLabelText(/OpenAI API Key/i);
expect(apiKeyLabel).toBeInTheDocument();
const apiBaseInput = screen.getByPlaceholderText("https://api.openai.com/v1");
@ -167,14 +172,14 @@ describe("ProviderSpecificFields", () => {
const queryClient = createQueryClient();
render(
<QueryClientProvider client={queryClient}>
<Form>
<Wrapper>
<ProviderSpecificFields selectedProvider={"Hosted_Vllm" as Providers} />
</Form>
</Wrapper>
</QueryClientProvider>,
);
await waitFor(() => {
const apiKeyLabel = screen.getByLabelText("vLLM API Key");
const apiKeyLabel = screen.getByLabelText(/vLLM API Key/i);
expect(apiKeyLabel).toBeInTheDocument();
const apiBaseInput = screen.getByPlaceholderText("https://...");
@ -187,19 +192,19 @@ describe("ProviderSpecificFields", () => {
const queryClient = createQueryClient();
render(
<QueryClientProvider client={queryClient}>
<Form>
<Wrapper>
<ProviderSpecificFields selectedProvider={Providers.Azure} />
</Form>
</Wrapper>
</QueryClientProvider>,
);
await waitFor(() => {
const apiKeyInput = screen.getByLabelText("Azure API Key");
const apiKeyInput = screen.getByLabelText(/Azure API Key/i);
expect(apiKeyInput).toBeInTheDocument();
expect(apiKeyInput).toHaveAttribute("type", "password");
expect(apiKeyInput).toHaveAttribute("placeholder", "Enter your Azure API Key");
const azureAdTokenInput = screen.getByLabelText("Azure AD Token");
const azureAdTokenInput = screen.getByLabelText(/Azure AD Token/i);
expect(azureAdTokenInput).toBeInTheDocument();
expect(azureAdTokenInput).toHaveAttribute("type", "password");
expect(azureAdTokenInput).toHaveAttribute("placeholder", "Enter your Azure AD Token");

View file

@ -2,17 +2,22 @@ import { useProviderFields } from "@/app/(dashboard)/hooks/providers/useProvider
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { Textarea } from "@/components/ui/textarea";
import { Upload as UploadIcon } from "lucide-react";
import { Label } from "@/components/ui/label";
import {
Col,
Form,
Row,
Select,
Upload,
UploadProps,
} from "antd";
SelectContent,
SelectItem,
SelectTrigger,
SelectValue,
} from "@/components/ui/select";
import { Upload as UploadIcon } from "lucide-react";
import React from "react";
import { CredentialItem, ProviderCredentialFieldMetadata } from "../networking";
import { Controller, useFormContext } from "react-hook-form";
import type { UploadProps } from "./add_model_upload_types";
import {
CredentialItem,
ProviderCredentialFieldMetadata,
} from "../networking";
import { provider_map, Providers } from "../provider_info_helpers";
interface ProviderSpecificFieldsProps {
@ -36,7 +41,9 @@ export interface CredentialValues {
value: string;
}
const mapFieldMetadataToUiField = (field: ProviderCredentialFieldMetadata): ProviderCredentialField => {
const mapFieldMetadataToUiField = (
field: ProviderCredentialFieldMetadata,
): ProviderCredentialField => {
const type: ProviderCredentialField["type"] =
field.field_type === "password"
? "password"
@ -63,7 +70,8 @@ const mapFieldMetadataToUiField = (field: ProviderCredentialFieldMetadata): Prov
// In-memory cache of provider credential fields keyed by provider display name.
// This lets us reuse the data across multiple mounts and also supports
// non-React helpers like createCredentialFromModel.
const providerFieldsByDisplayName: Record<string, ProviderCredentialField[]> = {};
const providerFieldsByDisplayName: Record<string, ProviderCredentialField[]> =
{};
export const createCredentialFromModel = (
provider: string,
@ -99,11 +107,20 @@ export const createCredentialFromModel = (
return credential;
};
const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({ selectedProvider, uploadProps }) => {
const selectedProviderEnum = Providers[selectedProvider as keyof typeof Providers] as Providers;
const form = Form.useFormInstance(); // Get form instance from context
const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({
selectedProvider,
uploadProps,
}) => {
const selectedProviderEnum = Providers[
selectedProvider as keyof typeof Providers
] as Providers;
const { control, setValue, formState } = useFormContext();
const { data: providerMetadata, isLoading, error: loadError } = useProviderFields();
const {
data: providerMetadata,
isLoading,
error: loadError,
} = useProviderFields();
// Memoize the expensive cache computation
const cacheEntries = React.useMemo(() => {
@ -111,16 +128,15 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({ selecte
return null;
}
// Compute cache entries keyed by provider display name and identifiers
const entries: Record<string, ProviderCredentialField[]> = {};
providerMetadata.forEach((providerInfo) => {
const displayName = providerInfo.provider_display_name;
const mappedFields = providerInfo.credential_fields.map(mapFieldMetadataToUiField);
const mappedFields = providerInfo.credential_fields.map(
mapFieldMetadataToUiField,
);
// Primary key: human-readable display name
entries[displayName] = mappedFields;
// Also cache by backend identifiers so lookups by provider slug work
if (providerInfo.provider) {
entries[providerInfo.provider] = mappedFields;
}
@ -131,20 +147,17 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({ selecte
return entries;
}, [providerMetadata]);
// Sync memoized cache entries to module-level cache
React.useEffect(() => {
if (!cacheEntries) {
return;
}
Object.assign(providerFieldsByDisplayName, cacheEntries);
}, [cacheEntries]);
const allFields = React.useMemo(() => {
// First try to resolve from the in-memory cache. We support both the
// enum/display-name form and the raw provider slug (e.g. "petals").
const cachedFields =
providerFieldsByDisplayName[selectedProviderEnum] ?? providerFieldsByDisplayName[selectedProvider];
providerFieldsByDisplayName[selectedProviderEnum] ??
providerFieldsByDisplayName[selectedProvider];
if (cachedFields) {
return cachedFields;
}
@ -174,125 +187,185 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({ selecte
return mapped;
}, [selectedProviderEnum, selectedProvider, providerMetadata]);
const handleUpload = {
name: "file",
accept: ".json",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
beforeUpload: (file: any) => {
if (file.type === "application/json") {
const reader = new FileReader();
reader.onload = (e) => {
if (e.target) {
const jsonStr = e.target.result as string;
form.setFieldsValue({ vertex_credentials: jsonStr });
}
};
reader.readAsText(file);
}
return false;
},
const fileInputRef = React.useRef<HTMLInputElement | null>(null);
const handleFilePick = async (e: React.ChangeEvent<HTMLInputElement>) => {
const file = e.target.files?.[0];
if (!file) return;
if (file.type === "application/json") {
const text = await file.text();
setValue("vertex_credentials", text, { shouldValidate: true });
}
// Delegate to consumer-provided upload props if present (kept for parity
// with the antd Upload onChange callback).
if (uploadProps?.onChange) {
uploadProps.onChange({
file: {
name: file.name,
status: "done",
type: file.type,
} as unknown as Parameters<NonNullable<UploadProps["onChange"]>>[0]["file"],
});
}
// Reset the native input so the same file can be selected again.
e.target.value = "";
};
return (
<>
{isLoading && allFields.length === 0 && (
<Row>
<Col span={24}>
<p className="mb-2">Loading provider fields...</p>
</Col>
</Row>
<div className="mb-2">Loading provider fields...</div>
)}
{loadError && allFields.length === 0 && (
<Row>
<Col span={24}>
<p className="mb-2 text-destructive">
{loadError instanceof Error
? loadError.message
: "Failed to load provider credential fields"}
</p>
</Col>
</Row>
<div className="mb-2 text-destructive">
{loadError instanceof Error
? loadError.message
: "Failed to load provider credential fields"}
</div>
)}
{allFields.map((field) => (
<React.Fragment key={field.key}>
<Form.Item
label={field.label}
name={field.key}
rules={field.required ? [{ required: true, message: "Required" }] : undefined}
tooltip={field.tooltip}
className={field.key === "vertex_credentials" ? "mb-0" : undefined}
>
{field.type === "select" ? (
<Select placeholder={field.placeholder} defaultValue={field.defaultValue}>
{field.options?.map((option) => (
<Select.Option key={option} value={option}>
{option}
</Select.Option>
))}
</Select>
) : field.type === "upload" ? (
<Upload
{...handleUpload}
onChange={(info) => {
if (uploadProps?.onChange) {
uploadProps.onChange(info);
}
}}
{allFields.map((field) => {
const fieldError =
(formState.errors as Record<string, { message?: string }>)[field.key];
const requiredRule = field.required
? { required: "Required" as const }
: {};
return (
<React.Fragment key={field.key}>
<div
className={`grid grid-cols-24 gap-2 mb-4 ${
field.key === "vertex_credentials" ? "mb-0" : ""
}`}
>
<Label
className="col-span-10 pt-2"
title={field.tooltip}
htmlFor={`provider-field-${field.key}`}
>
<Button type="button" variant="outline">
<UploadIcon className="h-4 w-4" />
Click to Upload
</Button>
</Upload>
) : field.type === "textarea" ? (
<Textarea
placeholder={field.placeholder}
defaultValue={field.defaultValue}
rows={6}
className="font-mono text-xs"
/>
) : (
<Input
placeholder={field.placeholder}
type={field.type === "password" ? "password" : "text"}
defaultValue={field.defaultValue}
/>
{field.label}
{field.required && (
<span className="text-destructive ml-1">*</span>
)}
</Label>
<div className="col-span-14">
{field.type === "select" ? (
<Controller
control={control}
name={field.key}
defaultValue={field.defaultValue ?? ""}
rules={requiredRule}
render={({ field: controllerField }) => (
<Select
value={(controllerField.value as string) || ""}
onValueChange={controllerField.onChange}
>
<SelectTrigger id={`provider-field-${field.key}`}>
<SelectValue placeholder={field.placeholder} />
</SelectTrigger>
<SelectContent>
{field.options?.map((option) => (
<SelectItem key={option} value={option}>
{option}
</SelectItem>
))}
</SelectContent>
</Select>
)}
/>
) : field.type === "upload" ? (
<>
<input
ref={fileInputRef}
type="file"
accept=".json"
hidden
onChange={handleFilePick}
/>
<Button
type="button"
variant="outline"
onClick={() => fileInputRef.current?.click()}
>
<UploadIcon className="h-4 w-4" />
Click to Upload
</Button>
</>
) : field.type === "textarea" ? (
<Controller
control={control}
name={field.key}
defaultValue={field.defaultValue ?? ""}
rules={requiredRule}
render={({ field: controllerField }) => (
<Textarea
id={`provider-field-${field.key}`}
placeholder={field.placeholder}
rows={6}
className="font-mono text-xs"
{...controllerField}
value={(controllerField.value as string) ?? ""}
/>
)}
/>
) : (
<Controller
control={control}
name={field.key}
defaultValue={field.defaultValue ?? ""}
rules={requiredRule}
render={({ field: controllerField }) => (
<Input
id={`provider-field-${field.key}`}
placeholder={field.placeholder}
type={field.type === "password" ? "password" : "text"}
{...controllerField}
value={(controllerField.value as string) ?? ""}
/>
)}
/>
)}
{fieldError?.message && (
<p className="text-sm text-destructive mt-1">
{String(fieldError.message)}
</p>
)}
</div>
</div>
{/* Special case for Vertex Credentials help text */}
{field.key === "vertex_credentials" && (
<div className="grid grid-cols-24 gap-2">
<div className="col-span-24">
<p className="mb-3 mt-1">
Give a gcp service account(.json file)
</p>
</div>
</div>
)}
</Form.Item>
{/* Special case for Vertex Credentials help text */}
{field.key === "vertex_credentials" && (
<Row>
<Col>
<p className="mb-3 mt-1">
Give a gcp service account(.json file)
</p>
</Col>
</Row>
)}
{/* Special case for Azure Base Model help text */}
{field.key === "base_model" && (
<Row>
<Col span={10}></Col>
<Col span={10}>
<p className="mb-2">
The actual model your azure deployment uses. Used for
accurate cost tracking. Select name from{" "}
<a
href="https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json"
target="_blank"
rel="noopener noreferrer"
className="text-primary hover:text-primary/80 underline"
>
here
</a>
</p>
</Col>
</Row>
)}
</React.Fragment>
))}
{/* Special case for Azure Base Model help text */}
{field.key === "base_model" && (
<div className="grid grid-cols-24 gap-2">
<div className="col-span-10" />
<div className="col-span-14">
<p className="mb-2">
The actual model your azure deployment uses. Used for
accurate cost tracking. Select name from{" "}
<a
href="https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json"
target="_blank"
rel="noopener noreferrer"
className="text-primary hover:text-primary/80 underline"
>
here
</a>
</p>
</div>
</div>
)}
</React.Fragment>
);
})}
</>
);
};

View file

@ -22,3 +22,19 @@ export const formItemValidateJSON = (_: any, value: string) => {
return Promise.reject("Please enter valid JSON");
}
};
/**
* React Hook Form-compatible JSON validator. Returns `true` when `value`
* is empty or parses as valid JSON, otherwise the error message string.
*/
export const validateJsonValue = (value: unknown): true | string => {
if (value === undefined || value === null || value === "") {
return true;
}
try {
JSON.parse(value as string);
return true;
} catch {
return "Please enter valid JSON";
}
};