@assistant-ui/ai-sdk 0.0.8 → 0.0.10

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (51) hide show
  1. package/LICENSE +1 -1
  2. package/dist/converters/convertMessage.d.ts +2 -0
  3. package/dist/converters/convertMessage.d.ts.map +1 -1
  4. package/dist/converters/convertMessage.js +15 -4
  5. package/dist/converters/convertMessage.js.map +1 -1
  6. package/dist/converters/toCreateMessage.d.ts.map +1 -1
  7. package/dist/converters/toCreateMessage.js +2 -1
  8. package/dist/converters/toCreateMessage.js.map +1 -1
  9. package/dist/runtime/AISDKChat.d.ts.map +1 -1
  10. package/dist/runtime/AISDKChat.js.map +1 -1
  11. package/dist/runtime/AISDKThreads.js.map +1 -1
  12. package/dist/runtime/sdkIdentity.js +1 -1
  13. package/dist/runtime/useAISDKRuntime.d.ts +8 -1
  14. package/dist/runtime/useAISDKRuntime.d.ts.map +1 -1
  15. package/dist/runtime/useAISDKRuntime.js +212 -34
  16. package/dist/runtime/useAISDKRuntime.js.map +1 -1
  17. package/dist/runtime/useChatRuntime.js +1 -1
  18. package/dist/runtime/useChatRuntime.js.map +1 -1
  19. package/dist/runtime/useChatThread.d.ts.map +1 -1
  20. package/dist/runtime/useChatThread.js +40 -10
  21. package/dist/runtime/useChatThread.js.map +1 -1
  22. package/dist/runtime/useExternalHistory.d.ts.map +1 -1
  23. package/dist/runtime/useExternalHistory.js +2 -1
  24. package/dist/runtime/useExternalHistory.js.map +1 -1
  25. package/dist/runtime/useResourceCleanup.js.map +1 -1
  26. package/dist/runtime/useStreamingTiming.js.map +1 -1
  27. package/dist/transport/AssistantChatTransport.d.ts.map +1 -1
  28. package/dist/transport/AssistantChatTransport.js +9 -2
  29. package/dist/transport/AssistantChatTransport.js.map +1 -1
  30. package/dist/transport/resumable.js.map +1 -1
  31. package/dist/usage.js.map +1 -1
  32. package/package.json +15 -12
  33. package/src/converters/convertMessage.test.ts +99 -0
  34. package/src/converters/convertMessage.ts +27 -2
  35. package/src/converters/toCreateMessage.test.ts +13 -0
  36. package/src/converters/toCreateMessage.ts +1 -0
  37. package/src/runtime/AISDKChat.ts +0 -4
  38. package/src/runtime/AISDKThreads.test.ts +26 -0
  39. package/src/runtime/useAISDKRuntime.approval.integration.test.tsx +1588 -19
  40. package/src/runtime/useAISDKRuntime.approval.test.tsx +27 -0
  41. package/src/runtime/useAISDKRuntime.fast-refresh.test.tsx +182 -0
  42. package/src/runtime/useAISDKRuntime.ts +397 -28
  43. package/src/runtime/useChatRuntime.fast-refresh.test.tsx +94 -0
  44. package/src/runtime/useChatRuntime.integration.test.tsx +134 -112
  45. package/src/runtime/useChatRuntime.test.ts +7 -7
  46. package/src/runtime/useChatThread.test.ts +166 -2
  47. package/src/runtime/useChatThread.transport.test.tsx +5 -2
  48. package/src/runtime/useChatThread.ts +56 -17
  49. package/src/runtime/useExternalHistory.ts +12 -1
  50. package/src/transport/AssistantChatTransport.test.ts +164 -0
  51. package/src/transport/AssistantChatTransport.ts +22 -2
@@ -181,6 +181,8 @@ type ChatCallbacks<UI_MESSAGE extends UIMessage> = Pick<
181
181
  "onToolCall" | "onData" | "onFinish" | "onError" | "sendAutomaticallyWhen"
182
182
  >;
183
183
 
