@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.
- package/README.md +1 -1
- package/dist/converters/convertMessage.d.ts.map +1 -1
- package/dist/converters/convertMessage.js +23 -2
- package/dist/converters/convertMessage.js.map +1 -1
- package/dist/converters/toCreateMessage.d.ts.map +1 -1
- package/dist/converters/toCreateMessage.js +6 -2
- package/dist/converters/toCreateMessage.js.map +1 -1
- package/dist/model-context/injectInteractableContext.d.ts +1 -1
- package/dist/model-context/injectInteractableContext.js +1 -1
- package/dist/model-context/injectInteractableContext.js.map +1 -1
- package/dist/model-context/injectQuoteContext.d.ts +1 -1
- package/dist/model-context/injectQuoteContext.js +1 -1
- package/dist/model-context/injectQuoteContext.js.map +1 -1
- package/dist/runtime/AISDKChat.js +1 -1
- package/dist/runtime/AISDKThreads.d.ts +10 -8
- package/dist/runtime/AISDKThreads.d.ts.map +1 -1
- package/dist/runtime/AISDKThreads.js +34 -27
- package/dist/runtime/AISDKThreads.js.map +1 -1
- package/dist/runtime/useAISDKRuntime.d.ts +18 -3
- package/dist/runtime/useAISDKRuntime.d.ts.map +1 -1
- package/dist/runtime/useAISDKRuntime.js +199 -73
- package/dist/runtime/useAISDKRuntime.js.map +1 -1
- package/dist/runtime/useChatRuntime.js +1 -1
- package/dist/runtime/useChatThread.d.ts +14 -1
- package/dist/runtime/useChatThread.d.ts.map +1 -1
- package/dist/runtime/useChatThread.js +11 -4
- package/dist/runtime/useChatThread.js.map +1 -1
- package/dist/runtime/useExternalHistory.d.ts.map +1 -1
- package/dist/runtime/useExternalHistory.js +52 -47
- package/dist/runtime/useExternalHistory.js.map +1 -1
- package/dist/runtime/useResourceCleanup.js +1 -1
- package/dist/usage.d.ts +1 -2
- package/dist/usage.d.ts.map +1 -1
- package/dist/usage.js +5 -7
- package/dist/usage.js.map +1 -1
- package/package.json +12 -12
- package/src/converters/convertMessage.test.ts +22 -0
- package/src/converters/convertMessage.ts +26 -2
- package/src/converters/toCreateMessage.ts +6 -5
- package/src/model-context/injectInteractableContext.ts +1 -1
- package/src/model-context/injectQuoteContext.ts +1 -1
- package/src/runtime/AISDKThreads.cloud.test.ts +11 -1
- package/src/runtime/AISDKThreads.test.ts +102 -0
- package/src/runtime/AISDKThreads.ts +16 -9
- package/src/runtime/useAISDKRuntime.denied-tool.test.tsx +122 -0
- package/src/runtime/useAISDKRuntime.test.ts +479 -10
- package/src/runtime/useAISDKRuntime.ts +199 -55
- package/src/runtime/useChatRuntime.integration.test.tsx +46 -0
- package/src/runtime/useChatRuntime.test.ts +16 -0
- package/src/runtime/useChatThread.ts +26 -0
- package/src/runtime/useExternalHistory.test.ts +143 -1
- package/src/runtime/useExternalHistory.ts +53 -56
- package/src/usage.test.ts +26 -8
- 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
|
|
53
|
-
|
|
54
|
-
|
|
55
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
|
|
337
|
+
const persisted = persistedInnerMessages.current.get(innerId);
|
|
338
|
+
if (!persisted) {
|
|
348
339
|
await adapter.append(item);
|
|
349
|
-
|
|
350
|
-
|
|
351
|
-
|
|
352
|
-
|
|
353
|
-
|
|
354
|
-
|
|
355
|
-
|
|
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
|
-
|
|
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 {
|
|
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("
|
|
161
|
-
it("
|
|
162
|
-
const
|
|
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("
|
|
172
|
-
const
|
|
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
|
|
145
|
-
if (
|
|
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
|
}
|