diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.test.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.test.tsx new file mode 100644 index 00000000000..225c308af91 --- /dev/null +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.test.tsx @@ -0,0 +1,87 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import EditFallbacks, { Fallbacks } from "./EditFallbacks"; +import * as fetchModelsModule from "@/components/llm_calls/fetch_models"; + +vi.mock("@/components/llm_calls/fetch_models", () => ({ + fetchAvailableModels: vi.fn(), +})); + +const renderWithQueryClient = (ui: React.ReactElement) => { + const queryClient = new QueryClient({ + defaultOptions: { queries: { retry: false } }, + }); + return render({ui}); +}; + +describe("EditFallbacks", () => { + const accessToken = "test-token"; + const fallbackEntry = { "gpt-4": ["gpt-3.5-turbo", "claude-3-opus"] }; + const value: Fallbacks = [{ "gpt-4": ["gpt-3.5-turbo", "claude-3-opus"] }, { "claude-3-opus": ["gpt-4"] }]; + + const setup = (overrides: Partial> = {}) => { + const onChange = overrides.onChange ?? vi.fn().mockResolvedValue(undefined); + const onClose = overrides.onClose ?? vi.fn(); + renderWithQueryClient( + , + ); + return { onChange, onClose }; + }; + + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(fetchModelsModule.fetchAvailableModels).mockResolvedValue([ + { model_group: "gpt-4", mode: "chat" }, + { model_group: "gpt-3.5-turbo", mode: "chat" }, + { model_group: "claude-3-opus", mode: "chat" }, + { model_group: "gemini-pro", mode: "chat" }, + ]); + }); + + it("prefills the existing fallback chain for the primary model", async () => { + setup(); + await waitFor(() => { + expect(screen.getByText("gpt-3.5-turbo")).toBeInTheDocument(); + expect(screen.getByText("claude-3-opus")).toBeInTheDocument(); + }); + }); + + it("removes a fallback model and saves only the edited entry", async () => { + const user = userEvent.setup(); + const onChange = vi.fn().mockResolvedValue(undefined); + const onClose = vi.fn(); + setup({ onChange, onClose }); + + await screen.findByText("gpt-3.5-turbo"); + await user.click(screen.getByTestId("remove-fallback-gpt-3.5-turbo")); + + await user.click(screen.getByRole("button", { name: /save changes/i })); + + await waitFor(() => { + expect(onChange).toHaveBeenCalledWith([{ "gpt-4": ["claude-3-opus"] }, { "claude-3-opus": ["gpt-4"] }]); + }); + await waitFor(() => expect(onClose).toHaveBeenCalled()); + }); + + it("blocks saving with an empty fallback chain", async () => { + const user = userEvent.setup(); + const onChange = vi.fn().mockResolvedValue(undefined); + setup({ fallbackEntry: { "gpt-4": ["gpt-3.5-turbo"] }, onChange }); + + await screen.findByText("gpt-3.5-turbo"); + await user.click(screen.getByTestId("remove-fallback-gpt-3.5-turbo")); + + const saveButton = screen.getByRole("button", { name: /save changes/i }); + expect(saveButton).toBeDisabled(); + expect(onChange).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.tsx new file mode 100644 index 00000000000..938e1104301 --- /dev/null +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.tsx @@ -0,0 +1,106 @@ +/** + * Modal for editing an existing fallback entry + * Lets the user add/remove models from a primary model's fallback chain + * Reuses FallbackGroupConfig with the primary model locked + */ + +import { Button } from "antd"; +import { useQuery } from "@tanstack/react-query"; +import { Pencil } from "lucide-react"; +import React, { useMemo, useState } from "react"; +import { fetchAvailableModels } from "@/components/llm_calls/fetch_models"; +import NotificationManager from "../../../molecules/notifications_manager"; +import { AddFallbacksModal } from "./AddFallbacksModal"; +import { FallbackGroup, FallbackGroupConfig } from "./FallbackGroupConfig"; + +export type FallbackEntry = { [modelName: string]: string[] }; +export type Fallbacks = FallbackEntry[]; + +interface EditFallbacksProps { + accessToken: string; + fallbackEntry: FallbackEntry; + value: Fallbacks; + onChange: (fallbacks: Fallbacks) => Promise; + onClose: () => void; + maxFallbacks?: number; +} + +const toGroup = (entry: FallbackEntry): FallbackGroup => { + const primaryModel = Object.keys(entry)[0] ?? null; + return { + id: "edit", + primaryModel, + fallbackModels: primaryModel ? [...(entry[primaryModel] ?? [])] : [], + }; +}; + +export default function EditFallbacks({ + accessToken, + fallbackEntry, + value, + onChange, + onClose, + maxFallbacks = 10, +}: EditFallbacksProps) { + const [group, setGroup] = useState(() => toGroup(fallbackEntry)); + const [isSaving, setIsSaving] = useState(false); + + const { data: modelGroups = [] } = useQuery({ + queryKey: ["availableModels", "fallbacks"], + queryFn: () => fetchAvailableModels(accessToken), + enabled: Boolean(accessToken), + }); + + const availableModels = useMemo( + () => Array.from(new Set(modelGroups.map((option) => option.model_group))).sort(), + [modelGroups], + ); + + const handleSave = async () => { + const primaryModel = group.primaryModel; + if (!primaryModel) { + return; + } + + const updatedFallbacks = (value || []).map((entry) => + primaryModel in entry ? { ...entry, [primaryModel]: group.fallbackModels } : entry, + ); + + setIsSaving(true); + try { + await onChange(updatedFallbacks); + NotificationManager.success(`Fallbacks for ${primaryModel} updated successfully!`); + onClose(); + } catch (error) { + console.error("Error updating fallbacks:", error); + } finally { + setIsSaving(false); + } + }; + + return ( + + + + + Cancel + + } + onClick={handleSave} + disabled={isSaving || group.fallbackModels.length === 0} + loading={isSaving} + > + {isSaving ? "Saving Changes..." : "Save Changes"} + + + + ); +} diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx index e91c938b87a..381818a53f5 100644 --- a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx @@ -18,9 +18,16 @@ interface FallbackGroupConfigProps { onChange: (updatedGroup: FallbackGroup) => void; availableModels: string[]; maxFallbacks: number; + disablePrimaryModel?: boolean; } -export function FallbackGroupConfig({ group, onChange, availableModels, maxFallbacks }: FallbackGroupConfigProps) { +export function FallbackGroupConfig({ + group, + onChange, + availableModels, + maxFallbacks, + disablePrimaryModel = false, +}: FallbackGroupConfigProps) { // Filter available options for fallbacks (exclude primary only, allow already selected to be shown for deselection) const availableFallbackOptions = availableModels.filter((m) => m !== group.primaryModel); @@ -70,12 +77,13 @@ export function FallbackGroupConfig({ group, onChange, availableModels, maxFallb placeholder="Select primary model" value={group.primaryModel} onChange={handlePrimaryChange} + disabled={disablePrimaryModel} showSearch getPopupContainer={(trigger) => trigger.parentElement || document.body} filterOption={(input, option) => (option?.label ?? "").toLowerCase().includes(input.toLowerCase())} options={availableModels.map((m) => ({ label: m, value: m }))} /> - {!group.primaryModel && ( + {!disablePrimaryModel && !group.primaryModel && ( Select a model to begin configuring fallbacks @@ -176,6 +184,7 @@ export function FallbackGroupConfig({ group, onChange, availableModels, maxFallb removeFallback(index)} className="opacity-0 group-hover:opacity-100 transition-opacity text-gray-400 hover:text-red-500 p-1" > diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.test.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.test.tsx index 91a5b305999..a397b061801 100644 --- a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.test.tsx +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.test.tsx @@ -1,3 +1,4 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { render, screen, waitFor } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; @@ -94,6 +95,13 @@ describe("Fallbacks", () => { return deleteButtons.length > 0 ? deleteButtons[0] : null; }; + const renderWithQueryClient = (ui: React.ReactElement) => { + const queryClient = new QueryClient({ + defaultOptions: { queries: { retry: false } }, + }); + return render({ui}); + }; + beforeEach(() => { vi.clearAllMocks(); vi.mocked(networkingModule.getCallbacksCall).mockResolvedValue({ @@ -108,7 +116,7 @@ describe("Fallbacks", () => { }); it("should render the component", async () => { - render(); + renderWithQueryClient(); await waitFor(() => { expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument(); @@ -116,12 +124,12 @@ describe("Fallbacks", () => { }); it("should not render when accessToken is null", () => { - const { container } = render(); + const { container } = renderWithQueryClient(); expect(container.firstChild).toBeNull(); }); it("should fetch router settings on mount", async () => { - render(); + renderWithQueryClient(); await waitFor(() => { expect(networkingModule.getCallbacksCall).toHaveBeenCalledWith(mockAccessToken, mockUserID, mockUserRole); @@ -129,7 +137,7 @@ describe("Fallbacks", () => { }); it("should display fallback entries in table", async () => { - render(); + renderWithQueryClient(); await waitFor(() => { expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); @@ -139,7 +147,7 @@ describe("Fallbacks", () => { }); it("should show delete button for each fallback row when fallbacks exist", async () => { - render(); + renderWithQueryClient(); await waitFor(() => { expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); @@ -149,9 +157,27 @@ describe("Fallbacks", () => { expect(deleteButtons.length).toBe(2); }); + it("should show an edit button for each fallback row and open the edit modal", async () => { + const user = userEvent.setup(); + renderWithQueryClient(); + + await waitFor(() => { + expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); + }); + + const editButtons = screen.getAllByTestId("edit-fallback-button"); + expect(editButtons.length).toBe(2); + + await user.click(editButtons[0]); + + await waitFor(() => { + expect(screen.getByText("Configure Model Fallbacks")).toBeInTheDocument(); + }); + }); + it("should open delete modal when delete icon is clicked", async () => { const user = userEvent.setup(); - render(); + renderWithQueryClient(); await waitFor(() => { expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); @@ -170,7 +196,7 @@ describe("Fallbacks", () => { it("should delete fallback when confirmed", async () => { const user = userEvent.setup(); - render(); + renderWithQueryClient(); await waitFor(() => { expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); @@ -198,7 +224,7 @@ describe("Fallbacks", () => { it("should close delete modal when cancel is clicked", async () => { const user = userEvent.setup(); - render(); + renderWithQueryClient(); await waitFor(() => { expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); @@ -225,7 +251,7 @@ describe("Fallbacks", () => { const user = userEvent.setup(); const error = new Error("Delete failed"); vi.mocked(networkingModule.setCallbacksCall).mockRejectedValueOnce(error); - render(); + renderWithQueryClient(); await waitFor(() => { expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); @@ -252,7 +278,7 @@ describe("Fallbacks", () => { const user = userEvent.setup(); const error = new Error("Delete failed"); vi.mocked(networkingModule.setCallbacksCall).mockRejectedValueOnce(error); - render(); + renderWithQueryClient(); await waitFor(() => { expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); @@ -280,7 +306,7 @@ describe("Fallbacks", () => { vi.mocked(networkingModule.getCallbacksCall).mockResolvedValueOnce({ router_settings: { fallbacks: [] }, }); - render(); + renderWithQueryClient(); await waitFor(() => { expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument(); @@ -296,7 +322,7 @@ describe("Fallbacks", () => { vi.mocked(networkingModule.getCallbacksCall).mockResolvedValueOnce({ router_settings: {}, }); - render(); + renderWithQueryClient(); await waitFor(() => { expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument(); @@ -313,7 +339,7 @@ describe("Fallbacks", () => { model_group_retry_policy: { some: "policy" }, }, }); - render(); + renderWithQueryClient(); await waitFor(() => { expect(networkingModule.getCallbacksCall).toHaveBeenCalled(); @@ -322,7 +348,7 @@ describe("Fallbacks", () => { it("should update fallbacks when AddFallbacks onChange is called", async () => { const user = userEvent.setup(); - render(); + renderWithQueryClient(); await waitFor(() => { expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument(); @@ -343,7 +369,7 @@ describe("Fallbacks", () => { vi.mocked(networkingModule.getCallbacksCall).mockResolvedValue({ router_settings: mockRouterSettings, }); - render(); + renderWithQueryClient(); await waitFor(() => { expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.tsx index c1140350466..4aa9fb15705 100644 --- a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.tsx +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.tsx @@ -1,5 +1,5 @@ import { useModelCostMap } from "@/app/(dashboard)/hooks/models/useModelCostMap"; -import { ArrowRightIcon, PlayIcon, TrashIcon } from "@heroicons/react/outline"; +import { ArrowRightIcon, PencilAltIcon, PlayIcon, TrashIcon } from "@heroicons/react/outline"; import { Icon, Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow } from "@tremor/react"; import { Tooltip, Typography } from "antd"; import openai from "openai"; @@ -10,6 +10,7 @@ import NotificationsManager from "../../../molecules/notifications_manager"; import { getCallbacksCall, setCallbacksCall } from "../../../networking"; import { isProxyAdminRole } from "@/utils/roles"; import AddFallbacks from "./AddFallbacks"; +import EditFallbacks from "./EditFallbacks"; type FallbackEntry = { [modelName: string]: string[] }; type Fallbacks = FallbackEntry[]; @@ -119,6 +120,7 @@ const Fallbacks: React.FC = ({ accessToken, userRole, userID }) const [isDeleting, setIsDeleting] = useState(false); const [fallbackToDelete, setFallbackToDelete] = useState(null); const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); + const [fallbackToEdit, setFallbackToEdit] = useState(null); const { data: modelCostMapData } = useModelCostMap(); const getProviderFromModel = (model: string): string => { @@ -146,6 +148,14 @@ const Fallbacks: React.FC = ({ accessToken, userRole, userID }) setIsDeleteModalOpen(true); }; + const handleEditClick = (fallbackEntry: FallbackEntry) => { + setFallbackToEdit(fallbackEntry); + }; + + const handleEditClose = () => { + setFallbackToEdit(null); + }; + const handleDeleteConfirm = async () => { if (!fallbackToDelete || !accessToken) { return; @@ -281,6 +291,18 @@ const Fallbacks: React.FC = ({ accessToken, userRole, userID }) className="cursor-pointer hover:text-blue-600" /> + + handleEditClick(item)} + onKeyDown={(e) => e.key === "Enter" && handleEditClick(item)} + className="cursor-pointer inline-flex" + > + + + = ({ accessToken, userRole, userID }) )} + {canModify && fallbackToEdit && ( + + )}