184
+ const requestsByChat = new WeakMap<object, symbol>();
185
+
184
186
  /**
185
187
  * Constructs a `Chat` whose callbacks read the latest options through
186
188
  * `callbacksRef`, the forwarding `useChat` applies only to a chat it
@@ -189,9 +191,22 @@ type ChatCallbacks<UI_MESSAGE extends UIMessage> = Pick<
189
191
  export const createChat = <UI_MESSAGE extends UIMessage>(
190
192
  init: ChatInit<UI_MESSAGE>,
191
193
  callbacksRef: { readonly current: ChatCallbacks<UI_MESSAGE> | undefined },
192
- ): Chat<UI_MESSAGE> =>
193
- new Chat<UI_MESSAGE>({
194
+ ): Chat<UI_MESSAGE> => {
195
+ const transport = init.transport;
196
+ const chat = new Chat<UI_MESSAGE>({
194
197
  ...init,
198
+ ...(transport && {
199
+ transport: {
200
+ sendMessages: (options) => {
201
+ requestsByChat.set(chat, Symbol());
202
+ return transport.sendMessages(options);
203
+ },
204
+ reconnectToStream: (options) => {
205
+ requestsByChat.set(chat, Symbol());
206
+ return transport.reconnectToStream(options);
207
+ },
208
+ },
209
+ }),
195
210
  onToolCall: (arg) => callbacksRef.current?.onToolCall?.(arg),
196
211
  onData: (arg) => callbacksRef.current?.onData?.(arg),
197
212
  onFinish: (arg) => callbacksRef.current?.onFinish?.(arg),
@@ -199,6 +214,8 @@ export const createChat = <UI_MESSAGE extends UIMessage>(
199
214
  sendAutomaticallyWhen: (arg) =>
200
215
  callbacksRef.current?.sendAutomaticallyWhen?.(arg) ?? false,
201
216
  });
217
+ return chat;
218
+ };
202
219
 
203
220
  export const useChatThread = <UI_MESSAGE extends UIMessage = UIMessage>(
204
221
  options: ChatThreadOptions<UI_MESSAGE> | undefined,
@@ -271,7 +288,10 @@ export const useChatThread = <UI_MESSAGE extends UIMessage = UIMessage>(
271
288
  );
272
289
 
273
290
  const runtime = useAISDKRuntime(chat, {
274
- adapters,
291
+ adapters: {
292
+ ...adapters,
293
+ threadList: { threadId: id, ...adapters?.threadList },
294
+ },
275
295
  ...pickExternalStoreSharedOptions(options ?? {}),
276
296
  ...(toCreateMessage && { toCreateMessage }),
277
297
  ...(onResume && { onResume }),
@@ -282,6 +302,11 @@ export const useChatThread = <UI_MESSAGE extends UIMessage = UIMessage>(
282
302
  ...(messageRepositoryInstance && {
283
303
  unstable_messageRepositoryInstance: messageRepositoryInstance,
284
304
  }),
305
+ // The chat outlives this runtime when a host mounts only the visible
306
+ // thread, so a host approval answer is kept with it. This is the Chat
307
+ // instance, not the useChat helpers, which are re-minted every render and
308
+ // would be a dead WeakMap key by the next one.
309
+ unstable_hostApprovalOwner: externalChat ?? ownedChat,
285
310
  ...(unstable_onBranchChange && { unstable_onBranchChange }),
286
311
  });
287
312
 
@@ -350,27 +375,41 @@ export const useChatThread = <UI_MESSAGE extends UIMessage = UIMessage>(
350
375
  }
351
376
  if (isLoadingHistory) return;
352
377
  resumedStreamIds.add(pendingStreamId);
353
- chat.resumeStream().catch((err: unknown) => {
354
- console.warn("[assistant-ui] resumable: resume failed", err);
355
- try {
356
- onResumeErrorRef.current?.(err);
357
- } catch (callbackError) {
358
- console.error(
359
- "[assistant-ui] resumable: onResumeError callback failed",
360
- callbackError,
361
- );
362
- } finally {
363
- if (resumableStorage?.getStreamId(id) === pendingStreamId) {
364
- resumableStorage.clear(id);
378
+ const activeChat = externalChat ?? ownedChat;
379
+ activeChat.clearError();
380
+ const pending = chat.resumeStream();
381
+ const request = requestsByChat.get(activeChat);
382
+ pending
383
+ .then(() => {
384
+ // Chat.error is shared with sends and resumes that can start before
385
+ // this promise settles, including inside the caller's onFinish.
386
+ if (requestsByChat.get(activeChat) === request && activeChat.error) {
387
+ throw activeChat.error;
388
+ }
389
+ })
390
+ .catch((err: unknown) => {
391
+ console.warn("[assistant-ui] resumable: resume failed", err);
392
+ try {
393
+ onResumeErrorRef.current?.(err);
394
+ } catch (callbackError) {
395
+ console.error(
396
+ "[assistant-ui] resumable: onResumeError callback failed",
397
+ callbackError,
398
+ );
399
+ } finally {
400
+ if (resumableStorage?.getStreamId(id) === pendingStreamId) {
401
+ resumableStorage.clear(id);
402
+ }
365
403
  }
366
- }
367
- });
404
+ });
368
405
  }, [
369
406
  chat,
407
+ externalChat,
370
408
  id,
371
409
  isChatRunning,
372
410
  isLoadingHistory,
373
411
  pendingStreamId,
412
+ ownedChat,
374
413
  resumableStorage,
375
414
  resumedStreamIds,
376
415
  ]);
@@ -51,9 +51,20 @@ export const toExportedMessageRepository = <TMessage>(
51
51
  };
52
52
  };
53
53
 
54
+ const hasUnansweredApproval = (message: ThreadMessage) =>
55
+ message.content.some(
56
+ (part) =>
57
+ part.type === "tool-call" &&
58
+ part.approval != null &&
59
+ part.approval.approved === undefined &&
60
+ part.approval.resolution === undefined,
61
+ );
62
+
63
+ // Core reports a run paused on an unanswered approval as an interrupt, not as tool calls.
54
64
  const isAwaitingToolApproval = (message: ThreadMessage) =>
55
65
  message.status?.type === "requires-action" &&
56
- message.status.reason === "tool-calls";
66
+ (message.status.reason === "tool-calls" ||
67
+ (message.status.reason === "interrupt" && hasUnansweredApproval(message)));
57
68
 
58
69
  const isTerminalMessage = (message: ThreadMessage) =>
59
70
  message.status === undefined ||
@@ -205,6 +205,170 @@ const wrappedFetchOf = (
205
205
  ).fetch;
206
206
 
207
207
  describe("AssistantChatTransport resumable fetch wrapper", () => {
208
+ it("keeps a replacement checkpoint set while preparing an older reconnect", async () => {
209
+ const storage = createMemoryStorage("stream-old");
210
+ let finishPrepare!: () => void;
211
+ const prepareReconnectToStreamRequest = vi.fn(async () => {
212
+ await new Promise<void>((resolve) => {
213
+ finishPrepare = resolve;
214
+ });
215
+ return { headers: { "x-custom": "retained" } };
216
+ });
217
+ const fetch = vi
218
+ .fn<typeof globalThis.fetch>()
219
+ .mockResolvedValue(new Response(null, { status: 204 }));
220
+ const transport = new AssistantChatTransport({
221
+ fetch,
222
+ prepareReconnectToStreamRequest,
223
+ resumable: { storage, resumeApi: (id) => `/api/resume/${id}` },
224
+ });
225
+ const pending = transport.reconnectToStream({ chatId: "thread" });
226
+ await vi.waitFor(() =>
227
+ expect(prepareReconnectToStreamRequest).toHaveBeenCalledOnce(),
228
+ );
229
+ storage.setStreamId("stream-new");
230
+ finishPrepare();
231
+ await expect(pending).resolves.toBeNull();
232
+ expect(fetch.mock.calls[0]?.[0]).toBe("/api/resume/stream-old");
233
+ expect(
234
+ Array.from(new Headers(fetch.mock.calls[0]?.[1]?.headers).entries()),
235
+ ).toEqual([["x-custom", "retained"]]);
236
+ expect(storage.getStreamId()).toBe("stream-new");
237
+ });
238
+
239
+ it.each([
240
+ { status: 204, replaceCheckpoint: false },
241
+ { status: 204, replaceCheckpoint: true },
242
+ { status: 404, replaceCheckpoint: false },
243
+ { status: 404, replaceCheckpoint: true },
244
+ ])(
245
+ "clears only the matching checkpoint after $status (replacement: $replaceCheckpoint)",
246
+ async ({ status, replaceCheckpoint }) => {
247
+ const storage = createMemoryStorage("stream-old");
248
+ let respond!: (response: Response) => void;
249
+ const fetch = vi.fn(
250
+ () =>
251
+ new Promise<Response>((resolve) => {
252
+ respond = resolve;
253
+ }),
254
+ );
255
+ const transport = new AssistantChatTransport({
256
+ fetch,
257
+ resumable: { storage, resumeApi: "/api/resume" },
258
+ });
259
+ const pending = transport.reconnectToStream({ chatId: "thread" });
260
+ await vi.waitFor(() => expect(fetch).toHaveBeenCalledOnce());
261
+ if (replaceCheckpoint) storage.setStreamId("stream-new");
262
+ respond(
263
+ new Response(status === 404 ? "stream expired" : null, { status }),
264
+ );
265
+ if (status === 404) {
266
+ await expect(pending).rejects.toThrow("stream expired");
267
+ } else {
268
+ await expect(pending).resolves.toBeNull();
269
+ }
270
+ expect(storage.getStreamId()).toBe(
271
+ replaceCheckpoint ? "stream-new" : null,
272
+ );
273
+ },
274
+ );
275
+
276
+ it.each([undefined, "stream-old", "response-id"])(
277
+ "preserves a replacement checkpoint on a delayed successful reconnect (%s)",
278
+ async (responseId) => {
279
+ const storage = createMemoryStorage("stream-old");
280
+ let respond!: (response: Response) => void;
281
+ const fetch = vi.fn(
282
+ () =>
283
+ new Promise<Response>((resolve) => {
284
+ respond = resolve;
285
+ }),
286
+ );
287
+ const transport = new AssistantChatTransport({
288
+ fetch,
289
+ resumable: { storage, resumeApi: "/api/resume" },
290
+ });
291
+ const pending = transport.reconnectToStream({ chatId: "thread" });
292
+ await vi.waitFor(() => expect(fetch).toHaveBeenCalledOnce());
293
+ storage.setStreamId("stream-new");
294
+ respond(
295
+ new Response('data: {"type":"finish"}\n\n', {
296
+ headers: {
297
+ "content-type": "text/event-stream",
298
+ ...(responseId && { [RESUMABLE_STREAM_ID_HEADER]: responseId }),
299
+ },
300
+ }),
301
+ );
302
+ const stream = await pending;
303
+ expect(storage.getStreamId()).toBe("stream-new");
304
+ const reader = stream!.getReader();
305
+ while (!(await reader.read()).done) {}
306
+ expect(storage.getStreamId()).toBe("stream-new");
307
+ },
308
+ );
309
+
310
+ it.each([undefined, "stream-old", "response-id"])(
311
+ "preserves a checkpoint replaced while consuming a reconnect (%s)",
312
+ async (responseId) => {
313
+ const storage = createMemoryStorage("stream-old");
314
+ let controller!: ReadableStreamDefaultController<Uint8Array>;
315
+ const transport = new AssistantChatTransport({
316
+ fetch: vi.fn(
317
+ async () =>
318
+ new Response(
319
+ new ReadableStream({
320
+ start(value) {
321
+ controller = value;
322
+ },
323
+ }),
324
+ {
325
+ headers: {
326
+ "content-type": "text/event-stream",
327
+ ...(responseId && {
328
+ [RESUMABLE_STREAM_ID_HEADER]: responseId,
329
+ }),
330
+ },
331
+ },
332
+ ),
333
+ ),
334
+ resumable: { storage, resumeApi: "/api/resume" },
335
+ });
336
+ const stream = await transport.reconnectToStream({ chatId: "thread" });
337
+ expect(storage.getStreamId()).toBe(responseId ?? "stream-old");
338
+ storage.setStreamId("stream-new");
339
+ controller.enqueue(
340
+ new TextEncoder().encode('data: {"type":"finish"}\n\n'),
341
+ );
342
+ controller.close();
343
+ const reader = stream!.getReader();
344
+ while (!(await reader.read()).done) {}
345
+ expect(storage.getStreamId()).toBe("stream-new");
346
+ },
347
+ );
348
+
349
+ it.each([undefined, "stream-old", "response-id"])(
350
+ "clears the checkpoint owned by a completed reconnect (%s)",
351
+ async (responseId) => {
352
+ const storage = createMemoryStorage("stream-old");
353
+ const transport = new AssistantChatTransport({
354
+ fetch: vi.fn(
355
+ async () =>
356
+ new Response('data: {"type":"finish"}\n\n', {
357
+ headers: {
358
+ "content-type": "text/event-stream",
359
+ ...(responseId && { [RESUMABLE_STREAM_ID_HEADER]: responseId }),
360
+ },
361
+ }),
362
+ ),
363
+ resumable: { storage, resumeApi: "/api/resume" },
364
+ });
365
+ const stream = await transport.reconnectToStream({ chatId: "thread" });
366
+ const reader = stream!.getReader();
367
+ while (!(await reader.read()).done) {}
368
+ expect(storage.getStreamId()).toBeNull();
369
+ },
370
+ );
371
+
208
372
  it("passes a 204 with a non-null empty body through untouched (WebKit)", async () => {
209
373
  const response = nullBodyStatusWithBody(204);
210
374
  const fetchMock = vi.fn(async () => response);
@@ -22,6 +22,7 @@ const FINISH_MARKER = '"type":"finish"';
22
22
  const FINISH_BUFFER_LIMIT = 4096;
23
23
  const FINISH_BUFFER_TAIL = 1024;
24
24
  const RESUMABLE_THREAD_ID_HEADER = "x-assistant-ui-resumable-thread-id";
25
+ const RESUMABLE_RECONNECT_ID_HEADER = "x-assistant-ui-resumable-reconnect-id";
25
26
 
26
27
  // 101/204/205/304 are null-body statuses per the fetch spec: `new Response(body, { status })`
27
28
  // throws for them, and WebKit returns a non-null empty body, so the body check alone does not guard it.
@@ -137,10 +138,24 @@ function wrapFetchWithResumable(
137
138
  return async (input, init) => {
138
139
  const headers = new Headers(init?.headers);
139
140
  const threadId = headers.get(RESUMABLE_THREAD_ID_HEADER) ?? undefined;
141
+ const reconnectingStreamId = headers.get(RESUMABLE_RECONNECT_ID_HEADER);
142
+ const checkpointAtRequest =
143
+ reconnectingStreamId ?? resumable.storage.getStreamId(threadId);
140
144
  headers.delete(RESUMABLE_THREAD_ID_HEADER);
145
+ headers.delete(RESUMABLE_RECONNECT_ID_HEADER);
141
146
  const res = await baseFetch(input, { ...init, headers });
142
147
  const id = res.headers.get(RESUMABLE_STREAM_ID_HEADER);
143
- if (id) resumable.storage.setStreamId(id, threadId);
148
+ const ownsCheckpoint =
149
+ !reconnectingStreamId ||
150
+ resumable.storage.getStreamId(threadId) === reconnectingStreamId;
151
+ if (id && ownsCheckpoint) resumable.storage.setStreamId(id, threadId);
152
+ if (
153
+ (res.status === 204 || res.status === 404) &&
154
+ reconnectingStreamId &&
155
+ resumable.storage.getStreamId(threadId) === reconnectingStreamId
156
+ ) {
157
+ resumable.storage.clear(threadId);
158
+ }
144
159
  if (!res.body || NULL_BODY_STATUSES.has(res.status)) return res;
145
160
 
146
161
  const detectFinish = resumable.isFinishEvent ?? defaultIsFinishEvent;
@@ -153,7 +168,11 @@ function wrapFetchWithResumable(
153
168
  controller.enqueue(chunk);
154
169
  accumulator += decoder.decode(chunk, { stream: true });
155
170
  if (detectFinish(chunk, accumulator)) {
156
- if (!id || resumable.storage.getStreamId(threadId) === id) {
171
+ if (
172
+ ownsCheckpoint &&
173
+ resumable.storage.getStreamId(threadId) ===
174
+ (id ?? checkpointAtRequest)
175
+ ) {
157
176
  resumable.storage.clear(threadId);
158
177
  }
159
178
  accumulator = "";
@@ -195,6 +214,7 @@ function wrapPrepareReconnect(
195
214
  const userPrepared = await userPrepareReconnect?.({ ...options, api });
196
215
  const headers = new Headers(userPrepared?.headers ?? options.headers);
197
216
  headers.set(RESUMABLE_THREAD_ID_HEADER, options.id);
217
+ headers.set(RESUMABLE_RECONNECT_ID_HEADER, streamId);
198
218
  return {
199
219
  ...userPrepared,
200
220
  headers,