@assistant-ui/ai-sdk 0.0.3 → 0.0.5

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 (49) hide show
  1. package/README.md +1 -1
  2. package/dist/converters/convertMessage.d.ts.map +1 -1
  3. package/dist/converters/convertMessage.js +31 -6
  4. package/dist/converters/convertMessage.js.map +1 -1
  5. package/dist/runtime/AISDKChat.js +1 -1
  6. package/dist/runtime/AISDKThreads.d.ts +10 -8
  7. package/dist/runtime/AISDKThreads.d.ts.map +1 -1
  8. package/dist/runtime/AISDKThreads.js +34 -27
  9. package/dist/runtime/AISDKThreads.js.map +1 -1
  10. package/dist/runtime/useAISDKRuntime.d.ts +2 -0
  11. package/dist/runtime/useAISDKRuntime.d.ts.map +1 -1
  12. package/dist/runtime/useAISDKRuntime.js +171 -68
  13. package/dist/runtime/useAISDKRuntime.js.map +1 -1
  14. package/dist/runtime/useChatRuntime.d.ts.map +1 -1
  15. package/dist/runtime/useChatRuntime.js +3 -2
  16. package/dist/runtime/useChatRuntime.js.map +1 -1
  17. package/dist/runtime/useChatThread.d.ts +7 -0
  18. package/dist/runtime/useChatThread.d.ts.map +1 -1
  19. package/dist/runtime/useChatThread.js +2 -1
  20. package/dist/runtime/useChatThread.js.map +1 -1
  21. package/dist/runtime/useExternalHistory.d.ts.map +1 -1
  22. package/dist/runtime/useExternalHistory.js +42 -47
  23. package/dist/runtime/useExternalHistory.js.map +1 -1
  24. package/dist/runtime/useResourceCleanup.js +1 -1
  25. package/dist/tools/generativeTools.js +7 -4
  26. package/dist/tools/generativeTools.js.map +1 -1
  27. package/dist/usage.d.ts +9 -2
  28. package/dist/usage.d.ts.map +1 -1
  29. package/dist/usage.js +13 -7
  30. package/dist/usage.js.map +1 -1
  31. package/package.json +11 -11
  32. package/src/converters/convertMessage.test.ts +124 -0
  33. package/src/converters/convertMessage.ts +64 -20
  34. package/src/runtime/AISDKChat.integration.test.tsx +48 -27
  35. package/src/runtime/AISDKThreads.cloud.test.ts +11 -1
  36. package/src/runtime/AISDKThreads.test.ts +102 -0
  37. package/src/runtime/AISDKThreads.ts +16 -9
  38. package/src/runtime/__tests__/controlled-transport.ts +21 -0
  39. package/src/runtime/useAISDKRuntime.test.ts +288 -3
  40. package/src/runtime/useAISDKRuntime.ts +90 -5
  41. package/src/runtime/useChatRuntime.integration.test.tsx +183 -4
  42. package/src/runtime/useChatRuntime.ts +1 -0
  43. package/src/runtime/useChatThread.ts +11 -0
  44. package/src/runtime/useExternalHistory.test.ts +64 -0
  45. package/src/runtime/useExternalHistory.ts +43 -56
  46. package/src/tools/generativeTools.test.ts +79 -0
  47. package/src/tools/generativeTools.ts +7 -8
  48. package/src/usage.test.ts +26 -8
  49. package/src/usage.ts +12 -9
@@ -1,13 +1,21 @@
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 { UIMessage } from "ai";
7
- import { StrictMode, useState } from "react";
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";
18
+ import { useThreadTokenUsage } from "../usage";
11
19
 
