import { useMemo, useState } from 'react'; import { useFilter } from 'react-aria'; import type { Collection, Key, ListProps, ListState, Node } from 'react-stately'; import { useListState } from 'react-stately'; import { ListCollection } from '@react-stately/list'; /** Props for the useMultiSelectListState hook. */ export interface MultiSelectListProps extends Omit, 'selectionMode'> { /** List of items to display */ items?: Iterable; /** Handler that is called when the search query text changes. */ onQueryChange?: (query: string) => void; /** The query to search for (controlled). */ query?: string; /** The default query to search for (uncontrolled). */ defaultQuery?: string; } /** State returned from the useMultiSelectListState hook. */ export interface MultiSelectListState extends ListState { /** The unfiltered collection */ originalCollection: Collection>; /** Search query for collection */ searchQuery: string; /** Sets the search query */ setSearchQuery: (query: string) => void; /** Toggles selection of filtered items only */ toggleFiltered: () => void; /** Whether all filtered items are selected */ isFilteredSelectAll: boolean; } export function useMultiSelectListState({ items, query: queryProp, defaultQuery, onQueryChange, ...listStateProps }: MultiSelectListProps): MultiSelectListState { const { collection, disabledKeys, selectionManager } = useListState({ ...listStateProps, items: items ?? [], selectionMode: 'multiple', }); const [localQuery, setLocalQuery] = useState(queryProp ?? defaultQuery ?? ''); const searchQuery = queryProp ?? localQuery; const { contains } = useFilter({ sensitivity: 'base' }); const isFilterUncontrolled = queryProp === undefined; const displayCollection = useMemo( () => (isFilterUncontrolled ? filterCollection(collection, searchQuery, contains) : collection), [collection, searchQuery, contains, isFilterUncontrolled] ); return { collection: displayCollection, originalCollection: collection, disabledKeys, selectionManager, searchQuery, setSearchQuery: query => { setLocalQuery(query); selectionManager.setFocusedKey(null); onQueryChange?.(query); }, toggleFiltered: () => { const selectedKeys = selectionManager.selectedKeys; const filteredKeys = getAllItemKeys(displayCollection).filter(key => !disabledKeys.has(key)); // If all filtered items are selected, deselect them, but keep the rest unchanged if (filteredKeys.every(key => selectedKeys.has(key))) { selectionManager.setSelectedKeys([...selectedKeys].filter(key => !filteredKeys.includes(key))); } else { // Otherwise, add all filtered items to the selection, but keep the rest unchanged selectionManager.setSelectedKeys(new Set([...selectedKeys, ...filteredKeys])); } }, get isFilteredSelectAll() { const selectedKeys = selectionManager.selectedKeys; const filteredKeys = getAllItemKeys(displayCollection).filter(key => !disabledKeys.has(key)); return filteredKeys.every(key => selectedKeys.has(key)); }, }; } type Filter = (textValue: string, inputValue: string) => boolean; function filterNodes( collection: Collection>, nodes: Iterable>, searchValue: string, filter: Filter ): Iterable> { const filteredNodes: Node[] = []; for (const node of nodes) { if (node.type === 'section' && node.hasChildNodes) { const filtered = filterNodes(collection, collection.getChildren?.(node.key) ?? [], searchValue, filter); if ([...filtered].some(node => node.type === 'item')) { filteredNodes.push({ ...node, childNodes: filtered }); } } else if (node.type === 'item' && filter(node.textValue, searchValue)) { filteredNodes.push({ ...node }); } else if (node.type !== 'item') { filteredNodes.push({ ...node }); } } return filteredNodes; } function filterCollection( collection: Collection>, searchValue: string, filter: Filter ): Collection> { return new ListCollection(filterNodes(collection, collection, searchValue, filter)); } function getAllItemKeys(collection: Collection>, key?: Key): Key[] { const keys: Key[] = []; const nodes = key ? (collection.getChildren?.(key) ?? []) : collection; for (const node of nodes) { if (node.type === 'section' && node.hasChildNodes) { keys.push(...getAllItemKeys(collection, node.key)); } else if (node.type === 'item') { keys.push(node.key); } } return keys; }