From 388687470b4f459674feaf94a828e922c70d31bf Mon Sep 17 00:00:00 2001 From: Sean Teramae Date: Tue, 28 Jul 2026 11:19:31 -0700 Subject: [PATCH 1/3] fix(studio): Performance optimization for ModelSelect model fetch Signed-off-by: Sean Teramae --- .../common/src/api/models/useModelEntity.ts | 31 +++ .../src/api/models/useModelSearch.test.tsx | 118 ++++++++++++ .../common/src/api/models/useModelSearch.ts | 131 +++++++++++++ .../ModelSelectV2/ModelDropdown.test.tsx | 119 ++++++++++++ .../ModelSelectV2/ModelDropdown.tsx | 145 +++++++------- .../ModelSelectV2/ModelDropdownItem.tsx | 10 +- .../ModelSelectV2/ModelDropdownList.tsx | 181 ++++++++++++++++++ .../ModelSelectV2/ModelSelectV2.stories.tsx | 65 +++++++ .../ModelSelectV2/ModelSelectV2.tsx | 29 +-- .../ModelSelectV2/WorkspaceModelSelect.tsx | 68 +++++++ .../src/components/ModelSelectV2/index.tsx | 2 + .../src/components/ModelSelectV2/types.ts | 43 ++++- web/packages/common/src/utils/models.ts | 3 + .../AddModelPalette.stories.tsx | 2 +- .../src/components/AddModelPalette/index.tsx | 20 +- .../ModelChatPanel/ModelChatPanel.test.tsx | 40 +--- .../src/components/ModelChatPanel/index.tsx | 17 +- .../src/components/ModelCompareChat/index.tsx | 7 - .../ModelComparePrompts/ModelColumnSelect.tsx | 18 +- .../ModelComparePrompts/ModelCompareTable.tsx | 10 +- .../components/ModelComparePrompts/index.tsx | 5 +- .../components/ModelComparePrompts/types.ts | 3 - .../src/components/ModelConfigPanel/index.tsx | 18 +- .../ModelSelectionSection.tsx | 10 +- .../ModelDetailsSection/index.tsx | 12 +- .../evaluation/JudgeModelSelect.tsx | 31 +-- .../sidePanels/MetricRunSidePanel/index.tsx | 29 +-- .../BuilderConfigPane.tsx | 10 +- .../BuilderPalette.tsx | 10 +- .../DataDesignerJobBuildRoute/index.tsx | 25 +-- .../DataDesignerJobBuildRoute/models.test.ts | 18 +- .../DataDesignerJobBuildRoute/models.ts | 72 +++++-- .../useJobBuilder.ts | 47 +++-- .../WorkspaceSourceFields.tsx | 24 +-- .../src/routes/ModelCompareRoute/index.tsx | 66 +++---- 35 files changed, 1080 insertions(+), 359 deletions(-) create mode 100644 web/packages/common/src/api/models/useModelEntity.ts create mode 100644 web/packages/common/src/api/models/useModelSearch.test.tsx create mode 100644 web/packages/common/src/api/models/useModelSearch.ts create mode 100644 web/packages/common/src/components/ModelSelectV2/ModelDropdown.test.tsx create mode 100644 web/packages/common/src/components/ModelSelectV2/ModelDropdownList.tsx create mode 100644 web/packages/common/src/components/ModelSelectV2/WorkspaceModelSelect.tsx diff --git a/web/packages/common/src/api/models/useModelEntity.ts b/web/packages/common/src/api/models/useModelEntity.ts new file mode 100644 index 0000000000..65f53d8258 --- /dev/null +++ b/web/packages/common/src/api/models/useModelEntity.ts @@ -0,0 +1,31 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import { getPartsFromReference } from '@nemo/common/src/namedEntity'; +import { useModelsGetModel } from '@nemo/sdk/generated/platform/api'; +import type { ModelEntity } from '@nemo/sdk/generated/platform/schema'; + +export interface UseModelEntityOptions { + enabled?: boolean; +} + +/** + * Resolves a single model URN to its entity. + * + * Companion to `useModelSearch`: a paged dropdown only holds the models it has loaded, so a + * selection restored from a URL or a form default has no entity attached. Callers that need + * fields off the entity — `model_providers`, adapters, deployment state — fetch just that one + * model instead of walking the catalogue to find it. + * + * Endpoint: GET /apis/models/v2/workspaces/{workspace}/models/{name} + */ +export const useModelEntity = ( + modelUrn: string | null | undefined, + { enabled = true }: UseModelEntityOptions = {} +): ModelEntity | undefined => { + const parts = modelUrn ? getPartsFromReference(modelUrn) : undefined; + const { data } = useModelsGetModel(parts?.workspace ?? '', parts?.name ?? '', undefined, { + query: { enabled: enabled && !!parts?.workspace && !!parts?.name }, + }); + return data; +}; diff --git a/web/packages/common/src/api/models/useModelSearch.test.tsx b/web/packages/common/src/api/models/useModelSearch.test.tsx new file mode 100644 index 0000000000..5c7b80dd93 --- /dev/null +++ b/web/packages/common/src/api/models/useModelSearch.test.tsx @@ -0,0 +1,118 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import { useModelSearch } from '@nemo/common/src/api/models/useModelSearch'; +import { modelsListModels } from '@nemo/sdk/generated/platform/api'; +import type { ModelEntity, ModelEntitysPage } from '@nemo/sdk/generated/platform/schema'; +import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; +import { act, renderHook, waitFor } from '@testing-library/react'; +import type { ReactNode } from 'react'; + +vi.mock('@nemo/sdk/generated/platform/api', async (importOriginal) => { + const actual = await importOriginal(); + return { ...actual, modelsListModels: vi.fn() }; +}); + +const mockListModels = vi.mocked(modelsListModels); + +const createWrapper = () => { + const queryClient = new QueryClient({ + defaultOptions: { queries: { retry: false, gcTime: 0 } }, + }); + return ({ children }: { children: ReactNode }) => ( + {children} + ); +}; + +const makeModel = (name: string, overrides: Partial = {}): ModelEntity => + ({ id: name, name, workspace: 'ws1', ...overrides }) as ModelEntity; + +const makePage = (data: ModelEntity[], page: number, totalPages: number): ModelEntitysPage => + ({ data, pagination: { page, total_pages: totalPages } }) as ModelEntitysPage; + +const renderSearch = (options: Partial[0]> = {}) => + renderHook(() => useModelSearch({ workspace: 'ws1', ...options }), { wrapper: createWrapper() }); + +beforeEach(() => { + mockListModels.mockReset(); +}); + +describe('useModelSearch', () => { + it('stays idle while disabled', () => { + renderSearch({ enabled: false }); + expect(mockListModels).not.toHaveBeenCalled(); + }); + + it('stays idle without a workspace', () => { + renderHook(() => useModelSearch({ workspace: null }), { wrapper: createWrapper() }); + expect(mockListModels).not.toHaveBeenCalled(); + }); + + it('groups the first page and reports that more remain', async () => { + mockListModels.mockResolvedValue(makePage([makeModel('a'), makeModel('b')], 1, 2)); + + const { result } = renderSearch(); + + await waitFor(() => expect(result.current.groups).toHaveLength(1)); + expect(result.current.groups[0].models.map((m) => m.name)).toEqual(['a', 'b']); + expect(result.current.hasMore).toBe(true); + }); + + it('sends the search term as a case-insensitive substring filter', async () => { + mockListModels.mockResolvedValue(makePage([makeModel('a')], 1, 1)); + const { result } = renderSearch(); + await waitFor(() => expect(mockListModels).toHaveBeenCalled()); + + act(() => result.current.onSearchChange(' llama ')); + + await waitFor(() => + expect(mockListModels).toHaveBeenLastCalledWith( + 'ws1', + expect.objectContaining({ filter: expect.objectContaining({ name: { $like: 'llama' } }) }) + ) + ); + }); + + it('merges caller filters with the search term', async () => { + mockListModels.mockResolvedValue(makePage([makeModel('a')], 1, 1)); + + renderSearch({ filter: { lora_enabled: true } }); + + await waitFor(() => + expect(mockListModels).toHaveBeenCalledWith( + 'ws1', + expect.objectContaining({ filter: { lora_enabled: true } }) + ) + ); + }); + + it('keeps paging when include filters a whole page down to nothing', async () => { + mockListModels + .mockResolvedValueOnce(makePage([makeModel('no-provider')], 1, 2)) + .mockResolvedValueOnce( + makePage([makeModel('served', { model_providers: ['ws1/build'] })], 2, 2) + ); + + const { result } = renderSearch({ include: (model) => !!model.model_providers?.length }); + + await waitFor(() => expect(result.current.models.map((m) => m.name)).toEqual(['served'])); + expect(mockListModels).toHaveBeenCalledTimes(2); + expect(result.current.hasMore).toBe(false); + }); + + it('appends the next page on demand', async () => { + mockListModels + .mockResolvedValueOnce(makePage([makeModel('a')], 1, 2)) + .mockResolvedValueOnce(makePage([makeModel('b')], 2, 2)); + + const { result } = renderSearch(); + await waitFor(() => expect(result.current.hasMore).toBe(true)); + + await act(async () => { + await result.current.onLoadMore(); + }); + + await waitFor(() => expect(result.current.models.map((m) => m.name)).toEqual(['a', 'b'])); + expect(result.current.hasMore).toBe(false); + }); +}); diff --git a/web/packages/common/src/api/models/useModelSearch.ts b/web/packages/common/src/api/models/useModelSearch.ts new file mode 100644 index 0000000000..1d848d5e13 --- /dev/null +++ b/web/packages/common/src/api/models/useModelSearch.ts @@ -0,0 +1,131 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import type { WithFilterOperators } from '@nemo/common/src/api/filterOperators'; +import { useModelsInfinite, type ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels'; +import { groupModelsByWorkspace } from '@nemo/common/src/utils/models'; +import { + type ModelEntity, + ModelEntitySortField, + type ModelEntityFilter, +} from '@nemo/sdk/generated/platform/schema'; +import { useCallback, useEffect, useMemo, useState } from 'react'; + +/** + * Page size for search-as-you-type model lists. Small on purpose: the dropdown pulls the next + * page as the user scrolls, so the first page needs to arrive fast, not be complete. + */ +export const MODEL_SEARCH_PAGE_SIZE = 25; + +export type ModelSearchFilter = WithFilterOperators; + +export interface UseModelSearchOptions { + /** Workspace to search. The query stays idle while this is null. */ + workspace: string | null; + /** Extra filters merged into the request (e.g. `lora_enabled`, `base_model`). */ + filter?: ModelSearchFilter; + sort?: ModelEntitySortField; + pageSize?: number; + enabled?: boolean; + /** + * Client-side predicate applied to every page — for conditions the API cannot express, such as + * "has a ready deployment" (`model_providers.length > 0`). The hook keeps paging while a page + * filters down to nothing, so an excluded page never stalls the list. + */ + include?: (model: ModelEntity) => boolean; +} + +/** + * Props for `ModelSelectV2`, ready to spread. Every field lines up with a prop name so a caller + * that needs nothing custom is a single line. + */ +export interface ModelSearchProps { + groups: ModelWorkspaceGroup[]; + loading: boolean; + onSearchChange: (search: string) => void; + onLoadMore: () => Promise; + hasMore: boolean; + isLoadingMore: boolean; +} + +export interface UseModelSearchResult extends ModelSearchProps { + models: ModelEntity[]; + search: string; + error: Error | null; +} + +/** + * Server-side model search with progressive paging — the counterpart to `useAllModels`, which + * walks every page up front. Filtering happens in the API and pages arrive as the user scrolls, + * so a workspace with thousands of models costs one small request at a time. + * + * @example + * const [open, setOpen] = useState(false); + * const models = useModelSearch({ workspace, enabled: open }); + * return ; + */ +export const useModelSearch = ({ + workspace, + filter, + sort = ModelEntitySortField.name, + pageSize = MODEL_SEARCH_PAGE_SIZE, + enabled = true, + include, +}: UseModelSearchOptions): UseModelSearchResult => { + const [search, setSearch] = useState(''); + + const query = useMemo(() => { + const trimmed = search.trim(); + const merged: ModelSearchFilter = { + ...filter, + ...(trimmed ? { name: { $like: trimmed } } : {}), + }; + return { + page_size: pageSize, + sort, + ...(Object.keys(merged).length > 0 ? { filter: merged as ModelEntityFilter } : {}), + }; + }, [filter, pageSize, search, sort]); + + const isEnabled = enabled && !!workspace; + const { data, error, fetchNextPage, hasNextPage, isFetchingNextPage, isLoading } = + useModelsInfinite({ + workspace: workspace ?? undefined, + query, + queryOptions: { enabled: isEnabled }, + }); + + const models = useMemo(() => { + const loaded = data?.pages.flatMap((page) => page.data ?? []) ?? []; + return include ? loaded.filter(include) : loaded; + }, [data?.pages, include]); + + const groups = useMemo(() => groupModelsByWorkspace(models, { sort: true }), [models]); + + const hasMore = !!hasNextPage; + + const onLoadMore = useCallback(async () => { + if (!hasNextPage || isFetchingNextPage) return; + await fetchNextPage(); + }, [fetchNextPage, hasNextPage, isFetchingNextPage]); + + // `include` can empty a whole page, leaving the list with no rows to scroll and therefore no way + // to ask for the next one. Keep paging until something survives the filter. + useEffect(() => { + if (isEnabled && models.length === 0 && hasNextPage && !isFetchingNextPage && !isLoading) { + void fetchNextPage(); + } + }, [fetchNextPage, hasNextPage, isEnabled, isFetchingNextPage, isLoading, models.length]); + + return { + models, + groups, + search, + error, + loading: isLoading, + onSearchChange: setSearch, + onLoadMore, + hasMore, + isLoadingMore: isFetchingNextPage, + }; +}; diff --git a/web/packages/common/src/components/ModelSelectV2/ModelDropdown.test.tsx b/web/packages/common/src/components/ModelSelectV2/ModelDropdown.test.tsx new file mode 100644 index 0000000000..f8125ca0f3 --- /dev/null +++ b/web/packages/common/src/components/ModelSelectV2/ModelDropdown.test.tsx @@ -0,0 +1,119 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import type { ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels'; +import { ModelDropdown } from '@nemo/common/src/components/ModelSelectV2/ModelDropdown'; +import type { ModelEntity } from '@nemo/sdk/generated/platform/schema'; +import { fireEvent, render, screen, waitFor } from '@testing-library/react'; + +const makeModel = (name: string, workspace = 'nvidia'): ModelEntity => + ({ id: name, name, workspace }) as unknown as ModelEntity; + +const groups: ModelWorkspaceGroup[] = [ + { workspace: 'nvidia', models: [makeModel('nemotron-8b'), makeModel('llama-3.1-8b')] }, +]; + +type Props = React.ComponentProps; + +const renderOpen = (props: Partial = {}) => + render( + + ); + +const typeFilter = (text: string) => + fireEvent.change(screen.getByTestId('model-select-v2-filter'), { target: { value: text } }); + +/** Names of the rows in the list, ignoring the details panel each row also renders. */ +const listedModels = () => + screen.queryAllByTestId('model-dropdown-item').map((item) => item.textContent); + +describe('ModelDropdown', () => { + describe('search', () => { + it('filters the given groups itself when no onSearchChange is provided', async () => { + renderOpen(); + await waitFor(() => expect(listedModels()).toHaveLength(2)); + + typeFilter('llama'); + + await waitFor(() => expect(listedModels()).toEqual(['llama-3.1-8b'])); + }); + + it('reports the debounced term and leaves the groups alone when onSearchChange is provided', async () => { + const onSearchChange = vi.fn(); + renderOpen({ onSearchChange, searchDebounceMs: 0 }); + + typeFilter('llama'); + + await waitFor(() => expect(onSearchChange).toHaveBeenCalledWith('llama')); + // The caller owns the query, so both models stay listed until new groups arrive. + expect(listedModels()).toEqual(['nemotron-8b', 'llama-3.1-8b']); + }); + }); + + describe('paging', () => { + it('asks for the next page while more remain', async () => { + const onLoadMore = vi.fn(); + renderOpen({ onLoadMore, hasMore: true }); + + await waitFor(() => expect(onLoadMore).toHaveBeenCalled()); + }); + + it('does not ask for more once the last page has loaded', async () => { + const onLoadMore = vi.fn(); + renderOpen({ onLoadMore, hasMore: false }); + + await waitFor(() => expect(listedModels()).toHaveLength(2)); + expect(onLoadMore).not.toHaveBeenCalled(); + }); + + it('shows the done message only after the last page', async () => { + const { rerender } = renderOpen({ + onLoadMore: vi.fn(), + hasMore: true, + doneLoadingMessage: 'No more models', + }); + await waitFor(() => expect(listedModels()).toHaveLength(2)); + expect(screen.queryByText('No more models')).not.toBeInTheDocument(); + + rerender( + + ); + + expect(await screen.findByText('No more models')).toBeInTheDocument(); + }); + }); + + describe('trigger', () => { + it('falls back to the name in the URN when the selection is not in the loaded pages', () => { + renderOpen({ open: false, value: { model: 'nvidia/not-yet-loaded' }, groups: [] }); + + expect(screen.getByTestId('model-select-v2-trigger')).toHaveTextContent('not-yet-loaded'); + }); + + it('prefers the entity the selection carries', () => { + renderOpen({ + open: false, + value: { model: 'nvidia/nemotron-8b', entity: makeModel('nemotron-8b') }, + groups: [], + }); + + expect(screen.getByTestId('model-select-v2-trigger')).toHaveTextContent('nemotron-8b'); + }); + }); +}); diff --git a/web/packages/common/src/components/ModelSelectV2/ModelDropdown.tsx b/web/packages/common/src/components/ModelSelectV2/ModelDropdown.tsx index a49831bb4c..a6bd3286de 100644 --- a/web/packages/common/src/components/ModelSelectV2/ModelDropdown.tsx +++ b/web/packages/common/src/components/ModelSelectV2/ModelDropdown.tsx @@ -1,61 +1,64 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -import type { ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels'; -import { ModelDropdownItem } from '@nemo/common/src/components/ModelSelectV2/ModelDropdownItem'; +import { ModelDropdownList } from '@nemo/common/src/components/ModelSelectV2/ModelDropdownList'; import { ModelDropdownSearch } from '@nemo/common/src/components/ModelSelectV2/ModelDropdownSearch'; -import type { ModelSelection, ModelType } from '@nemo/common/src/components/ModelSelectV2/types'; +import type { + ModelSelectV2Props, + ModelSelection, + ModelType, +} from '@nemo/common/src/components/ModelSelectV2/types'; import { creatorToIcon } from '@nemo/common/src/constants/modelMetadata'; -import { getURNFromNamedEntityRef } from '@nemo/common/src/namedEntity'; +import { getPartsFromReference, getURNFromNamedEntityRef } from '@nemo/common/src/namedEntity'; import { filterModel, isBaseModel } from '@nemo/common/src/utils/models'; import type { ModelEntity } from '@nemo/sdk/generated/platform/schema'; import { Button, DropdownContent, - DropdownHeading, DropdownRoot, - DropdownSection, DropdownTrigger, Flex, SegmentedControl, - Stack, Text, } from '@nvidia/foundations-react-core'; import { ChevronDown, LoaderCircle } from 'lucide-react'; -import { useMemo, useState, type FC } from 'react'; +import { useEffect, useMemo, useState, type FC } from 'react'; +import { useDebounce } from 'use-debounce'; const MODEL_TYPE_ITEMS = [ { value: 'custom', children: 'Custom Models' }, { value: 'base', children: 'Base Models' }, ]; +const DEFAULT_SEARCH_DEBOUNCE_MS = 300; + const isCustomModel = (model: ModelEntity): boolean => !isBaseModel(model); -interface ModelDropdownProps { - value: ModelSelection | null; - onValueChange: (selection: ModelSelection) => void; - groups: ModelWorkspaceGroup[]; - loading?: boolean; - disabled?: boolean; - placeholder?: string; - showModelTypeToggle?: boolean; - defaultModelType?: ModelType; - hideAdapters?: boolean; - fullWidth?: boolean; - dropdownSide?: 'top' | 'bottom'; +type ModelDropdownProps = Omit< + ModelSelectV2Props, + 'showParams' | 'inferenceParams' | 'onInferenceParamsChange' | 'onOpenChange' | 'aria-label' +> & { open: boolean; onOpenChange: (open: boolean) => void; -} +}; export const ModelDropdown: FC = ({ value, onValueChange, groups, + onSearchChange, + searchDebounceMs = DEFAULT_SEARCH_DEBOUNCE_MS, + onLoadMore, + hasMore = false, + isLoadingMore = false, + doneLoadingMessage, + emptyMessage, loading = false, disabled = false, placeholder = 'Select a model', showModelTypeToggle = false, defaultModelType = 'custom', + onModelTypeChange, hideAdapters = false, fullWidth = false, dropdownSide = 'bottom', @@ -64,37 +67,51 @@ export const ModelDropdown: FC = ({ }) => { const [search, setSearch] = useState(''); const [modelType, setModelType] = useState(defaultModelType); + const [debouncedSearch] = useDebounce(search, searchDebounceMs); + + useEffect(() => { + onSearchChange?.(debouncedSearch); + }, [debouncedSearch, onSearchChange]); + + const localGroups = useMemo(() => groups ?? [], [groups]); const selectedModel = useMemo(() => { if (!value) return undefined; - return groups.flatMap((g) => g.models).find((m) => getURNFromNamedEntityRef(m) === value.model); - }, [groups, value]); + if (value.entity) return value.entity; + return localGroups + .flatMap((g) => g.models) + .find((m) => getURNFromNamedEntityRef(m) === value.model); + }, [localGroups, value]); const filteredGroups = useMemo(() => { - return groups + const filterType = showModelTypeToggle && !onModelTypeChange; + const filterSearch = !onSearchChange && search.length > 0; + if (!filterType && !filterSearch) return localGroups; + + return localGroups .map((group) => { let models = group.models; - - // Apply model type filter - if (showModelTypeToggle) { + if (filterType) { models = modelType === 'base' ? models.filter(isBaseModel) : models.filter(isCustomModel); } - - // Apply search filter - if (search) { + if (filterSearch) { models = models.filter((m) => filterModel(m, search)); } - return { ...group, models }; }) .filter((group) => group.models.length > 0); - }, [groups, showModelTypeToggle, modelType, search]); + }, [localGroups, modelType, onModelTypeChange, onSearchChange, search, showModelTypeToggle]); const handleSelect = (selection: ModelSelection) => { onValueChange(selection); onOpenChange(false); }; + const handleModelTypeChange = (val: string) => { + setModelType(val as ModelType); + onModelTypeChange?.(val as ModelType); + }; + const handleOpenChange = (nextOpen: boolean) => { onOpenChange(nextOpen); if (!nextOpen) { @@ -102,9 +119,10 @@ export const ModelDropdown: FC = ({ } }; - const triggerLabel = selectedModel - ? (selectedModel.name?.split('@')[0] ?? selectedModel.name) - : placeholder; + const selectedParts = value?.model ? getPartsFromReference(value.model) : undefined; + const selectedName = selectedModel?.name ?? selectedParts?.name; + const selectedWorkspace = selectedModel?.workspace ?? selectedParts?.workspace; + const triggerLabel = selectedName ? (selectedName.split('@')[0] ?? selectedName) : placeholder; return ( @@ -118,7 +136,7 @@ export const ModelDropdown: FC = ({ disabled={disabled} aria-label="Select a model" data-testid="model-select-v2-trigger" - className="overflow-hidden [&[data-state=open]]:border-[var(--border-color-feedback-success)] [&[data-state=open]]:bg-[var(--background-color-interaction-base)]" + className="overflow-hidden data-[state=open]:border-(--border-color-feedback-success) data-[state=open]:bg-(--background-color-interaction-base)" > = ({ className={`min-w-0 ${fullWidth ? 'w-full justify-between' : ''}`} > - {selectedModel && - creatorToIcon(selectedModel.workspace ?? '', { - className: 'text-base flex-shrink-0', - })} - {loading ? ( + {selectedWorkspace && + creatorToIcon(selectedWorkspace, { className: 'text-base flex-shrink-0' })} + {loading && !selectedName ? ( <> - + {placeholder} ) : ( {triggerLabel} )} - + @@ -157,39 +173,22 @@ export const ModelDropdown: FC = ({ className="w-full" value={modelType} items={MODEL_TYPE_ITEMS} - onValueChange={(val: string) => setModelType(val as ModelType)} + onValueChange={handleModelTypeChange} /> )} - - {filteredGroups.length > 0 ? ( - filteredGroups.map((group) => ( - - - - {creatorToIcon(group.workspace, { className: 'text-base' })} - {group.workspace} - - - {group.models.map((model) => ( - - ))} - - )) - ) : ( - - - {loading ? 'Loading models...' : 'No models found'} - - - )} - + ); diff --git a/web/packages/common/src/components/ModelSelectV2/ModelDropdownItem.tsx b/web/packages/common/src/components/ModelSelectV2/ModelDropdownItem.tsx index 5aa03d9876..511c0a50d8 100644 --- a/web/packages/common/src/components/ModelSelectV2/ModelDropdownItem.tsx +++ b/web/packages/common/src/components/ModelSelectV2/ModelDropdownItem.tsx @@ -62,11 +62,11 @@ const AdapterItem: FC<{ onSelect({ model: modelUrn, adapter: adapter.name })} + onSelect={() => onSelect({ model: modelUrn, adapter: adapter.name, entity: model })} > - {isSelected && } + {isSelected && } {adapter.name} {adapter.created_at && ( @@ -99,7 +99,7 @@ export const ModelDropdownItem: FC = ({ onSelect({ model: modelUrn })} + onClick={() => onSelect({ model: modelUrn, entity: model })} > @@ -128,10 +128,10 @@ export const ModelDropdownItem: FC = ({ Base Model onSelect({ model: modelUrn })} + onSelect={() => onSelect({ model: modelUrn, entity: model })} > - {isBaseSelected && } + {isBaseSelected && } {modelUrn} diff --git a/web/packages/common/src/components/ModelSelectV2/ModelDropdownList.tsx b/web/packages/common/src/components/ModelSelectV2/ModelDropdownList.tsx new file mode 100644 index 0000000000..5a107be1be --- /dev/null +++ b/web/packages/common/src/components/ModelSelectV2/ModelDropdownList.tsx @@ -0,0 +1,181 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import type { ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels'; +import { ModelDropdownItem } from '@nemo/common/src/components/ModelSelectV2/ModelDropdownItem'; +import type { ModelSelection } from '@nemo/common/src/components/ModelSelectV2/types'; +import { creatorToIcon } from '@nemo/common/src/constants/modelMetadata'; +import { getURNFromNamedEntityRef } from '@nemo/common/src/namedEntity'; +import type { ModelEntity } from '@nemo/sdk/generated/platform/schema'; +import { DropdownHeading, Flex, Text } from '@nvidia/foundations-react-core'; +import { useVirtualizer } from '@tanstack/react-virtual'; +import { LoaderCircle } from 'lucide-react'; +import { useCallback, useEffect, useLayoutEffect, useMemo, useRef, useState, type FC } from 'react'; + +/** Rows the virtualizer measures: group headings and models share one flat index space. */ +type ModelRow = + | { kind: 'heading'; key: string; workspace: string } + | { kind: 'model'; key: string; model: ModelEntity }; + +const HEADING_HEIGHT = 32; +const ITEM_HEIGHT = 36; +const LIST_MAX_HEIGHT = 300; +const OVERSCAN = 8; +const LOAD_MORE_THRESHOLD = 5; + +/** Assumed viewport before the first measurement, so the initial render is already windowed. */ +const INITIAL_RECT = { width: 360, height: LIST_MAX_HEIGHT }; + +export interface ModelDropdownListProps { + groups: ModelWorkspaceGroup[]; + value: ModelSelection | null; + onSelect: (selection: ModelSelection) => void; + hideAdapters?: boolean; + loading?: boolean; + /** Called as the user scrolls near the end; no-op when {@link hasMore} is false. */ + onLoadMore?: () => void | Promise; + hasMore?: boolean; + isLoadingMore?: boolean; + /** Footer copy once every page has loaded. Omit to render nothing. */ + doneLoadingMessage?: string; + emptyMessage?: string; +} + +/** + * The scrollable body of the model dropdown, virtualized so only the visible slice of a + * workspace's models is in the DOM. Each item carries a submenu with a details panel, so an + * unvirtualized list of a few hundred models mounts thousands of nodes on open — this keeps + * that cost proportional to the viewport instead of the catalogue. + */ +export const ModelDropdownList: FC = ({ + groups, + value, + onSelect, + hideAdapters = false, + loading = false, + onLoadMore, + hasMore = false, + isLoadingMore = false, + doneLoadingMessage, + emptyMessage = 'No models found', +}) => { + const scrollRef = useRef(null); + const [hasViewport, setHasViewport] = useState(true); + const [isLoadingMoreLocal, setIsLoadingMoreLocal] = useState(false); + + useLayoutEffect(() => { + setHasViewport((scrollRef.current?.clientHeight ?? 0) > 0); + }, []); + + const rows = useMemo( + () => + groups.flatMap((group) => [ + { kind: 'heading' as const, key: `heading-${group.workspace}`, workspace: group.workspace }, + ...group.models.map((model) => ({ + kind: 'model' as const, + key: getURNFromNamedEntityRef(model) ?? `${group.workspace}/${model.name}`, + model, + })), + ]), + [groups] + ); + + const virtualizer = useVirtualizer({ + count: rows.length, + getScrollElement: () => scrollRef.current, + estimateSize: (index) => (rows[index]?.kind === 'heading' ? HEADING_HEIGHT : ITEM_HEIGHT), + overscan: OVERSCAN, + initialRect: INITIAL_RECT, + }); + + const virtualItems = virtualizer.getVirtualItems(); + + const loadMore = useCallback(async () => { + if (!onLoadMore || !hasMore || isLoadingMore || isLoadingMoreLocal) return; + setIsLoadingMoreLocal(true); + try { + await onLoadMore(); + } finally { + setIsLoadingMoreLocal(false); + } + }, [hasMore, isLoadingMore, isLoadingMoreLocal, onLoadMore]); + + useEffect(() => { + if (virtualItems.length === 0) return; + const lastRow = virtualItems[virtualItems.length - 1]; + if (lastRow.index >= rows.length - LOAD_MORE_THRESHOLD) { + void loadMore(); + } + }, [loadMore, rows.length, virtualItems]); + + const loadingMore = isLoadingMore || isLoadingMoreLocal; + const showDoneMessage = Boolean(doneLoadingMessage) && !hasMore && !loading && !loadingMore; + + const renderRow = (row: ModelRow) => + row.kind === 'heading' ? ( + + + {creatorToIcon(row.workspace, { className: 'text-base' })} + {row.workspace} + + + ) : ( + + ); + + if (rows.length === 0) { + return ( + + {loading ? 'Loading models...' : emptyMessage} + + ); + } + + return ( +
+ {hasViewport ? ( +
+ {virtualItems.map((virtualRow) => ( +
+ {renderRow(rows[virtualRow.index])} +
+ ))} +
+ ) : ( + rows.map((row) =>
{renderRow(row)}
) + )} + + {loadingMore && ( + + + + Loading more… + + + )} + {showDoneMessage && ( + + + {doneLoadingMessage} + + + )} +
+ ); +}; diff --git a/web/packages/common/src/components/ModelSelectV2/ModelSelectV2.stories.tsx b/web/packages/common/src/components/ModelSelectV2/ModelSelectV2.stories.tsx index a755107a08..e3041d318a 100644 --- a/web/packages/common/src/components/ModelSelectV2/ModelSelectV2.stories.tsx +++ b/web/packages/common/src/components/ModelSelectV2/ModelSelectV2.stories.tsx @@ -97,6 +97,18 @@ const mockGroups: ModelWorkspaceGroup[] = [ }, ]; +const groupByWorkspace = (models: ModelWorkspaceGroup['models']): ModelWorkspaceGroup[] => { + const byWorkspace = new Map(); + for (const model of models) { + const workspace = model.workspace ?? 'default'; + byWorkspace.set(workspace, [...(byWorkspace.get(workspace) ?? []), model]); + } + return Array.from(byWorkspace, ([workspace, groupModels]) => ({ + workspace, + models: groupModels, + })); +}; + const meta: Meta = { component: ModelSelectV2, title: 'Studio Common/ModelSelectV2', @@ -347,3 +359,56 @@ const HideAdaptersRender = (args: ModelSelectV2Props) => { export const HideAdapters: Story = { render: HideAdaptersRender, }; + +const allProgressiveModels = manyModelsGroups.flatMap((group) => group.models); +const PROGRESSIVE_PAGE_SIZE = 6; +const PROGRESSIVE_LATENCY_MS = 600; + +/** + * The paged contract `useModelSearch` implements, stubbed with a local list: the filter box + * reports its debounced value instead of filtering in place, and pages arrive as the list scrolls. + */ +const ProgressiveRender = (args: ModelSelectV2Props) => { + const [value, setValue] = useState(null); + const [search, setSearch] = useState(''); + const [pageCount, setPageCount] = useState(1); + const [isLoadingMore, setIsLoadingMore] = useState(false); + + const matches = allProgressiveModels.filter((model) => + (model.name ?? '').toLowerCase().includes(search.toLowerCase()) + ); + const loaded = matches.slice(0, pageCount * PROGRESSIVE_PAGE_SIZE); + const hasMore = loaded.length < matches.length; + + const handleSearchChange = (next: string) => { + setSearch(next); + setPageCount(1); + }; + + const handleLoadMore = async () => { + setIsLoadingMore(true); + await new Promise((resolve) => setTimeout(resolve, PROGRESSIVE_LATENCY_MS)); + setPageCount((count) => count + 1); + setIsLoadingMore(false); + }; + + return ( + + ); +}; + +export const ProgressiveLoading: Story = { + render: ProgressiveRender, +}; diff --git a/web/packages/common/src/components/ModelSelectV2/ModelSelectV2.tsx b/web/packages/common/src/components/ModelSelectV2/ModelSelectV2.tsx index 47c4e20d13..6df2eb8a07 100644 --- a/web/packages/common/src/components/ModelSelectV2/ModelSelectV2.tsx +++ b/web/packages/common/src/components/ModelSelectV2/ModelSelectV2.tsx @@ -8,23 +8,14 @@ import { Group } from '@nvidia/foundations-react-core'; import { FC, useState } from 'react'; export const ModelSelectV2: FC = ({ - value, - onValueChange, - groups, - loading, - disabled, - placeholder, - showModelTypeToggle, - defaultModelType, showParams = false, - hideAdapters = false, - fullWidth = false, - dropdownSide, inferenceParams, onInferenceParamsChange, onOpenChange, 'aria-label': ariaLabel, + ...dropdownProps }) => { + const { disabled, fullWidth = false } = dropdownProps; const [modelOpen, setModelOpen] = useState(false); const [paramsOpen, setParamsOpen] = useState(false); @@ -40,21 +31,7 @@ export const ModelSelectV2: FC = ({ }; const modelDropdown = ( - + ); if (!showParams) return modelDropdown; diff --git a/web/packages/common/src/components/ModelSelectV2/WorkspaceModelSelect.tsx b/web/packages/common/src/components/ModelSelectV2/WorkspaceModelSelect.tsx new file mode 100644 index 0000000000..82a1ae9e55 --- /dev/null +++ b/web/packages/common/src/components/ModelSelectV2/WorkspaceModelSelect.tsx @@ -0,0 +1,68 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import { + useModelSearch, + type ModelSearchFilter, + type ModelSearchProps, +} from '@nemo/common/src/api/models/useModelSearch'; +import { ModelSelectV2 } from '@nemo/common/src/components/ModelSelectV2/ModelSelectV2'; +import type { ModelSelectV2Props } from '@nemo/common/src/components/ModelSelectV2/types'; +import type { ModelEntity } from '@nemo/sdk/generated/platform/schema'; +import { useState, type FC } from 'react'; + +export interface WorkspaceModelSelectProps extends Omit< + ModelSelectV2Props, + keyof ModelSearchProps +> { + /** Workspace to search. Nothing is fetched while this is null. */ + workspace: string | null; + /** Extra filters merged into the search request (e.g. `lora_enabled`). */ + filter?: ModelSearchFilter; + /** Client-side predicate for conditions the API cannot express (e.g. "has a deployment"). */ + include?: (model: ModelEntity) => boolean; + /** Hold the request back for reasons of the caller's own, on top of the open check. */ + enabled?: boolean; +} + +/** + * `ModelSelectV2` wired to its own paged search: nothing is requested until the user opens the + * menu, the filter box queries the API, and further pages arrive as the list scrolls. + * + * Reach for `useModelSearch` directly only when the caller needs the models themselves — to merge + * in a pinned entry, or to surface the query error. + */ +export const WorkspaceModelSelect: FC = ({ + workspace, + filter, + include, + enabled = true, + onOpenChange, + ...selectProps +}) => { + const [open, setOpen] = useState(false); + const { groups, loading, onSearchChange, onLoadMore, hasMore, isLoadingMore } = useModelSearch({ + workspace, + filter, + include, + enabled: enabled && open, + }); + + const handleOpenChange = (nextOpen: boolean) => { + setOpen(nextOpen); + onOpenChange?.(nextOpen); + }; + + return ( + + ); +}; diff --git a/web/packages/common/src/components/ModelSelectV2/index.tsx b/web/packages/common/src/components/ModelSelectV2/index.tsx index 0104a8df58..e7f386fb80 100644 --- a/web/packages/common/src/components/ModelSelectV2/index.tsx +++ b/web/packages/common/src/components/ModelSelectV2/index.tsx @@ -2,6 +2,8 @@ // SPDX-License-Identifier: Apache-2.0 export { ModelSelectV2 } from '@nemo/common/src/components/ModelSelectV2/ModelSelectV2'; +export { WorkspaceModelSelect } from '@nemo/common/src/components/ModelSelectV2/WorkspaceModelSelect'; +export type { WorkspaceModelSelectProps } from '@nemo/common/src/components/ModelSelectV2/WorkspaceModelSelect'; export type { ModelSelectV2Props, ModelSelection, diff --git a/web/packages/common/src/components/ModelSelectV2/types.ts b/web/packages/common/src/components/ModelSelectV2/types.ts index 3747d4e35a..6f6ab53513 100644 --- a/web/packages/common/src/components/ModelSelectV2/types.ts +++ b/web/packages/common/src/components/ModelSelectV2/types.ts @@ -2,13 +2,19 @@ // SPDX-License-Identifier: Apache-2.0 import type { ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels'; -import type { InferenceParams } from '@nemo/sdk/generated/platform/schema'; +import type { InferenceParams, ModelEntity } from '@nemo/sdk/generated/platform/schema'; export interface ModelSelection { /** Model URN (e.g. "workspace/model_name") */ model: string; /** Adapter name, if an adapter was selected instead of the base model */ adapter?: string; + /** + * The catalogue entry the user picked, when the selection came from the dropdown. Saves + * callers a lookup (or a second request) for fields like `model_providers`; absent when the + * selection was restored from a URL, a form default, or any other source outside the list. + */ + entity?: ModelEntity; } export interface ModelSelectV2Props { @@ -16,9 +22,31 @@ export interface ModelSelectV2Props { value: ModelSelection | null; /** Called when user selects a model or adapter */ onValueChange: (selection: ModelSelection) => void; - /** Models grouped by workspace */ - groups: ModelWorkspaceGroup[]; - /** Whether models are still loading */ + /** + * The page of models to show, grouped by workspace. The dropdown renders exactly what it is + * given — pair with {@link onSearchChange} and {@link onLoadMore} (see `useModelSearch`) so a + * large catalogue arrives one page at a time instead of all at once. + */ + groups?: ModelWorkspaceGroup[]; + /** + * Server-side search. When set, the dropdown stops filtering {@link groups} itself and instead + * reports the debounced filter text; the caller re-queries and passes back new `groups`. + * Leave unset to filter the given groups client-side. + */ + onSearchChange?: (search: string) => void; + /** How long the filter box settles before {@link onSearchChange} fires. */ + searchDebounceMs?: number; + /** Called as the list scrolls near its end. Ignored unless {@link hasMore} is true. */ + onLoadMore?: () => void | Promise; + /** Whether another page is available for {@link onLoadMore} to fetch. */ + hasMore?: boolean; + /** Whether the next page is in flight. */ + isLoadingMore?: boolean; + /** Footer copy shown once every page has loaded. Omit to render nothing. */ + doneLoadingMessage?: string; + /** Copy shown when no models match. */ + emptyMessage?: string; + /** Whether the first page of models is still loading */ loading?: boolean; /** Whether the component is disabled */ disabled?: boolean; @@ -27,7 +55,12 @@ export interface ModelSelectV2Props { /** Show the Custom/Base segmented control toggle */ showModelTypeToggle?: boolean; /** Default active segment when toggle is shown */ - defaultModelType?: 'custom' | 'base'; + defaultModelType?: ModelType; + /** + * Called when the model-type segment changes. Same contract as {@link onSearchChange}: when set, + * the caller owns the filter and the dropdown stops applying it to {@link groups}. + */ + onModelTypeChange?: (modelType: ModelType) => void; /** Show the params button alongside the model button */ showParams?: boolean; /** diff --git a/web/packages/common/src/utils/models.ts b/web/packages/common/src/utils/models.ts index 681ca824c0..62d379291f 100644 --- a/web/packages/common/src/utils/models.ts +++ b/web/packages/common/src/utils/models.ts @@ -35,6 +35,9 @@ export const groupModelsByWorkspace = ( return entries.map(([ws, ms]) => buildWorkspaceGroup(ws, ms)); }; +export const hasModelProvider = (model: ModelEntity): boolean => + Array.isArray(model.model_providers) && model.model_providers.length > 0; + /** * Returns true if the model is a base model, false otherwise. This is determined by * checking if the model has a base_model property. diff --git a/web/packages/studio/src/components/AddModelPalette/AddModelPalette.stories.tsx b/web/packages/studio/src/components/AddModelPalette/AddModelPalette.stories.tsx index 87764a58b8..99fa2f415c 100644 --- a/web/packages/studio/src/components/AddModelPalette/AddModelPalette.stories.tsx +++ b/web/packages/studio/src/components/AddModelPalette/AddModelPalette.stories.tsx @@ -12,7 +12,7 @@ const meta = { layout: 'fullscreen', }, args: { - modelGroups: [], + workspace: 'default', onAddModel: () => {}, onSelectModel: () => {}, }, diff --git a/web/packages/studio/src/components/AddModelPalette/index.tsx b/web/packages/studio/src/components/AddModelPalette/index.tsx index 6cd5fbbb4c..15001b72b9 100644 --- a/web/packages/studio/src/components/AddModelPalette/index.tsx +++ b/web/packages/studio/src/components/AddModelPalette/index.tsx @@ -1,14 +1,13 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -import type { ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels'; -import { ModelSelectV2 } from '@nemo/common/src/components/ModelSelectV2/ModelSelectV2'; import type { ModelSelection } from '@nemo/common/src/components/ModelSelectV2/types'; +import { WorkspaceModelSelect } from '@nemo/common/src/components/ModelSelectV2/WorkspaceModelSelect'; import { Stack, Text } from '@nvidia/foundations-react-core'; import { CardIconBadge, SelectableCard } from '@studio/components/common/SelectableCard'; import { type BuilderModel, - providerForModel, + providerForSelection, } from '@studio/routes/DataDesignerJobBuildRoute/models'; import { Cpu } from 'lucide-react'; import type { FC } from 'react'; @@ -16,8 +15,7 @@ import type { FC } from 'react'; export interface AddModelPaletteProps { models: BuilderModel[]; selectedId?: string | null; - modelGroups: ModelWorkspaceGroup[]; - isLoadingModels?: boolean; + workspace: string; onAddModel: (selection: ModelSelection, provider: string) => void; onSelectModel: (id: string) => void; className?: string; @@ -25,8 +23,7 @@ export interface AddModelPaletteProps { export const AddModelPalette: FC = ({ models, selectedId, - modelGroups, - isLoadingModels, + workspace, onAddModel, onSelectModel, className, @@ -40,13 +37,10 @@ export const AddModelPalette: FC = ({
- - onAddModel(selection, providerForModel(modelGroups, selection.model)) - } - groups={modelGroups} - loading={isLoadingModels} + onValueChange={(selection) => onAddModel(selection, providerForSelection(selection))} placeholder="Add a model" fullWidth dropdownSide="bottom" diff --git a/web/packages/studio/src/components/ModelChatPanel/ModelChatPanel.test.tsx b/web/packages/studio/src/components/ModelChatPanel/ModelChatPanel.test.tsx index da2903a6a5..2f8ea01a7d 100644 --- a/web/packages/studio/src/components/ModelChatPanel/ModelChatPanel.test.tsx +++ b/web/packages/studio/src/components/ModelChatPanel/ModelChatPanel.test.tsx @@ -1,8 +1,6 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -import type { ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels'; -import type { ModelEntity } from '@nemo/sdk/generated/platform/schema'; import { ModelChatPanel } from '@studio/components/ModelChatPanel'; import { TestProviders } from '@studio/tests/util/TestProviders'; import { render } from '@testing-library/react'; @@ -19,25 +17,12 @@ vi.mock('@studio/components/ModelChat', () => ({ }, })); -// ModelSelectV2 internals are not what we're testing here. +// The model select fetches its own options; its internals are not what we're testing here. vi.mock('@nemo/common/src/components/ModelSelectV2', () => ({ - ModelSelectV2: () =>
, + WorkspaceModelSelect: () =>
, })); -const makeModel = (workspace: string, name: string): ModelEntity => - ({ workspace, name }) as unknown as ModelEntity; - -const makeGroups = (models: ModelEntity[]): ModelWorkspaceGroup[] => { - const byWorkspace = new Map(); - for (const m of models) { - const ws = m.workspace ?? ''; - if (!byWorkspace.has(ws)) byWorkspace.set(ws, []); - byWorkspace.get(ws)!.push(m); - } - return Array.from(byWorkspace.entries()).map(([workspace, models]) => ({ workspace, models })); -}; - -const renderPanel = (modelURN: string | null, modelGroups: ModelWorkspaceGroup[]) => { +const renderPanel = (modelURN: string | null) => { return render( @@ -52,8 +37,6 @@ const renderPanel = (modelURN: string | null, modelGroups: ModelWorkspaceGroup[] locked: false, }} fallbackWorkspace="route-workspace" - modelGroups={modelGroups} - isLoadingModels={false} onToggle={vi.fn()} onRemove={vi.fn()} onModelChange={vi.fn()} @@ -69,10 +52,7 @@ describe('ModelChatPanel — URN routing', () => { }); it("routes inference to the model's own workspace (not the route workspace)", () => { - renderPanel( - 'nvidia/llama-70b', - makeGroups([makeModel('abacusai', 'llama-70b'), makeModel('nvidia', 'llama-70b')]) - ); + renderPanel('nvidia/llama-70b'); expect(modelChatSpy).toHaveBeenCalledWith( expect.objectContaining({ workspace: 'nvidia', model: 'llama-70b' }) @@ -80,13 +60,9 @@ describe('ModelChatPanel — URN routing', () => { }); it('picks the correct workspace even when two models share the same name', () => { - // The previous name-based lookup would have silently bound this panel to - // whichever workspace's model came first in the list. With URNs end-to-end, - // the workspace selected in the URN is used. - renderPanel( - 'abacusai/llama-70b', - makeGroups([makeModel('nvidia', 'llama-70b'), makeModel('abacusai', 'llama-70b')]) - ); + // A name-based lookup would have silently bound this panel to whichever workspace's model + // came first in the list. With URNs end-to-end, the workspace in the URN is used. + renderPanel('abacusai/llama-70b'); expect(modelChatSpy).toHaveBeenCalledWith( expect.objectContaining({ workspace: 'abacusai', model: 'llama-70b' }) @@ -94,7 +70,7 @@ describe('ModelChatPanel — URN routing', () => { }); it('falls back to the route workspace and disables chat when no model is assigned', () => { - renderPanel(null, []); + renderPanel(null); // ModelChat still renders, but disabled and showing an empty state; with no // model URN it uses the route fallback workspace and an empty model id. expect(modelChatSpy).toHaveBeenCalledWith( diff --git a/web/packages/studio/src/components/ModelChatPanel/index.tsx b/web/packages/studio/src/components/ModelChatPanel/index.tsx index 6a2ddec242..7ced399986 100644 --- a/web/packages/studio/src/components/ModelChatPanel/index.tsx +++ b/web/packages/studio/src/components/ModelChatPanel/index.tsx @@ -1,9 +1,12 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -import type { ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels'; -import { ModelSelectV2, type ModelSelection } from '@nemo/common/src/components/ModelSelectV2'; +import { + WorkspaceModelSelect, + type ModelSelection, +} from '@nemo/common/src/components/ModelSelectV2'; import { getPartsFromReference } from '@nemo/common/src/namedEntity'; +import { hasModelProvider } from '@nemo/common/src/utils/models'; import { Flex, Stack, Text } from '@nvidia/foundations-react-core'; import { DEFAULT_INFERENCE_PARAMS, type InferenceParams } from '@studio/components/chat/params'; import { ParamsPopover } from '@studio/components/chat/ParamsPopover'; @@ -20,8 +23,6 @@ interface ModelChatPanelProps extends PanelChatControls { panel: PanelState; /** Fallback workspace used only if a panel has no model assigned yet. */ fallbackWorkspace: string; - modelGroups: ModelWorkspaceGroup[]; - isLoadingModels: boolean; onToggle: (id: number) => void; onRemove: (id: number) => void; /** Receives the full URN ("workspace/name"), or null when cleared. */ @@ -33,8 +34,6 @@ interface ModelChatPanelProps extends PanelChatControls { export const ModelChatPanel: FC = ({ panel, fallbackWorkspace, - modelGroups, - isLoadingModels, onToggle, onRemove, onModelChange, @@ -110,11 +109,11 @@ export const ModelChatPanel: FC = ({ {/* Model picker + inference params — shared across single and compare modes. */}
- void; onSetModel: (id: number, modelURN: string | null) => void; @@ -32,8 +29,6 @@ interface ModelCompareChatProps extends PanelChatControls { export const ModelCompareChat: FC = ({ workspace, - modelGroups, - isLoadingModels, models, onRemoveModel, onSetModel, @@ -113,8 +108,6 @@ export const ModelCompareChat: FC = ({ key={`${panel.id}-${chatResetCount ?? 0}`} panel={panel} fallbackWorkspace={workspace} - modelGroups={modelGroups} - isLoadingModels={isLoadingModels} onToggle={togglePanel} onRemove={onRemoveModel} onModelChange={onSetModel} diff --git a/web/packages/studio/src/components/ModelComparePrompts/ModelColumnSelect.tsx b/web/packages/studio/src/components/ModelComparePrompts/ModelColumnSelect.tsx index dd4fa5991b..f6bde49e3e 100644 --- a/web/packages/studio/src/components/ModelComparePrompts/ModelColumnSelect.tsx +++ b/web/packages/studio/src/components/ModelComparePrompts/ModelColumnSelect.tsx @@ -1,18 +1,20 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -import type { ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels'; -import { ModelSelectV2, type ModelSelection } from '@nemo/common/src/components/ModelSelectV2'; +import { + WorkspaceModelSelect, + type ModelSelection, +} from '@nemo/common/src/components/ModelSelectV2'; +import { hasModelProvider } from '@nemo/common/src/utils/models'; import { type FC, useCallback } from 'react'; /** Thin wrapper around ModelSelectV2 for table header use */ export const ModelColumnSelect: FC<{ - modelGroups: ModelWorkspaceGroup[]; - isLoadingModels: boolean; + workspace: string; value: string | null; disabled?: boolean; onChange: (ref: string) => void; -}> = ({ modelGroups, isLoadingModels, value, disabled, onChange }) => { +}> = ({ workspace, value, disabled, onChange }) => { const selectedModel: ModelSelection | null = value ? { model: value } : null; const handleValueChange = useCallback( @@ -23,11 +25,11 @@ export const ModelColumnSelect: FC<{ ); return ( - = ({ models, - modelGroups, - isLoadingModels, + workspace, promptRows, fileResult, sampleMethod, @@ -202,8 +199,7 @@ export const ModelCompareTable: FC = ({ className={`${hasPrompts ? 'border-b ' : ''}${idx < models.length - 1 ? 'border-r ' : ''}border-base px-2 py-2 align-top`} > { diff --git a/web/packages/studio/src/components/ModelComparePrompts/index.tsx b/web/packages/studio/src/components/ModelComparePrompts/index.tsx index 25530818f2..e219e18682 100644 --- a/web/packages/studio/src/components/ModelComparePrompts/index.tsx +++ b/web/packages/studio/src/components/ModelComparePrompts/index.tsx @@ -12,8 +12,6 @@ import { type FC } from 'react'; export const ModelComparePrompts: FC = ({ workspace, - modelGroups, - isLoadingModels, models, onRemoveModel, onSetModel, @@ -58,8 +56,7 @@ export const ModelComparePrompts: FC = ({
void; onSetModel: (id: number, modelURN: string | null) => void; diff --git a/web/packages/studio/src/components/ModelConfigPanel/index.tsx b/web/packages/studio/src/components/ModelConfigPanel/index.tsx index 40c07b2913..a51a331546 100644 --- a/web/packages/studio/src/components/ModelConfigPanel/index.tsx +++ b/web/packages/studio/src/components/ModelConfigPanel/index.tsx @@ -1,15 +1,14 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -import type { ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels'; import { ControlledTextInput } from '@nemo/common/src/components/form/ControlledTextInput'; -import { ModelSelectV2 } from '@nemo/common/src/components/ModelSelectV2/ModelSelectV2'; import type { ModelSelection } from '@nemo/common/src/components/ModelSelectV2/types'; +import { WorkspaceModelSelect } from '@nemo/common/src/components/ModelSelectV2/WorkspaceModelSelect'; import type { InferenceParams } from '@nemo/sdk/generated/platform/schema'; import { Button, Flex, FormField, Stack, Text } from '@nvidia/foundations-react-core'; import { CardIconBadge } from '@studio/components/common/SelectableCard'; import { - providerForModel, + providerForSelection, validateModelAlias, } from '@studio/routes/DataDesignerJobBuildRoute/models'; import type { JobBuilderFormValues } from '@studio/routes/DataDesignerJobBuildRoute/useJobBuilder'; @@ -21,8 +20,7 @@ const EMPTY_INFERENCE_PARAMS: Partial = {}; export interface ModelConfigPanelProps { modelId: string; - modelGroups: ModelWorkspaceGroup[]; - isLoadingModels?: boolean; + workspace: string; onRemove: () => void; onClose: () => void; } @@ -30,8 +28,7 @@ export interface ModelConfigPanelProps { /** Right-hand config panel for a model, with subscriptions scoped to its individual fields. */ export const ModelConfigPanel: FC = ({ modelId, - modelGroups, - isLoadingModels, + workspace, onRemove, onClose, }) => { @@ -61,7 +58,7 @@ export const ModelConfigPanel: FC = ({ const handleModelChange = (selection: ModelSelection) => { modelField.onChange(selection.model); - providerField.onChange(providerForModel(modelGroups, selection.model)); + providerField.onChange(providerForSelection(selection)); }; return ( @@ -114,11 +111,10 @@ export const ModelConfigPanel: FC = ({ /> - { @@ -23,7 +24,8 @@ export const ModelSelectionSection = () => { name: modelFieldName, }); - const { groups, isFetching } = useModelsFromWorkspace({ workspace }); + const [open, setOpen] = useState(false); + const modelSearch = useModelSearch({ workspace, enabled: open }); const selectedValue: ModelSelection | null = modelField.value ? { model: modelField.value as string } @@ -43,10 +45,10 @@ export const ModelSelectionSection = () => { slotError={modelFieldState.error?.message} > + getModelEntityChatStatus(model) !== 'disabled'; + export interface JudgeModelSelectProps { required?: boolean; placeholder?: string; @@ -24,8 +29,8 @@ export interface JudgeModelSelectProps({ required = false, @@ -49,12 +54,16 @@ export const JudgeModelSelect = rules: required ? { required: requiredMessage } : undefined, }); - const { data: judgeModels, isLoading, error } = useJudgeModels({ enabled: !disabled }); + const workspace = useWorkspaceFromPath(); + const [open, setOpen] = useState(false); + const { error, ...modelSearch } = useModelSearch({ + workspace: workspace ?? null, + enabled: open && !disabled, + include: isChatCapable, + }); useSetFieldErrorOnApiError(formFieldName, error); - const groups = useMemo(() => groupModelsByWorkspace(judgeModels ?? []), [judgeModels]); - const value: ModelSelection | null = field.value ? { model: field.value as string } : null; const handleValueChange = (selection: ModelSelection) => { @@ -62,6 +71,7 @@ export const JudgeModelSelect = }; const handleOpenChange = (isOpen: boolean) => { + setOpen(isOpen); if (!isOpen) field.onBlur(); }; @@ -73,10 +83,9 @@ export const JudgeModelSelect = required={required} > = ({ const [datasetVariables, setDatasetVariables] = useState([]); const [selectedJobType, setSelectedJobType] = useState('online'); - const { groups: modelGroups, isFetching: isLoadingModels } = useModelsFromWorkspace({ - workspace, - query: BASIC_ALL_MODELS_DROPDOWN_FILTER, - queryOptions: { enabled: open }, - }); - const evaluationModels = useMemo(() => modelGroups.flatMap((g) => g.models), [modelGroups]); + const [modelSelectOpen, setModelSelectOpen] = useState(false); + const modelSearch = useModelSearch({ workspace, enabled: open && modelSelectOpen }); const { mutateAsync: createEvaluateJob, isPending } = useEvaluatorCreateEvaluateJob(); const modelSearchParam = searchParams.get(QUERY_PARAMETERS.model); @@ -198,6 +195,12 @@ export const MetricRunSidePanel: FC = ({ remove: removePromptMessage, } = useFieldArray({ control: form.control, name: 'promptMessages' }); + const selectedModel = useWatch({ control: form.control, name: 'model' }); + const fetchedModel = useModelEntity(selectedModel?.model, { + enabled: open && !!selectedModel?.adapter && !selectedModel?.entity, + }); + const selectedModelEntity = selectedModel?.entity ?? fetchedModel; + useEffect(() => { if (open) { form.reset(defaultFormValues); @@ -259,7 +262,11 @@ export const MetricRunSidePanel: FC = ({ if (formData.jobType === 'online') { const { model, adapter } = formData.model!; const modelValue = adapter ? `${model}::${adapter}` : model; - const result = buildModelPayload(modelValue, evaluationModels, PLATFORM_BASE_URL); + const result = buildModelPayload( + modelValue, + selectedModelEntity ? [selectedModelEntity] : [], + PLATFORM_BASE_URL + ); if (!result.ok) { toast.error(result.error); return; @@ -383,11 +390,11 @@ export const MetricRunSidePanel: FC = ({ slotError={fieldState.error?.message} > void; onColumnClose: () => void; onModelRemove: () => void; @@ -21,8 +19,7 @@ export interface BuilderConfigPaneProps { export const BuilderConfigPane: FC = ({ selectedColumnId, selectedModelId, - modelGroups, - isLoadingModels, + workspace, onColumnRemove, onColumnClose, onModelRemove, @@ -38,8 +35,7 @@ export const BuilderConfigPane: FC = ({ ) : selectedModelId ? ( diff --git a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/BuilderPalette.tsx b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/BuilderPalette.tsx index d749515416..f6b2ecda49 100644 --- a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/BuilderPalette.tsx +++ b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/BuilderPalette.tsx @@ -1,7 +1,6 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -import type { ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels'; import type { ModelSelection } from '@nemo/common/src/components/ModelSelectV2/types'; import { SegmentedControl } from '@nvidia/foundations-react-core'; import { AddColumnPalette } from '@studio/components/AddColumnPalette'; @@ -18,8 +17,7 @@ export interface BuilderPaletteProps { tab: PaletteTab; onTabChange: (tab: PaletteTab) => void; selectedModelId: string | null; - modelGroups: ModelWorkspaceGroup[]; - isLoadingModels?: boolean; + workspace: string; onAddColumn: (selection: AddColumnSelection) => void; onAddModel: (selection: ModelSelection, provider: string) => void; onSelectModel: (id: string | null) => void; @@ -30,8 +28,7 @@ export const BuilderPalette: FC = ({ tab, onTabChange, selectedModelId, - modelGroups, - isLoadingModels, + workspace, onAddColumn, onAddModel, onSelectModel, @@ -64,8 +61,7 @@ export const BuilderPalette: FC = ({ diff --git a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/index.tsx b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/index.tsx index d8d2f4dc4a..7cf9a85552 100644 --- a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/index.tsx +++ b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/index.tsx @@ -1,9 +1,7 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -import { useAllModels } from '@nemo/common/src/api/models/useModels'; import { DEFAULT_LARGE_PAGE_SIZE } from '@nemo/common/src/constants/api'; -import { groupModelsByWorkspace } from '@nemo/common/src/utils/models'; import { useDataDesignerCreateJob } from '@nemo/sdk/generated/data-designer/api'; import { useModelsListProviders } from '@nemo/sdk/generated/platform/api'; import { Flex, Stack } from '@nvidia/foundations-react-core'; @@ -89,22 +87,7 @@ export const DataDesignerJobBuildRoute: FC = () => { ], }); - const { - data: modelsData, - isLoading: isLoadingModels, - hasNextPage, - isFetchingNextPage, - } = useAllModels({ workspace }); - const modelGroups = useMemo( - () => - groupModelsByWorkspace(modelsData?.pages.flatMap((page) => page.data ?? []) ?? [], { - sort: true, - }), - [modelsData?.pages] - ); - const modelsSettled = !isLoadingModels && !hasNextPage && !isFetchingNextPage; - - const builder = useJobBuilder(template, modelGroups, modelsSettled, cloneSeed); + const builder = useJobBuilder(template, workspace, cloneSeed); const { data: providersPage } = useModelsListProviders( workspace, { page_size: DEFAULT_LARGE_PAGE_SIZE }, @@ -209,8 +192,7 @@ export const DataDesignerJobBuildRoute: FC = () => { tab={builder.paletteTab} onTabChange={builder.setPaletteTab} selectedModelId={builder.selectedModelId} - modelGroups={modelGroups} - isLoadingModels={isLoadingModels} + workspace={workspace} onAddColumn={builder.handleAddColumn} onAddModel={builder.handleAddModel} onSelectModel={builder.selectModel} @@ -235,8 +217,7 @@ export const DataDesignerJobBuildRoute: FC = () => { builder.selectedColumnId && builder.removeColumn(builder.selectedColumnId) } diff --git a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/models.test.ts b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/models.test.ts index 8649fa8994..b25c69e862 100644 --- a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/models.test.ts +++ b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/models.test.ts @@ -13,7 +13,7 @@ import { defaultModelAlias, firstAvailableModel, modelIdForModel, - providerForModel, + providerForSelection, resolveTemplateModel, validateModelAlias, validateModels, @@ -35,7 +35,7 @@ describe('defaultModelAlias', () => { }); }); -describe('providerForModel', () => { +describe('model resolution', () => { const groups = [ { workspace: 'steramae', @@ -51,13 +51,17 @@ describe('providerForModel', () => { }, ] as unknown as ModelWorkspaceGroup[]; - it('returns the model’s first provider ref', () => { - expect(providerForModel(groups, 'steramae/nemotron-oss')).toBe('steramae/build'); + it('providerForSelection returns the picked entity’s first provider ref', () => { + expect( + providerForSelection({ model: 'steramae/nemotron-oss', entity: groups[0].models[0] }) + ).toBe('steramae/build'); }); - it('returns empty string when the model or its provider is missing', () => { - expect(providerForModel(groups, 'steramae/no-provider')).toBe(''); - expect(providerForModel(groups, 'steramae/unknown')).toBe(''); + it('providerForSelection returns empty string without an entity or provider', () => { + expect( + providerForSelection({ model: 'steramae/no-provider', entity: groups[0].models[2] }) + ).toBe(''); + expect(providerForSelection({ model: 'steramae/unknown' })).toBe(''); }); it('firstAvailableModel picks the first model and its provider', () => { diff --git a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/models.ts b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/models.ts index 18a4056f2b..45e5c2a84b 100644 --- a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/models.ts +++ b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/models.ts @@ -1,17 +1,25 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 +import { withOperators } from '@nemo/common/src/api/filterOperators'; import type { ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels'; import type { ModelSelection } from '@nemo/common/src/components/ModelSelectV2/types'; import { MAX_COMPLETION_TOKENS_DEFAULT } from '@nemo/common/src/constants/inferenceParameters'; import { getURNFromNamedEntityRef } from '@nemo/common/src/namedEntity'; +import { groupModelsByWorkspace } from '@nemo/common/src/utils/models'; import type { ChatCompletionInferenceParams, EmbeddingInferenceParams, EmbeddingInferenceParamsExtraBody, ModelConfig, } from '@nemo/sdk/generated/data-designer/schema'; -import type { InferenceParams, ModelProvider } from '@nemo/sdk/generated/platform/schema'; +import { modelsListModels } from '@nemo/sdk/generated/platform/api'; +import type { + InferenceParams, + ModelEntity, + ModelEntityFilter, + ModelProvider, +} from '@nemo/sdk/generated/platform/schema'; import type { TemplateModelSpec } from '@studio/components/CreateFilesetStart/types'; /** Mirrors the SDK ModelConfig shape; `alias` is what LLM columns reference via `model_alias`. */ @@ -27,19 +35,57 @@ export interface BuilderModel { export type BuilderModelPatch = Partial>; /** - * Resolves the provider for a model URN from the platform model list: the model's first - * `model_providers` entry (a `workspace/provider-name` resource ref). Data Designer needs - * an explicit provider on each model config — an unset provider is deprecated and the job - * fails with "the model does not have a provider". Returns '' when the model isn't found - * or has no provider (the user can still fill it in manually). + * The provider a dropdown selection carries: the picked model's first `model_providers` entry + * (a `workspace/provider-name` resource ref). Data Designer needs an explicit provider on each + * model config — an unset provider is deprecated and the job fails with "the model does not have + * a provider". Returns '' when the entry has no provider, or when the selection arrived without + * its entity (the user can still fill it in manually). */ -export const providerForModel = (modelGroups: ModelWorkspaceGroup[], model: string): string => { - for (const group of modelGroups) { - for (const entity of group.models) { - if (getURNFromNamedEntityRef(entity) === model) return entity.model_providers?.[0] ?? ''; - } - } - return ''; +export const providerForSelection = (selection: ModelSelection): string => + selection.entity?.model_providers?.[0] ?? ''; + +/** One page is plenty: auto-fill only ever needs a name match or a first choice. */ +const AUTO_FILL_PAGE_SIZE = 25; + +/** + * The models {@link resolveTemplateModel} should consider: those whose name matches `preferred`, + * plus the first page of the workspace as a fallback. Two small requests instead of walking the + * whole catalogue, which is all the resolver needs to make its choice. + */ +export const fetchAutoFillCandidates = async ( + workspace: string, + preferred?: string +): Promise => { + const listPage = async (filter?: ModelEntityFilter): Promise => { + const page = await modelsListModels(workspace, { + page_size: AUTO_FILL_PAGE_SIZE, + sort: 'name', + ...(filter ? { filter } : {}), + }); + return page.data ?? []; + }; + + // Template specs name a model without its workspace prefix or version suffix; match on that. + const preferredName = preferred + ? (preferred.split('/').pop() ?? preferred).split('@')[0] + : undefined; + + const [matches, firstPage] = await Promise.all([ + preferredName + ? listPage(withOperators({ name: { $like: preferredName } })) + : Promise.resolve([]), + listPage(), + ]); + + const seen = new Set(); + const models = [...matches, ...firstPage].filter((entity) => { + const urn = getURNFromNamedEntityRef(entity); + if (!urn || seen.has(urn)) return false; + seen.add(urn); + return true; + }); + + return groupModelsByWorkspace(models); }; export const buildServedModelNames = (providers: ModelProvider[]): Map => { diff --git a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/useJobBuilder.ts b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/useJobBuilder.ts index 8878f4eb9c..e2f03ccde9 100644 --- a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/useJobBuilder.ts +++ b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/useJobBuilder.ts @@ -1,7 +1,6 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -import type { ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels'; import type { ModelSelection } from '@nemo/common/src/components/ModelSelectV2/types'; import type { AddColumnSelection } from '@studio/components/AddColumnPalette/types'; import type { FilesetTemplate } from '@studio/components/CreateFilesetStart/types'; @@ -15,6 +14,7 @@ import { type BuilderModel, buildModelsFromTemplate, builderModelFromSelection, + fetchAutoFillCandidates, resolveTemplateModel, } from '@studio/routes/DataDesignerJobBuildRoute/models'; import { useCallback, useEffect, useRef, useState } from 'react'; @@ -56,16 +56,15 @@ export interface JobBuilderSeed { * Job-level concerns (name, row count, validation, preview, submit) live in the route so * this hook stays a pure graph-editing store. * - * `modelGroups` auto-fills a template's seeded models once the platform model list loads. - * `modelsSettled` gates that auto-fill on the full (all-pages) model list being available. + * A template's seeded models are auto-filled once from `workspace`, resolving each spec with a + * targeted lookup rather than the whole model catalogue. * * `seed`, when provided (cloning a job), takes precedence over the template and pre-fills the * form with the source job's columns, models, name, and row count. */ export const useJobBuilder = ( template: FilesetTemplate | null, - modelGroups: ModelWorkspaceGroup[], - modelsSettled: boolean, + workspace: string, seed: JobBuilderSeed | null = null ) => { // Seed once from the clone source, else the template (if any). `useForm` keeps these values @@ -107,16 +106,36 @@ export const useJobBuilder = ( const autoFilled = useRef(false); useEffect(() => { - if (autoFilled.current || !modelsSettled || modelGroups.length === 0) return; + if (autoFilled.current || !workspace) return; + const pending = getValues('models').filter((model) => !model.provider); + if (pending.length === 0) return; autoFilled.current = true; - const models = getValues('models'); - const nextModels = models.map((model) => { - if (model.provider) return model; - const resolved = resolveTemplateModel(modelGroups, model.model || undefined); - return resolved ? { ...model, ...resolved } : { ...model, model: '' }; - }); - setValue('models', nextModels); - }, [getValues, modelGroups, modelsSettled, setValue]); + + let cancelled = false; + void (async () => { + const resolutions = await Promise.all( + pending.map(async (model) => { + const preferred = model.model || undefined; + const candidates = await fetchAutoFillCandidates(workspace, preferred); + return [model.id, resolveTemplateModel(candidates, preferred)] as const; + }) + ); + if (cancelled) return; + const byId = new Map(resolutions); + setValue( + 'models', + getValues('models').map((model) => { + if (!byId.has(model.id)) return model; + const resolved = byId.get(model.id); + return resolved ? { ...model, ...resolved } : { ...model, model: '' }; + }) + ); + })(); + + return () => { + cancelled = true; + }; + }, [getValues, setValue, workspace]); const selectColumn = useCallback((id: string | null) => { setSelectedId(id); diff --git a/web/packages/studio/src/routes/DeploymentsListRoute/CreateDeploymentSidePanel/WorkspaceSourceFields.tsx b/web/packages/studio/src/routes/DeploymentsListRoute/CreateDeploymentSidePanel/WorkspaceSourceFields.tsx index 791c611d53..16c29adba4 100644 --- a/web/packages/studio/src/routes/DeploymentsListRoute/CreateDeploymentSidePanel/WorkspaceSourceFields.tsx +++ b/web/packages/studio/src/routes/DeploymentsListRoute/CreateDeploymentSidePanel/WorkspaceSourceFields.tsx @@ -3,7 +3,7 @@ * SPDX-License-Identifier: Apache-2.0 */ -import { buildWorkspaceGroup, useAllModels } from '@nemo/common/src/api/models/useModels'; +import { useModelSearch } from '@nemo/common/src/api/models/useModelSearch'; import { ControlledTextInput } from '@nemo/common/src/components/form/ControlledTextInput'; import { type ModelSelection, ModelSelectV2 } from '@nemo/common/src/components/ModelSelectV2'; import { RadioCard } from '@nemo/common/src/components/RadioCard'; @@ -14,7 +14,7 @@ import { WORKSPACE_PICKER_MODEL, type WizardFormValues, } from '@studio/routes/DeploymentsListRoute/CreateDeploymentSidePanel/schema'; -import { useMemo, type FC } from 'react'; +import { useState, type FC } from 'react'; import { useController, useWatch, type Control, type FieldErrors } from 'react-hook-form'; export type WorkspaceSourceFieldsProps = { @@ -109,17 +109,9 @@ const WorkspaceModelPicker: FC = ({ errorMessage, }) => { const { field } = useController({ control, name: 'modelRef' }); + const [open, setOpen] = useState(false); - const { data, isLoading } = useAllModels({ - workspace, - query: { page_size: 100, sort: 'name' }, - queryOptions: { enabled: queryEnabled && !!workspace }, - }); - - const groups = useMemo(() => { - const models = data?.pages.flatMap((page) => page.data) ?? []; - return models.length > 0 ? [buildWorkspaceGroup(workspace, models)] : []; - }, [data?.pages, workspace]); + const modelSearch = useModelSearch({ workspace, enabled: queryEnabled && open }); const value: ModelSelection | null = field.value ? { model: field.value as string } : null; @@ -130,15 +122,15 @@ const WorkspaceModelPicker: FC = ({ slotError={errorMessage} > field.onChange(selection.model)} - groups={groups} - loading={isLoading} placeholder="Select a model" hideAdapters fullWidth - onOpenChange={(open) => { - if (!open) field.onBlur(); + onOpenChange={(nextOpen) => { + setOpen(nextOpen); + if (!nextOpen) field.onBlur(); }} /> diff --git a/web/packages/studio/src/routes/ModelCompareRoute/index.tsx b/web/packages/studio/src/routes/ModelCompareRoute/index.tsx index 83d77f45a4..2e1d712c92 100644 --- a/web/packages/studio/src/routes/ModelCompareRoute/index.tsx +++ b/web/packages/studio/src/routes/ModelCompareRoute/index.tsx @@ -1,14 +1,11 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -import { - BASIC_ALL_MODELS_DROPDOWN_FILTER, - buildWorkspaceGroup, - type ModelWorkspaceGroup, - useAllModels, -} from '@nemo/common/src/api/models/useModels'; +import { useModelEntity } from '@nemo/common/src/api/models/useModelEntity'; +import { useModelSearch } from '@nemo/common/src/api/models/useModelSearch'; import { ComposerMode } from '@nemo/common/src/components/AssistantChat'; import type { BroadcastSignal } from '@nemo/common/src/components/AssistantChat/types'; +import { hasModelProvider } from '@nemo/common/src/utils/models'; import { PageHeader, Tabs, Tooltip } from '@nvidia/foundations-react-core'; import { ChatEmptyState } from '@studio/components/chat/ChatEmptyState'; import { CompareComposer } from '@studio/components/chat/CompareComposer'; @@ -37,18 +34,10 @@ const makeDefaultEntry = ( export const ModelCompareRoute: FC = () => { const workspace = useWorkspaceFromPath(); - const { data, isFetching: isLoadingModels } = useAllModels({ - workspace: workspace ?? undefined, - query: BASIC_ALL_MODELS_DROPDOWN_FILTER, - }); - const modelGroups = useMemo((): ModelWorkspaceGroup[] => { - if (!workspace) return []; - const allModels = data?.pages.flatMap((p) => (Array.isArray(p.data) ? p.data : [])) ?? []; - const available = allModels.filter( - (m) => Array.isArray(m.model_providers) && m.model_providers.length > 0 - ); - return available.length > 0 ? [buildWorkspaceGroup(workspace, available)] : []; - }, [data, workspace]); + const availableModels = useModelSearch({ workspace, include: hasModelProvider }); + const hasNoModels = + !availableModels.loading && !availableModels.hasMore && availableModels.models.length === 0; + const [searchParams] = useSearchParams(); const [activeView, setActiveView] = useState('compare'); const [perPanelInput, setPerPanelInput] = useState(false); @@ -59,24 +48,31 @@ export const ModelCompareRoute: FC = () => { const nextIdRef = useRef(2); const didPreselectRef = useRef(false); - // Preselect panel 0 from ?model= query param once models load. + // Preselect panel 0 from ?model=, which carries either a full URN or a bare name in this + // workspace. Resolved with a single lookup rather than by scanning the whole catalogue. + const modelParam = searchParams.get('model'); + const preselectUrn = + modelParam && workspace + ? modelParam.includes('/') + ? modelParam + : `${workspace}/${modelParam}` + : null; + const preselectModel = useModelEntity(preselectUrn, { enabled: !didPreselectRef.current }); + useEffect(() => { - if (didPreselectRef.current || isLoadingModels || modelGroups.length === 0) return; - const param = searchParams.get('model'); - if (!param) { + if (didPreselectRef.current) return; + if (!modelParam) { didPreselectRef.current = true; return; } - const match = modelGroups - .flatMap((g) => g.models) - .find((m) => `${m.workspace}/${m.name}` === param || m.name === param); - if (match) { - setModels((prev) => - prev.map((m, i) => (i === 0 ? { ...m, modelURN: `${match.workspace}/${match.name}` } : m)) - ); - } + if (!preselectModel) return; + setModels((prev) => + prev.map((m, i) => + i === 0 ? { ...m, modelURN: `${preselectModel.workspace}/${preselectModel.name}` } : m + ) + ); didPreselectRef.current = true; - }, [isLoadingModels, modelGroups, searchParams]); + }, [modelParam, preselectModel]); // Seed transfer: broadcast→panels and panel→broadcast. const compareComposerDraftRef = useRef(''); @@ -162,8 +158,8 @@ export const ModelCompareRoute: FC = () => { const atMaxModels = models.length >= MAX_MODELS; const readyPanelCount = models.filter((m) => !!m.modelURN).length; - // Empty state when the workspace has zero models and we're not still loading. - if (!isLoadingModels && modelGroups.length === 0) { + // Empty state when the workspace has no servable models and we're not still loading. + if (hasNoModels) { return ; } @@ -187,8 +183,6 @@ export const ModelCompareRoute: FC = () => {
{
Date: Tue, 28 Jul 2026 11:57:19 -0700 Subject: [PATCH 2/3] fix issues with long lists in virtualizer Signed-off-by: Sean Teramae --- .../ModelSelectV2/ModelDropdown.test.tsx | 154 +++++++++++++++++- .../ModelSelectV2/ModelDropdown.tsx | 79 +++++---- .../ModelSelectV2/ModelDropdownItem.tsx | 115 ++++++++----- .../ModelSelectV2/ModelDropdownList.tsx | 41 +++-- .../useJobBuilder.test.tsx | 83 ++++++++++ .../useJobBuilder.ts | 8 +- 6 files changed, 385 insertions(+), 95 deletions(-) create mode 100644 web/packages/studio/src/routes/DataDesignerJobBuildRoute/useJobBuilder.test.tsx diff --git a/web/packages/common/src/components/ModelSelectV2/ModelDropdown.test.tsx b/web/packages/common/src/components/ModelSelectV2/ModelDropdown.test.tsx index f8125ca0f3..b95ad97bee 100644 --- a/web/packages/common/src/components/ModelSelectV2/ModelDropdown.test.tsx +++ b/web/packages/common/src/components/ModelSelectV2/ModelDropdown.test.tsx @@ -4,15 +4,82 @@ import type { ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels'; import { ModelDropdown } from '@nemo/common/src/components/ModelSelectV2/ModelDropdown'; import type { ModelEntity } from '@nemo/sdk/generated/platform/schema'; -import { fireEvent, render, screen, waitFor } from '@testing-library/react'; +import { act, fireEvent, render, screen, waitFor } from '@testing-library/react'; -const makeModel = (name: string, workspace = 'nvidia'): ModelEntity => - ({ id: name, name, workspace }) as unknown as ModelEntity; +const makeModel = (name: string, overrides: Partial = {}): ModelEntity => + ({ id: name, name, workspace: 'nvidia', ...overrides }) as unknown as ModelEntity; const groups: ModelWorkspaceGroup[] = [ { workspace: 'nvidia', models: [makeModel('nemotron-8b'), makeModel('llama-3.1-8b')] }, ]; +const withAdapters: ModelWorkspaceGroup[] = [ + { + workspace: 'nvidia', + models: [ + makeModel('nemotron-8b', { + adapters: [ + { + name: 'support-v1', + created_at: '2026-01-10T00:00:00Z', + workspace: 'nvidia', + fileset: 'nvidia/support-v1', + finetuning_type: 'lora', + }, + ], + }), + ], + }, +]; + +const MANY_MODELS = 200; + +const manyGroups: ModelWorkspaceGroup[] = [ + { + workspace: 'nvidia', + models: Array.from({ length: MANY_MODELS }, (_, i) => makeModel(`model-${i}`)), + }, +]; + +/** + * jsdom has no layout, so the viewport is stubbed to make the virtualizer engage. Returns a + * handle for driving the resize callback the list subscribes to. + */ +const stubViewport = () => { + type ResizeCallback = (entries: ResizeObserverEntry[]) => void; + const callbacks: ResizeCallback[] = []; + let height = 0; + + vi.spyOn(HTMLElement.prototype, 'clientHeight', 'get').mockImplementation(() => height); + // Rows measure themselves; jsdom would report every one as zero-height, collapsing the window. + vi.spyOn(HTMLElement.prototype, 'getBoundingClientRect').mockReturnValue({ + width: 360, + height: 36, + } as DOMRect); + vi.stubGlobal( + 'ResizeObserver', + class { + constructor(callback: ResizeCallback) { + callbacks.push(callback); + } + observe() {} + unobserve() {} + disconnect() {} + } + ); + + return { + resizeTo: (next: number) => { + height = next; + // The virtualizer reads `borderBoxSize` off the entry; the list ignores the argument. + const entries = [ + { borderBoxSize: [{ inlineSize: 360, blockSize: next }] }, + ] as unknown as ResizeObserverEntry[]; + callbacks.forEach((callback) => callback(entries)); + }, + }; +}; + type Props = React.ComponentProps; const renderOpen = (props: Partial = {}) => @@ -99,6 +166,72 @@ describe('ModelDropdown', () => { }); }); + describe('virtualization', () => { + afterEach(() => { + vi.restoreAllMocks(); + vi.unstubAllGlobals(); + }); + + it('renders nothing at all while the menu is closed', () => { + renderOpen({ open: false, groups: manyGroups }); + + // DropdownContent is a native popover and stays mounted, so a closed menu would otherwise + // keep a row (and its submenu popover) in the DOM for every model in the catalogue. + expect(listedModels()).toHaveLength(0); + expect(screen.queryByTestId('model-select-v2-filter')).not.toBeInTheDocument(); + }); + + // Regression: the viewport used to be measured once on mount. The menu is a popover, so that + // measurement read zero height and disabled windowing for good — every model in the workspace + // rendered at once, which is what ground the page to a halt on large catalogues. + it('windows the list once the popover gains height after mount', async () => { + const viewport = stubViewport(); + renderOpen({ groups: manyGroups }); + + await act(async () => viewport.resizeTo(300)); + + // A 300px viewport over 36px rows is ~9 visible, plus overscan either side — nowhere near + // the 200 in the catalogue. + const rendered = listedModels().length; + expect(rendered).toBeGreaterThan(0); + expect(rendered).toBeLessThan(40); + }); + + it('falls back to rendering every row when there is no viewport to window into', async () => { + renderOpen({ groups: manyGroups }); + + await waitFor(() => expect(listedModels()).toHaveLength(MANY_MODELS)); + }); + }); + + describe('row details', () => { + // KUI's DropdownSubContent is always in the DOM, so rendering its body eagerly would mount a + // details panel (and every adapter's) for each row in view — the cost that made typing lag. + it('does not mount a row’s details until its submenu is opened', async () => { + renderOpen({ groups: withAdapters }); + await waitFor(() => + expect(screen.getAllByTestId('model-dropdown-item-with-adapters')).toHaveLength(1) + ); + + expect(screen.queryByTestId('model-dropdown-adapter-option')).not.toBeInTheDocument(); + expect(screen.queryByText('Fine-tuning Type')).not.toBeInTheDocument(); + }); + + it('mounts the details on hover and keeps them mounted afterwards', async () => { + renderOpen({ groups: withAdapters }); + const sub = (await screen.findAllByTestId('nv-dropdown-sub'))[0]; + + fireEvent.pointerEnter(sub); + + expect(await screen.findByTestId('model-dropdown-adapter-option')).toBeInTheDocument(); + + fireEvent.pointerLeave(sub, { relatedTarget: document.body }); + + // Latched: leaving does not throw the panel away, so a second hover is instant. + expect(screen.getByTestId('model-dropdown-adapter-option')).toBeInTheDocument(); + }); + }); + describe('trigger', () => { it('falls back to the name in the URN when the selection is not in the loaded pages', () => { renderOpen({ open: false, value: { model: 'nvidia/not-yet-loaded' }, groups: [] }); @@ -106,6 +239,21 @@ describe('ModelDropdown', () => { expect(screen.getByTestId('model-select-v2-trigger')).toHaveTextContent('not-yet-loaded'); }); + // Data Designer templates seed a bare model name, which only becomes a URN once auto-fill + // resolves it. Showing the placeholder over it reads as "nothing selected". + it('shows a reference that has no workspace prefix rather than the placeholder', () => { + renderOpen({ + open: false, + value: { model: 'nvidia-llama-3-3-nemotron-super-49b-v1' }, + groups: [], + placeholder: 'Select a model', + }); + + expect(screen.getByTestId('model-select-v2-trigger')).toHaveTextContent( + 'nvidia-llama-3-3-nemotron-super-49b-v1' + ); + }); + it('prefers the entity the selection carries', () => { renderOpen({ open: false, diff --git a/web/packages/common/src/components/ModelSelectV2/ModelDropdown.tsx b/web/packages/common/src/components/ModelSelectV2/ModelDropdown.tsx index a6bd3286de..4625457913 100644 --- a/web/packages/common/src/components/ModelSelectV2/ModelDropdown.tsx +++ b/web/packages/common/src/components/ModelSelectV2/ModelDropdown.tsx @@ -22,7 +22,7 @@ import { Text, } from '@nvidia/foundations-react-core'; import { ChevronDown, LoaderCircle } from 'lucide-react'; -import { useEffect, useMemo, useState, type FC } from 'react'; +import { useCallback, useEffect, useMemo, useState, type FC } from 'react'; import { useDebounce } from 'use-debounce'; const MODEL_TYPE_ITEMS = [ @@ -85,7 +85,7 @@ export const ModelDropdown: FC = ({ const filteredGroups = useMemo(() => { const filterType = showModelTypeToggle && !onModelTypeChange; - const filterSearch = !onSearchChange && search.length > 0; + const filterSearch = !onSearchChange && debouncedSearch.length > 0; if (!filterType && !filterSearch) return localGroups; return localGroups @@ -95,17 +95,27 @@ export const ModelDropdown: FC = ({ models = modelType === 'base' ? models.filter(isBaseModel) : models.filter(isCustomModel); } if (filterSearch) { - models = models.filter((m) => filterModel(m, search)); + models = models.filter((m) => filterModel(m, debouncedSearch)); } return { ...group, models }; }) .filter((group) => group.models.length > 0); - }, [localGroups, modelType, onModelTypeChange, onSearchChange, search, showModelTypeToggle]); - - const handleSelect = (selection: ModelSelection) => { - onValueChange(selection); - onOpenChange(false); - }; + }, [ + debouncedSearch, + localGroups, + modelType, + onModelTypeChange, + onSearchChange, + showModelTypeToggle, + ]); + + const handleSelect = useCallback( + (selection: ModelSelection) => { + onValueChange(selection); + onOpenChange(false); + }, + [onOpenChange, onValueChange] + ); const handleModelTypeChange = (val: string) => { setModelType(val as ModelType); @@ -120,8 +130,9 @@ export const ModelDropdown: FC = ({ }; const selectedParts = value?.model ? getPartsFromReference(value.model) : undefined; - const selectedName = selectedModel?.name ?? selectedParts?.name; - const selectedWorkspace = selectedModel?.workspace ?? selectedParts?.workspace; + const selectedName = selectedModel?.name ?? (selectedParts?.name || value?.model); + const selectedWorkspace = + selectedModel?.workspace ?? (selectedParts?.name ? selectedParts.workspace : undefined); const triggerLabel = selectedName ? (selectedName.split('@')[0] ?? selectedName) : placeholder; return ( @@ -166,29 +177,33 @@ export const ModelDropdown: FC = ({ className="min-w-[360px]" style={{ width: 360 }} // eslint-disable-line no-restricted-syntax -- KUI DropdownContent needs explicit width > - - {showModelTypeToggle && ( - - + + {showModelTypeToggle && ( + + + + )} + - + )} - ); diff --git a/web/packages/common/src/components/ModelSelectV2/ModelDropdownItem.tsx b/web/packages/common/src/components/ModelSelectV2/ModelDropdownItem.tsx index 511c0a50d8..9d2fd5c5dd 100644 --- a/web/packages/common/src/components/ModelSelectV2/ModelDropdownItem.tsx +++ b/web/packages/common/src/components/ModelSelectV2/ModelDropdownItem.tsx @@ -3,7 +3,6 @@ import { ModelDetailsPanel } from '@nemo/common/src/components/ModelSelectV2/ModelDetailsPanel'; import type { ModelSelection } from '@nemo/common/src/components/ModelSelectV2/types'; -import { getURNFromNamedEntityRef } from '@nemo/common/src/namedEntity'; import type { Adapter, ModelEntity } from '@nemo/sdk/generated/platform/schema'; import { Divider, @@ -16,11 +15,16 @@ import { Text, } from '@nvidia/foundations-react-core'; import { Check } from 'lucide-react'; -import type { FC } from 'react'; +import { memo, useState, type FC } from 'react'; interface ModelDropdownItemProps { model: ModelEntity; - value: ModelSelection | null; + /** Precomputed by the list, so the URN is parsed once per model rather than once per render. */ + modelUrn: string; + /** Whether this model is the current selection. Primitive so the row can memoize. */ + isSelected: boolean; + /** The selected adapter, when it belongs to this model. */ + selectedAdapter?: string; onSelect: (selection: ModelSelection) => void; hideAdapters?: boolean; } @@ -50,6 +54,23 @@ const ModelName: FC<{ name: string | undefined }> = ({ name }) => { ); }; +/** + * Tracks whether a submenu has ever been opened. + * + * KUI's `DropdownSubContent` is a native `popover="auto"` surface with no presence gating — it + * sits in the DOM whether or not the submenu is open. Rendering its body unconditionally means + * every row in view mounts a whole details panel (two live relative-time tickers each), and every + * adapter mounts another one below that. Deferring the body until first hover keeps the cost + * proportional to what the user actually looks at; latching keeps repeat hovers instant. + */ +const useOpenedOnce = () => { + const [hasOpened, setHasOpened] = useState(false); + const handleOpenChange = (open: boolean) => { + if (open) setHasOpened(true); + }; + return [hasOpened, handleOpenChange] as const; +}; + const AdapterItem: FC<{ adapter: Adapter; model: ModelEntity; @@ -57,8 +78,10 @@ const AdapterItem: FC<{ isSelected: boolean; onSelect: (selection: ModelSelection) => void; }> = ({ adapter, model, modelUrn, isSelected, onSelect }) => { + const [hasOpened, handleOpenChange] = useOpenedOnce(); + return ( - + {/* eslint-disable-next-line no-restricted-syntax -- KUI ignores Tailwind width classes */} - + {hasOpened && } ); }; -export const ModelDropdownItem: FC = ({ +const sortAdaptersByNewest = (adapters: Adapter[]): Adapter[] => + [...adapters].sort((a, b) => { + if (!a.created_at || !b.created_at) return 0; + return new Date(b.created_at).getTime() - new Date(a.created_at).getTime(); + }); + +const ModelDropdownItemImpl: FC = ({ model, - value, + modelUrn, + isSelected, + selectedAdapter, onSelect, hideAdapters = false, }) => { - const modelUrn = getURNFromNamedEntityRef(model)!; + const [hasOpened, handleOpenChange] = useOpenedOnce(); const hasAdapters = !hideAdapters && model.adapters && model.adapters.length > 0; if (!hasAdapters) { return ( - + = ({ {/* eslint-disable-next-line no-restricted-syntax -- KUI ignores Tailwind width classes */} - + {hasOpened && } ); } - const sortedAdapters = [...model.adapters!].sort((a, b) => { - if (!a.created_at || !b.created_at) return 0; - return new Date(b.created_at).getTime() - new Date(a.created_at).getTime(); - }); - - const isBaseSelected = value?.model === modelUrn && !value?.adapter; + const isBaseSelected = isSelected && !selectedAdapter; return ( - + {/* eslint-disable-next-line no-restricted-syntax -- KUI ignores Tailwind width classes */} - Base Model - onSelect({ model: modelUrn, entity: model })} - > - - {isBaseSelected && } - {modelUrn} - - - - Adapters - {sortedAdapters.map((adapter) => ( - - ))} + {hasOpened && ( + <> + Base Model + onSelect({ model: modelUrn, entity: model })} + > + + {isBaseSelected && } + {modelUrn} + + + + Adapters + {sortAdaptersByNewest(model.adapters!).map((adapter) => ( + + ))} + + )} ); }; + +/** + * Memoized because the filter box's state lives above the list: without this, every keystroke + * re-renders every row in view. All props are primitives or stable references, so the shallow + * compare actually holds. + */ +export const ModelDropdownItem = memo(ModelDropdownItemImpl); diff --git a/web/packages/common/src/components/ModelSelectV2/ModelDropdownList.tsx b/web/packages/common/src/components/ModelSelectV2/ModelDropdownList.tsx index 5a107be1be..198268c6c0 100644 --- a/web/packages/common/src/components/ModelSelectV2/ModelDropdownList.tsx +++ b/web/packages/common/src/components/ModelSelectV2/ModelDropdownList.tsx @@ -15,7 +15,7 @@ import { useCallback, useEffect, useLayoutEffect, useMemo, useRef, useState, typ /** Rows the virtualizer measures: group headings and models share one flat index space. */ type ModelRow = | { kind: 'heading'; key: string; workspace: string } - | { kind: 'model'; key: string; model: ModelEntity }; + | { kind: 'model'; key: string; model: ModelEntity; urn: string }; const HEADING_HEIGHT = 32; const ITEM_HEIGHT = 36; @@ -60,22 +60,32 @@ export const ModelDropdownList: FC = ({ emptyMessage = 'No models found', }) => { const scrollRef = useRef(null); - const [hasViewport, setHasViewport] = useState(true); + const [viewportHeight, setViewportHeight] = useState(0); const [isLoadingMoreLocal, setIsLoadingMoreLocal] = useState(false); useLayoutEffect(() => { - setHasViewport((scrollRef.current?.clientHeight ?? 0) > 0); + const element = scrollRef.current; + if (!element) return; + + const measure = () => setViewportHeight(element.clientHeight); + measure(); + + if (typeof ResizeObserver === 'undefined') return; + const observer = new ResizeObserver(measure); + observer.observe(element); + return () => observer.disconnect(); }, []); + const isVirtualized = viewportHeight > 0; + const rows = useMemo( () => groups.flatMap((group) => [ { kind: 'heading' as const, key: `heading-${group.workspace}`, workspace: group.workspace }, - ...group.models.map((model) => ({ - kind: 'model' as const, - key: getURNFromNamedEntityRef(model) ?? `${group.workspace}/${model.name}`, - model, - })), + ...group.models.map((model) => { + const urn = getURNFromNamedEntityRef(model) ?? `${group.workspace}/${model.name}`; + return { kind: 'model' as const, key: urn, model, urn }; + }), ]), [groups] ); @@ -89,6 +99,8 @@ export const ModelDropdownList: FC = ({ }); const virtualItems = virtualizer.getVirtualItems(); + const lastVisibleIndex = + virtualItems.length > 0 ? virtualItems[virtualItems.length - 1].index : -1; const loadMore = useCallback(async () => { if (!onLoadMore || !hasMore || isLoadingMore || isLoadingMoreLocal) return; @@ -101,12 +113,11 @@ export const ModelDropdownList: FC = ({ }, [hasMore, isLoadingMore, isLoadingMoreLocal, onLoadMore]); useEffect(() => { - if (virtualItems.length === 0) return; - const lastRow = virtualItems[virtualItems.length - 1]; - if (lastRow.index >= rows.length - LOAD_MORE_THRESHOLD) { + if (lastVisibleIndex < 0) return; + if (lastVisibleIndex >= rows.length - LOAD_MORE_THRESHOLD) { void loadMore(); } - }, [loadMore, rows.length, virtualItems]); + }, [lastVisibleIndex, loadMore, rows.length]); const loadingMore = isLoadingMore || isLoadingMoreLocal; const showDoneMessage = Boolean(doneLoadingMessage) && !hasMore && !loading && !loadingMore; @@ -122,7 +133,9 @@ export const ModelDropdownList: FC = ({ ) : ( @@ -138,7 +151,7 @@ export const ModelDropdownList: FC = ({ return (
- {hasViewport ? ( + {isVirtualized ? (
{ + const actual = await importOriginal(); + return { ...actual, modelsListModels: vi.fn() }; +}); + +const mockListModels = vi.mocked(modelsListModels); + +const NEMOTRON = 'nvidia-llama-3-3-nemotron-super-49b-v1'; + +const makePage = (data: ModelEntity[]): ModelEntitysPage => + ({ data, pagination: { page: 1, total_pages: 1 } }) as ModelEntitysPage; + +const model = (name: string, providers: string[] = ['ws1/build']): ModelEntity => + ({ id: name, name, workspace: 'ws1', model_providers: providers }) as unknown as ModelEntity; + +/** A template that seeds a bare model name, the way the real templates do. */ +const template = { + id: 'text-to-python', + models: [{ alias: 'default', model: NEMOTRON }], + columns: [], +} as unknown as FilesetTemplate; + +beforeEach(() => { + mockListModels.mockReset(); + mockListModels.mockResolvedValue(makePage([model(NEMOTRON), model('some-other-model')])); +}); + +describe('useJobBuilder template auto-fill', () => { + it('applies under StrictMode, where effects mount twice', async () => { + const { result } = renderHook(() => useJobBuilder(template, 'ws1'), { wrapper: StrictMode }); + + await waitFor(() => + expect(result.current.getBuilderValues().models[0]).toMatchObject({ + model: `ws1/${NEMOTRON}`, + provider: 'ws1/build', + }) + ); + }); + + it('resolves each seeded model exactly once', async () => { + const { result } = renderHook(() => useJobBuilder(template, 'ws1'), { wrapper: StrictMode }); + + await waitFor(() => + expect(result.current.getBuilderValues().models[0].provider).toBe('ws1/build') + ); + + // One preferred-name lookup plus one first-page fallback, for the single seeded model. + expect(mockListModels).toHaveBeenCalledTimes(2); + }); + + it('clears the model when the workspace has nothing to resolve to', async () => { + mockListModels.mockResolvedValue(makePage([])); + + const { result } = renderHook(() => useJobBuilder(template, 'ws1'), { wrapper: StrictMode }); + + await waitFor(() => expect(result.current.getBuilderValues().models[0].model).toBe('')); + }); + + it('leaves models that already carry a provider alone', async () => { + const seeded = { + ...template, + models: [{ alias: 'default', model: NEMOTRON }], + } as unknown as FilesetTemplate; + + const { result } = renderHook( + () => useJobBuilder(seeded, 'ws1', { name: 'clone', rows: '10', columns: [], models: [] }), + { wrapper: StrictMode } + ); + + await waitFor(() => expect(result.current.getBuilderValues().models).toEqual([])); + expect(mockListModels).not.toHaveBeenCalled(); + }); +}); diff --git a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/useJobBuilder.ts b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/useJobBuilder.ts index e2f03ccde9..9f1e9ea642 100644 --- a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/useJobBuilder.ts +++ b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/useJobBuilder.ts @@ -111,16 +111,14 @@ export const useJobBuilder = ( if (pending.length === 0) return; autoFilled.current = true; - let cancelled = false; void (async () => { const resolutions = await Promise.all( pending.map(async (model) => { const preferred = model.model || undefined; - const candidates = await fetchAutoFillCandidates(workspace, preferred); + const candidates = await fetchAutoFillCandidates(workspace, preferred).catch(() => []); return [model.id, resolveTemplateModel(candidates, preferred)] as const; }) ); - if (cancelled) return; const byId = new Map(resolutions); setValue( 'models', @@ -131,10 +129,6 @@ export const useJobBuilder = ( }) ); })(); - - return () => { - cancelled = true; - }; }, [getValues, setValue, workspace]); const selectColumn = useCallback((id: string | null) => { From bb7a44a76e017ea40e4a76245ca9a37e59e6eb2d Mon Sep 17 00:00:00 2001 From: Sean Teramae Date: Tue, 28 Jul 2026 12:16:33 -0700 Subject: [PATCH 3/3] fix pr comments Signed-off-by: Sean Teramae --- .../ModelDetailsSection/index.tsx | 5 +- .../sidePanels/MetricRunSidePanel/index.tsx | 6 + .../DataDesignerJobBuildRoute/models.ts | 7 +- .../useJobBuilder.test.tsx | 8 ++ .../routes/ModelCompareRoute/index.test.tsx | 133 ++++++++++++++++++ .../src/routes/ModelCompareRoute/index.tsx | 38 +++-- 6 files changed, 180 insertions(+), 17 deletions(-) create mode 100644 web/packages/studio/src/routes/ModelCompareRoute/index.test.tsx diff --git a/web/packages/studio/src/components/PromptTuningForm/ModelDetailsSection/index.tsx b/web/packages/studio/src/components/PromptTuningForm/ModelDetailsSection/index.tsx index e87cee68c8..60fd243795 100644 --- a/web/packages/studio/src/components/PromptTuningForm/ModelDetailsSection/index.tsx +++ b/web/packages/studio/src/components/PromptTuningForm/ModelDetailsSection/index.tsx @@ -12,6 +12,7 @@ import { compileSystemPrompt } from '@nemo/common/src/models/utils'; import { getURNFromNamedEntityRef } from '@nemo/common/src/namedEntity'; import { useModelsGetModel as useGetModel } from '@nemo/sdk/generated/platform/api'; import { FormField, Stack } from '@nvidia/foundations-react-core'; +import { useSetFieldErrorOnApiError } from '@studio/hooks/evaluation/useSetFieldErrorOnApiError'; import { useWorkspaceFromPath } from '@studio/hooks/useWorkspaceFromPath'; import type { PromptTuningFormFields, @@ -38,12 +39,14 @@ export const ModelDetailsSection: FC< rules: { required: 'Base model is required' }, }); - const { groups, ...modelSearch } = useModelSearch({ + const { groups, error, ...modelSearch } = useModelSearch({ workspace: workspace ?? null, filter: QUERY_PROMPT_TUNEABLE_MODELS.filter, enabled: open, }); + useSetFieldErrorOnApiError('baseModel', error); + const iclFewShotExamples = useWatch({ control, name: 'iclFewShotExamples' }); const baseModelFullName = getValues('baseModel'); diff --git a/web/packages/studio/src/components/sidePanels/MetricRunSidePanel/index.tsx b/web/packages/studio/src/components/sidePanels/MetricRunSidePanel/index.tsx index 08b8a2b2c7..95b88af946 100644 --- a/web/packages/studio/src/components/sidePanels/MetricRunSidePanel/index.tsx +++ b/web/packages/studio/src/components/sidePanels/MetricRunSidePanel/index.tsx @@ -201,6 +201,12 @@ export const MetricRunSidePanel: FC = ({ }); const selectedModelEntity = selectedModel?.entity ?? fetchedModel; + useEffect(() => { + if (modelSearch.error) { + form.setError('model', { message: modelSearch.error.message }); + } + }, [form, modelSearch.error]); + useEffect(() => { if (open) { form.reset(defaultFormValues); diff --git a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/models.ts b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/models.ts index 45e5c2a84b..b20226b5de 100644 --- a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/models.ts +++ b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/models.ts @@ -6,7 +6,7 @@ import type { ModelWorkspaceGroup } from '@nemo/common/src/api/models/useModels' import type { ModelSelection } from '@nemo/common/src/components/ModelSelectV2/types'; import { MAX_COMPLETION_TOKENS_DEFAULT } from '@nemo/common/src/constants/inferenceParameters'; import { getURNFromNamedEntityRef } from '@nemo/common/src/namedEntity'; -import { groupModelsByWorkspace } from '@nemo/common/src/utils/models'; +import { groupModelsByWorkspace, hasModelProvider } from '@nemo/common/src/utils/models'; import type { ChatCompletionInferenceParams, EmbeddingInferenceParams, @@ -51,6 +51,10 @@ const AUTO_FILL_PAGE_SIZE = 25; * The models {@link resolveTemplateModel} should consider: those whose name matches `preferred`, * plus the first page of the workspace as a fallback. Two small requests instead of walking the * whole catalogue, which is all the resolver needs to make its choice. + * + * Provider-less models are dropped: auto-fill happens without the user asking, so seeding one + * would hand them a recipe that fails at submit with "the model does not have a provider" + * (see {@link providerForSelection}) — better to leave the field empty and let them pick. */ export const fetchAutoFillCandidates = async ( workspace: string, @@ -79,6 +83,7 @@ export const fetchAutoFillCandidates = async ( const seen = new Set(); const models = [...matches, ...firstPage].filter((entity) => { + if (!hasModelProvider(entity)) return false; const urn = getURNFromNamedEntityRef(entity); if (!urn || seen.has(urn)) return false; seen.add(urn); diff --git a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/useJobBuilder.test.tsx b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/useJobBuilder.test.tsx index 2545d2bca0..d8a5910a41 100644 --- a/web/packages/studio/src/routes/DataDesignerJobBuildRoute/useJobBuilder.test.tsx +++ b/web/packages/studio/src/routes/DataDesignerJobBuildRoute/useJobBuilder.test.tsx @@ -66,6 +66,14 @@ describe('useJobBuilder template auto-fill', () => { await waitFor(() => expect(result.current.getBuilderValues().models[0].model).toBe('')); }); + it('skips provider-less models rather than seeding one that cannot run', async () => { + mockListModels.mockResolvedValue(makePage([model(NEMOTRON, [])])); + + const { result } = renderHook(() => useJobBuilder(template, 'ws1'), { wrapper: StrictMode }); + + await waitFor(() => expect(result.current.getBuilderValues().models[0].model).toBe('')); + }); + it('leaves models that already carry a provider alone', async () => { const seeded = { ...template, diff --git a/web/packages/studio/src/routes/ModelCompareRoute/index.test.tsx b/web/packages/studio/src/routes/ModelCompareRoute/index.test.tsx new file mode 100644 index 0000000000..771a2fc852 --- /dev/null +++ b/web/packages/studio/src/routes/ModelCompareRoute/index.test.tsx @@ -0,0 +1,133 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import { useModelEntity } from '@nemo/common/src/api/models/useModelEntity'; +import { useModelSearch } from '@nemo/common/src/api/models/useModelSearch'; +import type { ModelEntity } from '@nemo/sdk/generated/platform/schema'; +import { ModelCompareRoute } from '@studio/routes/ModelCompareRoute'; +import type { SharedModelEntry } from '@studio/routes/ModelCompareRoute/types'; +import { fireEvent, renderRoute, screen, waitFor } from '@studio/tests/util/render'; +import { useNavigate } from 'react-router-dom'; + +vi.mock('@nemo/common/src/api/models/useModelSearch', async (importOriginal) => { + const actual = + await importOriginal(); + return { ...actual, useModelSearch: vi.fn() }; +}); + +vi.mock('@nemo/common/src/api/models/useModelEntity', async (importOriginal) => { + const actual = + await importOriginal(); + return { ...actual, useModelEntity: vi.fn() }; +}); + +vi.mock('@studio/hooks/useWorkspaceFromPath', () => ({ + useWorkspaceFromPath: () => 'ws1', +})); + +vi.mock('@studio/components/ModelCompareChat', () => ({ + ModelCompareChat: ({ models }: { models: SharedModelEntry[] }) => ( +
{models.map((entry) => entry.modelURN ?? 'empty').join(',')}
+ ), +})); + +vi.mock('@studio/components/ModelComparePrompts', () => ({ + ModelComparePrompts: () => null, +})); + +const mockUseModelSearch = vi.mocked(useModelSearch); +const mockUseModelEntity = vi.mocked(useModelEntity); + +const entity = (workspace: string, name: string, providers?: string[]): ModelEntity => + ({ id: `${workspace}/${name}`, workspace, name, model_providers: providers }) as ModelEntity; + +const ENTITIES: Record = { + 'ws1/alpha': entity('ws1', 'alpha', ['ws1/build']), + 'ws1/beta': entity('ws1', 'beta', ['ws1/build']), + 'ws1/unserved': entity('ws1', 'unserved'), + 'other/foreign': entity('other', 'foreign', ['other/build']), +}; + +const searchResult = (overrides: Partial> = {}) => + ({ + models: [ENTITIES['ws1/alpha']], + groups: [{ workspace: 'ws1', models: [ENTITIES['ws1/alpha']] }], + search: '', + error: null, + loading: false, + onSearchChange: vi.fn(), + onLoadMore: vi.fn(), + hasMore: false, + isLoadingMore: false, + ...overrides, + }) as ReturnType; + +beforeEach(() => { + mockUseModelSearch.mockReset(); + mockUseModelEntity.mockReset(); + mockUseModelSearch.mockReturnValue(searchResult()); + mockUseModelEntity.mockImplementation((urn, options) => + options?.enabled === false || !urn ? undefined : ENTITIES[urn] + ); +}); + +describe('ModelCompareRoute availability', () => { + it('shows the no-models state only for an exhausted, successful search', async () => { + mockUseModelSearch.mockReturnValue(searchResult({ models: [], groups: [] })); + + renderRoute(, { history: '/compare' }); + + expect(await screen.findByText('No models available')).toBeInTheDocument(); + }); + + it('shows an error state instead of the no-models state when the search fails', async () => { + mockUseModelSearch.mockReturnValue( + searchResult({ models: [], groups: [], error: new Error('gateway unreachable') }) + ); + + renderRoute(, { history: '/compare' }); + + expect(await screen.findByText('gateway unreachable')).toBeInTheDocument(); + expect(screen.queryByText('No models available')).not.toBeInTheDocument(); + }); +}); + +describe('ModelCompareRoute ?model= preselection', () => { + it('preselects a model this workspace serves', async () => { + renderRoute(, { history: '/compare?model=alpha' }); + + await waitFor(() => expect(screen.getByTestId('panels')).toHaveTextContent('ws1/alpha,empty')); + }); + + it('ignores a model from another workspace', async () => { + renderRoute(, { history: '/compare?model=other/foreign' }); + + await waitFor(() => expect(screen.getByTestId('panels')).toHaveTextContent('empty,empty')); + }); + + it('ignores a model with no provider', async () => { + renderRoute(, { history: '/compare?model=unserved' }); + + await waitFor(() => expect(screen.getByTestId('panels')).toHaveTextContent('empty,empty')); + }); + + it('applies ?model= added after the first render', async () => { + const Harness = () => { + const navigate = useNavigate(); + return ( + <> + + + + ); + }; + + renderRoute(, { history: '/compare' }); + + await waitFor(() => expect(screen.getByTestId('panels')).toHaveTextContent('empty,empty')); + + fireEvent.click(screen.getByRole('button', { name: 'select beta' })); + + await waitFor(() => expect(screen.getByTestId('panels')).toHaveTextContent('ws1/beta,empty')); + }); +}); diff --git a/web/packages/studio/src/routes/ModelCompareRoute/index.tsx b/web/packages/studio/src/routes/ModelCompareRoute/index.tsx index 2e1d712c92..724cc707f1 100644 --- a/web/packages/studio/src/routes/ModelCompareRoute/index.tsx +++ b/web/packages/studio/src/routes/ModelCompareRoute/index.tsx @@ -5,6 +5,7 @@ import { useModelEntity } from '@nemo/common/src/api/models/useModelEntity'; import { useModelSearch } from '@nemo/common/src/api/models/useModelSearch'; import { ComposerMode } from '@nemo/common/src/components/AssistantChat'; import type { BroadcastSignal } from '@nemo/common/src/components/AssistantChat/types'; +import { ErrorMessage } from '@nemo/common/src/components/ErrorMessage'; import { hasModelProvider } from '@nemo/common/src/utils/models'; import { PageHeader, Tabs, Tooltip } from '@nvidia/foundations-react-core'; import { ChatEmptyState } from '@studio/components/chat/ChatEmptyState'; @@ -36,7 +37,10 @@ export const ModelCompareRoute: FC = () => { const workspace = useWorkspaceFromPath(); const availableModels = useModelSearch({ workspace, include: hasModelProvider }); const hasNoModels = - !availableModels.loading && !availableModels.hasMore && availableModels.models.length === 0; + !availableModels.loading && + !availableModels.error && + !availableModels.hasMore && + availableModels.models.length === 0; const [searchParams] = useSearchParams(); const [activeView, setActiveView] = useState('compare'); @@ -46,7 +50,7 @@ export const ModelCompareRoute: FC = () => { makeDefaultEntry(1), ]); const nextIdRef = useRef(2); - const didPreselectRef = useRef(false); + const preselectedUrnRef = useRef(null); // Preselect panel 0 from ?model=, which carries either a full URN or a bare name in this // workspace. Resolved with a single lookup rather than by scanning the whole catalogue. @@ -57,22 +61,22 @@ export const ModelCompareRoute: FC = () => { ? modelParam : `${workspace}/${modelParam}` : null; - const preselectModel = useModelEntity(preselectUrn, { enabled: !didPreselectRef.current }); + const preselectModel = useModelEntity(preselectUrn, { + enabled: !!preselectUrn && preselectedUrnRef.current !== preselectUrn, + }); useEffect(() => { - if (didPreselectRef.current) return; - if (!modelParam) { - didPreselectRef.current = true; - return; - } + if (!preselectUrn || preselectedUrnRef.current === preselectUrn) return; if (!preselectModel) return; - setModels((prev) => - prev.map((m, i) => - i === 0 ? { ...m, modelURN: `${preselectModel.workspace}/${preselectModel.name}` } : m - ) - ); - didPreselectRef.current = true; - }, [modelParam, preselectModel]); + if (preselectModel.workspace === workspace && hasModelProvider(preselectModel)) { + setModels((prev) => + prev.map((m, i) => + i === 0 ? { ...m, modelURN: `${preselectModel.workspace}/${preselectModel.name}` } : m + ) + ); + } + preselectedUrnRef.current = preselectUrn; + }, [preselectModel, preselectUrn, workspace]); // Seed transfer: broadcast→panels and panel→broadcast. const compareComposerDraftRef = useRef(''); @@ -158,6 +162,10 @@ export const ModelCompareRoute: FC = () => { const atMaxModels = models.length >= MAX_MODELS; const readyPanelCount = models.filter((m) => !!m.modelURN).length; + if (availableModels.error) { + return ; + } + // Empty state when the workspace has no servable models and we're not still loading. if (hasNoModels) { return ;