@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.
Files changed (35) hide show
  1. package/dist/AdkEventAccumulator.d.ts.map +1 -1
  2. package/dist/AdkEventAccumulator.js +4 -1
  3. package/dist/AdkEventAccumulator.js.map +1 -1
  4. package/dist/adkToolApproval.js +1 -1
  5. package/dist/adkToolApproval.js.map +1 -1
  6. package/dist/convertAdkMessages.d.ts.map +1 -1
  7. package/dist/convertAdkMessages.js +30 -19
  8. package/dist/convertAdkMessages.js.map +1 -1
  9. package/dist/convertToAdkMessages.d.ts.map +1 -1
  10. package/dist/convertToAdkMessages.js +2 -2
  11. package/dist/convertToAdkMessages.js.map +1 -1
  12. package/dist/sdkIdentity.js +1 -1
  13. package/dist/useAdkMessages.d.ts +23 -0
  14. package/dist/useAdkMessages.d.ts.map +1 -1
  15. package/dist/useAdkMessages.js +34 -9
  16. package/dist/useAdkMessages.js.map +1 -1
  17. package/dist/useAdkRuntime.d.ts.map +1 -1
  18. package/dist/useAdkRuntime.js +60 -9
  19. package/dist/useAdkRuntime.js.map +1 -1
  20. package/package.json +7 -6
  21. package/src/AdkEventAccumulator.test.ts +32 -0
  22. package/src/AdkEventAccumulator.ts +1 -0
  23. package/src/adkToolApproval.test.ts +19 -0
  24. package/src/adkToolApproval.ts +1 -1
  25. package/src/convertAdkMessages.test.ts +129 -0
  26. package/src/convertAdkMessages.ts +19 -1
  27. package/src/convertToAdkMessages.test.ts +14 -0
  28. package/src/convertToAdkMessages.ts +4 -1
  29. package/src/useAdkMessages.test.ts +55 -0
  30. package/src/useAdkMessages.ts +65 -6
  31. package/src/useAdkRuntime.cancellation.test.tsx +544 -0
  32. package/src/useAdkRuntime.fast-refresh.test.tsx +166 -0
  33. package/src/useAdkRuntime.refetch.test.tsx +1 -0
  34. package/src/useAdkRuntime.ts +91 -9
  35. package/src/useAdkRuntimeApproval.test.tsx +305 -36
@@ -42,18 +42,30 @@ const contentToParts = (
42
42
  text: typeof part.text === "string" ? part.text : "",
43
43
  };
44
44
  case "reasoning":
45
+ if (role === "user") return null;
45
46
  return {
46
47
  type: "reasoning",
47
48
  text: typeof part.text === "string" ? part.text : "",
48
49
  };
49
50
  case "image":
51
+ if (
52
+ typeof part.mimeType !== "string" ||
53
+ typeof part.data !== "string"
54
+ )
55
+ return null;
50
56
  return {
51
57
  type: "image",
52
58
  image: `data:${part.mimeType};base64,${part.data}`,
53
59
  };
54
60
  case "image_url":
61
+ if (typeof part.url !== "string") return null;
55
62
  return { type: "image", image: part.url };
56
63
  case "file":
