@assistant-ui/ai-sdk 0.0.3 → 0.0.4

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 (40) hide show
  1. package/README.md +1 -1
  2. package/dist/converters/convertMessage.d.ts.map +1 -1
  3. package/dist/converters/convertMessage.js +23 -2
  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.js +1 -1
  15. package/dist/runtime/useChatThread.d.ts +7 -0
  16. package/dist/runtime/useChatThread.d.ts.map +1 -1
  17. package/dist/runtime/useChatThread.js +2 -1
  18. package/dist/runtime/useChatThread.js.map +1 -1
  19. package/dist/runtime/useExternalHistory.d.ts.map +1 -1
  20. package/dist/runtime/useExternalHistory.js +42 -47
  21. package/dist/runtime/useExternalHistory.js.map +1 -1
  22. package/dist/runtime/useResourceCleanup.js +1 -1
  23. package/dist/usage.d.ts +1 -2
  24. package/dist/usage.d.ts.map +1 -1
  25. package/dist/usage.js +5 -7
  26. package/dist/usage.js.map +1 -1
  27. package/package.json +10 -10
  28. package/src/converters/convertMessage.test.ts +22 -0
  29. package/src/converters/convertMessage.ts +26 -2
  30. package/src/runtime/AISDKThreads.cloud.test.ts +11 -1
  31. package/src/runtime/AISDKThreads.test.ts +102 -0
  32. package/src/runtime/AISDKThreads.ts +16 -9
  33. package/src/runtime/useAISDKRuntime.test.ts +288 -3
  34. package/src/runtime/useAISDKRuntime.ts +90 -5
  35. package/src/runtime/useChatRuntime.integration.test.tsx +46 -0
  36. package/src/runtime/useChatThread.ts +11 -0
  37. package/src/runtime/useExternalHistory.test.ts +64 -0
  38. package/src/runtime/useExternalHistory.ts +40 -56
  39. package/src/usage.test.ts +26 -8
  40. package/src/usage.ts +4 -9
@@ -8,6 +8,7 @@ import { StrictMode, useState } from "react";
8
8
  import { describe, expect, it } from "vitest";
9
9
  import { AssistantChatTransport } from "../transport/AssistantChatTransport";
10
10
  import { useChatRuntime } from "./useChatRuntime";
11
+ import { useThreadTokenUsage } from "../usage";
11
12
 
12
13
  const messages: UIMessage[] = [
13
14
  {
@@ -65,3 +66,48 @@ describe("useChatRuntime integration", () => {
65
66
  });
66
67
  });
67
68
  });
