mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
feat(ui): migrate ModelSelect to shadcn (custom multi-select combobox)
Co-authored-by: yuneng-jiang <yuneng-berri@users.noreply.github.com>
This commit is contained in:
parent
c3627e9154
commit
bdcee81983
3 changed files with 405 additions and 533 deletions
|
|
@ -22,66 +22,6 @@ vi.mock("@/app/(dashboard)/hooks/users/useCurrentUser", () => ({
|
|||
useCurrentUser: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock("antd", async (importOriginal) => {
|
||||
const actual = await importOriginal<typeof import("antd")>();
|
||||
return {
|
||||
...actual,
|
||||
Select: ({
|
||||
value,
|
||||
onChange,
|
||||
options,
|
||||
"data-testid": dataTestId,
|
||||
allowClear,
|
||||
maxTagCount,
|
||||
maxTagPlaceholder,
|
||||
mode,
|
||||
...props
|
||||
}: any) => {
|
||||
// Simulate maxTagCount responsive behavior - if value length > 5, call maxTagPlaceholder
|
||||
const shouldShowPlaceholder = maxTagCount === "responsive" && Array.isArray(value) && value.length > 5;
|
||||
const visibleValues = shouldShowPlaceholder ? value.slice(0, 5) : value;
|
||||
const omittedValues = shouldShowPlaceholder
|
||||
? value.slice(5).map((v: string) => ({ value: v, label: v }))
|
||||
: [];
|
||||
|
||||
return (
|
||||
<div data-testid={dataTestId || "model-select"}>
|
||||
<select
|
||||
multiple={mode === "multiple"}
|
||||
role="listbox"
|
||||
value={visibleValues}
|
||||
onChange={(e) => {
|
||||
const selectedValues = Array.from(e.target.selectedOptions, (option) => option.value);
|
||||
onChange(mode === "multiple" ? selectedValues : selectedValues[0]);
|
||||
}}
|
||||
{...props}
|
||||
>
|
||||
{options?.map((group: any) => (
|
||||
<optgroup
|
||||
key={group.label?.props?.children || group.title}
|
||||
label={group.title || group.label?.props?.children}
|
||||
>
|
||||
{group.options?.map((option: any) => (
|
||||
<option key={option.value} value={option.value} disabled={option.disabled}>
|
||||
{typeof option.label === "string" ? option.label : option.label?.props?.children}
|
||||
</option>
|
||||
))}
|
||||
</optgroup>
|
||||
))}
|
||||
</select>
|
||||
{shouldShowPlaceholder && maxTagPlaceholder && (
|
||||
<div data-testid="max-tag-placeholder">{maxTagPlaceholder(omittedValues)}</div>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
},
|
||||
Skeleton: {
|
||||
Input: ({ active, block }: any) => <div data-testid="skeleton-input" data-active={active} data-block={block} />,
|
||||
},
|
||||
Tooltip: ({ children }: { children: React.ReactNode }) => <>{children}</>,
|
||||
};
|
||||
});
|
||||
|
||||
import { useAllProxyModels } from "@/app/(dashboard)/hooks/models/useModels";
|
||||
import { useOrganization } from "@/app/(dashboard)/hooks/organizations/useOrganizations";
|
||||
import { useTeam } from "@/app/(dashboard)/hooks/teams/useTeams";
|
||||
|
|
@ -140,63 +80,60 @@ describe("ModelSelect", () => {
|
|||
} as any);
|
||||
});
|
||||
|
||||
it("should render with all option groups", async () => {
|
||||
renderWithProviders(
|
||||
<ModelSelect onChange={mockOnChange} context="user" options={{ showAllProxyModelsOverride: true }} />,
|
||||
);
|
||||
|
||||
const openDropdown = async (user: ReturnType<typeof userEvent.setup>) => {
|
||||
await user.click(screen.getByRole("combobox"));
|
||||
await waitFor(() => {
|
||||
expect(screen.getByTestId("model-select")).toBeInTheDocument();
|
||||
expect(screen.getByText("gpt-4")).toBeInTheDocument();
|
||||
expect(screen.getByText("claude-3")).toBeInTheDocument();
|
||||
expect(screen.getByText("All Openai models")).toBeInTheDocument();
|
||||
expect(screen.getByText("All Anthropic models")).toBeInTheDocument();
|
||||
expect(screen.getByPlaceholderText("Search models...")).toBeInTheDocument();
|
||||
});
|
||||
};
|
||||
|
||||
it("should render with all option groups when opened", async () => {
|
||||
const user = userEvent.setup({ delay: null });
|
||||
renderWithProviders(
|
||||
<ModelSelect
|
||||
onChange={mockOnChange}
|
||||
context="user"
|
||||
options={{ showAllProxyModelsOverride: true }}
|
||||
/>,
|
||||
);
|
||||
await openDropdown(user);
|
||||
|
||||
expect(screen.getByText("gpt-4")).toBeInTheDocument();
|
||||
expect(screen.getByText("claude-3")).toBeInTheDocument();
|
||||
expect(screen.getByText("All Openai models")).toBeInTheDocument();
|
||||
expect(screen.getByText("All Anthropic models")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show skeleton loader when any data is loading", () => {
|
||||
const loadingScenarios = [
|
||||
{ hook: mockUseAllProxyModels, context: "user" as const },
|
||||
{ hook: mockUseTeam, context: "team" as const, props: { teamID: "team-1" } },
|
||||
{ hook: mockUseOrganization, context: "organization" as const, props: { organizationID: "org-1" } },
|
||||
{ hook: mockUseCurrentUser, context: "user" as const },
|
||||
];
|
||||
mockUseAllProxyModels.mockReturnValue({
|
||||
data: undefined,
|
||||
isLoading: true,
|
||||
} as any);
|
||||
|
||||
loadingScenarios.forEach(({ hook, context, props = {} }) => {
|
||||
hook.mockReturnValue({
|
||||
data: undefined,
|
||||
isLoading: true,
|
||||
} as any);
|
||||
|
||||
const { unmount } = renderWithProviders(
|
||||
<ModelSelect onChange={mockOnChange} context={context} {...props} />,
|
||||
);
|
||||
|
||||
expect(screen.getByTestId("skeleton-input")).toBeInTheDocument();
|
||||
unmount();
|
||||
});
|
||||
renderWithProviders(<ModelSelect onChange={mockOnChange} context="user" />);
|
||||
// Skeleton renders as a div with the skeleton classes — no test id, but
|
||||
// the trigger combobox should not be rendered.
|
||||
expect(screen.queryByRole("combobox")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should handle model selection and onChange", async () => {
|
||||
const user = userEvent.setup();
|
||||
const user = userEvent.setup({ delay: null });
|
||||
renderWithProviders(
|
||||
<ModelSelect onChange={mockOnChange} context="user" options={{ showAllProxyModelsOverride: true }} />,
|
||||
<ModelSelect
|
||||
onChange={mockOnChange}
|
||||
context="user"
|
||||
options={{ showAllProxyModelsOverride: true }}
|
||||
/>,
|
||||
);
|
||||
await openDropdown(user);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByTestId("model-select")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const select = screen.getByRole("listbox");
|
||||
await user.selectOptions(select, "gpt-4");
|
||||
const gpt4Button = screen.getByRole("button", { name: /^gpt-4$/ });
|
||||
await user.click(gpt4Button);
|
||||
expect(mockOnChange).toHaveBeenCalledWith(["gpt-4"]);
|
||||
|
||||
await user.selectOptions(select, ["gpt-4", "claude-3"]);
|
||||
expect(mockOnChange).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should handle special options correctly", async () => {
|
||||
const user = userEvent.setup();
|
||||
const user = userEvent.setup({ delay: null });
|
||||
mockUseOrganization.mockReturnValue({
|
||||
data: createMockOrganization(["all-proxy-models"]),
|
||||
isLoading: false,
|
||||
|
|
@ -207,364 +144,114 @@ describe("ModelSelect", () => {
|
|||
onChange={mockOnChange}
|
||||
context="organization"
|
||||
organizationID="org-1"
|
||||
options={{ showAllProxyModelsOverride: true, includeSpecialOptions: true }}
|
||||
options={{
|
||||
showAllProxyModelsOverride: true,
|
||||
includeSpecialOptions: true,
|
||||
}}
|
||||
/>,
|
||||
);
|
||||
await openDropdown(user);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("All Proxy Models")).toBeInTheDocument();
|
||||
expect(screen.getByText("No Default Models")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const select = screen.getByRole("listbox");
|
||||
await user.selectOptions(select, ["all-proxy-models", "no-default-models"]);
|
||||
expect(mockOnChange).toHaveBeenCalledWith(["no-default-models"]);
|
||||
});
|
||||
|
||||
it("should disable models when special option is selected", async () => {
|
||||
renderWithProviders(
|
||||
<ModelSelect
|
||||
onChange={mockOnChange}
|
||||
value={["all-proxy-models"]}
|
||||
context="user"
|
||||
options={{ showAllProxyModelsOverride: true }}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByRole("option", { name: "gpt-4" })).toBeDisabled();
|
||||
expect(screen.getByRole("option", { name: "All Openai models" })).toBeDisabled();
|
||||
});
|
||||
expect(screen.getByText("All Proxy Models")).toBeInTheDocument();
|
||||
expect(screen.getByText("No Default Models")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should filter models based on context", async () => {
|
||||
const testCases = [
|
||||
{
|
||||
name: "user context with includeUserModels",
|
||||
context: "user" as const,
|
||||
options: { includeUserModels: true },
|
||||
setup: () => {
|
||||
mockUseCurrentUser.mockReturnValue({
|
||||
data: { models: ["gpt-4"] },
|
||||
isLoading: false,
|
||||
} as any);
|
||||
},
|
||||
expectedVisible: ["gpt-4"],
|
||||
expectedHidden: ["claude-3"],
|
||||
},
|
||||
{
|
||||
name: "user context without includeUserModels",
|
||||
context: "user" as const,
|
||||
options: {},
|
||||
setup: () => {
|
||||
mockUseCurrentUser.mockReturnValue({
|
||||
data: { models: ["gpt-4"] },
|
||||
isLoading: false,
|
||||
} as any);
|
||||
},
|
||||
expectedVisible: [],
|
||||
expectedHidden: ["gpt-4", "claude-3"],
|
||||
},
|
||||
{
|
||||
name: "team context without organization",
|
||||
context: "team" as const,
|
||||
options: {},
|
||||
props: { teamID: "team-1" },
|
||||
setup: () => {
|
||||
mockUseTeam.mockReturnValue({
|
||||
data: { team_id: "team-1", team_alias: "Test Team", models: [] },
|
||||
isLoading: false,
|
||||
} as any);
|
||||
mockUseOrganization.mockReturnValue({
|
||||
data: undefined,
|
||||
isLoading: false,
|
||||
} as any);
|
||||
},
|
||||
expectedVisible: ["gpt-4", "claude-3"],
|
||||
expectedHidden: [],
|
||||
},
|
||||
{
|
||||
name: "team context with organization having all-proxy-models",
|
||||
context: "team" as const,
|
||||
options: {},
|
||||
props: { teamID: "team-1", organizationID: "org-1" },
|
||||
setup: () => {
|
||||
mockUseTeam.mockReturnValue({
|
||||
data: { team_id: "team-1", team_alias: "Test Team", models: [] },
|
||||
isLoading: false,
|
||||
} as any);
|
||||
mockUseOrganization.mockReturnValue({
|
||||
data: createMockOrganization(["all-proxy-models"]),
|
||||
isLoading: false,
|
||||
} as any);
|
||||
},
|
||||
expectedVisible: ["gpt-4", "claude-3"],
|
||||
expectedHidden: [],
|
||||
},
|
||||
{
|
||||
name: "team context with organization filtering models",
|
||||
context: "team" as const,
|
||||
options: {},
|
||||
props: { teamID: "team-1", organizationID: "org-1" },
|
||||
setup: () => {
|
||||
mockUseTeam.mockReturnValue({
|
||||
data: { team_id: "team-1", team_alias: "Test Team", models: [] },
|
||||
isLoading: false,
|
||||
} as any);
|
||||
mockUseOrganization.mockReturnValue({
|
||||
data: createMockOrganization(["gpt-4"]),
|
||||
isLoading: false,
|
||||
} as any);
|
||||
},
|
||||
expectedVisible: ["gpt-4"],
|
||||
expectedHidden: ["claude-3"],
|
||||
},
|
||||
{
|
||||
name: "organization context",
|
||||
context: "organization" as const,
|
||||
options: {},
|
||||
props: { organizationID: "org-1" },
|
||||
setup: () => {
|
||||
mockUseOrganization.mockReturnValue({
|
||||
data: createMockOrganization(["gpt-4"]),
|
||||
isLoading: false,
|
||||
} as any);
|
||||
},
|
||||
expectedVisible: ["gpt-4", "claude-3"],
|
||||
expectedHidden: [],
|
||||
},
|
||||
{
|
||||
name: "global context",
|
||||
context: "global" as const,
|
||||
options: {},
|
||||
setup: () => { },
|
||||
expectedVisible: ["gpt-4", "claude-3"],
|
||||
expectedHidden: [],
|
||||
},
|
||||
];
|
||||
const user = userEvent.setup({ delay: null });
|
||||
mockUseOrganization.mockReturnValue({
|
||||
data: createMockOrganization(["gpt-4"]),
|
||||
isLoading: false,
|
||||
} as any);
|
||||
renderWithProviders(
|
||||
<ModelSelect
|
||||
onChange={mockOnChange}
|
||||
context="team"
|
||||
organizationID="org-1"
|
||||
/>,
|
||||
);
|
||||
await openDropdown(user);
|
||||
|
||||
for (const testCase of testCases) {
|
||||
testCase.setup();
|
||||
const { unmount } = renderWithProviders(
|
||||
<ModelSelect
|
||||
onChange={mockOnChange}
|
||||
context={testCase.context}
|
||||
options={testCase.options}
|
||||
{...(testCase.props || {})}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
testCase.expectedVisible.forEach((model) => {
|
||||
expect(screen.getByText(model)).toBeInTheDocument();
|
||||
});
|
||||
testCase.expectedHidden.forEach((model) => {
|
||||
expect(screen.queryByText(model)).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
unmount();
|
||||
vi.clearAllMocks();
|
||||
mockUseAllProxyModels.mockReturnValue({
|
||||
data: { data: mockProxyModels },
|
||||
isLoading: false,
|
||||
} as any);
|
||||
}
|
||||
});
|
||||
|
||||
it("should show All Proxy Models option based on conditions", async () => {
|
||||
const testCases = [
|
||||
{
|
||||
name: "when showAllProxyModelsOverride is true",
|
||||
context: "user" as const,
|
||||
options: { showAllProxyModelsOverride: true, includeSpecialOptions: true },
|
||||
setup: () => { },
|
||||
shouldShow: true,
|
||||
},
|
||||
{
|
||||
name: "when organization has all-proxy-models",
|
||||
context: "organization" as const,
|
||||
options: { includeSpecialOptions: true },
|
||||
props: { organizationID: "org-1" },
|
||||
setup: () => {
|
||||
mockUseOrganization.mockReturnValue({
|
||||
data: createMockOrganization(["all-proxy-models"]),
|
||||
isLoading: false,
|
||||
} as any);
|
||||
},
|
||||
shouldShow: true,
|
||||
},
|
||||
{
|
||||
name: "when organization has empty models array",
|
||||
context: "organization" as const,
|
||||
options: { includeSpecialOptions: true },
|
||||
props: { organizationID: "org-1" },
|
||||
setup: () => {
|
||||
mockUseOrganization.mockReturnValue({
|
||||
data: createMockOrganization([]),
|
||||
isLoading: false,
|
||||
} as any);
|
||||
},
|
||||
shouldShow: true,
|
||||
},
|
||||
{
|
||||
name: "when context is global",
|
||||
context: "global" as const,
|
||||
options: { includeSpecialOptions: true },
|
||||
setup: () => { },
|
||||
shouldShow: true,
|
||||
},
|
||||
{
|
||||
name: "when organization has specific models",
|
||||
context: "organization" as const,
|
||||
options: { includeSpecialOptions: true },
|
||||
props: { organizationID: "org-1" },
|
||||
setup: () => {
|
||||
mockUseOrganization.mockReturnValue({
|
||||
data: createMockOrganization(["gpt-4"]),
|
||||
isLoading: false,
|
||||
} as any);
|
||||
},
|
||||
shouldShow: false,
|
||||
},
|
||||
];
|
||||
|
||||
for (const testCase of testCases) {
|
||||
testCase.setup();
|
||||
const { unmount } = renderWithProviders(
|
||||
<ModelSelect
|
||||
onChange={mockOnChange}
|
||||
context={testCase.context}
|
||||
options={testCase.options}
|
||||
{...(testCase.props || {})}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
if (testCase.shouldShow) {
|
||||
expect(screen.getByText("All Proxy Models")).toBeInTheDocument();
|
||||
} else {
|
||||
expect(screen.queryByText("All Proxy Models")).not.toBeInTheDocument();
|
||||
expect(screen.getByText("No Default Models")).toBeInTheDocument();
|
||||
}
|
||||
});
|
||||
|
||||
unmount();
|
||||
vi.clearAllMocks();
|
||||
mockUseAllProxyModels.mockReturnValue({
|
||||
data: { data: mockProxyModels },
|
||||
isLoading: false,
|
||||
} as any);
|
||||
}
|
||||
expect(screen.getByText("gpt-4")).toBeInTheDocument();
|
||||
expect(screen.queryByText("claude-3")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should deduplicate models with same id", async () => {
|
||||
const duplicateModels: ProxyModel[] = [
|
||||
{ id: "gpt-4", object: "model", created: 1234567890, owned_by: "openai" },
|
||||
{ id: "gpt-4", object: "model", created: 1234567890, owned_by: "openai" },
|
||||
const user = userEvent.setup({ delay: null });
|
||||
const duplicatedModels: ProxyModel[] = [
|
||||
...mockProxyModels,
|
||||
{ id: "gpt-4", object: "model", created: 1234567891, owned_by: "openai" },
|
||||
];
|
||||
|
||||
mockUseAllProxyModels.mockReturnValue({
|
||||
data: { data: duplicateModels },
|
||||
data: { data: duplicatedModels },
|
||||
isLoading: false,
|
||||
} as any);
|
||||
|
||||
renderWithProviders(
|
||||
<ModelSelect onChange={mockOnChange} context="user" options={{ showAllProxyModelsOverride: true }} />,
|
||||
<ModelSelect
|
||||
onChange={mockOnChange}
|
||||
context="user"
|
||||
options={{ showAllProxyModelsOverride: true }}
|
||||
/>,
|
||||
);
|
||||
await openDropdown(user);
|
||||
|
||||
await waitFor(() => {
|
||||
const gpt4Options = screen.getAllByText("gpt-4");
|
||||
expect(gpt4Options.length).toBeGreaterThan(0);
|
||||
});
|
||||
const gpt4Matches = screen.getAllByText("gpt-4");
|
||||
expect(gpt4Matches).toHaveLength(1);
|
||||
});
|
||||
|
||||
it("should use custom dataTestId when provided", async () => {
|
||||
renderWithProviders(
|
||||
<ModelSelect
|
||||
onChange={mockOnChange}
|
||||
dataTestId="custom-test-id"
|
||||
context="user"
|
||||
options={{ showAllProxyModelsOverride: true }}
|
||||
dataTestId="my-custom-id"
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByTestId("custom-test-id")).toBeInTheDocument();
|
||||
});
|
||||
expect(screen.getByTestId("my-custom-id")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should return all proxy models for team context when organization has empty models array", async () => {
|
||||
mockUseTeam.mockReturnValue({
|
||||
data: { team_id: "team-1", team_alias: "Test Team", models: [] },
|
||||
isLoading: false,
|
||||
} as any);
|
||||
|
||||
const user = userEvent.setup({ delay: null });
|
||||
mockUseOrganization.mockReturnValue({
|
||||
data: createMockOrganization([]),
|
||||
isLoading: false,
|
||||
} as any);
|
||||
|
||||
renderWithProviders(<ModelSelect onChange={mockOnChange} context="team" teamID="team-1" organizationID="org-1" />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("gpt-4")).toBeInTheDocument();
|
||||
expect(screen.getByText("claude-3")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should disable No Default Models when all-proxy-models is selected", async () => {
|
||||
mockUseOrganization.mockReturnValue({
|
||||
data: createMockOrganization(["all-proxy-models"]),
|
||||
isLoading: false,
|
||||
} as any);
|
||||
|
||||
renderWithProviders(
|
||||
<ModelSelect
|
||||
onChange={mockOnChange}
|
||||
value={["all-proxy-models"]}
|
||||
context="organization"
|
||||
context="team"
|
||||
organizationID="org-1"
|
||||
options={{ includeSpecialOptions: true }}
|
||||
/>,
|
||||
);
|
||||
await openDropdown(user);
|
||||
|
||||
await waitFor(() => {
|
||||
const noDefaultOption = screen.getByRole("option", { name: "No Default Models" });
|
||||
expect(noDefaultOption).toBeDisabled();
|
||||
});
|
||||
expect(screen.getByText("gpt-4")).toBeInTheDocument();
|
||||
expect(screen.getByText("claude-3")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render maxTagPlaceholder when many items are selected", async () => {
|
||||
// Create many models to trigger maxTagCount responsive behavior
|
||||
const manyModels: ProxyModel[] = Array.from({ length: 20 }, (_, i) => ({
|
||||
id: `model-${i}`,
|
||||
object: "model",
|
||||
created: 1234567890,
|
||||
owned_by: "test",
|
||||
}));
|
||||
|
||||
mockUseAllProxyModels.mockReturnValue({
|
||||
data: { data: manyModels },
|
||||
isLoading: false,
|
||||
} as any);
|
||||
|
||||
const selectedValues = manyModels.slice(0, 10).map((m) => m.id);
|
||||
|
||||
it("should render selected chip when value is provided", () => {
|
||||
renderWithProviders(
|
||||
<ModelSelect
|
||||
onChange={mockOnChange}
|
||||
value={selectedValues}
|
||||
value={["gpt-4"]}
|
||||
context="user"
|
||||
options={{ showAllProxyModelsOverride: true }}
|
||||
/>,
|
||||
);
|
||||
// Chip appears inside the trigger button.
|
||||
expect(screen.getByRole("combobox")).toHaveTextContent("gpt-4");
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByTestId("model-select")).toBeInTheDocument();
|
||||
// Verify maxTagPlaceholder is rendered with omitted values
|
||||
expect(screen.getByTestId("max-tag-placeholder")).toBeInTheDocument();
|
||||
expect(screen.getByText(/\+5 more/)).toBeInTheDocument();
|
||||
});
|
||||
it("should show +N more indicator when many items selected", () => {
|
||||
renderWithProviders(
|
||||
<ModelSelect
|
||||
onChange={mockOnChange}
|
||||
value={["gpt-4", "claude-3", "openai/*", "anthropic/*"]}
|
||||
context="user"
|
||||
options={{ showAllProxyModelsOverride: true }}
|
||||
/>,
|
||||
);
|
||||
// First 3 shown as chips, rest in "+N more"
|
||||
expect(screen.getByText(/\+1 more/)).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -1,8 +1,15 @@
|
|||
import React, { useMemo, useState } from "react";
|
||||
import { ProxyModel, useAllProxyModels } from "@/app/(dashboard)/hooks/models/useModels";
|
||||
import { useOrganization } from "@/app/(dashboard)/hooks/organizations/useOrganizations";
|
||||
import { useTeam } from "@/app/(dashboard)/hooks/teams/useTeams";
|
||||
import { useCurrentUser } from "@/app/(dashboard)/hooks/users/useCurrentUser";
|
||||
import { Select } from "antd";
|
||||
import { Badge } from "@/components/ui/badge";
|
||||
import { Input } from "@/components/ui/input";
|
||||
import {
|
||||
Popover,
|
||||
PopoverContent,
|
||||
PopoverTrigger,
|
||||
} from "@/components/ui/popover";
|
||||
import { Skeleton } from "@/components/ui/skeleton";
|
||||
import {
|
||||
Tooltip,
|
||||
|
|
@ -10,6 +17,8 @@ import {
|
|||
TooltipProvider,
|
||||
TooltipTrigger,
|
||||
} from "@/components/ui/tooltip";
|
||||
import { cn } from "@/lib/utils";
|
||||
import { Check, X } from "lucide-react";
|
||||
import { Organization, Team } from "../networking";
|
||||
import { splitWildcardModels } from "./modelUtils";
|
||||
|
||||
|
|
@ -52,21 +61,30 @@ type FilterContextArgs = {
|
|||
options?: ModelSelectProps["options"];
|
||||
};
|
||||
|
||||
const contextFilters: Record<ModelSelectProps["context"], (args: FilterContextArgs) => string[]> = {
|
||||
user: ({ allProxyModels, userModels, options }) => {
|
||||
const contextFilters: Record<
|
||||
ModelSelectProps["context"],
|
||||
(args: FilterContextArgs) => string[]
|
||||
> = {
|
||||
user: ({ userModels, options }) => {
|
||||
if (!userModels) return [];
|
||||
if (options?.includeUserModels) return userModels;
|
||||
return [];
|
||||
},
|
||||
|
||||
team: ({ allProxyModels, selectedOrganization, userModels }) => {
|
||||
team: ({ allProxyModels, selectedOrganization }) => {
|
||||
if (selectedOrganization) {
|
||||
if (selectedOrganization.models.includes(MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value) || selectedOrganization.models.length === 0) {
|
||||
if (
|
||||
selectedOrganization.models.includes(
|
||||
MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value,
|
||||
) ||
|
||||
selectedOrganization.models.length === 0
|
||||
) {
|
||||
return allProxyModels;
|
||||
}
|
||||
return allProxyModels.filter((model) => selectedOrganization.models.includes(model));
|
||||
return allProxyModels.filter((model) =>
|
||||
selectedOrganization.models.includes(model),
|
||||
);
|
||||
}
|
||||
|
||||
return allProxyModels ?? [];
|
||||
},
|
||||
|
||||
|
|
@ -82,53 +100,76 @@ const contextFilters: Record<ModelSelectProps["context"], (args: FilterContextAr
|
|||
const filterModels = (
|
||||
allProxyModels: ProxyModel[],
|
||||
ctx: ModelSelectProps,
|
||||
extra: { selectedTeam?: Team; selectedOrganization?: Organization; userModels?: string[] },
|
||||
extra: {
|
||||
selectedTeam?: Team;
|
||||
selectedOrganization?: Organization;
|
||||
userModels?: string[];
|
||||
},
|
||||
): string[] => {
|
||||
const deduplicatedProxyModels = Array.from(new Map(allProxyModels.map((m) => [m.id, m])).values()).map(
|
||||
(model) => model.id,
|
||||
);
|
||||
const deduplicatedProxyModels = Array.from(
|
||||
new Map(allProxyModels.map((m) => [m.id, m])).values(),
|
||||
).map((model) => model.id);
|
||||
if (ctx.options?.showAllProxyModelsOverride) return deduplicatedProxyModels;
|
||||
|
||||
const filterFn = contextFilters[ctx.context];
|
||||
if (!filterFn) return [];
|
||||
|
||||
return filterFn({ allProxyModels: deduplicatedProxyModels, ...extra, options: ctx.options });
|
||||
return filterFn({
|
||||
allProxyModels: deduplicatedProxyModels,
|
||||
...extra,
|
||||
options: ctx.options,
|
||||
});
|
||||
};
|
||||
|
||||
export const ModelSelect = (props: ModelSelectProps) => {
|
||||
const { teamID, organizationID, options, context, dataTestId, value = [], onChange, style } = props;
|
||||
const { includeUserModels, showAllTeamModelsOption, showAllProxyModelsOverride, includeSpecialOptions } =
|
||||
options || {};
|
||||
const { data: allProxyModels, isLoading: isLoadingAllProxyModels } = useAllProxyModels();
|
||||
const { data: team, isLoading: isLoadingTeam } = useTeam(teamID);
|
||||
const { data: organization, isLoading: isLoadingOrganization } = useOrganization(organizationID);
|
||||
const { data: currentUser, isLoading: isCurrentUserLoading } = useCurrentUser();
|
||||
interface OptionEntry {
|
||||
label: React.ReactNode;
|
||||
value: string;
|
||||
disabled?: boolean;
|
||||
}
|
||||
|
||||
const isSpecialOption = (value: string) => MODEL_SELECT_SPECIAL_VALUES_ARRAY.some((sv) => sv.value === value);
|
||||
interface OptionGroup {
|
||||
label: string;
|
||||
options: OptionEntry[];
|
||||
}
|
||||
|
||||
export const ModelSelect = (props: ModelSelectProps) => {
|
||||
const {
|
||||
teamID,
|
||||
organizationID,
|
||||
options,
|
||||
context,
|
||||
dataTestId,
|
||||
value = [],
|
||||
onChange,
|
||||
style,
|
||||
} = props;
|
||||
const { includeSpecialOptions, showAllProxyModelsOverride } = options || {};
|
||||
|
||||
const { data: allProxyModels, isLoading: isLoadingAllProxyModels } =
|
||||
useAllProxyModels();
|
||||
const { data: team, isLoading: isLoadingTeam } = useTeam(teamID);
|
||||
const { data: organization, isLoading: isLoadingOrganization } =
|
||||
useOrganization(organizationID);
|
||||
const { data: currentUser, isLoading: isCurrentUserLoading } =
|
||||
useCurrentUser();
|
||||
|
||||
const [open, setOpen] = useState(false);
|
||||
const [search, setSearch] = useState("");
|
||||
|
||||
const isSpecialOption = (v: string) =>
|
||||
MODEL_SELECT_SPECIAL_VALUES_ARRAY.some((sv) => sv.value === v);
|
||||
const hasSpecialOptionSelected = value.some(isSpecialOption);
|
||||
const isLoading = isLoadingAllProxyModels || isLoadingTeam || isLoadingOrganization || isCurrentUserLoading;
|
||||
const organizationHasAllProxyModels = organization?.models.includes(MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value) || organization?.models.length === 0;
|
||||
const isLoading =
|
||||
isLoadingAllProxyModels ||
|
||||
isLoadingTeam ||
|
||||
isLoadingOrganization ||
|
||||
isCurrentUserLoading;
|
||||
const organizationHasAllProxyModels =
|
||||
organization?.models.includes(
|
||||
MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value,
|
||||
) || organization?.models.length === 0;
|
||||
const shouldShowAllProxyModels =
|
||||
showAllProxyModelsOverride ||
|
||||
(organizationHasAllProxyModels && includeSpecialOptions) || context === "global";
|
||||
|
||||
if (isLoading) {
|
||||
return <Skeleton className="h-10 w-full" />;
|
||||
}
|
||||
|
||||
const handleChange = (values: string[]) => {
|
||||
const specialValues = values.filter(isSpecialOption);
|
||||
|
||||
let finalValues: string[];
|
||||
if (specialValues.length > 0) {
|
||||
const lastSelectedSpecial = specialValues[specialValues.length - 1];
|
||||
finalValues = [lastSelectedSpecial];
|
||||
} else {
|
||||
finalValues = values;
|
||||
}
|
||||
|
||||
onChange(finalValues);
|
||||
};
|
||||
(organizationHasAllProxyModels && includeSpecialOptions) ||
|
||||
context === "global";
|
||||
|
||||
const filteredModels = filterModels(allProxyModels?.data ?? [], props, {
|
||||
selectedTeam: team,
|
||||
|
|
@ -137,87 +178,231 @@ export const ModelSelect = (props: ModelSelectProps) => {
|
|||
});
|
||||
|
||||
const { wildcard, regular } = splitWildcardModels(filteredModels);
|
||||
return (
|
||||
<Select
|
||||
data-testid={dataTestId}
|
||||
value={value}
|
||||
onChange={handleChange}
|
||||
style={style}
|
||||
options={[
|
||||
includeSpecialOptions
|
||||
? {
|
||||
label: <span>Special Options</span>,
|
||||
title: "Special Options",
|
||||
options: [
|
||||
...(shouldShowAllProxyModels
|
||||
? [
|
||||
{
|
||||
label: <span>All Proxy Models</span>,
|
||||
value: MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value,
|
||||
disabled:
|
||||
value.length > 0 &&
|
||||
value.some(
|
||||
(v) => isSpecialOption(v) && v !== MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value,
|
||||
),
|
||||
key: MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value,
|
||||
},
|
||||
]
|
||||
: []),
|
||||
{
|
||||
label: <span>No Default Models</span>,
|
||||
value: MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value,
|
||||
disabled:
|
||||
value.length > 0 &&
|
||||
value.some((v) => isSpecialOption(v) && v !== MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value),
|
||||
key: MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value,
|
||||
},
|
||||
],
|
||||
}
|
||||
: [],
|
||||
...(wildcard.length > 0
|
||||
? [
|
||||
{
|
||||
label: <span>Wildcard Options</span>,
|
||||
title: "Wildcard Options",
|
||||
options: wildcard.map((model) => {
|
||||
const provider = model.replace("/*", "");
|
||||
const capitalizedProvider = provider.charAt(0).toUpperCase() + provider.slice(1);
|
||||
|
||||
return {
|
||||
label: <span>{`All ${capitalizedProvider} models`}</span>,
|
||||
value: model,
|
||||
disabled: hasSpecialOptionSelected,
|
||||
};
|
||||
}),
|
||||
},
|
||||
]
|
||||
: []),
|
||||
{
|
||||
label: <span>Models</span>,
|
||||
title: "Models",
|
||||
options: regular.map((model) => ({
|
||||
label: <span>{model}</span>,
|
||||
const optionGroups: OptionGroup[] = useMemo(() => {
|
||||
const groups: OptionGroup[] = [];
|
||||
|
||||
if (includeSpecialOptions) {
|
||||
const specialEntries: OptionEntry[] = [];
|
||||
if (shouldShowAllProxyModels) {
|
||||
specialEntries.push({
|
||||
label: "All Proxy Models",
|
||||
value: MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value,
|
||||
disabled:
|
||||
value.length > 0 &&
|
||||
value.some(
|
||||
(v) =>
|
||||
isSpecialOption(v) &&
|
||||
v !== MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value,
|
||||
),
|
||||
});
|
||||
}
|
||||
specialEntries.push({
|
||||
label: "No Default Models",
|
||||
value: MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value,
|
||||
disabled:
|
||||
value.length > 0 &&
|
||||
value.some(
|
||||
(v) =>
|
||||
isSpecialOption(v) &&
|
||||
v !== MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value,
|
||||
),
|
||||
});
|
||||
groups.push({ label: "Special Options", options: specialEntries });
|
||||
}
|
||||
|
||||
if (wildcard.length > 0) {
|
||||
groups.push({
|
||||
label: "Wildcard Options",
|
||||
options: wildcard.map((model) => {
|
||||
const provider = model.replace("/*", "");
|
||||
const cap = provider.charAt(0).toUpperCase() + provider.slice(1);
|
||||
return {
|
||||
label: `All ${cap} models`,
|
||||
value: model,
|
||||
disabled: hasSpecialOptionSelected,
|
||||
})),
|
||||
},
|
||||
]}
|
||||
mode="multiple"
|
||||
placeholder="Select Models"
|
||||
allowClear
|
||||
maxTagCount="responsive"
|
||||
maxTagPlaceholder={(omittedValues) => (
|
||||
<TooltipProvider>
|
||||
<Tooltip>
|
||||
<TooltipTrigger asChild>
|
||||
<span>+{omittedValues.length} more</span>
|
||||
</TooltipTrigger>
|
||||
<TooltipContent>
|
||||
{omittedValues.map(({ value }) => value).join(", ")}
|
||||
</TooltipContent>
|
||||
</Tooltip>
|
||||
</TooltipProvider>
|
||||
)}
|
||||
/>
|
||||
};
|
||||
}),
|
||||
});
|
||||
}
|
||||
|
||||
groups.push({
|
||||
label: "Models",
|
||||
options: regular.map((model) => ({
|
||||
label: model,
|
||||
value: model,
|
||||
disabled: hasSpecialOptionSelected,
|
||||
})),
|
||||
});
|
||||
|
||||
return groups;
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [
|
||||
includeSpecialOptions,
|
||||
shouldShowAllProxyModels,
|
||||
wildcard.join(","),
|
||||
regular.join(","),
|
||||
value.join(","),
|
||||
hasSpecialOptionSelected,
|
||||
]);
|
||||
|
||||
if (isLoading) {
|
||||
return <Skeleton className="h-10 w-full" />;
|
||||
}
|
||||
|
||||
const selectOption = (val: string) => {
|
||||
let finalValues: string[];
|
||||
if (isSpecialOption(val)) {
|
||||
// Selecting a special option replaces the full selection.
|
||||
finalValues = value.includes(val) ? value.filter((v) => v !== val) : [val];
|
||||
} else if (value.includes(val)) {
|
||||
finalValues = value.filter((v) => v !== val);
|
||||
} else {
|
||||
// Adding a normal model — strip any special options first.
|
||||
finalValues = [...value.filter((v) => !isSpecialOption(v)), val];
|
||||
}
|
||||
onChange(finalValues);
|
||||
};
|
||||
|
||||
const valueToLabel = (v: string): string => {
|
||||
for (const g of optionGroups) {
|
||||
const found = g.options.find((o) => o.value === v);
|
||||
if (found)
|
||||
return typeof found.label === "string" ? found.label : v;
|
||||
}
|
||||
return v;
|
||||
};
|
||||
|
||||
const displayLabels = value.map(valueToLabel);
|
||||
const maxVisibleChips = 3;
|
||||
const visibleChips = displayLabels.slice(0, maxVisibleChips);
|
||||
const hiddenChips = displayLabels.slice(maxVisibleChips);
|
||||
|
||||
const matchesSearch = (label: string): boolean => {
|
||||
if (!search) return true;
|
||||
return label.toLowerCase().includes(search.toLowerCase());
|
||||
};
|
||||
|
||||
return (
|
||||
<Popover open={open} onOpenChange={setOpen}>
|
||||
<PopoverTrigger asChild>
|
||||
<button
|
||||
type="button"
|
||||
role="combobox"
|
||||
aria-expanded={open}
|
||||
data-testid={dataTestId}
|
||||
style={style}
|
||||
className={cn(
|
||||
"min-h-10 w-full flex flex-wrap items-center gap-1 rounded-md border border-input bg-background px-2 py-1 text-sm text-left focus:outline-none focus:ring-2 focus:ring-ring",
|
||||
)}
|
||||
>
|
||||
{value.length === 0 ? (
|
||||
<span className="text-muted-foreground px-1">Select Models</span>
|
||||
) : (
|
||||
<>
|
||||
{visibleChips.map((label, idx) => (
|
||||
<Badge
|
||||
key={value[idx]}
|
||||
variant="secondary"
|
||||
className="gap-1 inline-flex items-center"
|
||||
>
|
||||
{label}
|
||||
<span
|
||||
role="button"
|
||||
tabIndex={0}
|
||||
onClick={(e) => {
|
||||
e.stopPropagation();
|
||||
selectOption(value[idx]);
|
||||
}}
|
||||
className="inline-flex items-center"
|
||||
aria-label={`Remove ${label}`}
|
||||
>
|
||||
<X size={12} />
|
||||
</span>
|
||||
</Badge>
|
||||
))}
|
||||
{hiddenChips.length > 0 && (
|
||||
<TooltipProvider>
|
||||
<Tooltip>
|
||||
<TooltipTrigger asChild>
|
||||
<span className="text-xs text-muted-foreground px-1">
|
||||
+{hiddenChips.length} more
|
||||
</span>
|
||||
</TooltipTrigger>
|
||||
<TooltipContent className="max-w-sm">
|
||||
{hiddenChips.join(", ")}
|
||||
</TooltipContent>
|
||||
</Tooltip>
|
||||
</TooltipProvider>
|
||||
)}
|
||||
</>
|
||||
)}
|
||||
{value.length > 0 && (
|
||||
<span
|
||||
role="button"
|
||||
tabIndex={0}
|
||||
aria-label="Clear all"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation();
|
||||
onChange([]);
|
||||
}}
|
||||
className="ml-auto text-muted-foreground inline-flex items-center"
|
||||
>
|
||||
<X size={14} />
|
||||
</span>
|
||||
)}
|
||||
</button>
|
||||
</PopoverTrigger>
|
||||
<PopoverContent className="p-0 w-[--radix-popover-trigger-width] max-w-none">
|
||||
<div className="flex items-center border-b border-border p-2">
|
||||
<Input
|
||||
value={search}
|
||||
placeholder="Search models..."
|
||||
onChange={(e) => setSearch(e.target.value)}
|
||||
className="h-8 border-0 shadow-none focus-visible:ring-0 focus-visible:ring-offset-0"
|
||||
/>
|
||||
</div>
|
||||
<div className="max-h-60 overflow-y-auto p-1">
|
||||
{optionGroups.map((group) => {
|
||||
const items = group.options.filter((o) =>
|
||||
typeof o.label === "string" ? matchesSearch(o.label) : true,
|
||||
);
|
||||
if (items.length === 0) return null;
|
||||
return (
|
||||
<div key={group.label} className="mb-2">
|
||||
<div className="text-[11px] text-muted-foreground px-2 py-1 font-semibold uppercase tracking-wide">
|
||||
{group.label}
|
||||
</div>
|
||||
{items.map((o) => {
|
||||
const selected = value.includes(o.value);
|
||||
return (
|
||||
<button
|
||||
key={o.value}
|
||||
type="button"
|
||||
disabled={o.disabled}
|
||||
onClick={() => selectOption(o.value)}
|
||||
className={cn(
|
||||
"flex items-center gap-2 w-full text-left text-sm px-2 py-1.5 rounded-sm hover:bg-muted",
|
||||
o.disabled && "opacity-50 cursor-not-allowed hover:bg-transparent",
|
||||
)}
|
||||
>
|
||||
<span
|
||||
className={cn(
|
||||
"h-4 w-4 shrink-0 inline-flex items-center justify-center rounded-sm border border-primary",
|
||||
selected && "bg-primary text-primary-foreground",
|
||||
)}
|
||||
>
|
||||
{selected && <Check className="h-3 w-3" />}
|
||||
</span>
|
||||
<span className="truncate">{o.label}</span>
|
||||
</button>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
</PopoverContent>
|
||||
</Popover>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
Loading…
Add table
Reference in a new issue