@assistant-ui/ai-sdk 0.0.2 → 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 (54) 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/converters/toCreateMessage.d.ts.map +1 -1
  6. package/dist/converters/toCreateMessage.js +6 -2
  7. package/dist/converters/toCreateMessage.js.map +1 -1
  8. package/dist/model-context/injectInteractableContext.d.ts +1 -1
  9. package/dist/model-context/injectInteractableContext.js +1 -1
  10. package/dist/model-context/injectInteractableContext.js.map +1 -1
  11. package/dist/model-context/injectQuoteContext.d.ts +1 -1
  12. package/dist/model-context/injectQuoteContext.js +1 -1
  13. package/dist/model-context/injectQuoteContext.js.map +1 -1
  14. package/dist/runtime/AISDKChat.js +1 -1
  15. package/dist/runtime/AISDKThreads.d.ts +10 -8
  16. package/dist/runtime/AISDKThreads.d.ts.map +1 -1
  17. package/dist/runtime/AISDKThreads.js +34 -27
  18. package/dist/runtime/AISDKThreads.js.map +1 -1
  19. package/dist/runtime/useAISDKRuntime.d.ts +18 -3
  20. package/dist/runtime/useAISDKRuntime.d.ts.map +1 -1
  21. package/dist/runtime/useAISDKRuntime.js +199 -73
  22. package/dist/runtime/useAISDKRuntime.js.map +1 -1
  23. package/dist/runtime/useChatRuntime.js +1 -1
  24. package/dist/runtime/useChatThread.d.ts +14 -1
  25. package/dist/runtime/useChatThread.d.ts.map +1 -1
  26. package/dist/runtime/useChatThread.js +11 -4
  27. package/dist/runtime/useChatThread.js.map +1 -1
  28. package/dist/runtime/useExternalHistory.d.ts.map +1 -1
  29. package/dist/runtime/useExternalHistory.js +52 -47
  30. package/dist/runtime/useExternalHistory.js.map +1 -1
  31. package/dist/runtime/useResourceCleanup.js +1 -1
  32. package/dist/usage.d.ts +1 -2
  33. package/dist/usage.d.ts.map +1 -1
  34. package/dist/usage.js +5 -7
  35. package/dist/usage.js.map +1 -1
  36. package/package.json +12 -12
  37. package/src/converters/convertMessage.test.ts +22 -0
  38. package/src/converters/convertMessage.ts +26 -2
  39. package/src/converters/toCreateMessage.ts +6 -5
  40. package/src/model-context/injectInteractableContext.ts +1 -1
  41. package/src/model-context/injectQuoteContext.ts +1 -1
  42. package/src/runtime/AISDKThreads.cloud.test.ts +11 -1
  43. package/src/runtime/AISDKThreads.test.ts +102 -0
  44. package/src/runtime/AISDKThreads.ts +16 -9
  45. package/src/runtime/useAISDKRuntime.denied-tool.test.tsx +122 -0
  46. package/src/runtime/useAISDKRuntime.test.ts +479 -10
  47. package/src/runtime/useAISDKRuntime.ts +199 -55
  48. package/src/runtime/useChatRuntime.integration.test.tsx +46 -0
  49. package/src/runtime/useChatRuntime.test.ts +16 -0
  50. package/src/runtime/useChatThread.ts +26 -0
  51. package/src/runtime/useExternalHistory.test.ts +143 -1
  52. package/src/runtime/useExternalHistory.ts +53 -56
  53. package/src/usage.test.ts +26 -8
  54. package/src/usage.ts +4 -9
