fix(ui): scope shadow eval models to configured chat groups (#40488)

Co-authored-by: Claude Code <noreply@anthropic.com>
This commit is contained in:
tin-berri 2026-09-09 18:09:28 -07:00 committed by GitHub
parent 317b29e69d
commit e34c4c8edc
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 317 additions and 157 deletions

View file

@ -1,3 +1,4 @@
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { fireEvent, render, screen, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import React from "react";
@ -14,7 +15,9 @@ vi.mock("./useShadowEval", () => ({
}));
const authorizedRoleMock = vi.fn(() => ({ accessToken: "token", isViewOnly: false }));
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: () => authorizedRoleMock() }));
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
default: () => ({ userId: "test-user-id", userRole: "Admin", ...authorizedRoleMock() }),
}));
vi.mock("@/app/(dashboard)/hooks/keys/useKeys", () => ({
useInfiniteKeys: vi.fn(() => ({
@ -68,27 +71,33 @@ vi.mock("@/app/(dashboard)/hooks/users/useUsers", () => ({
})),
}));
vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({
vi.mock("@/app/(dashboard)/hooks/models/useModels", async (importOriginal) => ({
...(await importOriginal<typeof import("@/app/(dashboard)/hooks/models/useModels")>()),
useAutoRouters: vi.fn(() => ({
data: [
{ model_name: "claude-auto", litellm_params: { model: "auto_router/claude-auto" } },
{ model_name: "gpt-auto", litellm_params: { model: "auto_router/gpt-auto" } },
],
})),
usePlainModelGroups: vi.fn(() => new Set(["prod-claude"])),
usePlainModelGroups: vi.fn(() => new Set(["prod-claude", "prod-judge"])),
usePlainChatModelGroups: vi.fn(() => new Set(["prod-claude", "prod-judge"])),
usePlainChatModelDeployments: vi.fn(() => [
{
model_name: "prod-judge",
litellm_params: { model: "anthropic/claude-sonnet-5" },
model_info: { mode: "chat" },
},
]),
}));
vi.mock("@/app/(dashboard)/hooks/models/useModelCostMap", () => ({
useModelCostMap: vi.fn(() => ({
data: {
"claude-sonnet-5": { litellm_provider: "anthropic", mode: "chat" },
"gpt-4o": { litellm_provider: "openai", mode: "chat" },
"gemini/gemini-2.5-pro": { litellm_provider: "gemini", mode: "chat" },
"text-embedding-3-large": { litellm_provider: "openai", mode: "embedding" },
},
})),
vi.mock("@/components/networking", async (importOriginal) => ({
...(await importOriginal<typeof import("@/components/networking")>()),
modelInfoCall: vi.fn(),
}));
import { usePlainChatModelGroups, usePlainModelGroups } from "@/app/(dashboard)/hooks/models/useModels";
import { modelInfoCall } from "@/components/networking";
import ShadowEvalSection, { shadowedTargetLabel } from "./ShadowEvalSection";
import {
useShadowEvalJob,
@ -107,7 +116,7 @@ const job = (overrides: Partial<ShadowEvalJob> = {}): ShadowEvalJob => ({
models: [],
direction: "forward",
baseline_model: null,
judge_model: "anthropic/claude-sonnet-5",
judge_model: "prod-judge",
shadow_percentage: 10,
targets: [
{
@ -249,6 +258,85 @@ describe("ShadowEvalSection", () => {
if (defaultKeysImpl) vi.mocked(useInfiniteKeys).mockImplementation(defaultKeysImpl);
});
it("labels only configured judge recommendations", async () => {
const user = userEvent.setup();
mockHooks({});
render(<ShadowEvalSection />);
await user.click(screen.getByPlaceholderText("Select a judge model"));
expect(screen.getByRole("option", { name: /prod-judge.*Recommended/ })).toBeInTheDocument();
expect(screen.queryByRole("option", { name: /openai\/gpt-4o/ })).not.toBeInTheDocument();
await user.keyboard("{Escape}");
await chooseSelectOption(
user,
screen.getByText("Adoption check: key's traffic vs the router"),
"Regression check: router's picks vs a baseline",
);
await user.click(screen.getByPlaceholderText("Select a baseline model"));
expect(screen.getByRole("option", { name: "prod-judge", exact: true })).toBeInTheDocument();
expect(screen.queryByText("Recommended")).not.toBeInTheDocument();
});
it("keeps custom models selectable through the real model hooks without widening chat choices to traffic filters", async () => {
const hooks = await vi.importActual<typeof import("@/app/(dashboard)/hooks/models/useModels")>(
"@/app/(dashboard)/hooks/models/useModels",
);
const client = new QueryClient({ defaultOptions: { queries: { retry: false } } });
const deployments = [
{ model_name: "custom-chat", litellm_params: { model: "openai/private-chat" } },
{ model_name: "custom-judge", litellm_params: { model: "openai/private-judge" }, model_info: { mode: null } },
{
model_name: "embedding",
litellm_params: { model: "openai/private-embedding" },
model_info: { mode: "embedding" },
},
{
model_name: "responses-only",
litellm_params: { model: "openai/private-responses" },
model_info: { mode: "responses" },
},
{ model_name: "auto-router", litellm_params: { model: "auto_router/complexity_router" } },
];
vi.mocked(modelInfoCall).mockResolvedValue({ data: deployments, total_pages: 1 });
const user = userEvent.setup();
const { start } = mockHooks({});
await vi.mocked(usePlainModelGroups).withImplementation(hooks.usePlainModelGroups, async () => {
await vi.mocked(usePlainChatModelGroups).withImplementation(hooks.usePlainChatModelGroups, async () => {
render(
<QueryClientProvider client={client}>
<ShadowEvalSection />
</QueryClientProvider>,
);
await chooseSelectOption(user, screen.getByPlaceholderText("Every model the targets use"), "responses-only");
await chooseSelectOption(user, screen.getByPlaceholderText("Every model the targets use"), "custom-chat");
await chooseSelectOption(
user,
screen.getByText("Adoption check: key's traffic vs the router"),
"Regression check: router's picks vs a baseline",
);
await user.click(screen.getByPlaceholderText("Search keys by alias"));
await user.click(within(await screen.findByTestId("paginated-multi-select-list")).getByText("prod-alpha"));
await chooseSelectOption(user, screen.getByPlaceholderText("Select up to 4 auto-routers"), "gpt-auto");
await user.click(screen.getByPlaceholderText("Select a judge model"));
expect(screen.getAllByRole("option")).toHaveLength(2);
expect(screen.getByRole("option", { name: "custom-chat", exact: true })).toBeInTheDocument();
expect(screen.getByRole("option", { name: "custom-judge", exact: true })).toBeInTheDocument();
await user.click(screen.getByRole("option", { name: "custom-judge", exact: true }));
await user.click(screen.getByPlaceholderText("Select a baseline model"));
expect(screen.getAllByRole("option")).toHaveLength(2);
expect(screen.getByRole("option", { name: "custom-chat", exact: true })).toBeInTheDocument();
expect(screen.getByRole("option", { name: "custom-judge", exact: true })).toBeInTheDocument();
await user.click(screen.getByRole("option", { name: "custom-chat", exact: true }));
await user.click(screen.getByText("Start shadow eval"));
expect(start.mutate).toHaveBeenCalledWith(
expect.objectContaining({ judge_model: "custom-judge", baseline_model: "custom-chat", models: [] }),
);
});
});
client.clear();
});
it("offers the start form while the list is still loading", () => {
mockHooks({ isPending: true });
render(<ShadowEvalSection />);
@ -444,7 +532,8 @@ describe("ShadowEvalSection", () => {
expect(screen.getByText("Start shadow eval")).toBeDisabled();
await user.click(screen.getByPlaceholderText("Select a judge model"));
await user.click(await screen.findByRole("option", { name: /anthropic\/claude-sonnet-5/ }));
expect(screen.queryByRole("option", { name: /openai\/gpt-4o/ })).not.toBeInTheDocument();
await user.click(await screen.findByRole("option", { name: /prod-judge/ }));
await user.click(screen.getByText("Start shadow eval"));
const expectedBody = {
@ -457,7 +546,7 @@ describe("ShadowEvalSection", () => {
shadow_percentage: 10,
duration_days: 7,
max_budget: 10,
judge_model: "anthropic/claude-sonnet-5",
judge_model: "prod-judge",
};
expect(start.mutate).toHaveBeenCalledWith(expectedBody);
});
@ -474,7 +563,7 @@ describe("ShadowEvalSection", () => {
await user.click(within(teamList).getByText("engineering"));
await chooseSelectOption(user, screen.getByPlaceholderText("Select up to 4 auto-routers"), "gpt-auto");
await user.click(screen.getByPlaceholderText("Select a judge model"));
await user.click(await screen.findByRole("option", { name: /anthropic\/claude-sonnet-5/ }));
await user.click(await screen.findByRole("option", { name: /prod-judge/ }));
await user.click(screen.getByText("Start shadow eval"));
const expectedBody = {
@ -487,7 +576,7 @@ describe("ShadowEvalSection", () => {
shadow_percentage: 10,
duration_days: 7,
max_budget: 10,
judge_model: "anthropic/claude-sonnet-5",
judge_model: "prod-judge",
};
expect(start.mutate).toHaveBeenCalledWith(expectedBody);
});
@ -503,7 +592,7 @@ describe("ShadowEvalSection", () => {
await chooseSelectOption(user, screen.getByPlaceholderText("Every model the targets use"), "prod-claude");
await chooseSelectOption(user, screen.getByPlaceholderText("Select up to 4 auto-routers"), "gpt-auto");
await user.click(screen.getByPlaceholderText("Select a judge model"));
await user.click(await screen.findByRole("option", { name: /anthropic\/claude-sonnet-5/ }));
await user.click(await screen.findByRole("option", { name: /prod-judge/ }));
await user.click(screen.getByText("Start shadow eval"));
expect(start.mutate).toHaveBeenCalledWith(
@ -524,20 +613,23 @@ describe("ShadowEvalSection", () => {
expect(screen.queryByPlaceholderText("Select a baseline model")).not.toBeInTheDocument();
expect(screen.getByPlaceholderText("Every model the targets use")).toBeInTheDocument();
await user.click(screen.getByText("Adoption check: key's traffic vs the router"));
await user.click(await screen.findByText("Regression check: router's picks vs a baseline"));
await chooseSelectOption(
user,
screen.getByText("Adoption check: key's traffic vs the router"),
"Regression check: router's picks vs a baseline",
);
expect(screen.queryByPlaceholderText("Every model the targets use")).not.toBeInTheDocument();
await user.click(screen.getByPlaceholderText("Search keys by alias"));
const keyList = await screen.findByTestId("paginated-multi-select-list");
await user.click(within(keyList).getByText("prod-alpha"));
await chooseSelectOption(user, screen.getByPlaceholderText("Select up to 4 auto-routers"), "gpt-auto");
await user.click(screen.getByPlaceholderText("Select a judge model"));
await user.click(await screen.findByRole("option", { name: /anthropic\/claude-sonnet-5/ }));
await user.click(await screen.findByRole("option", { name: /prod-judge/ }));
expect(screen.getByText("Start shadow eval")).toBeDisabled();
await user.click(screen.getByPlaceholderText("Select a baseline model"));
expect(await screen.findByRole("option", { name: /openai\/gpt-4o/ })).toBeInTheDocument();
expect(screen.queryByRole("option", { name: /openai\/gpt-4o/ })).not.toBeInTheDocument();
await user.click(screen.getByRole("option", { name: /prod-claude/ }));
await user.click(screen.getByText("Start shadow eval"));
@ -552,7 +644,7 @@ describe("ShadowEvalSection", () => {
shadow_percentage: 10,
duration_days: 7,
max_budget: 10,
judge_model: "anthropic/claude-sonnet-5",
judge_model: "prod-judge",
};
expect(start.mutate).toHaveBeenCalledWith(expectedBody);
});
@ -574,7 +666,7 @@ describe("ShadowEvalSection", () => {
screen.getByText("Every router sees the same sampled requests, judged against the same live responses"),
).toBeInTheDocument();
await user.click(screen.getByPlaceholderText("Select a judge model"));
await user.click(await screen.findByRole("option", { name: /anthropic\/claude-sonnet-5/ }));
await user.click(await screen.findByRole("option", { name: /prod-judge/ }));
await user.click(screen.getByText("Start shadow eval"));
const expectedBody = {
@ -587,7 +679,7 @@ describe("ShadowEvalSection", () => {
shadow_percentage: 10,
duration_days: 7,
max_budget: 10,
judge_model: "anthropic/claude-sonnet-5",
judge_model: "prod-judge",
};
expect(start.mutate).toHaveBeenCalledWith(expectedBody);
});
@ -605,10 +697,13 @@ describe("ShadowEvalSection", () => {
await user.click(await screen.findByText("gpt-auto"));
await user.click(routerInput);
await user.click(await screen.findByText("claude-auto"));
await user.click(screen.getByText("Adoption check: key's traffic vs the router"));
await user.click(await screen.findByText("Regression check: router's picks vs a baseline"));
await chooseSelectOption(
user,
screen.getByText("Adoption check: key's traffic vs the router"),
"Regression check: router's picks vs a baseline",
);
await user.click(screen.getByPlaceholderText("Select a judge model"));
await user.click(await screen.findByRole("option", { name: /anthropic\/claude-sonnet-5/ }));
await user.click(await screen.findByRole("option", { name: /prod-judge/ }));
await user.click(screen.getByPlaceholderText("Select a baseline model"));
await user.click(screen.getByRole("option", { name: /prod-claude/ }));

View file

@ -5,8 +5,13 @@ import React, { useMemo, useState } from "react";
import { useInfiniteKeys } from "@/app/(dashboard)/hooks/keys/useKeys";
import { useInfiniteUsers } from "@/app/(dashboard)/hooks/users/useUsers";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { useModelCostMap } from "@/app/(dashboard)/hooks/models/useModelCostMap";
import { useAutoRouters, usePlainModelGroups } from "@/app/(dashboard)/hooks/models/useModels";
import {
useAutoRouters,
usePlainChatModelDeployments,
usePlainChatModelGroups,
usePlainModelGroups,
} from "@/app/(dashboard)/hooks/models/useModels";
import { buildModelAvailability, deploymentRefsFromModelInfo, resolveAvailableModels } from "@/lib/autorouter_presets";
import { MultiSelect } from "@/components/shared/MultiSelect";
import { PaginatedMultiSelect } from "@/components/shared/PaginatedMultiSelect";
import TeamMultiSelect from "@/components/common_components/team_multi_select";
@ -24,53 +29,8 @@ type ShadowEvalDirection = ShadowEvalJob["direction"];
const MAX_ROUTERS = 4;
const MAX_MODELS = 100;
const RECOMMENDED_JUDGE_MODELS = ["anthropic/claude-sonnet-5", "openai/gpt-4o", "gemini/gemini-2.5-pro"] as const;
interface CostMapEntry {
litellm_provider?: string;
mode?: string;
}
const useChatModelNames = (): string[] => {
const { data: costMap } = useModelCostMap();
return useMemo(() => {
if (!costMap) return [];
const chatModels = Object.entries(costMap as Record<string, CostMapEntry>)
.filter(([, value]) => value?.mode === "chat" && value?.litellm_provider)
.map(([key, value]) => (key.startsWith(`${value.litellm_provider}/`) ? key : `${value.litellm_provider}/${key}`));
return [...new Set(chatModels)].toSorted((a, b) => a.localeCompare(b));
}, [costMap]);
};
const useJudgeModelOptions = (): SearchSelectOption[] => {
const chatModels = useChatModelNames();
return useMemo(() => {
const pinned: SearchSelectOption[] = RECOMMENDED_JUDGE_MODELS.map((model) => ({
label: model,
value: model,
sublabel: "Recommended",
}));
const pinnedNames = new Set<string>(RECOMMENDED_JUDGE_MODELS);
const rest = chatModels.filter((model) => !pinnedNames.has(model)).map((model) => ({ label: model, value: model }));
return [...pinned, ...rest];
}, [chatModels]);
};
const useBaselineModelOptions = (): SearchSelectOption[] => {
const configuredGroups = usePlainModelGroups();
const chatModels = useChatModelNames();
return useMemo(() => {
const configured = [...configuredGroups]
.toSorted((a, b) => a.localeCompare(b))
.map((model) => ({ label: model, value: model, sublabel: "Configured on this gateway" }));
const rest = chatModels
.filter((model) => !configuredGroups.has(model))
.map((model) => ({ label: model, value: model }));
return [...configured, ...rest];
}, [configuredGroups, chatModels]);
};
const DIRECTION_OPTIONS: readonly { value: ShadowEvalDirection; label: string }[] = [
{ value: "forward", label: "Adoption check: key's traffic vs the router" },
{ value: "reverse", label: "Regression check: router's picks vs a baseline" },
@ -276,13 +236,32 @@ export const StartForm: React.FC = () => {
const [judgeModel, setJudgeModel] = useState("");
const [maxBudget, setMaxBudget] = useState("10");
const { data: autoRouters } = useAutoRouters();
const judgeModelOptions = useJudgeModelOptions();
const baselineModelOptions = useBaselineModelOptions();
const configuredGroups = usePlainModelGroups();
const chatGroups = usePlainChatModelGroups();
const chatDeployments = usePlainChatModelDeployments();
const modelOptions = useMemo<SearchSelectOption[]>(
() => [...configuredGroups].toSorted((a, b) => a.localeCompare(b)).map((name) => ({ label: name, value: name })),
[configuredGroups],
);
const chatOptions = useMemo(
() => modelOptions.filter((option) => chatGroups.has(option.value)),
[modelOptions, chatGroups],
);
const chatAvailability = useMemo(
() => buildModelAvailability(chatGroups, deploymentRefsFromModelInfo(chatDeployments)),
[chatDeployments, chatGroups],
);
const recommendedJudgeModels = useMemo(
() => new Set(RECOMMENDED_JUDGE_MODELS.flatMap((model) => resolveAvailableModels(model, chatAvailability))),
[chatAvailability],
);
const judgeOptions = useMemo(
() =>
chatOptions.map((option) =>
recommendedJudgeModels.has(option.value) ? { ...option, sublabel: "Recommended" } : option,
),
[chatOptions, recommendedJudgeModels],
);
const start = useStartShadowEval();
const routerOptions = useMemo<SearchSelectOption[]>(() => {
@ -434,7 +413,7 @@ export const StartForm: React.FC = () => {
{direction === "reverse" && (
<Field label="Baseline model">
<SearchSelect
options={baselineModelOptions}
options={chatOptions}
value={baselineModel}
onValueChange={setBaselineModel}
placeholder="Select a baseline model"
@ -444,7 +423,7 @@ export const StartForm: React.FC = () => {
)}
<Field label="Judge model" className="sm:col-span-2">
<SearchSelect
options={judgeModelOptions}
options={judgeOptions}
value={judgeModel}
onValueChange={setJudgeModel}
placeholder="Select a judge model"

View file

@ -5,17 +5,19 @@ import { beforeEach, describe, expect, it, vi } from "vitest";
import {
isAutoRouterDeployment,
selectAutoRouterModelGroups,
selectPlainModelGroups,
selectPlainChatModelGroups,
useAllProxyModels,
useAutoRouterModelGroups,
useAutoRouters,
useInfiniteModelInfo,
useModelHub,
useModelsInfo,
usePlainChatModelGroups,
useSelectedTeamModels,
useUserModels,
type AllProxyModelsResponse,
type AutoRouterCandidateDeployment,
type AutoRouterDeployment,
type PaginatedModelInfoResponse,
type ProxyModel,
} from "./useModels";
@ -984,29 +986,45 @@ describe("selectAutoRouterModelGroups", () => {
});
});
describe("selectPlainModelGroups", () => {
it("keeps only non-auto-router model groups", () => {
const deployments: AutoRouterCandidateDeployment[] = [
{ model_name: "smart-router", litellm_params: { model: "auto_router/complexity_router" } },
{ model_name: "claude-haiku", litellm_params: { model: "anthropic/claude-haiku-4-5" } },
{ model_name: "claude-sonnet", litellm_params: { model: "anthropic/claude-sonnet-4-5" } },
{ model_name: "cheap-router", litellm_params: { model: "auto_router/adaptive_router" } },
describe("selectPlainChatModelGroups", () => {
it("keeps chat-capable groups when mode metadata is absent or any sibling is compatible", () => {
const deployments: AutoRouterDeployment[] = [
{ model_name: "no-info" },
{ model_name: "null-info", model_info: null },
{ model_name: "empty-info", model_info: {} },
{ model_name: "missing-mode", model_info: { db_model: false } },
{ model_name: "null-mode", model_info: { mode: null } },
{ model_name: "empty-mode", model_info: { mode: "" } },
{ model_name: "chat", model_info: { mode: "chat", db_model: true } },
{ model_name: "completion", model_info: { mode: "completion" } },
{ model_name: "chat-and-missing", model_info: { mode: "chat" } },
{ model_name: "chat-and-missing" },
{ model_name: "chat-then-embedding", model_info: { mode: "chat" } },
{ model_name: "chat-then-embedding", model_info: { mode: "embedding" } },
{ model_name: "embedding-then-chat", model_info: { mode: "embedding" } },
{ model_name: "embedding-then-chat", model_info: { mode: "chat" } },
{ model_name: "embedding-only", model_info: { mode: "embedding" } },
{ model_name: "speech-only", model_info: { mode: "speech" } },
{ model_name: "shared-router", litellm_params: { model: "openai/gpt-4o" } },
{ model_name: "shared-router", litellm_params: { model: "auto_router/complexity_router" } },
{ model_name: "", model_info: { mode: "chat" } },
];
expect(selectPlainModelGroups(deployments)).toEqual(new Set(["claude-haiku", "claude-sonnet"]));
});
it("drops a group name that also fronts an auto-router deployment", () => {
const deployments: AutoRouterCandidateDeployment[] = [
{ model_name: "shared-name", litellm_params: { model: "auto_router/complexity_router" } },
{ model_name: "shared-name", litellm_params: { model: "anthropic/claude-sonnet-4-5" } },
];
expect(selectPlainModelGroups(deployments)).toEqual(new Set());
});
it("drops deployments that have no public model_name", () => {
expect(selectPlainModelGroups([{ model_name: "", litellm_params: { model: "openai/gpt-4o" } }])).toEqual(new Set());
expect(selectPlainChatModelGroups(deployments)).toEqual(
new Set([
"no-info",
"null-info",
"empty-info",
"missing-mode",
"null-mode",
"empty-mode",
"chat",
"completion",
"chat-and-missing",
"chat-then-embedding",
"embedding-then-chat",
]),
);
});
});
@ -1103,6 +1121,47 @@ describe("useAutoRouterModelGroups", () => {
expect(modelInfoCall).toHaveBeenCalledWith("test-access-token", "test-user-id", "Admin", 3, 1000);
});
it("uses every page for configured chat groups and keeps custom deployments without mode metadata", async () => {
(modelInfoCall as any).mockImplementation((_t: string, _u: string, _r: string, page: number) =>
Promise.resolve(
page === 1
? {
data: [
{ model_name: "configured-chat", model_info: { mode: "chat" } },
{ model_name: "embedding-only", model_info: { mode: "embedding" } },
],
total_pages: 2,
}
: {
data: [
{ model_name: "custom-no-mode", model_info: { db_model: true } },
{ model_name: "speech-only", model_info: { mode: "speech" } },
],
total_pages: 2,
},
),
);
const { result } = renderHook(() => usePlainChatModelGroups(), { wrapper });
await waitFor(() => expect(result.current.size).toBe(2));
expect(result.current).toEqual(new Set(["configured-chat", "custom-no-mode"]));
expect(modelInfoCall).toHaveBeenCalledTimes(2);
});
it("returns an empty chat group set while loading and after failure", async () => {
(modelInfoCall as any).mockReturnValueOnce(new Promise(() => {}));
const loading = renderHook(() => usePlainChatModelGroups(), { wrapper });
expect(loading.result.current).toEqual(new Set());
loading.unmount();
queryClient.clear();
(modelInfoCall as any).mockRejectedValueOnce(new Error("boom"));
const failed = renderHook(() => usePlainChatModelGroups(), { wrapper });
await waitFor(() => expect(modelInfoCall).toHaveBeenCalledTimes(2));
expect(failed.result.current).toEqual(new Set());
});
it("returns an empty set before the model list resolves", () => {
(modelInfoCall as any).mockReturnValue(new Promise(() => {}));

View file

@ -2,6 +2,7 @@ import { useQuery, useInfiniteQuery, useQueryClient, UseQueryResult } from "@tan
import { createQueryKeys } from "../common/queryKeysFactory";
import { modelInfoCall, modelHubCall, modelAvailableCall } from "@/components/networking";
import useAuthorized from "../useAuthorized";
import { EndpointType, isModeCompatibleWithEndpoint } from "@/components/chat_ui/mode_endpoint_mapping";
export interface ProxyModel {
id: string;
@ -87,6 +88,7 @@ export const useModelsInfo = (
const AUTO_ROUTER_MODEL_PREFIX = "auto_router/";
const AUTO_ROUTER_LOOKUP_PAGE_SIZE = 1000;
const NO_AUTO_ROUTERS: ReadonlySet<string> = new Set<string>();
const NO_DEPLOYMENTS: AutoRouterDeployment[] = [];
export interface AutoRouterCandidateDeployment {
model_name?: string | null;
@ -96,6 +98,7 @@ export interface AutoRouterCandidateDeployment {
export interface AutoRouterDeployment extends AutoRouterCandidateDeployment {
litellm_params?: {
model?: string | null;
base_model?: string | null;
complexity_router_config?: unknown;
complexity_router_default_model?: string | null;
auto_router_config?: unknown;
@ -111,6 +114,7 @@ export interface AutoRouterDeployment extends AutoRouterCandidateDeployment {
/** False for config.yaml-defined deployments, which the update and delete routes refuse. */
db_model?: boolean | null;
base_model?: string | null;
mode?: string | null;
created_at?: string | null;
updated_at?: string | null;
team_id?: string | null;
@ -142,6 +146,22 @@ export const selectPlainModelGroups = (deployments: AutoRouterCandidateDeploymen
);
};
export const selectPlainChatModelDeployments = (deployments: AutoRouterDeployment[]): AutoRouterDeployment[] => {
const plainGroups = selectPlainModelGroups(deployments);
return deployments.filter(
(deployment) =>
plainGroups.has(deployment.model_name ?? "") &&
isModeCompatibleWithEndpoint(deployment.model_info?.mode, EndpointType.CHAT),
);
};
export const selectPlainChatModelGroups = (deployments: AutoRouterDeployment[]): ReadonlySet<string> =>
new Set(
selectPlainChatModelDeployments(deployments)
.map((deployment) => deployment.model_name)
.filter((name): name is string => Boolean(name)),
);
export const fetchAllModelDeployments = async (
accessToken: string,
userId: string,
@ -180,37 +200,32 @@ export const autoRouterListKey = (userId: string | null, userRole: string | null
},
});
export const useAutoRouterModelGroups = (): ReadonlySet<string> => {
const useDeployments = <TSelected>(
select: (deployments: AutoRouterDeployment[]) => TSelected,
): UseQueryResult<TSelected, Error> => {
const { accessToken, userId, userRole } = useAuthorized();
const { data } = useQuery<AutoRouterDeployment[], Error, ReadonlySet<string>>({
return useQuery<AutoRouterDeployment[], Error, TSelected>({
queryKey: autoRouterListKey(userId, userRole),
queryFn: async () => await fetchAllModelDeployments(accessToken!, userId!, userRole!),
enabled: Boolean(accessToken && userId && userRole),
select: selectAutoRouterModelGroups,
select,
});
return data ?? NO_AUTO_ROUTERS;
};
export const usePlainModelGroups = (): ReadonlySet<string> => {
const { accessToken, userId, userRole } = useAuthorized();
const { data } = useQuery<AutoRouterDeployment[], Error, ReadonlySet<string>>({
queryKey: autoRouterListKey(userId, userRole),
queryFn: async () => await fetchAllModelDeployments(accessToken!, userId!, userRole!),
enabled: Boolean(accessToken && userId && userRole),
select: selectPlainModelGroups,
});
return data ?? NO_AUTO_ROUTERS;
};
export const useAutoRouterModelGroups = (): ReadonlySet<string> =>
useDeployments(selectAutoRouterModelGroups).data ?? NO_AUTO_ROUTERS;
export const useAutoRouters = (): UseQueryResult<AutoRouterDeployment[], Error> => {
const { accessToken, userId, userRole } = useAuthorized();
return useQuery<AutoRouterDeployment[], Error, AutoRouterDeployment[]>({
queryKey: autoRouterListKey(userId, userRole),
queryFn: async () => await fetchAllModelDeployments(accessToken!, userId!, userRole!),
enabled: Boolean(accessToken && userId && userRole),
select: selectAutoRouterDeployments,
});
};
export const usePlainModelGroups = (): ReadonlySet<string> =>
useDeployments(selectPlainModelGroups).data ?? NO_AUTO_ROUTERS;
export const usePlainChatModelGroups = (): ReadonlySet<string> =>
useDeployments(selectPlainChatModelGroups).data ?? NO_AUTO_ROUTERS;
export const usePlainChatModelDeployments = (): AutoRouterDeployment[] =>
useDeployments(selectPlainChatModelDeployments).data ?? NO_DEPLOYMENTS;
export const useAutoRouters = (): UseQueryResult<AutoRouterDeployment[], Error> =>
useDeployments(selectAutoRouterDeployments);
export const useInvalidateAutoRouters = (): (() => Promise<void>) => {
const queryClient = useQueryClient();

View file

@ -1,7 +1,9 @@
import { ModelGroup } from "@/components/llm_calls/fetch_models";
import { EndpointType, getEndpointType, ModelMode } from "@/components/chat_ui/mode_endpoint_mapping";
const KNOWN_MODEL_MODES = new Set<string>(Object.values(ModelMode));
import {
EndpointType,
getEndpointType,
isModeCompatibleWithEndpoint,
} from "@/components/chat_ui/mode_endpoint_mapping";
export const determineEndpointType = (selectedModel: string, modelInfo: ModelGroup[]): EndpointType => {
const selectedModelInfo = modelInfo.find((option) => option.model_group === selectedModel);
@ -13,31 +15,8 @@ export const determineEndpointType = (selectedModel: string, modelInfo: ModelGro
return EndpointType.CHAT;
};
export const isModelCompatibleWithEndpoint = (model: ModelGroup, endpointType: EndpointType): boolean => {
if (!model.mode) {
return true;
}
if (!KNOWN_MODEL_MODES.has(model.mode)) {
return false;
}
const optionEndpoint = getEndpointType(model.mode);
if (
endpointType === EndpointType.RESPONSES ||
endpointType === EndpointType.ANTHROPIC_MESSAGES ||
endpointType === EndpointType.INTERACTIONS
) {
return optionEndpoint === endpointType || optionEndpoint === EndpointType.CHAT;
}
if (endpointType === EndpointType.IMAGE_EDITS) {
return optionEndpoint === endpointType || optionEndpoint === EndpointType.IMAGE;
}
return optionEndpoint === endpointType;
};
export const isModelCompatibleWithEndpoint = (model: ModelGroup, endpointType: EndpointType): boolean =>
isModeCompatibleWithEndpoint(model.mode, endpointType);
export const filterModelsForEndpoint = (models: ModelGroup[], endpointType: EndpointType): ModelGroup[] =>
models.filter((model) => isModelCompatibleWithEndpoint(model, endpointType));

View file

@ -57,3 +57,20 @@ export const getEndpointType = (mode: string): EndpointType => {
// else default to chat
return EndpointType.CHAT;
};
export const isModeCompatibleWithEndpoint = (mode: string | null | undefined, endpointType: EndpointType): boolean => {
if (!mode) return true;
if (!Object.values(ModelMode).includes(mode as ModelMode)) return false;
const optionEndpoint = getEndpointType(mode);
if (
endpointType === EndpointType.RESPONSES ||
endpointType === EndpointType.ANTHROPIC_MESSAGES ||
endpointType === EndpointType.INTERACTIONS
) {
return optionEndpoint === endpointType || optionEndpoint === EndpointType.CHAT;
}
if (endpointType === EndpointType.IMAGE_EDITS) {
return optionEndpoint === endpointType || optionEndpoint === EndpointType.IMAGE;
}
return optionEndpoint === endpointType;
};

View file

@ -13,6 +13,7 @@ import {
buildModelAvailability,
deploymentRefsFromModelInfo,
normalizeModelName,
resolveAvailableModels,
} from "./autorouter_presets";
import { DEFAULT_MATCH_THRESHOLD } from "@/components/add_model/SemanticKeywordMatching";
import { DEFAULT_ESCALATION_KEYWORDS } from "@/components/add_model/EscalationKeywords";
@ -380,6 +381,18 @@ describe("autorouter_presets", () => {
expect(availability.underlyingIndex.size).toBe(0);
});
it("returns every configured group serving the same underlying model", () => {
const availability = buildModelAvailability(
["z-group", "a-group"],
[
{ modelGroup: "z-group", underlyingModels: ["anthropic/claude-sonnet-5"] },
{ modelGroup: "a-group", underlyingModels: ["bedrock/us.anthropic.claude-sonnet-5-v1:0"] },
],
);
expect(resolveAvailableModels("anthropic/claude-sonnet-5", availability)).toEqual(["a-group", "z-group"]);
});
it("breaks ties between groups serving the same model deterministically, alphabetically", () => {
const availability = buildModelAvailability(
["z-group", "a-group"],

View file

@ -159,16 +159,19 @@ export const deploymentRefsFromModelInfo = (
return row.model_name && underlyingModels.length > 0 ? [{ modelGroup: row.model_name, underlyingModels }] : [];
});
export const resolveAvailableModel = (requiredModel: string, availability: ModelAvailability): string | undefined => {
export const resolveAvailableModels = (requiredModel: string, availability: ModelAvailability): readonly string[] => {
const { modelGroups, underlyingIndex } = availability;
if (modelGroups.has(requiredModel)) return requiredModel;
if (modelGroups.has(requiredModel)) return [requiredModel];
const normalized = normalizeModelName(requiredModel);
const groupMatch = Array.from(modelGroups).find((available) => normalizeModelName(available) === normalized);
if (groupMatch !== undefined) return groupMatch;
const groupMatches = Array.from(modelGroups).filter((available) => normalizeModelName(available) === normalized);
if (groupMatches.length > 0) return groupMatches;
const key = normalizeUnderlyingModel(requiredModel);
return key === null ? undefined : underlyingIndex.get(key)?.[0];
return key === null ? [] : underlyingIndex.get(key) ?? [];
};
export const resolveAvailableModel = (requiredModel: string, availability: ModelAvailability): string | undefined =>
resolveAvailableModels(requiredModel, availability)[0];
export const getMissingModels = (
config: Parameters<typeof getRequiredModels>[0],
availability: ModelAvailability,