From 966f35da6c5229ff3abecbd53c4487c910425d68 Mon Sep 17 00:00:00 2001 From: Bohdan Triapitsyn Date: Wed, 20 May 2026 16:26:40 +0300 Subject: [PATCH] refactor: unify model picker behavior Share model picker UI across chat, agents, and multi-run Prefer live provider limits with metadata fallback --- .../ui/src/components/chat/ModelControls.tsx | 947 +++--------------- .../model-picker/ModelPickerList.tsx | 708 +++++++++++++ .../components/multirun/ModelMultiSelect.tsx | 418 ++------ .../sections/agents/ModelSelector.tsx | 773 +++----------- packages/ui/src/lib/modelMetadata.ts | 35 + 5 files changed, 1082 insertions(+), 1799 deletions(-) create mode 100644 packages/ui/src/components/model-picker/ModelPickerList.tsx create mode 100644 packages/ui/src/lib/modelMetadata.ts diff --git a/packages/ui/src/components/chat/ModelControls.tsx b/packages/ui/src/components/chat/ModelControls.tsx index 7a52a731..4fc705cb 100644 --- a/packages/ui/src/components/chat/ModelControls.tsx +++ b/packages/ui/src/components/chat/ModelControls.tsx @@ -1,14 +1,4 @@ import React from 'react'; -import { - DndContext, - PointerSensor, - closestCenter, - useSensor, - useSensors, - type DragEndEvent, -} from '@dnd-kit/core'; -import { SortableContext, useSortable, verticalListSortingStrategy } from '@dnd-kit/sortable'; -import { CSS as DndCSS } from '@dnd-kit/utilities'; import type { EditPermissionMode } from '@/stores/types/sessionTypes'; import type { ModelMetadata } from '@/types'; import { @@ -23,14 +13,15 @@ import { Input } from '@/components/ui/input'; import { MobileOverlayPanel } from '@/components/ui/MobileOverlayPanel'; import { ProviderLogo } from '@/components/ui/ProviderLogo'; import { ScrollableOverlay } from '@/components/ui/ScrollableOverlay'; -import { TextLoop } from '@/components/ui/TextLoop'; import { Tooltip, TooltipContent, TooltipTrigger } from '@/components/ui/tooltip'; import { Icon } from "@/components/icon/Icon"; import type { IconName } from "@/components/icon/icons"; +import { ModelPickerList, type ModelPickerEntry, type ModelPickerProvider } from '@/components/model-picker/ModelPickerList'; import { useIsVSCodeRuntime } from '@/hooks/useRuntimeAPIs'; import { isDesktopShell } from '@/lib/desktop'; import { getAgentColor } from '@/lib/agentColors'; import { useDeviceInfo } from '@/lib/device'; +import { mergeModelMetadataWithLiveModel } from '@/lib/modelMetadata'; import { getEditModeColors } from '@/lib/permissions/editModeColors'; import { cn, fuzzyMatch } from '@/lib/utils'; import { useContextStore } from '@/stores/contextStore'; @@ -46,7 +37,7 @@ import { useIsTextTruncated } from '@/hooks/useIsTextTruncated'; import { formatEffortLabel, getCycledPrimaryAgentName, type MobileControlsPanel } from './mobileControlsUtils'; import { useI18n } from '@/lib/i18n'; import { useOpenCodeReadiness } from '@/hooks/useOpenCodeReadiness'; -import { eventMatchesShortcut, formatShortcutForDisplay, getEffectiveShortcutCombo, normalizeCombo } from '@/lib/shortcuts'; +import { eventMatchesShortcut, getEffectiveShortcutCombo, normalizeCombo } from '@/lib/shortcuts'; type IconComponent = IconName; @@ -55,46 +46,11 @@ type ProviderModel = Record & { id?: string; name?: string }; type PermissionAction = 'allow' | 'ask' | 'deny'; type PermissionRule = { permission: string; pattern: string; action: PermissionAction }; -type SortableFavoriteHandleProps = { - attributes: ReturnType['attributes']; - listeners: ReturnType['listeners']; - setActivatorNodeRef: ReturnType['setActivatorNodeRef']; - isDragging: boolean; -}; type MobileVariantTarget = { providerId: string; modelId: string }; const buildModelRefKey = (providerID: string, modelID: string) => `${providerID}:${modelID}`; const MAX_INLINE_MOBILE_VARIANT_OPTIONS = 6; -const SortableFavoriteModelRow: React.FC<{ - id: string; - disabled?: boolean; - children: (dragHandleProps: SortableFavoriteHandleProps) => React.ReactNode; -}> = ({ id, disabled = false, children }) => { - const { - attributes, - listeners, - setNodeRef, - setActivatorNodeRef, - transform, - transition, - isDragging, - } = useSortable({ id, disabled }); - - return ( -
- {children({ attributes, listeners, setActivatorNodeRef, isDragging })} -
- ); -}; - const asPermissionRuleset = (value: unknown): PermissionRule[] | null => { if (!Array.isArray(value)) { return null; @@ -281,28 +237,6 @@ const formatCost = (value?: number | null) => { return CURRENCY_FORMATTER.format(value); }; -const formatCompactPrice = (metadata?: ModelMetadata): string | null => { - if (!metadata?.cost) { - return null; - } - - const inputCost = metadata.cost.input; - const outputCost = metadata.cost.output; - const hasInput = typeof inputCost === 'number' && Number.isFinite(inputCost); - const hasOutput = typeof outputCost === 'number' && Number.isFinite(outputCost); - - if (hasInput && hasOutput) { - return `In ${formatCost(inputCost)} · Out ${formatCost(outputCost)}`; - } - if (hasInput) { - return `In ${formatCost(inputCost)}`; - } - if (hasOutput) { - return `Out ${formatCost(outputCost)}`; - } - return null; -}; - const getCapabilityIcons = (metadata?: ModelMetadata) => { const result: { key: string; icon: IconComponent; label: string }[] = []; for (const definition of CAPABILITY_DEFINITIONS) { @@ -419,9 +353,6 @@ export const ModelControls: React.FC = ({ const toggleFavoriteModel = useUIStore((state) => state.toggleFavoriteModel); const reorderFavoriteModel = useUIStore((state) => state.reorderFavoriteModel); const isFavoriteModel = useUIStore((state) => state.isFavoriteModel); - const collapsedModelProviders = useUIStore((state) => state.collapsedModelProviders); - const toggleModelProviderCollapsed = useUIStore((state) => state.toggleModelProviderCollapsed); - const setModelProvidersCollapsed = useUIStore((state) => state.setModelProvidersCollapsed); const addRecentModel = useUIStore((state) => state.addRecentModel); const addRecentAgent = useUIStore((state) => state.addRecentAgent); const addRecentEffort = useUIStore((state) => state.addRecentEffort); @@ -434,24 +365,12 @@ export const ModelControls: React.FC = ({ const cycleAgentShortcut = React.useMemo(() => ( getEffectiveShortcutCombo('cycle_agent', cycleAgentShortcutOverride ? { cycle_agent: cycleAgentShortcutOverride } : undefined) ), [cycleAgentShortcutOverride]); - const cycleAgentShortcutLabel = React.useMemo(() => formatShortcutForDisplay(cycleAgentShortcut), [cycleAgentShortcut]); - const collapsedProviderSet = React.useMemo(() => { - const result = new Set(); - for (const providerId of collapsedModelProviders) { - const trimmed = providerId.trim(); - if (trimmed) { - result.add(trimmed); - } - } - return result; - }, [collapsedModelProviders]); // Separate state for agent selector to avoid conflict with model selector const [isAgentSelectorOpen, setIsAgentSelectorOpen] = React.useState(false); const { favoriteModelsList, recentModelsList } = useModelLists(); - const { isMobile, isTablet } = useDeviceInfo(); - const alwaysShowHoverDetails = isMobile || isTablet; + const { isMobile } = useDeviceInfo(); const isDesktop = React.useMemo(() => isDesktopShell(), []); const isVSCodeRuntime = useIsVSCodeRuntime(); // Only use mobile panels on actual mobile devices, VSCode uses desktop dropdowns @@ -486,15 +405,12 @@ export const ModelControls: React.FC = ({ closeMobilePanel(); }, [setSelectedProvider, setSettingsPage, setSettingsDialogOpen, setAgentMenuOpen, closeMobilePanel]); const [desktopModelQuery, setDesktopModelQuery] = React.useState(''); - const [modelSelectedIndex, setModelSelectedIndex] = React.useState(0); - const modelItemRefs = React.useRef<(HTMLDivElement | null)[]>([]); const keyboardOwnsModelSelectionRef = React.useRef(false); const lastModelPointerPositionRef = React.useRef<{ x: number; y: number } | null>(null); + const activeModelPickerEntryRef = React.useRef(undefined); const [pendingThinkingVariants, setPendingThinkingVariants] = React.useState>(new Map()); const [adjustedThinkingModels, setAdjustedThinkingModels] = React.useState>(new Set()); - const favoriteRowSensors = useSensors( - useSensor(PointerSensor, { activationConstraint: { distance: 8 } }), - ); + const [modelPickerRenderVersion, setModelPickerRenderVersion] = React.useState(0); React.useEffect(() => { if (activeMobilePanel === 'model') { @@ -532,7 +448,6 @@ export const ModelControls: React.FC = ({ if (!isModelSelectorOpen) { setDesktopModelQuery(''); - setModelSelectedIndex(0); keyboardOwnsModelSelectionRef.current = false; lastModelPointerPositionRef.current = null; setPendingThinkingVariants(new Map()); @@ -653,82 +568,14 @@ export const ModelControls: React.FC = ({ ); }, [normalizeModelSearchValue]); - const getDesktopModelPickerSelectedIndex = React.useCallback((query: string) => { - const normalizedQuery = query.trim(); - const forceExpandProviders = normalizedQuery.length > 0; - const matchesQuery = (modelName: string, providerName: string) => { - if (!normalizedQuery) return true; - return matchesModelSearch(modelName, normalizedQuery) || matchesModelSearch(providerName, normalizedQuery); - }; - const providersById = new Map(providers.map((p) => [p.id, p])); - - let flatIndex = 0; - - for (const { model, providerID, modelID } of favoriteModelsList) { - const provider = providersById.get(providerID); - const providerName = provider?.name || providerID; - const modelName = getModelDisplayName(model); - if (!matchesQuery(modelName, providerName)) { - continue; - } - if (providerID === currentProviderId && modelID === currentModelId) { - return flatIndex; - } - flatIndex += 1; - } - - for (const { model, providerID, modelID } of recentModelsList) { - const provider = providersById.get(providerID); - const providerName = provider?.name || providerID; - const modelName = getModelDisplayName(model); - if (!matchesQuery(modelName, providerName)) { - continue; - } - if (providerID === currentProviderId && modelID === currentModelId) { - return flatIndex; - } - flatIndex += 1; - } - - for (const provider of visibleProviders) { - const providerId = typeof provider.id === 'string' ? provider.id : ''; - const providerName = provider.name || providerId; - const providerModels = Array.isArray(provider.models) ? (provider.models as ProviderModel[]) : []; - const filteredModels = providerModels.filter((model) => matchesQuery(getModelDisplayName(model), providerName)); - const isExpanded = forceExpandProviders || !collapsedProviderSet.has(providerId); - if (!isExpanded) { - continue; - } - for (const model of filteredModels) { - const modelId = typeof model.id === 'string' ? model.id : ''; - if (providerId === currentProviderId && modelId === currentModelId) { - return flatIndex; - } - flatIndex += 1; - } - } - - return 0; - }, [ - collapsedProviderSet, - currentModelId, - currentProviderId, - favoriteModelsList, - matchesModelSearch, - providers, - recentModelsList, - visibleProviders, - ]); - - React.useEffect(() => { - if (!isModelSelectorOpen) { - return; - } - setModelSelectedIndex(getDesktopModelPickerSelectedIndex(desktopModelQuery)); - }, [desktopModelQuery, getDesktopModelPickerSelectedIndex, isModelSelectorOpen]); - - const currentMetadata = - currentProviderId && currentModelId ? getModelMetadata(currentProviderId, currentModelId) : undefined; + const currentModelForMetadata = currentModelId + ? models.find((model: ProviderModel) => model.id === currentModelId) + : undefined; + const currentMetadata = currentProviderId && currentModelId && currentModelForMetadata + ? mergeModelMetadataWithLiveModel(currentProviderId, currentModelForMetadata, getModelMetadata(currentProviderId, currentModelId)) + : currentProviderId && currentModelId + ? getModelMetadata(currentProviderId, currentModelId) + : undefined; const localizeMetaLabel = React.useCallback((label: string) => { if (label === 'Tool calling') return t('chat.modelControls.capability.toolCalling'); if (label === 'Reasoning') return t('chat.modelControls.capability.reasoning'); @@ -1735,7 +1582,7 @@ export const ModelControls: React.FC = ({ }) => { const rowKey = buildModelRefKey(providerId, modelId); const isSelected = providerId === currentProviderId && modelId === currentModelId; - const metadata = getModelMetadata(providerId, modelId); + const metadata = mergeModelMetadataWithLiveModel(providerId, model, getModelMetadata(providerId, modelId)); const variantOptions = getModelVariantOptions(providerId, modelId); const hasVariants = variantOptions.length > 0; const resolvedVariant = resolveModelVariantSelection(providerId, modelId); @@ -2252,416 +2099,66 @@ export const ModelControls: React.FC = ({ ); - // Helper to render a single model row in the flat dropdown - const renderModelRow = ( - model: ProviderModel, - providerID: string, - modelID: string, - keyPrefix: string, - flatIndex: number, - isHighlighted: boolean, - dragHandleProps?: SortableFavoriteHandleProps | null, - ) => { - const metadata = getModelMetadata(providerID, modelID); - const capabilityIcons = getCapabilityIcons(metadata).map((icon) => ({ - ...icon, - label: localizeMetaLabel(icon.label), - id: `cap-${icon.key}`, - })); - const modalityIcons = [ - ...getModalityIcons(metadata, 'input').map((icon) => ({ ...icon, label: localizeMetaLabel(icon.label) })), - ...getModalityIcons(metadata, 'output').map((icon) => ({ ...icon, label: localizeMetaLabel(icon.label) })), - ]; - const uniqueModalityIcons = Array.from( - new Map(modalityIcons.map((icon) => [icon.key, icon])).values() - ).map((icon) => ({ ...icon, id: `mod-${icon.key}` })); - const indicatorIcons = [...capabilityIcons, ...uniqueModalityIcons]; - const contextTokens = formatTokens(metadata?.limit?.context); - const isSelected = currentProviderId === providerID && currentModelId === modelID; - const isFavorite = isFavoriteModel(providerID, modelID); - - const showProviderLogo = keyPrefix === 'fav' || keyPrefix === 'recent'; - - // Check if model supports thinking variants - variants are on the model object, not metadata - const modelVariants = (model as { variants?: Record } | undefined)?.variants; - const hasThinkingVariants = modelVariants && Object.keys(modelVariants).length > 0; - const mapKey = buildModelRefKey(providerID, modelID); - const wasAdjusted = adjustedThinkingModels.has(mapKey); - const pendingVariant = pendingThinkingVariants.get(mapKey); - const effectiveVariant = pendingVariant ?? (isSelected ? currentVariant : undefined); - - // Build thinking variant display - only show for models that were adjusted with arrow keys - let thinkingDisplay: React.ReactNode = null; - if (hasThinkingVariants && wasAdjusted && (isHighlighted || isSelected)) { - const displayLabel = effectiveVariant - ? effectiveVariant.charAt(0).toUpperCase() + effectiveVariant.slice(1) - : 'Default'; - thinkingDisplay = ( - - Thinking: {displayLabel} - - ); - } - - // Build animated metadata slides for desktop (price/capabilities) - only shown when not showing thinking - const priceText = formatCompactPrice(metadata); - const hasPrice = priceText !== null; - const hasCapabilities = indicatorIcons.length > 0; - - const slides: React.ReactNode[] = []; - if (hasPrice) { - slides.push( - - {priceText} - - ); - } - if (hasCapabilities) { - slides.push( -
- {indicatorIcons.map(({ id, icon: iconName, label }) => ( - - - - ))} -
- ); - } - - const supportsRotatingMetadata = !isVSCodeRuntime; - const shouldShowThinking = hasThinkingVariants && wasAdjusted; - const shouldAnimate = supportsRotatingMetadata && slides.length > 1 && (isHighlighted || isSelected) && !shouldShowThinking; - const staticSlideIndex = !supportsRotatingMetadata && hasCapabilities && hasPrice ? 1 : 0; - const staticMetadataSlide = slides[staticSlideIndex]; - - const handlePointerActivity = (event: React.MouseEvent) => { - const nextPosition = { x: event.clientX, y: event.clientY }; - const previousPosition = lastModelPointerPositionRef.current; - const pointerMoved = !previousPosition - || previousPosition.x !== nextPosition.x - || previousPosition.y !== nextPosition.y; - - lastModelPointerPositionRef.current = nextPosition; - - if (keyboardOwnsModelSelectionRef.current && !pointerMoved) { - return; - } - - if (keyboardOwnsModelSelectionRef.current && pointerMoved) { - keyboardOwnsModelSelectionRef.current = false; - } - - setModelSelectedIndex(flatIndex); - }; - - return ( -
{ modelItemRefs.current[flatIndex] = el; }} - role="option" - aria-selected={isSelected} - className={cn( - "typography-meta group flex items-center gap-2 px-2 py-1.5 rounded-md cursor-pointer", - isHighlighted ? "bg-interactive-selection" : "hover:bg-interactive-hover/50" - )} - onClick={() => handleProviderAndModelChange(providerID, modelID)} - onKeyDown={(e) => { - if (e.key === 'Enter' || e.key === ' ') { - e.preventDefault(); - handleProviderAndModelChange(providerID, modelID); - } - }} - onMouseEnter={handlePointerActivity} - onMouseMove={handlePointerActivity} - > - {dragHandleProps ? ( - - ) : null} -
- {showProviderLogo && ( - - )} - - {getModelDisplayName(model)} - - {metadata?.limit?.context ? ( - - {contextTokens} - - ) : null} -
-
- {/* Metadata slot: thinking variant for adjusted models, otherwise price/capabilities carousel */} - {shouldShowThinking && (isHighlighted || isSelected) ? ( -
- {thinkingDisplay} -
- ) : slides.length > 0 ? ( -
- {shouldAnimate ? ( - - {slides} - - ) : ( - <> - {/* In static runtimes (VS Code), prefer capabilities over price when both exist. */} - {staticMetadataSlide} - - )} -
- ) : null} - {isSelected && ( - - )} - -
-
- ); - }; - - type FlatModelItem = { model: ProviderModel; providerID: string; modelID: string; section: string }; - - const modelSelectorData = React.useMemo(() => { - const filterByQuery = (modelName: string, providerName: string, query: string) => { - if (!query.trim()) return true; - return ( - matchesModelSearch(modelName, query) || - matchesModelSearch(providerName, query) - ); - }; - - const normalizedDesktopQuery = desktopModelQuery.trim(); - const forceExpandProviders = normalizedDesktopQuery.length > 0; - - const filteredFavorites = favoriteModelsList.filter(({ model, providerID }) => { - const provider = providers.find(p => p.id === providerID); - const providerName = provider?.name || providerID; - const modelName = getModelDisplayName(model); - return filterByQuery(modelName, providerName, desktopModelQuery); - }); - const favoriteSortingEnabled = normalizedDesktopQuery.length === 0 && filteredFavorites.length > 1; - - const filteredRecents = recentModelsList.filter(({ model, providerID }) => { - const provider = providers.find(p => p.id === providerID); - const providerName = provider?.name || providerID; - const modelName = getModelDisplayName(model); - return filterByQuery(modelName, providerName, desktopModelQuery); - }); - - const filteredProviders: typeof visibleProviders = []; - for (const provider of visibleProviders) { - const providerModels = Array.isArray(provider.models) ? provider.models : []; - const filteredModels = providerModels.filter((model: ProviderModel) => { - const modelName = getModelDisplayName(model); - return filterByQuery(modelName, provider.name || provider.id || '', desktopModelQuery); - }); - if (filteredModels.length > 0) { - filteredProviders.push({ ...provider, models: filteredModels }); - } - } - - const providerSections = filteredProviders.map((provider) => { - const providerId = typeof provider.id === 'string' ? provider.id : ''; - const isExpanded = forceExpandProviders || !collapsedProviderSet.has(providerId); - const models = Array.isArray(provider.models) ? (provider.models as ProviderModel[]) : []; - return { - provider, - isExpanded, - models, - visibleModels: isExpanded ? models : [], - }; - }); - - const hasResults = - filteredFavorites.length > 0 || - filteredRecents.length > 0 || - filteredProviders.length > 0; - - const filteredProviderIds: string[] = []; - for (const provider of filteredProviders) { - if (typeof provider.id === 'string' && provider.id) { - filteredProviderIds.push(provider.id); - } - } - - const favoriteModelLookup = new Map( - filteredFavorites.map(({ providerID, modelID }) => [buildModelRefKey(providerID, modelID), { providerID, modelID }]) - ); - const flatModelList: FlatModelItem[] = []; - - filteredFavorites.forEach(({ model, providerID, modelID }) => { - flatModelList.push({ model, providerID, modelID, section: 'fav' }); - }); - filteredRecents.forEach(({ model, providerID, modelID }) => { - flatModelList.push({ model, providerID, modelID, section: 'recent' }); - }); - providerSections.forEach(({ provider, visibleModels }) => { - visibleModels.forEach((model) => { - flatModelList.push({ model, providerID: provider.id as string, modelID: model.id as string, section: 'provider' }); - }); - }); - - return { - filteredFavorites, - filteredRecents, - filteredProviders, - providerSections, - flatModelList, - hasResults, - forceExpandProviders, - favoriteSortingEnabled, - filteredProviderIds, - favoriteModelLookup, - }; - }, [desktopModelQuery, favoriteModelsList, recentModelsList, visibleProviders, providers, collapsedProviderSet, matchesModelSearch]); - const renderModelSelector = () => { - const { - filteredFavorites, - filteredRecents, - filteredProviders, - providerSections, - flatModelList, - hasResults, - forceExpandProviders, - favoriteSortingEnabled, - filteredProviderIds, - favoriteModelLookup, - } = modelSelectorData; - - const totalItems = flatModelList.length; - - // Check if currently highlighted model supports thinking variants - const highlightedItem = flatModelList[modelSelectedIndex]; - const highlightedSupportsThinking = highlightedItem ? (() => { - const modelVariants = (highlightedItem.model as { variants?: Record } | undefined)?.variants; - return modelVariants && Object.keys(modelVariants).length > 0; - })() : false; - - // Handle keyboard navigation - const handleModelKeyDown = (e: React.KeyboardEvent) => { - e.stopPropagation(); + const handleThinkingVariantKey = (e: React.KeyboardEvent, selectedItem: ModelPickerEntry) => { keyboardOwnsModelSelectionRef.current = true; + if (e.key !== 'ArrowLeft' && e.key !== 'ArrowRight') return false; + const { providerID, modelID } = selectedItem; + const canonicalProvider = useConfigStore.getState().providers.find((provider) => provider.id === providerID); + const canonicalModel = canonicalProvider?.models.find((model) => model.id === modelID) as { variants?: Record } | undefined; + const variantKeys = canonicalModel?.variants ? Object.keys(canonicalModel.variants) : getModelVariantOptions(providerID, modelID); + if (variantKeys.length === 0) return false; + + e.preventDefault(); + e.stopPropagation(); + + const mapKey = buildModelRefKey(providerID, modelID); + const hasPendingVariant = pendingThinkingVariants.has(mapKey); + const currentPending = pendingThinkingVariants.get(mapKey); + const activeModelVariant = hasPendingVariant ? currentPending : (currentProviderId === providerID && currentModelId === modelID ? currentVariant : undefined); + + const variantsWithDefault: Array = [undefined, ...variantKeys]; + const currentVariantIndex = variantsWithDefault.indexOf(activeModelVariant); + const safeCurrentIndex = currentVariantIndex >= 0 ? currentVariantIndex : 0; + const direction = e.key === 'ArrowRight' ? 1 : -1; + const nextVariantIndex = (safeCurrentIndex + direction + variantsWithDefault.length) % variantsWithDefault.length; + const nextVariant = variantsWithDefault[nextVariantIndex]; + + setPendingThinkingVariants((prev) => { + const next = new Map(prev); + next.set(mapKey, nextVariant); + return next; + }); + setAdjustedThinkingModels((prev) => { + const next = new Set(prev); + next.add(mapKey); + return next; + }); + setModelPickerRenderVersion((version) => version + 1); + return true; + }; + + const handleModelPickerKeyDown = (e: React.KeyboardEvent, selectedItem: ModelPickerEntry | undefined) => { const cycleAgentDirection = getCycleAgentDirectionFromEvent(e); if (cycleAgentDirection) { e.preventDefault(); handleCycleAgentFromModelPicker(cycleAgentDirection); - } else if (e.key === 'ArrowDown') { - e.preventDefault(); - setModelSelectedIndex((prev) => (prev + 1) % Math.max(1, totalItems)); - // Scroll into view - setTimeout(() => { - const nextIndex = (modelSelectedIndex + 1) % Math.max(1, totalItems); - modelItemRefs.current[nextIndex]?.scrollIntoView({ behavior: 'smooth', block: 'nearest' }); - }, 0); - } else if (e.key === 'ArrowUp') { - e.preventDefault(); - setModelSelectedIndex((prev) => (prev - 1 + Math.max(1, totalItems)) % Math.max(1, totalItems)); - // Scroll into view - setTimeout(() => { - const prevIndex = (modelSelectedIndex - 1 + Math.max(1, totalItems)) % Math.max(1, totalItems); - modelItemRefs.current[prevIndex]?.scrollIntoView({ behavior: 'smooth', block: 'nearest' }); - }, 0); - } else if (e.key === 'ArrowLeft' || e.key === 'ArrowRight') { - e.preventDefault(); - const selectedItem = flatModelList[modelSelectedIndex]; - if (!selectedItem) return; - - const { providerID, modelID, model } = selectedItem; - const modelVariants = (model as { variants?: Record } | undefined)?.variants; - if (!modelVariants) return; - - const variantKeys = Object.keys(modelVariants); - if (variantKeys.length === 0) return; - - const mapKey = buildModelRefKey(providerID, modelID); - const currentPending = pendingThinkingVariants.get(mapKey); - const activeModelVariant = currentPending ?? (currentProviderId === providerID && currentModelId === modelID ? currentVariant : undefined); - - const variantsWithDefault: Array = [undefined, ...variantKeys]; - const currentVariantIndex = variantsWithDefault.indexOf(activeModelVariant); - const safeCurrentIndex = currentVariantIndex >= 0 ? currentVariantIndex : 0; - const direction = e.key === 'ArrowRight' ? 1 : -1; - const nextVariantIndex = Math.min( - variantsWithDefault.length - 1, - Math.max(0, safeCurrentIndex + direction), - ); - const nextVariant = variantsWithDefault[nextVariantIndex]; - - setPendingThinkingVariants((prev) => { - const next = new Map(prev); - next.set(mapKey, nextVariant); - return next; - }); - setAdjustedThinkingModels((prev) => { - const next = new Set(prev); - next.add(mapKey); - return next; - }); - } else if (e.key === 'Enter') { - e.preventDefault(); - const selectedItem = flatModelList[modelSelectedIndex]; - if (selectedItem) { - const { providerID, modelID } = selectedItem; - const mapKey = buildModelRefKey(providerID, modelID); - const pendingVariant = pendingThinkingVariants.get(mapKey); - const wasAdjusted = adjustedThinkingModels.has(mapKey); - const effectiveAgentName = resolveLiveAgentName(); - - handleProviderAndModelChange(providerID, modelID, wasAdjusted - ? { applyVariant: true, variant: pendingVariant, agentName: effectiveAgentName } - : { agentName: effectiveAgentName }); - } - } else if (e.key === 'Escape') { - e.preventDefault(); - setAgentMenuOpen(false); + return; } + + if (selectedItem) handleThinkingVariantKey(e, selectedItem); + }; + + const handleSharedModelSelect = (entry: ModelPickerEntry) => { + const mapKey = buildModelRefKey(entry.providerID, entry.modelID); + const pendingVariant = pendingThinkingVariants.get(mapKey); + const wasAdjusted = adjustedThinkingModels.has(mapKey); + const effectiveAgentName = resolveLiveAgentName(); + + handleProviderAndModelChange(entry.providerID, entry.modelID, wasAdjusted + ? { applyVariant: true, variant: pendingVariant, agentName: effectiveAgentName } + : { agentName: effectiveAgentName }); }; const handleModelShortcutKeyDownCapture = (e: React.KeyboardEvent) => { @@ -2676,40 +2173,45 @@ export const ModelControls: React.FC = ({ handleCycleAgentFromModelPicker(cycleAgentDirection); }; - const handleFavoriteDragEnd = (event: DragEndEvent) => { - const { active, over } = event; - if (!over || active.id === over.id) { - return; - } - - const activeFavorite = favoriteModelLookup.get(String(active.id)); - const overFavorite = favoriteModelLookup.get(String(over.id)); - if (!activeFavorite || !overFavorite) { - return; - } - - reorderFavoriteModel( - activeFavorite.providerID, - activeFavorite.modelID, - overFavorite.providerID, - overFavorite.modelID, - ); - }; - - const handleProviderSectionToggle = (expand: boolean) => { - if (filteredProviderIds.length === 0) { - return; - } - setModelProvidersCollapsed(filteredProviderIds, !expand); - setModelSelectedIndex(0); - }; - const handleModelMenuOpenChange = (nextOpen: boolean) => { setAgentMenuOpen(nextOpen); }; - // Build index mapping for rendering - let currentFlatIndex = 0; + const modelPickerLabels = { + searchPlaceholder: t('chat.modelControls.searchModels'), + noResults: t('chat.modelControls.noModelsFound'), + favorites: t('chat.modelControls.favorites'), + recent: t('chat.modelControls.recent'), + keyboardHint: t('chat.modelControls.keyboardHintNavigate'), + favorite: t('chat.modelControls.favoriteAria'), + unfavorite: t('chat.modelControls.unfavoriteAria'), + capabilities: t('chat.modelControls.capabilities'), + capabilityToolCalling: t('chat.modelControls.capability.toolCalling'), + capabilityReasoning: t('chat.modelControls.capability.reasoning'), + input: t('chat.modelControls.input'), + output: t('chat.modelControls.output'), + costPerMillion: t('chat.modelControls.costPerMillion'), + }; + + const renderThinkingSlot = (entry: ModelPickerEntry, { isHighlighted, isSelected }: { isHighlighted: boolean; isSelected: boolean }) => { + const hasThinkingVariants = getModelVariantOptions(entry.providerID, entry.modelID).length > 0; + const mapKey = buildModelRefKey(entry.providerID, entry.modelID); + const wasAdjusted = adjustedThinkingModels.has(mapKey); + if (!hasThinkingVariants || (!isHighlighted && !isSelected)) return null; + + const hasPendingVariant = pendingThinkingVariants.has(mapKey); + const pendingVariant = pendingThinkingVariants.get(mapKey); + const effectiveVariant = hasPendingVariant ? pendingVariant : (isSelected ? currentVariant : undefined); + const displayLabel = effectiveVariant + ? effectiveVariant.charAt(0).toUpperCase() + effectiveVariant.slice(1) + : 'Default'; + + return ( + + Thinking: {displayLabel} + + ); + }; return ( @@ -2770,201 +2272,60 @@ export const ModelControls: React.FC = ({ alignOffset={-40} onKeyDownCapture={handleModelShortcutKeyDownCapture} > - {/* Search Input */} -
-
- - setDesktopModelQuery(e.target.value)} - onKeyDown={handleModelKeyDown} - className="pl-8 h-8 typography-meta" - /> -
-
- - {/* Scrollable content */} - -
-
{ - if (e.key === 'Enter' || e.key === ' ') { - e.preventDefault(); - openAddProviderSettings(); - } - }} - className="typography-meta group flex items-center gap-1 rounded-md px-2 py-1.5 cursor-pointer hover:bg-interactive-hover/50" - > - - - - {t('chat.modelControls.addNewProvider')} -
- - - - {!hasResults && ( -
- {t('chat.modelControls.noModelsFound')} -
- )} - - {/* Favorites Section */} - {filteredFavorites.length > 0 && ( -
- - - {t('chat.modelControls.favorites')} - - {favoriteSortingEnabled ? ( - - buildModelRefKey(providerID, modelID))} - strategy={verticalListSortingStrategy} - > - {filteredFavorites.map(({ model, providerID, modelID }) => { - const idx = currentFlatIndex++; - return ( - - {(dragHandleProps) => renderModelRow( - model, - providerID, - modelID, - 'fav', - idx, - modelSelectedIndex === idx, - dragHandleProps, - )} - - ); - })} - - - ) : ( - filteredFavorites.map(({ model, providerID, modelID }) => { - const idx = currentFlatIndex++; - return renderModelRow(model, providerID, modelID, 'fav', idx, modelSelectedIndex === idx); - }) - )} -
- )} - - {/* Recents Section */} - {filteredRecents.length > 0 && ( -
- {filteredFavorites.length > 0 && } - - - {t('chat.modelControls.recent')} - - {filteredRecents.map(({ model, providerID, modelID }) => { - const idx = currentFlatIndex++; - return renderModelRow(model, providerID, modelID, 'recent', idx, modelSelectedIndex === idx); - })} -
- )} - - {/* Separator before providers */} - {(filteredFavorites.length > 0 || filteredRecents.length > 0) && filteredProviders.length > 0 && ( - - )} - - {/* All Providers - Flat List */} - {providerSections.map(({ provider, isExpanded, visibleModels }, index) => ( -
- {index > 0 && } -
{ - if (forceExpandProviders) { - return; - } - - if (event.metaKey || event.ctrlKey) { - handleProviderSectionToggle(!isExpanded); - return; - } - - toggleModelProviderCollapsed(String(provider.id)); - setModelSelectedIndex(0); - }} - onKeyDown={(event) => { - if (forceExpandProviders) { - return; - } - if (event.key === 'Enter' || event.key === ' ') { - event.preventDefault(); - toggleModelProviderCollapsed(String(provider.id)); - setModelSelectedIndex(0); - } - }} - className={cn( - 'typography-micro font-semibold text-muted-foreground uppercase tracking-wider flex w-full items-center gap-2 -mx-1 px-3 py-1.5 border-b border-border/30', - 'text-left transition-colors', - forceExpandProviders ? 'cursor-default' : 'cursor-pointer' - )} - aria-expanded={isExpanded} - title={forceExpandProviders - ? undefined - : (isExpanded - ? t('chat.modelControls.collapseProvider') - : t('chat.modelControls.expandProvider'))} - > -
- - {provider.name} - - {isExpanded ? ( - - ) : ( - - )} - -
-
- {isExpanded && visibleModels.map((model: ProviderModel) => { - const idx = currentFlatIndex++; - return renderModelRow(model, provider.id as string, model.id as string, 'provider', idx, modelSelectedIndex === idx); - })} -
- ))} -
-
- - {/* Keyboard hints footer */} -
-
- {t('chat.modelControls.keyboardHintNavigate')} - {t('chat.modelControls.keyboardHintSwitchAgent', { shortcut: cycleAgentShortcutLabel })} - - {t('chat.modelControls.keyboardHintThinking')} +
+
+ {t('chat.modelControls.addNewProvider')} +
+ { activeModelPickerEntryRef.current = entry; }} + onVariantKey={handleThinkingVariantKey} + isFavorite={(entry) => isFavoriteModel(entry.providerID, entry.modelID)} + onToggleFavorite={(entry) => toggleFavoriteModel(entry.providerID, entry.modelID)} + renderRowEnd={renderThinkingSlot} + renderVersion={modelPickerRenderVersion} + onReorderFavorite={(active, over) => reorderFavoriteModel( + active.providerID, + active.modelID, + over.providerID, + over.modelID, + )} + reorderFavoriteAriaLabel={t('chat.modelControls.reorderFavoriteAria')} + reorderFavoriteTitle={t('chat.modelControls.reorderFavoriteTitle')} + footerContent={(activeEntry) => { + const activeHasThinkingVariants = activeEntry + ? getModelVariantOptions(activeEntry.providerID, activeEntry.modelID).length > 0 + : false; + + return ( +
+ {t('chat.modelControls.keyboardHintNavigate')} + {t('chat.modelControls.keyboardHintSwitchAgent', { shortcut: 'Tab' })} + {activeHasThinkingVariants ? {t('chat.modelControls.keyboardHintThinking')} : null} +
+ ); + }} + tooltipsEnabled={agentMenuOpen} + onEscape={() => setAgentMenuOpen(false)} + /> ) : ( diff --git a/packages/ui/src/components/model-picker/ModelPickerList.tsx b/packages/ui/src/components/model-picker/ModelPickerList.tsx new file mode 100644 index 00000000..555a9eb6 --- /dev/null +++ b/packages/ui/src/components/model-picker/ModelPickerList.tsx @@ -0,0 +1,708 @@ +import React from 'react'; +import { + DndContext, + PointerSensor, + closestCenter, + useSensor, + useSensors, + type DragEndEvent, +} from '@dnd-kit/core'; +import { SortableContext, useSortable, verticalListSortingStrategy } from '@dnd-kit/sortable'; +import { CSS as DndCSS } from '@dnd-kit/utilities'; +import { Icon } from '@/components/icon/Icon'; +import { Input } from '@/components/ui/input'; +import { ProviderLogo } from '@/components/ui/ProviderLogo'; +import { ScrollableOverlay } from '@/components/ui/ScrollableOverlay'; +import { Tooltip, TooltipContent, TooltipTrigger } from '@/components/ui/tooltip'; +import { mergeModelMetadataWithLiveModel } from '@/lib/modelMetadata'; +import { cn } from '@/lib/utils'; +import type { ModelMetadata } from '@/types'; + +export type ProviderModel = Record & { id?: string; name?: string }; + +export type ModelPickerProvider = { + id: string; + name?: string; + models?: ProviderModel[]; +}; + +export type ModelPickerEntry = { + model: ProviderModel; + providerID: string; + modelID: string; +}; + +export type ModelPickerFavoriteEntry = ModelPickerEntry; + +type HiddenModel = { providerID: string; modelID: string }; + +type IndexSelectionStore = { + getSnapshot: () => number; + subscribe: (listener: () => void) => () => void; + subscribeIndex: (index: number, listener: () => void) => () => void; + set: (value: number) => void; +}; + +const COMPACT_NUMBER_FORMATTER = new Intl.NumberFormat('en-US', { + notation: 'compact', + compactDisplay: 'short', + maximumFractionDigits: 1, + minimumFractionDigits: 0, +}); + +const CURRENCY_FORMATTER = new Intl.NumberFormat('en-US', { + style: 'currency', + currency: 'USD', + maximumFractionDigits: 4, + minimumFractionDigits: 2, +}); + +const getModelDisplayName = (model: Record) => { + const name = model?.name || model?.id || ''; + const nameStr = String(name); + if (nameStr.length > 40) return `${nameStr.substring(0, 37)}...`; + return nameStr; +}; + +const formatModelContextTokens = (value?: number | null) => { + if (typeof value !== 'number' || Number.isNaN(value)) return ''; + if (value === 0) return '0'; + const formatted = COMPACT_NUMBER_FORMATTER.format(value); + return formatted.endsWith('.0') ? formatted.slice(0, -2) : formatted; +}; + +const formatCost = (value?: number | null) => { + if (typeof value !== 'number' || !Number.isFinite(value)) return '—'; + return CURRENCY_FORMATTER.format(value); +}; + +const hasTooltipMetadata = (metadata?: ModelMetadata) => { + if (!metadata) return false; + return Boolean( + metadata.tool_call || + metadata.reasoning || + metadata.cost?.input !== undefined || + metadata.cost?.output !== undefined || + (metadata.modalities?.input?.length ?? 0) > 0 || + (metadata.modalities?.output?.length ?? 0) > 0, + ); +}; + +const ModelPickerRowTooltip: React.FC<{ + metadata?: ModelMetadata; + active: boolean; + labels: ModelPickerListProps['labels']; + children: React.ReactElement; +}> = ({ metadata, active, labels, children }) => { + const [delayedActive, setDelayedActive] = React.useState(false); + + React.useEffect(() => { + if (!active) { + setDelayedActive(false); + return; + } + const timeout = window.setTimeout(() => setDelayedActive(true), 450); + return () => window.clearTimeout(timeout); + }, [active]); + + if (!hasTooltipMetadata(metadata)) return children; + + const inputModalities = metadata?.modalities?.input ?? []; + const outputModalities = metadata?.modalities?.output ?? []; + const capabilities = [ + metadata?.tool_call ? labels.capabilityToolCalling : null, + metadata?.reasoning ? labels.capabilityReasoning : null, + ].filter(Boolean); + + return ( + {}}> + {children} + {active && delayedActive ? ( + +
+ {capabilities.length > 0 ? ( +
+ {labels.capabilities} + {capabilities.join(', ')} +
+ ) : null} + {inputModalities.length > 0 ? ( +
+ {labels.input} + {inputModalities.join(', ')} +
+ ) : null} + {outputModalities.length > 0 ? ( +
+ {labels.output} + {outputModalities.join(', ')} +
+ ) : null} + {(metadata?.cost?.input !== undefined || metadata?.cost?.output !== undefined) ? ( +
+ {labels.costPerMillion} + In {formatCost(metadata?.cost?.input)} · Out {formatCost(metadata?.cost?.output)} +
+ ) : null} +
+
+ ) : null} +
+ ); +}; + +const createIndexSelectionStore = (): IndexSelectionStore => { + let value = 0; + const listeners = new Set<() => void>(); + const listenersByIndex = new Map void>>(); + const notify = (index: number) => { + const listeners = listenersByIndex.get(index); + if (!listeners) return; + for (const listener of listeners) listener(); + }; + + return { + getSnapshot: () => value, + subscribe: (listener) => { + listeners.add(listener); + return () => listeners.delete(listener); + }, + subscribeIndex: (index, listener) => { + let listeners = listenersByIndex.get(index); + if (!listeners) { + listeners = new Set(); + listenersByIndex.set(index, listeners); + } + listeners.add(listener); + return () => { + listeners.delete(listener); + if (listeners.size === 0) listenersByIndex.delete(index); + }; + }, + set: (nextValue) => { + if (value === nextValue) return; + const previousValue = value; + value = nextValue; + notify(previousValue); + notify(nextValue); + for (const listener of listeners) listener(); + }, + }; +}; + +const ModelPickerRowHighlight: React.FC<{ + store: IndexSelectionStore; + index: number; + renderVersion?: number; + children: (isHighlighted: boolean) => React.ReactNode; +}> = React.memo(({ store, index, children }) => { + const [isHighlighted, setIsHighlighted] = React.useState(() => store.getSnapshot() === index); + + React.useEffect(() => { + const sync = () => setIsHighlighted(store.getSnapshot() === index); + sync(); + return store.subscribeIndex(index, sync); + }, [index, store]); + + return <>{children(isHighlighted)}; +}); + +const ModelPickerFooter: React.FC<{ + store: IndexSelectionStore; + flatModelList: ModelPickerEntry[]; + footerContent: ModelPickerListProps['footerContent']; + fallback: React.ReactNode; +}> = ({ store, flatModelList, footerContent, fallback }) => { + const [selectedIndex, setSelectedIndex] = React.useState(() => store.getSnapshot()); + + React.useEffect(() => store.subscribe(() => setSelectedIndex(store.getSnapshot())), [store]); + + const activeEntry = flatModelList[selectedIndex]; + return <>{typeof footerContent === 'function' ? footerContent(activeEntry) : (footerContent ?? fallback)}; +}; + +type SortableFavoriteHandleProps = { + attributes: ReturnType['attributes']; + listeners: ReturnType['listeners']; + setActivatorNodeRef: ReturnType['setActivatorNodeRef']; + isDragging: boolean; +}; + +const SortableFavoriteModelRow: React.FC<{ + id: string; + disabled?: boolean; + children: (dragHandleProps: SortableFavoriteHandleProps) => React.ReactNode; +}> = ({ id, disabled = false, children }) => { + const { + attributes, + listeners, + setNodeRef, + setActivatorNodeRef, + transform, + transition, + isDragging, + } = useSortable({ id, disabled }); + + return ( +
+ {children({ attributes, listeners, setActivatorNodeRef, isDragging })} +
+ ); +}; + +const STICKY_HEADER_OFFSET = 32; + +const scrollIntoView = (container: HTMLElement | null, node: HTMLElement | null) => { + if (!node) return; + if (!container) { + node.scrollIntoView({ block: 'nearest' }); + return; + } + + const containerRect = container.getBoundingClientRect(); + const nodeRect = node.getBoundingClientRect(); + const top = nodeRect.top - containerRect.top + container.scrollTop; + const bottom = top + nodeRect.height; + const viewTop = container.scrollTop; + const viewBottom = viewTop + container.clientHeight; + const viewTopWithHeader = viewTop + STICKY_HEADER_OFFSET; + const target = top < viewTopWithHeader + ? top - STICKY_HEADER_OFFSET + : bottom > viewBottom + ? bottom - container.clientHeight + : viewTop; + const max = Math.max(0, container.scrollHeight - container.clientHeight); + container.scrollTop = Math.max(0, Math.min(target, max)); +}; + +interface ModelPickerListProps { + providers: ModelPickerProvider[]; + favoriteModels: ModelPickerFavoriteEntry[]; + recentModels: ModelPickerFavoriteEntry[]; + modelsMetadata: Map; + searchQuery: string; + onSearchQueryChange: (value: string) => void; + onSelect: (entry: ModelPickerEntry) => void; + labels: { + searchPlaceholder: string; + noResults: string; + favorites: string; + recent: string; + keyboardHint: string; + notSelected?: string; + favorite?: string; + unfavorite?: string; + capabilities?: string; + capabilityToolCalling?: string; + capabilityReasoning?: string; + input?: string; + output?: string; + costPerMillion?: string; + }; + selectedModel?: { providerID: string; modelID: string } | null; + hiddenModels?: HiddenModel[]; + allowedProviderIds?: string[]; + includeNotSelected?: boolean; + onSelectNone?: () => void; + selectionCount?: (entry: ModelPickerEntry) => number; + disabled?: boolean; + maxHeightClassName?: string; + maxHeightStyle?: React.CSSProperties; + sectionHeaderClassName?: string; + rowClassName?: string; + stickyHeaders?: boolean; + autoFocus?: boolean; + onEscape?: () => void; + isFavorite?: (entry: ModelPickerEntry) => boolean; + onToggleFavorite?: (entry: ModelPickerEntry) => void; + renderRowEnd?: (entry: ModelPickerEntry, state: { isHighlighted: boolean; isSelected: boolean }) => React.ReactNode; + onActiveKeyDown?: (event: React.KeyboardEvent, entry: ModelPickerEntry | undefined) => void; + onActiveEntryChange?: (entry: ModelPickerEntry | undefined) => void; + onVariantKey?: (event: React.KeyboardEvent, entry: ModelPickerEntry) => boolean; + onReorderFavorite?: (active: ModelPickerEntry, over: ModelPickerEntry) => void; + reorderFavoriteAriaLabel?: string; + reorderFavoriteTitle?: string; + footerContent?: React.ReactNode | ((activeEntry: ModelPickerEntry | undefined) => React.ReactNode); + renderVersion?: number; + tooltipsEnabled?: boolean; +} + +export const ModelPickerList: React.FC = ({ + providers, + favoriteModels, + recentModels, + modelsMetadata, + searchQuery, + onSearchQueryChange, + onSelect, + labels, + selectedModel, + hiddenModels = [], + allowedProviderIds, + includeNotSelected = false, + onSelectNone, + selectionCount, + disabled = false, + maxHeightClassName = 'max-h-[min(400px,calc(100dvh-12rem))] flex-1', + maxHeightStyle, + sectionHeaderClassName, + rowClassName, + stickyHeaders = true, + autoFocus = true, + onEscape, + isFavorite, + onToggleFavorite, + renderRowEnd, + onActiveKeyDown, + onActiveEntryChange, + onVariantKey, + onReorderFavorite, + reorderFavoriteAriaLabel, + reorderFavoriteTitle, + footerContent, + renderVersion, + tooltipsEnabled = true, +}) => { + const selectionStoreRef = React.useRef(null); + if (!selectionStoreRef.current) selectionStoreRef.current = createIndexSelectionStore(); + const selectionStore = selectionStoreRef.current; + const itemRefs = React.useRef<(HTMLDivElement | null)[]>([]); + const scrollRef = React.useRef(null); + const keyboardOwnsSelectionRef = React.useRef(false); + const lastMousePositionRef = React.useRef<{ x: number; y: number } | null>(null); + const [collapsedSections, setCollapsedSections] = React.useState>(() => new Set()); + const favoriteRowSensors = useSensors( + useSensor(PointerSensor, { activationConstraint: { distance: 8 } }), + ); + + const allowedProviderSet = React.useMemo(() => { + if (!allowedProviderIds || allowedProviderIds.length === 0) return null; + return new Set(allowedProviderIds); + }, [allowedProviderIds]); + + const providerById = React.useMemo(() => new Map(providers.map((provider) => [provider.id, provider])), [providers]); + + const isHidden = React.useCallback((providerID: string, modelID: string) => { + return hiddenModels.some((hidden) => hidden.providerID === providerID && hidden.modelID === modelID); + }, [hiddenModels]); + + const matchesQuery = React.useCallback((modelName: string, providerName: string) => { + const query = searchQuery.trim().toLowerCase(); + if (!query) return true; + return modelName.toLowerCase().includes(query) || providerName.toLowerCase().includes(query); + }, [searchQuery]); + + const filteredFavorites = React.useMemo(() => favoriteModels.filter(({ model, providerID, modelID }) => { + if (allowedProviderSet && !allowedProviderSet.has(providerID)) return false; + if (isHidden(providerID, modelID)) return false; + const providerName = providerById.get(providerID)?.name || providerID; + return matchesQuery(getModelDisplayName(model), providerName); + }), [allowedProviderSet, favoriteModels, isHidden, matchesQuery, providerById]); + + const filteredRecents = React.useMemo(() => recentModels.filter(({ model, providerID, modelID }) => { + if (allowedProviderSet && !allowedProviderSet.has(providerID)) return false; + if (isHidden(providerID, modelID)) return false; + const providerName = providerById.get(providerID)?.name || providerID; + return matchesQuery(getModelDisplayName(model), providerName); + }), [allowedProviderSet, isHidden, matchesQuery, providerById, recentModels]); + + const filteredProviders = React.useMemo(() => providers + .filter((provider) => !allowedProviderSet || allowedProviderSet.has(provider.id)) + .map((provider) => { + const models = Array.isArray(provider.models) ? provider.models : []; + const filteredModels = models.filter((model) => { + const modelID = typeof model.id === 'string' ? model.id : ''; + if (!modelID || isHidden(provider.id, modelID)) return false; + return matchesQuery(getModelDisplayName(model), provider.name || provider.id); + }); + return { ...provider, models: filteredModels }; + }) + .filter((provider) => provider.models.length > 0), [allowedProviderSet, isHidden, matchesQuery, providers]); + + const flatModelList = React.useMemo(() => { + const items: ModelPickerEntry[] = []; + if (!collapsedSections.has('favorites')) filteredFavorites.forEach((entry) => items.push(entry)); + if (!collapsedSections.has('recent')) filteredRecents.forEach((entry) => items.push(entry)); + filteredProviders.forEach((provider) => { + if (collapsedSections.has(`provider:${provider.id}`)) return; + provider.models.forEach((model) => items.push({ model, providerID: provider.id, modelID: model.id as string })); + }); + return items; + }, [collapsedSections, filteredFavorites, filteredProviders, filteredRecents]); + + const hasResults = flatModelList.length > 0; + const favoriteSortingEnabled = Boolean(onReorderFavorite) && searchQuery.trim().length === 0 && filteredFavorites.length > 1; + const favoriteLookup: Map = React.useMemo(() => new Map( + filteredFavorites.map((entry) => [`${entry.providerID}:${entry.modelID}`, entry] as const), + ), [filteredFavorites]); + + React.useEffect(() => { + selectionStore.set(0); + }, [searchQuery, selectionStore]); + + const selectIndex = React.useCallback((index: number) => { + selectionStore.set(index); + onActiveEntryChange?.(flatModelList[index]); + }, [flatModelList, onActiveEntryChange, selectionStore]); + + const moveSelection = React.useCallback((direction: 1 | -1) => { + const total = flatModelList.length; + if (total === 0) return; + keyboardOwnsSelectionRef.current = true; + lastMousePositionRef.current = null; + const currentIndex = selectionStore.getSnapshot(); + const nextIndex = (currentIndex + direction + total) % total; + selectionStore.set(nextIndex); + onActiveEntryChange?.(flatModelList[nextIndex]); + requestAnimationFrame(() => scrollIntoView(scrollRef.current, itemRefs.current[nextIndex])); + }, [flatModelList, onActiveEntryChange, selectionStore]); + + React.useEffect(() => { + onActiveEntryChange?.(flatModelList[selectionStore.getSnapshot()]); + }, [flatModelList, onActiveEntryChange, selectionStore]); + + const handleKeyDown = React.useCallback((event: React.KeyboardEvent) => { + if (event.defaultPrevented) return; + event.stopPropagation(); + if ((event.key === 'ArrowLeft' || event.key === 'ArrowRight')) { + const selected = flatModelList[selectionStore.getSnapshot()]; + if (selected && onVariantKey?.(event, selected)) return; + } + onActiveKeyDown?.(event, flatModelList[selectionStore.getSnapshot()]); + if (event.defaultPrevented) return; + if (event.key === 'ArrowDown') { + event.preventDefault(); + moveSelection(1); + return; + } + if (event.key === 'ArrowUp') { + event.preventDefault(); + moveSelection(-1); + return; + } + if (event.key === 'Enter') { + event.preventDefault(); + const selected = flatModelList[selectionStore.getSnapshot()]; + if (selected && !disabled) onSelect(selected); + return; + } + if (event.key === 'Escape') { + event.preventDefault(); + onEscape?.(); + } + }, [disabled, flatModelList, moveSelection, onActiveKeyDown, onEscape, onSelect, onVariantKey, selectionStore]); + + const headerClassName = cn( + 'typography-micro font-semibold text-muted-foreground uppercase tracking-wider flex items-center gap-2 -mx-1 px-3 py-1.5 border-b border-border/30', + stickyHeaders && 'sticky top-0 z-10 [background:linear-gradient(var(--surface-elevated),var(--surface-elevated)),linear-gradient(var(--surface-background),var(--surface-background))]', + sectionHeaderClassName, + ); + + let currentFlatIndex = 0; + + const renderRow = (entry: ModelPickerEntry, keyPrefix: string, showProviderLogo: boolean, rowIndex: number, dragHandleProps?: SortableFavoriteHandleProps | null) => { + const metadata = mergeModelMetadataWithLiveModel(entry.providerID, entry.model, modelsMetadata.get(`${entry.providerID}/${entry.modelID}`)); + const contextTokens = formatModelContextTokens(metadata?.limit?.context); + const count = selectionCount?.(entry) ?? 0; + const isSelected = selectedModel?.providerID === entry.providerID && selectedModel.modelID === entry.modelID; + const favorite = isFavorite?.(entry) ?? false; + + const handleMouseActivity = (event: React.MouseEvent) => { + const nextPosition = { x: event.clientX, y: event.clientY }; + const previousPosition = lastMousePositionRef.current; + const pointerMoved = !previousPosition || previousPosition.x !== nextPosition.x || previousPosition.y !== nextPosition.y; + lastMousePositionRef.current = nextPosition; + + if (keyboardOwnsSelectionRef.current && !previousPosition) return; + if (keyboardOwnsSelectionRef.current && !pointerMoved) return; + if (keyboardOwnsSelectionRef.current && pointerMoved) keyboardOwnsSelectionRef.current = false; + selectIndex(rowIndex); + }; + + return ( + + {(isHighlighted) => { + const rowElement = ( +
{ itemRefs.current[rowIndex] = el; }} + role="option" + aria-selected={isSelected} + aria-disabled={disabled || undefined} + tabIndex={-1} + onClick={() => { if (!disabled) onSelect(entry); }} + onKeyDown={(event) => { + if (disabled) return; + if (event.key === 'Enter' || event.key === ' ') { + event.preventDefault(); + onSelect(entry); + } + }} + onMouseEnter={handleMouseActivity} + onMouseMove={handleMouseActivity} + className={cn( + 'w-full text-left px-2 py-1.5 rounded-md typography-meta flex items-center gap-2 cursor-pointer', + !disabled && (isHighlighted ? 'bg-interactive-selection' : 'hover:bg-interactive-hover/50'), + disabled && 'cursor-not-allowed opacity-60', + rowClassName, + )} + > +
+ {dragHandleProps ? ( + + ) : null} + {showProviderLogo ? : null} + {getModelDisplayName(entry.model)} + {contextTokens ? {contextTokens} : null} +
+ {count > 0 ? x{count} : null} + {renderRowEnd?.(entry, { isHighlighted, isSelected })} + {isSelected ? : null} + {onToggleFavorite ? ( + + ) : null} +
+ ); + + return {rowElement}; + }} +
+ ); + }; + + const handleFavoriteDragEnd = (event: DragEndEvent) => { + if (!onReorderFavorite) return; + const { active, over } = event; + if (!over || active.id === over.id) return; + + const activeFavorite = favoriteLookup.get(String(active.id)); + const overFavorite = favoriteLookup.get(String(over.id)); + if (!activeFavorite || !overFavorite) return; + + onReorderFavorite(activeFavorite, overFavorite); + }; + + const isSectionCollapsed = (key: string) => collapsedSections.has(key); + const toggleSectionCollapsed = (key: string) => { + setCollapsedSections((prev) => { + const next = new Set(prev); + if (next.has(key)) next.delete(key); + else next.add(key); + return next; + }); + }; + + const renderSectionHeader = (key: string, icon: React.ReactNode, label: React.ReactNode) => { + const collapsed = isSectionCollapsed(key); + return ( + + ); + }; + + return ( + <> +
+
+ + onSearchQueryChange(event.target.value)} + onKeyDown={handleKeyDown} + className="h-7 rounded-none bg-transparent pl-8 pr-0 typography-meta ring-0 hover:[&:not(:focus)]:bg-transparent focus:ring-0 focus-visible:ring-0" + autoFocus={autoFocus} + /> +
+
+ + +
+ {includeNotSelected ? ( + <> + +
+ + ) : null} + + {!hasResults ? ( +
{labels.noResults}
+ ) : null} + + {filteredFavorites.length > 0 ? ( +
+ {renderSectionHeader('favorites', , labels.favorites)} + {!isSectionCollapsed('favorites') && (favoriteSortingEnabled ? ( + + `${entry.providerID}:${entry.modelID}`)} strategy={verticalListSortingStrategy}> + {filteredFavorites.map((entry) => { + const rowIndex = currentFlatIndex++; + return ( + + {(dragHandleProps) => renderRow(entry, 'fav', true, rowIndex, dragHandleProps)} + + ); + })} + + + ) : filteredFavorites.map((entry) => renderRow(entry, 'fav', true, currentFlatIndex++)))} +
+ ) : null} + + {filteredRecents.length > 0 ? ( +
+ {filteredFavorites.length > 0 ?
: null} + {renderSectionHeader('recent', , labels.recent)} + {!isSectionCollapsed('recent') ? filteredRecents.map((entry) => renderRow(entry, 'recent', true, currentFlatIndex++)) : null} +
+ ) : null} + + {(filteredFavorites.length > 0 || filteredRecents.length > 0) && filteredProviders.length > 0 ?
: null} + + {filteredProviders.map((provider, providerIndex) => ( +
+ {providerIndex > 0 ?
: null} + {renderSectionHeader(`provider:${provider.id}`, , provider.name || provider.id)} + {!isSectionCollapsed(`provider:${provider.id}`) + ? provider.models.map((model) => renderRow({ model, providerID: provider.id, modelID: model.id as string }, 'provider', false, currentFlatIndex++)) + : null} +
+ ))} +
+ + +
+ +
+ + ); +}; diff --git a/packages/ui/src/components/multirun/ModelMultiSelect.tsx b/packages/ui/src/components/multirun/ModelMultiSelect.tsx index a668c3e5..64c6e95f 100644 --- a/packages/ui/src/components/multirun/ModelMultiSelect.tsx +++ b/packages/ui/src/components/multirun/ModelMultiSelect.tsx @@ -1,16 +1,14 @@ import React from 'react'; import { Button } from '@/components/ui/button'; -import { Input } from '@/components/ui/input'; import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from '@/components/ui/select'; -import { ScrollableOverlay } from '@/components/ui/ScrollableOverlay'; import { ProviderLogo } from '@/components/ui/ProviderLogo'; import { Icon } from "@/components/icon/Icon"; import { cn } from '@/lib/utils'; -import { isIMECompositionEvent } from '@/lib/ime'; import { useConfigStore } from '@/stores/useConfigStore'; +import { useUIStore } from '@/stores/useUIStore'; import { useModelLists } from '@/hooks/useModelLists'; -import type { ModelMetadata } from '@/types'; import { useI18n } from '@/lib/i18n'; +import { ModelPickerList, type ModelPickerEntry, type ModelPickerProvider } from '@/components/model-picker/ModelPickerList'; /** Chip height class - shared between chips and add button */ const CHIP_HEIGHT_CLASS = 'h-7'; @@ -37,24 +35,6 @@ export const generateInstanceId = (): string => { return `${Date.now()}-${Math.random().toString(36).substring(2, 9)}`; }; -const COMPACT_NUMBER_FORMATTER = new Intl.NumberFormat('en-US', { - notation: 'compact', - compactDisplay: 'short', - maximumFractionDigits: 1, - minimumFractionDigits: 0, -}); - -const formatTokens = (value?: number | null) => { - if (typeof value !== 'number' || Number.isNaN(value)) { - return ''; - } - if (value === 0) { - return '0'; - } - const formatted = COMPACT_NUMBER_FORMATTER.format(value); - return formatted.endsWith('.0') ? formatted.slice(0, -2) : formatted; -}; - /** * Model selection chip with remove button. * Shows instance index (e.g., "(2)") when same model is selected multiple times. @@ -129,17 +109,16 @@ export const ModelMultiSelect: React.FC = ({ triggerIcon, }) => { const { t } = useI18n(); - const providers = useConfigStore((state) => state.providers); + const providers = useConfigStore((state) => state.providers) as ModelPickerProvider[]; const modelsMetadata = useConfigStore((state) => state.modelsMetadata); + const toggleFavoriteModel = useUIStore((state) => state.toggleFavoriteModel); + const isFavoriteModel = useUIStore((state) => state.isFavoriteModel); const { favoriteModelsList, recentModelsList } = useModelLists(); const [isOpen, setIsOpen] = React.useState(false); const [searchQuery, setSearchQuery] = React.useState(''); - const [selectedIndex, setSelectedIndex] = React.useState(0); const [availableHeight, setAvailableHeight] = React.useState(null); - const searchInputRef = React.useRef(null); const dropdownRef = React.useRef(null); const triggerRef = React.useRef(null); - const itemRefs = React.useRef<(HTMLButtonElement | null)[]>([]); const isSingleSelect = maxModels === 1; const canAddModel = maxModels === undefined || selectedModels.length < maxModels || isSingleSelect; @@ -161,66 +140,6 @@ export const ModelMultiSelect: React.FC = ({ return sameModels.findIndex(m => m.instanceId === model.instanceId) + 1; }, [selectedModels]); - const getModelMetadata = (provId: string, modId: string): ModelMetadata | undefined => { - const key = `${provId}/${modId}`; - return modelsMetadata.get(key); - }; - - const getModelDisplayName = (model: Record) => { - const name = model?.name || model?.id || ''; - const nameStr = String(name); - if (nameStr.length > 40) { - return nameStr.substring(0, 37) + '...'; - } - return nameStr; - }; - - // Filter helper - const filterByQuery = React.useCallback((modelName: string, providerName: string) => { - if (!searchQuery.trim()) return true; - const lowerQuery = searchQuery.toLowerCase(); - return ( - modelName.toLowerCase().includes(lowerQuery) || - providerName.toLowerCase().includes(lowerQuery) - ); - }, [searchQuery]); - - // Filter favorites - const filteredFavorites = React.useMemo(() => { - return favoriteModelsList.filter(({ model, providerID }) => { - const provider = providers.find(p => p.id === providerID); - const providerName = provider?.name || providerID; - const modelName = getModelDisplayName(model); - return filterByQuery(modelName, providerName); - }); - }, [favoriteModelsList, providers, filterByQuery]); - - // Filter recents - const filteredRecents = React.useMemo(() => { - return recentModelsList.filter(({ model, providerID }) => { - const provider = providers.find(p => p.id === providerID); - const providerName = provider?.name || providerID; - const modelName = getModelDisplayName(model); - return filterByQuery(modelName, providerName); - }); - }, [recentModelsList, providers, filterByQuery]); - - // Filter providers - const filteredProviders = React.useMemo(() => { - return providers - .map((provider) => { - const models = Array.isArray(provider.models) ? provider.models : []; - const filteredModels = models.filter((model) => { - const modelName = getModelDisplayName(model); - return filterByQuery(modelName, provider.name || provider.id || ''); - }); - return { ...provider, models: filteredModels }; - }) - .filter((provider) => provider.models.length > 0); - }, [providers, filterByQuery]); - - const hasResults = filteredFavorites.length > 0 || filteredRecents.length > 0 || filteredProviders.length > 0; - // Calculate available height: multi-run opens upward inside a scroller; fusion opens downward and may extend past the dialog. React.useEffect(() => { if (!isOpen || !triggerRef.current) return; @@ -255,18 +174,10 @@ export const ModelMultiSelect: React.FC = ({ setAvailableHeight(Math.max(150, Math.min(300, spaceAbove))); }, [dropdownSide, isOpen]); - // Focus search input when opened - React.useEffect(() => { - if (isOpen && searchInputRef.current) { - searchInputRef.current.focus(); - } - }, [isOpen]); - React.useEffect(() => { if (!canAddModel && isOpen) { setIsOpen(false); setSearchQuery(''); - setSelectedIndex(0); } }, [canAddModel, isOpen]); @@ -278,7 +189,6 @@ export const ModelMultiSelect: React.FC = ({ if (dropdownRef.current && !dropdownRef.current.contains(event.target as Node)) { setIsOpen(false); setSearchQuery(''); - setSelectedIndex(0); } }; @@ -286,85 +196,39 @@ export const ModelMultiSelect: React.FC = ({ return () => document.removeEventListener('mousedown', handleClickOutside); }, [isOpen]); - // Reset selection when search query changes - React.useEffect(() => { - setSelectedIndex(0); - }, [searchQuery]); + const handleSelectModel = React.useCallback((entry: ModelPickerEntry) => { + const nextModel = { + providerID: entry.providerID, + modelID: entry.modelID, + displayName: (entry.model.name as string) || entry.modelID, + instanceId: generateInstanceId(), + }; + if (isSingleSelect && selectedModels.length > 0 && onUpdate) { + onUpdate(0, nextModel); + } else { + onAdd(nextModel); + } + if (isSingleSelect) { + setIsOpen(false); + setSearchQuery(''); + } + }, [isSingleSelect, onAdd, onUpdate, selectedModels.length]); - // Render a model row - const renderModelRow = ( - model: Record, - providerID: string, - modelID: string, - keyPrefix: string, - flatIndex: number, - isHighlighted: boolean - ) => { - const key = `${providerID}:${modelID}`; - const selectionCount = modelCounts.get(key) || 0; - const metadata = getModelMetadata(providerID, modelID); - const contextTokens = formatTokens(metadata?.limit?.context); - - const showProviderLogo = keyPrefix === 'fav' || keyPrefix === 'recent'; - - return ( - - ); - }; + const labels = React.useMemo(() => ({ + searchPlaceholder: t('multirun.modelMultiSelect.search.placeholder'), + noResults: t('multirun.modelMultiSelect.search.noResults'), + favorites: t('multirun.modelMultiSelect.sections.favorites'), + recent: t('multirun.modelMultiSelect.sections.recent'), + keyboardHint: t('multirun.modelMultiSelect.keyboard.hint'), + favorite: t('settings.agents.modelSelector.actions.favorite'), + unfavorite: t('settings.agents.modelSelector.actions.unfavorite'), + capabilities: t('chat.modelControls.capabilities'), + capabilityToolCalling: t('chat.modelControls.capability.toolCalling'), + capabilityReasoning: t('chat.modelControls.capability.reasoning'), + input: t('chat.modelControls.input'), + output: t('chat.modelControls.output'), + costPerMillion: t('chat.modelControls.costPerMillion'), + }), [t]); return (
@@ -390,179 +254,41 @@ export const ModelMultiSelect: React.FC = ({ {addButtonLabel ?? t('multirun.modelMultiSelect.actions.addModel')} - {isOpen && (() => { - // Build flat list for keyboard navigation - type FlatModelItem = { model: Record; providerID: string; modelID: string; section: string }; - const flatModelList: FlatModelItem[] = []; - - filteredFavorites.forEach(({ model, providerID, modelID }) => { - flatModelList.push({ model, providerID, modelID, section: 'fav' }); - }); - filteredRecents.forEach(({ model, providerID, modelID }) => { - flatModelList.push({ model, providerID, modelID, section: 'recent' }); - }); - filteredProviders.forEach((provider) => { - provider.models.forEach((model) => { - flatModelList.push({ model, providerID: provider.id, modelID: model.id as string, section: 'provider' }); - }); - }); - - const totalItems = flatModelList.length; - - // Handle keyboard navigation - const handleKeyDown = (e: React.KeyboardEvent) => { - if (isIMECompositionEvent(e)) { - return; - } - if (e.key === 'ArrowDown') { - e.preventDefault(); - e.stopPropagation(); - const nextIndex = (selectedIndex + 1) % Math.max(1, totalItems); - setSelectedIndex(nextIndex); - setTimeout(() => { - itemRefs.current[nextIndex]?.scrollIntoView({ behavior: 'smooth', block: 'nearest' }); - }, 0); - } else if (e.key === 'ArrowUp') { - e.preventDefault(); - e.stopPropagation(); - const prevIndex = (selectedIndex - 1 + Math.max(1, totalItems)) % Math.max(1, totalItems); - setSelectedIndex(prevIndex); - setTimeout(() => { - itemRefs.current[prevIndex]?.scrollIntoView({ behavior: 'smooth', block: 'nearest' }); - }, 0); - } else if (e.key === 'Enter') { - e.preventDefault(); - e.stopPropagation(); - const selectedItem = flatModelList[selectedIndex]; - if (selectedItem && canAddModel) { - const nextModel = { - providerID: selectedItem.providerID, - modelID: selectedItem.modelID, - displayName: (selectedItem.model.name as string) || selectedItem.modelID, - instanceId: generateInstanceId(), - }; - if (isSingleSelect && selectedModels.length > 0 && onUpdate) { - onUpdate(0, nextModel); - } else { - onAdd(nextModel); - } - if (isSingleSelect) { - setIsOpen(false); - setSearchQuery(''); - setSelectedIndex(0); - } - } - } else if (e.key === 'Escape') { - e.preventDefault(); - e.stopPropagation(); - setIsOpen(false); - setSearchQuery(''); - setSelectedIndex(0); - } - }; - - let currentFlatIndex = 0; - - return ( -
+ modelCounts.get(`${entry.providerID}:${entry.modelID}`) || 0} + disabled={!canAddModel} + maxHeightClassName="flex-1" + maxHeightStyle={{ maxHeight: availableHeight ? `${availableHeight}px` : '300px' }} + stickyHeaders + tooltipsEnabled={isOpen} + isFavorite={(entry) => isFavoriteModel(entry.providerID, entry.modelID)} + onToggleFavorite={(entry) => toggleFavoriteModel(entry.providerID, entry.modelID)} + onEscape={() => { + setIsOpen(false); + setSearchQuery(''); }} - > - {/* Search input */} -
-
- - setSearchQuery(e.target.value)} - onKeyDown={handleKeyDown} - className="h-8 pl-8 typography-meta" - /> -
-
- - {/* Models list */} - -
- {!hasResults && ( -
- {t('multirun.modelMultiSelect.search.noResults')} -
- )} - - {/* Favorites Section */} - {filteredFavorites.length > 0 && ( - <> -
- - {t('multirun.modelMultiSelect.sections.favorites')} -
- {filteredFavorites.map(({ model, providerID, modelID }) => { - const idx = currentFlatIndex++; - return renderModelRow(model, providerID, modelID, 'fav', idx, selectedIndex === idx); - })} - - )} - - {/* Recents Section */} - {filteredRecents.length > 0 && ( - <> - {filteredFavorites.length > 0 &&
} -
- - {t('multirun.modelMultiSelect.sections.recent')} -
- {filteredRecents.map(({ model, providerID, modelID }) => { - const idx = currentFlatIndex++; - return renderModelRow(model, providerID, modelID, 'recent', idx, selectedIndex === idx); - })} - - )} - - {/* Separator before providers */} - {(filteredFavorites.length > 0 || filteredRecents.length > 0) && filteredProviders.length > 0 && ( -
- )} - - {/* All Providers - Flat List */} - {filteredProviders.map((provider, index) => ( - - {index > 0 &&
} -
- - {provider.name} -
- {provider.models.map((model) => { - const idx = currentFlatIndex++; - return renderModelRow(model, provider.id, model.id as string, 'provider', idx, selectedIndex === idx); - })} - - ))} -
- - - {/* Keyboard hints footer */} -
- {t('multirun.modelMultiSelect.keyboard.hint')} -
-
- ); - })()} + /> +
+ ) : null}
{/* Selected models */} @@ -574,7 +300,7 @@ export const ModelMultiSelect: React.FC = ({ const instanceIndex = getInstanceIndex(model); const provider = providers.find((p) => p.id === model.providerID); - const providerModel = provider?.models.find((m: Record) => (m as { id?: string }).id === model.modelID) as + const providerModel = provider?.models?.find((m: Record) => (m as { id?: string }).id === model.modelID) as | { variants?: Record } | undefined; const variantKeys = providerModel?.variants ? Object.keys(providerModel.variants) : []; diff --git a/packages/ui/src/components/sections/agents/ModelSelector.tsx b/packages/ui/src/components/sections/agents/ModelSelector.tsx index e0e61666..ae2175fb 100644 --- a/packages/ui/src/components/sections/agents/ModelSelector.tsx +++ b/packages/ui/src/components/sections/agents/ModelSelector.tsx @@ -2,26 +2,19 @@ import React from 'react'; import { DropdownMenu, DropdownMenuContent, - DropdownMenuItem, - DropdownMenuLabel, - DropdownMenuSeparator, DropdownMenuTrigger, } from '@/components/ui/dropdown-menu'; -import { Input } from '@/components/ui/input'; -import { useConfigStore } from '@/stores/useConfigStore'; -import { useUIStore } from '@/stores/useUIStore'; -import { useDeviceInfo } from '@/lib/device'; -import { cn } from '@/lib/utils'; import { MobileOverlayPanel } from '@/components/ui/MobileOverlayPanel'; import { ProviderLogo } from '@/components/ui/ProviderLogo'; -import { ScrollableOverlay } from '@/components/ui/ScrollableOverlay'; -import { Icon } from "@/components/icon/Icon"; +import { Icon } from '@/components/icon/Icon'; import { useModelLists } from '@/hooks/useModelLists'; -import type { ModelMetadata } from '@/types'; -import { useI18n } from '@/lib/i18n'; import { useOpenCodeReadiness } from '@/hooks/useOpenCodeReadiness'; - -type ProviderModel = Record & { id?: string; name?: string }; +import { useDeviceInfo } from '@/lib/device'; +import { useI18n } from '@/lib/i18n'; +import { cn } from '@/lib/utils'; +import { useConfigStore } from '@/stores/useConfigStore'; +import { useUIStore } from '@/stores/useUIStore'; +import { ModelPickerList, type ModelPickerEntry, type ModelPickerProvider } from '@/components/model-picker/ModelPickerList'; interface ModelSelectorProps { providerId: string; @@ -32,38 +25,20 @@ interface ModelSelectorProps { placeholder?: string; } -const COMPACT_NUMBER_FORMATTER = new Intl.NumberFormat('en-US', { - notation: 'compact', - compactDisplay: 'short', - maximumFractionDigits: 1, - minimumFractionDigits: 0, -}); - -const formatTokens = (value?: number | null) => { - if (typeof value !== 'number' || Number.isNaN(value)) { - return ''; - } - if (value === 0) { - return '0'; - } - const formatted = COMPACT_NUMBER_FORMATTER.format(value); - return formatted.endsWith('.0') ? formatted.slice(0, -2) : formatted; -}; - export const ModelSelector: React.FC = ({ providerId, modelId, onChange, className, allowedProviderIds, - placeholder + placeholder, }) => { const { t } = useI18n(); const { isReady, isUnavailable } = useOpenCodeReadiness(); - const providers = useConfigStore((state) => state.providers); + const providers = useConfigStore((state) => state.providers) as ModelPickerProvider[]; const modelsMetadata = useConfigStore((state) => state.modelsMetadata); - const isMobile = useUIStore(state => state.isMobile); - const hiddenModels = useUIStore(state => state.hiddenModels); + const isMobile = useUIStore((state) => state.isMobile); + const hiddenModels = useUIStore((state) => state.hiddenModels); const toggleFavoriteModel = useUIStore((state) => state.toggleFavoriteModel); const isFavoriteModel = useUIStore((state) => state.isFavoriteModel); const addRecentModel = useUIStore((state) => state.addRecentModel); @@ -72,429 +47,71 @@ export const ModelSelector: React.FC = ({ const isActuallyMobile = isMobile || deviceIsMobile; const [isMobilePanelOpen, setIsMobilePanelOpen] = React.useState(false); - const [expandedMobileProviders, setExpandedMobileProviders] = React.useState>(new Set()); const [isDropdownOpen, setIsDropdownOpen] = React.useState(false); const [searchQuery, setSearchQuery] = React.useState(''); - const [selectedIndex, setSelectedIndex] = React.useState(0); - const itemRefs = React.useRef<(HTMLElement | null)[]>([]); - const allowedProviderSet = React.useMemo(() => { - if (!Array.isArray(allowedProviderIds) || allowedProviderIds.length === 0) { - return null; - } - return new Set(allowedProviderIds); - }, [allowedProviderIds]); - - const visibleProviders = React.useMemo(() => { - const baseProviders = allowedProviderSet - ? providers.filter((provider) => allowedProviderSet.has(String(provider.id))) - : providers; - - return baseProviders - .map((provider) => { - const providerModels = Array.isArray(provider.models) ? provider.models : []; - const filteredModels = providerModels.filter((model: ProviderModel) => { - const modelId = typeof model?.id === 'string' ? model.id : ''; - return !hiddenModels.some( - (hidden) => hidden.providerID === String(provider.id) && hidden.modelID === modelId - ); - }); - return { ...provider, models: filteredModels }; - }) - .filter((provider) => provider.models.length > 0); - }, [providers, allowedProviderSet, hiddenModels]); - - const closeMobilePanel = () => setIsMobilePanelOpen(false); - const toggleMobileProviderExpansion = (provId: string) => { - setExpandedMobileProviders(prev => { - const newSet = new Set(prev); - if (newSet.has(provId)) { - newSet.delete(provId); - } else { - newSet.add(provId); - } - return newSet; - }); - }; - - // Reset search and selection when dropdown closes - React.useEffect(() => { - if (!isDropdownOpen) { - setSearchQuery(''); - setSelectedIndex(0); - } - }, [isDropdownOpen]); - - // Reset selection when search query changes - React.useEffect(() => { - setSelectedIndex(0); - }, [searchQuery]); - - const getModelDisplayName = (model: Record) => { - const name = model?.name || model?.id || ''; - const nameStr = String(name); - if (nameStr.length > 40) { - return nameStr.substring(0, 37) + '...'; - } - return nameStr; - }; - - const getModelMetadata = (provId: string, modId: string): ModelMetadata | undefined => { - const key = `${provId}/${modId}`; - return modelsMetadata.get(key); - }; - - const handleProviderAndModelChange = (newProviderId: string, newModelId: string) => { - onChange(newProviderId, newModelId); - if (newProviderId && newModelId) { - addRecentModel(newProviderId, newModelId); - } + const closePicker = React.useCallback(() => { + setIsMobilePanelOpen(false); setIsDropdownOpen(false); - }; + setSearchQuery(''); + }, []); - // Filter helper - const filterByQuery = (modelName: string, providerName: string) => { - if (!searchQuery.trim()) return true; - const lowerQuery = searchQuery.toLowerCase(); + const handleSelect = React.useCallback((entry: ModelPickerEntry) => { + onChange(entry.providerID, entry.modelID); + addRecentModel(entry.providerID, entry.modelID); + closePicker(); + }, [addRecentModel, closePicker, onChange]); + + const handleSelectNone = React.useCallback(() => { + onChange('', ''); + closePicker(); + }, [closePicker, onChange]); + + const labels = React.useMemo(() => ({ + searchPlaceholder: t('settings.agents.modelSelector.searchPlaceholder'), + noResults: t('settings.agents.modelSelector.state.noModelsFound'), + favorites: t('settings.agents.modelSelector.section.favorites'), + recent: t('settings.agents.modelSelector.section.recent'), + keyboardHint: t('settings.agents.modelSelector.keyboardHints'), + notSelected: placeholder || t('settings.agents.modelSelector.notSelected'), + favorite: t('settings.agents.modelSelector.actions.favorite'), + unfavorite: t('settings.agents.modelSelector.actions.unfavorite'), + capabilities: t('chat.modelControls.capabilities'), + capabilityToolCalling: t('chat.modelControls.capability.toolCalling'), + capabilityReasoning: t('chat.modelControls.capability.reasoning'), + input: t('chat.modelControls.input'), + output: t('chat.modelControls.output'), + costPerMillion: t('chat.modelControls.costPerMillion'), + }), [placeholder, t]); + + const selectedModel = providerId && modelId ? { providerID: providerId, modelID: modelId } : null; + const triggerLabel = providerId && modelId ? `${providerId}/${modelId}` : (placeholder || t('settings.agents.modelSelector.notSelected')); + + const picker = ( + isFavoriteModel(entry.providerID, entry.modelID)} + onToggleFavorite={(entry) => toggleFavoriteModel(entry.providerID, entry.modelID)} + /> + ); + + if (isActuallyMobile) { return ( - modelName.toLowerCase().includes(lowerQuery) || - providerName.toLowerCase().includes(lowerQuery) - ); - }; - - // Render a model row for desktop dropdown - const renderModelRow = ( - model: ProviderModel, - provID: string, - modID: string, - keyPrefix: string, - flatIndex: number, - isHighlighted: boolean - ) => { - const metadata = getModelMetadata(provID, modID); - const contextTokens = formatTokens(metadata?.limit?.context); - const isSelected = providerId === provID && modelId === modID; - const isFavorite = isFavoriteModel(provID, modID); - - const showProviderLogo = keyPrefix === 'fav' || keyPrefix === 'recent'; - - return ( - { itemRefs.current[flatIndex] = el; }} - className={cn( - "group flex items-center gap-2", - isHighlighted && "bg-interactive-selection" - )} - onSelect={() => handleProviderAndModelChange(provID, modID)} - onMouseEnter={() => setSelectedIndex(flatIndex)} - > -
- {showProviderLogo && ( - - )} - - {getModelDisplayName(model)} - - {contextTokens ? ( - - {contextTokens} - - ) : null} -
-
- {isSelected && ( - - )} - -
-
- ); - }; - - // Filter data for desktop dropdown - const filteredFavorites = favoriteModelsList.filter(({ model, providerID }) => { - if (allowedProviderSet && !allowedProviderSet.has(providerID)) { - return false; - } - const provider = providers.find(p => p.id === providerID); - const providerName = provider?.name || providerID; - const modelName = getModelDisplayName(model); - return filterByQuery(modelName, providerName); - }); - - const filteredRecents = recentModelsList.filter(({ model, providerID }) => { - if (allowedProviderSet && !allowedProviderSet.has(providerID)) { - return false; - } - const provider = providers.find(p => p.id === providerID); - const providerName = provider?.name || providerID; - const modelName = getModelDisplayName(model); - return filterByQuery(modelName, providerName); - }); - - const filteredProviders = visibleProviders - .map((provider) => { - const providerModels = Array.isArray(provider.models) ? provider.models : []; - const filteredModels = providerModels.filter((model: ProviderModel) => { - const modelName = getModelDisplayName(model); - return filterByQuery(modelName, provider.name || provider.id || ''); - }); - return { ...provider, models: filteredModels }; - }) - .filter((provider) => provider.models.length > 0); - - const hasResults = filteredFavorites.length > 0 || filteredRecents.length > 0 || filteredProviders.length > 0; - - const renderMobileModelPanel = () => { - if (!isActuallyMobile) return null; - - return ( - -
- {/* Favorites Section for Mobile */} - {favoriteModelsList.length > 0 && ( -
-
- {t('settings.agents.modelSelector.section.favorites')} -
-
- {favoriteModelsList.map(({ model, providerID, modelID }) => { - const isSelectedModel = providerID === providerId && modelID === modelId; - - return ( -
- - - -
- ); - })} -
-
- )} - - {/* Recents Section for Mobile */} - {recentModelsList.length > 0 && ( -
-
- {t('settings.agents.modelSelector.section.recents')} -
-
- {recentModelsList.map(({ model, providerID, modelID }) => { - const isSelectedModel = providerID === providerId && modelID === modelId; - - return ( -
- - - -
- ); - })} -
-
- )} - - {visibleProviders.map((provider) => { - const providerModels = Array.isArray(provider.models) ? provider.models : []; - if (providerModels.length === 0) return null; - - const isActiveProvider = provider.id === providerId; - const isExpanded = expandedMobileProviders.has(provider.id); - - return ( -
- - - {isExpanded && ( -
- {providerModels.map((modelItem: ProviderModel) => { - const isSelectedModel = provider.id === providerId && modelItem.id === modelId; - - return ( -
- - -
- - - {isSelectedModel && ( -
- )} -
-
- ); - })} -
- )} -
- ); - })} - - -
- - ); - }; - - return ( - <> - {isActuallyMobile ? ( + <> - ) : ( - - -
- {!isReady ? ( - <> - - - {isUnavailable ? t('common.unavailable') : t('common.loading')} - - - ) : ( - <> - {providerId ? ( - <> - - - - ) : ( - - )} - - {providerId && modelId ? `${providerId}/${modelId}` : (placeholder || t('settings.agents.modelSelector.notSelected'))} - - - )} - -
-
- - {(() => { - // Build flat list for keyboard navigation - type FlatModelItem = { model: ProviderModel; providerID: string; modelID: string; section: string }; - const flatModelList: FlatModelItem[] = []; - - filteredFavorites.forEach(({ model, providerID, modelID }) => { - flatModelList.push({ model, providerID, modelID, section: 'fav' }); - }); - filteredRecents.forEach(({ model, providerID, modelID }) => { - flatModelList.push({ model, providerID, modelID, section: 'recent' }); - }); - filteredProviders.forEach((provider) => { - (provider.models as ProviderModel[]).forEach((model) => { - flatModelList.push({ model, providerID: provider.id as string, modelID: model.id as string, section: 'provider' }); - }); - }); + + {picker} + + + ); + } - const totalItems = flatModelList.length; - - // Handle keyboard navigation - const handleKeyDown = (e: React.KeyboardEvent) => { - e.stopPropagation(); - - if (e.key === 'ArrowDown') { - e.preventDefault(); - const nextIndex = (selectedIndex + 1) % Math.max(1, totalItems); - setSelectedIndex(nextIndex); - setTimeout(() => { - itemRefs.current[nextIndex]?.scrollIntoView({ behavior: 'smooth', block: 'nearest' }); - }, 0); - } else if (e.key === 'ArrowUp') { - e.preventDefault(); - const prevIndex = (selectedIndex - 1 + Math.max(1, totalItems)) % Math.max(1, totalItems); - setSelectedIndex(prevIndex); - setTimeout(() => { - itemRefs.current[prevIndex]?.scrollIntoView({ behavior: 'smooth', block: 'nearest' }); - }, 0); - } else if (e.key === 'Enter') { - e.preventDefault(); - const selectedItem = flatModelList[selectedIndex]; - if (selectedItem) { - handleProviderAndModelChange(selectedItem.providerID, selectedItem.modelID); - } - } else if (e.key === 'Escape') { - e.preventDefault(); - setIsDropdownOpen(false); - } - }; - - let currentFlatIndex = 0; - - return ( - <> - {/* Search Input */} -
-
- - setSearchQuery(e.target.value)} - onKeyDown={handleKeyDown} - className="pl-8 h-8 typography-meta" - autoFocus - /> -
-
- - {/* Scrollable content */} - -
- {/* Not selected option */} - handleProviderAndModelChange('', '')} - > - - {placeholder || t('settings.agents.modelSelector.notSelected')} - {!providerId && !modelId && ( - - )} - - - - - {!hasResults && searchQuery && ( -
- {t('settings.agents.modelSelector.state.noModelsFound')} -
- )} - - {/* Favorites Section */} - {filteredFavorites.length > 0 && ( -
- - - {t('settings.agents.modelSelector.section.favorites')} - - {filteredFavorites.map(({ model, providerID, modelID }) => { - const idx = currentFlatIndex++; - return renderModelRow(model, providerID, modelID, 'fav', idx, selectedIndex === idx); - })} -
- )} - - {/* Recents Section */} - {filteredRecents.length > 0 && ( -
- {filteredFavorites.length > 0 && } - - - {t('settings.agents.modelSelector.section.recent')} - - {filteredRecents.map(({ model, providerID, modelID }) => { - const idx = currentFlatIndex++; - return renderModelRow(model, providerID, modelID, 'recent', idx, selectedIndex === idx); - })} -
- )} - - {/* Separator before providers */} - {(filteredFavorites.length > 0 || filteredRecents.length > 0) && filteredProviders.length > 0 && ( - - )} - - {/* All Providers - Flat List */} - {filteredProviders.map((provider, index) => ( -
- {index > 0 && } - - - {provider.name} - - {(provider.models as ProviderModel[]).map((model: ProviderModel) => { - const idx = currentFlatIndex++; - return renderModelRow(model, provider.id as string, model.id as string, 'provider', idx, selectedIndex === idx); - })} -
- ))} -
-
- - {/* Keyboard hints footer */} -
- {t('settings.agents.modelSelector.keyboardHints')} -
- - ); - })()} -
-
- )} - {renderMobileModelPanel()} - + return ( + + +
+ {!isReady ? ( + <> + + + {isUnavailable ? t('common.unavailable') : t('common.loading')} + + + ) : ( + <> + {providerId ? : } + {triggerLabel} + + )} + +
+
+ + {picker} + +
); }; diff --git a/packages/ui/src/lib/modelMetadata.ts b/packages/ui/src/lib/modelMetadata.ts new file mode 100644 index 00000000..98178e2d --- /dev/null +++ b/packages/ui/src/lib/modelMetadata.ts @@ -0,0 +1,35 @@ +import type { ModelMetadata } from '@/types'; + +type LiveProviderModel = Record & { id?: string; name?: string }; + +const getNumericLimit = (limit: unknown, key: 'context' | 'output') => { + if (!limit || typeof limit !== 'object') return undefined; + const value = (limit as Record)[key]; + return typeof value === 'number' && Number.isFinite(value) ? value : undefined; +}; + +export const mergeModelMetadataWithLiveModel = ( + providerId: string, + model: LiveProviderModel, + metadata?: ModelMetadata, +): ModelMetadata | undefined => { + const liveContextLimit = getNumericLimit(model.limit, 'context'); + const liveOutputLimit = getNumericLimit(model.limit, 'output'); + const contextLimit = liveContextLimit ?? metadata?.limit?.context; + const outputLimit = liveOutputLimit ?? metadata?.limit?.output; + + if (contextLimit === undefined && outputLimit === undefined) return metadata; + + return { + ...(metadata ?? { + id: typeof model.id === 'string' ? model.id : '', + providerId, + name: typeof model.name === 'string' ? model.name : undefined, + }), + limit: { + ...metadata?.limit, + ...(contextLimit !== undefined ? { context: contextLimit } : {}), + ...(outputLimit !== undefined ? { output: outputLimit } : {}), + }, + }; +};