@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.
- package/README.md +1 -1
- package/dist/converters/convertMessage.d.ts.map +1 -1
- package/dist/converters/convertMessage.js +31 -6
- package/dist/converters/convertMessage.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 +2 -0
- package/dist/runtime/useAISDKRuntime.d.ts.map +1 -1
- package/dist/runtime/useAISDKRuntime.js +171 -68
- package/dist/runtime/useAISDKRuntime.js.map +1 -1
- package/dist/runtime/useChatRuntime.d.ts.map +1 -1
- package/dist/runtime/useChatRuntime.js +3 -2
- package/dist/runtime/useChatRuntime.js.map +1 -1
- package/dist/runtime/useChatThread.d.ts +7 -0
- package/dist/runtime/useChatThread.d.ts.map +1 -1
- package/dist/runtime/useChatThread.js +2 -1
- package/dist/runtime/useChatThread.js.map +1 -1
- package/dist/runtime/useExternalHistory.d.ts.map +1 -1
- package/dist/runtime/useExternalHistory.js +42 -47
- package/dist/runtime/useExternalHistory.js.map +1 -1
- package/dist/runtime/useResourceCleanup.js +1 -1
- package/dist/tools/generativeTools.js +7 -4
- package/dist/tools/generativeTools.js.map +1 -1
- package/dist/usage.d.ts +9 -2
- package/dist/usage.d.ts.map +1 -1
- package/dist/usage.js +13 -7
- package/dist/usage.js.map +1 -1
- package/package.json +11 -11
- package/src/converters/convertMessage.test.ts +124 -0
- package/src/converters/convertMessage.ts +64 -20
- package/src/runtime/AISDKChat.integration.test.tsx +48 -27
- 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/__tests__/controlled-transport.ts +21 -0
- package/src/runtime/useAISDKRuntime.test.ts +288 -3
- package/src/runtime/useAISDKRuntime.ts +90 -5
- package/src/runtime/useChatRuntime.integration.test.tsx +183 -4
- package/src/runtime/useChatRuntime.ts +1 -0
- package/src/runtime/useChatThread.ts +11 -0
- package/src/runtime/useExternalHistory.test.ts +64 -0
- package/src/runtime/useExternalHistory.ts +43 -56
- package/src/tools/generativeTools.test.ts +79 -0
- package/src/tools/generativeTools.ts +7 -8
- package/src/usage.test.ts +26 -8
- 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 {
|
|
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";
|
|
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
|
});
|
|
@@ -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
|
|
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);
|
|
@@ -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
|
-
|
|
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
|
-
|
|
340
|
+
const persisted = persistedInnerMessages.current.get(innerId);
|
|
341
|
+
if (!persisted) {
|
|
361
342
|
await adapter.append(item);
|
|
362
|
-
|
|
363
|
-
|
|
364
|
-
|
|
365
|
-
|
|
366
|
-
|
|
367
|
-
|
|
368
|
-
|
|
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
|
-
|
|
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
|
|
231
|
-
|
|
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.
|
|
249
|
-
|
|
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
|
|
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 {
|
|
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,
|