import { ThreadPrimitive, useAuiEvent, useAuiState } from '@assistant-ui/react' import { useVirtualizer, type Virtualizer } from '@tanstack/react-virtual' import { type ComponentProps, type FC, type ReactNode, useCallback, useEffect, useMemo, useRef } from 'react' import { cn } from '@/lib/utils' import { setThreadScrolledUp } from '@/store/thread-scroll' const ESTIMATED_ITEM_HEIGHT = 220 const OVERSCAN = 4 const AT_BOTTOM_THRESHOLD = 4 type ThreadMessageComponents = ComponentProps['components'] type MessageGroup = { id: string; index: number; kind: 'standalone' } | { id: string; indices: number[]; kind: 'turn' } interface VirtualizedThreadProps { clampToComposer: boolean components: ThreadMessageComponents emptyPlaceholder?: ReactNode loadingIndicator?: ReactNode sessionKey?: string | null } function buildGroups(signature: string): MessageGroup[] { if (!signature) { return [] } const messages = signature.split('\n').map(row => { const [index, id, role] = row.split(':') return { id, index: Number(index), role } }) const groups: MessageGroup[] = [] for (let i = 0; i < messages.length; i++) { const message = messages[i] if (message.role !== 'user') { groups.push({ id: message.id, index: message.index, kind: 'standalone' }) continue } const indices = [message.index] while (i + 1 < messages.length && messages[i + 1].role !== 'user') { indices.push(messages[++i].index) } groups.push({ id: message.id, indices, kind: 'turn' }) } return groups } export const VirtualizedThread: FC = ({ clampToComposer, components, emptyPlaceholder, loadingIndicator, sessionKey }) => { const messageSignature = useAuiState(s => s.thread.messages.map((message, index) => `${index}:${message.id}:${message.role}`).join('\n') ) const groups = useMemo(() => buildGroups(messageSignature), [messageSignature]) const renderEmpty = groups.length === 0 && Boolean(emptyPlaceholder) const scrollerRef = useRef(null) const virtualizer = useVirtualizer({ count: groups.length, estimateSize: () => ESTIMATED_ITEM_HEIGHT, getItemKey: index => groups[index]?.id ?? index, getScrollElement: () => scrollerRef.current, // Seed the rect so the initial range mounts something before // `observeElementRect` reports the real layout (it overrides this). initialRect: { height: 600, width: 800 }, overscan: OVERSCAN }) useThreadScrollAnchor({ enabled: !renderEmpty, groupCount: groups.length, scrollerRef, sessionKey: sessionKey ?? null, virtualizer }) const virtualItems = virtualizer.getVirtualItems() const totalSize = virtualizer.getTotalSize() const paddingTop = virtualItems[0]?.start ?? 0 const paddingBottom = Math.max(0, totalSize - (virtualItems.at(-1)?.end ?? 0)) return (
{renderEmpty ? (
{emptyPlaceholder}
) : (
{/* Natural-flow virtualization: mounted items render as normal flex siblings so `position: sticky` on the human bubble resolves against the scroller without transform interference. Padding spacers reserve scroll space for unmounted items. */}
{virtualItems.map(virtualItem => { const group = groups[virtualItem.index] if (!group) { return null } return (
{group.kind === 'turn' ? (
{group.indices.map(index => ( ))}
) : ( )}
) })}
{loadingIndicator} {clampToComposer && ( )}
) } interface ScrollAnchorOptions { enabled: boolean groupCount: number scrollerRef: React.RefObject sessionKey: string | null virtualizer: Virtualizer } function useThreadScrollAnchor({ enabled, groupCount, scrollerRef, sessionKey, virtualizer }: ScrollAnchorOptions) { // `armed` = parked at bottom, content growth should follow. Cleared on // user-driven upward scroll; re-armed when they reach bottom again. const armedRef = useRef(true) const lastTopRef = useRef(0) const prevSessionKeyRef = useRef(sessionKey) const prevGroupCountRef = useRef(0) const pinToBottom = useCallback(() => { const el = scrollerRef.current if (!el) { return } el.scrollTop = el.scrollHeight lastTopRef.current = el.scrollTop }, [scrollerRef]) const jumpToBottom = useCallback(() => { armedRef.current = true if (groupCount > 0) { virtualizer.scrollToIndex(groupCount - 1, { align: 'end', behavior: 'auto' }) } requestAnimationFrame(() => { if (armedRef.current) { pinToBottom() } }) }, [groupCount, pinToBottom, virtualizer]) useEffect(() => () => setThreadScrolledUp(false), []) // Track at-bottom state, dim composer when scrolled up, disarm on user // scroll/wheel/touch. useEffect(() => { const el = scrollerRef.current if (!el) { return undefined } const disarm = () => { armedRef.current = false } const onScroll = () => { const top = el.scrollTop if (top + 1 < lastTopRef.current) { armedRef.current = false } lastTopRef.current = top const atBottom = el.scrollHeight - (top + el.clientHeight) <= AT_BOTTOM_THRESHOLD if (atBottom) { armedRef.current = true } setThreadScrolledUp(!atBottom) } const onWheel = (event: WheelEvent) => { if (event.deltaY < 0) { disarm() } } el.addEventListener('scroll', onScroll, { passive: true }) el.addEventListener('wheel', onWheel, { passive: true }) el.addEventListener('touchmove', disarm, { passive: true }) return () => { el.removeEventListener('scroll', onScroll) el.removeEventListener('wheel', onWheel) el.removeEventListener('touchmove', disarm) } }, [scrollerRef]) // Follow content growth (streaming, item measurements, loading indicator) // while armed. useEffect(() => { if (!enabled) { return undefined } const el = scrollerRef.current if (!el) { return undefined } const observer = new ResizeObserver(() => { if (armedRef.current) { pinToBottom() } }) observer.observe(el) if (el.firstElementChild) { observer.observe(el.firstElementChild) } return () => observer.disconnect() }, [enabled, pinToBottom, scrollerRef]) // Jump to bottom on session change OR when an empty thread first gets // content. Both share the same intent and the same effect. useEffect(() => { const sessionChanged = prevSessionKeyRef.current !== sessionKey const becameNonEmpty = prevGroupCountRef.current === 0 && groupCount > 0 prevSessionKeyRef.current = sessionKey prevGroupCountRef.current = groupCount if (enabled && (sessionChanged || becameNonEmpty)) { jumpToBottom() } }, [enabled, groupCount, jumpToBottom, sessionKey]) useAuiEvent('thread.runStart', jumpToBottom) }