@assistant-ui/react-google-adk 0.0.28 → 0.0.30
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/AdkClient.d.ts.map +1 -1
- package/dist/AdkClient.js +2 -1
- package/dist/AdkClient.js.map +1 -1
- package/dist/AdkEventAccumulator.d.ts +1 -1
- package/dist/AdkEventAccumulator.d.ts.map +1 -1
- package/dist/AdkEventAccumulator.js +33 -12
- package/dist/AdkEventAccumulator.js.map +1 -1
- package/dist/convertToAdkMessages.d.ts +10 -1
- package/dist/convertToAdkMessages.d.ts.map +1 -1
- package/dist/convertToAdkMessages.js +20 -6
- package/dist/convertToAdkMessages.js.map +1 -1
- package/dist/hooks.d.ts.map +1 -1
- package/dist/hooks.js +6 -8
- package/dist/hooks.js.map +1 -1
- package/dist/parseAdkEvent.d.ts.map +1 -1
- package/dist/parseAdkEvent.js +15 -2
- package/dist/parseAdkEvent.js.map +1 -1
- package/dist/sdkIdentity.d.ts +6 -0
- package/dist/sdkIdentity.d.ts.map +1 -0
- package/dist/sdkIdentity.js +9 -0
- package/dist/sdkIdentity.js.map +1 -0
- package/dist/server/parseAdkRequest.d.ts.map +1 -1
- package/dist/server/parseAdkRequest.js +2 -1
- package/dist/server/parseAdkRequest.js.map +1 -1
- package/dist/toAdkFunctionResponse.d.ts +6 -0
- package/dist/toAdkFunctionResponse.d.ts.map +1 -0
- package/dist/toAdkFunctionResponse.js +11 -0
- package/dist/toAdkFunctionResponse.js.map +1 -0
- package/dist/useAdkMessages.d.ts.map +1 -1
- package/dist/useAdkMessages.js +22 -10
- package/dist/useAdkMessages.js.map +1 -1
- package/dist/useAdkRuntime.d.ts.map +1 -1
- package/dist/useAdkRuntime.js +6 -2
- package/dist/useAdkRuntime.js.map +1 -1
- package/package.json +6 -6
- package/src/AdkClient.test.ts +192 -2
- package/src/AdkClient.ts +2 -1
- package/src/AdkEventAccumulator.test.ts +297 -0
- package/src/AdkEventAccumulator.ts +35 -5
- package/src/AdkSessionAdapter.test.ts +88 -0
- package/src/convertToAdkMessages.test.ts +70 -0
- package/src/convertToAdkMessages.ts +19 -4
- package/src/hooks.render.test.tsx +74 -0
- package/src/hooks.test.tsx +33 -0
- package/src/hooks.ts +17 -22
- package/src/parseAdkEvent.ts +36 -7
- package/src/sdkIdentity.ts +9 -0
- package/src/server/parseAdkRequest.test.ts +63 -0
- package/src/server/parseAdkRequest.ts +2 -1
- package/src/toAdkFunctionResponse.test.ts +46 -0
- package/src/toAdkFunctionResponse.ts +18 -0
- package/src/useAdkMessages.test.ts +247 -0
- package/src/useAdkMessages.ts +51 -15
- package/src/useAdkRuntime.replacement.test.tsx +138 -0
- package/src/useAdkRuntime.ts +5 -1
- package/src/useAdkRuntimeApproval.test.tsx +0 -1
|
@@ -22,6 +22,7 @@ import {
|
|
|
22
22
|
import { projectAdkToolApprovals } from "./adkToolApproval";
|
|
23
23
|
import { createAdkStream } from "./AdkClient";
|
|
24
24
|
import { AdkEventAccumulator } from "./AdkEventAccumulator";
|
|
25
|
+
import { getPendingCancellations } from "./convertToAdkMessages";
|
|
25
26
|
import type { AdkEvent, AdkMessage, AdkStreamCallback } from "./types";
|
|
26
27
|
|
|
27
28
|
afterEach(() => {
|
|
@@ -29,6 +30,53 @@ afterEach(() => {
|
|
|
29
30
|
vi.unstubAllGlobals();
|
|
30
31
|
});
|
|
31
32
|
|
|
33
|
+
describe("optimistic tool outcomes", () => {
|
|
34
|
+
it.each([false, true])(
|
|
35
|
+
"preserves failures alongside successful results (batch: %s)",
|
|
36
|
+
(batch) => {
|
|
37
|
+
const failed: AdkMessage = {
|
|
38
|
+
id: "failed",
|
|
39
|
+
type: "tool",
|
|
40
|
+
name: "search",
|
|
41
|
+
tool_call_id: "tc-error",
|
|
42
|
+
content: "denied",
|
|
43
|
+
status: "error",
|
|
44
|
+
};
|
|
45
|
+
const succeeded: AdkMessage = {
|
|
46
|
+
id: "succeeded",
|
|
47
|
+
type: "tool",
|
|
48
|
+
name: "search",
|
|
49
|
+
tool_call_id: "tc-ok",
|
|
50
|
+
content: "found",
|
|
51
|
+
status: "success",
|
|
52
|
+
};
|
|
53
|
+
const events = batch
|
|
54
|
+
? messagesToEvents([
|
|
55
|
+
failed,
|
|
56
|
+
succeeded,
|
|
57
|
+
{ id: "human", type: "human", content: "continue" },
|
|
58
|
+
])
|
|
59
|
+
: [messageToEvent(failed), messageToEvent(succeeded)];
|
|
60
|
+
const acc = new AdkEventAccumulator();
|
|
61
|
+
for (const event of events) acc.processEvent(event);
|
|
62
|
+
expect(
|
|
63
|
+
acc.getMessages().filter((message) => message.type === "tool"),
|
|
64
|
+
).toMatchObject([
|
|
65
|
+
{
|
|
66
|
+
tool_call_id: "tc-error",
|
|
67
|
+
status: "error",
|
|
68
|
+
content: JSON.stringify({ error: "denied" }),
|
|
69
|
+
},
|
|
70
|
+
{
|
|
71
|
+
tool_call_id: "tc-ok",
|
|
72
|
+
status: "success",
|
|
73
|
+
content: JSON.stringify({ result: "found" }),
|
|
74
|
+
},
|
|
75
|
+
]);
|
|
76
|
+
},
|
|
77
|
+
);
|
|
78
|
+
});
|
|
79
|
+
|
|
32
80
|
describe("ADK runtime callbacks", () => {
|
|
33
81
|
it.each(["onAgentTransfer", "onCustomEvent", "onError"] as const)(
|
|
34
82
|
"continues streaming when %s throws",
|
|
@@ -131,6 +179,105 @@ describe("ADK runtime callbacks", () => {
|
|
|
131
179
|
});
|
|
132
180
|
|
|
133
181
|
describe("ADK stream lifecycle", () => {
|
|
182
|
+
it("settles a superseded send while its stream is still opening", async () => {
|
|
183
|
+
const signals: AbortSignal[] = [];
|
|
184
|
+
const parked = new Promise<AsyncGenerator<AdkEvent>>(() => {});
|
|
185
|
+
let calls = 0;
|
|
186
|
+
const stream = vi.fn(function (_messages, { abortSignal }) {
|
|
187
|
+
signals.push(abortSignal);
|
|
188
|
+
if (calls++ === 0) return parked;
|
|
189
|
+
return (async function* () {
|
|
190
|
+
yield {
|
|
191
|
+
id: "event-1",
|
|
192
|
+
invocationId: "run-1",
|
|
193
|
+
author: "agent",
|
|
194
|
+
content: { role: "model", parts: [{ text: "done-1" }] },
|
|
195
|
+
};
|
|
196
|
+
})();
|
|
197
|
+
}) satisfies AdkStreamCallback;
|
|
198
|
+
const { result } = renderHook(() => useAdkMessages({ stream }));
|
|
199
|
+
|
|
200
|
+
let firstSend!: Promise<void>;
|
|
201
|
+
act(() => {
|
|
202
|
+
firstSend = result.current.sendMessage(
|
|
203
|
+
[{ id: "user-1", type: "human", content: "first" }],
|
|
204
|
+
{},
|
|
205
|
+
);
|
|
206
|
+
});
|
|
207
|
+
await vi.waitFor(() => expect(stream).toHaveBeenCalledOnce());
|
|
208
|
+
|
|
209
|
+
let secondSend!: Promise<void>;
|
|
210
|
+
act(() => {
|
|
211
|
+
secondSend = result.current.sendMessage(
|
|
212
|
+
[{ id: "user-2", type: "human", content: "second" }],
|
|
213
|
+
{},
|
|
214
|
+
);
|
|
215
|
+
});
|
|
216
|
+
|
|
217
|
+
await act(async () => {
|
|
218
|
+
await Promise.all([firstSend, secondSend]);
|
|
219
|
+
});
|
|
220
|
+
expect(signals[0]?.aborted).toBe(true);
|
|
221
|
+
expect(signals[1]?.aborted).toBe(false);
|
|
222
|
+
expect(result.current.messages.at(-1)).toMatchObject({
|
|
223
|
+
type: "ai",
|
|
224
|
+
content: [{ type: "text", text: "done-1" }],
|
|
225
|
+
});
|
|
226
|
+
});
|
|
227
|
+
|
|
228
|
+
it("aborts and settles a superseded stream that stops yielding", async () => {
|
|
229
|
+
const signals: AbortSignal[] = [];
|
|
230
|
+
const parked = new Promise<void>(() => {});
|
|
231
|
+
let calls = 0;
|
|
232
|
+
const stream = vi.fn(function (
|
|
233
|
+
_messages,
|
|
234
|
+
{ abortSignal },
|
|
235
|
+
): AsyncGenerator<AdkEvent> {
|
|
236
|
+
signals.push(abortSignal);
|
|
237
|
+
if (calls++ === 0) {
|
|
238
|
+
return (async function* () {
|
|
239
|
+
await parked;
|
|
240
|
+
})();
|
|
241
|
+
}
|
|
242
|
+
return (async function* () {
|
|
243
|
+
yield {
|
|
244
|
+
id: "event-1",
|
|
245
|
+
invocationId: "run-1",
|
|
246
|
+
author: "agent",
|
|
247
|
+
content: { role: "model", parts: [{ text: "done-1" }] },
|
|
248
|
+
};
|
|
249
|
+
})();
|
|
250
|
+
}) satisfies AdkStreamCallback;
|
|
251
|
+
const { result } = renderHook(() => useAdkMessages({ stream }));
|
|
252
|
+
|
|
253
|
+
let firstSend!: Promise<void>;
|
|
254
|
+
act(() => {
|
|
255
|
+
firstSend = result.current.sendMessage(
|
|
256
|
+
[{ id: "user-1", type: "human", content: "first" }],
|
|
257
|
+
{},
|
|
258
|
+
);
|
|
259
|
+
});
|
|
260
|
+
await vi.waitFor(() => expect(stream).toHaveBeenCalledOnce());
|
|
261
|
+
|
|
262
|
+
let secondSend!: Promise<void>;
|
|
263
|
+
act(() => {
|
|
264
|
+
secondSend = result.current.sendMessage(
|
|
265
|
+
[{ id: "user-2", type: "human", content: "second" }],
|
|
266
|
+
{},
|
|
267
|
+
);
|
|
268
|
+
});
|
|
269
|
+
|
|
270
|
+
await act(async () => {
|
|
271
|
+
await Promise.all([firstSend, secondSend]);
|
|
272
|
+
});
|
|
273
|
+
expect(signals[0]?.aborted).toBe(true);
|
|
274
|
+
expect(signals[1]?.aborted).toBe(false);
|
|
275
|
+
expect(result.current.messages.at(-1)).toMatchObject({
|
|
276
|
+
type: "ai",
|
|
277
|
+
content: [{ type: "text", text: "done-1" }],
|
|
278
|
+
});
|
|
279
|
+
});
|
|
280
|
+
|
|
134
281
|
it("aborts the active stream when the hook unmounts", async () => {
|
|
135
282
|
let runSignal: AbortSignal | undefined;
|
|
136
283
|
let resolveStarted!: () => void;
|
|
@@ -216,6 +363,86 @@ describe("optimistic confirmation replies", () => {
|
|
|
216
363
|
return result;
|
|
217
364
|
};
|
|
218
365
|
|
|
366
|
+
it("preserves an unanswered gate across a reply run", async () => {
|
|
367
|
+
let run = 0;
|
|
368
|
+
const stream: AdkStreamCallback = async function* () {
|
|
369
|
+
run += 1;
|
|
370
|
+
if (run === 1) {
|
|
371
|
+
yield {
|
|
372
|
+
id: "gates",
|
|
373
|
+
author: "agent",
|
|
374
|
+
longRunningToolIds: ["conf-a", "conf-b"],
|
|
375
|
+
content: {
|
|
376
|
+
role: "model",
|
|
377
|
+
parts: [
|
|
378
|
+
{
|
|
379
|
+
functionCall: {
|
|
380
|
+
id: "conf-a",
|
|
381
|
+
name: "adk_request_confirmation",
|
|
382
|
+
args: {},
|
|
383
|
+
},
|
|
384
|
+
},
|
|
385
|
+
{
|
|
386
|
+
functionCall: {
|
|
387
|
+
id: "conf-b",
|
|
388
|
+
name: "adk_request_confirmation",
|
|
389
|
+
args: {},
|
|
390
|
+
},
|
|
391
|
+
},
|
|
392
|
+
],
|
|
393
|
+
},
|
|
394
|
+
} satisfies AdkEvent;
|
|
395
|
+
} else {
|
|
396
|
+
yield {
|
|
397
|
+
id: "rerun",
|
|
398
|
+
author: "agent",
|
|
399
|
+
content: {
|
|
400
|
+
role: "user",
|
|
401
|
+
parts: [
|
|
402
|
+
{
|
|
403
|
+
functionResponse: {
|
|
404
|
+
id: "orig-conf-a",
|
|
405
|
+
name: "delete_file",
|
|
406
|
+
response: { result: "deleted" },
|
|
407
|
+
},
|
|
408
|
+
},
|
|
409
|
+
],
|
|
410
|
+
},
|
|
411
|
+
} satisfies AdkEvent;
|
|
412
|
+
}
|
|
413
|
+
};
|
|
414
|
+
const { result } = renderHook(() => useAdkMessages({ stream }));
|
|
415
|
+
|
|
416
|
+
await act(async () => {
|
|
417
|
+
await result.current.sendMessage(
|
|
418
|
+
[{ id: "user-1", type: "human", content: "start" }],
|
|
419
|
+
{},
|
|
420
|
+
);
|
|
421
|
+
});
|
|
422
|
+
expect(result.current.longRunningToolIds).toEqual(["conf-a", "conf-b"]);
|
|
423
|
+
|
|
424
|
+
await act(async () => {
|
|
425
|
+
await result.current.sendMessage(
|
|
426
|
+
[
|
|
427
|
+
confirmationReply(
|
|
428
|
+
"reply-a",
|
|
429
|
+
"conf-a",
|
|
430
|
+
JSON.stringify({ confirmed: true }),
|
|
431
|
+
),
|
|
432
|
+
],
|
|
433
|
+
{},
|
|
434
|
+
);
|
|
435
|
+
});
|
|
436
|
+
|
|
437
|
+
expect(result.current.longRunningToolIds).toEqual(["conf-b"]);
|
|
438
|
+
expect(
|
|
439
|
+
getPendingCancellations(
|
|
440
|
+
result.current.messages,
|
|
441
|
+
result.current.longRunningToolIds,
|
|
442
|
+
),
|
|
443
|
+
).toEqual([]);
|
|
444
|
+
});
|
|
445
|
+
|
|
219
446
|
it("keeps both gates pending when one send carries an unreadable reply", async () => {
|
|
220
447
|
const result = await renderWithGates();
|
|
221
448
|
|
|
@@ -439,6 +666,26 @@ describe("optimistic multi-message sends", () => {
|
|
|
439
666
|
});
|
|
440
667
|
|
|
441
668
|
describe("messageToEvent (contentToParts)", () => {
|
|
669
|
+
it.each([
|
|
670
|
+
["scalar", "false", { result: false }],
|
|
671
|
+
["array", "[1,2]", { results: [1, 2] }],
|
|
672
|
+
])(
|
|
673
|
+
"normalizes an optimistic %s tool response",
|
|
674
|
+
(_label, content, response) => {
|
|
675
|
+
const event = messageToEvent({
|
|
676
|
+
id: "tool-1",
|
|
677
|
+
type: "tool",
|
|
678
|
+
content,
|
|
679
|
+
tool_call_id: "call-1",
|
|
680
|
+
name: "search",
|
|
681
|
+
});
|
|
682
|
+
|
|
683
|
+
expect(event.content?.parts[0]?.functionResponse?.response).toEqual(
|
|
684
|
+
response,
|
|
685
|
+
);
|
|
686
|
+
},
|
|
687
|
+
);
|
|
688
|
+
|
|
442
689
|
it("serializes a file content part as inlineData", () => {
|
|
443
690
|
const msg: AdkMessage = {
|
|
444
691
|
id: "m1",
|
package/src/useAdkMessages.ts
CHANGED
|
@@ -8,9 +8,14 @@ import {
|
|
|
8
8
|
} from "react";
|
|
9
9
|
import { generateId } from "@assistant-ui/core";
|
|
10
10
|
import { useAui } from "@assistant-ui/store";
|
|
11
|
-
import {
|
|
11
|
+
import {
|
|
12
|
+
abortableIterable,
|
|
13
|
+
invokeUserCallback,
|
|
14
|
+
openAbortableIterable,
|
|
15
|
+
} from "@assistant-ui/core/internal";
|
|
12
16
|
import { AdkEventAccumulator } from "./AdkEventAccumulator";
|
|
13
17
|
import { contentToParts } from "./contentToParts";
|
|
18
|
+
import { toAdkFunctionResponse } from "./toAdkFunctionResponse";
|
|
14
19
|
import type {
|
|
15
20
|
AdkEvent,
|
|
16
21
|
AdkMessage,
|
|
@@ -54,7 +59,7 @@ export const useAdkMessages = ({
|
|
|
54
59
|
name?: string | undefined;
|
|
55
60
|
branch?: string | undefined;
|
|
56
61
|
}>({});
|
|
57
|
-
const [longRunningToolIds,
|
|
62
|
+
const [longRunningToolIds, _setLongRunningToolIds] = useState<string[]>([]);
|
|
58
63
|
const [artifactDelta, setArtifactDelta] = useState<Record<string, number>>(
|
|
59
64
|
{},
|
|
60
65
|
);
|
|
@@ -67,9 +72,9 @@ export const useAdkMessages = ({
|
|
|
67
72
|
Map<string, AdkMessageMetadata>
|
|
68
73
|
>(new Map());
|
|
69
74
|
const lastTransferToAgentRef = useRef<string | undefined>(undefined);
|
|
70
|
-
// setMessagesImmediate
|
|
71
|
-
// this ref with it, so the ref never trails a commit.
|
|
75
|
+
// setMessagesImmediate and setLongRunningToolIds are the only writers of their state and publish these refs with it, so neither ref trails a commit.
|
|
72
76
|
const messagesRef = useRef(messages);
|
|
77
|
+
const longRunningToolIdsRef = useRef(longRunningToolIds);
|
|
73
78
|
const stateDeltaRef = useRef(stateDelta);
|
|
74
79
|
useInsertionEffect(() => {
|
|
75
80
|
stateDeltaRef.current = stateDelta;
|
|
@@ -87,6 +92,10 @@ export const useAdkMessages = ({
|
|
|
87
92
|
messagesRef.current = msgs;
|
|
88
93
|
_setMessages(msgs);
|
|
89
94
|
}, []);
|
|
95
|
+
const setLongRunningToolIds = useCallback((ids: string[]) => {
|
|
96
|
+
longRunningToolIdsRef.current = ids;
|
|
97
|
+
_setLongRunningToolIds(ids);
|
|
98
|
+
}, []);
|
|
90
99
|
|
|
91
100
|
/**
|
|
92
101
|
* Swap the thread over to a loaded snapshot in one commit. Unlike
|
|
@@ -106,7 +115,7 @@ export const useAdkMessages = ({
|
|
|
106
115
|
setArtifactDelta(snapshot.artifactDelta ?? {});
|
|
107
116
|
setAgentInfo(snapshot.agentInfo ?? {});
|
|
108
117
|
},
|
|
109
|
-
[setMessagesImmediate],
|
|
118
|
+
[setLongRunningToolIds, setMessagesImmediate],
|
|
110
119
|
);
|
|
111
120
|
|
|
112
121
|
// Replace the message list AND reset derived per-turn HITL state.
|
|
@@ -122,7 +131,7 @@ export const useAdkMessages = ({
|
|
|
122
131
|
setEscalated(false);
|
|
123
132
|
setMessageMetadata(new Map());
|
|
124
133
|
},
|
|
125
|
-
[setMessagesImmediate],
|
|
134
|
+
[setLongRunningToolIds, setMessagesImmediate],
|
|
126
135
|
);
|
|
127
136
|
|
|
128
137
|
const abortControllerRef = useRef<AbortController | null>(null);
|
|
@@ -144,27 +153,52 @@ export const useAdkMessages = ({
|
|
|
144
153
|
// with the originals would leave every later staged id beside the merged
|
|
145
154
|
// copy of itself.
|
|
146
155
|
const resentIds = new Set(newMessagesWithId.map((m) => m.id));
|
|
156
|
+
// The optimistic event for a tool-only batch carries no author, so the accumulator cannot settle the calls this send answers.
|
|
157
|
+
const answeredToolCallIds = new Set(
|
|
158
|
+
newMessagesWithId.flatMap((m) =>
|
|
159
|
+
m.type === "tool" ? [m.tool_call_id] : [],
|
|
160
|
+
),
|
|
161
|
+
);
|
|
147
162
|
const accumulator = new AdkEventAccumulator(
|
|
148
163
|
messagesRef.current.filter((m) => !resentIds.has(m.id)),
|
|
164
|
+
longRunningToolIdsRef.current.filter(
|
|
165
|
+
(id) => !answeredToolCallIds.has(id),
|
|
166
|
+
),
|
|
149
167
|
);
|
|
150
168
|
for (const event of messagesToEvents(newMessagesWithId)) {
|
|
151
169
|
accumulator.processEvent(event);
|
|
152
170
|
}
|
|
153
171
|
setMessagesImmediate(accumulator.getMessages());
|
|
172
|
+
setLongRunningToolIds(accumulator.getLongRunningToolIds());
|
|
154
173
|
|
|
174
|
+
// Google ADK replaces active runs, while React LangGraph queues sends.
|
|
175
|
+
abortControllerRef.current?.abort();
|
|
155
176
|
const abortController = new AbortController();
|
|
156
177
|
abortControllerRef.current = abortController;
|
|
157
178
|
|
|
158
179
|
try {
|
|
159
|
-
const response = await
|
|
160
|
-
|
|
161
|
-
|
|
162
|
-
|
|
163
|
-
|
|
164
|
-
|
|
165
|
-
|
|
180
|
+
const response = await openAbortableIterable(
|
|
181
|
+
stream(newMessagesWithId, {
|
|
182
|
+
...config,
|
|
183
|
+
abortSignal: abortController.signal,
|
|
184
|
+
initialize: async () => {
|
|
185
|
+
return await aui.threadListItem.initialize();
|
|
186
|
+
},
|
|
187
|
+
}),
|
|
188
|
+
abortController.signal,
|
|
189
|
+
);
|
|
190
|
+
if (!response) return;
|
|
166
191
|
|
|
167
|
-
for await (const event of
|
|
192
|
+
for await (const event of abortableIterable(
|
|
193
|
+
response,
|
|
194
|
+
abortController.signal,
|
|
195
|
+
)) {
|
|
196
|
+
if (
|
|
197
|
+
abortController.signal.aborted ||
|
|
198
|
+
abortControllerRef.current !== abortController
|
|
199
|
+
) {
|
|
200
|
+
break;
|
|
201
|
+
}
|
|
168
202
|
const updatedMessages = accumulator.processEvent(event);
|
|
169
203
|
setMessagesImmediate(updatedMessages);
|
|
170
204
|
setStateDelta({
|
|
@@ -222,6 +256,7 @@ export const useAdkMessages = ({
|
|
|
222
256
|
} catch (error) {
|
|
223
257
|
if (
|
|
224
258
|
!abortController.signal.aborted &&
|
|
259
|
+
abortControllerRef.current === abortController &&
|
|
225
260
|
!(error instanceof Error && error.name === "AbortError")
|
|
226
261
|
) {
|
|
227
262
|
throw error;
|
|
@@ -235,6 +270,7 @@ export const useAdkMessages = ({
|
|
|
235
270
|
[
|
|
236
271
|
aui,
|
|
237
272
|
setMessagesImmediate,
|
|
273
|
+
setLongRunningToolIds,
|
|
238
274
|
stream,
|
|
239
275
|
onError,
|
|
240
276
|
onCustomEvent,
|
|
@@ -343,7 +379,7 @@ export const messageToEvent = (msg: AdkMessage): AdkEvent => {
|
|
|
343
379
|
functionResponse: {
|
|
344
380
|
name: msg.name,
|
|
345
381
|
id: msg.tool_call_id,
|
|
346
|
-
response,
|
|
382
|
+
response: toAdkFunctionResponse(response, msg.status === "error"),
|
|
347
383
|
},
|
|
348
384
|
},
|
|
349
385
|
],
|
|
@@ -0,0 +1,138 @@
|
|
|
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 type { AdkEvent } from "./types";
|
|
13
|
+
|
|
14
|
+
const makeThreadListAdapter = (): RemoteThreadListAdapter => ({
|
|
15
|
+
list: vi.fn(async () => ({
|
|
16
|
+
threads: [
|
|
17
|
+
{
|
|
18
|
+
status: "regular" as const,
|
|
19
|
+
remoteId: "adk-1",
|
|
20
|
+
externalId: "adk-1",
|
|
21
|
+
title: "ADK session",
|
|
22
|
+
},
|
|
23
|
+
],
|
|
24
|
+
})),
|
|
25
|
+
initialize: vi.fn(async () => ({
|
|
26
|
+
remoteId: "adk-1",
|
|
27
|
+
externalId: "adk-1",
|
|
28
|
+
})),
|
|
29
|
+
rename: vi.fn(async () => {}),
|
|
30
|
+
archive: vi.fn(async () => {}),
|
|
31
|
+
unarchive: vi.fn(async () => {}),
|
|
32
|
+
delete: vi.fn(async () => {}),
|
|
33
|
+
generateTitle: vi.fn(async () => new ReadableStream() as never),
|
|
34
|
+
fetch: vi.fn(async () => ({
|
|
35
|
+
status: "regular" as const,
|
|
36
|
+
remoteId: "adk-1",
|
|
37
|
+
externalId: "adk-1",
|
|
38
|
+
title: "ADK session",
|
|
39
|
+
})),
|
|
40
|
+
});
|
|
41
|
+
|
|
42
|
+
const deferred = () => {
|
|
43
|
+
let resolve!: () => void;
|
|
44
|
+
const promise = new Promise<void>((r) => {
|
|
45
|
+
resolve = r;
|
|
46
|
+
});
|
|
47
|
+
return { promise, resolve };
|
|
48
|
+
};
|
|
49
|
+
|
|
50
|
+
describe("useAdkRuntime replacement runs", () => {
|
|
51
|
+
it.each([
|
|
52
|
+
{ label: "events after cancellation", cancelFirst: true, failFirst: false },
|
|
53
|
+
{
|
|
54
|
+
label: "events without cancellation",
|
|
55
|
+
cancelFirst: false,
|
|
56
|
+
failFirst: false,
|
|
57
|
+
},
|
|
58
|
+
{
|
|
59
|
+
label: "errors without cancellation",
|
|
60
|
+
cancelFirst: false,
|
|
61
|
+
failFirst: true,
|
|
62
|
+
},
|
|
63
|
+
])("ignores superseded run $label", async ({ cancelFirst, failFirst }) => {
|
|
64
|
+
const gates = [deferred(), deferred()];
|
|
65
|
+
let calls = 0;
|
|
66
|
+
const stream = vi.fn(async function* (): AsyncGenerator<AdkEvent> {
|
|
67
|
+
const call = calls++;
|
|
68
|
+
await gates[call]!.promise;
|
|
69
|
+
if (call === 0 && failFirst) throw new Error("stale run failed");
|
|
70
|
+
yield {
|
|
71
|
+
id: `event-${call}`,
|
|
72
|
+
invocationId: `run-${call}`,
|
|
73
|
+
author: "agent",
|
|
74
|
+
content: { role: "model", parts: [{ text: `done-${call}` }] },
|
|
75
|
+
};
|
|
76
|
+
});
|
|
77
|
+
const sessionAdapter = makeThreadListAdapter();
|
|
78
|
+
const capture: { runtime: AssistantRuntime | null } = { runtime: null };
|
|
79
|
+
|
|
80
|
+
const Inner: FC = () => {
|
|
81
|
+
const runtime = useAdkRuntime({
|
|
82
|
+
stream,
|
|
83
|
+
sessionAdapter,
|
|
84
|
+
unstable_allowCancellation: true,
|
|
85
|
+
});
|
|
86
|
+
capture.runtime = runtime;
|
|
87
|
+
return <AssistantRuntimeProvider runtime={runtime} />;
|
|
88
|
+
};
|
|
89
|
+
|
|
90
|
+
await act(async () => {
|
|
91
|
+
render(<Inner />);
|
|
92
|
+
});
|
|
93
|
+
await waitFor(() => expect(capture.runtime).not.toBeNull());
|
|
94
|
+
await act(async () => {
|
|
95
|
+
await capture.runtime!.threads.switchToThread("adk-1");
|
|
96
|
+
});
|
|
97
|
+
|
|
98
|
+
let firstSend!: Promise<void>;
|
|
99
|
+
act(() => {
|
|
100
|
+
firstSend = capture.runtime!.thread.append({
|
|
101
|
+
role: "user",
|
|
102
|
+
content: [{ type: "text", text: "first" }],
|
|
103
|
+
});
|
|
104
|
+
});
|
|
105
|
+
await waitFor(() => expect(stream).toHaveBeenCalledTimes(1));
|
|
106
|
+
|
|
107
|
+
let secondSend!: Promise<void>;
|
|
108
|
+
await act(async () => {
|
|
109
|
+
if (cancelFirst) await capture.runtime!.thread.cancelRun();
|
|
110
|
+
secondSend = capture.runtime!.thread.append({
|
|
111
|
+
role: "user",
|
|
112
|
+
content: [{ type: "text", text: "second" }],
|
|
113
|
+
});
|
|
114
|
+
});
|
|
115
|
+
await waitFor(() => expect(stream).toHaveBeenCalledTimes(2));
|
|
116
|
+
|
|
117
|
+
await act(async () => {
|
|
118
|
+
gates[0]!.resolve();
|
|
119
|
+
await firstSend;
|
|
120
|
+
});
|
|
121
|
+
|
|
122
|
+
const messagesAfterFirstSettles = JSON.stringify(
|
|
123
|
+
capture.runtime!.thread.getState().messages,
|
|
124
|
+
);
|
|
125
|
+
expect(messagesAfterFirstSettles).toContain("second");
|
|
126
|
+
expect(messagesAfterFirstSettles).not.toContain("done-0");
|
|
127
|
+
expect(capture.runtime!.thread.getState().isRunning).toBe(true);
|
|
128
|
+
|
|
129
|
+
await act(async () => {
|
|
130
|
+
gates[1]!.resolve();
|
|
131
|
+
await secondSend;
|
|
132
|
+
});
|
|
133
|
+
expect(
|
|
134
|
+
JSON.stringify(capture.runtime!.thread.getState().messages),
|
|
135
|
+
).toContain("done-1");
|
|
136
|
+
expect(capture.runtime!.thread.getState().isRunning).toBe(false);
|
|
137
|
+
});
|
|
138
|
+
});
|
package/src/useAdkRuntime.ts
CHANGED
|
@@ -57,6 +57,7 @@ import {
|
|
|
57
57
|
toAdkToolConfirmationReply,
|
|
58
58
|
} from "./adkToolApproval";
|
|
59
59
|
import { adkExtras } from "./adkExtras";
|
|
60
|
+
import { ADK_SDK } from "./sdkIdentity";
|
|
60
61
|
|
|
61
62
|
export type UseAdkRuntimeOptions = ExternalStoreSharedOptions & {
|
|
62
63
|
stream: AdkStreamCallback;
|
|
@@ -168,16 +169,18 @@ const useAdkRuntimeImpl = (options: UseAdkRuntimeOptions) => {
|
|
|
168
169
|
useInsertionEffect(() => {
|
|
169
170
|
isRunningRef.current = effectiveIsRunning;
|
|
170
171
|
}, [effectiveIsRunning]);
|
|
172
|
+
const runGenerationRef = useRef(0);
|
|
171
173
|
|
|
172
174
|
const handleSendMessage = async (
|
|
173
175
|
msgs: AdkMessage[],
|
|
174
176
|
config: AdkSendMessageConfig,
|
|
175
177
|
) => {
|
|
178
|
+
const generation = ++runGenerationRef.current;
|
|
176
179
|
try {
|
|
177
180
|
setIsRunning(true);
|
|
178
181
|
await sendMessage(msgs, config);
|
|
179
182
|
} finally {
|
|
180
|
-
setIsRunning(false);
|
|
183
|
+
if (runGenerationRef.current === generation) setIsRunning(false);
|
|
181
184
|
}
|
|
182
185
|
};
|
|
183
186
|
|
|
@@ -502,6 +505,7 @@ export const useAdkRuntime = ({
|
|
|
502
505
|
}: UseAdkRuntimeOptions) => {
|
|
503
506
|
const aui = useAui();
|
|
504
507
|
const cloudAdapter = useCloudThreadListAdapter({
|
|
508
|
+
sdk: ADK_SDK,
|
|
505
509
|
cloud,
|
|
506
510
|
create: createCloudThreadListAdapterCreateFallback(
|
|
507
511
|
create,
|