Skip to content

Commit d84a49e

Browse files
authored
fix(ai): send a server tool call once on the next turn with withPersistence (#1605)
* fix(ai): keep each model call's message id in StreamProcessor When a model call starts with a server tool call and no text, the processor made a placeholder message and let the next call's TEXT_MESSAGE_START rename it. One UIMessage then held both calls under the second call's id, so withPersistence replaced the stored answer with a copy of the tool call and the next turn sent it twice. A message made with a server id is no longer a placeholder, and a placeholder made by thinking takes the tool call's parentMessageId. * fix(ai): put a tool call in its parent message and keep queued structured-output updates on rename
1 parent 8bd5f07 commit d84a49e

8 files changed

Lines changed: 691 additions & 35 deletions

File tree

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
---
2+
'@tanstack/ai': patch
3+
---
4+
5+
Fix a server tool call that the next turn sent twice when you use `withPersistence`. When a model call starts with a tool call and no text (often after thinking), `StreamProcessor` gave that message the id of the next model call. The next turn then replaced the stored answer with a copy that also held the tool call, so the provider got the same `tool_use` twice. The client now keeps the id of the message the tool call belongs to.
Lines changed: 199 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,199 @@
1+
import { describe, expect, it } from 'vitest'
2+
import {
3+
EventType,
4+
StreamProcessor,
5+
chat,
6+
chatParamsFromRequestBody,
7+
uiMessagesToWire,
8+
} from '@tanstack/ai'
9+
import type {
10+
AdapterYieldChunk,
11+
AnyTextAdapter,
12+
ModelMessage,
13+
StreamChunk,
14+
UIMessage,
15+
} from '@tanstack/ai'
16+
import { memoryPersistence } from '../src/memory'
17+
import { withPersistence } from '../src/middleware'
18+
import { threadMessages } from './persistence-fixtures'
19+
20+
function mockAdapter(iterations: Array<Array<AdapterYieldChunk>>) {
21+
const prompts: Array<Array<ModelMessage>> = []
22+
let i = 0
23+
const adapter = {
24+
kind: 'text',
25+
name: 'mock',
26+
model: 'test-model',
27+
'~types': {},
28+
chatStream: (opts: { messages: Array<ModelMessage> }) => {
29+
prompts.push(structuredClone(opts.messages))
30+
const chunks = iterations[i] ?? []
31+
i++
32+
return (async function* () {
33+
for (const c of chunks) yield c
34+
})()
35+
},
36+
structuredOutput: async () => ({ data: {}, rawText: '{}' }),
37+
} as unknown as AnyTextAdapter
38+
return { adapter, prompts }
39+
}
40+
41+
const t = 1
42+
const run = (
43+
runId: string,
44+
chunks: Array<AdapterYieldChunk>,
45+
finishReason: string,
46+
) =>
47+
[
48+
{ type: EventType.RUN_STARTED, runId, threadId: 't1', timestamp: t },
49+
...chunks,
50+
{
51+
type: EventType.RUN_FINISHED,
52+
runId,
53+
threadId: 't1',
54+
finishReason,
55+
timestamp: t,
56+
},
57+
] as Array<AdapterYieldChunk>
58+
59+
const thinking = (id: string): Array<AdapterYieldChunk> => [
60+
{ type: EventType.REASONING_START, messageId: id, timestamp: t },
61+
{
62+
type: EventType.REASONING_MESSAGE_START,
63+
messageId: id,
64+
role: 'reasoning',
65+
timestamp: t,
66+
},
67+
{
68+
type: EventType.REASONING_MESSAGE_CONTENT,
69+
messageId: id,
70+
delta: 'plan',
71+
timestamp: t,
72+
},
73+
{ type: EventType.REASONING_MESSAGE_END, messageId: id, timestamp: t },
74+
{ type: EventType.REASONING_END, messageId: id, timestamp: t },
75+
]
76+
77+
const text = (messageId: string, delta: string): Array<AdapterYieldChunk> => [
78+
{
79+
type: EventType.TEXT_MESSAGE_START,
80+
messageId,
81+
role: 'assistant',
82+
timestamp: t,
83+
},
84+
{ type: EventType.TEXT_MESSAGE_CONTENT, messageId, delta, timestamp: t },
85+
{ type: EventType.TEXT_MESSAGE_END, messageId, timestamp: t },
86+
]
87+
88+
// Turn 1: call 1 = (thinking) + server tool call with no text, call 2 =
89+
// (thinking) + text. Turn 2: a plain answer.
90+
function model(withThinking: boolean) {
91+
return [
92+
run(
93+
'r1',
94+
[
95+
...(withThinking ? thinking('think-1') : []),
96+
{
97+
type: EventType.TOOL_CALL_START,
98+
toolCallId: 'call_1',
99+
toolCallName: 'lookup',
100+
parentMessageId: 'm1',
101+
timestamp: t,
102+
},
103+
{
104+
type: EventType.TOOL_CALL_ARGS,
105+
toolCallId: 'call_1',
106+
delta: '{}',
107+
timestamp: t,
108+
},
109+
{ type: EventType.TOOL_CALL_END, toolCallId: 'call_1', timestamp: t },
110+
],
111+
'tool_calls',
112+
),
113+
run(
114+
'r1',
115+
[...(withThinking ? thinking('think-2') : []), ...text('m2', 'done')],
116+
'stop',
117+
),
118+
run('r2', text('m3', 'ok'), 'stop'),
119+
]
120+
}
121+
122+
const call1Count = (messages: ReadonlyArray<ModelMessage>) =>
123+
messages.flatMap((m) => m.toolCalls ?? []).filter((c) => c.id === 'call_1')
124+
.length
125+
126+
/** What `/api/chat` gets: the client's wire messages after the HTTP hop. */
127+
async function serverMessages(ui: Array<UIMessage>, runId: string) {
128+
const params = await chatParamsFromRequestBody({
129+
threadId: 't1',
130+
runId,
131+
messages: JSON.parse(JSON.stringify(uiMessagesToWire(ui))),
132+
tools: [],
133+
context: [],
134+
})
135+
return params.messages
136+
}
137+
138+
describe('useChat + withPersistence: message ids across model calls', () => {
139+
for (const withThinking of [false, true]) {
140+
it(`sends and stores a server tool call once (thinking first: ${withThinking})`, async () => {
141+
const persistence = memoryPersistence()
142+
const { adapter, prompts } = mockAdapter(model(withThinking))
143+
const tools = [
144+
{
145+
name: 'lookup',
146+
description: 'Lookup',
147+
execute: () => ({ ok: true }),
148+
},
149+
]
150+
const middleware = [withPersistence(persistence)]
151+
152+
// Client side, like ChatClient: one processor, every chunk in order.
153+
const processor = new StreamProcessor()
154+
const user1: UIMessage = {
155+
id: 'u1',
156+
role: 'user',
157+
parts: [{ type: 'text', content: 'hi' }],
158+
}
159+
processor.setMessages([user1])
160+
processor.prepareAssistantMessage()
161+
for await (const chunk of chat({
162+
adapter,
163+
messages: await serverMessages([user1], 'r1'),
164+
tools,
165+
runId: 'r1',
166+
threadId: 't1',
167+
middleware,
168+
}) as AsyncIterable<StreamChunk>) {
169+
processor.processChunk(chunk)
170+
}
171+
processor.finalizeStream()
172+
173+
const user2: UIMessage = {
174+
id: 'u2',
175+
role: 'user',
176+
parts: [{ type: 'text', content: 'next' }],
177+
}
178+
for await (const _ of chat({
179+
adapter,
180+
messages: await serverMessages(
181+
[...processor.getMessages(), user2],
182+
'r2',
183+
),
184+
tools,
185+
runId: 'r2',
186+
threadId: 't1',
187+
middleware,
188+
}) as AsyncIterable<StreamChunk>) {
189+
// drain
190+
}
191+
192+
expect(call1Count(prompts[2]!)).toBe(1)
193+
const stored = threadMessages(
194+
await persistence.stores.messages!.loadThread('t1'),
195+
)
196+
expect(call1Count(stored)).toBe(1)
197+
})
198+
}
199+
})

‎packages/ai/docs/chat-architecture.md‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -552,6 +552,8 @@ Content-bearing chunks that trigger `ensureAssistantMessage()`:
552552
- `STEP_FINISHED`
553553
- `RUN_ERROR`
554554

555+
The message takes the server id on the chunk when it has one (`TOOL_CALL_START.parentMessageId`, `TEXT_MESSAGE_CONTENT.messageId`). Thinking chunks carry no assistant id, so they make a placeholder with a client id. The placeholder takes the first server id that arrives in `TOOL_CALL_START.parentMessageId` or `TEXT_MESSAGE_START.messageId`, and keeps it. A later model call's `TEXT_MESSAGE_START` with a new id then starts a new UIMessage. It does not rename this one, so each UIMessage keeps the id the server stored for it. If a message with that `parentMessageId` already exists, the placeholder joins it, so the thinking stays with its tool call. A tool call always goes to the message its `parentMessageId` names, not to the active message.
556+
555557
---
556558

557559
## UIMessage Part Ordering Invariants

‎packages/ai/src/activities/chat/stream/processor.ts‎

Lines changed: 80 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -1147,7 +1147,12 @@ export class StreamProcessor {
11471147
* MESSAGES_SNAPSHOT) but whose transient state was cleared. In that case we
11481148
* hydrate state from the existing message rather than creating a duplicate.
11491149
*/
1150-
private ensureAssistantMessage(preferredId?: string): {
1150+
private ensureAssistantMessage(
1151+
preferredId?: string,
1152+
// False when `preferredId` names its own message (a tool call's
1153+
// `parentMessageId`), so a different active message never takes it.
1154+
fallbackToActive = true,
1155+
): {
11511156
messageId: string
11521157
state: MessageStreamState
11531158
} {
@@ -1161,7 +1166,9 @@ export class StreamProcessor {
11611166
}
11621167

11631168
// Try active assistant message
1164-
const activeId = this.getActiveAssistantMessageId()
1169+
const activeId = fallbackToActive
1170+
? this.getActiveAssistantMessageId()
1171+
: null
11651172
if (activeId) {
11661173
const state = this.getMessageState(activeId)
11671174
if (state) {
@@ -1209,7 +1216,8 @@ export class StreamProcessor {
12091216
this.messages = [...this.messages, assistantMessage]
12101217
const state = this.createMessageState(id, 'assistant')
12111218
this.activeMessageIds.add(id)
1212-
this.pendingManualMessageId = id
1219+
// Only a client id is a placeholder. A server id is already final.
1220+
if (!preferredId) this.pendingManualMessageId = id
12131221
this.events.onStreamStart?.()
12141222
this.emitMessagesChange()
12151223
return { messageId: id, state }
@@ -1509,6 +1517,47 @@ export class StreamProcessor {
15091517
return true
15101518
}
15111519

1520+
private renameMessage(fromId: string, toId: string): void {
1521+
// Update the message's ID in the messages array
1522+
this.messages = this.messages.map((msg) =>
1523+
msg.id === fromId ? { ...msg, id: toId } : msg,
1524+
)
1525+
1526+
// Move state to the new key
1527+
const existingState = this.messageStates.get(fromId)
1528+
if (existingState) {
1529+
existingState.id = toId
1530+
this.messageStates.delete(fromId)
1531+
this.messageStates.set(toId, existingState)
1532+
}
1533+
1534+
// Update activeMessageIds
1535+
this.activeMessageIds.delete(fromId)
1536+
this.activeMessageIds.add(toId)
1537+
1538+
// TOOL_CALL_ARGS/END route through toolCallToMessage. Keep those
1539+
// entries on the remapped id so later args still accumulate
1540+
// (interleaved text can arrive as a full START/CONTENT/END block).
1541+
for (const [toolCallId, mappedMessageId] of this.toolCallToMessage) {
1542+
if (mappedMessageId === fromId) {
1543+
this.toolCallToMessage.set(toolCallId, toId)
1544+
}
1545+
}
1546+
1547+
// `structured-output.start` can arrive before this event and mark
1548+
// the pending id. Keep the mark on the remapped id so the JSON deltas
1549+
// go into the structured-output part, not a text part.
1550+
if (this.structuredMessageIds.delete(fromId)) {
1551+
this.structuredMessageIds.add(toId)
1552+
}
1553+
// Queued deltas move too, so the next flush of `toId` sends them.
1554+
const batch = this.structuredOutputUpdateBatches.get(fromId)
1555+
if (batch) {
1556+
this.structuredOutputUpdateBatches.delete(fromId)
1557+
this.structuredOutputUpdateBatches.set(toId, batch)
1558+
}
1559+
}
1560+
15121561
/**
15131562
* Handle TEXT_MESSAGE_START event
15141563
*/
@@ -1537,38 +1586,7 @@ export class StreamProcessor {
15371586
this.pendingManualMessageId = null
15381587

15391588
if (pendingId !== messageId) {
1540-
// Update the message's ID in the messages array
1541-
this.messages = this.messages.map((msg) =>
1542-
msg.id === pendingId ? { ...msg, id: messageId } : msg,
1543-
)
1544-
1545-
// Move state to the new key
1546-
const existingState = this.messageStates.get(pendingId)
1547-
if (existingState) {
1548-
existingState.id = messageId
1549-
this.messageStates.delete(pendingId)
1550-
this.messageStates.set(messageId, existingState)
1551-
}
1552-
1553-
// Update activeMessageIds
1554-
this.activeMessageIds.delete(pendingId)
1555-
this.activeMessageIds.add(messageId)
1556-
1557-
// TOOL_CALL_ARGS/END route through toolCallToMessage. Keep those
1558-
// entries on the remapped id so later args still accumulate
1559-
// (interleaved text can arrive as a full START/CONTENT/END block).
1560-
for (const [toolCallId, mappedMessageId] of this.toolCallToMessage) {
1561-
if (mappedMessageId === pendingId) {
1562-
this.toolCallToMessage.set(toolCallId, messageId)
1563-
}
1564-
}
1565-
1566-
// `structured-output.start` can arrive before this event and mark
1567-
// the pending id. Keep the mark on the remapped id so the JSON deltas
1568-
// go into the structured-output part, not a text part.
1569-
if (this.structuredMessageIds.delete(pendingId)) {
1570-
this.structuredMessageIds.add(messageId)
1571-
}
1589+
this.renameMessage(pendingId, messageId)
15721590
}
15731591

15741592
// Ensure state exists
@@ -2205,11 +2223,38 @@ export class StreamProcessor {
22052223
private handleToolCallStartEvent(
22062224
chunk: Extract<StreamChunk, { type: 'TOOL_CALL_START' }>,
22072225
): void {
2226+
// A placeholder (e.g. made by thinking) takes the tool call's parent id.
2227+
// Otherwise the next model call's TEXT_MESSAGE_START renames it, and one
2228+
// message holds two calls under the wrong id.
2229+
const pendingId = this.pendingManualMessageId
2230+
const parentId = chunk.parentMessageId
2231+
if (pendingId && parentId && parentId !== pendingId) {
2232+
const parent = this.messages.find((m) => m.id === parentId)
2233+
const placeholder = this.messages.find((m) => m.id === pendingId)
2234+
if (!parent) {
2235+
this.pendingManualMessageId = null
2236+
this.renameMessage(pendingId, parentId)
2237+
} else if (parent.role === 'assistant' && placeholder) {
2238+
// The parent already exists. Fold the placeholder into it, so the
2239+
// thinking stays in the same message as the tool call.
2240+
this.pendingManualMessageId = null
2241+
this.messages = this.messages
2242+
.filter((m) => m.id !== pendingId)
2243+
.map((m) =>
2244+
m.id === parentId
2245+
? { ...m, parts: [...m.parts, ...placeholder.parts] }
2246+
: m,
2247+
)
2248+
this.renameMessage(pendingId, parentId)
2249+
}
2250+
}
2251+
22082252
// Determine the message this tool call belongs to
22092253
const targetMessageId =
22102254
chunk.parentMessageId ?? this.getActiveAssistantMessageId()
22112255
const { messageId, state } = this.ensureAssistantMessage(
22122256
targetMessageId ?? undefined,
2257+
chunk.parentMessageId === undefined,
22132258
)
22142259

22152260
// Mark that we've seen tool calls since the last text segment

0 commit comments

Comments
 (0)