diff --git a/apps/web/src/components/settings/ChatGptConnectDialog.tsx b/apps/web/src/components/settings/ChatGptConnectDialog.tsx index 03fc398d7..05263ea36 100644 --- a/apps/web/src/components/settings/ChatGptConnectDialog.tsx +++ b/apps/web/src/components/settings/ChatGptConnectDialog.tsx @@ -81,6 +81,9 @@ export function ChatGptConnectDialog({ await queryClient.invalidateQueries({ queryKey: trpc.taskModels.providerSetup.queryKey(), }); + await queryClient.invalidateQueries({ + queryKey: trpc.taskModels.launchOptions.queryKey(), + }); await queryClient.invalidateQueries({ queryKey: trpc.chatgptSubscription.status.queryKey(), }); diff --git a/apps/web/src/components/settings/InferenceProviderSection.tsx b/apps/web/src/components/settings/InferenceProviderSection.tsx index b72cd5647..efb9b00c6 100644 --- a/apps/web/src/components/settings/InferenceProviderSection.tsx +++ b/apps/web/src/components/settings/InferenceProviderSection.tsx @@ -595,12 +595,17 @@ export function InferenceProviderSection({ trpc.chatgptSubscription.disconnect.mutationOptions({ onSuccess: async () => { toast.success('Disconnected ChatGPT subscription.'); - await queryClient.invalidateQueries({ - queryKey: trpc.taskModels.providerSetup.queryKey(), - }); - await queryClient.invalidateQueries({ - queryKey: trpc.chatgptSubscription.status.queryKey(), - }); + await Promise.all([ + queryClient.invalidateQueries({ + queryKey: trpc.taskModels.providerSetup.queryKey(), + }), + queryClient.invalidateQueries({ + queryKey: trpc.taskModels.launchOptions.queryKey(), + }), + queryClient.invalidateQueries({ + queryKey: trpc.chatgptSubscription.status.queryKey(), + }), + ]); }, onError: (error) => { toast.error(error.message); diff --git a/apps/web/src/components/settings/ModelSettingsSection.tsx b/apps/web/src/components/settings/ModelSettingsSection.tsx index fd207d3cb..8c7fc4093 100644 --- a/apps/web/src/components/settings/ModelSettingsSection.tsx +++ b/apps/web/src/components/settings/ModelSettingsSection.tsx @@ -55,13 +55,12 @@ import { CHATGPT_SUBSCRIPTION_PROVIDER_ID, DEFAULT_MODEL_ROLE_REASONING_EFFORTS, REASONING_EFFORT_OPTIONS, - SETUP_MODEL_PROVIDER_CATALOG, - getDisplayModelProviderId, - getModelProviderLabel, + groupModelsByDisplayProvider, getSetupModelProvider, normalizeOptionalReasoningEffort, } from '@roomote/types'; import type { + DisplayModelProviderGroup, ReasoningEffort, SetupModelProviderId, SetupModelProviderStatus, @@ -301,7 +300,7 @@ function TaskModelRoleEditor({ managedByEnv: boolean; reasoningManagedByEnv: boolean; selectValue: string; - optionGroups: ProviderModelGroup[]; + optionGroups: DisplayModelProviderGroup[]; supportsReasoning: boolean; reasoningEffort: ReasoningEffort | null; onModelChange: (value: string) => void; @@ -320,6 +319,7 @@ function TaskModelRoleEditor({ : managedByEnv ? `Set by ${config.modelEnvVarName}, not changeable in the UI.` : `Set by ${config.reasoningEnvVarName}, not changeable in the UI.`; + const showProviderHeaders = optionGroups.length > 1; return (
@@ -356,16 +356,24 @@ function TaskModelRoleEditor({ Same as coding model )} - {optionGroups.map((group) => ( - - {group.label} - {group.items.map((option) => ( - - {option.displayName} - - ))} - - ))} + {showProviderHeaders + ? optionGroups.map((group) => ( + + {group.label} + {group.items.map((option) => ( + + {option.displayName} + + ))} + + )) + : optionGroups.flatMap((group) => + group.items.map((option) => ( + + {option.displayName} + + )), + )} {supportsReasoning && ( @@ -383,52 +391,6 @@ function TaskModelRoleEditor({ ); } -type ProviderModelGroup = { - providerId: string; - label: string; - items: T[]; -}; - -const KNOWN_MODEL_PROVIDER_ORDER = SETUP_MODEL_PROVIDER_CATALOG.map( - (provider) => provider.id as string, -); - -function groupByModelProvider( - items: T[], - options?: { - chatgptConnected?: boolean; - }, -): ProviderModelGroup[] { - const groups = new Map(); - - for (const item of items) { - const providerId = - getDisplayModelProviderId(item.id, { - chatgptConnected: options?.chatgptConnected, - }) ?? 'other'; - const groupItems = groups.get(providerId) ?? []; - groupItems.push(item); - groups.set(providerId, groupItems); - } - - return [...groups.entries()] - .sort(([left], [right]) => { - const leftIndex = KNOWN_MODEL_PROVIDER_ORDER.indexOf(left); - const rightIndex = KNOWN_MODEL_PROVIDER_ORDER.indexOf(right); - const leftOrder = - leftIndex === -1 ? KNOWN_MODEL_PROVIDER_ORDER.length : leftIndex; - const rightOrder = - rightIndex === -1 ? KNOWN_MODEL_PROVIDER_ORDER.length : rightIndex; - - return leftOrder - rightOrder || left.localeCompare(right); - }) - .map(([providerId, groupItems]) => ({ - providerId, - label: getModelProviderLabel(providerId), - items: groupItems, - })); -} - // The ChatGPT subscription provider has no model-id prefix of its own: its // models keep the `openai/` prefix so they are selected and billed like other // OpenAI models at runtime. Map it to `openai` wherever a model id is @@ -1012,20 +974,22 @@ export function ModelSettingsSection({ return options; }, [settingsData]); const codingModelGroups = useMemo( - () => groupByModelProvider(codingModelOptions, { chatgptConnected }), + () => + groupModelsByDisplayProvider(codingModelOptions, { chatgptConnected }), [codingModelOptions, chatgptConnected], ); const helperModelGroups = useMemo( - () => groupByModelProvider(helperModelOptions, { chatgptConnected }), + () => + groupModelsByDisplayProvider(helperModelOptions, { chatgptConnected }), [helperModelOptions, chatgptConnected], ); const modelGroups = useMemo( - () => groupByModelProvider(models, { chatgptConnected }), + () => groupModelsByDisplayProvider(models, { chatgptConnected }), [models, chatgptConnected], ); const roleOptionGroups: Record< TaskModelRole, - ProviderModelGroup[] + DisplayModelProviderGroup[] > = { coding: codingModelGroups, helper: helperModelGroups, diff --git a/apps/web/src/components/tasks/ModelSelect.client.test.tsx b/apps/web/src/components/tasks/ModelSelect.client.test.tsx new file mode 100644 index 000000000..26ebf685b --- /dev/null +++ b/apps/web/src/components/tasks/ModelSelect.client.test.tsx @@ -0,0 +1,119 @@ +import { render, screen } from '@testing-library/react'; +import { describe, expect, it, vi } from 'vitest'; +import type { ReactNode } from 'react'; + +const launchModelsData = vi.hoisted(() => ({ + current: null as { + defaultModelId: string; + chatgptConnected: boolean; + models: Array<{ + id: string; + displayName: string; + isDefault?: boolean; + }>; + } | null, +})); + +vi.mock('@/hooks/task-models/useLaunchTaskModels', () => ({ + useLaunchTaskModels: () => ({ + data: launchModelsData.current, + isPending: launchModelsData.current === null, + }), +})); + +vi.mock('@/components/system', async () => { + const actual = await vi.importActual( + '@/components/system', + ); + + return { + ...actual, + Select: ({ children }: { children: ReactNode }) =>
{children}
, + SelectTrigger: ({ children, ...props }: { children: ReactNode }) => ( + + ), + SelectValue: ({ placeholder }: { placeholder?: string }) => ( + {placeholder} + ), + SelectContent: ({ children }: { children: ReactNode }) => ( +
{children}
+ ), + SelectGroup: ({ children }: { children: ReactNode }) => ( +
{children}
+ ), + SelectLabel: ({ children }: { children: ReactNode }) => ( +
{children}
+ ), + SelectItem: ({ + children, + value, + }: { + children: ReactNode; + value: string; + }) => ( +
+ {children} +
+ ), + }; +}); + +import { ModelSelect } from './ModelSelect'; + +describe('ModelSelect', () => { + it('shows provider headers when multiple providers are represented', () => { + launchModelsData.current = { + defaultModelId: 'openrouter/x-ai/grok-4.5', + chatgptConnected: true, + models: [ + { + id: 'openrouter/x-ai/grok-4.5', + displayName: 'Grok 4.5', + isDefault: true, + }, + { + id: 'openai/gpt-5.6-terra', + displayName: 'GPT 5.6 Terra', + }, + ], + }; + + render( + , + ); + + expect( + screen.getAllByTestId('select-label').map((node) => node.textContent), + ).toEqual(['OpenRouter', 'ChatGPT (subscription)']); + expect(screen.getByText('Grok 4.5 (Default)')).toBeTruthy(); + expect(screen.getByText('GPT 5.6 Terra')).toBeTruthy(); + }); + + it('omits provider headers when only one provider group is present', () => { + launchModelsData.current = { + defaultModelId: 'openrouter/x-ai/grok-4.5', + chatgptConnected: false, + models: [ + { + id: 'openrouter/x-ai/grok-4.5', + displayName: 'Grok 4.5', + isDefault: true, + }, + { + id: 'openrouter/anthropic/claude-sonnet-5', + displayName: 'Claude Sonnet 5', + }, + ], + }; + + render( + , + ); + + expect(screen.queryByTestId('select-label')).toBeNull(); + expect(screen.getByText('Grok 4.5 (Default)')).toBeTruthy(); + expect(screen.getByText('Claude Sonnet 5')).toBeTruthy(); + }); +}); diff --git a/apps/web/src/components/tasks/ModelSelect.tsx b/apps/web/src/components/tasks/ModelSelect.tsx index c37756954..5ae227bf3 100644 --- a/apps/web/src/components/tasks/ModelSelect.tsx +++ b/apps/web/src/components/tasks/ModelSelect.tsx @@ -1,12 +1,14 @@ 'use client'; import { useMemo } from 'react'; -import { getModelProviderLabel, getTaskModelProviderId } from '@roomote/types'; +import { groupModelsByDisplayProvider } from '@roomote/types'; import { Select, SelectContent, + SelectGroup, SelectItem, + SelectLabel, SelectTrigger, SelectValue, } from '@/components/system'; @@ -21,6 +23,13 @@ type ModelSelectProps = { ariaLabel?: string; }; +function modelOptionLabel(model: { + displayName: string; + isDefault?: boolean; +}): string { + return `${model.displayName}${model.isDefault ? ' (Default)' : ''}`; +} + export function ModelSelect({ value, onValueChange, @@ -29,23 +38,16 @@ export function ModelSelect({ ariaLabel = 'Model', }: ModelSelectProps) { const { data, isPending } = useLaunchTaskModels(); - const sortedModels = useMemo( - () => - [...(data?.models ?? [])].sort((left, right) => { - const leftProvider = getModelProviderLabel( - getTaskModelProviderId(left.id) ?? 'other', - ); - const rightProvider = getModelProviderLabel( - getTaskModelProviderId(right.id) ?? 'other', - ); + const modelGroups = useMemo(() => { + const sortedModels = [...(data?.models ?? [])].sort((left, right) => + left.displayName.localeCompare(right.displayName), + ); - return ( - leftProvider.localeCompare(rightProvider) || - left.displayName.localeCompare(right.displayName) - ); - }), - [data?.models], - ); + return groupModelsByDisplayProvider(sortedModels, { + chatgptConnected: data?.chatgptConnected, + }); + }, [data?.chatgptConnected, data?.models]); + const showProviderHeaders = modelGroups.length > 1; return ( ); diff --git a/apps/web/src/trpc/commands/task-models/index.ts b/apps/web/src/trpc/commands/task-models/index.ts index 71e8f5eb1..ca9646b86 100644 --- a/apps/web/src/trpc/commands/task-models/index.ts +++ b/apps/web/src/trpc/commands/task-models/index.ts @@ -703,12 +703,16 @@ export async function deleteTaskModelProviderCommand( } export async function getLaunchTaskModelsCommand(_auth: UserAuthSuccess) { - const settings = await getDeploymentTaskModelSettings(); + const [settings, chatgptConnected] = await Promise.all([ + getDeploymentTaskModelSettings(), + isChatGptSubscriptionConnected(), + ]); const enabledModels = getEnabledTaskModels(settings); const defaultModel = getDefaultTaskModel(settings); return { defaultModelId: defaultModel.id, + chatgptConnected, models: enabledModels.map((option) => ({ ...option, isDefault: option.id === defaultModel.id, diff --git a/packages/types/src/model-provider-config.test.ts b/packages/types/src/model-provider-config.test.ts index 54df7ee8d..9f50cf3ad 100644 --- a/packages/types/src/model-provider-config.test.ts +++ b/packages/types/src/model-provider-config.test.ts @@ -5,6 +5,7 @@ import { DEFAULT_MODEL_PROVIDER_ENV_KEYS, DEFAULT_TASK_MODEL_ID, getDisplayModelProviderId, + groupModelsByDisplayProvider, getModelProviderEnvKeyCandidates, getReasoningEffortLabel, isInlineGoogleCredentialsValue, @@ -348,6 +349,39 @@ describe('SETUP_MODEL_PROVIDER_CATALOG', () => { ).toBe('chatgpt'); }); + it('groups model chooser options by display provider and catalog order', () => { + const groups = groupModelsByDisplayProvider( + [ + { id: 'openai/gpt-5.6-terra', displayName: 'GPT 5.6 Terra' }, + { + id: 'openrouter/x-ai/grok-4.5', + displayName: 'Grok 4.5', + }, + { + id: 'openrouter/anthropic/claude-sonnet-5', + displayName: 'Claude Sonnet 5', + }, + ], + { chatgptConnected: true }, + ); + + expect(groups.map((group) => group.providerId)).toEqual([ + 'openrouter', + 'chatgpt', + ]); + expect(groups[0]).toMatchObject({ + label: 'OpenRouter', + items: [ + { id: 'openrouter/x-ai/grok-4.5' }, + { id: 'openrouter/anthropic/claude-sonnet-5' }, + ], + }); + expect(groups[1]).toMatchObject({ + label: 'ChatGPT (subscription)', + items: [{ id: 'openai/gpt-5.6-terra' }], + }); + }); + it('maps Requesty to the REQUESTY_API_KEY env var', () => { const requestyProvider = SETUP_MODEL_PROVIDER_CATALOG.find( (provider) => provider.id === 'requesty', diff --git a/packages/types/src/model-provider-config.ts b/packages/types/src/model-provider-config.ts index 889d5b33c..b6e7b7827 100644 --- a/packages/types/src/model-provider-config.ts +++ b/packages/types/src/model-provider-config.ts @@ -611,6 +611,56 @@ export function getDisplayModelProviderId( return runtimeProviderId; } +export type DisplayModelProviderGroup = { + providerId: string; + label: string; + items: T[]; +}; + +const KNOWN_MODEL_PROVIDER_ORDER = SETUP_MODEL_PROVIDER_CATALOG.map( + (provider) => provider.id as string, +); + +/** + * Groups model options by display provider for chooser UIs. Preserves input + * order within each group and sorts groups by the setup provider catalog. + */ +export function groupModelsByDisplayProvider( + items: T[], + options?: { + chatgptConnected?: boolean; + }, +): DisplayModelProviderGroup[] { + const groups = new Map(); + + for (const item of items) { + const providerId = + getDisplayModelProviderId(item.id, { + chatgptConnected: options?.chatgptConnected, + }) ?? 'other'; + const groupItems = groups.get(providerId) ?? []; + groupItems.push(item); + groups.set(providerId, groupItems); + } + + return [...groups.entries()] + .sort(([left], [right]) => { + const leftIndex = KNOWN_MODEL_PROVIDER_ORDER.indexOf(left); + const rightIndex = KNOWN_MODEL_PROVIDER_ORDER.indexOf(right); + const leftOrder = + leftIndex === -1 ? KNOWN_MODEL_PROVIDER_ORDER.length : leftIndex; + const rightOrder = + rightIndex === -1 ? KNOWN_MODEL_PROVIDER_ORDER.length : rightIndex; + + return leftOrder - rightOrder || left.localeCompare(right); + }) + .map(([providerId, groupItems]) => ({ + providerId, + label: getModelProviderLabel(providerId), + items: groupItems, + })); +} + export function getSetupModelProviderForEnvVarName( envVarName: string | null | undefined, ): SetupModelProviderDescriptor | undefined {