12
20
  const messages: UIMessage[] = [
13
21
  {
@@ -64,4 +72,175 @@ describe("useChatRuntime integration", () => {
64
72
  );
65
73
  });
66
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
+ });
186
+ });
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
+
203
+ const UsageProbe = () => {
204
+ const usage = useThreadTokenUsage();
205
+ return (
206
+ <output data-testid="total-tokens">{usage?.totalTokens ?? "none"}</output>
207
+ );
208
+ };
209
+
210
+ const UsageApp = () => {
211
+ const [transport] = useState(
212
+ () => new AssistantChatTransport({ api: "/api/chat" }),
213
+ );
214
+ const runtime = useChatRuntime({
215
+ messages: [
216
+ ...messages,
217
+ {
218
+ id: "assistant-with-usage",
219
+ role: "assistant",
220
+ parts: [{ type: "text", text: "Hi" }],
221
+ metadata: { usage: { inputTokens: 40, outputTokens: 2 } },
222
+ },
223
+ ],
224
+ transport,
225
+ });
226
+
227
+ return (
228
+ <AssistantRuntimeProvider runtime={runtime}>
229
+ <UsageProbe />
230
+ </AssistantRuntimeProvider>
231
+ );
232
+ };
233
+
234
+ describe("useThreadTokenUsage through useChatRuntime", () => {
235
+ it("reads usage from the message metadata a server attached", async () => {
236
+ render(
237
+ <StrictMode>
238
+ <UsageApp />
239
+ </StrictMode>,
240
+ );
241
+
242
+ await waitFor(() => {
243
+ expect(screen.getByTestId("total-tokens").textContent).toBe("42");
244
+ });
245
+ });
67
246
  });
@@ -29,6 +29,7 @@ const useChatThreadRuntime = <UI_MESSAGE extends UIMessage = UIMessage>(
29
29
  isMainThread,
30
30
  getThreadListItem: () =>
31
31
  aui.threadListItem.source ? aui.threadListItem : undefined,
32
+ stopOnClientDestroy: true,
32
33
  });
33
34
  };
34
35
 
@@ -1,6 +1,7 @@
1
1
  "use client";
2
2
 
3
3
  import { useChat, type Chat, type UIMessage } from "@ai-sdk/react";
