@assistant-ui/react-google-adk 0.0.35 → 0.0.36
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 +12 -2
- package/dist/AdkClient.d.ts +4 -0
- package/dist/AdkClient.d.ts.map +1 -1
- package/dist/AdkClient.js +12 -7
- package/dist/AdkClient.js.map +1 -1
- package/dist/AdkSessionAdapter.d.ts.map +1 -1
- package/dist/AdkSessionAdapter.js +7 -7
- package/dist/AdkSessionAdapter.js.map +1 -1
- package/dist/AdkThreadController.d.ts +15 -0
- package/dist/AdkThreadController.d.ts.map +1 -0
- package/dist/AdkThreadController.js +35 -0
- package/dist/AdkThreadController.js.map +1 -0
- package/dist/adkThreadState.d.ts +54 -0
- package/dist/adkThreadState.d.ts.map +1 -0
- package/dist/adkThreadState.js +93 -0
- package/dist/adkThreadState.js.map +1 -0
- package/dist/convertToAdkMessages.js +1 -1
- package/dist/convertToAdkMessages.js.map +1 -1
- package/dist/sdkIdentity.js +1 -1
- package/dist/server/createAdkApiRoute.d.ts +37 -6
- package/dist/server/createAdkApiRoute.d.ts.map +1 -1
- package/dist/server/createAdkApiRoute.js +55 -5
- package/dist/server/createAdkApiRoute.js.map +1 -1
- package/dist/server/parseAdkRequest.d.ts +4 -1
- package/dist/server/parseAdkRequest.d.ts.map +1 -1
- package/dist/server/parseAdkRequest.js +5 -1
- package/dist/server/parseAdkRequest.js.map +1 -1
- package/dist/useAdkMessages.d.ts +9 -7
- package/dist/useAdkMessages.d.ts.map +1 -1
- package/dist/useAdkMessages.js +55 -77
- package/dist/useAdkMessages.js.map +1 -1
- package/dist/useAdkRuntime.d.ts +7 -1
- package/dist/useAdkRuntime.d.ts.map +1 -1
- package/dist/useAdkRuntime.js +134 -56
- package/dist/useAdkRuntime.js.map +1 -1
- package/package.json +5 -5
- package/src/AdkClient.test.ts +78 -2
- package/src/AdkClient.ts +24 -6
- package/src/AdkSessionAdapter.ts +1 -1
- package/src/AdkThreadController.test.ts +90 -0
- package/src/AdkThreadController.ts +45 -0
- package/src/adkThreadState.test.ts +207 -0
- package/src/adkThreadState.ts +124 -0
- package/src/convertToAdkMessages.test.ts +19 -0
- package/src/convertToAdkMessages.ts +1 -1
- package/src/hooks.test.tsx +1 -0
- package/src/server/createAdkApiRoute.controls.test.ts +66 -0
- package/src/server/createAdkApiRoute.test.ts +282 -0
- package/src/server/createAdkApiRoute.ts +119 -11
- package/src/server/parseAdkRequest.test.ts +11 -3
- package/src/server/parseAdkRequest.ts +7 -1
- package/src/useAdkMessages.test.ts +1 -0
- package/src/useAdkMessages.ts +61 -96
- package/src/useAdkRuntime.cancellation.test.tsx +4 -3
- package/src/useAdkRuntime.cloud-options.test.tsx +59 -0
- package/src/useAdkRuntime.refetch.test.tsx +548 -4
- package/src/useAdkRuntime.replacement.test.tsx +718 -1
- package/src/useAdkRuntime.ts +169 -73
- package/src/useAdkRuntimeApproval.test.tsx +87 -1
- package/dist/raceWithAbortSignal.d.ts +0 -2
- package/dist/raceWithAbortSignal.d.ts.map +0 -1
- package/dist/raceWithAbortSignal.js +0 -45
- package/dist/raceWithAbortSignal.js.map +0 -1
- package/src/raceWithAbortSignal.test.ts +0 -73
- package/src/raceWithAbortSignal.ts +0 -48
|
@@ -9,7 +9,7 @@ import type {
|
|
|
9
9
|
RemoteThreadListAdapter,
|
|
10
10
|
} from "@assistant-ui/core";
|
|
11
11
|
import { useAdkRuntime } from "./useAdkRuntime";
|
|
12
|
-
import type { AdkEvent } from "./types";
|
|
12
|
+
import type { AdkEvent, AdkMessage, AdkThreadSnapshot } from "./types";
|
|
13
13
|
import { settleOutsideAct } from "./tests/settleOutsideAct";
|
|
14
14
|
|
|
15
15
|
const makeThreadListAdapter = (): RemoteThreadListAdapter => ({
|
|
@@ -140,4 +140,721 @@ describe("useAdkRuntime replacement runs", () => {
|
|
|
140
140
|
).toContain("done-1");
|
|
141
141
|
expect(capture.runtime!.thread.getState().isRunning).toBe(false);
|
|
142
142
|
});
|
|
143
|
+
|
|
144
|
+
const mountWithCheckpoint = async (
|
|
145
|
+
stream: (...args: never[]) => AsyncGenerator<AdkEvent>,
|
|
146
|
+
getCheckpointId: () => Promise<string | null>,
|
|
147
|
+
allowCancellation = false,
|
|
148
|
+
load?: () => Promise<AdkThreadSnapshot>,
|
|
149
|
+
) => {
|
|
150
|
+
const capture: { runtime: AssistantRuntime | null } = { runtime: null };
|
|
151
|
+
const Inner: FC = () => {
|
|
152
|
+
const runtime = useAdkRuntime({
|
|
153
|
+
stream: stream as never,
|
|
154
|
+
sessionAdapter: makeThreadListAdapter(),
|
|
155
|
+
getCheckpointId,
|
|
156
|
+
unstable_allowCancellation: allowCancellation,
|
|
157
|
+
...(load ? { load } : {}),
|
|
158
|
+
});
|
|
159
|
+
capture.runtime = runtime;
|
|
160
|
+
return (
|
|
161
|
+
<AssistantRuntimeProvider runtime={runtime}>
|
|
162
|
+
{null}
|
|
163
|
+
</AssistantRuntimeProvider>
|
|
164
|
+
);
|
|
165
|
+
};
|
|
166
|
+
await act(async () => {
|
|
167
|
+
render(<Inner />);
|
|
168
|
+
});
|
|
169
|
+
await waitFor(() => expect(capture.runtime).not.toBeNull());
|
|
170
|
+
await settleOutsideAct(() =>
|
|
171
|
+
capture.runtime!.threads.switchToThread("adk-1"),
|
|
172
|
+
);
|
|
173
|
+
return capture.runtime!;
|
|
174
|
+
};
|
|
175
|
+
|
|
176
|
+
it("starts one run with both staged messages", async () => {
|
|
177
|
+
const stream = vi.fn(async function* (
|
|
178
|
+
_messages: AdkMessage[],
|
|
179
|
+
): AsyncGenerator<AdkEvent> {});
|
|
180
|
+
const runtime = await mountWithCheckpoint(stream, async () => null);
|
|
181
|
+
|
|
182
|
+
await act(async () => {
|
|
183
|
+
runtime.thread.append({
|
|
184
|
+
role: "user",
|
|
185
|
+
content: [{ type: "text", text: "first" }],
|
|
186
|
+
startRun: false,
|
|
187
|
+
});
|
|
188
|
+
});
|
|
189
|
+
await act(async () => {
|
|
190
|
+
runtime.thread.append({
|
|
191
|
+
role: "user",
|
|
192
|
+
content: [{ type: "text", text: "second" }],
|
|
193
|
+
startRun: false,
|
|
194
|
+
});
|
|
195
|
+
});
|
|
196
|
+
const parentId = runtime.thread.getState().messages[1]!.id;
|
|
197
|
+
|
|
198
|
+
await act(async () => {
|
|
199
|
+
await runtime.thread.startRun({ parentId });
|
|
200
|
+
});
|
|
201
|
+
|
|
202
|
+
expect(stream).toHaveBeenCalledTimes(1);
|
|
203
|
+
expect(stream.mock.calls[0]![0]).toMatchObject([
|
|
204
|
+
{ type: "human", content: "first" },
|
|
205
|
+
{ type: "human", content: "second" },
|
|
206
|
+
]);
|
|
207
|
+
});
|
|
208
|
+
|
|
209
|
+
it("starts an edit made while a run streams from the truncated thread", async () => {
|
|
210
|
+
const releaseStale = deferred();
|
|
211
|
+
const checkpoint = deferred();
|
|
212
|
+
let calls = 0;
|
|
213
|
+
const stream = vi.fn(async function* (): AsyncGenerator<AdkEvent> {
|
|
214
|
+
const call = calls++;
|
|
215
|
+
if (call === 0) {
|
|
216
|
+
yield {
|
|
217
|
+
id: "stale-1",
|
|
218
|
+
invocationId: "run-0",
|
|
219
|
+
author: "agent",
|
|
220
|
+
content: { role: "model", parts: [{ text: "stale partial" }] },
|
|
221
|
+
};
|
|
222
|
+
await releaseStale.promise;
|
|
223
|
+
yield {
|
|
224
|
+
id: "stale-2",
|
|
225
|
+
invocationId: "run-0",
|
|
226
|
+
author: "agent",
|
|
227
|
+
content: { role: "model", parts: [{ text: "stale late" }] },
|
|
228
|
+
};
|
|
229
|
+
return;
|
|
230
|
+
}
|
|
231
|
+
yield {
|
|
232
|
+
id: "fresh",
|
|
233
|
+
invocationId: "run-1",
|
|
234
|
+
author: "agent",
|
|
235
|
+
content: { role: "model", parts: [{ text: "fresh answer" }] },
|
|
236
|
+
};
|
|
237
|
+
});
|
|
238
|
+
const runtime = await mountWithCheckpoint(stream, async () => {
|
|
239
|
+
await checkpoint.promise;
|
|
240
|
+
return "cp-1";
|
|
241
|
+
});
|
|
242
|
+
|
|
243
|
+
act(() => {
|
|
244
|
+
runtime.thread.append({
|
|
245
|
+
role: "user",
|
|
246
|
+
content: [{ type: "text", text: "original question" }],
|
|
247
|
+
});
|
|
248
|
+
});
|
|
249
|
+
await waitFor(() =>
|
|
250
|
+
expect(JSON.stringify(runtime.thread.getState().messages)).toContain(
|
|
251
|
+
"stale partial",
|
|
252
|
+
),
|
|
253
|
+
);
|
|
254
|
+
const original = runtime.thread.getState().messages[0]!;
|
|
255
|
+
|
|
256
|
+
await act(async () => {
|
|
257
|
+
runtime.thread.append({
|
|
258
|
+
role: "user",
|
|
259
|
+
parentId: null,
|
|
260
|
+
sourceId: original.id,
|
|
261
|
+
content: [{ type: "text", text: "edited question" }],
|
|
262
|
+
});
|
|
263
|
+
});
|
|
264
|
+
await act(async () => {
|
|
265
|
+
releaseStale.resolve();
|
|
266
|
+
await Promise.resolve();
|
|
267
|
+
});
|
|
268
|
+
await act(async () => {
|
|
269
|
+
checkpoint.resolve();
|
|
270
|
+
});
|
|
271
|
+
await waitFor(() =>
|
|
272
|
+
expect(runtime.thread.getState().isRunning).toBe(false),
|
|
273
|
+
);
|
|
274
|
+
|
|
275
|
+
const messages = JSON.stringify(runtime.thread.getState().messages);
|
|
276
|
+
expect(messages).toContain("edited question");
|
|
277
|
+
expect(messages).toContain("fresh answer");
|
|
278
|
+
expect(messages).not.toContain("original question");
|
|
279
|
+
expect(messages).not.toContain("stale");
|
|
280
|
+
});
|
|
281
|
+
|
|
282
|
+
it("starts a reload made while a run streams from the truncated thread", async () => {
|
|
283
|
+
const releaseStale = deferred();
|
|
284
|
+
const checkpoint = deferred();
|
|
285
|
+
let calls = 0;
|
|
286
|
+
const stream = vi.fn(async function* (): AsyncGenerator<AdkEvent> {
|
|
287
|
+
const call = calls++;
|
|
288
|
+
if (call === 0) {
|
|
289
|
+
yield {
|
|
290
|
+
id: "stale-1",
|
|
291
|
+
invocationId: "run-0",
|
|
292
|
+
author: "agent",
|
|
293
|
+
content: { role: "model", parts: [{ text: "stale partial" }] },
|
|
294
|
+
};
|
|
295
|
+
await releaseStale.promise;
|
|
296
|
+
yield {
|
|
297
|
+
id: "stale-2",
|
|
298
|
+
invocationId: "run-0",
|
|
299
|
+
author: "agent",
|
|
300
|
+
content: { role: "model", parts: [{ text: "stale late" }] },
|
|
301
|
+
};
|
|
302
|
+
return;
|
|
303
|
+
}
|
|
304
|
+
yield {
|
|
305
|
+
id: "fresh",
|
|
306
|
+
invocationId: "run-1",
|
|
307
|
+
author: "agent",
|
|
308
|
+
content: { role: "model", parts: [{ text: "fresh answer" }] },
|
|
309
|
+
};
|
|
310
|
+
});
|
|
311
|
+
const runtime = await mountWithCheckpoint(stream, async () => {
|
|
312
|
+
await checkpoint.promise;
|
|
313
|
+
return "cp-1";
|
|
314
|
+
});
|
|
315
|
+
|
|
316
|
+
act(() => {
|
|
317
|
+
runtime.thread.append({
|
|
318
|
+
role: "user",
|
|
319
|
+
content: [{ type: "text", text: "original question" }],
|
|
320
|
+
});
|
|
321
|
+
});
|
|
322
|
+
await waitFor(() =>
|
|
323
|
+
expect(JSON.stringify(runtime.thread.getState().messages)).toContain(
|
|
324
|
+
"stale partial",
|
|
325
|
+
),
|
|
326
|
+
);
|
|
327
|
+
const answer = runtime.thread.getState().messages[1]!;
|
|
328
|
+
|
|
329
|
+
await act(async () => {
|
|
330
|
+
runtime.thread.getMessageById(answer.id).reload();
|
|
331
|
+
});
|
|
332
|
+
expect(runtime.thread.getState().isRunning).toBe(true);
|
|
333
|
+
await act(async () => {
|
|
334
|
+
releaseStale.resolve();
|
|
335
|
+
await Promise.resolve();
|
|
336
|
+
});
|
|
337
|
+
const duringLookup = JSON.stringify(runtime.thread.getState().messages);
|
|
338
|
+
expect(duringLookup).toContain("original question");
|
|
339
|
+
expect(duringLookup).not.toContain("stale");
|
|
340
|
+
await act(async () => {
|
|
341
|
+
checkpoint.resolve();
|
|
342
|
+
});
|
|
343
|
+
await waitFor(() =>
|
|
344
|
+
expect(runtime.thread.getState().isRunning).toBe(false),
|
|
345
|
+
);
|
|
346
|
+
|
|
347
|
+
expect(stream).toHaveBeenCalledTimes(2);
|
|
348
|
+
expect(stream.mock.calls[1]).toEqual([
|
|
349
|
+
[],
|
|
350
|
+
expect.objectContaining({ checkpointId: "cp-1" }),
|
|
351
|
+
]);
|
|
352
|
+
const messages = JSON.stringify(runtime.thread.getState().messages);
|
|
353
|
+
expect(messages).toContain("original question");
|
|
354
|
+
expect(messages).toContain("fresh answer");
|
|
355
|
+
expect(messages).not.toContain("stale");
|
|
356
|
+
});
|
|
357
|
+
|
|
358
|
+
it("reports the thread running while an edit looks up its checkpoint", async () => {
|
|
359
|
+
const checkpoint = deferred();
|
|
360
|
+
const stream = vi.fn(async function* (): AsyncGenerator<AdkEvent> {
|
|
361
|
+
yield {
|
|
362
|
+
id: "answer",
|
|
363
|
+
invocationId: "run",
|
|
364
|
+
author: "agent",
|
|
365
|
+
content: { role: "model", parts: [{ text: "answer" }] },
|
|
366
|
+
};
|
|
367
|
+
});
|
|
368
|
+
const runtime = await mountWithCheckpoint(stream, async () => {
|
|
369
|
+
await checkpoint.promise;
|
|
370
|
+
return "cp-1";
|
|
371
|
+
});
|
|
372
|
+
|
|
373
|
+
await act(async () => {
|
|
374
|
+
runtime.thread.append({
|
|
375
|
+
role: "user",
|
|
376
|
+
content: [{ type: "text", text: "question" }],
|
|
377
|
+
});
|
|
378
|
+
});
|
|
379
|
+
await waitFor(() =>
|
|
380
|
+
expect(runtime.thread.getState().isRunning).toBe(false),
|
|
381
|
+
);
|
|
382
|
+
const original = runtime.thread.getState().messages[0]!;
|
|
383
|
+
|
|
384
|
+
await act(async () => {
|
|
385
|
+
runtime.thread.append({
|
|
386
|
+
role: "user",
|
|
387
|
+
parentId: null,
|
|
388
|
+
sourceId: original.id,
|
|
389
|
+
content: [{ type: "text", text: "edited question" }],
|
|
390
|
+
});
|
|
391
|
+
});
|
|
392
|
+
expect(runtime.thread.getState().isRunning).toBe(true);
|
|
393
|
+
expect(runtime.thread.getState().messages[0]!.content).toEqual([
|
|
394
|
+
expect.objectContaining({ text: "edited question" }),
|
|
395
|
+
]);
|
|
396
|
+
|
|
397
|
+
await act(async () => {
|
|
398
|
+
checkpoint.resolve();
|
|
399
|
+
});
|
|
400
|
+
await waitFor(() =>
|
|
401
|
+
expect(runtime.thread.getState().isRunning).toBe(false),
|
|
402
|
+
);
|
|
403
|
+
expect(stream).toHaveBeenCalledTimes(2);
|
|
404
|
+
expect(
|
|
405
|
+
runtime.thread
|
|
406
|
+
.getState()
|
|
407
|
+
.messages.map((m) => (m.content[0] as { text: string }).text),
|
|
408
|
+
).toEqual(["edited question", "answer"]);
|
|
409
|
+
});
|
|
410
|
+
|
|
411
|
+
it("stops an edit that is still looking up its checkpoint", async () => {
|
|
412
|
+
const checkpoint = deferred();
|
|
413
|
+
const stream = vi.fn(async function* (): AsyncGenerator<AdkEvent> {
|
|
414
|
+
yield {
|
|
415
|
+
id: "answer",
|
|
416
|
+
invocationId: "run",
|
|
417
|
+
author: "agent",
|
|
418
|
+
content: { role: "model", parts: [{ text: "answer" }] },
|
|
419
|
+
};
|
|
420
|
+
});
|
|
421
|
+
const runtime = await mountWithCheckpoint(
|
|
422
|
+
stream,
|
|
423
|
+
async () => {
|
|
424
|
+
await checkpoint.promise;
|
|
425
|
+
return "cp-1";
|
|
426
|
+
},
|
|
427
|
+
true,
|
|
428
|
+
);
|
|
429
|
+
|
|
430
|
+
await act(async () => {
|
|
431
|
+
runtime.thread.append({
|
|
432
|
+
role: "user",
|
|
433
|
+
content: [{ type: "text", text: "question" }],
|
|
434
|
+
});
|
|
435
|
+
});
|
|
436
|
+
await waitFor(() =>
|
|
437
|
+
expect(runtime.thread.getState().isRunning).toBe(false),
|
|
438
|
+
);
|
|
439
|
+
const original = runtime.thread.getState().messages[0]!;
|
|
440
|
+
|
|
441
|
+
await act(async () => {
|
|
442
|
+
runtime.thread.append({
|
|
443
|
+
role: "user",
|
|
444
|
+
parentId: null,
|
|
445
|
+
sourceId: original.id,
|
|
446
|
+
content: [{ type: "text", text: "edited question" }],
|
|
447
|
+
});
|
|
448
|
+
});
|
|
449
|
+
expect(runtime.thread.getState().isRunning).toBe(true);
|
|
450
|
+
|
|
451
|
+
await act(async () => {
|
|
452
|
+
runtime.thread.cancelRun();
|
|
453
|
+
});
|
|
454
|
+
expect(runtime.thread.getState().isRunning).toBe(false);
|
|
455
|
+
await act(async () => {
|
|
456
|
+
checkpoint.resolve();
|
|
457
|
+
});
|
|
458
|
+
expect(stream).toHaveBeenCalledTimes(1);
|
|
459
|
+
expect(
|
|
460
|
+
runtime.thread
|
|
461
|
+
.getState()
|
|
462
|
+
.messages.map((m) => (m.content[0] as { text: string }).text),
|
|
463
|
+
).toEqual(["edited question"]);
|
|
464
|
+
});
|
|
465
|
+
|
|
466
|
+
const reloadAndStop = async () => {
|
|
467
|
+
const checkpoint = deferred();
|
|
468
|
+
let calls = 0;
|
|
469
|
+
const stream = vi.fn(async function* (): AsyncGenerator<AdkEvent> {
|
|
470
|
+
const call = calls++;
|
|
471
|
+
yield {
|
|
472
|
+
id: `answer-${call}`,
|
|
473
|
+
invocationId: `run-${call}`,
|
|
474
|
+
author: "agent",
|
|
475
|
+
content: { role: "model", parts: [{ text: `answer ${call}` }] },
|
|
476
|
+
};
|
|
477
|
+
});
|
|
478
|
+
const getCheckpointId = vi.fn(async () => {
|
|
479
|
+
await checkpoint.promise;
|
|
480
|
+
return "cp-1";
|
|
481
|
+
});
|
|
482
|
+
const runtime = await mountWithCheckpoint(stream, getCheckpointId, true);
|
|
483
|
+
|
|
484
|
+
await act(async () => {
|
|
485
|
+
runtime.thread.append({
|
|
486
|
+
role: "user",
|
|
487
|
+
content: [{ type: "text", text: "question" }],
|
|
488
|
+
});
|
|
489
|
+
});
|
|
490
|
+
await waitFor(() =>
|
|
491
|
+
expect(runtime.thread.getState().isRunning).toBe(false),
|
|
492
|
+
);
|
|
493
|
+
const answer = runtime.thread.getState().messages[1]!;
|
|
494
|
+
|
|
495
|
+
await act(async () => {
|
|
496
|
+
runtime.thread.getMessageById(answer.id).reload();
|
|
497
|
+
});
|
|
498
|
+
expect(JSON.stringify(runtime.thread.getState().messages)).not.toContain(
|
|
499
|
+
"answer 0",
|
|
500
|
+
);
|
|
501
|
+
expect(runtime.thread.getState().isRunning).toBe(true);
|
|
502
|
+
|
|
503
|
+
await act(async () => {
|
|
504
|
+
runtime.thread.cancelRun();
|
|
505
|
+
});
|
|
506
|
+
return { runtime, stream, checkpoint, getCheckpointId };
|
|
507
|
+
};
|
|
508
|
+
|
|
509
|
+
const texts = (runtime: AssistantRuntime) =>
|
|
510
|
+
runtime.thread
|
|
511
|
+
.getState()
|
|
512
|
+
.messages.map((m) => (m.content[0] as { text: string }).text);
|
|
513
|
+
|
|
514
|
+
it("restores the thread when a reload is stopped while it looks up its checkpoint", async () => {
|
|
515
|
+
const { runtime, stream, checkpoint } = await reloadAndStop();
|
|
516
|
+
|
|
517
|
+
expect(runtime.thread.getState().isRunning).toBe(false);
|
|
518
|
+
expect(texts(runtime)).toEqual(["question", "answer 0"]);
|
|
519
|
+
|
|
520
|
+
await act(async () => {
|
|
521
|
+
checkpoint.resolve();
|
|
522
|
+
});
|
|
523
|
+
expect(stream).toHaveBeenCalledTimes(1);
|
|
524
|
+
expect(texts(runtime)).toEqual(["question", "answer 0"]);
|
|
525
|
+
});
|
|
526
|
+
|
|
527
|
+
it("sends the next message after a stopped reload without a checkpoint", async () => {
|
|
528
|
+
const { runtime, stream, checkpoint, getCheckpointId } =
|
|
529
|
+
await reloadAndStop();
|
|
530
|
+
|
|
531
|
+
await act(async () => {
|
|
532
|
+
runtime.thread.append({
|
|
533
|
+
role: "user",
|
|
534
|
+
content: [{ type: "text", text: "follow-up" }],
|
|
535
|
+
});
|
|
536
|
+
});
|
|
537
|
+
await act(async () => {
|
|
538
|
+
checkpoint.resolve();
|
|
539
|
+
});
|
|
540
|
+
await waitFor(() =>
|
|
541
|
+
expect(runtime.thread.getState().isRunning).toBe(false),
|
|
542
|
+
);
|
|
543
|
+
|
|
544
|
+
expect(stream).toHaveBeenCalledTimes(2);
|
|
545
|
+
const [, followUpConfig] = stream.mock.calls[1] as unknown as [
|
|
546
|
+
unknown,
|
|
547
|
+
object,
|
|
548
|
+
];
|
|
549
|
+
expect(followUpConfig).not.toHaveProperty("checkpointId");
|
|
550
|
+
expect(getCheckpointId).toHaveBeenCalledTimes(1);
|
|
551
|
+
expect(texts(runtime)).toEqual([
|
|
552
|
+
"question",
|
|
553
|
+
"answer 0",
|
|
554
|
+
"follow-up",
|
|
555
|
+
"answer 1",
|
|
556
|
+
]);
|
|
557
|
+
});
|
|
558
|
+
|
|
559
|
+
it("keeps a message sent during a reload's checkpoint lookup when Stop follows", async () => {
|
|
560
|
+
const checkpoint = deferred();
|
|
561
|
+
const followUpGate = deferred();
|
|
562
|
+
let calls = 0;
|
|
563
|
+
const stream = vi.fn(async function* (): AsyncGenerator<AdkEvent> {
|
|
564
|
+
const call = calls++;
|
|
565
|
+
if (call === 1) await followUpGate.promise;
|
|
566
|
+
yield {
|
|
567
|
+
id: `answer-${call}`,
|
|
568
|
+
invocationId: `run-${call}`,
|
|
569
|
+
author: "agent",
|
|
570
|
+
content: { role: "model", parts: [{ text: `answer ${call}` }] },
|
|
571
|
+
};
|
|
572
|
+
});
|
|
573
|
+
const runtime = await mountWithCheckpoint(
|
|
574
|
+
stream,
|
|
575
|
+
async () => {
|
|
576
|
+
await checkpoint.promise;
|
|
577
|
+
return "cp-1";
|
|
578
|
+
},
|
|
579
|
+
true,
|
|
580
|
+
);
|
|
581
|
+
|
|
582
|
+
await act(async () => {
|
|
583
|
+
runtime.thread.append({
|
|
584
|
+
role: "user",
|
|
585
|
+
content: [{ type: "text", text: "question" }],
|
|
586
|
+
});
|
|
587
|
+
});
|
|
588
|
+
await waitFor(() =>
|
|
589
|
+
expect(runtime.thread.getState().isRunning).toBe(false),
|
|
590
|
+
);
|
|
591
|
+
const answer = runtime.thread.getState().messages[1]!;
|
|
592
|
+
|
|
593
|
+
await act(async () => {
|
|
594
|
+
runtime.thread.getMessageById(answer.id).reload();
|
|
595
|
+
});
|
|
596
|
+
await act(async () => {
|
|
597
|
+
runtime.thread.append({
|
|
598
|
+
role: "user",
|
|
599
|
+
content: [{ type: "text", text: "follow-up" }],
|
|
600
|
+
});
|
|
601
|
+
});
|
|
602
|
+
await waitFor(() => expect(stream).toHaveBeenCalledTimes(2));
|
|
603
|
+
await act(async () => {
|
|
604
|
+
runtime.thread.cancelRun();
|
|
605
|
+
});
|
|
606
|
+
await act(async () => {
|
|
607
|
+
followUpGate.resolve();
|
|
608
|
+
checkpoint.resolve();
|
|
609
|
+
});
|
|
610
|
+
|
|
611
|
+
expect(stream).toHaveBeenCalledTimes(2);
|
|
612
|
+
const messages = JSON.stringify(runtime.thread.getState().messages);
|
|
613
|
+
expect(messages).toContain("follow-up");
|
|
614
|
+
expect(messages).not.toContain("answer 0");
|
|
615
|
+
});
|
|
616
|
+
|
|
617
|
+
it("keeps an edit made while a reload looks up its checkpoint", async () => {
|
|
618
|
+
const lookups = [deferred(), deferred()];
|
|
619
|
+
const waiting = [...lookups];
|
|
620
|
+
const stream = vi.fn(async function* (): AsyncGenerator<AdkEvent> {
|
|
621
|
+
yield {
|
|
622
|
+
id: "answer",
|
|
623
|
+
invocationId: "run",
|
|
624
|
+
author: "agent",
|
|
625
|
+
content: { role: "model", parts: [{ text: "answer 0" }] },
|
|
626
|
+
};
|
|
627
|
+
});
|
|
628
|
+
const runtime = await mountWithCheckpoint(
|
|
629
|
+
stream,
|
|
630
|
+
async () => {
|
|
631
|
+
await waiting.shift()!.promise;
|
|
632
|
+
return "cp-1";
|
|
633
|
+
},
|
|
634
|
+
true,
|
|
635
|
+
);
|
|
636
|
+
|
|
637
|
+
await act(async () => {
|
|
638
|
+
runtime.thread.append({
|
|
639
|
+
role: "user",
|
|
640
|
+
content: [{ type: "text", text: "question" }],
|
|
641
|
+
});
|
|
642
|
+
});
|
|
643
|
+
await waitFor(() =>
|
|
644
|
+
expect(runtime.thread.getState().isRunning).toBe(false),
|
|
645
|
+
);
|
|
646
|
+
const [original, answer] = runtime.thread.getState().messages;
|
|
647
|
+
|
|
648
|
+
await act(async () => {
|
|
649
|
+
runtime.thread.getMessageById(answer!.id).reload();
|
|
650
|
+
});
|
|
651
|
+
await act(async () => {
|
|
652
|
+
runtime.thread.append({
|
|
653
|
+
role: "user",
|
|
654
|
+
parentId: null,
|
|
655
|
+
sourceId: original!.id,
|
|
656
|
+
content: [{ type: "text", text: "edited question" }],
|
|
657
|
+
});
|
|
658
|
+
});
|
|
659
|
+
await act(async () => {
|
|
660
|
+
runtime.thread.cancelRun();
|
|
661
|
+
});
|
|
662
|
+
await act(async () => {
|
|
663
|
+
lookups[0]!.resolve();
|
|
664
|
+
lookups[1]!.resolve();
|
|
665
|
+
});
|
|
666
|
+
|
|
667
|
+
expect(stream).toHaveBeenCalledTimes(1);
|
|
668
|
+
expect(texts(runtime)).toEqual(["edited question"]);
|
|
669
|
+
});
|
|
670
|
+
|
|
671
|
+
it("keeps a history load that lands while a reload looks up its checkpoint when Stop follows", async () => {
|
|
672
|
+
const checkpoint = deferred();
|
|
673
|
+
const initialLoad = deferred();
|
|
674
|
+
const loaded = deferred();
|
|
675
|
+
let loadCount = 0;
|
|
676
|
+
const stream = vi.fn(async function* (): AsyncGenerator<AdkEvent> {
|
|
677
|
+
yield {
|
|
678
|
+
id: "answer",
|
|
679
|
+
invocationId: "run",
|
|
680
|
+
author: "agent",
|
|
681
|
+
content: { role: "model", parts: [{ text: "answer 0" }] },
|
|
682
|
+
};
|
|
683
|
+
});
|
|
684
|
+
const getCheckpointId = vi.fn(async () => {
|
|
685
|
+
await checkpoint.promise;
|
|
686
|
+
return "cp-1";
|
|
687
|
+
});
|
|
688
|
+
const load = vi.fn(async (): Promise<AdkThreadSnapshot> => {
|
|
689
|
+
if (++loadCount === 1) {
|
|
690
|
+
await initialLoad.promise;
|
|
691
|
+
return { messages: [] };
|
|
692
|
+
}
|
|
693
|
+
await loaded.promise;
|
|
694
|
+
return {
|
|
695
|
+
messages: [
|
|
696
|
+
{
|
|
697
|
+
id: "h-1",
|
|
698
|
+
type: "ai",
|
|
699
|
+
content: [{ type: "text", text: "loaded" }],
|
|
700
|
+
},
|
|
701
|
+
],
|
|
702
|
+
};
|
|
703
|
+
});
|
|
704
|
+
const runtime = await mountWithCheckpoint(
|
|
705
|
+
stream,
|
|
706
|
+
getCheckpointId,
|
|
707
|
+
true,
|
|
708
|
+
load,
|
|
709
|
+
);
|
|
710
|
+
|
|
711
|
+
await waitFor(() => expect(load).toHaveBeenCalledTimes(1));
|
|
712
|
+
await act(async () => {
|
|
713
|
+
initialLoad.resolve();
|
|
714
|
+
});
|
|
715
|
+
|
|
716
|
+
await act(async () => {
|
|
717
|
+
runtime.thread.append({
|
|
718
|
+
role: "user",
|
|
719
|
+
content: [{ type: "text", text: "question" }],
|
|
720
|
+
});
|
|
721
|
+
});
|
|
722
|
+
await waitFor(() =>
|
|
723
|
+
expect(JSON.stringify(runtime.thread.getState().messages)).toContain(
|
|
724
|
+
"answer 0",
|
|
725
|
+
),
|
|
726
|
+
);
|
|
727
|
+
const answer = runtime.thread
|
|
728
|
+
.getState()
|
|
729
|
+
.messages.find((m) => JSON.stringify(m).includes("answer 0"))!;
|
|
730
|
+
|
|
731
|
+
let refetch!: Promise<void>;
|
|
732
|
+
await act(async () => {
|
|
733
|
+
runtime.thread.getMessageById(answer.id).reload();
|
|
734
|
+
await Promise.resolve();
|
|
735
|
+
expect(getCheckpointId).toHaveBeenCalledTimes(1);
|
|
736
|
+
refetch = runtime.threads.reloadMainThread();
|
|
737
|
+
expect(load).toHaveBeenCalledTimes(2);
|
|
738
|
+
loaded.resolve();
|
|
739
|
+
await refetch;
|
|
740
|
+
});
|
|
741
|
+
await waitFor(() => expect(texts(runtime)).toEqual(["loaded"]));
|
|
742
|
+
await act(async () => {
|
|
743
|
+
runtime.thread.cancelRun();
|
|
744
|
+
});
|
|
745
|
+
await act(async () => {
|
|
746
|
+
checkpoint.resolve();
|
|
747
|
+
});
|
|
748
|
+
|
|
749
|
+
expect(stream).toHaveBeenCalledTimes(1);
|
|
750
|
+
expect(texts(runtime)).toEqual(["loaded"]);
|
|
751
|
+
});
|
|
752
|
+
|
|
753
|
+
it("sends a message sent while an edit looks up its checkpoint instead of the edit", async () => {
|
|
754
|
+
const checkpoint = deferred();
|
|
755
|
+
const stream = vi.fn(async function* (): AsyncGenerator<AdkEvent> {
|
|
756
|
+
yield {
|
|
757
|
+
id: "answer",
|
|
758
|
+
invocationId: "run",
|
|
759
|
+
author: "agent",
|
|
760
|
+
content: { role: "model", parts: [{ text: "answer" }] },
|
|
761
|
+
};
|
|
762
|
+
});
|
|
763
|
+
const runtime = await mountWithCheckpoint(stream, async () => {
|
|
764
|
+
await checkpoint.promise;
|
|
765
|
+
return "cp-1";
|
|
766
|
+
});
|
|
767
|
+
|
|
768
|
+
await act(async () => {
|
|
769
|
+
runtime.thread.append({
|
|
770
|
+
role: "user",
|
|
771
|
+
content: [{ type: "text", text: "question" }],
|
|
772
|
+
});
|
|
773
|
+
});
|
|
774
|
+
await waitFor(() =>
|
|
775
|
+
expect(runtime.thread.getState().isRunning).toBe(false),
|
|
776
|
+
);
|
|
777
|
+
const original = runtime.thread.getState().messages[0]!;
|
|
778
|
+
|
|
779
|
+
await act(async () => {
|
|
780
|
+
runtime.thread.append({
|
|
781
|
+
role: "user",
|
|
782
|
+
parentId: null,
|
|
783
|
+
sourceId: original.id,
|
|
784
|
+
content: [{ type: "text", text: "edited question" }],
|
|
785
|
+
});
|
|
786
|
+
});
|
|
787
|
+
await act(async () => {
|
|
788
|
+
runtime.thread.append({
|
|
789
|
+
role: "user",
|
|
790
|
+
content: [{ type: "text", text: "follow-up" }],
|
|
791
|
+
});
|
|
792
|
+
});
|
|
793
|
+
await act(async () => {
|
|
794
|
+
checkpoint.resolve();
|
|
795
|
+
});
|
|
796
|
+
await waitFor(() =>
|
|
797
|
+
expect(runtime.thread.getState().isRunning).toBe(false),
|
|
798
|
+
);
|
|
799
|
+
|
|
800
|
+
expect(stream).toHaveBeenCalledTimes(2);
|
|
801
|
+
expect(JSON.stringify(stream.mock.calls[1])).toContain("follow-up");
|
|
802
|
+
expect(JSON.stringify(stream.mock.calls)).not.toContain("edited question");
|
|
803
|
+
});
|
|
804
|
+
|
|
805
|
+
it("stops an edit that is still looking up its checkpoint when a staged edit replaces it", async () => {
|
|
806
|
+
const checkpoint = deferred();
|
|
807
|
+
const stream = vi.fn(async function* (): AsyncGenerator<AdkEvent> {
|
|
808
|
+
yield {
|
|
809
|
+
id: "answer",
|
|
810
|
+
invocationId: "run",
|
|
811
|
+
author: "agent",
|
|
812
|
+
content: { role: "model", parts: [{ text: "answer" }] },
|
|
813
|
+
};
|
|
814
|
+
});
|
|
815
|
+
const runtime = await mountWithCheckpoint(stream, async () => {
|
|
816
|
+
await checkpoint.promise;
|
|
817
|
+
return "cp-1";
|
|
818
|
+
});
|
|
819
|
+
|
|
820
|
+
await act(async () => {
|
|
821
|
+
runtime.thread.append({
|
|
822
|
+
role: "user",
|
|
823
|
+
content: [{ type: "text", text: "question" }],
|
|
824
|
+
});
|
|
825
|
+
});
|
|
826
|
+
await waitFor(() =>
|
|
827
|
+
expect(runtime.thread.getState().isRunning).toBe(false),
|
|
828
|
+
);
|
|
829
|
+
const original = runtime.thread.getState().messages[0]!;
|
|
830
|
+
|
|
831
|
+
await act(async () => {
|
|
832
|
+
runtime.thread.append({
|
|
833
|
+
role: "user",
|
|
834
|
+
parentId: null,
|
|
835
|
+
sourceId: original.id,
|
|
836
|
+
content: [{ type: "text", text: "edited question" }],
|
|
837
|
+
});
|
|
838
|
+
});
|
|
839
|
+
await act(async () => {
|
|
840
|
+
runtime.thread.append({
|
|
841
|
+
role: "user",
|
|
842
|
+
parentId: null,
|
|
843
|
+
sourceId: original.id,
|
|
844
|
+
content: [{ type: "text", text: "staged question" }],
|
|
845
|
+
startRun: false,
|
|
846
|
+
});
|
|
847
|
+
});
|
|
848
|
+
expect(runtime.thread.getState().isRunning).toBe(false);
|
|
849
|
+
await act(async () => {
|
|
850
|
+
checkpoint.resolve();
|
|
851
|
+
});
|
|
852
|
+
|
|
853
|
+
expect(stream).toHaveBeenCalledTimes(1);
|
|
854
|
+
expect(
|
|
855
|
+
runtime.thread
|
|
856
|
+
.getState()
|
|
857
|
+
.messages.map((m) => (m.content[0] as { text: string }).text),
|
|
858
|
+
).toEqual(["staged question"]);
|
|
859
|
+
});
|
|
143
860
|
});
|