@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.
Files changed (65) hide show
  1. package/README.md +12 -2
  2. package/dist/AdkClient.d.ts +4 -0
  3. package/dist/AdkClient.d.ts.map +1 -1
  4. package/dist/AdkClient.js +12 -7
  5. package/dist/AdkClient.js.map +1 -1
  6. package/dist/AdkSessionAdapter.d.ts.map +1 -1
  7. package/dist/AdkSessionAdapter.js +7 -7
  8. package/dist/AdkSessionAdapter.js.map +1 -1
  9. package/dist/AdkThreadController.d.ts +15 -0
  10. package/dist/AdkThreadController.d.ts.map +1 -0
  11. package/dist/AdkThreadController.js +35 -0
  12. package/dist/AdkThreadController.js.map +1 -0
  13. package/dist/adkThreadState.d.ts +54 -0
  14. package/dist/adkThreadState.d.ts.map +1 -0
  15. package/dist/adkThreadState.js +93 -0
  16. package/dist/adkThreadState.js.map +1 -0
  17. package/dist/convertToAdkMessages.js +1 -1
  18. package/dist/convertToAdkMessages.js.map +1 -1
  19. package/dist/sdkIdentity.js +1 -1
  20. package/dist/server/createAdkApiRoute.d.ts +37 -6
  21. package/dist/server/createAdkApiRoute.d.ts.map +1 -1
  22. package/dist/server/createAdkApiRoute.js +55 -5
  23. package/dist/server/createAdkApiRoute.js.map +1 -1
  24. package/dist/server/parseAdkRequest.d.ts +4 -1
  25. package/dist/server/parseAdkRequest.d.ts.map +1 -1
  26. package/dist/server/parseAdkRequest.js +5 -1
  27. package/dist/server/parseAdkRequest.js.map +1 -1
  28. package/dist/useAdkMessages.d.ts +9 -7
  29. package/dist/useAdkMessages.d.ts.map +1 -1
  30. package/dist/useAdkMessages.js +55 -77
  31. package/dist/useAdkMessages.js.map +1 -1
  32. package/dist/useAdkRuntime.d.ts +7 -1
  33. package/dist/useAdkRuntime.d.ts.map +1 -1
  34. package/dist/useAdkRuntime.js +134 -56
  35. package/dist/useAdkRuntime.js.map +1 -1
  36. package/package.json +5 -5
  37. package/src/AdkClient.test.ts +78 -2
  38. package/src/AdkClient.ts +24 -6
  39. package/src/AdkSessionAdapter.ts +1 -1
  40. package/src/AdkThreadController.test.ts +90 -0
  41. package/src/AdkThreadController.ts +45 -0
  42. package/src/adkThreadState.test.ts +207 -0
  43. package/src/adkThreadState.ts +124 -0
  44. package/src/convertToAdkMessages.test.ts +19 -0
  45. package/src/convertToAdkMessages.ts +1 -1
  46. package/src/hooks.test.tsx +1 -0
  47. package/src/server/createAdkApiRoute.controls.test.ts +66 -0
  48. package/src/server/createAdkApiRoute.test.ts +282 -0
  49. package/src/server/createAdkApiRoute.ts +119 -11
  50. package/src/server/parseAdkRequest.test.ts +11 -3
  51. package/src/server/parseAdkRequest.ts +7 -1
  52. package/src/useAdkMessages.test.ts +1 -0
  53. package/src/useAdkMessages.ts +61 -96
  54. package/src/useAdkRuntime.cancellation.test.tsx +4 -3
  55. package/src/useAdkRuntime.cloud-options.test.tsx +59 -0
  56. package/src/useAdkRuntime.refetch.test.tsx +548 -4
  57. package/src/useAdkRuntime.replacement.test.tsx +718 -1
  58. package/src/useAdkRuntime.ts +169 -73
  59. package/src/useAdkRuntimeApproval.test.tsx +87 -1
  60. package/dist/raceWithAbortSignal.d.ts +0 -2
  61. package/dist/raceWithAbortSignal.d.ts.map +0 -1
  62. package/dist/raceWithAbortSignal.js +0 -45
  63. package/dist/raceWithAbortSignal.js.map +0 -1
  64. package/src/raceWithAbortSignal.test.ts +0 -73
  65. package/src/raceWithAbortSignal.ts +0 -48
@@ -4,6 +4,7 @@ import {
4
4
  useMemo,
5
5
  useRef,
6
6
  useState,
7
+ useSyncExternalStore,
7
8
  } from "react";