@@ -392,6 +392,7 @@ describe("useExternalHistory persistence", () => {
392
392
  options?: {
393
393
  loadMessages?: MessageFormatRepository<InnerMessage>;
394
394
  toThreadMessages?: (messages: InnerMessage[]) => ThreadMessage[];
395
+ initialIsRunning?: boolean;
395
396
  },
396
397
  ) => {
397
398
  const append = vi.fn(
@@ -426,7 +427,7 @@ describe("useExternalHistory persistence", () => {
426
427
  };
427
428
 
428
429
  let listener: (() => void) | undefined;
429
- let isRunning = false;
430
+ let isRunning = options?.initialIsRunning ?? false;
430
431
  let messages: ThreadMessage[] = [];
431
432
  const getState = vi.fn(() => ({ isRunning, messages }));
432
433
  const thread = {
@@ -497,6 +498,83 @@ describe("useExternalHistory persistence", () => {
497
498
  };
498
499
  };
499
500
 
501
+ it("persists a settled turn when the history adapter becomes active after it", async () => {
502
+ let listener: (() => void) | undefined;
503
+ let isRunning = false;
504
+ let messages: ThreadMessage[] = [];
505
+ const append = vi.fn(async () => {});
506
+ const formattedAdapter = {
507
+ load: vi.fn().mockResolvedValue({ messages: [] }),
508
+ append,
509
+ };
510
+ const historyAdapter: ThreadHistoryAdapter = {
511
+ load: vi.fn(),
512
+ append: vi.fn(),
513
+ withFormat: vi.fn().mockReturnValue(formattedAdapter),
514
+ };
515
+ let activeAdapter: ThreadHistoryAdapter | undefined;
516
+ const thread = {
517
+ subscribe: (next: () => void) => {
518
+ listener = next;
519
+ return () => {};
520
+ },
521
+ getState: () => ({ isRunning, messages }),
522
+ import: vi.fn(),
523
+ export: vi.fn(() => ({ headId: null, messages: [] })),
524
+ } as unknown as AssistantRuntime["thread"];
525
+ const persistenceRuntimeRef = {
526
+ current: { thread } as AssistantRuntime,
527
+ };
528
+ const message = createAssistantMessage(
529
+ { type: "complete", reason: "stop" },
530
+ [{ id: "inner-1", parts: ["answer"] }],
531
+ );
532
+
533
+ mocks.hasThreadListItem = true;
534
+ mocks.remoteId = "remote-thread";
535
+
536
+ const { rerender } = renderHook(() =>
537
+ useExternalHistory(
538
+ persistenceRuntimeRef,
539
+ activeAdapter,
540
+ () => [],
541
+ persistenceStorageFormat,
542
+ () => {},
543
+ ),
544
+ );
545
+
546
+ await act(async () => {
547
+ messages = [message];
548
+ isRunning = true;
549
+ listener?.();
550
+ isRunning = false;
551
+ listener?.();
552
+ });
553
+
554
+ activeAdapter = historyAdapter;
555
+ await act(async () => rerender());
556
+
557
+ await waitFor(() => expect(append).toHaveBeenCalledTimes(1));
558
+ expect(append).toHaveBeenCalledWith({
559
+ parentId: null,
560
+ message: { id: "inner-1", parts: ["answer"] },
561
+ });
562
+ });
563
+
564
+ it("persists a turn that is already running when the subscription starts", async () => {
565
+ const { append, step } = createPersistenceHarness(false, {
566
+ initialIsRunning: true,
567
+ });
568
+ const message = createAssistantMessage(
569
+ { type: "complete", reason: "stop" },
570
+ [{ id: "inner-1", parts: ["answer"] }],
571
+ );
572
+
573
+ await step({ isRunning: false, messages: [message] });
574
+
575
+ await waitFor(() => expect(append).toHaveBeenCalledTimes(1));
576
+ });
577
+
500
578
  it("retries a failed append on the next persistence pass", async () => {
501
579
  const consoleError = vi
502
580
  .spyOn(console, "error")
@@ -1102,6 +1180,70 @@ describe("useExternalHistory persistence", () => {
1102
1180
  expect(reportTelemetry).not.toHaveBeenCalled();
1103
1181
  });
1104
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
+
1105
1247
  it("absorbs agentic flickers without losing change detection", async () => {
1106
1248
  const { append, update, reportTelemetry, runCycle, flush, step } =
1107
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);
@@ -234,6 +232,19 @@ export const useExternalHistory = <TMessage>(
234
232
  }, 0);
235
233
  });
236
234
 
235
+ const initialThreadState = runtimeRef.current.thread.getState();
236
+ wasRunningRef.current = initialThreadState.isRunning;
237
+ if (initialThreadState.isRunning) {
238
+ if (runStartRef.current == null) {
239
+ runStartRef.current = Date.now();
240
+ stepBoundariesRef.current = [];
241
+ toolCallCountRef.current = 0;
242
+ adapter.pin?.();
243
+ }
244
+ } else if (initialThreadState.messages.length > 0) {
245
+ persistSettled(true);
246
+ }
247
+
237
248
  function persistSettled(ignoreRunning: boolean) {
238
249
  persistTimerRef.current = null;
239
250
  const latest = runtimeRef.current.thread.getState();
@@ -281,23 +292,8 @@ export const useExternalHistory = <TMessage>(
281
292
 
282
293
  persistInFlightRef.current = persistInFlightRef.current
283
294
  .then(async () => {
284
- const changedRunMessageIds = new Set<string>();
285
- for (const message of latest.messages) {
286
- const externalMessages =
287
- getExternalStoreMessages<TMessage>(message);
288
- const previous = persistedExternalMessages.current.get(message.id);
289
- if (
290
- previous === undefined ||
291
- previous.length !== externalMessages.length ||
292
- externalMessages.some((item, index) => item !== previous[index])
293
- ) {
294
- changedRunMessageIds.add(message.id);
295
- }
296
- }
297
-
298
295
  const { messages } = latest;
299
296
  let lastInnerMessageId: string | null = null;
300
- const failedUpdateIds = new Set<string>();
301
297
 
302
298
  const getLastInnerId = (msgs: TMessage[]): string | null =>
303
299
  msgs.length > 0 ? storageFormatAdapter.getId(msgs.at(-1)!) : null;
@@ -330,13 +326,7 @@ export const useExternalHistory = <TMessage>(
330
326
  continue;
331
327
  }
332
328
 
333
- const isPersistedMessage = historyIds.current.has(message.id);
334
- if (isPersistedMessage && !changedRunMessageIds.has(message.id)) {
335
- lastInnerMessageId =
336
- getLastInnerId(innerMessages) ?? lastInnerMessageId;
337
- continue;
338
- }
339
- if (!isPersistedMessage) {
329
+ if (!historyIds.current.has(message.id)) {
340
330
  historyIds.current.add(message.id);
341
331
  deferredTelemetryIds.current.add(message.id);
342
332
  }
@@ -344,15 +334,31 @@ export const useExternalHistory = <TMessage>(
344
334
  const batchItems = toBatchItems(innerMessages);
345
335
  for (const item of batchItems) {
346
336
  const innerId = storageFormatAdapter.getId(item.message);
347
- if (!persistedInnerIds.current.has(innerId)) {
337
+ const persisted = persistedInnerMessages.current.get(innerId);
338
+ if (!persisted) {
348
339
  await adapter.append(item);
349
- persistedInnerIds.current.add(innerId);
350
- } else if (durationMs !== undefined) {
351
- try {
352
- await adapter.update?.(item, innerId);
353
- } catch {
354
- // A failed update drops the message from the refreshed baseline so it retries on the next run stop.
355
- 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
+ }
356
362
  }
357
363
  }
358
364
  }
@@ -365,14 +371,6 @@ export const useExternalHistory = <TMessage>(
365
371
  adapter.reportTelemetry?.(batchItems, telemetryOptions);
366
372
  }
367
373
  }
368
-
369
- const nextSnapshot = snapshotExternalMessages<TMessage>(
370
- latest.messages,
371
- );
372
- for (const id of failedUpdateIds) {
373
- nextSnapshot.delete(id);
374
- }
375
- persistedExternalMessages.current = nextSnapshot;
376
374
  })
377
375
  .catch((error) => {
378
376
  console.error("Failed to persist message history:", error);
@@ -417,9 +415,8 @@ export const useExternalHistory = <TMessage>(
417
415
 
418
416
  historyIds.current.delete(messageId);
419
417
  deferredTelemetryIds.current.delete(messageId);
420
- persistedExternalMessages.current.delete(messageId);
421
418
  for (const item of itemsToDelete) {
422
- persistedInnerIds.current.delete(
419
+ persistedInnerMessages.current.delete(
423
420
  storageFormatAdapter.getId(item.message),
424
421
  );
425
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
  }