4
+ import type { MessageRepository } from "@assistant-ui/core/internal";
4
5
  import {
5
6
  pickExternalStoreSharedOptions,
6
7
  type AssistantRuntime,
@@ -60,6 +61,12 @@ export type ChatThreadEnvironment<UI_MESSAGE extends UIMessage = UIMessage> = {
60
61
  * from the instance.
61
62
  */
62
63
  chat?: Chat<UI_MESSAGE> | undefined;
64
+ /**
65
+ * An externally owned per-thread message repository. Hosts that route
66
+ * multiple threads through one mounting pass a distinct instance per
67
+ * thread so histories and branches stay isolated.
68
+ */
69
+ messageRepositoryInstance?: MessageRepository | undefined;
63
70
  };
64
71
 
65
72
  const useDynamicChatTransport = <UI_MESSAGE extends UIMessage = UIMessage>(
@@ -183,6 +190,7 @@ export const useChatThread = <UI_MESSAGE extends UIMessage = UIMessage>(
183
190
  getThreadListItem,
184
191
  stopOnClientDestroy = false,
185
192
  chat: externalChat,
193
+ messageRepositoryInstance,
186
194
  } = env;
187
195
 
188
196
  const defaultTransport = useMemo(() => new AssistantChatTransport(), []);
@@ -209,6 +217,9 @@ export const useChatThread = <UI_MESSAGE extends UIMessage = UIMessage>(
209
217
  ...(onResumeToolCall && { onResumeToolCall }),
210
218
  ...(joinStrategy && { joinStrategy }),
211
219
  ...(messageRepository && { messageRepository }),
220
+ ...(messageRepositoryInstance && {
221
+ unstable_messageRepositoryInstance: messageRepositoryInstance,
222
+ }),
212
223
  ...(unstable_onBranchChange && { unstable_onBranchChange }),
213
224
  });
214
225
 
@@ -1180,6 +1180,70 @@ describe("useExternalHistory persistence", () => {
1180
1180
  expect(reportTelemetry).not.toHaveBeenCalled();
1181
1181
  });
1182
1182
 
1183
+ it("skips updates for a persisted message restored by a branch switch", async () => {
1184
+ const { append, update, runCycle, flush } = createPersistenceHarness(true);
1185
+ const completeStatus: ThreadAssistantMessage["status"] = {
1186
+ type: "complete",
1187
+ reason: "stop",
1188
+ };
1189
+ const createUserMessage = (
1190
+ id: string,
1191
+ inner: InnerMessage,
1192
+ ): ThreadMessage => {
1193
+ const message: ThreadMessage = {
1194
+ id,
1195
+ role: "user",
1196
+ content: [{ type: "text", text: "hi" }],
1197
+ attachments: [],
1198
+ createdAt: new Date(),
1199
+ metadata: { custom: {} },
1200
+ };
1201
+ bindExternalStoreMessage(message, [inner]);
1202
+ return message;
1203
+ };
1204
+ const branchA = () => [
1205
+ createUserMessage("user-a", { id: "inner-user-a", parts: ["question"] }),
1206
+ createAssistantMessage(
1207
+ completeStatus,
1208
+ [{ id: "inner-assistant-a", parts: ["answer"] }],
1209
+ "assistant-a",
1210
+ ),
1211
+ ];
1212
+
1213
+ await runCycle(branchA());
1214
+ await waitFor(() => expect(append).toHaveBeenCalledTimes(2));
1215
+
1216
+ await runCycle([
1217
+ createUserMessage("user-b", { id: "inner-user-b", parts: ["edited"] }),
1218
+ createAssistantMessage(
1219
+ completeStatus,
1220
+ [{ id: "inner-assistant-b", parts: ["answer"] }],
1221
+ "assistant-b",
1222
+ ),
1223
+ ]);
1224
+ await waitFor(() => expect(append).toHaveBeenCalledTimes(4));
1225
+
1226
+ append.mockClear();
1227
+ update.mockClear();
1228
+
1229
+ await runCycle([
1230
+ ...branchA(),
1231
+ createUserMessage("user-c", { id: "inner-user-c", parts: ["follow-up"] }),
1232
+ createAssistantMessage(
1233
+ completeStatus,
1234
+ [{ id: "inner-assistant-c", parts: ["answer"] }],
1235
+ "assistant-c",
1236
+ ),
1237
+ ]);
1238
+ await flush();
1239
+
1240
+ expect(update).not.toHaveBeenCalled();
1241
+ expect(append.mock.calls.map(([item]) => item.message.id)).toEqual([
1242
+ "inner-user-c",
1243
+ "inner-assistant-c",
1244
+ ]);
1245
+ });
1246
+
1183
1247
  it("absorbs agentic flickers without losing change detection", async () => {
1184
1248
  const { append, update, reportTelemetry, runCycle, flush, step } =
1185
1249
  createPersistenceHarness(true);
@@ -5,6 +5,7 @@ import type {
5
5
  ThreadHistoryAdapter,
6
6
  ThreadMessage,
7
7
  MessageFormatAdapter,
8
+ MessageFormatItem,
8
9
  MessageFormatRepository,
9
10
  ExportedMessageRepository,
10
11
  } from "@assistant-ui/core";
@@ -49,15 +50,10 @@ const isAwaitingToolApproval = (message: ThreadMessage) =>
49
50
  message.status?.type === "requires-action" &&
50
51
  message.status.reason === "tool-calls";
51
52
 
52
- const snapshotExternalMessages = <TMessage>(
53
- messages: readonly ThreadMessage[],
54
- ) =>
55
- new Map<string, TMessage[]>(
56
- messages.map((message) => [
57
- message.id,
58
- [...getExternalStoreMessages<TMessage>(message)],
59
- ]),
60
- );
53
+ const encodeContent = <TMessage>(
54
+ storageFormatAdapter: MessageFormatAdapter<TMessage, any>,
55
+ item: MessageFormatItem<TMessage>,
56
+ ) => JSON.stringify(storageFormatAdapter.encode(item));
61
57
 
62
58
  export const useExternalHistory = <TMessage>(
63
59
  runtimeRef: RefObject<AssistantRuntime>,
@@ -78,9 +74,11 @@ export const useExternalHistory = <TMessage>(
78
74
  const [hasLoaded, setHasLoaded] = useState(false);
79
75
 
80
76
  const historyIds = useRef(new Set<string>());
81
- const persistedInnerIds = useRef(new Set<string>());
82
77
  const deferredTelemetryIds = useRef(new Set<string>());
83
- const persistedExternalMessages = useRef(new Map<string, TMessage[]>());
78
+ // `content` is a snapshot taken at write time rather than a re-encode of `source`, because a retained message object can be mutated in place by the runtime that produced it.
79
+ const persistedInnerMessages = useRef(
80
+ new Map<string, { source: TMessage; content: string }>(),
81
+ );
84
82
 
85
83
  const onSetMessagesRef = useRef(onSetMessages);
86
84
  useEffect(() => {
@@ -107,8 +105,12 @@ export const useExternalHistory = <TMessage>(
107
105
  const repo = await formatAdapter.load();
108
106
  if (repo && repo.messages.length > 0) {
109
107
  for (const m of repo.messages) {
110
- persistedInnerIds.current.add(
108
+ persistedInnerMessages.current.set(
111
109
  storageFormatAdapter.getId(m.message),
110
+ {
111
+ source: m.message,
112
+ content: encodeContent(storageFormatAdapter, m),
113
+ },
112
114
  );
113
115
  }
114
116
  const converted = toExportedMessageRepository(toThreadMessages, repo);
@@ -129,10 +131,6 @@ export const useExternalHistory = <TMessage>(
129
131
  deferredTelemetryIds.current.add(m.message.id);
130
132
  }
131
133
  }
132
- persistedExternalMessages.current =
133
- snapshotExternalMessages<TMessage>(
134
- converted.messages.map((m) => m.message),
135
- );
136
134
  }
137
135
  } catch (error) {
138
136
  console.error("Failed to load message history:", error);
@@ -145,6 +143,9 @@ export const useExternalHistory = <TMessage>(
145
143
 
146
144
  const remoteId = optionalThreadListItem()?.getState().remoteId;
147
145
  if (!remoteId) {
146
+ // History loads asynchronously against the thread list item; without a
147
+ // remote id there is nothing to await, so the flag settles here.
148
+ // eslint-disable-next-line react-hooks/set-state-in-effect
148
149
  setHasLoaded(true);
149
150
  return aui.subscribe(() => {
150
151
  if (optionalThreadListItem()?.getState().remoteId) {
@@ -294,23 +295,8 @@ export const useExternalHistory = <TMessage>(
294
295
 
295
296
  persistInFlightRef.current = persistInFlightRef.current
296
297
  .then(async () => {
297
- const changedRunMessageIds = new Set<string>();
298
- for (const message of latest.messages) {
299
- const externalMessages =
300
- getExternalStoreMessages<TMessage>(message);
301
- const previous = persistedExternalMessages.current.get(message.id);
302
- if (
303
- previous === undefined ||
304
- previous.length !== externalMessages.length ||
305
- externalMessages.some((item, index) => item !== previous[index])
306
- ) {
307
- changedRunMessageIds.add(message.id);
308
- }
309
- }
310
-
311
298
  const { messages } = latest;
312
299
  let lastInnerMessageId: string | null = null;
313
- const failedUpdateIds = new Set<string>();
314
300
 
315
301
  const getLastInnerId = (msgs: TMessage[]): string | null =>
316
302
  msgs.length > 0 ? storageFormatAdapter.getId(msgs.at(-1)!) : null;
@@ -343,13 +329,7 @@ export const useExternalHistory = <TMessage>(
343
329
  continue;
344
330
  }
345
331
 
346
- const isPersistedMessage = historyIds.current.has(message.id);
347
- if (isPersistedMessage && !changedRunMessageIds.has(message.id)) {
348
- lastInnerMessageId =
349
- getLastInnerId(innerMessages) ?? lastInnerMessageId;
350
- continue;
351
- }
352
- if (!isPersistedMessage) {
332
+ if (!historyIds.current.has(message.id)) {
353
333
  historyIds.current.add(message.id);
354
334
  deferredTelemetryIds.current.add(message.id);
355
335
  }
@@ -357,15 +337,31 @@ export const useExternalHistory = <TMessage>(
357
337
  const batchItems = toBatchItems(innerMessages);
358
338
  for (const item of batchItems) {
359
339
  const innerId = storageFormatAdapter.getId(item.message);
360
- if (!persistedInnerIds.current.has(innerId)) {
340
+ const persisted = persistedInnerMessages.current.get(innerId);
341
+ if (!persisted) {
361
342
  await adapter.append(item);
362
- persistedInnerIds.current.add(innerId);
363
- } else if (durationMs !== undefined) {
364
- try {
365
- await adapter.update?.(item, innerId);
366
- } catch {
367
- // A failed update drops the message from the refreshed baseline so it retries on the next run stop.
368
- failedUpdateIds.add(message.id);
343
+ persistedInnerMessages.current.set(innerId, {
344
+ source: item.message,
345
+ content: encodeContent(storageFormatAdapter, item),
346
+ });
347
+ } else if (
348
+ persisted.source !== item.message &&
349
+ durationMs !== undefined &&
350
+ adapter.update
351
+ ) {
352
+ const content = encodeContent(storageFormatAdapter, item);
353
+ if (content === persisted.content) {
354
+ persisted.source = item.message;
355
+ } else {
356
+ try {
357
+ await adapter.update(item, innerId);
358
+ persistedInnerMessages.current.set(innerId, {
359
+ source: item.message,
360
+ content,
361
+ });
362
+ } catch {
363
+ // A failed update leaves the stale baseline behind so the next run stop retries it.
364
+ }
369
365
  }
370
366
  }
371
367
  }
@@ -378,14 +374,6 @@ export const useExternalHistory = <TMessage>(
378
374
  adapter.reportTelemetry?.(batchItems, telemetryOptions);
379
375
  }
380
376
  }
381
-
382
- const nextSnapshot = snapshotExternalMessages<TMessage>(
383
- latest.messages,
384
- );
385
- for (const id of failedUpdateIds) {
386
- nextSnapshot.delete(id);
387
- }
388
- persistedExternalMessages.current = nextSnapshot;
389
377
  })
390
378
  .catch((error) => {
391
379
  console.error("Failed to persist message history:", error);
@@ -430,9 +418,8 @@ export const useExternalHistory = <TMessage>(
430
418
 
431
419
  historyIds.current.delete(messageId);
432
420
  deferredTelemetryIds.current.delete(messageId);
433
- persistedExternalMessages.current.delete(messageId);
434
421
  for (const item of itemsToDelete) {
435
- persistedInnerIds.current.delete(
422
+ persistedInnerMessages.current.delete(
436
423
  storageFormatAdapter.getId(item.message),
437
424
  );
438
425
  }
@@ -139,6 +139,30 @@ describe("AISDKToolkit", () => {
139
139
  mocks.createMCPClient.mockReset();
140
140
  });
141
141
 
142
+ it("preserves prototype-named MCP tools", async () => {
143
+ const prototypeTool = { inputSchema: {} };
144
+ mocks.tools.mockResolvedValue(
145
+ Object.fromEntries([["__proto__", prototypeTool]]),
146
+ );
147
+ mocks.createMCPClient.mockResolvedValue({
148
+ tools: mocks.tools,
149
+ close: mocks.close,
150
+ });
151
+
152
+ const toolkit = new AISDKToolkit({
153
+ toolkit: {
154
+ docs: {
155
+ type: "mcp",
156
+ server: { type: "http", url: "http://localhost:3001/mcp" },
157
+ },
158
+ },
159
+ });
160
+
161
+ const tools = await toolkit.tools();
162
+ expect(Object.hasOwn(tools, "__proto__")).toBe(true);
163
+ expect(tools["__proto__"]).toBe(prototypeTool);
164
+ });
165
+
142
166
  it("loads MCP tools through pooled clients", async () => {
143
167
  mocks.tools.mockResolvedValue({ echo: { inputSchema: {} } });
144
168
  mocks.createMCPClient.mockResolvedValue({
@@ -454,6 +478,61 @@ describe("AISDKToolkit", () => {
454
478
  }
455
479
  });
456
480
 
481
+ it("does not evict a replacement client after an older listing timeout", async () => {
482
+ vi.useFakeTimers();
483
+ const oldClient = {
484
+ tools: vi.fn(() => never()),
485
+ close: vi.fn().mockResolvedValue(undefined),
486
+ };
487
+ const replacementClient = {
488
+ tools: vi.fn().mockResolvedValue({ echo: { inputSchema: {} } }),
489
+ close: vi.fn().mockResolvedValue(undefined),
490
+ };
491
+ mocks.createMCPClient
492
+ .mockResolvedValueOnce(oldClient)
493
+ .mockResolvedValue(replacementClient);
494
+
495
+ const toolkit = new AISDKToolkit({
496
+ toolkit: {
497
+ docs: {
498
+ type: "mcp",
499
+ server: {
500
+ type: "http",
501
+ url: "http://localhost:3001/mcp",
502
+ connectionTimeout: 100,
503
+ },
504
+ },
505
+ },
506
+ });
507
+
508
+ try {
509
+ const first = toolkit.tools();
510
+ const firstRejection = expect(first).rejects.toThrow(
511
+ /timed out while listing tools/,
512
+ );
513
+ await vi.advanceTimersByTimeAsync(50);
514
+
515
+ const second = toolkit.tools();
516
+ const secondRejection = expect(second).rejects.toThrow(
517
+ /timed out while listing tools/,
518
+ );
519
+ await vi.advanceTimersByTimeAsync(50);
520
+ await firstRejection;
521
+
522
+ await expect(toolkit.tools()).resolves.toHaveProperty("echo");
523
+ expect(mocks.createMCPClient).toHaveBeenCalledTimes(2);
524
+
525
+ await vi.advanceTimersByTimeAsync(50);
526
+ await secondRejection;
527
+
528
+ await expect(toolkit.tools()).resolves.toHaveProperty("echo");
529
+ expect(mocks.createMCPClient).toHaveBeenCalledTimes(2);
530
+ expect(oldClient.close).toHaveBeenCalledTimes(1);
531
+ } finally {
532
+ vi.useRealTimers();
533
+ }
534
+ });
535
+
457
536
  it("includes the MCP toolkit entry name when listing tools fails", async () => {
458
537
  const error = new Error("list failed");
459
538
  mocks.tools.mockRejectedValue(error);
@@ -227,11 +227,8 @@ export class AISDKToolkit {
227
227
  )
228
228
  .map(async ([name, tool]) => {
229
229
  const startedAt = Date.now();
230
- const client = await this.#mcpClient(
231
- name,
232
- tool.server,
233
- startedAt,
234
- ).catch((error: unknown) => {
230
+ const clientPromise = this.#mcpClient(name, tool.server, startedAt);
231
+ const client = await clientPromise.catch((error: unknown) => {
235
232
  if (error instanceof MCPConnectionTimeoutError) throw error;
236
233
  throw toMcpToolkitError(name, "connect", error);
237
234
  });
@@ -245,8 +242,10 @@ export class AISDKToolkit {
245
242
  return [name, tool, tools] as const;
246
243
  } catch (error) {
247
244
  if (error instanceof MCPConnectionTimeoutError) {
248
- this.#mcpClients.delete(name);
249
- void client.close().catch(() => {});
245
+ if (this.#mcpClients.get(name) === clientPromise) {
246
+ this.#mcpClients.delete(name);
247
+ void client.close().catch(() => {});
248
+ }
250
249
  throw error;
251
250
  }
252
251
  throw toMcpToolkitError(name, "list tools", error);
@@ -254,7 +253,7 @@ export class AISDKToolkit {
254
253
  }),
255
254
  );
256
255
 
257
- const tools: ToolSet = {};
256
+ const tools = Object.create(null) as ToolSet;
258
257
  const toolSources = new Map<string, string>();
259
258
  for (const [serverName, mcpTool, toolSet] of toolSets) {
260
259
  for (const [toolName, tool] of Object.entries(toolSet)) {
package/src/usage.test.ts CHANGED
@@ -1,5 +1,5 @@
1
1
  import { describe, expect, it } from "vitest";
2
- import { getLatestThreadTokenUsage, getThreadMessageTokenUsage } from "./usage";
2
+ import { getThreadMessageTokenUsage } from "./usage";
3
3
 
4
4
  function msg(metadata: unknown): { role: "assistant"; metadata: unknown } {
5
5
  return {
@@ -9,6 +9,22 @@ function msg(metadata: unknown): { role: "assistant"; metadata: unknown } {
9
9
  }
10
10
 
11
11
  describe("getThreadMessageTokenUsage", () => {
12
+ it("reads usage from custom.usage", () => {
13
+ const usage = getThreadMessageTokenUsage(
14
+ msg({
15
+ custom: {
16
+ usage: { inputTokens: 4, outputTokens: 6 },
17
+ },
18
+ }),
19
+ );
20
+
21
+ expect(usage).toEqual({
22
+ totalTokens: 10,
23
+ inputTokens: 4,
24
+ outputTokens: 6,
25
+ });
26
+ });
27
+
12
28
  it("does not double-count reasoning/cached in fallback totalTokens", () => {
13
29
  const usage = getThreadMessageTokenUsage(
14
30
  msg({
@@ -157,25 +173,27 @@ describe("getThreadMessageTokenUsage", () => {
157
173
  });
158
174
  });
159
175
 
160
- describe("getLatestThreadTokenUsage", () => {
161
- it("falls back to the latest assistant message with usage", () => {
162
- const usage = getLatestThreadTokenUsage([
176
+ describe("getThreadMessageTokenUsage", () => {
177
+ it("returns token usage from an earlier assistant message", () => {
178
+ const messages = [
163
179
  { role: "assistant", metadata: { usage: { totalTokens: 100 } } },
164
180
  { role: "user", metadata: {} },
165
181
  { role: "assistant", metadata: {} },
166
- ]);
182
+ ];
183
+ const usage = getThreadMessageTokenUsage(messages[0]);
167
184
 
168
185
  expect(usage).toEqual({ totalTokens: 100 });
169
186
  });
170
187
 
171
- it("prefers the newest assistant message when it has usage", () => {
172
- const usage = getLatestThreadTokenUsage([
188
+ it("returns token usage from the newest assistant message", () => {
189
+ const messages = [
173
190
  { role: "assistant", metadata: { usage: { totalTokens: 100 } } },
174
191
  {
175
192
  role: "assistant",
176
193
  metadata: { usage: { inputTokens: 40, outputTokens: 2 } },
177
194
  },
178
- ]);
195
+ ];
196
+ const usage = getThreadMessageTokenUsage(messages[1]);
179
197
 
180
198
  expect(usage).toEqual({
181
199
  totalTokens: 42,