diff --git a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx index 93a43d5e533..7ab757cfed9 100644 --- a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx +++ b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx @@ -1,6 +1,6 @@ import type { ProxyModel } from "@/app/(dashboard)/hooks/models/useModels"; import type { Organization } from "@/components/networking"; -import { screen } from "@testing-library/react"; +import { screen, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders } from "../../../tests/test-utils"; @@ -660,4 +660,181 @@ describe("ModelSelect", () => { expect(screen.getByLabelText("model-4")).toBeInTheDocument(); expect(screen.queryByLabelText("model-5")).not.toBeInTheDocument(); }); + + it("should let a selected model that is no longer offered be found and deselected", async () => { + const user = userEvent.setup(); + const liveModels: ProxyModel[] = Array.from({ length: 6 }, (_, i) => ({ + id: `model-${i}`, + object: "model", + created: 1234567890, + owned_by: "test", + })); + mockUseAllProxyModels.mockReturnValue({ + data: { data: liveModels }, + isLoading: false, + } as unknown as ReturnType); + const liveIds = liveModels.map((m) => m.id); + + renderWithProviders(); + + await openModelList(user); + await user.type(screen.getAllByRole("combobox")[0], "retired"); + const retired = await screen.findByRole("option", { name: "retired-model" }); + expect(retired).toHaveAttribute("aria-selected", "true"); + + await user.click(retired); + + expect(mockOnChange).toHaveBeenCalledWith(liveIds); + }); + + it("should list selections that are no longer offered in an Unavailable group ahead of every other group", async () => { + const user = userEvent.setup(); + renderWithProviders( + , + ); + + await openModelList(user); + + const groups = within(screen.getByRole("listbox")).getAllByRole("group"); + expect( + groups.map((group) => + within(group) + .getAllByRole("option") + .map((option) => option.textContent), + ), + ).toEqual([ + ["retired-b", "retired-a", "retired-c", "retired-d", "retired-e", "retired-f"], + ["All Proxy Models", "No Default Models"], + ["All Openai models", "All Anthropic models"], + ["gpt-4", "claude-3"], + ]); + expect(screen.getByText("Unavailable")).toBeInTheDocument(); + }); + + it("should keep an unavailable selection removable while a special option is selected", async () => { + const user = userEvent.setup(); + renderWithProviders( + , + ); + + await openModelList(user); + const retired = screen.getByRole("option", { name: "retired-model" }); + expect(retired).not.toHaveAttribute("aria-disabled", "true"); + + await user.click(retired); + + expect(mockOnChange).toHaveBeenCalledWith(["all-proxy-models"]); + }); + + it("should not mark selections Unavailable when the model list could not be loaded", async () => { + const user = userEvent.setup(); + mockUseAllProxyModels.mockReturnValue({ + data: undefined, + isLoading: false, + } as unknown as ReturnType); + + renderWithProviders( + , + ); + + await openModelList(user); + + expectOffered("All Proxy Models"); + expectNotOffered("Unavailable"); + }); + + it("should list selections outside the organization's model ceiling as Unavailable", async () => { + const user = userEvent.setup(); + mockUseOrganization.mockReturnValue({ + data: createMockOrganization(["other-model"]), + isLoading: false, + } as unknown as ReturnType); + + renderWithProviders( + , + ); + + await openModelList(user); + await user.click(screen.getByRole("option", { name: "claude-3" })); + + expect(mockOnChange).toHaveBeenCalledWith(["gpt-4"]); + }); + + it("should keep the other selections when removing one unavailable model alongside a special option", async () => { + const user = userEvent.setup(); + renderWithProviders( + , + ); + + await openModelList(user); + await user.click(screen.getByRole("option", { name: "retired-a" })); + + expect(mockOnChange).toHaveBeenCalledWith(["all-proxy-models", "retired-b"]); + }); + + it("should not mark team selections Unavailable while the organization's model ceiling is unknown", async () => { + const user = userEvent.setup(); + mockUseOrganization.mockReturnValue({ + data: undefined, + isLoading: false, + } as unknown as ReturnType); + mockUseTeam.mockReturnValue({ + data: { team_id: "team-1", organization_models: null }, + isLoading: false, + isFetching: false, + } as unknown as ReturnType); + + renderWithProviders( + , + ); + + await openModelList(user); + + expectOffered("No Default Models"); + expectNotOffered("Unavailable"); + }); + + it("should not show an Unavailable group when every selection is offered", async () => { + const user = userEvent.setup(); + renderWithProviders( + , + ); + + await openModelList(user); + + expect(screen.getByRole("option", { name: "gpt-4" })).toHaveAttribute("aria-selected", "true"); + expectNotOffered("Unavailable"); + }); }); diff --git a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx index c29eba9d997..f169382b4ed 100644 --- a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx +++ b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx @@ -126,6 +126,24 @@ const filterModels = ( return filterFn(filterArgs); }; +const isOfferedListKnown = ( + proxyModelsLoaded: boolean, + context: ModelSelectProps["context"], + organizationID: string | undefined, + organizationModels: string[] | undefined, +) => proxyModelsLoaded && !(context === "team" && organizationID !== undefined && organizationModels === undefined); + +const unavailableGroups = ( + selectedOptions: ModelOption[], + offeredByValue: Map, + offeredListKnown: boolean, +): ModelOptionGroup[] => { + if (!offeredListKnown) return []; + const items = selectedOptions.filter((option) => !offeredByValue.has(option.value)); + if (items.length === 0) return []; + return [{ label: "Unavailable", items }]; +}; + export const ModelSelect = (props: ModelSelectProps) => { const anchor = useComboboxAnchor(); const { id, teamID, organizationID, options, context, dataTestId, value = [], onChange, style } = props; @@ -151,17 +169,10 @@ export const ModelSelect = (props: ModelSelectProps) => { const handleChange = (selected: ModelOption[]) => { const values = selected.map((option) => option.value); - const specialValues = values.filter(isSpecialOption); + const addedSpecialValues = values.filter((v) => isSpecialOption(v) && !value.includes(v)); + const addedSpecial = addedSpecialValues[addedSpecialValues.length - 1]; - let finalValues: string[]; - if (specialValues.length > 0) { - const lastSelectedSpecial = specialValues[specialValues.length - 1]; - finalValues = [lastSelectedSpecial]; - } else { - finalValues = values; - } - - onChange(finalValues); + onChange(addedSpecial === undefined ? values : [addedSpecial]); }; const filteredModels = filterModels(allProxyModels?.data ?? [], props, { @@ -171,7 +182,7 @@ export const ModelSelect = (props: ModelSelectProps) => { const { wildcard, regular } = splitWildcardModels(filteredModels); - const groups: ModelOptionGroup[] = [ + const offeredGroups: ModelOptionGroup[] = [ ...(includeSpecialOptions ? [ { @@ -228,8 +239,16 @@ export const ModelSelect = (props: ModelSelectProps) => { }, ]; - const optionsByValue = new Map(groups.flatMap((group) => group.items).map((option) => [option.value, option])); - const selectedOptions = value.map((v) => optionsByValue.get(v) ?? { label: v, value: v }); + const offeredByValue = new Map(offeredGroups.flatMap((group) => group.items).map((option) => [option.value, option])); + const selectedOptions = value.map((v) => offeredByValue.get(v) ?? { label: v, value: v }); + const groups: ModelOptionGroup[] = [ + ...unavailableGroups( + selectedOptions, + offeredByValue, + isOfferedListKnown(allProxyModels !== undefined, context, organizationID, organizationModels), + ), + ...offeredGroups, + ]; const overflowOptions = selectedOptions.slice(MAX_VISIBLE_MODEL_CHIPS); return (