64
+ if (
65
+ typeof part.mimeType !== "string" ||
66
+ typeof part.data !== "string"
67
+ )
68
+ return null;
57
69
  return {
58
70
  type: "file",
59
71
  data: part.data,
@@ -61,6 +73,7 @@ const contentToParts = (
61
73
  ...(part.filename != null && { filename: part.filename }),
62
74
  };
63
75
  case "file_url":
76
+ if (typeof part.url !== "string") return null;
64
77
  if (role === "user") {
65
78
  return {
66
79
  type: "file",
@@ -118,7 +131,8 @@ export const createAdkMessageConverter =
118
131
 
119
132
  case "ai": {
120
133
  const toolCallParts: ToolCallMessagePart[] =
121
- message.tool_calls?.map((tc) => {
134
+ message.tool_calls?.flatMap((tc) => {
135
+ if (typeof tc?.name !== "string" || tc.name.length === 0) return [];
122
136
  const approval = approvals.get(tc.id);
123
137
  return {
124
138
  type: "tool-call",
@@ -169,6 +183,10 @@ export const createAdkMessageConverter =
169
183
  isError: message.status === "error",
170
184
  };
171
185
  }
186
+
187
+ default:
188
+ message satisfies never;
189
+ return [];
172
190
  }
173
191
  };
174
192
 
@@ -81,6 +81,20 @@ describe("getPendingToolCalls", () => {
81
81
  });
82
82
 
83
83
  describe("getPendingCancellations", () => {
84
+ it("skips a null tool_calls entry and cancels the rest", () => {
85
+ const messages = [
86
+ {
87
+ id: "ai-1",
88
+ type: "ai",
89
+ content: [],
90
+ tool_calls: [null, { id: "tc-1", name: "tool_a", args: {} }],
91
+ },
92
+ ] as unknown as AdkMessage[];
93
+ expect(getPendingCancellations(messages, [])).toMatchObject([
94
+ { type: "tool", name: "tool_a", tool_call_id: "tc-1", status: "error" },
95
+ ]);
96
+ });
97
+
84
98
  it("emits a {cancelled:true} tool message for every pending tool call", () => {
85
99
  const messages: AdkMessage[] = [
86
100
  aiWithToolCalls("ai-1", [{ id: "tc-1", name: "tool_a" }]),
@@ -1,3 +1,4 @@
1
+ import { isRecord } from "@assistant-ui/core/internal";
1
2
  import {
2
3
  generateId,
3
4
  getExternalStoreMessages,
@@ -94,7 +95,9 @@ export const getPendingToolCalls = (messages: AdkMessage[]) => {
94
95
  messages,
95
96
  (message) => {
96
97
  if (message.type === "ai") {
97
- return { toolCalls: message.tool_calls ?? [] };
98
+ return {
99
+ toolCalls: (message.tool_calls ?? []).filter(isRecord),
100
+ };
98
101
  }
99
102
  if (message.type === "tool") {
100
103
  return { toolCallId: message.tool_call_id };
@@ -17,6 +17,7 @@ import {
17
17
  messageToEvent,
18
18
  messagesToEvents,
19
19
  useAdkMessages,
20
+ useAdkMessagesInternal,
20
21
  } from "./useAdkMessages";
21
22
  import { projectAdkToolApprovals } from "./adkToolApproval";
22
23
  import { createAdkStream } from "./AdkClient";
@@ -212,6 +213,47 @@ describe("ADK runtime callbacks", () => {
212
213
  });
213
214
 
214
215
  describe("ADK stream lifecycle", () => {
216
+ it("reports streamed tool calls with the run config that produced them", async () => {
217
+ const runConfig = { custom: { model: "model-a" } };
218
+ const onMessages = vi.fn();
219
+ const stream: AdkStreamCallback = async function* () {
220
+ yield {
221
+ id: "event-1",
222
+ author: "agent",
223
+ content: {
224
+ role: "model",
225
+ parts: [
226
+ { functionCall: { id: "tool-1", name: "lookup", args: {} } },
227
+ { functionCall: { id: "tool-2", name: "search", args: {} } },
228
+ ],
229
+ },
230
+ };
231
+ };
232
+ const { result } = renderHook(() =>
233
+ useAdkMessagesInternal({ stream, onMessages }),
234
+ );
235
+
236
+ await act(async () => {
237
+ await result.current.sendMessage(
238
+ [{ id: "user-1", type: "human", content: "look it up" }],
239
+ { runConfig },
240
+ );
241
+ });
242
+
243
+ expect(onMessages).toHaveBeenLastCalledWith(
244
+ expect.arrayContaining([
245
+ expect.objectContaining({
246
+ type: "ai",
247
+ tool_calls: [
248
+ expect.objectContaining({ id: "tool-1" }),
249
+ expect.objectContaining({ id: "tool-2" }),
250
+ ],
251
+ }),
252
+ ]),
253
+ runConfig,
254
+ );
255
+ });
256
+
215
257
  it("settles a superseded send while its stream is still opening", async () => {
216
258
  const signals: AbortSignal[] = [];
217
259
  const parked = new Promise<AsyncGenerator<AdkEvent>>(() => {});
@@ -800,6 +842,19 @@ describe("optimistic multi-message sends", () => {
800
842
  });
801
843
 
802
844
  describe("messageToEvent (contentToParts)", () => {
845
+ it("skips a null tool_calls entry", () => {
846
+ const event = messageToEvent({
847
+ id: "ai-1",
848
+ type: "ai",
849
+ content: [],
850
+ tool_calls: [null, { id: "tc-1", name: "search", args: { q: "x" } }],
851
+ } as unknown as AdkMessage);
852
+
853
+ expect(event.content?.parts).toEqual([
854
+ { functionCall: { name: "search", id: "tc-1", args: { q: "x" } } },
855
+ ]);
856
+ });
857
+
803
858
  it.each([
804
859
  ["scalar", "false", { result: false }],
805
860
  ["array", "[1,2]", { results: [1, 2] }],
@@ -1,3 +1,4 @@
1
+ import { isRecord } from "@assistant-ui/core/internal";
1
2
  import {
2
3
  useState,
3
4
  useCallback,
@@ -39,6 +40,10 @@ export type UseAdkMessagesOptions = {
39
40
  };
40
41
  };
41
42
 
43
+ type UseAdkMessagesInternalOptions = UseAdkMessagesOptions & {
44
+ onMessages?: (messages: AdkMessage[], runConfig: unknown) => void;
45
+ };
46
+
42
47
  type AdkRuntimeCallbackName = "onError" | "onCustomEvent" | "onAgentTransfer";
43
48
 
44
49
  const invokeAdkRuntimeCallback = <TArgs extends readonly unknown[]>(
@@ -49,10 +54,11 @@ const invokeAdkRuntimeCallback = <TArgs extends readonly unknown[]>(
49
54
  void invokeUserCallback("react-google-adk", name, callback, ...args);
50
55
  };
51
56
 
52
- export const useAdkMessages = ({
57
+ const useAdkMessagesInternal = ({
53
58
  stream,
54
59
  eventHandlers,
55
- }: UseAdkMessagesOptions) => {
60
+ onMessages,
61
+ }: UseAdkMessagesInternalOptions) => {
56
62
  const [messages, _setMessages] = useState<AdkMessage[]>([]);
57
63
  const [stateDelta, setStateDelta] = useState<Record<string, unknown>>({});
58
64
  const [agentInfo, setAgentInfo] = useState<{
@@ -167,8 +173,11 @@ export const useAdkMessages = ({
167
173
  for (const event of messagesToEvents(newMessagesWithId)) {
168
174
  accumulator.processEvent(event);
169
175
  }
170
- setMessagesImmediate(accumulator.getMessages());
171
- setLongRunningToolIds(accumulator.getLongRunningToolIds());
176
+ const initialMessages = accumulator.getMessages();
177
+ const initialMessageIds = new Set(initialMessages.map((m) => m.id));
178
+ const initialLongRunningToolIds = accumulator.getLongRunningToolIds();
179
+ setMessagesImmediate(initialMessages);
180
+ setLongRunningToolIds(initialLongRunningToolIds);
172
181
  setToolConfirmations(accumulator.getToolConfirmations());
173
182
  setAuthRequests(accumulator.getAuthRequests());
174
183
  let lastTransferToAgent: string | undefined;
@@ -202,6 +211,17 @@ export const useAdkMessages = ({
202
211
  break;
203
212
  }
204
213
  const updatedMessages = accumulator.processEvent(event);
214
+ // Each event part can append at most one message, and a function call
215
+ // stays on the current assistant message until a later part finalizes
216
+ // it, so every message touched by this event is within this tail.
217
+ const affectedMessageCount = Math.max(
218
+ event.content?.parts?.length ?? 0,
219
+ 1,
220
+ );
221
+ const affectedMessages = updatedMessages.slice(-affectedMessageCount);
222
+ if (affectedMessages.length > 0) {
223
+ onMessages?.(affectedMessages, config.runConfig);
224
+ }
205
225
  setMessagesImmediate(updatedMessages);
206
226
  setStateDelta({
207
227
  ...stateDeltaRef.current,
@@ -265,6 +285,33 @@ export const useAdkMessages = ({
265
285
  }
266
286
  } finally {
267
287
  if (abortControllerRef.current === abortController) {
288
+ if (abortController.signal.aborted) {
289
+ setLongRunningToolIds(
290
+ accumulator
291
+ .getLongRunningToolIds()
292
+ .filter((id) => initialLongRunningToolIds.includes(id)),
293
+ );
294
+ const updatedMessages = messagesRef.current;
295
+ const lastAssistantMessage = updatedMessages.findLast(
296
+ (m) => m.type === "ai",
297
+ );
298
+ if (
299
+ lastAssistantMessage &&
300
+ !initialMessageIds.has(lastAssistantMessage.id) &&
301
+ !lastAssistantMessage.status
302
+ ) {
303
+ setMessagesImmediate(
304
+ updatedMessages.map((m) =>
305
+ m === lastAssistantMessage
306
+ ? {
307
+ ...lastAssistantMessage,
308
+ status: { type: "incomplete", reason: "cancelled" },
309
+ }
310
+ : m,
311
+ ),
312
+ );
313
+ }
314
+ }
268
315
  abortControllerRef.current = null;
269
316
  }
270
317
  }
@@ -277,6 +324,7 @@ export const useAdkMessages = ({
277
324
  onError,
278
325
  onCustomEvent,
279
326
  onAgentTransfer,
327
+ onMessages,
280
328
  ],
281
329
  );
282
330
 
@@ -306,6 +354,17 @@ export const useAdkMessages = ({
306
354
  };
307
355
  };
308
356
 
357
+ export const useAdkMessages = ({
358
+ stream,
359
+ eventHandlers,
360
+ }: UseAdkMessagesOptions) =>
361
+ useAdkMessagesInternal({
362
+ stream,
363
+ ...(eventHandlers !== undefined && { eventHandlers }),
364
+ });
365
+
366
+ export { useAdkMessagesInternal };
367
+
309
368
  /**
310
369
  * Transport sends every human and tool message of one `send` call as a single
311
370
  * ADK `Content`, and ADK parses that event's function responses before running
@@ -395,9 +454,9 @@ export const messageToEvent = (msg: AdkMessage): AdkEvent => {
395
454
  role: "model",
396
455
  parts: [
397
456
  ...contentToParts(msg.content),
398
- ...(msg.tool_calls?.map((tc) => ({
457
+ ...(msg.tool_calls ?? []).filter(isRecord).map((tc) => ({
399
458
  functionCall: { name: tc.name, id: tc.id, args: { ...tc.args } },
400
- })) ?? []),
459
+ })),
401
460
  ],
402
461
  };
403
462
  return result;