69
+
70
+ const UsageProbe = () => {
71
+ const usage = useThreadTokenUsage();
72
+ return (
73
+ <output data-testid="total-tokens">{usage?.totalTokens ?? "none"}</output>
74
+ );
75
+ };
76
+
77
+ const UsageApp = () => {
78
+ const [transport] = useState(
79
+ () => new AssistantChatTransport({ api: "/api/chat" }),
80
+ );
81
+ const runtime = useChatRuntime({
82
+ messages: [
83
+ ...messages,
84
+ {
85
+ id: "assistant-with-usage",
86
+ role: "assistant",
87
+ parts: [{ type: "text", text: "Hi" }],
88
+ metadata: { usage: { inputTokens: 40, outputTokens: 2 } },
89
+ },
90
+ ],
91
+ transport,
92
+ });
93
+
94
+ return (
95
+ <AssistantRuntimeProvider runtime={runtime}>
96
+ <UsageProbe />
97
+ </AssistantRuntimeProvider>
98
+ );
99
+ };
100
+
101
+ describe("useThreadTokenUsage through useChatRuntime", () => {
102
+ it("reads usage from the message metadata a server attached", async () => {
103
+ render(
104
+ <StrictMode>
105
+ <UsageApp />
106
+ </StrictMode>,
107
+ );
108
+
109
+ await waitFor(() => {
110
+ expect(screen.getByTestId("total-tokens").textContent).toBe("42");
111
+ });
112
+ });
113
+ });
@@ -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);
@@ -294,23 +292,8 @@ export const useExternalHistory = <TMessage>(
294
292
 
295
293
  persistInFlightRef.current = persistInFlightRef.current
296
294
  .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
295
  const { messages } = latest;
312
296
  let lastInnerMessageId: string | null = null;
313
- const failedUpdateIds = new Set<string>();
314
297
 
315
298
  const getLastInnerId = (msgs: TMessage[]): string | null =>
316
299
  msgs.length > 0 ? storageFormatAdapter.getId(msgs.at(-1)!) : null;
@@ -343,13 +326,7 @@ export const useExternalHistory = <TMessage>(
343
326
  continue;
344
327
  }
345
328
 
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) {
329
+ if (!historyIds.current.has(message.id)) {
353
330
  historyIds.current.add(message.id);
354
331
  deferredTelemetryIds.current.add(message.id);
355
332
  }
@@ -357,15 +334,31 @@ export const useExternalHistory = <TMessage>(
357
334
  const batchItems = toBatchItems(innerMessages);
358
335
  for (const item of batchItems) {
359
336
  const innerId = storageFormatAdapter.getId(item.message);
360
- if (!persistedInnerIds.current.has(innerId)) {
337
+ const persisted = persistedInnerMessages.current.get(innerId);
338
+ if (!persisted) {
361
339
  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);
340
+ persistedInnerMessages.current.set(innerId, {
341
+ source: item.message,
342
+ content: encodeContent(storageFormatAdapter, item),
343
+ });
344
+ } else if (
345
+ persisted.source !== item.message &&
346
+ durationMs !== undefined &&
347
+ adapter.update
348
+ ) {
349
+ const content = encodeContent(storageFormatAdapter, item);
350
+ if (content === persisted.content) {
351
+ persisted.source = item.message;
352
+ } else {
353
+ try {
354
+ await adapter.update(item, innerId);
355
+ persistedInnerMessages.current.set(innerId, {
356
+ source: item.message,
357
+ content,
358
+ });
359
+ } catch {
360
+ // A failed update leaves the stale baseline behind so the next run stop retries it.
361
+ }
369
362
  }
370
363
  }
371
364
  }
@@ -378,14 +371,6 @@ export const useExternalHistory = <TMessage>(
378
371
  adapter.reportTelemetry?.(batchItems, telemetryOptions);
379
372
  }
380
373
  }
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
374
  })
390
375
  .catch((error) => {
391
376
  console.error("Failed to persist message history:", error);
@@ -430,9 +415,8 @@ export const useExternalHistory = <TMessage>(
430
415
 
431
416
  historyIds.current.delete(messageId);
432
417
  deferredTelemetryIds.current.delete(messageId);
433
- persistedExternalMessages.current.delete(messageId);
434
418
  for (const item of itemsToDelete) {
435
- persistedInnerIds.current.delete(
419
+ persistedInnerMessages.current.delete(
436
420
  storageFormatAdapter.getId(item.message),
437
421
  );
438
422
  }
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,
package/src/usage.ts CHANGED
@@ -1,4 +1,5 @@
1
1
  /// <reference types="@assistant-ui/core/react" />
2
+ import { useMemo } from "react";
2
3
  import { useAuiState } from "@assistant-ui/store";
3
4
 
4
5
  export type ThreadTokenUsage = {
@@ -141,18 +142,12 @@ export function getThreadMessageTokenUsage(
141
142
  const topLevelUsage = normalizeUsage(metadata.usage);
142
143
  if (topLevelUsage) return withComputedTotal(topLevelUsage);
143
144
 
144
- const legacyUsage = normalizeUsage(asRecord(metadata.custom)?.usage);
145
- if (legacyUsage) return withComputedTotal(legacyUsage);
145
+ const customUsage = normalizeUsage(asRecord(metadata.custom)?.usage);
146
+ if (customUsage) return withComputedTotal(customUsage);
146
147
 
147
148
  return usageFromSteps(metadata.steps);
148
149
  }
149
150
 
150
- export function getLatestThreadTokenUsage(
151
- messages: readonly TokenUsageExtractableMessage[] | undefined,
152
- ): ThreadTokenUsage | undefined {
153
- return getThreadMessageTokenUsage(findLatestMessageWithUsage(messages));
154
- }
155
-
156
151
  function findLatestMessageWithUsage(
157
152
  messages: readonly TokenUsageExtractableMessage[] | undefined,
158
153
  ): TokenUsageExtractableMessage | undefined {
@@ -170,5 +165,5 @@ function findLatestMessageWithUsage(
170
165
 
171
166
  export function useThreadTokenUsage(): ThreadTokenUsage | undefined {
172
167
  const msg = useAuiState((s) => findLatestMessageWithUsage(s.thread.messages));
173
- return getThreadMessageTokenUsage(msg);
168
+ return useMemo(() => getThreadMessageTokenUsage(msg), [msg]);
174
169
  }