8
9
  import {
9
10
  pickExternalStoreSharedOptions,
@@ -22,6 +23,8 @@ import {
22
23
  createAbortableThreadLoad,
23
24
  createCloudThreadListAdapterCreateFallback,
24
25
  isRecord,
26
+ RunLeases,
27
+ type RunLease,
25
28
  } from "@assistant-ui/core/internal";
26
29
  import {
27
30
  useCloudThreadListAdapter,
@@ -72,6 +75,7 @@ export type UseAdkRuntimeOptions = ExternalStoreSharedOptions & {
72
75
  */
73
76
  onThreadIdChange?: ((threadId: string | undefined) => void) | undefined;
74
77
  autoCancelPendingToolCalls?: boolean | undefined;
78
+ /** @deprecated Experimental since 2025-01-03. Not scheduled for removal; the API may change in any release. */
75
79
  unstable_allowCancellation?: boolean | undefined;
76
80
  getCheckpointId?: (
77
81
  threadId: string,
@@ -106,6 +110,11 @@ export type UseAdkRuntimeOptions = ExternalStoreSharedOptions & {
106
110
  }
107
111
  | undefined;
108
112
  cloud?: AssistantCloud | undefined;
113
+ /**
114
+ * Stable identity for the account or workspace owning Cloud runtime state.
115
+ * Provide it from the first render and change it when that scope changes.
116
+ */
117
+ scopeId?: string | undefined;
109
118
  /**
110
119
  * A `RemoteThreadListAdapter` to use instead of the cloud adapter.
111
120
  * Use with `createAdkSessionAdapter` for ADK session-backed persistence.
@@ -181,6 +190,7 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
181
190
  }, []);
182
191
 
183
192
  const {
193
+ controller,
184
194
  messages,
185
195
  stateDelta,
186
196
  agentInfo,
@@ -206,6 +216,22 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
206
216
  loadRef.current = load;
207
217
  }, [load]);
208
218
  const [loadController] = useState(createAbortableThreadLoad);
219
+ const initialLoadRef = useRef<{
220
+ promise: Promise<void>;
221
+ active: boolean;
222
+ snapshot: AdkThreadSnapshot | undefined;
223
+ } | null>(null);
224
+ const waitForInitialLoad = () => {
225
+ const load = initialLoadRef.current;
226
+ if (!load) return undefined;
227
+ return load.promise.then(() => ({
228
+ active:
229
+ load.active &&
230
+ (!threadListItem ||
231
+ aui.threads.getState().mainThreadId === threadListItem.getState().id),
232
+ snapshot: load.snapshot,
233
+ }));
234
+ };
209
235
  const messagesRef = useRef(messages);
210
236
  useInsertionEffect(() => {
211
237
  messagesRef.current = messages;
@@ -241,9 +267,25 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
241
267
  useInsertionEffect(() => {
242
268
  isRunningRef.current = effectiveIsRunning;
243
269
  }, [effectiveIsRunning]);
244
- const runGenerationRef = useRef(0);
270
+ const [runLeases] = useState(() => new RunLeases());
271
+ const reloadLookupRef = useRef<{
272
+ lease: RunLease;
273
+ beforeReload: AdkThreadSnapshot;
274
+ } | null>(null);
275
+
276
+ const runExclusive = async (
277
+ run: (isCurrent: () => boolean) => Promise<void>,
278
+ ) => {
279
+ const lease = runLeases.begin();
280
+ try {
281
+ setIsRunning(true);
282
+ await run(lease.isCurrent);
283
+ } finally {
284
+ if (lease.isCurrent()) setIsRunning(false);
285
+ }
286
+ };
245
287
 
246
- const handleSendMessage = async (
288
+ const handleSendMessage = (
247
289
  msgs: AdkMessage[],
248
290
  config: AdkSendMessageConfig,
249
291
  ) => {
@@ -257,13 +299,13 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
257
299
  }
258
300
  : config;
259
301
 
260
- const generation = ++runGenerationRef.current;
261
- try {
262
- setIsRunning(true);
263
- await sendMessage(msgs, continuationConfig);
264
- } finally {
265
- if (runGenerationRef.current === generation) setIsRunning(false);
266
- }
302
+ return runExclusive(() => sendMessage(msgs, continuationConfig));
303
+ };
304
+
305
+ const stopRun = () => {
306
+ runLeases.invalidate();
307
+ setIsRunning(false);
308
+ cancel();
267
309
  };
268
310
 
269
311
  const { approvals: toolApprovals, key: toolApprovalsKey } =
@@ -315,43 +357,20 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
315
357
  adkMessagesRef.current = messages;
316
358
  }, [messages]);
317
359
 
318
- const stagedMessagesRef = useRef(
319
- new Map<
320
- string,
321
- {
322
- message: AdkMessage & { id: string };
323
- runConfig: AppendMessage["runConfig"];
324
- }
325
- >(),
360
+ const stagedMessageCount = useSyncExternalStore(
361
+ controller.subscribe,
362
+ controller.getStagedMessageCount,
363
+ controller.getStagedMessageCount,
326
364
  );
327
- const [stagedMessageCount, setStagedMessageCount] = useState(0);
328
365
  const hasStagedMessages = stagedMessageCount > 0;
329
366
 
330
- const getStagedRun = (parentId: string | null) => {
331
- if (!parentId || !stagedMessagesRef.current.has(parentId)) return null;
332
-
333
- const staged: AdkMessage[] = [];
334
- for (const message of adkMessagesRef.current) {
335
- if (message.id && stagedMessagesRef.current.has(message.id)) {
336
- staged.push(stagedMessagesRef.current.get(message.id)!.message);
337
- }
338
- if (message.id === parentId) break;
339
- }
340
-
341
- return {
342
- messages: staged,
343
- runConfig: stagedMessagesRef.current.get(parentId)!.runConfig,
344
- };
345
- };
346
-
347
367
  const stageUserMessage = (msg: AppendMessage) => {
348
368
  const stagedMessage = toAdkUserMessage(msg);
349
- stagedMessagesRef.current.set(stagedMessage.id, {
350
- message: stagedMessage,
351
- runConfig: msg.runConfig,
369
+ controller.dispatch({
370
+ type: "staged.stage",
371
+ entry: { message: stagedMessage, runConfig: msg.runConfig },
352
372
  });
353
- setStagedMessageCount(stagedMessagesRef.current.size);
354
- const nextMessages = [...adkMessagesRef.current, stagedMessage];
373
+ const nextMessages = [...controller.getState().messages, stagedMessage];
355
374
  adkMessagesRef.current = nextMessages;
356
375
  setMessages(nextMessages);
357
376
  };
@@ -394,7 +413,14 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
394
413
  messagesRef.current !== messagesAtLoadStart)
395
414
  )
396
415
  return;
416
+ reloadLookupRef.current = null;
397
417
  applySnapshot(snapshot);
418
+ messagesRef.current = snapshot.messages;
419
+ adkMessagesRef.current = snapshot.messages;
420
+ longRunningToolIdsRef.current = snapshot.longRunningToolIds ?? [];
421
+ if (purpose === "initial" && initialLoadRef.current) {
422
+ initialLoadRef.current.snapshot = snapshot;
423
+ }
398
424
  },
399
425
  onSettled: () => {
400
426
  setIsLoadingThread(false);
@@ -408,11 +434,26 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
408
434
  );
409
435
 
410
436
  useReplaySafeEffect(() => {
411
- runLoad();
437
+ let release!: () => void;
438
+ const barrier = {
439
+ promise: new Promise<void>((resolve) => {
440
+ release = resolve;
441
+ }),
442
+ active: true,
443
+ snapshot: undefined as AdkThreadSnapshot | undefined,
444
+ };
445
+ initialLoadRef.current = barrier;
446
+ const settle = () => {
447
+ if (initialLoadRef.current === barrier) initialLoadRef.current = null;
448
+ release();
449
+ };
450
+ void runLoad().then(settle, settle);
412
451
  return () => {
452
+ barrier.active = false;
413
453
  // Whatever is current, not this effect's own controller: a refetch swaps
414
454
  // the ref, and one in flight at unmount must be aborted too.
415
455
  loadController.abort();
456
+ settle();
416
457
  setIsLoadingThread(false);
417
458
  };
418
459
  }, [threadListItem]);
@@ -435,9 +476,17 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
435
476
  authRequests,
436
477
  escalated,
437
478
  messageMetadata,
438
- send: handleSendMessage,
479
+ send: (messages, config) => {
480
+ const initialLoad = waitForInitialLoad();
481
+ if (!initialLoad) return handleSendMessage(messages, config);
482
+ return initialLoad.then(({ active }) =>
483
+ active ? handleSendMessage(messages, config) : undefined,
484
+ );
485
+ },
439
486
  }),
440
487
  onNew: async (msg) => {
488
+ const initialLoad = await waitForInitialLoad();
489
+ if (initialLoad && !initialLoad.active) return;
441
490
  if (!(msg.startRun ?? msg.role === "user")) {
442
491
  stageUserMessage(msg);
443
492
  return;
@@ -445,7 +494,11 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
445
494
 
446
495
  const cancellations =
447
496
  autoCancelPendingToolCalls !== false
448
- ? getPendingCancellations(messages, longRunningToolIds)
497
+ ? getPendingCancellations(
498
+ initialLoad?.snapshot?.messages ?? messagesRef.current,
499
+ initialLoad?.snapshot?.longRunningToolIds ??
500
+ longRunningToolIdsRef.current,
501
+ )
449
502
  : [];
450
503
 
451
504
  return handleSendMessage(
@@ -462,6 +515,9 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
462
515
  },
463
516
  onEdit: getCheckpointId
464
517
  ? async (msg) => {
518
+ const initialLoad = waitForInitialLoad();
519
+ if (initialLoad && !(await initialLoad).active) return;
520
+ stopRun();
465
521
  const truncated = truncateAdkMessages(
466
522
  threadMessagesRef.current,
467
523
  msg.parentId,
@@ -469,44 +525,41 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
469
525
  replaceMessages(truncated);
470
526
  if (!(msg.startRun ?? msg.role === "user")) {
471
527
  const stagedMessage = toAdkUserMessage(msg);
472
- stagedMessagesRef.current.set(stagedMessage.id, {
473
- message: stagedMessage,
474
- runConfig: msg.runConfig,
528
+ controller.dispatch({
529
+ type: "staged.stage",
530
+ entry: { message: stagedMessage, runConfig: msg.runConfig },
475
531
  });
476
- setStagedMessageCount(stagedMessagesRef.current.size);
477
532
  const nextMessages = [...truncated, stagedMessage];
478
533
  adkMessagesRef.current = nextMessages;
479
534
  setMessages(nextMessages);
480
535
  return;
481
536
  }
537
+ const editedMessage = toAdkUserMessage(msg);
538
+ setMessages([...truncated, editedMessage]);
482
539
  const externalId = aui.threadListItem.getState().externalId;
483
- const checkpointId = externalId
484
- ? await getCheckpointId(externalId, truncated)
485
- : null;
486
- return handleSendMessage(
487
- [
488
- {
489
- id: generateId(),
490
- type: "human",
491
- content: getMessageContent(msg),
492
- },
493
- ],
494
- {
540
+ return runExclusive(async (isCurrent) => {
541
+ const checkpointId = externalId
542
+ ? await getCheckpointId(externalId, truncated)
543
+ : null;
544
+ if (!isCurrent()) return;
545
+ await sendMessage([editedMessage], {
495
546
  runConfig: msg.runConfig,
496
547
  ...(checkpointId && { checkpointId }),
497
- },
498
- );
548
+ });
549
+ });
499
550
  }
500
551
  : undefined,
501
552
  ...(getCheckpointId || hasStagedMessages
502
553
  ? {
503
554
  onReload: async (parentId, config) => {
504
- const stagedRun = getStagedRun(parentId);
555
+ const initialLoad = waitForInitialLoad();
556
+ if (initialLoad && !(await initialLoad).active) return;
557
+ const stagedRun = controller.getStagedRun(parentId);
505
558
  if (stagedRun) {
506
- for (const message of stagedRun.messages) {
507
- stagedMessagesRef.current.delete(message.id);
508
- }
509
- setStagedMessageCount(stagedMessagesRef.current.size);
559
+ controller.dispatch({
560
+ type: "staged.unstage",
561
+ ids: stagedRun.messages.map((message) => message.id!),
562
+ });
510
563
  return handleSendMessage(stagedRun.messages, {
511
564
  runConfig: config.runConfig ?? stagedRun.runConfig,
512
565
  });
@@ -515,18 +568,48 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
515
568
  if (!getCheckpointId)
516
569
  throw new Error("Runtime does not support reloading messages.");
517
570
 
571
+ stopRun();
572
+ const beforeReload: AdkThreadSnapshot = {
573
+ messages: adkMessagesRef.current,
574
+ longRunningToolIds,
575
+ toolConfirmations,
576
+ authRequests,
577
+ escalated,
578
+ messageMetadata,
579
+ stateDelta,
580
+ artifactDelta,
581
+ agentInfo,
582
+ };
518
583
  const truncated = truncateAdkMessages(
519
584
  threadMessagesRef.current,
520
585
  parentId,
521
586
  );
522
587
  replaceMessages(truncated);
523
588
  const externalId = aui.threadListItem.getState().externalId;
524
- const checkpointId = externalId
525
- ? await getCheckpointId(externalId, truncated)
526
- : null;
527
- return handleSendMessage([], {
528
- runConfig: config.runConfig,
529
- ...(checkpointId && { checkpointId }),
589
+ return runExclusive(async (isCurrent) => {
590
+ const lookup = {
591
+ lease: runLeases.current(),
592
+ beforeReload,
593
+ };
594
+ reloadLookupRef.current = lookup;
595
+ let checkpointId: string | null;
596
+ try {
597
+ checkpointId = externalId
598
+ ? await getCheckpointId(externalId, truncated)
599
+ : null;
600
+ } catch (error) {
601
+ if (isCurrent() && reloadLookupRef.current === lookup)
602
+ applySnapshot(beforeReload);
603
+ throw error;
604
+ } finally {
605
+ if (reloadLookupRef.current === lookup)
606
+ reloadLookupRef.current = null;
607
+ }
608
+ if (!isCurrent()) return;
609
+ await sendMessage([], {
610
+ runConfig: config.runConfig,
611
+ ...(checkpointId && { checkpointId }),
612
+ });
530
613
  });
531
614
  },
532
615
  }
@@ -538,6 +621,8 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
538
621
  isError,
539
622
  artifact,
540
623
  }) => {
624
+ const initialLoad = waitForInitialLoad();
625
+ if (initialLoad && !(await initialLoad).active) return;
541
626
  await handleSendMessage(
542
627
  [
543
628
  {
@@ -554,6 +639,8 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
554
639
  );
555
640
  },
556
641
  onRespondToToolApproval: async (options) => {
642
+ const initialLoad = waitForInitialLoad();
643
+ if (initialLoad && !(await initialLoad).active) return;
557
644
  await handleSendMessage(
558
645
  [
559
646
  toAdkToolConfirmationReply(
@@ -566,7 +653,14 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
566
653
  },
567
654
  onCancel: unstable_allowCancellation
568
655
  ? async () => {
569
- cancel();
656
+ const lookup = reloadLookupRef.current;
657
+ const beforeReload = lookup?.lease.isCurrent()
658
+ ? lookup.beforeReload
659
+ : undefined;
660
+ stopRun();
661
+ // A reload stopped before it sent leaves the ADK session holding the
662
+ // turn it removed, so the thread shows that turn again.
663
+ if (beforeReload) applySnapshot(beforeReload);
570
664
  }
571
665
  : undefined,
572
666
  ...(load !== undefined && {
@@ -579,6 +673,7 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
579
673
 
580
674
  export const useAdkRuntime = ({
581
675
  cloud,
676
+ scopeId,
582
677
  sessionAdapter,
583
678
  create,
584
679
  delete: deleteFn,
@@ -589,6 +684,7 @@ export const useAdkRuntime = ({
589
684
  const cloudAdapter = useCloudThreadListAdapter({
590
685
  sdk: ADK_SDK,
591
686
  cloud,
687
+ scopeId,
592
688
  create: createCloudThreadListAdapterCreateFallback(
593
689
  create,
594
690
  aui.threadListItem,
@@ -1,3 +1,4 @@
1
+ /** @vitest-environment jsdom */
1
2
  import { act, renderHook } from "@testing-library/react";
2
3
  import { afterEach, describe, expect, it, vi } from "vitest";
3
4
  import type {
@@ -7,13 +8,19 @@ import type {
7
8
  ThreadMessage,
8
9
  ToolCallMessagePart,
9
10
  } from "@assistant-ui/core";
10
- import type { AdkMessage, AdkSendMessageConfig } from "./types";
11
+ import type {
12
+ AdkMessage,
13
+ AdkSendMessageConfig,
14
+ AdkThreadSnapshot,
15
+ } from "./types";
11
16
 
12
17
  const mocks = vi.hoisted(() => {
13
18
  const threadListItem = {
14
19
  source: null as object | null,
20
+ id: "thread-a",
15
21
  externalId: undefined as string | undefined,
16
22
  getState: () => ({
23
+ id: threadListItem.id,
17
24
  externalId: threadListItem.externalId,
18
25
  }),
19
26
  initialize: vi.fn(),
@@ -29,6 +36,14 @@ const mocks = vi.hoisted(() => {
29
36
  };
30
37
  });
31
38
 
39
+ const mockController = {
40
+ subscribe: () => () => {},
41
+ getStagedMessageCount: () => 0,
42
+ getState: () => ({ messages: mocks.messages }),
43
+ dispatch: vi.fn(),
44
+ getStagedRun: () => null,
45
+ };
46
+
32
47
  vi.mock("@assistant-ui/core/react", async (importOriginal) => ({
33
48
  ...(await importOriginal<typeof import("@assistant-ui/core/react")>()),
34
49
  useCloudThreadListAdapter: () => ({}),
@@ -44,6 +59,9 @@ vi.mock("@assistant-ui/store", async (importOriginal) => ({
44
59
  ...(await importOriginal<typeof import("@assistant-ui/store")>()),
45
60
  useAui: () => ({
46
61
  threadListItem: mocks.threadListItem,
62
+ threads: {
63
+ getState: () => ({ mainThreadId: mocks.threadListItem.id }),
64
+ },
47
65
  }),
48
66
  }));
49
67
 
@@ -62,6 +80,7 @@ vi.mock("./useAdkMessages", async (importOriginal) => {
62
80
  }
63
81
  };
64
82
  return {
83
+ controller: mockController,
65
84
  messages: mocks.messages,
66
85
  stateDelta: {},
67
86
  agentInfo: {},
@@ -105,6 +124,10 @@ type RuntimeAdapter = {
105
124
  onRespondToToolApproval?: (
106
125
  options: RespondToToolApprovalOptions,
107
126
  ) => Promise<void> | void;
127
+ onReload?: (
128
+ parentId: string | null,
129
+ config: { runConfig?: AppendMessage["runConfig"] },
130
+ ) => Promise<void> | void;
108
131
  onRefetchThread?: () => Promise<void> | void;
109
132
  };
110
133
 
@@ -161,6 +184,69 @@ afterEach(() => {
161
184
  });
162
185
 
163
186
  describe("useAdkRuntime tool approvals", () => {
187
+ it.each([
188
+ "reload",
189
+ "tool result",
190
+ "approval response",
191
+ "extras send",
192
+ ] as const)("waits for the initial load before %s", async (route) => {
193
+ let resolveLoad!: (snapshot: AdkThreadSnapshot) => void;
194
+ const pendingLoad = new Promise<AdkThreadSnapshot>((resolve) => {
195
+ resolveLoad = resolve;
196
+ });
197
+ const load = vi.fn(() => pendingLoad);
198
+ mocks.threadListItem.source = {};
199
+ mocks.threadListItem.externalId = "thread-a";
200
+ renderHook(() =>
201
+ useAdkRuntime({
202
+ stream: vi.fn(),
203
+ load,
204
+ getCheckpointId: vi.fn(async () => null),
205
+ }),
206
+ );
207
+ expect(load).toHaveBeenCalledOnce();
208
+
209
+ let action: Promise<void>;
210
+ switch (route) {
211
+ case "reload":
212
+ action = Promise.resolve(latestAdapter().onReload!(null, {}));
213
+ break;
214
+ case "tool result":
215
+ action = Promise.resolve(
216
+ latestAdapter().onAddToolResult!({
217
+ messageId: "ai-1",
218
+ toolCallId: "tool-a",
219
+ toolName: "lookup",
220
+ result: { value: "done" },
221
+ isError: false,
222
+ }),
223
+ );
224
+ break;
225
+ case "approval response":
226
+ action = Promise.resolve(
227
+ latestAdapter().onRespondToToolApproval!({
228
+ approvalId: CONFIRMATION_CALL,
229
+ approved: true,
230
+ }),
231
+ );
232
+ break;
233
+ case "extras send":
234
+ action = latestAdapter().extras.send(
235
+ [{ id: "new-user", type: "human", content: "new question" }],
236
+ {},
237
+ );
238
+ }
239
+ void action.catch(() => {});
240
+ await Promise.resolve();
241
+ expect(mocks.sendMessage).not.toHaveBeenCalled();
242
+
243
+ await act(async () => {
244
+ resolveLoad({ messages: [makeConfirmationRequest()] });
245
+ await action;
246
+ });
247
+ expect(mocks.sendMessage).toHaveBeenCalledOnce();
248
+ });
249
+
164
250
  it("resumes a delayed tool result with its originating run config", async () => {
165
251
  const runConfigA = { custom: { model: "model-a" } };
166
252
  const runConfigB = { custom: { model: "model-b" } };
@@ -1,2 +0,0 @@
1
- export declare const raceWithAbortSignal: <T>(signal: AbortSignal | undefined, operation: () => T | PromiseLike<T>) => Promise<T>;
2
- //# sourceMappingURL=raceWithAbortSignal.d.ts.map
@@ -1 +0,0 @@
1
- {"version":3,"file":"raceWithAbortSignal.d.ts","sourceRoot":"","sources":["../src/raceWithAbortSignal.ts"],"names":[],"mappings":"AAOA,eAAO,MAAM,mBAAmB,GAAI,CAAC,EACnC,QAAQ,WAAW,GAAG,SAAS,EAC/B,WAAW,MAAM,CAAC,GAAG,WAAW,CAAC,CAAC,CAAC,KAClC,OAAO,CAAC,CAAC,CAqCX,CAAC"}
@@ -1,45 +0,0 @@
1
- //#region src/raceWithAbortSignal.ts
2
- const getAbortReason = (signal) => {
3
- if (signal.reason !== void 0) return signal.reason;
4
- const error = /* @__PURE__ */ new Error("The operation was aborted");
5
- error.name = "AbortError";
6
- return error;
7
- };
8
- const raceWithAbortSignal = (signal, operation) => {
9
- if (!signal) try {
10
- return Promise.resolve(operation());
11
- } catch (error) {
12
- return Promise.reject(error);
13
- }
14
- if (signal.aborted) return Promise.reject(getAbortReason(signal));
15
- return new Promise((resolve, reject) => {
16
- let settled = false;
17
- const cleanup = () => signal.removeEventListener("abort", handleAbort);
18
- const resolveOnce = (value) => {
19
- if (settled) return;
20
- settled = true;
21
- cleanup();
22
- resolve(value);
23
- };
24
- const rejectOnce = (error) => {
25
- if (settled) return;
26
- settled = true;
27
- cleanup();
28
- reject(error);
29
- };
30
- const handleAbort = () => rejectOnce(getAbortReason(signal));
31
- signal.addEventListener("abort", handleAbort, { once: true });
32
- let result;
33
- try {
34
- result = operation();
35
- } catch (error) {
36
- rejectOnce(error);
37
- return;
38
- }
39
- Promise.resolve(result).then(resolveOnce, rejectOnce);
40
- });
41
- };
42
- //#endregion
43
- export { raceWithAbortSignal };
44
-
45
- //# sourceMappingURL=raceWithAbortSignal.js.map
@@ -1 +0,0 @@
1
- {"version":3,"file":"raceWithAbortSignal.js","names":[],"sources":["../src/raceWithAbortSignal.ts"],"sourcesContent":["const getAbortReason = (signal: AbortSignal): unknown => {\n if (signal.reason !== undefined) return signal.reason;\n const error = new Error(\"The operation was aborted\");\n error.name = \"AbortError\";\n return error;\n};\n\nexport const raceWithAbortSignal = <T>(\n signal: AbortSignal | undefined,\n operation: () => T | PromiseLike<T>,\n): Promise<T> => {\n if (!signal) {\n try {\n return Promise.resolve(operation());\n } catch (error) {\n return Promise.reject(error);\n }\n }\n if (signal.aborted) return Promise.reject(getAbortReason(signal));\n\n return new Promise<T>((resolve, reject) => {\n let settled = false;\n const cleanup = () => signal.removeEventListener(\"abort\", handleAbort);\n const resolveOnce = (value: T) => {\n if (settled) return;\n settled = true;\n cleanup();\n resolve(value);\n };\n const rejectOnce = (error: unknown) => {\n if (settled) return;\n settled = true;\n cleanup();\n reject(error);\n };\n const handleAbort = () => rejectOnce(getAbortReason(signal));\n\n signal.addEventListener(\"abort\", handleAbort, { once: true });\n let result: T | PromiseLike<T>;\n try {\n result = operation();\n } catch (error) {\n rejectOnce(error);\n return;\n }\n Promise.resolve(result).then(resolveOnce, rejectOnce);\n });\n};\n"],"mappings":";AAAA,MAAM,kBAAkB,WAAiC;CACvD,IAAI,OAAO,WAAW,KAAA,GAAW,OAAO,OAAO;CAC/C,MAAM,wBAAQ,IAAI,MAAM,2BAA2B;CACnD,MAAM,OAAO;CACb,OAAO;AACT;AAEA,MAAa,uBACX,QACA,cACe;CACf,IAAI,CAAC,QACH,IAAI;EACF,OAAO,QAAQ,QAAQ,UAAU,CAAC;CACpC,SAAS,OAAO;EACd,OAAO,QAAQ,OAAO,KAAK;CAC7B;CAEF,IAAI,OAAO,SAAS,OAAO,QAAQ,OAAO,eAAe,MAAM,CAAC;CAEhE,OAAO,IAAI,SAAY,SAAS,WAAW;EACzC,IAAI,UAAU;EACd,MAAM,gBAAgB,OAAO,oBAAoB,SAAS,WAAW;EACrE,MAAM,eAAe,UAAa;GAChC,IAAI,SAAS;GACb,UAAU;GACV,QAAQ;GACR,QAAQ,KAAK;EACf;EACA,MAAM,cAAc,UAAmB;GACrC,IAAI,SAAS;GACb,UAAU;GACV,QAAQ;GACR,OAAO,KAAK;EACd;EACA,MAAM,oBAAoB,WAAW,eAAe,MAAM,CAAC;EAE3D,OAAO,iBAAiB,SAAS,aAAa,EAAE,MAAM,KAAK,CAAC;EAC5D,IAAI;EACJ,IAAI;GACF,SAAS,UAAU;EACrB,SAAS,OAAO;GACd,WAAW,KAAK;GAChB;EACF;EACA,QAAQ,QAAQ,MAAM,CAAC,CAAC,KAAK,aAAa,UAAU;CACtD,CAAC;AACH"}
@@ -1,73 +0,0 @@
1
- import { describe, expect, it, vi } from "vitest";
2
- import { raceWithAbortSignal } from "./raceWithAbortSignal";
3
-
4
- describe("raceWithAbortSignal", () => {
5
- it("invokes the operation synchronously without a signal", async () => {
6
- const order: string[] = [];
7
-
8
- const result = raceWithAbortSignal(undefined, () => {
9
- order.push("operation");
10
- return "done";
11
- });
12
- order.push("after");
13
-
14
- expect(order).toEqual(["operation", "after"]);
15
- await expect(result).resolves.toBe("done");
16
- });
17
-
18
- it("converts a synchronous operation error to a rejection", async () => {
19
- const error = new Error("failed");
20
-
21
- const result = raceWithAbortSignal(undefined, () => {
22
- throw error;
23
- });
24
-
25
- await expect(result).rejects.toBe(error);
26
- });
27
-
28
- it("rejects a pending operation with the abort reason", async () => {
29
- const controller = new AbortController();
30
- const reason = new Error("cancelled");
31
- let resolveOperation!: (value: string) => void;
32
- const operation = new Promise<string>((resolve) => {
33
- resolveOperation = resolve;
34
- });
35
-
36
- const result = raceWithAbortSignal(controller.signal, () => operation);
37
- controller.abort(reason);
38
-
39
- await expect(result).rejects.toBe(reason);
40
- resolveOperation("late result");
41
- });
42
-
43
- it("rejects before invoking an operation for an already aborted signal", async () => {
44
- const controller = new AbortController();
45
- const reason = new Error("already cancelled");
46
- const operation = vi.fn(() => "done");
47
- controller.abort(reason);
48
-
49
- const result = raceWithAbortSignal(controller.signal, operation);
50
-
51
- await expect(result).rejects.toBe(reason);
52
- expect(operation).not.toHaveBeenCalled();
53
- });
54
-
55
- it("removes the abort listener after the operation settles", async () => {
56
- const controller = new AbortController();
57
- const removeEventListener = vi.spyOn(
58
- controller.signal,
59
- "removeEventListener",
60
- );
61
- const result = raceWithAbortSignal(controller.signal, () => "done");
62
-
63
- await expect(result).resolves.toBe("done");
64
- expect(removeEventListener).toHaveBeenCalledOnce();
65
- expect(removeEventListener).toHaveBeenCalledWith(
66
- "abort",
67
- expect.any(Function),
68
- );
69
-
70
- controller.abort(new Error("late abort"));
71
- await expect(result).resolves.toBe("done");
72
- });
73
- });