mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
ce1c604be2
commit
a2f8c87e9f
19 changed files with 2660 additions and 1421 deletions
|
|
@ -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"}`,
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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">✓ <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;
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
|
@ -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"));
|
||||
|
|
|
|||
|
|
@ -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 > 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 'Guardrails' 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>
|
||||
);
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
})}
|
||||
</>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
}
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue