Merge branch 'chore/assistant-performance' into chore/assistant-tool-nesting

Resolve conflicts with the memoized message feed:
- Include isActive in MessagePartSwitcher's memo comparator so the shimmer
  moves to each new tool call.
- Snapshot parts once in MessageDisplayContent before grouping, so rows
  inside tool groups get the same protection from SDK mutations.
- Expand the tool group in the reasoning mutation test, since reasoning
  rows now render inside a collapsed group.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
Saxon FletcherandClaude Opus 5.5 committed 2026-09-25 16:12:21 +10:00
commit 7fef8a145b
14 files changed
+646 -147

No files matched your search

@@ -33,6 +33,7 @@ import { Markdown } from '@/components/interfaces/Markdown'
import { useCheckOpenAIKeyQuery } from '@/data/ai/check-api-key-query'
import { useRateMessageMutation } from '@/data/ai/rate-message-mutation'
import { useTablesQuery } from '@/data/tables/tables-query'
import { useLatest } from '@/hooks/misc/useLatest'
import { useLocalStorageQuery } from '@/hooks/misc/useLocalStorage'
import { useOrgAiOptInLevel } from '@/hooks/misc/useOrgOptedIntoAi'
import { useSelectedOrganizationQuery } from '@/hooks/misc/useSelectedOrganization'
@@ -172,6 +173,8 @@ export const AssistantChat = ({
regenerate,
} = useChat<MessageType>({
id: chatId,
// Batch token updates without throttling the SDK's tool execution or approval state.
throttle: 50,
...(chatInstance ? { chat: chatInstance } : {}),
sendAutomaticallyWhen: lastAssistantMessageIsCompleteWithApprovalResponses,
onError: onErrorChat,
@@ -186,32 +189,36 @@ export const AssistantChat = ({
const isChatInputDisabled =
!isApiKeySet || disablePrompts || isLoadingOrganization || isSupportChatClosed
const messagesRef = useLatest(chatMessages)
const isChatLoadingRef = useLatest(isChatLoading)
const branchedFrom = currentChat?.branchedFrom
const branchedConversation = branchedFrom ? snap.chats[branchedFrom.chatId] : undefined
const deleteMessageFromHere = useCallback(
(messageId: string) => {
// Find the message index in current chatMessages
const messageIndex = chatMessages.findIndex((msg) => msg.id === messageId)
const messages = messagesRef.current
const messageIndex = messages.findIndex((msg) => msg.id === messageId)
if (messageIndex === -1) return
if (isChatLoading) stop()
if (isChatLoadingRef.current) stop()
snap.deleteMessagesAfter(messageId, { includeSelf: true, chatId })
state.deleteMessagesAfter(messageId, { includeSelf: true, chatId })
const updatedMessages = chatMessages.slice(0, messageIndex)
const updatedMessages = messages.slice(0, messageIndex)
setMessages(updatedMessages)
},
[snap, setMessages, chatMessages, isChatLoading, stop, chatId]
[state, setMessages, messagesRef, isChatLoadingRef, stop, chatId]
)
const editMessage = useCallback(
(messageId: string) => {
const messageIndex = chatMessages.findIndex((msg) => msg.id === messageId)
const messages = messagesRef.current
const messageIndex = messages.findIndex((msg) => msg.id === messageId)
if (messageIndex === -1) return
// Target message
const messageToEdit = chatMessages[messageIndex]
const messageToEdit = messages[messageIndex]
// Activate editing mode
setEditingMessageId(messageId)
@@ -233,7 +240,7 @@ export const AssistantChat = ({
}
}, 100)
},
[chatMessages, setValue]
[messagesRef, setValue]
)
const cancelEdit = useCallback(() => {
@@ -251,7 +258,7 @@ export const AssistantChat = ({
try {
const result = await rateMessage({
rating,
messages: chatMessages,
messages: messagesRef.current,
messageId,
projectRef: project.ref,
orgSlug: selectedOrganization.slug,
@@ -274,7 +281,7 @@ export const AssistantChat = ({
})
}
},
[chatMessages, project?.ref, selectedOrganization?.slug, rateMessage, track, state, chatId]
[messagesRef, project?.ref, selectedOrganization?.slug, rateMessage, track, state, chatId]
)
const isContextExceededError =
@@ -282,13 +289,15 @@ export const AssistantChat = ({
(error.message?.includes('context_length_exceeded') ||
error.message?.includes('exceeds the context window'))
const editedMessageIndex = editingMessageId
? chatMessages.findIndex((message) => message.id === editingMessageId)
: -1
const renderedMessages = useMemo(
() =>
chatMessages.map((message, index) => {
const isBeingEdited = editingMessageId === message.id
const isAfterEditedMessage = editingMessageId
? chatMessages.findIndex((m) => m.id === editingMessageId) < index
: false
const isAfterEditedMessage = !!editingMessageId && editedMessageIndex < index
const isLastMessage = index === chatMessages.length - 1
return (
@@ -334,6 +343,7 @@ export const AssistantChat = ({
editMessage,
cancelEdit,
editingMessageId,
editedMessageIndex,
chatStatus,
addToolApprovalResponse,
handleRateMessage,
@@ -516,7 +526,9 @@ export const AssistantChat = ({
onStop={() => {
stop()
// to save partial responses from the AI
const lastMessage = chatMessages[chatMessages.length - 1]
// Read the live SDK state: the rendered snapshot may trail the stream by 50ms.
const messages = chatInstance?.messages ?? chatMessages
const lastMessage = messages[messages.length - 1]
if (lastMessage && lastMessage.role === 'assistant') {
state.updateMessage(lastMessage, chatId)
}
@@ -1,4 +1,4 @@
import { useRef, useState } from 'react'
import { useMemo, useRef, useState } from 'react'
import { identifyQueryType } from './AIAssistant.utils'
import {
@@ -106,15 +106,18 @@ export const AssistantQueryCell = ({
}
const result = resultOverride === undefined ? initialResult : (resultOverride ?? undefined)
const display =
localDisplay ??
getAssistantQueryDisplay({
view,
xAxis,
yAxis,
sql: query.uncheckedSql,
rows: result?.rows,
})
const inferredDisplay = useMemo(
() =>
getAssistantQueryDisplay({
view,
xAxis,
yAxis,
sql: query.uncheckedSql,
rows: result?.rows,
}),
[view, xAxis, yAxis, query.uncheckedSql, result?.rows]
)
const display = localDisplay ?? inferredDisplay
const handleTitleChange = (value: string) => {
const nextTitle = value.trim()
@@ -165,11 +168,12 @@ export const AssistantQueryCell = ({
onCancel={onDeny}
onConfirm={onApprove}
>
{/* Keep editor state mounted; the fixed height preserves scroll geometry when skipped. */}
<QueryEditor
isReadOnly
id={id}
variant="viewport"
className="h-96"
className="h-96 [content-visibility:auto] [contain-intrinsic-block-size:auto_24rem]"
title={title}
query={query}
result={result}
@@ -1,5 +1,5 @@
import { UIMessage as VercelMessage } from '@ai-sdk/react'
import { type PropsWithChildren } from 'react'
import { memo, type PropsWithChildren } from 'react'
import { cn } from 'ui'
import { useMessageInfoContext } from './Message.Context'
@@ -42,7 +42,11 @@ function MessageDisplayMainArea({
return <div className={cn('flex gap-4 w-auto overflow-hidden group', className)}>{children}</div>
}
function MessageDisplayContent({ message }: { message: VercelMessage }) {
const MessageDisplayContent = memo(function MessageDisplayContent({
message,
}: {
message: VercelMessage
}) {
const { id, isLoading, isLastMessage, readOnly } = useMessageInfoContext()
const messageParts = message.parts
@@ -50,7 +54,9 @@ function MessageDisplayContent({ message }: { message: VercelMessage }) {
('content' in message && typeof message.content === 'string' && message.content.trim()) ||
undefined
const items = groupMessageParts(messageParts ?? [])
// The SDK exposes its mutable object on the first write, then publishes clones. Capture
// state/text now so later mutations cannot change memoized parts' previous props.
const items = groupMessageParts((messageParts ?? []).map((part) => ({ ...part })))
const isStreaming = isLoading && !!isLastMessage
return (
@@ -78,7 +84,7 @@ function MessageDisplayContent({ message }: { message: VercelMessage }) {
)}
</div>
)
}
})
function MessageDisplayTextMessage({
id,
@@ -105,6 +105,29 @@ describe('MessagePartToolGroup', () => {
expect(toolRows[1].querySelector('.shimmer')).toBeInTheDocument()
})
it('moves the shimmer to each new tool call as it arrives', async () => {
const user = userEvent.setup()
const { container, rerender } = renderInMessage(
<MessagePartToolGroup parts={[reasoningPart, toolPart]} isRunning={true} />
)
await user.click(screen.getByRole('button', { name: 'Ran load_knowledge' }))
const nextReasoningPart = { ...reasoningPart, state: 'streaming' as const }
rerender(
inMessage(
<MessagePartToolGroup
parts={[reasoningPart, toolPart, nextReasoningPart]}
isRunning={true}
/>
)
)
const toolRows = container.querySelectorAll('.tool-item')
expect(toolRows).toHaveLength(3)
expect(toolRows[1].querySelector('.shimmer')).toBeNull()
expect(toolRows[2].querySelector('.shimmer')).toBeInTheDocument()
})
it('shimmers the summary only while running', async () => {
const { rerender } = renderInMessage(
<MessagePartToolGroup parts={[reasoningPart, toolPart]} isRunning={true} />
@@ -1,7 +1,7 @@
import { UIMessage as VercelMessage } from '@ai-sdk/react'
import { type DynamicToolUIPart, type ReasoningUIPart, type TextUIPart, type ToolUIPart } from 'ai'
import { BrainIcon, CheckIcon, Loader2 } from 'lucide-react'
import { useState, type ReactNode } from 'react'
import { memo, useState, type ReactNode } from 'react'
import { cn } from 'ui'
import { AssistantQueryCell } from './AssistantQueryCell'
@@ -11,7 +11,12 @@ import { EdgeFunctionRenderer } from './EdgeFunctionRenderer'
import { Tool } from './elements/Tool'
import { ToolGroup } from './elements/ToolGroup'
import { useMessageActionsContext, useMessageInfoContext } from './Message.Context'
import { getCompactPartLabel, getMessagePartKind, getToolGroupHeader } from './Message.Parts.utils'
import {
areMessagePartsEqual,
getCompactPartLabel,
getMessagePartKind,
getToolGroupHeader,
} from './Message.Parts.utils'
import {
deployEdgeFunctionInputSchema,
deployEdgeFunctionOutputSchema,
@@ -324,69 +329,73 @@ const isWideMessagePart = (part: NonNullable<VercelMessage['parts']>[number]) =>
// Unlabelled code fences resolve to SQL in MessageMarkdown, too.
(part.type === 'text' && /```(?:sql)?(?:\s|$)/i.test(part.text))
export function MessagePartSwitcher({
part,
isActive,
}: {
part: NonNullable<VercelMessage['parts']>[number]
/** Marks the in-progress tool call within a running tool group. */
isActive?: boolean
}) {
const content = (() => {
switch (part.type) {
case 'dynamic-tool': {
if (part.toolName === 'query_logs') {
export const MessagePartSwitcher = memo(
function MessagePartSwitcher({
part,
isActive,
}: {
part: NonNullable<VercelMessage['parts']>[number]
/** Marks the in-progress tool call within a running tool group. */
isActive?: boolean
}) {
const content = (() => {
switch (part.type) {
case 'dynamic-tool': {
if (part.toolName === 'query_logs') {
return <MessagePart.QueryLogs toolPart={part} />
}
return <MessagePart.Tool toolPart={part} isActive={isActive} />
}
case 'tool-list_policies':
case 'tool-search_docs':
case 'tool-get_active_incidents':
case 'tool-load_knowledge': {
return <MessagePart.Tool toolPart={part} isActive={isActive} />
}
case 'reasoning':
return <MessagePart.Reasoning reasoningPart={part} isActive={isActive} />
case 'text':
return <MessagePart.Text textPart={part} />
case 'tool-execute_sql': {
return <MessagePart.ExecuteSql toolPart={part} />
}
case 'tool-query_logs': {
return <MessagePart.QueryLogs toolPart={part} />
}
return <MessagePart.Tool toolPart={part} isActive={isActive} />
}
case 'tool-list_policies':
case 'tool-search_docs':
case 'tool-get_active_incidents':
case 'tool-load_knowledge': {
return <MessagePart.Tool toolPart={part} isActive={isActive} />
}
case 'reasoning':
return <MessagePart.Reasoning reasoningPart={part} isActive={isActive} />
case 'text':
return <MessagePart.Text textPart={part} />
case 'tool-deploy_edge_function': {
return <MessagePart.DeployEdgeFunction toolPart={part} />
}
case 'tool-create_notebook': {
return <MessagePart.NotebookProposal toolPart={part} mode="create" />
}
case 'tool-update_notebook': {
return <MessagePart.NotebookProposal toolPart={part} mode="update" />
}
case 'tool-delete_notebook': {
return <MessagePart.NotebookProposal toolPart={part} mode="delete" />
}
case 'tool-run_notebook': {
return <MessagePart.NotebookRun toolPart={part} />
}
case 'tool-execute_sql': {
return <MessagePart.ExecuteSql toolPart={part} />
}
case 'tool-query_logs': {
return <MessagePart.QueryLogs toolPart={part} />
}
case 'tool-deploy_edge_function': {
return <MessagePart.DeployEdgeFunction toolPart={part} />
}
case 'tool-create_notebook': {
return <MessagePart.NotebookProposal toolPart={part} mode="create" />
}
case 'tool-update_notebook': {
return <MessagePart.NotebookProposal toolPart={part} mode="update" />
}
case 'tool-delete_notebook': {
return <MessagePart.NotebookProposal toolPart={part} mode="delete" />
}
case 'tool-run_notebook': {
return <MessagePart.NotebookRun toolPart={part} />
case 'source-url':
case 'source-document':
case 'file':
default:
return null
}
})()
case 'source-url':
case 'source-document':
case 'file':
default:
return null
}
})()
if (content === null) return null
// Tool rows depend on being direct siblings to share their compact spacing and dividers.
if (getMessagePartKind(part) === 'compact') return content
if (content === null) return null
// Tool rows depend on being direct siblings to share their compact spacing and dividers.
if (getMessagePartKind(part) === 'compact') return content
return <MessagePartContainer isWide={isWideMessagePart(part)}>{content}</MessagePartContainer>
}
return <MessagePartContainer isWide={isWideMessagePart(part)}>{content}</MessagePartContainer>
},
(previous, next) =>
previous.isActive === next.isActive && areMessagePartsEqual(previous.part, next.part)
)
export function MessagePartToolGroup({
parts,
@@ -1,13 +1,92 @@
import type { UIMessage } from 'ai'
import { describe, expect, it } from 'vitest'
import type { DynamicToolUIPart, ToolUIPart, UIMessage } from 'ai'
import { describe, expect, it, vi } from 'vitest'
import {
areMessagePartsEqual,
getCompactPartLabel,
getMessagePartKind,
getToolGroupHeader,
groupMessageParts,
} from './Message.Parts.utils'
const completedTool = {
type: 'tool-execute_sql',
toolCallId: 'query-1',
state: 'output-available',
input: { sql: 'select 1' },
output: [{ value: 1 }],
} satisfies ToolUIPart
describe('areMessagePartsEqual', () => {
it.each(['static', 'dynamic'])('does not traverse finalized %s tool output', (kind) => {
const readRows = vi.fn(() => [{ value: 1 }])
const output = () => ({
get rows() {
return readRows()
},
})
const tool =
kind === 'static'
? completedTool
: { ...completedTool, type: 'dynamic-tool' as const, toolName: 'query_logs' }
expect(areMessagePartsEqual({ ...tool, output: output() }, { ...tool, output: output() })).toBe(
true
)
expect(readRows).not.toHaveBeenCalled()
})
const changedTools: Array<[string, ToolUIPart | DynamicToolUIPart]> = [
['tool identity', { ...completedTool, toolCallId: 'query-2' }],
['tool type', { ...completedTool, type: 'tool-query_logs' }],
['input', { ...completedTool, input: { sql: 'select 2' } }],
['approval', { ...completedTool, approval: { id: 'approval-1', approved: true } }],
['metadata', { ...completedTool, toolMetadata: { title: 'Updated title' } }],
['preliminary flag', { ...completedTool, preliminary: true }],
[
'error state',
{
type: 'tool-execute_sql',
toolCallId: 'query-1',
state: 'output-error',
input: completedTool.input,
errorText: 'Query failed',
},
],
]
it.each(changedTools)('rerenders when %s changes', (_label, next) => {
expect(areMessagePartsEqual(completedTool, next)).toBe(false)
})
it('updates preliminary results before the tool state changes', () => {
const previous = { ...completedTool, preliminary: true }
expect(areMessagePartsEqual(previous, { ...previous, output: [{ value: 2 }] })).toBe(false)
})
it('renders the final result after preliminary output', () => {
expect(areMessagePartsEqual({ ...completedTool, preliminary: true }, completedTool)).toBe(false)
})
it('checks live text and reasoning state', () => {
expect(
areMessagePartsEqual({ type: 'text', text: 'Hello' }, { type: 'text', text: 'Hello again' })
).toBe(false)
expect(
areMessagePartsEqual(
{ type: 'reasoning', text: '', state: 'streaming' },
{ type: 'reasoning', text: '', state: 'done' }
)
).toBe(false)
})
it('keeps identical and cloned unchanged parts memoized', () => {
expect(areMessagePartsEqual(completedTool, completedTool)).toBe(true)
const text = { type: 'text' as const, text: 'Hello' }
expect(areMessagePartsEqual(text, { ...text })).toBe(true)
})
})
type MessagePart = UIMessage['parts'][number]
const reasoning = (text = 'Thinking about it'): MessagePart => ({
@@ -1,7 +1,31 @@
import { getToolName, isToolUIPart, type UIMessage } from 'ai'
import isEqual from 'lodash/isEqual'
type MessagePart = UIMessage['parts'][number]
export function areMessagePartsEqual(previous: MessagePart, next: MessagePart): boolean {
if (previous === next) return true
if (previous.type !== next.type) return false
if (
isToolUIPart(previous) &&
isToolUIPart(next) &&
previous.state === 'output-available' &&
next.state === 'output-available' &&
!previous.preliminary &&
!next.preliminary
) {
// Final output is fixed for a tool call, but the SDK clones it on every text update.
// Keep checking identity, input, approval and metadata without walking result rows.
const { output: _previousOutput, ...previousFields } = previous
const { output: _nextOutput, ...nextFields } = next
return isEqual(previousFields, nextFields)
}
// Preliminary output and live text/reasoning can still change without a state transition.
return isEqual(previous, next)
}
/**
* - `compact`: a one-line tool row (reasoning, lookups) that gets folded into a tool group
* - `block`: content that renders on its own (text, SQL results, notebooks, Edge Functions)
@@ -0,0 +1,234 @@
import { fireEvent, render, screen } from '@testing-library/react'
import type { UIMessage } from 'ai'
import { useState } from 'react'
import { beforeEach, describe, expect, it, vi } from 'vitest'
import { Message } from './Message'
const { renderQuery, renderFunction, renderMarkdown } = vi.hoisted(() => ({
renderQuery: vi.fn(),
renderFunction: vi.fn(),
renderMarkdown: vi.fn(),
}))
// Keep the message/context/part pipeline real, substituting stateful probes for the
// expensive leaves so we can detect both extra renders and lost local state.
vi.mock('./AssistantQueryCell', () => ({
AssistantQueryCell: (props: { initialResult?: { rows: unknown[] }; onApprove?: () => void }) => {
renderQuery(props)
const [runs, setRuns] = useState(0)
return (
<div>
<button tabIndex={0} onClick={() => setRuns((count) => count + 1)}>
Run count: {runs}
</button>
<button tabIndex={0} onClick={props.onApprove}>
Approve query
</button>
<output>{JSON.stringify(props.initialResult?.rows)}</output>
</div>
)
},
}))
vi.mock('./EdgeFunctionRenderer', () => ({
EdgeFunctionRenderer: (props: { code: string }) => {
renderFunction(props)
return <pre>{props.code}</pre>
},
}))
vi.mock('./MessageMarkdown', () => ({
MessageMarkdown: ({ children }: { children: string }) => {
renderMarkdown(children)
return <p>{children}</p>
},
}))
vi.mock('./Message.Actions', () => ({ MessageActions: () => null }))
const initialMessage: UIMessage = {
id: 'assistant-1',
role: 'assistant',
parts: [
{
type: 'tool-execute_sql',
toolCallId: 'query-1',
state: 'approval-requested',
approval: { id: 'approval-1' },
input: { sql: 'select 1', label: 'Query', view: 'table' },
},
{
type: 'tool-deploy_edge_function',
toolCallId: 'function-1',
state: 'input-available',
input: {
code: 'Deno.serve(() => new Response("hello"))',
label: 'Function',
functionName: 'hello',
},
},
{ type: 'text', text: 'Working' },
],
}
const callbacks = {
onDelete: vi.fn(),
onEdit: vi.fn(),
onBranch: vi.fn(),
onCancelEdit: vi.fn(),
addToolApprovalResponse: vi.fn(),
}
function FeedMessage({
message = initialMessage,
isLoading = true,
isLastMessage,
addToolApprovalResponse = callbacks.addToolApprovalResponse,
}: {
message?: UIMessage
isLoading?: boolean
isLastMessage?: boolean
addToolApprovalResponse?: typeof callbacks.addToolApprovalResponse
}) {
return (
<Message
{...callbacks}
id={message.id}
message={message}
isLoading={isLoading}
isLastMessage={isLastMessage}
isAfterEditedMessage={false}
isBeingEdited={false}
addToolApprovalResponse={addToolApprovalResponse}
/>
)
}
describe('assistant feed rendering', () => {
beforeEach(() => vi.clearAllMocks())
it('does not rerender unchanged history when the parent renders', () => {
const { rerender } = render(<FeedMessage />)
const counts = [
renderQuery.mock.calls.length,
renderFunction.mock.calls.length,
renderMarkdown.mock.calls.length,
]
rerender(<FeedMessage />)
expect([
renderQuery.mock.calls.length,
renderFunction.mock.calls.length,
renderMarkdown.mock.calls.length,
]).toEqual(counts)
})
it('updates streamed text without rerendering cloned tools or losing their local state', () => {
const { rerender } = render(<FeedMessage />)
fireEvent.click(screen.getByRole('button', { name: 'Run count: 0' }))
const queryRenders = renderQuery.mock.calls.length
const functionRenders = renderFunction.mock.calls.length
const updated = structuredClone(initialMessage)
updated.parts[2] = { type: 'text', text: 'Working on the next step' }
rerender(<FeedMessage message={updated} />)
expect(screen.getByText('Working on the next step')).toBeInTheDocument()
expect(screen.getByRole('button', { name: 'Run count: 1' })).toBeInTheDocument()
expect(renderQuery).toHaveBeenCalledTimes(queryRenders)
expect(renderFunction).toHaveBeenCalledTimes(functionRenders)
})
it('updates tool output and code when their content changes', () => {
const { rerender } = render(<FeedMessage />)
const updated = structuredClone(initialMessage)
updated.parts[0] = {
type: 'tool-execute_sql',
toolCallId: 'query-1',
state: 'output-available',
input: { sql: 'select 1', label: 'Query', view: 'table' },
output: [{ value: 1 }],
}
updated.parts[1] = {
type: 'tool-deploy_edge_function',
toolCallId: 'function-1',
state: 'input-available',
input: { code: 'updated code', label: 'Function', functionName: 'hello' },
}
rerender(<FeedMessage message={updated} />)
expect(screen.getByText('[{"value":1}]')).toBeInTheDocument()
expect(screen.getByText('updated code')).toBeInTheDocument()
})
it('uses the latest approval callback even when the tool part is unchanged', () => {
const { rerender } = render(<FeedMessage />)
const approve = vi.fn()
rerender(<FeedMessage addToolApprovalResponse={approve} />)
fireEvent.click(screen.getByRole('button', { name: 'Approve query' }))
expect(approve).toHaveBeenCalledWith({ id: 'approval-1', approved: true })
expect(callbacks.addToolApprovalResponse).not.toHaveBeenCalled()
})
it('retains query state when streaming completes', () => {
const { rerender } = render(<FeedMessage />)
fireEvent.click(screen.getByRole('button', { name: 'Run count: 0' }))
rerender(<FeedMessage isLoading={false} />)
expect(screen.getByRole('button', { name: 'Run count: 1' })).toBeInTheDocument()
})
it('finishes reasoning when the SDK mutates the first streamed part before publishing a snapshot', () => {
const reasoning = {
type: 'reasoning' as const,
text: '',
state: 'streaming' as 'streaming' | 'done',
}
const message: UIMessage = { id: 'reasoning-1', role: 'assistant', parts: [reasoning] }
const { rerender } = render(<FeedMessage message={message} isLastMessage />)
// Expand the tool group so the reasoning row itself is rendered
fireEvent.click(screen.getByRole('button', { name: 'Thinking...' }))
expect(screen.getByText('Thinking...')).toBeInTheDocument()
// Chat.pushMessage exposes the initial object; subsequent replaceMessage calls clone it.
reasoning.state = 'done'
rerender(<FeedMessage message={structuredClone(message)} isLoading={false} isLastMessage />)
expect(screen.queryByText('Thinking...')).not.toBeInTheDocument()
expect(screen.getByText('Reasoned')).toBeInTheDocument()
})
it('updates text when the SDK mutates the first streamed part', () => {
const text = { type: 'text' as const, text: 'First token', state: 'streaming' as const }
const message: UIMessage = { id: 'text-1', role: 'assistant', parts: [text] }
const { rerender } = render(<FeedMessage message={message} />)
text.text = 'First token and the rest of the response'
rerender(<FeedMessage message={structuredClone(message)} />)
expect(screen.getByText(text.text)).toBeInTheDocument()
})
it('updates a tool when its initial input-streaming part is mutated to a completed result', () => {
const tool = {
type: 'tool-execute_sql' as const,
toolCallId: 'query-1',
state: 'input-streaming' as const,
}
const message: UIMessage = { id: 'tool-1', role: 'assistant', parts: [tool] }
const { rerender } = render(<FeedMessage message={message} />)
expect(screen.getByText('Writing SQL...')).toBeInTheDocument()
Object.assign(tool, {
state: 'output-available',
input: { sql: 'select 1' },
output: [{ value: 1 }],
})
rerender(<FeedMessage message={structuredClone(message)} isLoading={false} />)
expect(screen.queryByText('Writing SQL...')).not.toBeInTheDocument()
expect(screen.getByText('[{"value":1}]')).toBeInTheDocument()
})
})
@@ -1,5 +1,5 @@
import { UIMessage as VercelMessage } from '@ai-sdk/react'
import { useState } from 'react'
import { memo, useMemo, useState } from 'react'
import { toast } from 'sonner'
import { cn, copyToClipboard } from 'ui'
@@ -114,38 +114,59 @@ interface MessageProps {
rating?: 'positive' | 'negative' | null
}
export function Message(props: MessageProps) {
export const Message = memo(function Message(props: MessageProps) {
const message = props.message
const { role } = message
const isUserMessage = role === 'user'
let messageState: MessageInfo['state'] = 'idle'
if (props.isBeingEdited) messageState = 'editing'
else if (props.isAfterEditedMessage) messageState = 'predecessor-editing'
const messageInfo = {
id: props.id,
isLoading: props.isLoading,
readOnly: props.readOnly,
variant: props.variant,
isUserMessage,
state: props.isBeingEdited
? 'editing'
: props.isAfterEditedMessage
? 'predecessor-editing'
: 'idle',
isLastMessage: props.isLastMessage,
rating: props.rating,
} satisfies MessageInfo
const messageInfo = useMemo<MessageInfo>(
() => ({
id: props.id,
isLoading: props.isLoading,
readOnly: props.readOnly,
variant: props.variant,
isUserMessage,
state: messageState,
isLastMessage: props.isLastMessage,
rating: props.rating,
}),
[
props.id,
props.isLoading,
props.readOnly,
props.variant,
isUserMessage,
messageState,
props.isLastMessage,
props.rating,
]
)
const messageActions = {
addToolApprovalResponse: props.addToolApprovalResponse,
onDelete: props.onDelete,
onEdit: props.onEdit,
onBranch: props.onBranch,
onCancelEdit: props.onCancelEdit,
onRate: props.onRate,
}
const messageActions = useMemo(
() => ({
addToolApprovalResponse: props.addToolApprovalResponse,
onDelete: props.onDelete,
onEdit: props.onEdit,
onBranch: props.onBranch,
onCancelEdit: props.onCancelEdit,
onRate: props.onRate,
}),
[
props.addToolApprovalResponse,
props.onDelete,
props.onEdit,
props.onBranch,
props.onCancelEdit,
props.onRate,
]
)
return (
<MessageProvider messageInfo={messageInfo} messageActions={messageActions}>
{isUserMessage ? <UserMessage message={message} /> : <AssistantMessage message={message} />}
</MessageProvider>
)
}
})
@@ -128,7 +128,7 @@ const baseMarkdownComponents = {
),
}
export function MessageMarkdown({
export const MessageMarkdown = memo(function MessageMarkdown({
id,
isLoading,
readOnly,
@@ -171,7 +171,7 @@ export function MessageMarkdown({
{markdownSource}
</Streamdown>
)
}
})
export const MarkdownPre = ({
children,
@@ -0,0 +1,42 @@
import { render, screen } from '@testing-library/react'
import { createRef } from 'react'
import type { StickToBottomContext } from 'use-stick-to-bottom'
import { describe, expect, it } from 'vitest'
import { Conversation, ConversationContent } from './Conversation'
describe('ConversationContent', () => {
it('keeps scroll viewport classes separate from content classes and DOM attributes', () => {
const context = createRef<StickToBottomContext>()
const { rerender } = render(
<Conversation contextRef={context}>
<ConversationContent scrollClassName="scroll-pt-4" className="space-y-4" id="messages">
Message
</ConversationContent>
</Conversation>
)
const viewport = context.current?.scrollRef.current
const content = context.current?.contentRef.current
expect(viewport).toHaveClass('scroll-pt-4', 'overscroll-y-contain')
expect(viewport).not.toHaveClass('space-y-4')
expect(content).toHaveClass('space-y-4')
expect(content).not.toHaveClass('scroll-pt-4')
expect(content).toHaveAttribute('id', 'messages')
expect(content).not.toHaveAttribute('scrollClassName')
rerender(
<Conversation contextRef={context}>
<ConversationContent scrollClassName="scroll-pt-8" className="space-y-4" id="messages">
{() => 'Updated message'}
</ConversationContent>
</Conversation>
)
expect(context.current?.scrollRef.current).toBe(viewport)
expect(context.current?.contentRef.current).toBe(content)
expect(viewport).toHaveClass('scroll-pt-8')
expect(viewport).not.toHaveClass('scroll-pt-4')
expect(screen.getByText('Updated message')).toBeInTheDocument()
})
})
@@ -7,7 +7,9 @@ import { StickToBottom, useStickToBottomContext } from 'use-stick-to-bottom'
type ConversationProps = Omit<ComponentProps<typeof StickToBottom>, 'children'> & {
children?: ReactNode
}
type ConversationContentProps = ComponentProps<typeof StickToBottom.Content>
type ConversationContentProps = ComponentProps<typeof StickToBottom.Content> & {
scrollClassName?: string
}
type ConversationScrollButtonProps = ComponentProps<typeof Button>
/**
@@ -20,7 +22,7 @@ const FADE_GUTTER = 'inset-x-7'
export const Conversation = ({ className, children, ...props }: ConversationProps) => (
<StickToBottom
className={cn('relative flex-1 overflow-y-auto', className)}
className={cn('relative min-h-0 flex-1 overflow-hidden', className)}
initial="smooth"
resize="smooth"
role="log"
@@ -44,9 +46,25 @@ export const Conversation = ({ className, children, ...props }: ConversationProp
</StickToBottom>
)
export const ConversationContent = ({ className, ...props }: ConversationContentProps) => (
<StickToBottom.Content className={cn(CONTENT_GUTTER, 'py-4', className)} {...props} />
)
export const ConversationContent = ({
className,
scrollClassName,
children,
...props
}: ConversationContentProps) => {
const context = useStickToBottomContext()
return (
<div
ref={context.scrollRef}
className={cn('h-full w-full overflow-auto overscroll-y-contain', scrollClassName)}
>
<div {...props} ref={context.contentRef} className={cn(CONTENT_GUTTER, 'py-4', className)}>
{typeof children === 'function' ? children(context) : children}
</div>
</div>
)
}
export const ConversationScrollButton = ({
className,
@@ -0,0 +1,26 @@
import { fireEvent, render, screen } from '@testing-library/react'
import { describe, expect, it, vi } from 'vitest'
import { CodeBlock } from './CodeBlock'
describe('CodeBlock', () => {
it('highlights updated code and copies it with the latest callback', () => {
const copy = vi.fn()
const nextCopy = vi.fn()
const { container, rerender } = render(
<CodeBlock value="select 1" language="sql" className="p-2" handleCopy={copy} />
)
expect(container.querySelector('code')?.textContent).toContain('select 1')
expect(screen.getByText('select', { selector: 'span' })).toHaveAttribute('style')
fireEvent.click(screen.getByRole('button', { name: 'Copy' }))
expect(copy).toHaveBeenCalledWith('select 1')
rerender(<CodeBlock value="select 2" language="sql" className="p-2" handleCopy={nextCopy} />)
expect(container.querySelector('code')?.textContent).toContain('select 2')
fireEvent.click(screen.getByRole('button', { name: 'Copied' }))
expect(nextCopy).toHaveBeenCalledWith('select 2')
expect(copy).toHaveBeenCalledTimes(1)
})
})
@@ -5,7 +5,7 @@ import curl from 'highlightjs-curl'
import { noop } from 'lodash'
import { Check, Copy } from 'lucide-react'
import { useTheme } from 'next-themes'
import { Children, ReactNode, useState } from 'react'
import { Children, memo, ReactNode, useState } from 'react'
import { Light as SyntaxHighlighter, SyntaxHighlighterProps } from 'react-syntax-highlighter'
import bash from 'react-syntax-highlighter/dist/cjs/languages/hljs/bash'
import csharp from 'react-syntax-highlighter/dist/cjs/languages/hljs/csharp'
@@ -32,6 +32,27 @@ import { Button, cn, copyToClipboard, FloatingPlate } from 'ui'
import { monokaiCustomTheme } from './CodeBlock.utils'
SyntaxHighlighter.registerLanguage('js', js)
SyntaxHighlighter.registerLanguage('ts', ts)
SyntaxHighlighter.registerLanguage('py', py)
SyntaxHighlighter.registerLanguage('sql', sql)
SyntaxHighlighter.registerLanguage('bash', bash)
SyntaxHighlighter.registerLanguage('dart', dart)
SyntaxHighlighter.registerLanguage('csharp', csharp)
SyntaxHighlighter.registerLanguage('json', json)
SyntaxHighlighter.registerLanguage('kotlin', kotlin)
SyntaxHighlighter.registerLanguage('curl', curl)
SyntaxHighlighter.registerLanguage('http', http)
SyntaxHighlighter.registerLanguage('php', php)
SyntaxHighlighter.registerLanguage('python', python)
SyntaxHighlighter.registerLanguage('go', go)
SyntaxHighlighter.registerLanguage('pgsql', pgsql)
SyntaxHighlighter.registerLanguage('swift', swift)
SyntaxHighlighter.registerLanguage('html', xml)
SyntaxHighlighter.registerLanguage('toml', ini)
SyntaxHighlighter.registerLanguage('yaml', yaml)
SyntaxHighlighter.registerLanguage('markdown', markdown)
const codeBlockLangs = [
'js',
'jsx',
@@ -106,7 +127,7 @@ export interface CodeBlockProps {
* @param {boolean} [props.focusable=true] - Whether the code block is focusable. When true, users can focus the code block to select text or use ⌘A (Cmd+A) to select all. This is so we don't need to load Monaco Editor.
* @param {function} [props.handleCopy] - Optional override behaviour for copying value. For e.g if the code block contains obfuscated values, but the copy behaviour should reveal those values instead.
*/
export const CodeBlock = ({
export const CodeBlock = memo(function CodeBlock({
title,
language,
linesToHighlight = [],
@@ -125,7 +146,7 @@ export const CodeBlock = ({
focusable = true,
onCopyCallback = noop,
handleCopy,
}: CodeBlockProps) => {
}: CodeBlockProps) {
const { resolvedTheme } = useTheme()
const isDarkTheme = resolvedTheme?.includes('dark')!
const monokaiTheme = theme ?? monokaiCustomTheme(isDarkTheme)
@@ -161,26 +182,6 @@ export const CodeBlock = ({
let lang = language ? language : className ? className.replace('language-', '') : 'js'
// force jsx to be js highlighted
if (lang === 'jsx') lang = 'js'
SyntaxHighlighter.registerLanguage('js', js)
SyntaxHighlighter.registerLanguage('ts', ts)
SyntaxHighlighter.registerLanguage('py', py)
SyntaxHighlighter.registerLanguage('sql', sql)
SyntaxHighlighter.registerLanguage('bash', bash)
SyntaxHighlighter.registerLanguage('dart', dart)
SyntaxHighlighter.registerLanguage('csharp', csharp)
SyntaxHighlighter.registerLanguage('json', json)
SyntaxHighlighter.registerLanguage('kotlin', kotlin)
SyntaxHighlighter.registerLanguage('curl', curl)
SyntaxHighlighter.registerLanguage('http', http)
SyntaxHighlighter.registerLanguage('php', php)
SyntaxHighlighter.registerLanguage('python', python)
SyntaxHighlighter.registerLanguage('go', go)
SyntaxHighlighter.registerLanguage('pgsql', pgsql)
SyntaxHighlighter.registerLanguage('swift', swift)
SyntaxHighlighter.registerLanguage('html', xml)
SyntaxHighlighter.registerLanguage('toml', ini)
SyntaxHighlighter.registerLanguage('yaml', yaml)
SyntaxHighlighter.registerLanguage('markdown', markdown)
const large = false
// don't show line numbers if bash == lang
@@ -289,4 +290,4 @@ export const CodeBlock = ({
)}
</>
)
}
})