Streamline model selector provider data

This commit is contained in:
Mohamed Boudra
2026-05-19 20:01:46 +07:00
parent bc7798af28
commit 5dd6b030f2
8 changed files with 298 additions and 364 deletions

View File

@@ -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}

View File

@@ -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");
});
});

View File

@@ -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}
/>

View File

@@ -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;
}

View File

@@ -19,7 +19,7 @@ describe("resolveStatusControlMode", () => {
selectedModel: "",
onSelectModel: () => undefined,
isModelLoading: false,
allProviderModels: new Map(),
modelSelectorProviders: [],
isAllModelsLoading: false,
onSelectProviderAndModel: () => undefined,
thinkingOptions: [],

View File

@@ -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,

View File

@@ -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,

View File

@@ -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" },