@assistant-ui/ai-sdk 0.0.4 → 0.0.6
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.
- package/dist/adapters/aiSDKFormatAdapter.d.ts +2 -8
- package/dist/adapters/aiSDKFormatAdapter.js +1 -25
- package/dist/converters/convertMessage.d.ts.map +1 -1
- package/dist/converters/convertMessage.js +9 -4
- package/dist/converters/convertMessage.js.map +1 -1
- package/dist/runtime/AISDKThreads.d.ts.map +1 -1
- package/dist/runtime/AISDKThreads.js +90 -65
- package/dist/runtime/AISDKThreads.js.map +1 -1
- package/dist/runtime/sdkIdentity.d.ts +6 -0
- package/dist/runtime/sdkIdentity.d.ts.map +1 -0
- package/dist/runtime/sdkIdentity.js +9 -0
- package/dist/runtime/sdkIdentity.js.map +1 -0
- package/dist/runtime/useAISDKRuntime.d.ts.map +1 -1
- package/dist/runtime/useAISDKRuntime.js +8 -1
- package/dist/runtime/useAISDKRuntime.js.map +1 -1
- package/dist/runtime/useChatRuntime.d.ts.map +1 -1
- package/dist/runtime/useChatRuntime.js +7 -2
- package/dist/runtime/useChatRuntime.js.map +1 -1
- package/dist/runtime/useChatThread.js +1 -1
- package/dist/runtime/useChatThread.js.map +1 -1
- package/dist/runtime/useExternalHistory.d.ts.map +1 -1
- package/dist/runtime/useExternalHistory.js +4 -1
- package/dist/runtime/useExternalHistory.js.map +1 -1
- package/dist/tools/generativeTools.js +7 -4
- package/dist/tools/generativeTools.js.map +1 -1
- package/dist/usage.d.ts +8 -0
- package/dist/usage.d.ts.map +1 -1
- package/dist/usage.js +8 -0
- package/dist/usage.js.map +1 -1
- package/package.json +12 -12
- package/src/adapters/aiSDKFormatAdapter.ts +4 -41
- package/src/converters/convertMessage.test.ts +116 -0
- package/src/converters/convertMessage.ts +39 -18
- package/src/runtime/AISDKChat.integration.test.tsx +48 -27
- package/src/runtime/AISDKThreads.cloud.test.ts +12 -3
- package/src/runtime/AISDKThreads.test.ts +131 -8
- package/src/runtime/AISDKThreads.ts +31 -3
- package/src/runtime/__tests__/controlled-transport.ts +21 -0
- package/src/runtime/sdkIdentity.ts +9 -0
- package/src/runtime/useAISDKRuntime.test.ts +39 -1
- package/src/runtime/useAISDKRuntime.ts +17 -1
- package/src/runtime/useChatRuntime.integration.test.tsx +137 -4
- package/src/runtime/useChatRuntime.test.ts +0 -1
- package/src/runtime/useChatRuntime.ts +3 -1
- package/src/runtime/useChatThread.test.ts +74 -0
- package/src/runtime/useChatThread.ts +1 -1
- package/src/runtime/useExternalHistory.test.ts +29 -0
- package/src/runtime/useExternalHistory.ts +7 -1
- package/src/tools/generativeTools.test.ts +79 -0
- package/src/tools/generativeTools.ts +7 -8
- package/src/transport/AssistantChatTransport.test.ts +1 -9
- package/src/usage.ts +8 -0
- package/dist/adapters/aiSDKFormatAdapter.d.ts.map +0 -1
- package/dist/adapters/aiSDKFormatAdapter.js.map +0 -1
|
@@ -54,23 +54,28 @@ const mocks = vi.hoisted(() => {
|
|
|
54
54
|
return { history };
|
|
55
55
|
},
|
|
56
56
|
};
|
|
57
|
-
return {
|
|
57
|
+
return {
|
|
58
|
+
adapter,
|
|
59
|
+
useCloudThreadListAdapter: vi.fn(() => adapter),
|
|
60
|
+
};
|
|
58
61
|
});
|
|
59
62
|
|
|
60
63
|
vi.mock("@assistant-ui/core/react", async (importOriginal) => ({
|
|
61
64
|
...(await importOriginal<typeof import("@assistant-ui/core/react")>()),
|
|
62
|
-
useCloudThreadListAdapter:
|
|
65
|
+
useCloudThreadListAdapter: mocks.useCloudThreadListAdapter,
|
|
63
66
|
}));
|
|
64
67
|
|
|
65
68
|
import { AISDKThreads } from "./AISDKThreads";
|
|
66
69
|
import { createCancellableTransport } from "./__tests__/controlled-transport";
|
|
70
|
+
import { AI_SDK_SDK } from "./sdkIdentity";
|
|
67
71
|
|
|
68
72
|
describe("AISDKThreads cloud", () => {
|
|
69
73
|
it("reloads history when switching a keyed cloud thread", async () => {
|
|
74
|
+
const cloud = {} as AssistantCloud;
|
|
70
75
|
const handle = createAssistantClient(
|
|
71
76
|
AuiConfig({
|
|
72
77
|
threads: AISDKThreads({
|
|
73
|
-
cloud
|
|
78
|
+
cloud,
|
|
74
79
|
threadId: "t1",
|
|
75
80
|
}),
|
|
76
81
|
}),
|
|
@@ -84,6 +89,10 @@ describe("AISDKThreads cloud", () => {
|
|
|
84
89
|
await vi.waitFor(() => {
|
|
85
90
|
expect(load).toHaveBeenCalled();
|
|
86
91
|
});
|
|
92
|
+
expect(mocks.useCloudThreadListAdapter).toHaveBeenCalledWith({
|
|
93
|
+
cloud,
|
|
94
|
+
sdk: AI_SDK_SDK,
|
|
95
|
+
});
|
|
87
96
|
const afterFirst = load.mock.calls.length;
|
|
88
97
|
flushTapSync(() => aui.threads.switchToThread("t2"));
|
|
89
98
|
await vi.waitFor(() => {
|
|
@@ -251,17 +251,28 @@ describe("AISDKThreads", () => {
|
|
|
251
251
|
}
|
|
252
252
|
});
|
|
253
253
|
|
|
254
|
-
it("forwards ChatInit callbacks to each thread's chat", async () => {
|
|
254
|
+
it("forwards ChatInit callbacks to each thread's chat from the latest render", async () => {
|
|
255
255
|
const { transport, emit, close } = createControlledTransport();
|
|
256
|
-
const
|
|
257
|
-
const
|
|
258
|
-
|
|
259
|
-
|
|
260
|
-
|
|
261
|
-
|
|
256
|
+
const onFinishA = vi.fn();
|
|
257
|
+
const onFinishB = vi.fn();
|
|
258
|
+
let onFinish = onFinishA;
|
|
259
|
+
const listeners = new Set<() => void>();
|
|
260
|
+
const handle = createAssistantClient({
|
|
261
|
+
getConfig: () =>
|
|
262
|
+
AuiConfig({
|
|
263
|
+
threads: AISDKThreads({ transport: () => transport, onFinish }),
|
|
264
|
+
}),
|
|
265
|
+
subscribe: (listener) => {
|
|
266
|
+
listeners.add(listener);
|
|
267
|
+
return () => listeners.delete(listener);
|
|
268
|
+
},
|
|
269
|
+
});
|
|
262
270
|
handle.subscribe(() => {});
|
|
263
271
|
const aui = handle.getClient();
|
|
264
272
|
|
|
273
|
+
onFinish = onFinishB;
|
|
274
|
+
flushTapSync(() => listeners.forEach((listener) => listener()));
|
|
275
|
+
|
|
265
276
|
flushTapSync(() => aui.composer.setText("hi"));
|
|
266
277
|
flushTapSync(() => aui.composer.send());
|
|
267
278
|
await vi.waitFor(() => {
|
|
@@ -271,11 +282,123 @@ describe("AISDKThreads", () => {
|
|
|
271
282
|
});
|
|
272
283
|
emit(...textReply("done"));
|
|
273
284
|
close();
|
|
274
|
-
await vi.waitFor(() => expect(
|
|
285
|
+
await vi.waitFor(() => expect(onFinishB).toHaveBeenCalledTimes(1));
|
|
286
|
+
expect(onFinishA).not.toHaveBeenCalled();
|
|
275
287
|
|
|
276
288
|
handle.destroy();
|
|
277
289
|
});
|
|
278
290
|
|
|
291
|
+
it("forwards the latest callbacks to a switched-away thread still streaming in the background", async () => {
|
|
292
|
+
const { transport, emit, close } = createControlledTransport();
|
|
293
|
+
const onFinishA = vi.fn();
|
|
294
|
+
const onFinishB = vi.fn();
|
|
295
|
+
let onFinish = onFinishA;
|
|
296
|
+
const listeners = new Set<() => void>();
|
|
297
|
+
const handle = createAssistantClient({
|
|
298
|
+
getConfig: () =>
|
|
299
|
+
AuiConfig({
|
|
300
|
+
threads: AISDKThreads({ transport: () => transport, onFinish }),
|
|
301
|
+
}),
|
|
302
|
+
subscribe: (listener) => {
|
|
303
|
+
listeners.add(listener);
|
|
304
|
+
return () => listeners.delete(listener);
|
|
305
|
+
},
|
|
306
|
+
});
|
|
307
|
+
handle.subscribe(() => {});
|
|
308
|
+
const aui = handle.getClient();
|
|
309
|
+
|
|
310
|
+
flushTapSync(() => aui.composer.setText("stream me"));
|
|
311
|
+
flushTapSync(() => aui.composer.send());
|
|
312
|
+
await vi.waitFor(() => {
|
|
313
|
+
expect(
|
|
314
|
+
handle.getClient().thread.getState().messages.length,
|
|
315
|
+
).toBeGreaterThan(0);
|
|
316
|
+
});
|
|
317
|
+
emit(
|
|
318
|
+
{ type: "start" },
|
|
319
|
+
{ type: "text-start", id: "t1" },
|
|
320
|
+
{ type: "text-delta", id: "t1", delta: "partial" },
|
|
321
|
+
);
|
|
322
|
+
|
|
323
|
+
flushTapSync(() => aui.threads.switchToNewThread());
|
|
324
|
+
onFinish = onFinishB;
|
|
325
|
+
flushTapSync(() => listeners.forEach((listener) => listener()));
|
|
326
|
+
|
|
327
|
+
emit({ type: "text-end", id: "t1" }, { type: "finish" });
|
|
328
|
+
close();
|
|
329
|
+
await vi.waitFor(() => expect(onFinishB).toHaveBeenCalledTimes(1));
|
|
330
|
+
expect(onFinishA).not.toHaveBeenCalled();
|
|
331
|
+
|
|
332
|
+
handle.destroy();
|
|
333
|
+
});
|
|
334
|
+
|
|
335
|
+
it("forwards the latest callbacks to a cloud thread's chat", async () => {
|
|
336
|
+
const cloudThread = (id: string) => ({
|
|
337
|
+
id,
|
|
338
|
+
title: id,
|
|
339
|
+
is_archived: false,
|
|
340
|
+
last_message_at: null,
|
|
341
|
+
external_id: null,
|
|
342
|
+
metadata: null,
|
|
343
|
+
});
|
|
344
|
+
const cloud = {
|
|
345
|
+
threads: {
|
|
346
|
+
list: vi.fn(async () => ({ threads: [cloudThread("t1")] })),
|
|
347
|
+
create: vi.fn(),
|
|
348
|
+
update: vi.fn(),
|
|
349
|
+
delete: vi.fn(),
|
|
350
|
+
get: vi.fn(async (id: string) => cloudThread(id)),
|
|
351
|
+
messages: {
|
|
352
|
+
list: vi.fn(async () => ({ messages: [] })),
|
|
353
|
+
create: vi.fn(async () => ({ message_id: "remote-message-1" })),
|
|
354
|
+
update: vi.fn(),
|
|
355
|
+
},
|
|
356
|
+
},
|
|
357
|
+
runs: { stream: vi.fn(), report: vi.fn() },
|
|
358
|
+
telemetry: { enabled: false },
|
|
359
|
+
} as unknown as AssistantCloud;
|
|
360
|
+
const { transport, emit, close } = createControlledTransport();
|
|
361
|
+
const onFinishA = vi.fn();
|
|
362
|
+
const onFinishB = vi.fn();
|
|
363
|
+
let onFinish = onFinishA;
|
|
364
|
+
const listeners = new Set<() => void>();
|
|
365
|
+
const handle = createAssistantClient({
|
|
366
|
+
getConfig: () =>
|
|
367
|
+
AuiConfig({
|
|
368
|
+
threads: AISDKThreads({ cloud, threadId: "t1", transport, onFinish }),
|
|
369
|
+
}),
|
|
370
|
+
subscribe: (listener) => {
|
|
371
|
+
listeners.add(listener);
|
|
372
|
+
return () => listeners.delete(listener);
|
|
373
|
+
},
|
|
374
|
+
});
|
|
375
|
+
handle.subscribe(() => {});
|
|
376
|
+
try {
|
|
377
|
+
await handle.getClient().threads.getLoadThreadsPromise();
|
|
378
|
+
await vi.waitFor(() => {
|
|
379
|
+
expect(handle.getClient().threads.getState().mainThreadId).toBe("t1");
|
|
380
|
+
});
|
|
381
|
+
await vi.waitFor(() => {
|
|
382
|
+
expect(handle.getClient().thread.getState().isLoading).toBe(false);
|
|
383
|
+
});
|
|
384
|
+
|
|
385
|
+
onFinish = onFinishB;
|
|
386
|
+
flushTapSync(() => listeners.forEach((listener) => listener()));
|
|
387
|
+
|
|
388
|
+
flushTapSync(() => handle.getClient().composer.setText("hi"));
|
|
389
|
+
flushTapSync(() => handle.getClient().composer.send());
|
|
390
|
+
await vi.waitFor(() => {
|
|
391
|
+
expect(handle.getClient().thread.getState().isRunning).toBe(true);
|
|
392
|
+
});
|
|
393
|
+
emit(...textReply("done"));
|
|
394
|
+
close();
|
|
395
|
+
await vi.waitFor(() => expect(onFinishB).toHaveBeenCalledTimes(1));
|
|
396
|
+
expect(onFinishA).not.toHaveBeenCalled();
|
|
397
|
+
} finally {
|
|
398
|
+
handle.destroy();
|
|
399
|
+
}
|
|
400
|
+
});
|
|
401
|
+
|
|
279
402
|
it("posts each thread's own id as the chat id", async () => {
|
|
280
403
|
const bodies: unknown[] = [];
|
|
281
404
|
const fetchStub = vi.fn(async (_url: unknown, init?: RequestInit) => {
|
|
@@ -26,6 +26,7 @@ import {
|
|
|
26
26
|
} from "./useChatThread";
|
|
27
27
|
import { MessageRepository } from "@assistant-ui/core/internal";
|
|
28
28
|
import { useResourceCleanup } from "./useResourceCleanup";
|
|
29
|
+
import { AI_SDK_SDK } from "./sdkIdentity";
|
|
29
30
|
|
|
30
31
|
export type AISDKThreadsOptions<UI_MESSAGE extends UIMessage = UIMessage> =
|
|
31
32
|
Omit<ChatThreadOptions<UI_MESSAGE>, "id" | "transport" | "messages"> & {
|
|
@@ -62,16 +63,22 @@ type AISDKThreadChatOptions<UI_MESSAGE extends UIMessage = UIMessage> = Omit<
|
|
|
62
63
|
"cloud" | "threadId" | "onThreadIdChange"
|
|
63
64
|
>;
|
|
64
65
|
|
|
66
|
+
type ChatOptionsRef<UI_MESSAGE extends UIMessage> = {
|
|
67
|
+
current: AISDKThreadChatOptions<UI_MESSAGE> | undefined;
|
|
68
|
+
};
|
|
69
|
+
|
|
65
70
|
type ChatEntry<UI_MESSAGE extends UIMessage> = {
|
|
66
71
|
chat: Chat<UI_MESSAGE>;
|
|
67
72
|
transport: ChatTransport<UI_MESSAGE>;
|
|
68
73
|
repository: MessageRepository;
|
|
74
|
+
optionsRef: ChatOptionsRef<UI_MESSAGE>;
|
|
69
75
|
};
|
|
70
76
|
|
|
71
77
|
const createChatEntry = <UI_MESSAGE extends UIMessage>(
|
|
72
78
|
threadId: string,
|
|
73
79
|
options: AISDKThreadChatOptions<UI_MESSAGE> | undefined,
|
|
74
80
|
): ChatEntry<UI_MESSAGE> => {
|
|
81
|
+
const optionsRef: ChatOptionsRef<UI_MESSAGE> = { current: options };
|
|
75
82
|
const { chatInit } = splitChatThreadOptions(
|
|
76
83
|
options as ChatThreadOptions<UI_MESSAGE> | undefined,
|
|
77
84
|
);
|
|
@@ -84,9 +91,20 @@ const createChatEntry = <UI_MESSAGE extends UIMessage>(
|
|
|
84
91
|
? options.transport.__internal_clone()
|
|
85
92
|
: options.transport;
|
|
86
93
|
return {
|
|
87
|
-
chat: new Chat<UI_MESSAGE>({
|
|
94
|
+
chat: new Chat<UI_MESSAGE>({
|
|
95
|
+
...chatInit,
|
|
96
|
+
id: threadId,
|
|
97
|
+
transport,
|
|
98
|
+
onToolCall: (arg) => optionsRef.current?.onToolCall?.(arg),
|
|
99
|
+
onData: (arg) => optionsRef.current?.onData?.(arg),
|
|
100
|
+
onFinish: (arg) => optionsRef.current?.onFinish?.(arg),
|
|
101
|
+
onError: (arg) => optionsRef.current?.onError?.(arg),
|
|
102
|
+
sendAutomaticallyWhen: (arg) =>
|
|
103
|
+
optionsRef.current?.sendAutomaticallyWhen?.(arg) ?? false,
|
|
104
|
+
}),
|
|
88
105
|
transport,
|
|
89
106
|
repository: new MessageRepository(),
|
|
107
|
+
optionsRef,
|
|
90
108
|
};
|
|
91
109
|
};
|
|
92
110
|
|
|
@@ -116,9 +134,13 @@ const useAISDKChatThread = <UI_MESSAGE extends UIMessage = UIMessage>({
|
|
|
116
134
|
const [owned] = useState(() =>
|
|
117
135
|
cloud ? createChatEntry(threadId, options) : undefined,
|
|
118
136
|
);
|
|
119
|
-
const { chat, transport, repository } =
|
|
137
|
+
const { chat, transport, repository, optionsRef } =
|
|
120
138
|
owned ?? getOrCreateChatEntry(threadId, options, chats);
|
|
121
139
|
|
|
140
|
+
useEffect(() => {
|
|
141
|
+
if (cloud) optionsRef.current = options;
|
|
142
|
+
});
|
|
143
|
+
|
|
122
144
|
useEffect(() => {
|
|
123
145
|
if (!cloud) return undefined;
|
|
124
146
|
return () => {
|
|
@@ -173,13 +195,19 @@ const useAISDKThreads = <UI_MESSAGE extends UIMessage = UIMessage>(
|
|
|
173
195
|
const [chats] = useState(() => new Map<string, ChatEntry<UI_MESSAGE>>());
|
|
174
196
|
const bindCloud = cloud !== undefined;
|
|
175
197
|
|
|
198
|
+
useEffect(() => {
|
|
199
|
+
for (const { optionsRef } of chats.values()) {
|
|
200
|
+
optionsRef.current = threadOptions;
|
|
201
|
+
}
|
|
202
|
+
});
|
|
203
|
+
|
|
176
204
|
useResourceCleanup(true, () => {
|
|
177
205
|
for (const { chat } of chats.values()) {
|
|
178
206
|
void chat.stop().catch(() => {});
|
|
179
207
|
}
|
|
180
208
|
});
|
|
181
209
|
|
|
182
|
-
const cloudAdapter = useCloudThreadListAdapter({ cloud });
|
|
210
|
+
const cloudAdapter = useCloudThreadListAdapter({ cloud, sdk: AI_SDK_SDK });
|
|
183
211
|
const thread = (id: string) => {
|
|
184
212
|
const element = AISDKChatThread({
|
|
185
213
|
threadId: id,
|
|
@@ -1,4 +1,6 @@
|
|
|
1
1
|
import type { ChatTransport, UIMessage, UIMessageChunk } from "ai";
|
|
2
|
+
import { useAui, type AssistantClient } from "@assistant-ui/store";
|
|
3
|
+
import { flushTapSync } from "@assistant-ui/tap";
|
|
2
4
|
|
|
3
5
|
export const createControlledTransport = () => {
|
|
4
6
|
let controller!: ReadableStreamDefaultController<UIMessageChunk>;
|
|
@@ -41,3 +43,22 @@ export const createCancellableTransport = () => {
|
|
|
41
43
|
close: () => controller.close(),
|
|
42
44
|
};
|
|
43
45
|
};
|
|
46
|
+
|
|
47
|
+
export const nextTask = () => new Promise((resolve) => setTimeout(resolve, 0));
|
|
48
|
+
|
|
49
|
+
export const createStreamHarness = () => {
|
|
50
|
+
let aui: AssistantClient | undefined;
|
|
51
|
+
const Probe = () => {
|
|
52
|
+
aui = useAui();
|
|
53
|
+
return null;
|
|
54
|
+
};
|
|
55
|
+
return {
|
|
56
|
+
Probe,
|
|
57
|
+
send: () => {
|
|
58
|
+
flushTapSync(() => aui!.composer.setText("keep streaming"));
|
|
59
|
+
flushTapSync(() => aui!.composer.send());
|
|
60
|
+
},
|
|
61
|
+
isRunning: () => aui?.thread.getState().isRunning === true,
|
|
62
|
+
client: () => aui!,
|
|
63
|
+
};
|
|
64
|
+
};
|
|
@@ -77,7 +77,6 @@ const textOf = (message: any): string =>
|
|
|
77
77
|
|
|
78
78
|
describe("useAISDKRuntime", () => {
|
|
79
79
|
beforeEach(() => {
|
|
80
|
-
vi.clearAllMocks();
|
|
81
80
|
vi.mocked(useExternalHistory).mockReturnValue({
|
|
82
81
|
isLoading: false,
|
|
83
82
|
deleteMessage: vi.fn().mockResolvedValue(undefined),
|
|
@@ -1487,4 +1486,43 @@ describe("useAISDKRuntime", () => {
|
|
|
1487
1486
|
error: chat.error,
|
|
1488
1487
|
});
|
|
1489
1488
|
});
|
|
1489
|
+
|
|
1490
|
+
it("keeps the error name and code on the failed message status", () => {
|
|
1491
|
+
const chat = createChatHelpers([
|
|
1492
|
+
{ id: "u1", role: "user", parts: [{ type: "text", text: "hi" }] },
|
|
1493
|
+
{ id: "a1", role: "assistant", parts: [{ type: "text", text: "" }] },
|
|
1494
|
+
]);
|
|
1495
|
+
chat.error = Object.assign(new Error("rate limited"), {
|
|
1496
|
+
name: "AI_APICallError",
|
|
1497
|
+
code: "rate_limited",
|
|
1498
|
+
});
|
|
1499
|
+
|
|
1500
|
+
const { result } = renderHook(() => useAISDKRuntime(chat));
|
|
1501
|
+
|
|
1502
|
+
expect(
|
|
1503
|
+
result.current.thread.getState().messages.at(-1)?.status,
|
|
1504
|
+
).toMatchObject({
|
|
1505
|
+
type: "incomplete",
|
|
1506
|
+
reason: "error",
|
|
1507
|
+
error: { code: "rate_limited", message: "rate limited" },
|
|
1508
|
+
});
|
|
1509
|
+
});
|
|
1510
|
+
|
|
1511
|
+
it("uses the error name as the code when the error carries none", () => {
|
|
1512
|
+
const chat = createChatHelpers([
|
|
1513
|
+
{ id: "u1", role: "user", parts: [{ type: "text", text: "hi" }] },
|
|
1514
|
+
{ id: "a1", role: "assistant", parts: [{ type: "text", text: "" }] },
|
|
1515
|
+
]);
|
|
1516
|
+
chat.error = Object.assign(new Error("upstream failed"), {
|
|
1517
|
+
name: "AI_APICallError",
|
|
1518
|
+
});
|
|
1519
|
+
|
|
1520
|
+
const { result } = renderHook(() => useAISDKRuntime(chat));
|
|
1521
|
+
|
|
1522
|
+
expect(
|
|
1523
|
+
result.current.thread.getState().messages.at(-1)?.status,
|
|
1524
|
+
).toMatchObject({
|
|
1525
|
+
error: { code: "AI_APICallError", message: "upstream failed" },
|
|
1526
|
+
});
|
|
1527
|
+
});
|
|
1490
1528
|
});
|
|
@@ -47,6 +47,7 @@ import {
|
|
|
47
47
|
MessageRepository,
|
|
48
48
|
} from "@assistant-ui/core/internal";
|
|
49
49
|
import type { ReadonlyJSONObject } from "assistant-stream/utils";
|
|
50
|
+
import type { AssistantError } from "@assistant-ui/core";
|
|
50
51
|
import { sliceMessagesUntil } from "../utils/sliceMessagesUntil";
|
|
51
52
|
import { toCreateMessage } from "../converters/toCreateMessage";
|
|
52
53
|
import { vercelAttachmentAdapter } from "../adapters/vercelAttachmentAdapter";
|
|
@@ -219,6 +220,19 @@ const useGeneratedSuggestions = (
|
|
|
219
220
|
|
|
220
221
|
const NO_CANCELLED_MESSAGE_IDS: ReadonlySet<string> = new Set();
|
|
221
222
|
|
|
223
|
+
const toChatError = (error: Error): AssistantError => {
|
|
224
|
+
const code = (error as { code?: unknown }).code;
|
|
225
|
+
return {
|
|
226
|
+
code:
|
|
227
|
+
typeof code === "string"
|
|
228
|
+
? code
|
|
229
|
+
: error.name !== "Error"
|
|
230
|
+
? error.name
|
|
231
|
+
: "unknown",
|
|
232
|
+
message: error.message,
|
|
233
|
+
};
|
|
234
|
+
};
|
|
235
|
+
|
|
222
236
|
export const useAISDKRuntime = <UI_MESSAGE extends UIMessage = UIMessage>(
|
|
223
237
|
chatHelpers: ReturnType<typeof useChat<UI_MESSAGE>>,
|
|
224
238
|
adapter: AISDKRuntimeAdapter<UI_MESSAGE> = {},
|
|
@@ -315,7 +329,9 @@ export const useAISDKRuntime = <UI_MESSAGE extends UIMessage = UIMessage>(
|
|
|
315
329
|
toolLastInputCache: toolLastInputCacheRef.current,
|
|
316
330
|
mcpAppMetadataCache: mcpAppMetadataCacheRef.current,
|
|
317
331
|
...(optimisticMessageId && { optimisticMessageId }),
|
|
318
|
-
...(chatHelpers.error && {
|
|
332
|
+
...(chatHelpers.error && {
|
|
333
|
+
error: toChatError(chatHelpers.error),
|
|
334
|
+
}),
|
|
319
335
|
...(cancelledMessageIds.size > 0 && { cancelledMessageIds }),
|
|
320
336
|
}),
|
|
321
337
|
[
|
|
@@ -1,12 +1,19 @@
|
|
|
1
1
|
// @vitest-environment jsdom
|
|
2
2
|
|
|
3
|
-
import { render, screen, waitFor } from "@testing-library/react";
|
|
3
|
+
import { act, render, screen, waitFor } from "@testing-library/react";
|
|
4
4
|
import { AssistantRuntimeProvider } from "@assistant-ui/core/react";
|
|
5
|
-
import { useAuiState } from "@assistant-ui/store";
|
|
6
|
-
import type {
|
|
7
|
-
import {
|
|
5
|
+
import { AuiConfig, AuiProvider, useAuiState } from "@assistant-ui/store";
|
|
6
|
+
import type { AssistantRuntime } from "@assistant-ui/core";
|
|
7
|
+
import { AISDKChat } from "./AISDKChat";
|
|
8
|
+
import type { ChatTransport, UIMessage } from "ai";
|
|
9
|
+
import { Activity, StrictMode, useState, type ReactNode } from "react";
|
|
8
10
|
import { describe, expect, it } from "vitest";
|
|
9
11
|
import { AssistantChatTransport } from "../transport/AssistantChatTransport";
|
|
12
|
+
import {
|
|
13
|
+
createCancellableTransport,
|
|
14
|
+
createStreamHarness,
|
|
15
|
+
nextTask,
|
|
16
|
+
} from "./__tests__/controlled-transport";
|
|
10
17
|
import { useChatRuntime } from "./useChatRuntime";
|
|
11
18
|
import { useThreadTokenUsage } from "../usage";
|
|
12
19
|
|
|
@@ -65,8 +72,134 @@ describe("useChatRuntime integration", () => {
|
|
|
65
72
|
);
|
|
66
73
|
});
|
|
67
74
|
});
|
|
75
|
+
|
|
76
|
+
it("aborts a deleted thread's stream while the host stays mounted", async () => {
|
|
77
|
+
const { transport, getCancelCount } = createCancellableTransport();
|
|
78
|
+
const { Probe, send, isRunning, client } = createStreamHarness();
|
|
79
|
+
|
|
80
|
+
const view = render(
|
|
81
|
+
<StrictMode>
|
|
82
|
+
<StreamingApp transport={transport} probe={<Probe />} />
|
|
83
|
+
</StrictMode>,
|
|
84
|
+
);
|
|
85
|
+
|
|
86
|
+
await act(async () => send());
|
|
87
|
+
await waitFor(() => expect(isRunning()).toBe(true));
|
|
88
|
+
|
|
89
|
+
await act(async () => client().threadListItem.delete());
|
|
90
|
+
|
|
91
|
+
await waitFor(() => expect(getCancelCount()).toBe(1));
|
|
92
|
+
view.unmount();
|
|
93
|
+
});
|
|
94
|
+
|
|
95
|
+
it("aborts the in-flight transport after a real unmount", async () => {
|
|
96
|
+
const { transport, getCancelCount } = createCancellableTransport();
|
|
97
|
+
const { Probe, send, isRunning } = createStreamHarness();
|
|
98
|
+
|
|
99
|
+
const view = render(
|
|
100
|
+
<StrictMode>
|
|
101
|
+
<StreamingApp transport={transport} probe={<Probe />} />
|
|
102
|
+
</StrictMode>,
|
|
103
|
+
);
|
|
104
|
+
|
|
105
|
+
await act(async () => send());
|
|
106
|
+
await waitFor(() => expect(isRunning()).toBe(true));
|
|
107
|
+
// the Strict Mode double mount already ran a host cleanup by now
|
|
108
|
+
expect(getCancelCount()).toBe(0);
|
|
109
|
+
|
|
110
|
+
view.unmount();
|
|
111
|
+
await waitFor(() => expect(getCancelCount()).toBe(1));
|
|
112
|
+
});
|
|
113
|
+
|
|
114
|
+
it("keeps streaming while hidden and aborts when the hidden host unmounts", async () => {
|
|
115
|
+
const { transport, getCancelCount } = createCancellableTransport();
|
|
116
|
+
const { Probe, send, isRunning } = createStreamHarness();
|
|
117
|
+
|
|
118
|
+
let setMode: ((mode: "visible" | "hidden") => void) | undefined;
|
|
119
|
+
const Shell = () => {
|
|
120
|
+
const [mode, set] = useState<"visible" | "hidden">("visible");
|
|
121
|
+
setMode = set;
|
|
122
|
+
return (
|
|
123
|
+
<Activity mode={mode}>
|
|
124
|
+
<StreamingApp transport={transport} probe={<Probe />} />
|
|
125
|
+
</Activity>
|
|
126
|
+
);
|
|
127
|
+
};
|
|
128
|
+
|
|
129
|
+
const view = render(
|
|
130
|
+
<StrictMode>
|
|
131
|
+
<Shell />
|
|
132
|
+
</StrictMode>,
|
|
133
|
+
);
|
|
134
|
+
|
|
135
|
+
await act(async () => send());
|
|
136
|
+
await waitFor(() => expect(isRunning()).toBe(true));
|
|
137
|
+
|
|
138
|
+
await act(async () => setMode?.("hidden"));
|
|
139
|
+
await act(nextTask);
|
|
140
|
+
expect(getCancelCount()).toBe(0);
|
|
141
|
+
expect(isRunning()).toBe(true);
|
|
142
|
+
|
|
143
|
+
await act(async () => setMode?.("visible"));
|
|
144
|
+
await act(nextTask);
|
|
145
|
+
expect(getCancelCount()).toBe(0);
|
|
146
|
+
expect(isRunning()).toBe(true);
|
|
147
|
+
|
|
148
|
+
await act(async () => setMode?.("hidden"));
|
|
149
|
+
view.unmount();
|
|
150
|
+
await waitFor(() => expect(getCancelCount()).toBe(1));
|
|
151
|
+
});
|
|
152
|
+
|
|
153
|
+
it("aborts a nested runtime's stream when the provider above it unmounts", async () => {
|
|
154
|
+
const outer = createCancellableTransport();
|
|
155
|
+
const { transport, getCancelCount } = createCancellableTransport();
|
|
156
|
+
let nested: AssistantRuntime | undefined;
|
|
157
|
+
|
|
158
|
+
// allowNesting: the inner useChatRuntime runs its thread hook directly, as
|
|
159
|
+
// a plain React hook under the provider rather than inside a tap resource.
|
|
160
|
+
const NestedChat = () => {
|
|
161
|
+
nested = useChatRuntime({ transport });
|
|
162
|
+
return null;
|
|
163
|
+
};
|
|
164
|
+
|
|
165
|
+
const view = render(
|
|
166
|
+
<StrictMode>
|
|
167
|
+
<AuiProvider
|
|
168
|
+
config={AuiConfig({
|
|
169
|
+
threads: AISDKChat({ transport: outer.transport }),
|
|
170
|
+
})}
|
|
171
|
+
>
|
|
172
|
+
<NestedChat />
|
|
173
|
+
</AuiProvider>
|
|
174
|
+
</StrictMode>,
|
|
175
|
+
);
|
|
176
|
+
|
|
177
|
+
await waitFor(() => expect(nested).toBeDefined());
|
|
178
|
+
await act(async () => {
|
|
179
|
+
await nested!.thread.append("keep streaming");
|
|
180
|
+
});
|
|
181
|
+
await waitFor(() => expect(nested!.thread.getState().isRunning).toBe(true));
|
|
182
|
+
|
|
183
|
+
view.unmount();
|
|
184
|
+
await waitFor(() => expect(getCancelCount()).toBe(1));
|
|
185
|
+
});
|
|
68
186
|
});
|
|
69
187
|
|
|
188
|
+
const StreamingApp = ({
|
|
189
|
+
transport,
|
|
190
|
+
probe,
|
|
191
|
+
}: {
|
|
192
|
+
transport: ChatTransport<UIMessage>;
|
|
193
|
+
probe: ReactNode;
|
|
194
|
+
}) => {
|
|
195
|
+
const runtime = useChatRuntime({ transport });
|
|
196
|
+
return (
|
|
197
|
+
<AssistantRuntimeProvider runtime={runtime}>
|
|
198
|
+
{probe}
|
|
199
|
+
</AssistantRuntimeProvider>
|
|
200
|
+
);
|
|
201
|
+
};
|
|
202
|
+
|
|
70
203
|
const UsageProbe = () => {
|
|
71
204
|
const usage = useThreadTokenUsage();
|
|
72
205
|
return (
|
|
@@ -9,6 +9,7 @@ import {
|
|
|
9
9
|
} from "@assistant-ui/core/react";
|
|
10
10
|
import { useAui, useAuiState } from "@assistant-ui/store";
|
|
11
11
|
import { useChatThread, type ChatThreadOptions } from "./useChatThread";
|
|
12
|
+
import { AI_SDK_SDK } from "./sdkIdentity";
|
|
12
13
|
|
|
13
14
|
export type UseChatRuntimeOptions<UI_MESSAGE extends UIMessage = UIMessage> =
|
|
14
15
|
ChatThreadOptions<UI_MESSAGE> & {
|
|
@@ -29,6 +30,7 @@ const useChatThreadRuntime = <UI_MESSAGE extends UIMessage = UIMessage>(
|
|
|
29
30
|
isMainThread,
|
|
30
31
|
getThreadListItem: () =>
|
|
31
32
|
aui.threadListItem.source ? aui.threadListItem : undefined,
|
|
33
|
+
stopOnClientDestroy: true,
|
|
32
34
|
});
|
|
33
35
|
};
|
|
34
36
|
|
|
@@ -37,7 +39,7 @@ export const useChatRuntime = <UI_MESSAGE extends UIMessage = UIMessage>({
|
|
|
37
39
|
onThreadIdChange,
|
|
38
40
|
...options
|
|
39
41
|
}: UseChatRuntimeOptions<UI_MESSAGE> = {}): AssistantRuntime => {
|
|
40
|
-
const cloudAdapter = useCloudThreadListAdapter({ cloud });
|
|
42
|
+
const cloudAdapter = useCloudThreadListAdapter({ cloud, sdk: AI_SDK_SDK });
|
|
41
43
|
return useRemoteThreadListRuntime({
|
|
42
44
|
runtimeHook: function RuntimeHook() {
|
|
43
45
|
return useChatThreadRuntime(options);
|
|
@@ -0,0 +1,74 @@
|
|
|
1
|
+
// @vitest-environment jsdom
|
|
2
|
+
|
|
3
|
+
import { describe, expect, it, vi } from "vitest";
|
|
4
|
+
import { resource, useResource, flushTapSync } from "@assistant-ui/tap";
|
|
5
|
+
import { useState } from "react";
|
|
6
|
+
import {
|
|
7
|
+
RuntimeAdapter,
|
|
8
|
+
runtimeAdapterTransformScopes,
|
|
9
|
+
} from "@assistant-ui/core/store";
|
|
10
|
+
import {
|
|
11
|
+
attachTransformScopes,
|
|
12
|
+
AuiConfig,
|
|
13
|
+
createAssistantClient,
|
|
14
|
+
} from "@assistant-ui/store/client";
|
|
15
|
+
import { useChatThread, type ChatThreadEnvironment } from "./useChatThread";
|
|
16
|
+
import {
|
|
17
|
+
createCancellableTransport,
|
|
18
|
+
nextTask,
|
|
19
|
+
} from "./__tests__/controlled-transport";
|
|
20
|
+
|
|
21
|
+
const createHost = (
|
|
22
|
+
env: Pick<ChatThreadEnvironment, "stopOnClientDestroy">,
|
|
23
|
+
) => {
|
|
24
|
+
const useHost = (options: Parameters<typeof useChatThread>[0]) => {
|
|
25
|
+
const [threadListItem] = useState(() => ({
|
|
26
|
+
initialize: async () => ({ remoteId: "main", externalId: undefined }),
|
|
27
|
+
}));
|
|
28
|
+
const runtime = useChatThread(options, {
|
|
29
|
+
id: "main",
|
|
30
|
+
isMainThread: true,
|
|
31
|
+
getThreadListItem: () => threadListItem,
|
|
32
|
+
...env,
|
|
33
|
+
});
|
|
34
|
+
return useResource(RuntimeAdapter(runtime));
|
|
35
|
+
};
|
|
36
|
+
attachTransformScopes(useHost, runtimeAdapterTransformScopes);
|
|
37
|
+
return resource(useHost);
|
|
38
|
+
};
|
|
39
|
+
|
|
40
|
+
const streamThenDestroy = async (
|
|
41
|
+
env: Pick<ChatThreadEnvironment, "stopOnClientDestroy">,
|
|
42
|
+
) => {
|
|
43
|
+
const { transport, getCancelCount, close } = createCancellableTransport();
|
|
44
|
+
const Host = createHost(env);
|
|
45
|
+
const handle = createAssistantClient(
|
|
46
|
+
AuiConfig({ threads: Host({ transport }) }),
|
|
47
|
+
);
|
|
48
|
+
handle.subscribe(() => {});
|
|
49
|
+
const aui = handle.getClient();
|
|
50
|
+
|
|
51
|
+
try {
|
|
52
|
+
flushTapSync(() => aui.composer.setText("stop me"));
|
|
53
|
+
flushTapSync(() => aui.composer.send());
|
|
54
|
+
await vi.waitFor(() => {
|
|
55
|
+
expect(aui.thread.getState().isRunning).toBe(true);
|
|
56
|
+
});
|
|
57
|
+
} finally {
|
|
58
|
+
handle.destroy();
|
|
59
|
+
}
|
|
60
|
+
await nextTask();
|
|
61
|
+
const cancelCount = getCancelCount();
|
|
62
|
+
if (cancelCount === 0) close();
|
|
63
|
+
return cancelCount;
|
|
64
|
+
};
|
|
65
|
+
|
|
66
|
+
describe("useChatThread", () => {
|
|
67
|
+
it("stops an in-flight chat on client destroy when stopOnClientDestroy is omitted", async () => {
|
|
68
|
+
expect(await streamThenDestroy({})).toBe(1);
|
|
69
|
+
});
|
|
70
|
+
|
|
71
|
+
it("leaves an in-flight chat running on client destroy when stopOnClientDestroy is false", async () => {
|
|
72
|
+
expect(await streamThenDestroy({ stopOnClientDestroy: false })).toBe(0);
|
|
73
|
+
});
|
|
74
|
+
});
|