From cfd13544bd9f2484688aa3ea189a936a947416b1 Mon Sep 17 00:00:00 2001 From: Bohdan Triapitsyn Date: Wed, 27 May 2026 14:01:25 +0300 Subject: [PATCH] perf: reduce chat rerenders during streaming Stabilizes unchanged chat turns while messages stream Keeps static chat history from rerendering unnecessarily Adds coverage for turn record reuse --- .../ui/src/components/chat/MessageList.tsx | 28 +++++---- .../components/chat/hooks/useTurnRecords.ts | 7 ++- .../chat/lib/turns/projectTurnRecords.test.ts | 32 ++++++++++ .../chat/lib/turns/projectTurnRecords.ts | 58 ++++++++++++++++++- 4 files changed, 113 insertions(+), 12 deletions(-) diff --git a/packages/ui/src/components/chat/MessageList.tsx b/packages/ui/src/components/chat/MessageList.tsx index a1400d95..95c2674b 100644 --- a/packages/ui/src/components/chat/MessageList.tsx +++ b/packages/ui/src/components/chat/MessageList.tsx @@ -19,6 +19,9 @@ import { normalizeParts } from './message/partUtils'; const MESSAGE_LIST_VIRTUALIZE_THRESHOLD = 5; const MESSAGE_LIST_OVERSCAN = 6; +const EMPTY_STATIC_ENTRY_MESSAGES: ChatMessageEntry[] = []; +const EMPTY_UNGROUPED_MESSAGE_IDS = new Set(); +const EMPTY_VIRTUAL_ROWS: VirtualItem[] = []; const estimateHistoryEntryHeight = (entry: RenderEntry | undefined): number => { if (!entry) { @@ -936,7 +939,7 @@ const MessageListEntry = React.memo(({ MessageListEntry.displayName = 'MessageListEntry'; // Inner component that renders staged turn entries. -const StaticHistoryList: React.FC<{ +type StaticHistoryListProps = { entries: RenderEntry[]; shouldVirtualize: boolean; virtualRows: VirtualItem[]; @@ -954,7 +957,9 @@ const StaticHistoryList: React.FC<{ shouldAnimateUserMessage: (message: ChatMessageEntry) => boolean; onUserAnimationConsumed: (messageId: string) => void; activeStreamingPhase?: StreamPhase | null; -}> = ({ entries, shouldVirtualize, virtualRows, totalSize, measureElement, contentRef, onMessageContentChange, getAnimationHandlers, scrollToBottom, stickyUserHeader, defaultActivityExpanded, turnUiStates, onToggleTurnGroup, chatRenderMode, shouldAnimateUserMessage, onUserAnimationConsumed, activeStreamingPhase }) => { +}; + +const StaticHistoryList = React.memo(({ entries, shouldVirtualize, virtualRows, totalSize, measureElement, contentRef, onMessageContentChange, getAnimationHandlers, scrollToBottom, stickyUserHeader, defaultActivityExpanded, turnUiStates, onToggleTurnGroup, chatRenderMode, shouldAnimateUserMessage, onUserAnimationConsumed, activeStreamingPhase }: StaticHistoryListProps) => { const renderEntry = React.useCallback((entry: RenderEntry) => { return ( 0 ? ); -}; +}); StaticHistoryList.displayName = 'StaticHistoryList'; @@ -1222,6 +1227,9 @@ const MessageList = React.forwardRef(({ sessionKey, showTextJustificationActivity: chatRenderMode === 'sorted', }); + const hasUngroupedStaticEntries = projection.ungroupedMessageIds.size > 0; + const staticEntryMessages = hasUngroupedStaticEntries ? displayMessages : EMPTY_STATIC_ENTRY_MESSAGES; + const staticEntryUngroupedIds = hasUngroupedStaticEntries ? projection.ungroupedMessageIds : EMPTY_UNGROUPED_MESSAGE_IDS; const staticRenderEntries = React.useMemo(() => streamPerfMeasure('ui.message_list.render_entries_ms', () => { const turnEntries = staticTurns.map((turn) => ({ kind: 'turn' as const, @@ -1230,7 +1238,7 @@ const MessageList = React.forwardRef(({ isLastTurn: turn.turnId === projection.lastTurnId, })); - if (projection.ungroupedMessageIds.size === 0) { + if (staticEntryUngroupedIds.size === 0) { return turnEntries; } @@ -1240,14 +1248,14 @@ const MessageList = React.forwardRef(({ }); const orderedEntries: RenderEntry[] = []; - displayMessages.forEach((message, index) => { + staticEntryMessages.forEach((message, index) => { const turnEntry = turnEntryByUserMessageId.get(message.info.id); if (turnEntry) { orderedEntries.push(turnEntry); return; } - if (!projection.ungroupedMessageIds.has(message.info.id)) { + if (!staticEntryUngroupedIds.has(message.info.id)) { return; } @@ -1255,13 +1263,13 @@ const MessageList = React.forwardRef(({ kind: 'ungrouped', key: `msg:${message.info.id}`, message, - previousMessage: index > 0 ? displayMessages[index - 1] : undefined, - nextMessage: index < displayMessages.length - 1 ? displayMessages[index + 1] : undefined, + previousMessage: index > 0 ? staticEntryMessages[index - 1] : undefined, + nextMessage: index < staticEntryMessages.length - 1 ? staticEntryMessages[index + 1] : undefined, }); }); return orderedEntries; - }), [displayMessages, projection.lastTurnId, projection.ungroupedMessageIds, staticTurns]); + }), [projection.lastTurnId, staticEntryMessages, staticEntryUngroupedIds, staticTurns]); const trailingStreamingEntry = React.useMemo(() => { if (streamingTurn) { @@ -1428,7 +1436,7 @@ const MessageList = React.forwardRef(({ }, []); const historyVirtualRows = React.useMemo( - () => (shouldVirtualizeHistory ? historyVirtualizer.getVirtualItems() : []), + () => (shouldVirtualizeHistory ? historyVirtualizer.getVirtualItems() : EMPTY_VIRTUAL_ROWS), [historyVirtualizer, shouldVirtualizeHistory], ); diff --git a/packages/ui/src/components/chat/hooks/useTurnRecords.ts b/packages/ui/src/components/chat/hooks/useTurnRecords.ts index 1aea08a6..2879107b 100644 --- a/packages/ui/src/components/chat/hooks/useTurnRecords.ts +++ b/packages/ui/src/components/chat/hooks/useTurnRecords.ts @@ -22,9 +22,14 @@ export const useTurnRecords = ( const staticTurnsRef = React.useRef([]); const streamingTurnRef = React.useRef(undefined); const previousSessionKeyRef = React.useRef(options.sessionKey); + const previousShowTextJustificationActivityRef = React.useRef(options.showTextJustificationActivity); - if (previousSessionKeyRef.current !== options.sessionKey) { + if ( + previousSessionKeyRef.current !== options.sessionKey + || previousShowTextJustificationActivityRef.current !== options.showTextJustificationActivity + ) { previousSessionKeyRef.current = options.sessionKey; + previousShowTextJustificationActivityRef.current = options.showTextJustificationActivity; previousProjectionRef.current = null; staticTurnsRef.current = []; streamingTurnRef.current = undefined; diff --git a/packages/ui/src/components/chat/lib/turns/projectTurnRecords.test.ts b/packages/ui/src/components/chat/lib/turns/projectTurnRecords.test.ts index 2806a42c..20f94859 100644 --- a/packages/ui/src/components/chat/lib/turns/projectTurnRecords.test.ts +++ b/packages/ui/src/components/chat/lib/turns/projectTurnRecords.test.ts @@ -86,4 +86,36 @@ describe('projectTurnRecords', () => { expect(projection.turns).toHaveLength(0); expect(projection.ungroupedMessageIds.has('s1')).toBe(true); }); + + test('reuses unchanged turn records from the previous projection', () => { + const user1 = createMessageEntry({ id: 'u1', role: 'user', createdAt: 1 }); + const assistant1 = createMessageEntry({ id: 'a1', role: 'assistant', parentID: 'u1', createdAt: 2 }); + const user2 = createMessageEntry({ id: 'u2', role: 'user', createdAt: 3 }); + const assistant2 = createMessageEntry({ id: 'a2', role: 'assistant', parentID: 'u2', createdAt: 4 }); + const initial = projectTurnRecords([user1, assistant1, user2, assistant2]); + const updatedAssistant2 = { + ...assistant2, + parts: [{ type: 'text', text: 'stream update' } as Part], + }; + + const next = projectTurnRecords([user1, assistant1, user2, updatedAssistant2], { + previousProjection: initial, + }); + + expect(next.turns[0]).toBe(initial.turns[0]); + expect(next.turns[1]).not.toBe(initial.turns[1]); + }); + + test('reuses the whole turns array when every turn is unchanged', () => { + const user = createMessageEntry({ id: 'u1', role: 'user', createdAt: 1 }); + const assistant = createMessageEntry({ id: 'a1', role: 'assistant', parentID: 'u1', createdAt: 2 }); + const initial = projectTurnRecords([user, assistant]); + + const next = projectTurnRecords([user, assistant], { + previousProjection: initial, + }); + + expect(next.turns).toBe(initial.turns); + expect(next.turns[0]).toBe(initial.turns[0]); + }); }); diff --git a/packages/ui/src/components/chat/lib/turns/projectTurnRecords.ts b/packages/ui/src/components/chat/lib/turns/projectTurnRecords.ts index 7358de2e..0fee35eb 100644 --- a/packages/ui/src/components/chat/lib/turns/projectTurnRecords.ts +++ b/packages/ui/src/components/chat/lib/turns/projectTurnRecords.ts @@ -90,6 +90,61 @@ const DEFAULT_OPTIONS: ProjectTurnRecordsOptions = { showTextJustificationActivity: false, }; +const areSameMessageRefs = (left: ChatMessageEntry[], right: ChatMessageEntry[]): boolean => { + if (left === right) { + return true; + } + if (left.length !== right.length) { + return false; + } + + for (let index = 0; index < left.length; index += 1) { + if (left[index] !== right[index]) { + return false; + } + } + + return true; +}; + +const canReusePreviousTurn = (previous: TurnRecord, next: TurnRecord): boolean => { + return previous.userMessage === next.userMessage + && previous.headerMessageId === next.headerMessageId + && areSameMessageRefs(previous.assistantMessages, next.assistantMessages); +}; + +const stabilizeTurnRecords = ( + turns: TurnRecord[], + previousProjection?: TurnProjectionResult | null, +): TurnRecord[] => { + if (!previousProjection || previousProjection.turns.length === 0 || turns.length === 0) { + return turns; + } + + let canReuseTurnArray = previousProjection.turns.length === turns.length; + let reusedAnyTurn = false; + + const nextTurns = turns.map((turn, index) => { + const previousTurn = previousProjection.indexes.turnById.get(turn.turnId); + if (previousTurn && canReusePreviousTurn(previousTurn, turn)) { + reusedAnyTurn = true; + if (previousProjection.turns[index] !== previousTurn) { + canReuseTurnArray = false; + } + return previousTurn; + } + + canReuseTurnArray = false; + return turn; + }); + + if (canReuseTurnArray && reusedAnyTurn) { + return previousProjection.turns; + } + + return reusedAnyTurn ? nextTurns : turns; +}; + export const projectTurnRecords = ( messages: ChatMessageEntry[], options?: Partial, @@ -179,7 +234,8 @@ export const projectTurnRecords = ( turn.durationMs = turn.stream.durationMs; }); - const projection = projectTurnIndexes(turns); + const stableTurns = stabilizeTurnRecords(turns, effectiveOptions.previousProjection); + const projection = projectTurnIndexes(stableTurns); const ungroupedMessageIds = new Set(); messages.forEach((message) => { if (resolveMessageRole(message) === 'assistant') {