Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions apps/web/src/components/settings/ChatGptConnectDialog.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -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(),
});
Expand Down
17 changes: 11 additions & 6 deletions apps/web/src/components/settings/InferenceProviderSection.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down
92 changes: 28 additions & 64 deletions apps/web/src/components/settings/ModelSettingsSection.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -301,7 +300,7 @@ function TaskModelRoleEditor({
managedByEnv: boolean;
reasoningManagedByEnv: boolean;
selectValue: string;
optionGroups: ProviderModelGroup<EditableRuntimeModelOption>[];
optionGroups: DisplayModelProviderGroup<EditableRuntimeModelOption>[];
supportsReasoning: boolean;
reasoningEffort: ReasoningEffort | null;
onModelChange: (value: string) => void;
Expand All @@ -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 (
<div className="flex items-start gap-3 py-3 first:pt-0 last:pb-0">
Expand Down Expand Up @@ -356,16 +356,24 @@ function TaskModelRoleEditor({
Same as coding model
</SelectItem>
)}
{optionGroups.map((group) => (
<SelectGroup key={group.providerId}>
<SelectLabel>{group.label}</SelectLabel>
{group.items.map((option) => (
<SelectItem key={option.id} value={option.id}>
{option.displayName}
</SelectItem>
))}
</SelectGroup>
))}
{showProviderHeaders
? optionGroups.map((group) => (
<SelectGroup key={group.providerId}>
<SelectLabel>{group.label}</SelectLabel>
{group.items.map((option) => (
<SelectItem key={option.id} value={option.id}>
{option.displayName}
</SelectItem>
))}
</SelectGroup>
))
: optionGroups.flatMap((group) =>
group.items.map((option) => (
<SelectItem key={option.id} value={option.id}>
{option.displayName}
</SelectItem>
)),
)}
</SelectContent>
</Select>
{supportsReasoning && (
Expand All @@ -383,52 +391,6 @@ function TaskModelRoleEditor({
);
}

type ProviderModelGroup<T extends { id: string }> = {
providerId: string;
label: string;
items: T[];
};

const KNOWN_MODEL_PROVIDER_ORDER = SETUP_MODEL_PROVIDER_CATALOG.map(
(provider) => provider.id as string,
);

function groupByModelProvider<T extends { id: string }>(
items: T[],
options?: {
chatgptConnected?: boolean;
},
): ProviderModelGroup<T>[] {
const groups = new Map<string, T[]>();

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
Expand Down Expand Up @@ -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<EditableRuntimeModelOption>[]
DisplayModelProviderGroup<EditableRuntimeModelOption>[]
> = {
coding: codingModelGroups,
helper: helperModelGroups,
Expand Down
119 changes: 119 additions & 0 deletions apps/web/src/components/tasks/ModelSelect.client.test.tsx

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

60 changes: 37 additions & 23 deletions apps/web/src/components/tasks/ModelSelect.tsx
Original file line number Diff line number Diff line change
@@ -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';
Expand All @@ -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,
Expand All @@ -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 (
<Select
Expand All @@ -61,12 +63,24 @@ export function ModelSelect({
<SelectValue placeholder="Model" />
</SelectTrigger>
<SelectContent>
{sortedModels.map((model) => (
<SelectItem key={model.id} value={model.id}>
{model.displayName}
{model.isDefault ? ' (Default)' : ''}
</SelectItem>
))}
{showProviderHeaders
? modelGroups.map((group) => (
<SelectGroup key={group.providerId}>
<SelectLabel>{group.label}</SelectLabel>
{group.items.map((model) => (
<SelectItem key={model.id} value={model.id}>
{modelOptionLabel(model)}
</SelectItem>
))}
</SelectGroup>
))
: modelGroups.flatMap((group) =>
group.items.map((model) => (
<SelectItem key={model.id} value={model.id}>
{modelOptionLabel(model)}
</SelectItem>
)),
)}
</SelectContent>
</Select>
);
Expand Down
Loading
Loading