mirror of
https://github.com/getpaseo/paseo.git
synced 2026-07-29 12:01:31 +00:00
Streamline model selector provider data
This commit is contained in:
@@ -32,6 +32,10 @@ import {
|
||||
} from "lucide-react-native";
|
||||
import { getProviderIcon } from "@/components/provider-icons";
|
||||
import { CombinedModelSelector } from "@/components/combined-model-selector";
|
||||
import {
|
||||
buildModelSelectorProviders,
|
||||
type ModelSelectorProvider,
|
||||
} from "@/components/combined-model-selector.utils";
|
||||
import { useSessionStore } from "@/stores/session-store";
|
||||
import { useProvidersSnapshot } from "@/hooks/use-providers-snapshot";
|
||||
import { resolveProviderDefinition } from "@/utils/provider-definitions";
|
||||
@@ -93,8 +97,7 @@ interface ControlledAgentStatusBarProps {
|
||||
disabled?: boolean;
|
||||
isModelLoading?: boolean;
|
||||
providerDefinitions: AgentProviderDefinition[];
|
||||
allProviderModels?: Map<string, AgentModelDefinition[]>;
|
||||
canSelectModelProvider?: (providerId: string) => boolean;
|
||||
modelSelectorProviders?: ModelSelectorProvider[];
|
||||
favoriteKeys?: Set<string>;
|
||||
onToggleFavoriteModel?: (provider: string, modelId: string) => void;
|
||||
features?: AgentFeature[];
|
||||
@@ -114,7 +117,7 @@ export interface DraftAgentStatusBarProps {
|
||||
selectedModel: string;
|
||||
onSelectModel: (modelId: string) => void;
|
||||
isModelLoading: boolean;
|
||||
allProviderModels: Map<string, AgentModelDefinition[]>;
|
||||
modelSelectorProviders: ModelSelectorProvider[];
|
||||
isAllModelsLoading: boolean;
|
||||
onSelectProviderAndModel: (provider: AgentProvider, modelId: string) => void;
|
||||
thinkingOptions: NonNullable<AgentModelDefinition["thinkingOptions"]>;
|
||||
@@ -184,10 +187,6 @@ const MODE_ICONS = {
|
||||
ShieldQuestionMark,
|
||||
} as const;
|
||||
|
||||
function alwaysTrue() {
|
||||
return true;
|
||||
}
|
||||
|
||||
function resolveDisplayModel(
|
||||
isModelLoading: boolean,
|
||||
modelOptions: StatusOption[] | undefined,
|
||||
@@ -235,23 +234,26 @@ function toComboboxOptions(options: StatusOption[] | undefined): ComboboxOption[
|
||||
return (options ?? []).map((o) => ({ id: o.id, label: o.label }));
|
||||
}
|
||||
|
||||
function buildFallbackAllProviderModels(
|
||||
function buildFallbackModelSelectorProviders(
|
||||
provider: string,
|
||||
modelOptions: StatusOption[] | undefined,
|
||||
): Map<string, AgentModelDefinition[]> {
|
||||
const map = new Map<string, AgentModelDefinition[]>();
|
||||
): ModelSelectorProvider[] {
|
||||
if (!modelOptions || modelOptions.length === 0) {
|
||||
return map;
|
||||
return [];
|
||||
}
|
||||
map.set(
|
||||
provider,
|
||||
modelOptions.map((option) => ({
|
||||
provider: provider,
|
||||
id: option.id,
|
||||
label: option.label,
|
||||
})),
|
||||
);
|
||||
return map;
|
||||
return [
|
||||
{
|
||||
id: provider,
|
||||
label: provider,
|
||||
rows: modelOptions.map((option) => ({
|
||||
favoriteKey: buildFavoriteModelKey({ provider, modelId: option.id }),
|
||||
provider,
|
||||
providerLabel: provider,
|
||||
modelId: option.id,
|
||||
modelLabel: option.label,
|
||||
})),
|
||||
},
|
||||
];
|
||||
}
|
||||
|
||||
function makeBadgePressableStyle(
|
||||
@@ -459,8 +461,7 @@ function ControlledStatusBar({
|
||||
disabled = false,
|
||||
isModelLoading = false,
|
||||
providerDefinitions,
|
||||
allProviderModels,
|
||||
canSelectModelProvider,
|
||||
modelSelectorProviders,
|
||||
favoriteKeys = new Set<string>(),
|
||||
onToggleFavoriteModel,
|
||||
features,
|
||||
@@ -521,13 +522,11 @@ function ControlledStatusBar({
|
||||
() => toComboboxOptions(modeOptions),
|
||||
[modeOptions],
|
||||
);
|
||||
const fallbackAllProviderModels = useMemo(
|
||||
() => buildFallbackAllProviderModels(provider, modelOptions),
|
||||
const fallbackModelSelectorProviders = useMemo(
|
||||
() => buildFallbackModelSelectorProviders(provider, modelOptions),
|
||||
[modelOptions, provider],
|
||||
);
|
||||
const effectiveProviderDefinitions = providerDefinitions;
|
||||
const effectiveAllProviderModels = allProviderModels ?? fallbackAllProviderModels;
|
||||
const canSelectProviderInModelMenu = canSelectModelProvider ?? alwaysTrue;
|
||||
const effectiveModelSelectorProviders = modelSelectorProviders ?? fallbackModelSelectorProviders;
|
||||
const comboboxThinkingOptions = useMemo<ComboboxOption[]>(
|
||||
() => toComboboxOptions(thinkingOptions),
|
||||
[thinkingOptions],
|
||||
@@ -701,10 +700,8 @@ function ControlledStatusBar({
|
||||
canSelectMode={canSelectMode}
|
||||
canSelectModel={canSelectModel}
|
||||
canSelectThinking={canSelectThinking}
|
||||
canSelectProviderInModelMenu={canSelectProviderInModelMenu}
|
||||
modelSelectorProviders={effectiveModelSelectorProviders}
|
||||
modelDisabled={modelDisabled}
|
||||
effectiveProviderDefinitions={effectiveProviderDefinitions}
|
||||
effectiveAllProviderModels={effectiveAllProviderModels}
|
||||
comboboxProviderOptions={comboboxProviderOptions}
|
||||
comboboxModeOptions={comboboxModeOptions}
|
||||
comboboxThinkingOptions={comboboxThinkingOptions}
|
||||
@@ -751,10 +748,8 @@ function ControlledStatusBar({
|
||||
canSelectMode={canSelectMode}
|
||||
canSelectModel={canSelectModel}
|
||||
canSelectThinking={canSelectThinking}
|
||||
canSelectProviderInModelMenu={canSelectProviderInModelMenu}
|
||||
modelSelectorProviders={effectiveModelSelectorProviders}
|
||||
modelDisabled={modelDisabled}
|
||||
effectiveProviderDefinitions={effectiveProviderDefinitions}
|
||||
effectiveAllProviderModels={effectiveAllProviderModels}
|
||||
comboboxModeOptions={comboboxModeOptions}
|
||||
comboboxThinkingOptions={comboboxThinkingOptions}
|
||||
ModeIconComponent={ModeIconComponent}
|
||||
@@ -799,10 +794,8 @@ interface DesktopStatusBarContentProps {
|
||||
canSelectMode: boolean;
|
||||
canSelectModel: boolean;
|
||||
canSelectThinking: boolean;
|
||||
canSelectProviderInModelMenu: (providerId: string) => boolean;
|
||||
modelSelectorProviders: ModelSelectorProvider[];
|
||||
modelDisabled: boolean;
|
||||
effectiveProviderDefinitions: AgentProviderDefinition[];
|
||||
effectiveAllProviderModels: Map<string, AgentModelDefinition[]>;
|
||||
comboboxProviderOptions: ComboboxOption[];
|
||||
comboboxModeOptions: ComboboxOption[];
|
||||
comboboxThinkingOptions: ComboboxOption[];
|
||||
@@ -868,10 +861,8 @@ function DesktopStatusBarContent(props: DesktopStatusBarContentProps) {
|
||||
canSelectMode,
|
||||
canSelectModel,
|
||||
canSelectThinking,
|
||||
canSelectProviderInModelMenu,
|
||||
modelSelectorProviders,
|
||||
modelDisabled,
|
||||
effectiveProviderDefinitions,
|
||||
effectiveAllProviderModels,
|
||||
comboboxProviderOptions,
|
||||
comboboxModeOptions,
|
||||
comboboxThinkingOptions,
|
||||
@@ -942,11 +933,9 @@ function DesktopStatusBarContent(props: DesktopStatusBarContentProps) {
|
||||
<TooltipTrigger asChild triggerRefProp="ref">
|
||||
<View>
|
||||
<CombinedModelSelector
|
||||
providerDefinitions={effectiveProviderDefinitions}
|
||||
allProviderModels={effectiveAllProviderModels}
|
||||
providers={modelSelectorProviders}
|
||||
selectedProvider={provider}
|
||||
selectedModel={selectedModelId ?? ""}
|
||||
canSelectProvider={canSelectProviderInModelMenu}
|
||||
onSelect={handleDesktopModelSelect}
|
||||
favoriteKeys={favoriteKeys}
|
||||
onToggleFavorite={onToggleFavoriteModel}
|
||||
@@ -1069,10 +1058,8 @@ interface SheetStatusBarContentProps {
|
||||
canSelectMode: boolean;
|
||||
canSelectModel: boolean;
|
||||
canSelectThinking: boolean;
|
||||
canSelectProviderInModelMenu: (providerId: string) => boolean;
|
||||
modelSelectorProviders: ModelSelectorProvider[];
|
||||
modelDisabled: boolean;
|
||||
effectiveProviderDefinitions: AgentProviderDefinition[];
|
||||
effectiveAllProviderModels: Map<string, AgentModelDefinition[]>;
|
||||
comboboxModeOptions: ComboboxOption[];
|
||||
comboboxThinkingOptions: ComboboxOption[];
|
||||
ModeIconComponent: (typeof MODE_ICONS)[keyof typeof MODE_ICONS] | null;
|
||||
@@ -1118,10 +1105,8 @@ function SheetStatusBarContent(props: SheetStatusBarContentProps) {
|
||||
canSelectMode,
|
||||
canSelectModel,
|
||||
canSelectThinking,
|
||||
canSelectProviderInModelMenu,
|
||||
modelSelectorProviders,
|
||||
modelDisabled,
|
||||
effectiveProviderDefinitions,
|
||||
effectiveAllProviderModels,
|
||||
comboboxModeOptions,
|
||||
comboboxThinkingOptions,
|
||||
ModeIconComponent,
|
||||
@@ -1214,11 +1199,9 @@ function SheetStatusBarContent(props: SheetStatusBarContentProps) {
|
||||
<>
|
||||
{canSelectModel ? (
|
||||
<CombinedModelSelector
|
||||
providerDefinitions={effectiveProviderDefinitions}
|
||||
allProviderModels={effectiveAllProviderModels}
|
||||
providers={modelSelectorProviders}
|
||||
selectedProvider={provider}
|
||||
selectedModel={selectedModelId ?? ""}
|
||||
canSelectProvider={canSelectProviderInModelMenu}
|
||||
onSelect={handleSheetModelSelect}
|
||||
favoriteKeys={favoriteKeys}
|
||||
onToggleFavorite={onToggleFavoriteModel}
|
||||
@@ -1678,6 +1661,10 @@ export const AgentStatusBar = memo(function AgentStatusBar({
|
||||
() => buildAgentProviderModels(agent?.provider, models),
|
||||
[agent?.provider, models],
|
||||
);
|
||||
const agentModelSelectorProviders = useMemo(
|
||||
() => buildModelSelectorProviders(agentProviderDefinitions, agentProviderModels),
|
||||
[agentProviderDefinitions, agentProviderModels],
|
||||
);
|
||||
|
||||
const displayMode = resolveAgentDisplayMode(availableModes, agent?.currentModeId);
|
||||
|
||||
@@ -1841,7 +1828,7 @@ export const AgentStatusBar = memo(function AgentStatusBar({
|
||||
modeOptions={fallbackModeOptions}
|
||||
selectedModeId={agent.currentModeId ?? undefined}
|
||||
providerDefinitions={agentProviderDefinitions}
|
||||
allProviderModels={agentProviderModels}
|
||||
modelSelectorProviders={agentModelSelectorProviders}
|
||||
onSelectMode={handleSelectMode}
|
||||
modelOptions={modelOptions}
|
||||
selectedModelId={modelSelection.activeModelId ?? undefined}
|
||||
@@ -1872,7 +1859,7 @@ export function DraftAgentStatusBar({
|
||||
selectedModel,
|
||||
onSelectModel,
|
||||
isModelLoading: _isModelLoading,
|
||||
allProviderModels,
|
||||
modelSelectorProviders,
|
||||
isAllModelsLoading,
|
||||
onSelectProviderAndModel,
|
||||
thinkingOptions,
|
||||
@@ -1937,8 +1924,7 @@ export function DraftAgentStatusBar({
|
||||
return (
|
||||
<View style={styles.container}>
|
||||
<CombinedModelSelector
|
||||
providerDefinitions={providerDefinitions}
|
||||
allProviderModels={allProviderModels}
|
||||
providers={modelSelectorProviders}
|
||||
selectedProvider={selectedProvider ?? ""}
|
||||
selectedModel={selectedModel}
|
||||
onSelect={onSelectProviderAndModel}
|
||||
@@ -1973,7 +1959,7 @@ export function DraftAgentStatusBar({
|
||||
<ControlledStatusBar
|
||||
provider={selectedProvider ?? ""}
|
||||
providerDefinitions={providerDefinitions}
|
||||
allProviderModels={allProviderModels}
|
||||
modelSelectorProviders={modelSelectorProviders}
|
||||
modeOptions={hasSelectedProvider ? mappedModeOptions : undefined}
|
||||
selectedModeId={effectiveSelectedMode}
|
||||
onSelectMode={onSelectMode}
|
||||
|
||||
@@ -1,79 +1,136 @@
|
||||
import { describe, expect, it } from "vitest";
|
||||
import type { AgentModelDefinition } from "@server/server/agent/agent-sdk-types";
|
||||
import type {
|
||||
AgentModelDefinition,
|
||||
ProviderSnapshotEntry,
|
||||
} from "@server/server/agent/agent-sdk-types";
|
||||
import type { AgentProviderDefinition } from "@server/server/agent/provider-manifest";
|
||||
import {
|
||||
buildModelRows,
|
||||
buildProviderGroups,
|
||||
buildModelSelectorProviders,
|
||||
buildSelectableModelSelectorProviders,
|
||||
buildSelectedTriggerLabel,
|
||||
filterAndRankModelRows,
|
||||
matchesSearch,
|
||||
resolveProviderLabel,
|
||||
} from "./combined-model-selector.utils";
|
||||
|
||||
describe("combined model selector helpers", () => {
|
||||
const providerDefinitions = [
|
||||
{
|
||||
id: "claude",
|
||||
label: "Claude",
|
||||
description: "Claude provider",
|
||||
defaultModeId: "default",
|
||||
modes: [],
|
||||
},
|
||||
{
|
||||
id: "codex",
|
||||
label: "Codex",
|
||||
description: "Codex provider",
|
||||
defaultModeId: "auto",
|
||||
modes: [],
|
||||
},
|
||||
{
|
||||
id: "deepseek-tui",
|
||||
label: "DeepSeek TUI",
|
||||
description: "DeepSeek TUI provider",
|
||||
defaultModeId: "default",
|
||||
modes: [],
|
||||
},
|
||||
];
|
||||
describe("combined model selector data", () => {
|
||||
const codexModel: AgentModelDefinition = {
|
||||
provider: "codex",
|
||||
id: "gpt-5.4",
|
||||
label: "GPT-5.4",
|
||||
};
|
||||
|
||||
const claudeModels: AgentModelDefinition[] = [
|
||||
{
|
||||
provider: "claude",
|
||||
id: "sonnet-4.6",
|
||||
label: "Sonnet 4.6",
|
||||
},
|
||||
];
|
||||
function snapshotEntry(
|
||||
overrides: Partial<ProviderSnapshotEntry> & Pick<ProviderSnapshotEntry, "provider">,
|
||||
): ProviderSnapshotEntry {
|
||||
return {
|
||||
...overrides,
|
||||
provider: overrides.provider,
|
||||
status: overrides.status ?? "ready",
|
||||
enabled: overrides.enabled ?? true,
|
||||
label: overrides.label ?? overrides.provider,
|
||||
description: overrides.description ?? `${overrides.provider} provider`,
|
||||
defaultModeId: overrides.defaultModeId ?? "default",
|
||||
modes: overrides.modes ?? [],
|
||||
models: overrides.models ?? [codexModel],
|
||||
};
|
||||
}
|
||||
|
||||
const codexModels: AgentModelDefinition[] = [
|
||||
{
|
||||
provider: "codex",
|
||||
id: "gpt-5.4",
|
||||
label: "GPT-5.4",
|
||||
},
|
||||
];
|
||||
|
||||
it("keeps enough data to search by model and provider name", async () => {
|
||||
const rows = buildModelRows(
|
||||
providerDefinitions,
|
||||
new Map([
|
||||
["claude", claudeModels],
|
||||
["codex", codexModels],
|
||||
it("builds selector providers from ready enabled snapshot entries", () => {
|
||||
expect(
|
||||
buildSelectableModelSelectorProviders([
|
||||
snapshotEntry({
|
||||
provider: "codex",
|
||||
label: "Codex",
|
||||
models: [codexModel],
|
||||
}),
|
||||
]),
|
||||
);
|
||||
|
||||
expect(rows).toEqual([
|
||||
expect.objectContaining({
|
||||
providerLabel: "Claude",
|
||||
modelLabel: "Sonnet 4.6",
|
||||
modelId: "sonnet-4.6",
|
||||
}),
|
||||
expect.objectContaining({
|
||||
providerLabel: "Codex",
|
||||
modelLabel: "GPT-5.4",
|
||||
modelId: "gpt-5.4",
|
||||
}),
|
||||
).toEqual([
|
||||
{
|
||||
id: "codex",
|
||||
label: "Codex",
|
||||
rows: [
|
||||
{
|
||||
favoriteKey: "codex:gpt-5.4",
|
||||
provider: "codex",
|
||||
providerLabel: "Codex",
|
||||
modelId: "gpt-5.4",
|
||||
modelLabel: "GPT-5.4",
|
||||
description: undefined,
|
||||
isDefault: undefined,
|
||||
},
|
||||
],
|
||||
},
|
||||
]);
|
||||
});
|
||||
|
||||
expect(matchesSearch(rows[0], "claude")).toBe(true);
|
||||
expect(matchesSearch(rows[1], "gpt-5.4")).toBe(true);
|
||||
it("keeps ready enabled providers with no models as model-less providers", () => {
|
||||
expect(
|
||||
buildSelectableModelSelectorProviders([
|
||||
snapshotEntry({
|
||||
provider: "deepseek-tui",
|
||||
label: "DeepSeek TUI",
|
||||
models: [],
|
||||
}),
|
||||
]),
|
||||
).toEqual([
|
||||
{
|
||||
id: "deepseek-tui",
|
||||
label: "DeepSeek TUI",
|
||||
rows: [],
|
||||
},
|
||||
]);
|
||||
});
|
||||
|
||||
it("excludes disabled providers from selector data", () => {
|
||||
expect(
|
||||
buildSelectableModelSelectorProviders([
|
||||
snapshotEntry({
|
||||
provider: "deepseek-tui",
|
||||
label: "DeepSeek TUI",
|
||||
enabled: false,
|
||||
models: [],
|
||||
}),
|
||||
]),
|
||||
).toEqual([]);
|
||||
});
|
||||
|
||||
it("excludes providers that are not ready", () => {
|
||||
expect(
|
||||
buildSelectableModelSelectorProviders([
|
||||
snapshotEntry({ provider: "loading-provider", status: "loading", models: [] }),
|
||||
snapshotEntry({ provider: "error-provider", status: "error", models: [] }),
|
||||
snapshotEntry({ provider: "unavailable-provider", status: "unavailable", models: [] }),
|
||||
]),
|
||||
).toEqual([]);
|
||||
});
|
||||
|
||||
it("builds selector providers from an already-curated provider list", () => {
|
||||
const providerDefinitions: AgentProviderDefinition[] = [
|
||||
{
|
||||
id: "codex",
|
||||
label: "Codex",
|
||||
description: "Codex provider",
|
||||
defaultModeId: "auto",
|
||||
modes: [],
|
||||
},
|
||||
];
|
||||
|
||||
expect(
|
||||
buildModelSelectorProviders(providerDefinitions, new Map([["codex", [codexModel]]])),
|
||||
).toEqual([
|
||||
{
|
||||
id: "codex",
|
||||
label: "Codex",
|
||||
rows: [
|
||||
expect.objectContaining({
|
||||
provider: "codex",
|
||||
providerLabel: "Codex",
|
||||
modelId: "gpt-5.4",
|
||||
modelLabel: "GPT-5.4",
|
||||
}),
|
||||
],
|
||||
},
|
||||
]);
|
||||
});
|
||||
|
||||
it("matches across label, provider, and description with multi-token fuzzy search", () => {
|
||||
@@ -120,55 +177,7 @@ describe("combined model selector helpers", () => {
|
||||
expect(filterAndRankModelRows(rows, "gpt54").map((row) => row.modelId)).toEqual(["gpt-5.4"]);
|
||||
});
|
||||
|
||||
it("includes providers that expose no models", () => {
|
||||
const rows = buildModelRows(
|
||||
providerDefinitions,
|
||||
new Map([
|
||||
["claude", claudeModels],
|
||||
["deepseek-tui", []],
|
||||
]),
|
||||
);
|
||||
|
||||
const groups = buildProviderGroups(
|
||||
providerDefinitions,
|
||||
new Map([
|
||||
["claude", claudeModels],
|
||||
["deepseek-tui", []],
|
||||
]),
|
||||
rows,
|
||||
"",
|
||||
);
|
||||
|
||||
expect(groups).toEqual([
|
||||
expect.objectContaining({
|
||||
providerId: "claude",
|
||||
hasNoModels: false,
|
||||
}),
|
||||
expect.objectContaining({
|
||||
providerId: "deepseek-tui",
|
||||
providerLabel: "DeepSeek TUI",
|
||||
rows: [],
|
||||
hasNoModels: true,
|
||||
}),
|
||||
]);
|
||||
});
|
||||
|
||||
it("matches model-less providers by provider name", () => {
|
||||
const groups = buildProviderGroups(
|
||||
providerDefinitions,
|
||||
new Map([
|
||||
["claude", claudeModels],
|
||||
["deepseek-tui", []],
|
||||
]),
|
||||
[],
|
||||
"deepseek",
|
||||
);
|
||||
|
||||
expect(groups.map((group) => group.providerId)).toEqual(["deepseek-tui"]);
|
||||
});
|
||||
|
||||
it("keeps the selected trigger label model-only", () => {
|
||||
expect(resolveProviderLabel(providerDefinitions, "codex")).toBe("Codex");
|
||||
expect(buildSelectedTriggerLabel("GPT-5.4")).toBe("GPT-5.4");
|
||||
});
|
||||
});
|
||||
|
||||
@@ -12,8 +12,7 @@ import { StyleSheet, useUnistyles } from "react-native-unistyles";
|
||||
import { useIsCompactFormFactor } from "@/constants/layout";
|
||||
import { isNative, isWeb as platformIsWeb } from "@/constants/platform";
|
||||
import { ChevronDown, ChevronRight, Search, Star } from "lucide-react-native";
|
||||
import type { AgentModelDefinition, AgentProvider } from "@server/server/agent/agent-sdk-types";
|
||||
import type { AgentProviderDefinition } from "@server/server/agent/provider-manifest";
|
||||
import type { AgentProvider } from "@server/server/agent/agent-sdk-types";
|
||||
import type { SheetHeader } from "@/components/adaptive-modal-sheet";
|
||||
const IS_WEB = platformIsWeb;
|
||||
|
||||
@@ -46,12 +45,9 @@ function drillDownRowStyle({
|
||||
}
|
||||
import { getProviderIcon } from "@/components/provider-icons";
|
||||
import {
|
||||
buildModelRows,
|
||||
buildProviderGroups,
|
||||
buildSelectedTriggerLabel,
|
||||
filterAndRankModelRows,
|
||||
resolveProviderLabel,
|
||||
type SelectorProviderGroup,
|
||||
type ModelSelectorProvider,
|
||||
type SelectorModelRow,
|
||||
} from "./combined-model-selector.utils";
|
||||
|
||||
@@ -63,13 +59,11 @@ type SelectorView =
|
||||
| { kind: "provider"; providerId: string; providerLabel: string };
|
||||
|
||||
interface CombinedModelSelectorProps {
|
||||
providerDefinitions: AgentProviderDefinition[];
|
||||
allProviderModels: Map<string, AgentModelDefinition[]>;
|
||||
providers: ModelSelectorProvider[];
|
||||
selectedProvider: string;
|
||||
selectedModel: string;
|
||||
onSelect: (provider: AgentProvider, modelId: string) => void;
|
||||
isLoading: boolean;
|
||||
canSelectProvider?: (provider: string) => boolean;
|
||||
favoriteKeys?: Set<string>;
|
||||
onToggleFavorite?: (provider: string, modelId: string) => void;
|
||||
renderTrigger?: (input: {
|
||||
@@ -85,23 +79,21 @@ interface CombinedModelSelectorProps {
|
||||
|
||||
interface SelectorContentProps {
|
||||
view: SelectorView;
|
||||
providerDefinitions: AgentProviderDefinition[];
|
||||
allProviderModels: Map<string, AgentModelDefinition[]>;
|
||||
providers: ModelSelectorProvider[];
|
||||
selectedProvider: string;
|
||||
selectedModel: string;
|
||||
searchQuery: string;
|
||||
favoriteKeys: Set<string>;
|
||||
onSelect: (provider: string, modelId: string) => void;
|
||||
canSelectProvider: (provider: string) => boolean;
|
||||
onToggleFavorite?: (provider: string, modelId: string) => void;
|
||||
onDrillDown: (providerId: string, providerLabel: string) => void;
|
||||
}
|
||||
|
||||
function resolveDefaultModelLabel(models: AgentModelDefinition[] | undefined): string {
|
||||
if (!models || models.length === 0) {
|
||||
function resolveDefaultModelLabel(rows: SelectorModelRow[] | undefined): string {
|
||||
if (!rows || rows.length === 0) {
|
||||
return "Select model";
|
||||
}
|
||||
return (models.find((model) => model.isDefault) ?? models[0])?.label ?? "Select model";
|
||||
return (rows.find((row) => row.isDefault) ?? rows[0])?.modelLabel ?? "Select model";
|
||||
}
|
||||
|
||||
function normalizeSearchQuery(value: string): string {
|
||||
@@ -128,7 +120,6 @@ function ModelRow({
|
||||
row,
|
||||
isSelected,
|
||||
isFavorite,
|
||||
disabled = false,
|
||||
elevated = false,
|
||||
onPress,
|
||||
onToggleFavorite,
|
||||
@@ -136,7 +127,6 @@ function ModelRow({
|
||||
row: SelectorModelRow;
|
||||
isSelected: boolean;
|
||||
isFavorite: boolean;
|
||||
disabled?: boolean;
|
||||
elevated?: boolean;
|
||||
onPress: () => void;
|
||||
onToggleFavorite?: (provider: string, modelId: string) => void;
|
||||
@@ -158,7 +148,7 @@ function ModelRow({
|
||||
);
|
||||
const trailingSlot = useMemo(
|
||||
() =>
|
||||
onToggleFavorite && !disabled ? (
|
||||
onToggleFavorite ? (
|
||||
<Pressable
|
||||
onPress={handleToggleFavorite}
|
||||
hitSlop={8}
|
||||
@@ -184,7 +174,6 @@ function ModelRow({
|
||||
) : null,
|
||||
[
|
||||
onToggleFavorite,
|
||||
disabled,
|
||||
handleToggleFavorite,
|
||||
isFavorite,
|
||||
row.provider,
|
||||
@@ -202,7 +191,6 @@ function ModelRow({
|
||||
label={row.modelLabel}
|
||||
description={showDescription ? row.description : undefined}
|
||||
selected={isSelected}
|
||||
disabled={disabled}
|
||||
elevated={elevated}
|
||||
onPress={onPress}
|
||||
leadingSlot={leadingSlot}
|
||||
@@ -215,7 +203,6 @@ interface SelectableModelRowProps {
|
||||
row: SelectorModelRow;
|
||||
isSelected: boolean;
|
||||
isFavorite: boolean;
|
||||
disabled?: boolean;
|
||||
elevated?: boolean;
|
||||
onSelect: (provider: string, modelId: string) => void;
|
||||
onToggleFavorite?: (provider: string, modelId: string) => void;
|
||||
@@ -225,7 +212,6 @@ function SelectableModelRow({
|
||||
row,
|
||||
isSelected,
|
||||
isFavorite,
|
||||
disabled,
|
||||
elevated,
|
||||
onSelect,
|
||||
onToggleFavorite,
|
||||
@@ -238,7 +224,6 @@ function SelectableModelRow({
|
||||
row={row}
|
||||
isSelected={isSelected}
|
||||
isFavorite={isFavorite}
|
||||
disabled={disabled}
|
||||
elevated={elevated}
|
||||
onPress={handlePress}
|
||||
onToggleFavorite={onToggleFavorite}
|
||||
@@ -252,7 +237,6 @@ function FavoritesSection({
|
||||
selectedModel,
|
||||
favoriteKeys,
|
||||
onSelect,
|
||||
canSelectProvider,
|
||||
onToggleFavorite,
|
||||
}: {
|
||||
favoriteRows: SelectorModelRow[];
|
||||
@@ -260,11 +244,8 @@ function FavoritesSection({
|
||||
selectedModel: string;
|
||||
favoriteKeys: Set<string>;
|
||||
onSelect: (provider: string, modelId: string) => void;
|
||||
canSelectProvider: (provider: string) => boolean;
|
||||
onToggleFavorite?: (provider: string, modelId: string) => void;
|
||||
}) {
|
||||
const { theme: _theme } = useUnistyles();
|
||||
|
||||
if (favoriteRows.length === 0) {
|
||||
return null;
|
||||
}
|
||||
@@ -280,7 +261,6 @@ function FavoritesSection({
|
||||
row={row}
|
||||
isSelected={row.provider === selectedProvider && row.modelId === selectedModel}
|
||||
isFavorite={favoriteKeys.has(row.favoriteKey)}
|
||||
disabled={!canSelectProvider(row.provider)}
|
||||
elevated
|
||||
onSelect={onSelect}
|
||||
onToggleFavorite={onToggleFavorite}
|
||||
@@ -295,7 +275,6 @@ interface GroupProviderButtonProps {
|
||||
providerLabel: string;
|
||||
rowCount: number;
|
||||
hasNoModels: boolean;
|
||||
disabled?: boolean;
|
||||
onDrillDown: (providerId: string, providerLabel: string) => void;
|
||||
onSelectDefault: (providerId: string) => void;
|
||||
}
|
||||
@@ -305,7 +284,6 @@ function GroupProviderButton({
|
||||
providerLabel,
|
||||
rowCount,
|
||||
hasNoModels,
|
||||
disabled,
|
||||
onDrillDown,
|
||||
onSelectDefault,
|
||||
}: GroupProviderButtonProps) {
|
||||
@@ -319,7 +297,7 @@ function GroupProviderButton({
|
||||
onDrillDown(providerId, providerLabel);
|
||||
}, [hasNoModels, onDrillDown, onSelectDefault, providerId, providerLabel]);
|
||||
return (
|
||||
<Pressable disabled={disabled} onPress={handlePress} style={drillDownRowStyle}>
|
||||
<Pressable onPress={handlePress} style={drillDownRowStyle}>
|
||||
<ProvIcon size={theme.iconSize.sm} color={theme.colors.foregroundMuted} />
|
||||
<Text style={styles.drillDownText}>{providerLabel}</Text>
|
||||
<View style={styles.drillDownTrailing}>
|
||||
@@ -335,28 +313,26 @@ function GroupProviderButton({
|
||||
}
|
||||
|
||||
function GroupedProviderRows({
|
||||
groupedRows,
|
||||
providers,
|
||||
onDrillDown,
|
||||
onSelectDefault,
|
||||
canSelectProvider,
|
||||
}: {
|
||||
groupedRows: SelectorProviderGroup[];
|
||||
providers: ModelSelectorProvider[];
|
||||
onDrillDown: (providerId: string, providerLabel: string) => void;
|
||||
onSelectDefault: (providerId: string) => void;
|
||||
canSelectProvider: (provider: string) => boolean;
|
||||
}) {
|
||||
return (
|
||||
<View>
|
||||
{groupedRows.map((group, index) => {
|
||||
{providers.map((provider, index) => {
|
||||
const hasNoModels = provider.rows.length === 0;
|
||||
return (
|
||||
<View key={group.providerId}>
|
||||
<View key={provider.id}>
|
||||
{index > 0 ? <View style={styles.separator} /> : null}
|
||||
<GroupProviderButton
|
||||
providerId={group.providerId}
|
||||
providerLabel={group.providerLabel}
|
||||
rowCount={group.rows.length}
|
||||
hasNoModels={group.hasNoModels}
|
||||
disabled={group.hasNoModels && !canSelectProvider(group.providerId)}
|
||||
providerId={provider.id}
|
||||
providerLabel={provider.label}
|
||||
rowCount={provider.rows.length}
|
||||
hasNoModels={hasNoModels}
|
||||
onDrillDown={onDrillDown}
|
||||
onSelectDefault={onSelectDefault}
|
||||
/>
|
||||
@@ -370,12 +346,10 @@ function GroupedProviderRows({
|
||||
function DefaultProviderRow({
|
||||
providerId,
|
||||
isSelected,
|
||||
disabled,
|
||||
onSelect,
|
||||
}: {
|
||||
providerId: string;
|
||||
isSelected: boolean;
|
||||
disabled?: boolean;
|
||||
onSelect: (provider: string, modelId: string) => void;
|
||||
}) {
|
||||
const { theme } = useUnistyles();
|
||||
@@ -392,7 +366,6 @@ function DefaultProviderRow({
|
||||
<ComboboxItem
|
||||
label="Default"
|
||||
selected={isSelected}
|
||||
disabled={disabled}
|
||||
onPress={handlePress}
|
||||
leadingSlot={leadingSlot}
|
||||
/>
|
||||
@@ -405,7 +378,6 @@ function ProviderModelRows({
|
||||
selectedModel,
|
||||
favoriteKeys,
|
||||
onSelect,
|
||||
canSelectProvider,
|
||||
onToggleFavorite,
|
||||
normalizedQuery,
|
||||
}: {
|
||||
@@ -414,7 +386,6 @@ function ProviderModelRows({
|
||||
selectedModel: string;
|
||||
favoriteKeys: Set<string>;
|
||||
onSelect: (provider: string, modelId: string) => void;
|
||||
canSelectProvider: (provider: string) => boolean;
|
||||
onToggleFavorite?: (provider: string, modelId: string) => void;
|
||||
normalizedQuery: string;
|
||||
}) {
|
||||
@@ -430,12 +401,11 @@ function ProviderModelRows({
|
||||
row={item}
|
||||
isSelected={item.provider === selectedProvider && item.modelId === selectedModel}
|
||||
isFavorite={favoriteKeys.has(item.favoriteKey)}
|
||||
disabled={!canSelectProvider(item.provider)}
|
||||
onSelect={onSelect}
|
||||
onToggleFavorite={onToggleFavorite}
|
||||
/>
|
||||
),
|
||||
[canSelectProvider, favoriteKeys, onSelect, onToggleFavorite, selectedModel, selectedProvider],
|
||||
[favoriteKeys, onSelect, onToggleFavorite, selectedModel, selectedProvider],
|
||||
);
|
||||
const keyExtractor = useCallback((row: SelectorModelRow) => row.favoriteKey, []);
|
||||
|
||||
@@ -464,45 +434,37 @@ function ProviderModelRows({
|
||||
|
||||
function SelectorContent({
|
||||
view,
|
||||
providerDefinitions,
|
||||
allProviderModels,
|
||||
providers,
|
||||
selectedProvider,
|
||||
selectedModel,
|
||||
searchQuery,
|
||||
favoriteKeys,
|
||||
onSelect,
|
||||
canSelectProvider,
|
||||
onToggleFavorite,
|
||||
onDrillDown,
|
||||
}: SelectorContentProps) {
|
||||
const { theme } = useUnistyles();
|
||||
const allRows = useMemo(
|
||||
() => buildModelRows(providerDefinitions, allProviderModels),
|
||||
[allProviderModels, providerDefinitions],
|
||||
);
|
||||
|
||||
const scopedRows = useMemo(() => {
|
||||
if (view.kind === "provider") {
|
||||
return allRows.filter((row) => row.provider === view.providerId);
|
||||
}
|
||||
return allRows;
|
||||
}, [allRows, view]);
|
||||
|
||||
const normalizedQuery = useMemo(() => normalizeSearchQuery(searchQuery), [searchQuery]);
|
||||
|
||||
const selectedViewProvider = useMemo(
|
||||
() =>
|
||||
view.kind === "provider"
|
||||
? providers.find((provider) => provider.id === view.providerId)
|
||||
: null,
|
||||
[providers, view],
|
||||
);
|
||||
const visibleRows = useMemo(
|
||||
() => filterAndRankModelRows(scopedRows, normalizedQuery),
|
||||
[normalizedQuery, scopedRows],
|
||||
() =>
|
||||
selectedViewProvider
|
||||
? filterAndRankModelRows(selectedViewProvider.rows, normalizedQuery)
|
||||
: [],
|
||||
[normalizedQuery, selectedViewProvider],
|
||||
);
|
||||
|
||||
const favoriteRows = useMemo(
|
||||
() => visibleRows.filter((row) => favoriteKeys.has(row.favoriteKey)),
|
||||
[favoriteKeys, visibleRows],
|
||||
);
|
||||
|
||||
const allGroupedRows = useMemo(
|
||||
() => buildProviderGroups(providerDefinitions, allProviderModels, visibleRows, normalizedQuery),
|
||||
[allProviderModels, normalizedQuery, providerDefinitions, visibleRows],
|
||||
() =>
|
||||
providers.flatMap((provider) =>
|
||||
provider.rows.filter((row) => favoriteKeys.has(row.favoriteKey)),
|
||||
),
|
||||
[favoriteKeys, providers],
|
||||
);
|
||||
const handleSelectDefaultProvider = useCallback(
|
||||
(providerId: string) => {
|
||||
@@ -510,7 +472,7 @@ function SelectorContent({
|
||||
},
|
||||
[onSelect],
|
||||
);
|
||||
const hasResults = favoriteRows.length > 0 || allGroupedRows.length > 0;
|
||||
const hasResults = favoriteRows.length > 0 || providers.length > 0;
|
||||
const emptyState = (
|
||||
<View style={styles.emptyState}>
|
||||
<Search size={theme.iconSize.md} color={theme.colors.foregroundMuted} />
|
||||
@@ -519,13 +481,15 @@ function SelectorContent({
|
||||
);
|
||||
|
||||
if (view.kind === "provider") {
|
||||
const providerModels = allProviderModels.get(view.providerId);
|
||||
if (providerModels && providerModels.length === 0 && !normalizedQuery) {
|
||||
if (!selectedViewProvider) {
|
||||
return emptyState;
|
||||
}
|
||||
|
||||
if (selectedViewProvider.rows.length === 0 && !normalizedQuery) {
|
||||
return (
|
||||
<DefaultProviderRow
|
||||
providerId={view.providerId}
|
||||
isSelected={view.providerId === selectedProvider && !selectedModel}
|
||||
disabled={!canSelectProvider(view.providerId)}
|
||||
onSelect={onSelect}
|
||||
/>
|
||||
);
|
||||
@@ -542,7 +506,6 @@ function SelectorContent({
|
||||
selectedModel={selectedModel}
|
||||
favoriteKeys={favoriteKeys}
|
||||
onSelect={onSelect}
|
||||
canSelectProvider={canSelectProvider}
|
||||
onToggleFavorite={onToggleFavorite}
|
||||
normalizedQuery={normalizedQuery}
|
||||
/>
|
||||
@@ -557,16 +520,14 @@ function SelectorContent({
|
||||
selectedModel={selectedModel}
|
||||
favoriteKeys={favoriteKeys}
|
||||
onSelect={onSelect}
|
||||
canSelectProvider={canSelectProvider}
|
||||
onToggleFavorite={onToggleFavorite}
|
||||
/>
|
||||
|
||||
{allGroupedRows.length > 0 ? (
|
||||
{providers.length > 0 ? (
|
||||
<GroupedProviderRows
|
||||
groupedRows={allGroupedRows}
|
||||
providers={providers}
|
||||
onDrillDown={onDrillDown}
|
||||
onSelectDefault={handleSelectDefaultProvider}
|
||||
canSelectProvider={canSelectProvider}
|
||||
/>
|
||||
) : null}
|
||||
|
||||
@@ -576,13 +537,11 @@ function SelectorContent({
|
||||
}
|
||||
|
||||
export function CombinedModelSelector({
|
||||
providerDefinitions,
|
||||
allProviderModels,
|
||||
providers,
|
||||
selectedProvider,
|
||||
selectedModel,
|
||||
onSelect,
|
||||
isLoading,
|
||||
canSelectProvider = () => true,
|
||||
favoriteKeys = new Set<string>(),
|
||||
onToggleFavorite,
|
||||
renderTrigger,
|
||||
@@ -598,26 +557,26 @@ export function CombinedModelSelector({
|
||||
const [searchQuery, setSearchQuery] = useState("");
|
||||
const [searchResetKey, bumpSearchResetKey] = useReducer((key: number) => key + 1, 0);
|
||||
|
||||
// Single-provider mode: only one provider with models → skip Level 1 entirely
|
||||
// Single-provider mode: only one provider → skip Level 1 entirely
|
||||
const singleProviderView = useMemo<SelectorView | null>(() => {
|
||||
const providers = Array.from(allProviderModels.keys());
|
||||
if (providers.length !== 1) return null;
|
||||
const providerId = providers[0];
|
||||
const label = resolveProviderLabel(providerDefinitions, providerId);
|
||||
return { kind: "provider", providerId, providerLabel: label };
|
||||
}, [allProviderModels, providerDefinitions]);
|
||||
const provider = providers[0];
|
||||
if (!provider) return null;
|
||||
return { kind: "provider", providerId: provider.id, providerLabel: provider.label };
|
||||
}, [providers]);
|
||||
|
||||
const computeInitialView = useCallback((): SelectorView => {
|
||||
if (singleProviderView) return singleProviderView;
|
||||
|
||||
const selectedFavoriteKey = `${selectedProvider}:${selectedModel}`;
|
||||
if (selectedProvider && selectedModel && !favoriteKeys.has(selectedFavoriteKey)) {
|
||||
const label = resolveProviderLabel(providerDefinitions, selectedProvider);
|
||||
return { kind: "provider", providerId: selectedProvider, providerLabel: label };
|
||||
const provider = providers.find((entry) => entry.id === selectedProvider);
|
||||
if (provider)
|
||||
return { kind: "provider", providerId: provider.id, providerLabel: provider.label };
|
||||
}
|
||||
|
||||
return { kind: "all" };
|
||||
}, [singleProviderView, selectedProvider, selectedModel, favoriteKeys, providerDefinitions]);
|
||||
}, [singleProviderView, selectedProvider, selectedModel, favoriteKeys, providers]);
|
||||
|
||||
const handleOpenChange = useCallback(
|
||||
(open: boolean) => {
|
||||
@@ -652,28 +611,28 @@ export function CombinedModelSelector({
|
||||
if (!hasSelectedProvider) {
|
||||
return "Select model";
|
||||
}
|
||||
const models = allProviderModels.get(selectedProvider);
|
||||
if (models && models.length === 0) {
|
||||
const provider = providers.find((entry) => entry.id === selectedProvider);
|
||||
if (provider?.rows.length === 0) {
|
||||
return "Default";
|
||||
}
|
||||
return isLoading ? "Loading..." : "Select model";
|
||||
}
|
||||
const models = allProviderModels.get(selectedProvider);
|
||||
if (!models) {
|
||||
const provider = providers.find((entry) => entry.id === selectedProvider);
|
||||
if (!provider) {
|
||||
return isLoading ? "Loading..." : "Select model";
|
||||
}
|
||||
const model = models.find((entry) => entry.id === selectedModel);
|
||||
return model?.label ?? resolveDefaultModelLabel(models);
|
||||
}, [allProviderModels, hasSelectedProvider, isLoading, selectedModel, selectedProvider]);
|
||||
const model = provider.rows.find((entry) => entry.modelId === selectedModel);
|
||||
return model?.modelLabel ?? resolveDefaultModelLabel(provider.rows);
|
||||
}, [hasSelectedProvider, isLoading, providers, selectedModel, selectedProvider]);
|
||||
|
||||
const desktopFixedHeight = useMemo(() => {
|
||||
if (view.kind !== "provider") {
|
||||
return undefined;
|
||||
}
|
||||
const models = allProviderModels.get(view.providerId);
|
||||
const modelCount = models?.length ?? 0;
|
||||
const modelCount =
|
||||
providers.find((provider) => provider.id === view.providerId)?.rows.length ?? 0;
|
||||
return Math.min(80 + modelCount * 40, 400);
|
||||
}, [allProviderModels, view]);
|
||||
}, [providers, view]);
|
||||
|
||||
const triggerLabel = useMemo(() => {
|
||||
if (selectedModelLabel === "Loading..." || selectedModelLabel === "Select model") {
|
||||
@@ -808,14 +767,12 @@ export function CombinedModelSelector({
|
||||
{isContentReady ? (
|
||||
<SelectorContent
|
||||
view={view}
|
||||
providerDefinitions={providerDefinitions}
|
||||
allProviderModels={allProviderModels}
|
||||
providers={providers}
|
||||
selectedProvider={selectedProvider}
|
||||
selectedModel={selectedModel}
|
||||
searchQuery={searchQuery}
|
||||
favoriteKeys={favoriteKeys}
|
||||
onSelect={handleSelect}
|
||||
canSelectProvider={canSelectProvider}
|
||||
onToggleFavorite={onToggleFavorite}
|
||||
onDrillDown={handleDrillDown}
|
||||
/>
|
||||
|
||||
@@ -1,54 +1,60 @@
|
||||
import type { AgentModelDefinition } from "@server/server/agent/agent-sdk-types";
|
||||
import type {
|
||||
AgentModelDefinition,
|
||||
ProviderSnapshotEntry,
|
||||
} from "@server/server/agent/agent-sdk-types";
|
||||
import type { AgentProviderDefinition } from "@server/server/agent/provider-manifest";
|
||||
import { buildFavoriteModelKey, type FavoriteModelRow } from "@/hooks/use-form-preferences";
|
||||
import { compareMatchScores, scoreTextFields } from "@/utils/score-match";
|
||||
|
||||
export type SelectorModelRow = FavoriteModelRow;
|
||||
export type SelectorModelRow = FavoriteModelRow & { isDefault?: boolean };
|
||||
|
||||
export interface SelectorProviderGroup {
|
||||
providerId: string;
|
||||
providerLabel: string;
|
||||
export interface ModelSelectorProvider {
|
||||
id: string;
|
||||
label: string;
|
||||
rows: SelectorModelRow[];
|
||||
hasNoModels: boolean;
|
||||
}
|
||||
|
||||
export function resolveProviderLabel(
|
||||
providerDefinitions: AgentProviderDefinition[],
|
||||
providerId: string,
|
||||
): string {
|
||||
return (
|
||||
providerDefinitions.find((definition) => definition.id === providerId)?.label ?? providerId
|
||||
);
|
||||
}
|
||||
|
||||
export function buildSelectedTriggerLabel(modelLabel: string): string {
|
||||
return modelLabel;
|
||||
}
|
||||
|
||||
export function buildModelRows(
|
||||
export function buildModelSelectorProviders(
|
||||
providerDefinitions: AgentProviderDefinition[],
|
||||
allProviderModels: Map<string, AgentModelDefinition[]>,
|
||||
): SelectorModelRow[] {
|
||||
const providerLabelMap = new Map(
|
||||
providerDefinitions.map((definition) => [definition.id, definition.label]),
|
||||
);
|
||||
const rows: SelectorModelRow[] = [];
|
||||
): ModelSelectorProvider[] {
|
||||
return providerDefinitions.map((definition) => ({
|
||||
id: definition.id,
|
||||
label: definition.label,
|
||||
rows: (allProviderModels.get(definition.id) ?? []).map((model) => ({
|
||||
favoriteKey: buildFavoriteModelKey({ provider: definition.id, modelId: model.id }),
|
||||
provider: definition.id,
|
||||
providerLabel: definition.label,
|
||||
modelId: model.id,
|
||||
modelLabel: model.label,
|
||||
description: model.description,
|
||||
isDefault: model.isDefault,
|
||||
})),
|
||||
}));
|
||||
}
|
||||
|
||||
for (const definition of providerDefinitions) {
|
||||
const providerLabel = providerLabelMap.get(definition.id) ?? definition.label;
|
||||
for (const model of allProviderModels.get(definition.id) ?? []) {
|
||||
rows.push({
|
||||
favoriteKey: buildFavoriteModelKey({ provider: definition.id, modelId: model.id }),
|
||||
provider: definition.id,
|
||||
providerLabel,
|
||||
export function buildSelectableModelSelectorProviders(
|
||||
entries: ProviderSnapshotEntry[] | undefined,
|
||||
): ModelSelectorProvider[] {
|
||||
return (entries ?? [])
|
||||
.filter((entry) => entry.enabled && entry.status === "ready")
|
||||
.map((entry) => ({
|
||||
id: entry.provider,
|
||||
label: entry.label ?? entry.provider,
|
||||
rows: (entry.models ?? []).map((model) => ({
|
||||
favoriteKey: buildFavoriteModelKey({ provider: entry.provider, modelId: model.id }),
|
||||
provider: entry.provider,
|
||||
providerLabel: entry.label ?? entry.provider,
|
||||
modelId: model.id,
|
||||
modelLabel: model.label,
|
||||
description: model.description,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
return rows;
|
||||
isDefault: model.isDefault,
|
||||
})),
|
||||
}));
|
||||
}
|
||||
|
||||
export function matchesSearch(row: SelectorModelRow, normalizedQuery: string): boolean {
|
||||
@@ -82,55 +88,3 @@ export function filterAndRankModelRows(
|
||||
|
||||
return scored.map((entry) => entry.row);
|
||||
}
|
||||
|
||||
export function buildProviderGroups(
|
||||
providerDefinitions: AgentProviderDefinition[],
|
||||
allProviderModels: Map<string, AgentModelDefinition[]>,
|
||||
rows: SelectorModelRow[],
|
||||
normalizedQuery: string,
|
||||
): SelectorProviderGroup[] {
|
||||
const rowsByProvider = new Map<string, SelectorModelRow[]>();
|
||||
for (const row of rows) {
|
||||
const providerRows = rowsByProvider.get(row.provider);
|
||||
if (providerRows) {
|
||||
providerRows.push(row);
|
||||
} else {
|
||||
rowsByProvider.set(row.provider, [row]);
|
||||
}
|
||||
}
|
||||
|
||||
const groups: SelectorProviderGroup[] = [];
|
||||
for (const definition of providerDefinitions) {
|
||||
const providerRows = rowsByProvider.get(definition.id) ?? [];
|
||||
if (providerRows.length > 0) {
|
||||
groups.push({
|
||||
providerId: definition.id,
|
||||
providerLabel: definition.label,
|
||||
rows: providerRows,
|
||||
hasNoModels: false,
|
||||
});
|
||||
continue;
|
||||
}
|
||||
|
||||
const models = allProviderModels.get(definition.id);
|
||||
if (!models || models.length > 0) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const providerMatches =
|
||||
!normalizedQuery ||
|
||||
scoreTextFields(normalizedQuery, [definition.label, definition.id]) !== null;
|
||||
if (!providerMatches) {
|
||||
continue;
|
||||
}
|
||||
|
||||
groups.push({
|
||||
providerId: definition.id,
|
||||
providerLabel: definition.label,
|
||||
rows: [],
|
||||
hasNoModels: true,
|
||||
});
|
||||
}
|
||||
|
||||
return groups;
|
||||
}
|
||||
|
||||
@@ -19,7 +19,7 @@ describe("resolveStatusControlMode", () => {
|
||||
selectedModel: "",
|
||||
onSelectModel: () => undefined,
|
||||
isModelLoading: false,
|
||||
allProviderModels: new Map(),
|
||||
modelSelectorProviders: [],
|
||||
isAllModelsLoading: false,
|
||||
onSelectProviderAndModel: () => undefined,
|
||||
thinkingOptions: [],
|
||||
|
||||
@@ -8,6 +8,10 @@ import type {
|
||||
} from "@server/server/agent/agent-sdk-types";
|
||||
import { useHosts } from "@/runtime/host-runtime";
|
||||
import { buildProviderDefinitions } from "@/utils/provider-definitions";
|
||||
import {
|
||||
buildSelectableModelSelectorProviders,
|
||||
type ModelSelectorProvider,
|
||||
} from "@/components/combined-model-selector.utils";
|
||||
import { useProvidersSnapshot } from "./use-providers-snapshot";
|
||||
import {
|
||||
useFormPreferences,
|
||||
@@ -63,6 +67,7 @@ export interface UseAgentFormStateResult {
|
||||
modeOptions: AgentMode[];
|
||||
availableModels: AgentModelDefinition[];
|
||||
allProviderModels: Map<string, AgentModelDefinition[]>;
|
||||
modelSelectorProviders: ModelSelectorProvider[];
|
||||
isAllModelsLoading: boolean;
|
||||
availableThinkingOptions: NonNullable<AgentModelDefinition["thinkingOptions"]>;
|
||||
isModelLoading: boolean;
|
||||
@@ -237,6 +242,10 @@ export function useAgentFormState(options: UseAgentFormStateOptions = {}): UseAg
|
||||
() => buildAllProviderModels(snapshotEntries),
|
||||
[snapshotEntries],
|
||||
);
|
||||
const snapshotModelSelectorProviders = useMemo(
|
||||
() => buildSelectableModelSelectorProviders(snapshotEntries),
|
||||
[snapshotEntries],
|
||||
);
|
||||
const snapshotSelectedEntry = useMemo(
|
||||
() =>
|
||||
formState.provider
|
||||
@@ -255,6 +264,7 @@ export function useAgentFormState(options: UseAgentFormStateOptions = {}): UseAg
|
||||
const providerDefinitionMap = snapshotProviderDefinitionMap;
|
||||
const selectableProviderDefinitionMap = snapshotSelectableProviderDefinitionMap;
|
||||
const allProviderModels = snapshotAllProviderModels;
|
||||
const modelSelectorProviders = snapshotModelSelectorProviders;
|
||||
const availableModels = snapshotSelectedProviderModels;
|
||||
const modeOptions = snapshotSelectedProviderModes;
|
||||
const isAllModelsLoading = snapshotIsLoading || selectedProviderIsLoading;
|
||||
@@ -516,6 +526,7 @@ export function useAgentFormState(options: UseAgentFormStateOptions = {}): UseAg
|
||||
modeOptions,
|
||||
availableModels: availableModels ?? [],
|
||||
allProviderModels,
|
||||
modelSelectorProviders,
|
||||
isAllModelsLoading,
|
||||
availableThinkingOptions,
|
||||
isModelLoading,
|
||||
@@ -548,6 +559,7 @@ export function useAgentFormState(options: UseAgentFormStateOptions = {}): UseAg
|
||||
modeOptions,
|
||||
availableModels,
|
||||
allProviderModels,
|
||||
modelSelectorProviders,
|
||||
isAllModelsLoading,
|
||||
availableThinkingOptions,
|
||||
isModelLoading,
|
||||
|
||||
@@ -88,7 +88,7 @@ export function buildDraftStatusControls(input: {
|
||||
selectedModel: formState.selectedModel,
|
||||
onSelectModel: formState.setModelFromUser,
|
||||
isModelLoading: formState.isModelLoading,
|
||||
allProviderModels: formState.allProviderModels,
|
||||
modelSelectorProviders: formState.modelSelectorProviders,
|
||||
isAllModelsLoading: formState.isAllModelsLoading,
|
||||
onSelectProviderAndModel: formState.setProviderAndModelFromUser,
|
||||
thinkingOptions: formState.availableThinkingOptions,
|
||||
|
||||
@@ -81,6 +81,22 @@ vi.mock("./use-agent-form-state", () => ({
|
||||
],
|
||||
],
|
||||
]),
|
||||
modelSelectorProviders: [
|
||||
{
|
||||
id: "codex",
|
||||
label: "Codex",
|
||||
rows: [
|
||||
{
|
||||
favoriteKey: "codex:gpt-5.4",
|
||||
provider: "codex",
|
||||
providerLabel: "Codex",
|
||||
modelId: "gpt-5.4",
|
||||
modelLabel: "gpt-5.4",
|
||||
isDefault: true,
|
||||
},
|
||||
],
|
||||
},
|
||||
],
|
||||
isAllModelsLoading: false,
|
||||
availableThinkingOptions: [
|
||||
{ id: "medium", label: "Medium" },
|
||||
|
||||
Reference in New Issue
Block a user