@assistant-ui/react-google-adk 0.0.33 → 0.0.35
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/dist/AdkEventAccumulator.d.ts.map +1 -1
- package/dist/AdkEventAccumulator.js +4 -1
- package/dist/AdkEventAccumulator.js.map +1 -1
- package/dist/adkToolApproval.js +1 -1
- package/dist/adkToolApproval.js.map +1 -1
- package/dist/convertAdkMessages.d.ts.map +1 -1
- package/dist/convertAdkMessages.js +30 -19
- package/dist/convertAdkMessages.js.map +1 -1
- package/dist/convertToAdkMessages.d.ts.map +1 -1
- package/dist/convertToAdkMessages.js +2 -2
- package/dist/convertToAdkMessages.js.map +1 -1
- package/dist/sdkIdentity.js +1 -1
- package/dist/useAdkMessages.d.ts +23 -0
- package/dist/useAdkMessages.d.ts.map +1 -1
- package/dist/useAdkMessages.js +34 -9
- package/dist/useAdkMessages.js.map +1 -1
- package/dist/useAdkRuntime.d.ts.map +1 -1
- package/dist/useAdkRuntime.js +60 -9
- package/dist/useAdkRuntime.js.map +1 -1
- package/package.json +7 -6
- package/src/AdkEventAccumulator.test.ts +32 -0
- package/src/AdkEventAccumulator.ts +1 -0
- package/src/adkToolApproval.test.ts +19 -0
- package/src/adkToolApproval.ts +1 -1
- package/src/convertAdkMessages.test.ts +129 -0
- package/src/convertAdkMessages.ts +19 -1
- package/src/convertToAdkMessages.test.ts +14 -0
- package/src/convertToAdkMessages.ts +4 -1
- package/src/useAdkMessages.test.ts +55 -0
- package/src/useAdkMessages.ts +65 -6
- package/src/useAdkRuntime.cancellation.test.tsx +544 -0
- package/src/useAdkRuntime.fast-refresh.test.tsx +166 -0
- package/src/useAdkRuntime.refetch.test.tsx +1 -0
- package/src/useAdkRuntime.ts +91 -9
- package/src/useAdkRuntimeApproval.test.tsx +305 -36
|
@@ -0,0 +1,544 @@
|
|
|
1
|
+
// @vitest-environment jsdom
|
|
2
|
+
|
|
3
|
+
import { act, render, waitFor } from "@testing-library/react";
|
|
4
|
+
import { type FC } from "react";
|
|
5
|
+
import { describe, expect, it, vi } from "vitest";
|
|
6
|
+
import { AssistantRuntimeProvider } from "@assistant-ui/core/react";
|
|
7
|
+
import type {
|
|
8
|
+
AssistantRuntime,
|
|
9
|
+
RemoteThreadListAdapter,
|
|
10
|
+
} from "@assistant-ui/core";
|
|
11
|
+
import { useAdkRuntime } from "./useAdkRuntime";
|
|
12
|
+
import { useAdkLongRunningToolIds, useAdkSend } from "./hooks";
|
|
13
|
+
import { AdkEventAccumulator } from "./AdkEventAccumulator";
|
|
14
|
+
import type { AdkEvent, AdkStreamCallback, AdkThreadSnapshot } from "./types";
|
|
15
|
+
|
|
16
|
+
const makeThreadListAdapter = (): RemoteThreadListAdapter => ({
|
|
17
|
+
list: vi.fn(async () => ({
|
|
18
|
+
threads: [
|
|
19
|
+
{
|
|
20
|
+
status: "regular" as const,
|
|
21
|
+
remoteId: "adk-1",
|
|
22
|
+
externalId: "adk-1",
|
|
23
|
+
title: "ADK session",
|
|
24
|
+
},
|
|
25
|
+
],
|
|
26
|
+
})),
|
|
27
|
+
initialize: vi.fn(async () => ({
|
|
28
|
+
remoteId: "adk-1",
|
|
29
|
+
externalId: "adk-1",
|
|
30
|
+
})),
|
|
31
|
+
rename: vi.fn(async () => {}),
|
|
32
|
+
archive: vi.fn(async () => {}),
|
|
33
|
+
unarchive: vi.fn(async () => {}),
|
|
34
|
+
delete: vi.fn(async () => {}),
|
|
35
|
+
generateTitle: vi.fn(async () => new ReadableStream() as never),
|
|
36
|
+
fetch: vi.fn(async () => ({
|
|
37
|
+
status: "regular" as const,
|
|
38
|
+
remoteId: "adk-1",
|
|
39
|
+
externalId: "adk-1",
|
|
40
|
+
title: "ADK session",
|
|
41
|
+
})),
|
|
42
|
+
});
|
|
43
|
+
|
|
44
|
+
const makePendingStream = (...events: AdkEvent[]): AdkStreamCallback =>
|
|
45
|
+
async function* (_messages, { abortSignal }) {
|
|
46
|
+
yield* events;
|
|
47
|
+
await new Promise<void>((resolve) => {
|
|
48
|
+
abortSignal.addEventListener("abort", () => resolve(), { once: true });
|
|
49
|
+
});
|
|
50
|
+
};
|
|
51
|
+
|
|
52
|
+
const renderAdkRuntime = async (
|
|
53
|
+
stream: AdkStreamCallback,
|
|
54
|
+
snapshot?: AdkThreadSnapshot,
|
|
55
|
+
) => {
|
|
56
|
+
const capture: {
|
|
57
|
+
runtime: AssistantRuntime | null;
|
|
58
|
+
longRunningToolIds: string[];
|
|
59
|
+
send: ReturnType<typeof useAdkSend> | null;
|
|
60
|
+
} = { runtime: null, longRunningToolIds: [], send: null };
|
|
61
|
+
const sessionAdapter = makeThreadListAdapter();
|
|
62
|
+
|
|
63
|
+
const CaptureExtras: FC = () => {
|
|
64
|
+
capture.longRunningToolIds = useAdkLongRunningToolIds();
|
|
65
|
+
capture.send = useAdkSend();
|
|
66
|
+
return null;
|
|
67
|
+
};
|
|
68
|
+
|
|
69
|
+
const Inner: FC = () => {
|
|
70
|
+
const runtime = useAdkRuntime({
|
|
71
|
+
stream,
|
|
72
|
+
sessionAdapter,
|
|
73
|
+
unstable_allowCancellation: true,
|
|
74
|
+
...(snapshot && { load: async () => snapshot }),
|
|
75
|
+
});
|
|
76
|
+
capture.runtime = runtime;
|
|
77
|
+
return (
|
|
78
|
+
<AssistantRuntimeProvider runtime={runtime}>
|
|
79
|
+
<CaptureExtras />
|
|
80
|
+
</AssistantRuntimeProvider>
|
|
81
|
+
);
|
|
82
|
+
};
|
|
83
|
+
|
|
84
|
+
await act(async () => {
|
|
85
|
+
render(<Inner />);
|
|
86
|
+
});
|
|
87
|
+
await waitFor(() => expect(capture.runtime).not.toBeNull());
|
|
88
|
+
await act(async () => {
|
|
89
|
+
await capture.runtime!.threads.switchToThread("adk-1");
|
|
90
|
+
});
|
|
91
|
+
|
|
92
|
+
return capture;
|
|
93
|
+
};
|
|
94
|
+
|
|
95
|
+
describe("useAdkRuntime cancellation", () => {
|
|
96
|
+
it.each(["none", "send", "stream"] as const)(
|
|
97
|
+
"preserves inherited confirmation IDs on Stop (answer: %s)",
|
|
98
|
+
async (answer) => {
|
|
99
|
+
const seed = new AdkEventAccumulator();
|
|
100
|
+
seed.processEvent({
|
|
101
|
+
id: "older-confirmations",
|
|
102
|
+
author: "agent",
|
|
103
|
+
longRunningToolIds: ["old-1", "old-2"],
|
|
104
|
+
content: {
|
|
105
|
+
role: "model",
|
|
106
|
+
parts: ["old-1", "old-2"].map((id) => ({
|
|
107
|
+
functionCall: {
|
|
108
|
+
id,
|
|
109
|
+
name: "adk_request_confirmation",
|
|
110
|
+
args: {
|
|
111
|
+
originalFunctionCall: {
|
|
112
|
+
id: `target-${id}`,
|
|
113
|
+
name: "lookup",
|
|
114
|
+
args: {},
|
|
115
|
+
},
|
|
116
|
+
toolConfirmation: { hint: "Allow lookup?" },
|
|
117
|
+
},
|
|
118
|
+
},
|
|
119
|
+
})),
|
|
120
|
+
},
|
|
121
|
+
});
|
|
122
|
+
const stream = vi.fn(
|
|
123
|
+
makePendingStream(
|
|
124
|
+
...(answer === "stream"
|
|
125
|
+
? [
|
|
126
|
+
{
|
|
127
|
+
id: "settled-confirmation",
|
|
128
|
+
author: "user",
|
|
129
|
+
content: {
|
|
130
|
+
role: "user",
|
|
131
|
+
parts: [
|
|
132
|
+
{
|
|
133
|
+
functionResponse: {
|
|
134
|
+
id: "old-1",
|
|
135
|
+
name: "adk_request_confirmation",
|
|
136
|
+
response: { confirmed: true },
|
|
137
|
+
},
|
|
138
|
+
},
|
|
139
|
+
],
|
|
140
|
+
},
|
|
141
|
+
} satisfies AdkEvent,
|
|
142
|
+
]
|
|
143
|
+
: []),
|
|
144
|
+
{
|
|
145
|
+
id: "new-tool-call",
|
|
146
|
+
author: "agent",
|
|
147
|
+
longRunningToolIds: ["new-call"],
|
|
148
|
+
content: {
|
|
149
|
+
role: "model",
|
|
150
|
+
parts: [
|
|
151
|
+
{
|
|
152
|
+
functionCall: {
|
|
153
|
+
id: "new-call",
|
|
154
|
+
name: "lookup",
|
|
155
|
+
args: {},
|
|
156
|
+
},
|
|
157
|
+
},
|
|
158
|
+
],
|
|
159
|
+
},
|
|
160
|
+
},
|
|
161
|
+
),
|
|
162
|
+
);
|
|
163
|
+
const capture = await renderAdkRuntime(stream, {
|
|
164
|
+
messages: seed.getMessages(),
|
|
165
|
+
longRunningToolIds: seed.getLongRunningToolIds(),
|
|
166
|
+
toolConfirmations: seed.getToolConfirmations(),
|
|
167
|
+
});
|
|
168
|
+
await waitFor(() =>
|
|
169
|
+
expect(capture.longRunningToolIds).toEqual(["old-1", "old-2"]),
|
|
170
|
+
);
|
|
171
|
+
|
|
172
|
+
act(() => {
|
|
173
|
+
if (answer === "send") {
|
|
174
|
+
void capture.send!(
|
|
175
|
+
[
|
|
176
|
+
{
|
|
177
|
+
id: "approval-reply",
|
|
178
|
+
type: "tool",
|
|
179
|
+
tool_call_id: "old-1",
|
|
180
|
+
name: "adk_request_confirmation",
|
|
181
|
+
content: '{"confirmed":true}',
|
|
182
|
+
},
|
|
183
|
+
],
|
|
184
|
+
{},
|
|
185
|
+
);
|
|
186
|
+
} else {
|
|
187
|
+
capture.runtime!.thread.append({
|
|
188
|
+
role: "user",
|
|
189
|
+
content: [{ type: "text", text: "continue" }],
|
|
190
|
+
});
|
|
191
|
+
}
|
|
192
|
+
});
|
|
193
|
+
const pendingInherited =
|
|
194
|
+
answer === "none" ? ["old-1", "old-2"] : ["old-2"];
|
|
195
|
+
await waitFor(() =>
|
|
196
|
+
expect(capture.longRunningToolIds).toEqual([
|
|
197
|
+
...pendingInherited,
|
|
198
|
+
"new-call",
|
|
199
|
+
]),
|
|
200
|
+
);
|
|
201
|
+
|
|
202
|
+
act(() => {
|
|
203
|
+
capture.runtime!.thread.cancelRun();
|
|
204
|
+
});
|
|
205
|
+
await waitFor(() =>
|
|
206
|
+
expect(capture.runtime!.thread.getState().isRunning).toBe(false),
|
|
207
|
+
);
|
|
208
|
+
expect(capture.longRunningToolIds).toEqual(pendingInherited);
|
|
209
|
+
|
|
210
|
+
act(() => {
|
|
211
|
+
capture.runtime!.thread.append({
|
|
212
|
+
role: "user",
|
|
213
|
+
content: [{ type: "text", text: "next turn" }],
|
|
214
|
+
});
|
|
215
|
+
});
|
|
216
|
+
await waitFor(() => expect(stream).toHaveBeenCalledTimes(2));
|
|
217
|
+
const cancellations = stream.mock.calls[1]![0].filter(
|
|
218
|
+
(m) => m.type === "tool",
|
|
219
|
+
);
|
|
220
|
+
expect(cancellations).toMatchObject([
|
|
221
|
+
{ tool_call_id: "new-call", content: '{"cancelled":true}' },
|
|
222
|
+
]);
|
|
223
|
+
},
|
|
224
|
+
);
|
|
225
|
+
|
|
226
|
+
it("preserves a user message staged during a cancelled run", async () => {
|
|
227
|
+
const capture = await renderAdkRuntime(
|
|
228
|
+
makePendingStream({
|
|
229
|
+
id: "partial",
|
|
230
|
+
author: "agent",
|
|
231
|
+
partial: true,
|
|
232
|
+
content: { role: "model", parts: [{ text: "partial answer" }] },
|
|
233
|
+
}),
|
|
234
|
+
);
|
|
235
|
+
|
|
236
|
+
act(() => {
|
|
237
|
+
capture.runtime!.thread.append({
|
|
238
|
+
role: "user",
|
|
239
|
+
content: [{ type: "text", text: "hello" }],
|
|
240
|
+
});
|
|
241
|
+
});
|
|
242
|
+
await waitFor(() =>
|
|
243
|
+
expect(capture.runtime!.thread.getState().messages.at(-1)).toMatchObject({
|
|
244
|
+
role: "assistant",
|
|
245
|
+
content: [{ type: "text", text: "partial answer" }],
|
|
246
|
+
}),
|
|
247
|
+
);
|
|
248
|
+
act(() => {
|
|
249
|
+
capture.runtime!.thread.append({
|
|
250
|
+
role: "user",
|
|
251
|
+
content: [{ type: "text", text: "staged follow-up" }],
|
|
252
|
+
startRun: false,
|
|
253
|
+
});
|
|
254
|
+
});
|
|
255
|
+
await waitFor(() =>
|
|
256
|
+
expect(
|
|
257
|
+
capture
|
|
258
|
+
.runtime!.thread.getState()
|
|
259
|
+
.messages.filter((m) => m.role === "user")
|
|
260
|
+
.at(-1),
|
|
261
|
+
).toMatchObject({
|
|
262
|
+
content: [{ type: "text", text: "staged follow-up" }],
|
|
263
|
+
}),
|
|
264
|
+
);
|
|
265
|
+
const stagedMessage = capture
|
|
266
|
+
.runtime!.thread.getState()
|
|
267
|
+
.messages.filter((m) => m.role === "user")
|
|
268
|
+
.at(-1);
|
|
269
|
+
expect(capture.runtime!.thread.getState().isRunning).toBe(true);
|
|
270
|
+
|
|
271
|
+
act(() => {
|
|
272
|
+
capture.runtime!.thread.cancelRun();
|
|
273
|
+
});
|
|
274
|
+
await waitFor(() =>
|
|
275
|
+
expect(capture.runtime!.thread.getState().isRunning).toBe(false),
|
|
276
|
+
);
|
|
277
|
+
expect(
|
|
278
|
+
capture
|
|
279
|
+
.runtime!.thread.getState()
|
|
280
|
+
.messages.find((m) => m.id === stagedMessage!.id),
|
|
281
|
+
).toEqual(stagedMessage);
|
|
282
|
+
expect(
|
|
283
|
+
capture
|
|
284
|
+
.runtime!.thread.getState()
|
|
285
|
+
.messages.find((m) => m.role === "assistant")?.status,
|
|
286
|
+
).toEqual({ type: "incomplete", reason: "cancelled" });
|
|
287
|
+
});
|
|
288
|
+
|
|
289
|
+
it.each([false, true])(
|
|
290
|
+
"cancels a finalized tool-call message (long-running: %s)",
|
|
291
|
+
async (longRunning) => {
|
|
292
|
+
const stream = vi.fn(
|
|
293
|
+
makePendingStream({
|
|
294
|
+
id: "tool-call",
|
|
295
|
+
author: "agent",
|
|
296
|
+
...(longRunning && { longRunningToolIds: ["call-1"] }),
|
|
297
|
+
content: {
|
|
298
|
+
role: "model",
|
|
299
|
+
parts: [
|
|
300
|
+
{ functionCall: { id: "call-1", name: "lookup", args: {} } },
|
|
301
|
+
],
|
|
302
|
+
},
|
|
303
|
+
}),
|
|
304
|
+
);
|
|
305
|
+
const capture = await renderAdkRuntime(stream);
|
|
306
|
+
|
|
307
|
+
act(() => {
|
|
308
|
+
capture.runtime!.thread.append({
|
|
309
|
+
role: "user",
|
|
310
|
+
content: [{ type: "text", text: "hello" }],
|
|
311
|
+
});
|
|
312
|
+
});
|
|
313
|
+
await waitFor(() =>
|
|
314
|
+
expect(
|
|
315
|
+
capture.runtime!.thread.getState().messages.at(-1),
|
|
316
|
+
).toMatchObject({
|
|
317
|
+
role: "assistant",
|
|
318
|
+
content: [{ type: "tool-call", toolCallId: "call-1" }],
|
|
319
|
+
}),
|
|
320
|
+
);
|
|
321
|
+
|
|
322
|
+
expect(capture.longRunningToolIds).toEqual(longRunning ? ["call-1"] : []);
|
|
323
|
+
|
|
324
|
+
act(() => {
|
|
325
|
+
capture.runtime!.thread.cancelRun();
|
|
326
|
+
});
|
|
327
|
+
await waitFor(() =>
|
|
328
|
+
expect(capture.runtime!.thread.getState().isRunning).toBe(false),
|
|
329
|
+
);
|
|
330
|
+
|
|
331
|
+
expect(
|
|
332
|
+
capture.runtime!.thread.getState().messages.at(-1)?.status,
|
|
333
|
+
).toEqual({
|
|
334
|
+
type: "incomplete",
|
|
335
|
+
reason: "cancelled",
|
|
336
|
+
});
|
|
337
|
+
expect(capture.longRunningToolIds).toEqual([]);
|
|
338
|
+
|
|
339
|
+
act(() => {
|
|
340
|
+
capture.runtime!.thread.append({
|
|
341
|
+
role: "user",
|
|
342
|
+
content: [{ type: "text", text: "next turn" }],
|
|
343
|
+
});
|
|
344
|
+
});
|
|
345
|
+
await waitFor(() => expect(stream).toHaveBeenCalledTimes(2));
|
|
346
|
+
expect(stream.mock.calls[1]![0]).toEqual(
|
|
347
|
+
expect.arrayContaining([
|
|
348
|
+
expect.objectContaining({
|
|
349
|
+
type: "tool",
|
|
350
|
+
tool_call_id: "call-1",
|
|
351
|
+
content: '{"cancelled":true}',
|
|
352
|
+
}),
|
|
353
|
+
]),
|
|
354
|
+
);
|
|
355
|
+
},
|
|
356
|
+
);
|
|
357
|
+
|
|
358
|
+
it("cancels the assistant message when the last event is a tool response", async () => {
|
|
359
|
+
const capture = await renderAdkRuntime(
|
|
360
|
+
makePendingStream(
|
|
361
|
+
{
|
|
362
|
+
id: "tool-call",
|
|
363
|
+
author: "agent",
|
|
364
|
+
content: {
|
|
365
|
+
role: "model",
|
|
366
|
+
parts: [
|
|
367
|
+
{ functionCall: { id: "call-1", name: "lookup", args: {} } },
|
|
368
|
+
],
|
|
369
|
+
},
|
|
370
|
+
},
|
|
371
|
+
{
|
|
372
|
+
id: "tool-response",
|
|
373
|
+
author: "agent",
|
|
374
|
+
content: {
|
|
375
|
+
role: "user",
|
|
376
|
+
parts: [
|
|
377
|
+
{
|
|
378
|
+
functionResponse: {
|
|
379
|
+
id: "call-1",
|
|
380
|
+
name: "lookup",
|
|
381
|
+
response: { result: "found" },
|
|
382
|
+
},
|
|
383
|
+
},
|
|
384
|
+
],
|
|
385
|
+
},
|
|
386
|
+
},
|
|
387
|
+
),
|
|
388
|
+
);
|
|
389
|
+
|
|
390
|
+
act(() => {
|
|
391
|
+
capture.runtime!.thread.append({
|
|
392
|
+
role: "user",
|
|
393
|
+
content: [{ type: "text", text: "hello" }],
|
|
394
|
+
});
|
|
395
|
+
});
|
|
396
|
+
await waitFor(() =>
|
|
397
|
+
expect(capture.runtime!.thread.getState().messages.at(-1)).toMatchObject({
|
|
398
|
+
role: "assistant",
|
|
399
|
+
content: [{ type: "tool-call", result: '{"result":"found"}' }],
|
|
400
|
+
}),
|
|
401
|
+
);
|
|
402
|
+
|
|
403
|
+
act(() => {
|
|
404
|
+
capture.runtime!.thread.cancelRun();
|
|
405
|
+
});
|
|
406
|
+
await waitFor(() =>
|
|
407
|
+
expect(capture.runtime!.thread.getState().isRunning).toBe(false),
|
|
408
|
+
);
|
|
409
|
+
expect(capture.runtime!.thread.getState().messages.at(-1)?.status).toEqual({
|
|
410
|
+
type: "incomplete",
|
|
411
|
+
reason: "cancelled",
|
|
412
|
+
});
|
|
413
|
+
});
|
|
414
|
+
|
|
415
|
+
it("keeps an earlier assistant message unchanged when stopped before new output", async () => {
|
|
416
|
+
let sends = 0;
|
|
417
|
+
const capture = await renderAdkRuntime(async function* (messages, config) {
|
|
418
|
+
if (sends++ === 0) {
|
|
419
|
+
yield {
|
|
420
|
+
id: "earlier-tool-call",
|
|
421
|
+
author: "agent",
|
|
422
|
+
content: {
|
|
423
|
+
role: "model",
|
|
424
|
+
parts: [
|
|
425
|
+
{ functionCall: { id: "call-1", name: "lookup", args: {} } },
|
|
426
|
+
],
|
|
427
|
+
},
|
|
428
|
+
};
|
|
429
|
+
} else {
|
|
430
|
+
yield* await makePendingStream()(messages, config);
|
|
431
|
+
}
|
|
432
|
+
});
|
|
433
|
+
|
|
434
|
+
act(() => {
|
|
435
|
+
capture.runtime!.thread.append({
|
|
436
|
+
role: "user",
|
|
437
|
+
content: [{ type: "text", text: "first" }],
|
|
438
|
+
});
|
|
439
|
+
});
|
|
440
|
+
await waitFor(() => {
|
|
441
|
+
expect(sends).toBe(1);
|
|
442
|
+
expect(capture.runtime!.thread.getState().isRunning).toBe(false);
|
|
443
|
+
});
|
|
444
|
+
act(() => {
|
|
445
|
+
capture.runtime!.thread.append({
|
|
446
|
+
role: "user",
|
|
447
|
+
content: [{ type: "text", text: "second" }],
|
|
448
|
+
});
|
|
449
|
+
});
|
|
450
|
+
await waitFor(() => expect(sends).toBe(2));
|
|
451
|
+
const earlierMessage = capture
|
|
452
|
+
.runtime!.thread.getState()
|
|
453
|
+
.messages.find((m) => m.role === "assistant");
|
|
454
|
+
expect(earlierMessage).toBeDefined();
|
|
455
|
+
|
|
456
|
+
act(() => {
|
|
457
|
+
capture.runtime!.thread.cancelRun();
|
|
458
|
+
});
|
|
459
|
+
await waitFor(() =>
|
|
460
|
+
expect(capture.runtime!.thread.getState().isRunning).toBe(false),
|
|
461
|
+
);
|
|
462
|
+
expect(
|
|
463
|
+
capture
|
|
464
|
+
.runtime!.thread.getState()
|
|
465
|
+
.messages.find((m) => m.id === earlierMessage!.id),
|
|
466
|
+
).toEqual(earlierMessage);
|
|
467
|
+
});
|
|
468
|
+
|
|
469
|
+
it("marks a partial assistant message as cancelled when Stop is pressed", async () => {
|
|
470
|
+
const capture = await renderAdkRuntime(
|
|
471
|
+
makePendingStream({
|
|
472
|
+
id: "partial",
|
|
473
|
+
invocationId: "run-1",
|
|
474
|
+
author: "agent",
|
|
475
|
+
partial: true,
|
|
476
|
+
content: { role: "model", parts: [{ text: "partial answer" }] },
|
|
477
|
+
}),
|
|
478
|
+
);
|
|
479
|
+
|
|
480
|
+
act(() => {
|
|
481
|
+
capture.runtime!.thread.append({
|
|
482
|
+
role: "user",
|
|
483
|
+
content: [{ type: "text", text: "hello" }],
|
|
484
|
+
});
|
|
485
|
+
});
|
|
486
|
+
await waitFor(() =>
|
|
487
|
+
expect(capture.runtime!.thread.getState().messages.at(-1)).toMatchObject({
|
|
488
|
+
role: "assistant",
|
|
489
|
+
content: [{ type: "text", text: "partial answer" }],
|
|
490
|
+
}),
|
|
491
|
+
);
|
|
492
|
+
|
|
493
|
+
act(() => {
|
|
494
|
+
capture.runtime!.thread.cancelRun();
|
|
495
|
+
});
|
|
496
|
+
await waitFor(() =>
|
|
497
|
+
expect(capture.runtime!.thread.getState().isRunning).toBe(false),
|
|
498
|
+
);
|
|
499
|
+
|
|
500
|
+
expect(capture.runtime!.thread.getState().messages.at(-1)?.status).toEqual({
|
|
501
|
+
type: "incomplete",
|
|
502
|
+
reason: "cancelled",
|
|
503
|
+
});
|
|
504
|
+
});
|
|
505
|
+
|
|
506
|
+
it("keeps a completed assistant message complete when Stop is pressed", async () => {
|
|
507
|
+
const capture = await renderAdkRuntime(
|
|
508
|
+
makePendingStream({
|
|
509
|
+
id: "complete",
|
|
510
|
+
invocationId: "run-1",
|
|
511
|
+
author: "agent",
|
|
512
|
+
content: { role: "model", parts: [{ text: "final answer" }] },
|
|
513
|
+
}),
|
|
514
|
+
);
|
|
515
|
+
|
|
516
|
+
act(() => {
|
|
517
|
+
capture.runtime!.thread.append({
|
|
518
|
+
role: "user",
|
|
519
|
+
content: [{ type: "text", text: "hello" }],
|
|
520
|
+
});
|
|
521
|
+
});
|
|
522
|
+
await waitFor(() =>
|
|
523
|
+
expect(
|
|
524
|
+
capture.runtime!.thread.getState().messages.at(-1)?.status,
|
|
525
|
+
).toEqual({
|
|
526
|
+
type: "complete",
|
|
527
|
+
reason: "stop",
|
|
528
|
+
}),
|
|
529
|
+
);
|
|
530
|
+
expect(capture.runtime!.thread.getState().isRunning).toBe(true);
|
|
531
|
+
|
|
532
|
+
act(() => {
|
|
533
|
+
capture.runtime!.thread.cancelRun();
|
|
534
|
+
});
|
|
535
|
+
await waitFor(() =>
|
|
536
|
+
expect(capture.runtime!.thread.getState().isRunning).toBe(false),
|
|
537
|
+
);
|
|
538
|
+
|
|
539
|
+
expect(capture.runtime!.thread.getState().messages.at(-1)?.status).toEqual({
|
|
540
|
+
type: "complete",
|
|
541
|
+
reason: "stop",
|
|
542
|
+
});
|
|
543
|
+
});
|
|
544
|
+
});
|
|
@@ -0,0 +1,166 @@
|
|
|
1
|
+
// @vitest-environment jsdom
|
|
2
|
+
|
|
3
|
+
import { act } from "react";
|
|
4
|
+
import { afterAll, afterEach, expect, it, vi } from "vitest";
|
|
5
|
+
import type { AdkThreadSnapshot } from "./types";
|
|
6
|
+
|
|
7
|
+
type Family = { current: unknown };
|
|
8
|
+
type RendererInternals = {
|
|
9
|
+
setRefreshHandler: (resolve: (type: unknown) => Family | undefined) => void;
|
|
10
|
+
scheduleRefresh: (
|
|
11
|
+
root: unknown,
|
|
12
|
+
update: { staleFamilies: Set<Family>; updatedFamilies: Set<Family> },
|
|
13
|
+
) => void;
|
|
14
|
+
};
|
|
15
|
+
|
|
16
|
+
const refreshHarness = vi.hoisted(() => {
|
|
17
|
+
const state: {
|
|
18
|
+
renderer: RendererInternals | undefined;
|
|
19
|
+
roots: Set<unknown>;
|
|
20
|
+
} = { renderer: undefined, roots: new Set() };
|
|
21
|
+
vi.stubGlobal("IS_REACT_ACT_ENVIRONMENT", true);
|
|
22
|
+
vi.stubGlobal("__REACT_DEVTOOLS_GLOBAL_HOOK__", {
|
|
23
|
+
supportsFiber: true,
|
|
24
|
+
inject: (internals: RendererInternals) => {
|
|
25
|
+
state.renderer = internals;
|
|
26
|
+
return 1;
|
|
27
|
+
},
|
|
28
|
+
onScheduleFiberRoot: () => {},
|
|
29
|
+
onCommitFiberRoot: (_id: number, root: unknown) => state.roots.add(root),
|
|
30
|
+
onCommitFiberUnmount: () => {},
|
|
31
|
+
});
|
|
32
|
+
return state;
|
|
33
|
+
});
|
|
34
|
+
|
|
35
|
+
const { aui } = vi.hoisted(() => ({
|
|
36
|
+
aui: {
|
|
37
|
+
threadListItem: {
|
|
38
|
+
source: {},
|
|
39
|
+
getState: () => ({ externalId: "adk-1" }),
|
|
40
|
+
},
|
|
41
|
+
},
|
|
42
|
+
}));
|
|
43
|
+
|
|
44
|
+
vi.mock("@assistant-ui/store", async (importOriginal) => ({
|
|
45
|
+
...(await importOriginal<typeof import("@assistant-ui/store")>()),
|
|
46
|
+
useAui: () => aui,
|
|
47
|
+
}));
|
|
48
|
+
|
|
49
|
+
vi.mock("@assistant-ui/core/react", async (importOriginal) => ({
|
|
50
|
+
...(await importOriginal<typeof import("@assistant-ui/core/react")>()),
|
|
51
|
+
useCloudThreadListAdapter: () => ({}),
|
|
52
|
+
useRemoteThreadListRuntime: ({
|
|
53
|
+
runtimeHook,
|
|
54
|
+
}: {
|
|
55
|
+
runtimeHook: () => unknown;
|
|
56
|
+
}) => runtimeHook(),
|
|
57
|
+
useExternalMessageConverter: ({ messages }: { messages: unknown }) =>
|
|
58
|
+
messages,
|
|
59
|
+
useExternalStoreRuntime: (options: unknown) => options,
|
|
60
|
+
}));
|
|
61
|
+
|
|
62
|
+
const { cleanup, render, waitFor } = await import("@testing-library/react");
|
|
63
|
+
const { useAdkRuntime } = await import("./useAdkRuntime");
|
|
64
|
+
|
|
65
|
+
afterEach(() => {
|
|
66
|
+
cleanup();
|
|
67
|
+
refreshHarness.renderer?.setRefreshHandler(() => undefined);
|
|
68
|
+
refreshHarness.roots.clear();
|
|
69
|
+
});
|
|
70
|
+
afterAll(() => vi.unstubAllGlobals());
|
|
71
|
+
|
|
72
|
+
const deferred = <T,>() => {
|
|
73
|
+
let resolve!: (value: T) => void;
|
|
74
|
+
const promise = new Promise<T>((res) => {
|
|
75
|
+
resolve = res;
|
|
76
|
+
});
|
|
77
|
+
return { promise, resolve };
|
|
78
|
+
};
|
|
79
|
+
|
|
80
|
+
const refresh = async (Before: unknown, After: unknown) => {
|
|
81
|
+
const family: Family = { current: After };
|
|
82
|
+
refreshHarness.renderer!.setRefreshHandler((type) =>
|
|
83
|
+
type === Before || type === After ? family : undefined,
|
|
84
|
+
);
|
|
85
|
+
await act(async () => {
|
|
86
|
+
for (const root of refreshHarness.roots) {
|
|
87
|
+
refreshHarness.renderer!.scheduleRefresh(root, {
|
|
88
|
+
staleFamilies: new Set(),
|
|
89
|
+
updatedFamilies: new Set([family]),
|
|
90
|
+
});
|
|
91
|
+
}
|
|
92
|
+
});
|
|
93
|
+
await act(async () => {});
|
|
94
|
+
};
|
|
95
|
+
|
|
96
|
+
it("keeps an initial thread load through Fast Refresh and aborts a later load on unmount", async () => {
|
|
97
|
+
const initial = deferred<AdkThreadSnapshot>();
|
|
98
|
+
const refetch = deferred<AdkThreadSnapshot>();
|
|
99
|
+
const signals: AbortSignal[] = [];
|
|
100
|
+
const load = vi.fn(
|
|
101
|
+
(_id: string, options?: { signal?: AbortSignal | undefined }) => {
|
|
102
|
+
signals.push(options!.signal!);
|
|
103
|
+
return signals.length === 1 ? initial.promise : refetch.promise;
|
|
104
|
+
},
|
|
105
|
+
);
|
|
106
|
+
let runtime:
|
|
107
|
+
| {
|
|
108
|
+
messages: unknown;
|
|
109
|
+
onRefetchThread: () => Promise<void>;
|
|
110
|
+
}
|
|
111
|
+
| undefined;
|
|
112
|
+
let rendered: string | undefined;
|
|
113
|
+
const host = (name: string) => () => {
|
|
114
|
+
rendered = name;
|
|
115
|
+
runtime = useAdkRuntime({
|
|
116
|
+
stream: async function* () {},
|
|
117
|
+
load,
|
|
118
|
+
}) as unknown as typeof runtime;
|
|
119
|
+
return null;
|
|
120
|
+
};
|
|
121
|
+
const Before = host("before");
|
|
122
|
+
const After = host("after");
|
|
123
|
+
const view = render(<Before />);
|
|
124
|
+
await waitFor(() => expect(load).toHaveBeenCalledTimes(1));
|
|
125
|
+
expect(signals[0]!.aborted).toBe(false);
|
|
126
|
+
|
|
127
|
+
await refresh(Before, After);
|
|
128
|
+
expect(rendered).toBe("after");
|
|
129
|
+
expect(signals[0]!.aborted).toBe(false);
|
|
130
|
+
expect(load).toHaveBeenCalledTimes(1);
|
|
131
|
+
|
|
132
|
+
let refetchSettled = false;
|
|
133
|
+
act(() => {
|
|
134
|
+
void runtime!.onRefetchThread().then(() => {
|
|
135
|
+
refetchSettled = true;
|
|
136
|
+
});
|
|
137
|
+
});
|
|
138
|
+
expect(load).toHaveBeenCalledTimes(1);
|
|
139
|
+
expect(refetchSettled).toBe(false);
|
|
140
|
+
|
|
141
|
+
await act(async () => {
|
|
142
|
+
initial.resolve({
|
|
143
|
+
messages: [
|
|
144
|
+
{
|
|
145
|
+
id: "history-1",
|
|
146
|
+
type: "ai",
|
|
147
|
+
content: [{ type: "text", text: "history landed" }],
|
|
148
|
+
},
|
|
149
|
+
],
|
|
150
|
+
});
|
|
151
|
+
});
|
|
152
|
+
await waitFor(() => expect(refetchSettled).toBe(true));
|
|
153
|
+
await waitFor(() =>
|
|
154
|
+
expect(JSON.stringify(runtime!.messages)).toContain("history landed"),
|
|
155
|
+
);
|
|
156
|
+
|
|
157
|
+
act(() => {
|
|
158
|
+
void runtime!.onRefetchThread();
|
|
159
|
+
});
|
|
160
|
+
await waitFor(() => expect(load).toHaveBeenCalledTimes(2));
|
|
161
|
+
expect(signals[1]!.aborted).toBe(false);
|
|
162
|
+
view.unmount();
|
|
163
|
+
await act(async () => {});
|
|
164
|
+
expect(signals[1]!.aborted).toBe(true);
|
|
165
|
+
refetch.resolve({ messages: [] });
|
|
166
|
+
});
|