@assistant-ui/ai-sdk 0.0.10 → 0.0.12

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 (124) hide show
  1. package/dist/converters/convertMessage.d.ts +8 -6
  2. package/dist/converters/convertMessage.d.ts.map +1 -1
  3. package/dist/converters/convertMessage.js +106 -28
  4. package/dist/converters/convertMessage.js.map +1 -1
  5. package/dist/converters/toolApprovalAnswers.d.ts +3 -0
  6. package/dist/converters/toolApprovalAnswers.d.ts.map +1 -0
  7. package/dist/converters/toolApprovalAnswers.js +17 -0
  8. package/dist/converters/toolApprovalAnswers.js.map +1 -0
  9. package/dist/index.d.ts +1 -1
  10. package/dist/index.d.ts.map +1 -1
  11. package/dist/index.js +1 -1
  12. package/dist/index.native.d.ts +1 -1
  13. package/dist/index.native.d.ts.map +1 -1
  14. package/dist/index.native.js +1 -1
  15. package/dist/model-context/injectInteractableContext.d.ts +3 -0
  16. package/dist/model-context/injectInteractableContext.d.ts.map +1 -1
  17. package/dist/model-context/injectInteractableContext.js +3 -0
  18. package/dist/model-context/injectInteractableContext.js.map +1 -1
  19. package/dist/model-context/injectQuoteContext.d.ts +1 -0
  20. package/dist/model-context/injectQuoteContext.d.ts.map +1 -1
  21. package/dist/model-context/injectQuoteContext.js +1 -0
  22. package/dist/model-context/injectQuoteContext.js.map +1 -1
  23. package/dist/runtime/AISDKChat.js +1 -1
  24. package/dist/runtime/AISDKThreads.d.ts +5 -0
  25. package/dist/runtime/AISDKThreads.d.ts.map +1 -1
  26. package/dist/runtime/AISDKThreads.js +37 -4
  27. package/dist/runtime/AISDKThreads.js.map +1 -1
  28. package/dist/runtime/DynamicChatTransport.d.ts +49 -0
  29. package/dist/runtime/DynamicChatTransport.d.ts.map +1 -0
  30. package/dist/runtime/DynamicChatTransport.js +147 -0
  31. package/dist/runtime/DynamicChatTransport.js.map +1 -0
  32. package/dist/runtime/getResumableAdapter.d.ts +5 -0
  33. package/dist/runtime/getResumableAdapter.d.ts.map +1 -0
  34. package/dist/runtime/getResumableAdapter.js +12 -0
  35. package/dist/runtime/getResumableAdapter.js.map +1 -0
  36. package/dist/runtime/sdkIdentity.js +1 -1
  37. package/dist/runtime/toolHistoryCodec.d.ts +20 -0
  38. package/dist/runtime/toolHistoryCodec.d.ts.map +1 -0
  39. package/dist/runtime/toolHistoryCodec.js +107 -0
  40. package/dist/runtime/toolHistoryCodec.js.map +1 -0
  41. package/dist/runtime/useAISDKRuntime.d.ts +4 -1
  42. package/dist/runtime/useAISDKRuntime.d.ts.map +1 -1
  43. package/dist/runtime/useAISDKRuntime.js +154 -155
  44. package/dist/runtime/useAISDKRuntime.js.map +1 -1
  45. package/dist/runtime/useChatRuntime.d.ts +14 -2
  46. package/dist/runtime/useChatRuntime.d.ts.map +1 -1
  47. package/dist/runtime/useChatRuntime.js +12 -3
  48. package/dist/runtime/useChatRuntime.js.map +1 -1
  49. package/dist/runtime/useChatThread.d.ts +3 -2
  50. package/dist/runtime/useChatThread.d.ts.map +1 -1
  51. package/dist/runtime/useChatThread.js +56 -36
  52. package/dist/runtime/useChatThread.js.map +1 -1
  53. package/dist/runtime/useDynamicChatTransport.d.ts +4 -0
  54. package/dist/runtime/useDynamicChatTransport.d.ts.map +1 -0
  55. package/dist/runtime/useDynamicChatTransport.js +64 -0
  56. package/dist/runtime/useDynamicChatTransport.js.map +1 -0
  57. package/dist/runtime/useExternalHistory.d.ts.map +1 -1
  58. package/dist/runtime/useExternalHistory.js +11 -105
  59. package/dist/runtime/useExternalHistory.js.map +1 -1
  60. package/dist/runtime/useResourceCleanup.js +1 -1
  61. package/dist/runtime/useStreamingTiming.js +2 -2
  62. package/dist/runtime/useStreamingTiming.js.map +1 -1
  63. package/dist/tools/generativeTools.d.ts +2 -1
  64. package/dist/tools/generativeTools.d.ts.map +1 -1
  65. package/dist/tools/generativeTools.js +17 -6
  66. package/dist/tools/generativeTools.js.map +1 -1
  67. package/dist/usage.js +1 -1
  68. package/dist/utils/sliceMessagesUntil.d.ts.map +1 -1
  69. package/dist/utils/sliceMessagesUntil.js +1 -2
  70. package/dist/utils/sliceMessagesUntil.js.map +1 -1
  71. package/package.json +12 -10
  72. package/src/converters/convertMessage.test.ts +350 -2
  73. package/src/converters/convertMessage.tool-args-status.test.tsx +185 -0
  74. package/src/converters/convertMessage.ts +154 -23
  75. package/src/converters/toCreateMessage.test.ts +27 -0
  76. package/src/converters/toolApprovalAnswers.ts +27 -0
  77. package/src/index.native.ts +1 -1
  78. package/src/index.ts +1 -1
  79. package/src/model-context/injectInteractableContext.ts +3 -0
  80. package/src/model-context/injectQuoteContext.ts +1 -0
  81. package/src/runtime/AISDKChat.integration.test.tsx +57 -2
  82. package/src/runtime/AISDKThreads.cloud.test.ts +3 -0
  83. package/src/runtime/AISDKThreads.test.ts +181 -0
  84. package/src/runtime/AISDKThreads.ts +31 -4
  85. package/src/runtime/DynamicChatTransport.test.ts +203 -0
  86. package/src/runtime/DynamicChatTransport.ts +273 -0
  87. package/src/runtime/__tests__/controlled-transport.ts +3 -0
  88. package/src/runtime/getResumableAdapter.ts +16 -0
  89. package/src/runtime/toolHistoryCodec.test.ts +161 -0
  90. package/src/runtime/toolHistoryCodec.ts +207 -0
  91. package/src/runtime/useAISDKRuntime.approval.test.tsx +12 -0
  92. package/src/runtime/useAISDKRuntime.reload.test.tsx +219 -0
  93. package/src/runtime/useAISDKRuntime.test.ts +538 -5
  94. package/src/runtime/useAISDKRuntime.ts +161 -51
  95. package/src/runtime/useChatRuntime.integration.test.tsx +319 -4
  96. package/src/runtime/useChatRuntime.local-storage.test.tsx +123 -0
  97. package/src/runtime/useChatRuntime.test.ts +107 -1
  98. package/src/runtime/useChatRuntime.ts +26 -4
  99. package/src/runtime/useChatThread.binding.test.tsx +143 -0
  100. package/src/runtime/useChatThread.ts +98 -81
  101. package/src/runtime/useDynamicChatTransport.ts +26 -0
  102. package/src/runtime/useExternalHistory.test.ts +205 -0
  103. package/src/runtime/useExternalHistory.ts +14 -206
  104. package/src/runtime/useStreamingTiming.ts +2 -2
  105. package/src/tools/generativeTools.test.ts +190 -2
  106. package/src/tools/generativeTools.ts +28 -8
  107. package/src/utils/sliceMessagesUntil.test.ts +2 -6
  108. package/src/utils/sliceMessagesUntil.ts +1 -5
  109. package/dist/converters/modelContentEnvelope.d.ts +0 -14
  110. package/dist/converters/modelContentEnvelope.d.ts.map +0 -1
  111. package/dist/converters/modelContentEnvelope.js +0 -22
  112. package/dist/converters/modelContentEnvelope.js.map +0 -1
  113. package/dist/converters/toolOutputConversion.d.ts +0 -26
  114. package/dist/converters/toolOutputConversion.d.ts.map +0 -1
  115. package/dist/converters/toolOutputConversion.js +0 -31
  116. package/dist/converters/toolOutputConversion.js.map +0 -1
  117. package/dist/tools/frontendTools.d.ts +0 -30
  118. package/dist/tools/frontendTools.d.ts.map +0 -1
  119. package/dist/tools/frontendTools.js +0 -33
  120. package/dist/tools/frontendTools.js.map +0 -1
  121. package/src/converters/modelContentEnvelope.ts +0 -41
  122. package/src/converters/toolOutputConversion.ts +0 -26
  123. package/src/tools/frontendTools.test.ts +0 -205
  124. package/src/tools/frontendTools.ts +0 -83
@@ -0,0 +1,273 @@
1
+ import type { AssistantRuntime } from "@assistant-ui/core";
2
+ import type { UIMessage } from "@ai-sdk/react";
3
+ import type { ChatTransport } from "ai";
4
+ import {
5
+ AssistantChatTransport,
6
+ type InitializableThreadListItem,
7
+ } from "../transport/AssistantChatTransport";
8
+ import type { ResumableClientStorage } from "../transport/resumable";
9
+ import { getResumableAdapter } from "./getResumableAdapter";
10
+
11
+ const resumedStreamIdsByStorage = new WeakMap<
12
+ ResumableClientStorage,
13
+ Set<string>
14
+ >();
15
+
16
+ export const getResumedStreamIds = (
17
+ storage: ResumableClientStorage | undefined,
18
+ ) => {
19
+ if (!storage) return new Set<string>();
20
+ let resumedStreamIds = resumedStreamIdsByStorage.get(storage);
21
+ if (!resumedStreamIds) {
22
+ resumedStreamIds = new Set();
23
+ resumedStreamIdsByStorage.set(storage, resumedStreamIds);
24
+ }
25
+ return resumedStreamIds;
26
+ };
27
+
28
+ type ResumableStorageSubscription = {
29
+ listener: () => void;
30
+ threadId: string | undefined;
31
+ unsubscribe: (() => void) | undefined;
32
+ };
33
+
34
+ class DynamicResumableStorage implements ResumableClientStorage {
35
+ private readonly subscriptions = new Set<ResumableStorageSubscription>();
36
+ private storage: ResumableClientStorage | undefined;
37
+ private resumedStreamIds: Set<string>;
38
+ private hasPendingNotification = false;
39
+
40
+ constructor(storage: ResumableClientStorage | undefined) {
41
+ this.storage = storage;
42
+ this.resumedStreamIds = getResumedStreamIds(storage);
43
+ }
44
+
45
+ public setStorage(storage: ResumableClientStorage | undefined) {
46
+ if (this.storage === storage) return;
47
+ this.storage = storage;
48
+ const nextResumedStreamIds = getResumedStreamIds(storage);
49
+ for (const streamId of this.resumedStreamIds) {
50
+ nextResumedStreamIds.add(streamId);
51
+ }
52
+ this.resumedStreamIds = nextResumedStreamIds;
53
+ this.hasPendingNotification = true;
54
+ }
55
+
56
+ public flushChange() {
57
+ if (!this.hasPendingNotification) return;
58
+ this.hasPendingNotification = false;
59
+ for (const subscription of this.subscriptions) {
60
+ subscription.unsubscribe?.();
61
+ subscription.unsubscribe = this.storage?.subscribe?.(
62
+ subscription.listener,
63
+ subscription.threadId,
64
+ );
65
+ subscription.listener();
66
+ }
67
+ }
68
+
69
+ public getResumedStreamIds() {
70
+ return this.resumedStreamIds;
71
+ }
72
+
73
+ public getStreamId(threadId?: string) {
74
+ return this.storage?.getStreamId(threadId) ?? null;
75
+ }
76
+
77
+ public setStreamId(id: string, threadId?: string) {
78
+ this.storage?.setStreamId(id, threadId);
79
+ }
80
+
81
+ public clear(threadId?: string) {
82
+ this.storage?.clear(threadId);
83
+ }
84
+
85
+ public subscribe(listener: () => void, threadId?: string) {
86
+ const subscription: ResumableStorageSubscription = {
87
+ listener,
88
+ threadId,
89
+ unsubscribe: this.storage?.subscribe?.(listener, threadId),
90
+ };
91
+ this.subscriptions.add(subscription);
92
+ return () => {
93
+ this.subscriptions.delete(subscription);
94
+ subscription.unsubscribe?.();
95
+ };
96
+ }
97
+ }
98
+
99
+ type ThreadTransportContext<UI_MESSAGE extends UIMessage> = {
100
+ owner: object;
101
+ sourceTransport?: ChatTransport<UI_MESSAGE> | undefined;
102
+ transport?: ChatTransport<UI_MESSAGE> | undefined;
103
+ runtime?: AssistantRuntime | undefined;
104
+ getThreadListItem?:
105
+ | (() => InitializableThreadListItem | undefined)
106
+ | undefined;
107
+ };
108
+
109
+ type ThreadTransportBinding = {
110
+ runtime: AssistantRuntime;
111
+ getThreadListItem: () => InitializableThreadListItem | undefined;
112
+ };
113
+
114
+ export class DynamicChatTransport<
115
+ UI_MESSAGE extends UIMessage,
116
+ > implements ChatTransport<UI_MESSAGE> {
117
+ private readonly threadContexts = new Map<
118
+ string,
119
+ ThreadTransportContext<UI_MESSAGE>
120
+ >();
121
+ private transport: ChatTransport<UI_MESSAGE>;
122
+ private readonly resumableStorage: DynamicResumableStorage;
123
+
124
+ constructor(transport: ChatTransport<UI_MESSAGE>) {
125
+ this.transport = transport;
126
+ this.resumableStorage = new DynamicResumableStorage(
127
+ getResumableAdapter(transport)?.storage,
128
+ );
129
+ }
130
+
131
+ public readonly sendMessages: ChatTransport<UI_MESSAGE>["sendMessages"] = (
132
+ options,
133
+ ) => this.getTransport(options.chatId).sendMessages(options);
134
+
135
+ public readonly reconnectToStream: ChatTransport<UI_MESSAGE>["reconnectToStream"] =
136
+ (options) => this.getTransport(options.chatId).reconnectToStream(options);
137
+
138
+ public readonly getCurrentTransport = (chatId: string) =>
139
+ this.getThreadTransport(this.threadContexts.get(chatId)) ?? this.transport;
140
+
141
+ public readonly getCurrentResumableStorage = () => this.resumableStorage;
142
+
143
+ public readonly getResumedStreamIds = () =>
144
+ this.resumableStorage.getResumedStreamIds();
145
+
146
+ public createThreadProxy(
147
+ owner: object,
148
+ getBinding: () => ThreadTransportBinding,
149
+ ): ChatTransport<UI_MESSAGE> {
150
+ return {
151
+ sendMessages: (options) =>
152
+ this.getOrCreateBoundTransport(
153
+ options.chatId,
154
+ owner,
155
+ getBinding,
156
+ ).sendMessages(options),
157
+ reconnectToStream: (options) =>
158
+ this.getOrCreateBoundTransport(
159
+ options.chatId,
160
+ owner,
161
+ getBinding,
162
+ ).reconnectToStream(options),
163
+ };
164
+ }
165
+
166
+ public setTransport(transport: ChatTransport<UI_MESSAGE>) {
167
+ if (this.transport === transport) return;
168
+ this.transport = transport;
169
+ this.resumableStorage.setStorage(getResumableAdapter(transport)?.storage);
170
+ }
171
+
172
+ public flushTransportChange() {
173
+ this.resumableStorage.flushChange();
174
+ }
175
+
176
+ public registerThread(chatId: string, owner: object) {
177
+ const existing = this.threadContexts.get(chatId);
178
+ if (existing?.owner === owner) return;
179
+ this.threadContexts.set(chatId, {
180
+ owner,
181
+ });
182
+ }
183
+
184
+ public setThreadContext(
185
+ chatId: string,
186
+ owner: object,
187
+ runtime: AssistantRuntime,
188
+ getThreadListItem: () => InitializableThreadListItem | undefined,
189
+ ) {
190
+ this.registerThread(chatId, owner);
191
+ const context = this.threadContexts.get(chatId)!;
192
+ context.runtime = runtime;
193
+ context.getThreadListItem = getThreadListItem;
194
+ if (context.transport === undefined) {
195
+ this.getThreadTransport(context);
196
+ return;
197
+ }
198
+ if (
199
+ context.sourceTransport === this.transport &&
200
+ context.transport !== undefined
201
+ ) {
202
+ this.wireTransport(context, context.transport);
203
+ }
204
+ }
205
+
206
+ public unregisterThread(chatId: string, owner: object) {
207
+ if (this.threadContexts.get(chatId)?.owner === owner) {
208
+ this.threadContexts.delete(chatId);
209
+ }
210
+ }
211
+
212
+ private getTransport(chatId: string) {
213
+ const context = this.threadContexts.get(chatId);
214
+ if (!context) {
215
+ throw new Error(
216
+ `DynamicChatTransport has no registered context for chat "${chatId}"`,
217
+ );
218
+ }
219
+ return this.getThreadTransport(context)!;
220
+ }
221
+
222
+ private getBoundTransport(
223
+ chatId: string,
224
+ owner: object,
225
+ binding: ThreadTransportBinding,
226
+ ) {
227
+ this.registerThread(chatId, owner);
228
+ const context = this.threadContexts.get(chatId)!;
229
+ context.runtime = binding.runtime;
230
+ context.getThreadListItem = binding.getThreadListItem;
231
+ return this.getThreadTransport(context)!;
232
+ }
233
+
234
+ private getOrCreateBoundTransport(
235
+ chatId: string,
236
+ owner: object,
237
+ getBinding: () => ThreadTransportBinding,
238
+ ) {
239
+ const context = this.threadContexts.get(chatId);
240
+ return context?.owner === owner
241
+ ? this.getThreadTransport(context)!
242
+ : this.getBoundTransport(chatId, owner, getBinding());
243
+ }
244
+
245
+ private getThreadTransport(
246
+ context: ThreadTransportContext<UI_MESSAGE> | undefined,
247
+ ) {
248
+ if (!context) return undefined;
249
+ if (context.sourceTransport !== this.transport) {
250
+ context.transport = this.createThreadTransport(this.transport);
251
+ context.sourceTransport = this.transport;
252
+ }
253
+ this.wireTransport(context, context.transport!);
254
+ return context.transport!;
255
+ }
256
+
257
+ private createThreadTransport(transport: ChatTransport<UI_MESSAGE>) {
258
+ return transport instanceof AssistantChatTransport
259
+ ? transport.__internal_clone()
260
+ : transport;
261
+ }
262
+
263
+ private wireTransport(
264
+ context: ThreadTransportContext<UI_MESSAGE>,
265
+ transport: ChatTransport<UI_MESSAGE>,
266
+ ) {
267
+ if (!(transport instanceof AssistantChatTransport)) return;
268
+ if (context.runtime) transport.setRuntime(context.runtime);
269
+ if (context.getThreadListItem) {
270
+ transport.__internal_setGetThreadListItem(context.getThreadListItem);
271
+ }
272
+ }
273
+ }
@@ -40,6 +40,9 @@ export const createCancellableTransport = () => {
40
40
  return {
41
41
  transport,
42
42
  getCancelCount: () => cancelCount,
43
+ emit: (...chunks: UIMessageChunk[]) => {
44
+ for (const chunk of chunks) controller.enqueue(chunk);
45
+ },
43
46
  close: () => controller.close(),
44
47
  };
45
48
  };
@@ -0,0 +1,16 @@
1
+ import type { UIMessage } from "@ai-sdk/react";
2
+ import type { ChatTransport } from "ai";
3
+ import { AssistantChatTransport } from "../transport/AssistantChatTransport";
4
+ import type { AssistantChatResumableOptions } from "../transport/resumable";
5
+
6
+ export const getResumableAdapter = <UI_MESSAGE extends UIMessage>(
7
+ transport: ChatTransport<UI_MESSAGE>,
8
+ ): AssistantChatResumableOptions | undefined => {
9
+ if (transport instanceof AssistantChatTransport) {
10
+ return transport.getResumableAdapter();
11
+ }
12
+ const candidate = (transport as { getResumableAdapter?: () => unknown })
13
+ .getResumableAdapter;
14
+ if (typeof candidate !== "function") return undefined;
15
+ return candidate.call(transport) as AssistantChatResumableOptions | undefined;
16
+ };
@@ -0,0 +1,161 @@
1
+ import type {
2
+ RespondToToolApprovalOptions,
3
+ ThreadMessage,
4
+ Unstable_ToolInteractionLog,
5
+ } from "@assistant-ui/core";
6
+ import { describe, expect, it } from "vitest";
7
+ import {
8
+ addToolData,
9
+ collectToolApprovalResponses,
10
+ collectToolArtifacts,
11
+ collectToolInteractions,
12
+ restoreToolData,
13
+ } from "./toolHistoryCodec";
14
+
15
+ describe("tool history codec", () => {
16
+ it("round trips artifacts, interactions, and approval responses", () => {
17
+ const log: Unstable_ToolInteractionLog = {
18
+ entries: [{ type: "action", occurredAt: 1, payload: { refresh: true } }],
19
+ };
20
+ const response: RespondToToolApprovalOptions = {
21
+ approvalId: "approval-1",
22
+ approved: true,
23
+ answers: { scope: { optionIds: ["tests"] } },
24
+ reason: "Approved",
25
+ };
26
+ const threadMessage: ThreadMessage = {
27
+ id: "assistant-1",
28
+ role: "assistant",
29
+ content: [
30
+ {
31
+ type: "tool-call",
32
+ toolCallId: "call-1",
33
+ toolName: "weather",
34
+ args: {},
35
+ argsText: "{}",
36
+ result: undefined,
37
+ isError: false,
38
+ approval: { id: "approval-1" },
39
+ },
40
+ ],
41
+ createdAt: new Date(),
42
+ status: { type: "complete", reason: "stop" },
43
+ metadata: {
44
+ unstable_state: null,
45
+ unstable_annotations: [],
46
+ unstable_data: [],
47
+ steps: [],
48
+ custom: {},
49
+ },
50
+ };
51
+ const innerMessage = {
52
+ id: "inner-1",
53
+ parts: [{ toolCallId: "call-1", approval: { id: "approval-1" } }],
54
+ metadata: { custom: "kept" },
55
+ };
56
+ const artifacts = new Map<string, unknown>([
57
+ ["call-1", { preview: "sunny" }],
58
+ ["unrelated", "ignored"],
59
+ ]);
60
+ const interactions = new Map<string, Unstable_ToolInteractionLog>([
61
+ ["call-1", log],
62
+ ["unrelated", log],
63
+ ]);
64
+ const responses = new Map<string, RespondToToolApprovalOptions>([
65
+ ["approval-1", response],
66
+ ["unrelated", response],
67
+ ]);
68
+
69
+ const encoded = addToolData(
70
+ innerMessage,
71
+ collectToolArtifacts(threadMessage, artifacts),
72
+ collectToolInteractions(threadMessage, interactions),
73
+ collectToolApprovalResponses(threadMessage, responses),
74
+ );
75
+
76
+ expect(encoded).toEqual({
77
+ ...innerMessage,
78
+ metadata: {
79
+ custom: "kept",
80
+ __aui_toolArtifacts: { "call-1": { preview: "sunny" } },
81
+ __aui_toolInteractions: { "call-1": log },
82
+ __aui_toolApprovalResponses: {
83
+ "approval-1": {
84
+ approved: true,
85
+ answers: response.answers,
86
+ reason: "Approved",
87
+ },
88
+ },
89
+ },
90
+ });
91
+ expect(innerMessage.metadata).toEqual({ custom: "kept" });
92
+
93
+ const restoredArtifacts = new Map<string, unknown>();
94
+ const restoredInteractions = new Map<string, Unstable_ToolInteractionLog>();
95
+ const restoredResponses = new Map<string, RespondToToolApprovalOptions>();
96
+ expect(
97
+ restoreToolData(
98
+ encoded,
99
+ restoredArtifacts,
100
+ restoredInteractions,
101
+ restoredResponses,
102
+ ),
103
+ ).toEqual(innerMessage);
104
+ expect(restoredArtifacts).toEqual(
105
+ new Map([["call-1", { preview: "sunny" }]]),
106
+ );
107
+ expect(restoredInteractions).toEqual(new Map([["call-1", log]]));
108
+ expect(restoredResponses).toEqual(new Map([["approval-1", response]]));
109
+ });
110
+
111
+ it("leaves messages without tool metadata unchanged", () => {
112
+ const message = {
113
+ parts: [{ toolCallId: "call-1" }],
114
+ metadata: { custom: 1 },
115
+ };
116
+ expect(addToolData(message, undefined, undefined, undefined)).toBe(message);
117
+ expect(restoreToolData(message, new Map(), new Map(), new Map())).toBe(
118
+ message,
119
+ );
120
+ expect(
121
+ restoreToolData({ parts: [] }, new Map(), new Map(), new Map()),
122
+ ).toEqual({ parts: [] });
123
+ });
124
+
125
+ it("ignores non-object messages and malformed metadata values", () => {
126
+ expect(
127
+ addToolData(null, { "call-1": "artifact" }, undefined, undefined),
128
+ ).toBeNull();
129
+ const invalidParts = { parts: "invalid" };
130
+ expect(
131
+ addToolData(invalidParts, { "call-1": "artifact" }, undefined, undefined),
132
+ ).toBe(invalidParts);
133
+ expect(restoreToolData(42, new Map(), new Map(), new Map())).toBe(42);
134
+ const invalidMetadata = { parts: [], metadata: "invalid" };
135
+ expect(
136
+ restoreToolData(invalidMetadata, new Map(), new Map(), new Map()),
137
+ ).toBe(invalidMetadata);
138
+
139
+ const artifacts = new Map<string, unknown>();
140
+ const interactions = new Map<string, Unstable_ToolInteractionLog>();
141
+ const responses = new Map<string, RespondToToolApprovalOptions>();
142
+ expect(
143
+ restoreToolData(
144
+ {
145
+ metadata: {
146
+ custom: "kept",
147
+ __aui_toolArtifacts: "invalid",
148
+ __aui_toolInteractions: 1,
149
+ __aui_toolApprovalResponses: [],
150
+ },
151
+ },
152
+ artifacts,
153
+ interactions,
154
+ responses,
155
+ ),
156
+ ).toEqual({ metadata: { custom: "kept" } });
157
+ expect(artifacts.size).toBe(0);
158
+ expect(interactions.size).toBe(0);
159
+ expect(responses.size).toBe(0);
160
+ });
161
+ });
@@ -0,0 +1,207 @@
1
+ import type {
2
+ RespondToToolApprovalOptions,
3
+ ThreadMessage,
4
+ Unstable_ToolInteractionLog,
5
+ } from "@assistant-ui/core";
6
+ import { isRecord, readToolInteractionLog } from "@assistant-ui/core/internal";
7
+ import { normalizeToolApprovalAnswers } from "../converters/toolApprovalAnswers";
8
+
9
+ const TOOL_ARTIFACTS_METADATA_KEY = "__aui_toolArtifacts";
10
+ const TOOL_INTERACTIONS_METADATA_KEY = "__aui_toolInteractions";
11
+ const TOOL_APPROVAL_RESPONSES_METADATA_KEY = "__aui_toolApprovalResponses";
12
+
13
+ export type StoredToolApprovalResponse = Omit<
14
+ RespondToToolApprovalOptions,
15
+ "approvalId"
16
+ >;
17
+
18
+ export const collectToolArtifacts = (
19
+ message: ThreadMessage,
20
+ toolArtifacts: ReadonlyMap<string, unknown> | undefined,
21
+ ) => {
22
+ if (!toolArtifacts) return undefined;
23
+ const entries = message.content.flatMap((part) => {
24
+ if (part.type !== "tool-call") return [];
25
+ const artifact = toolArtifacts.get(part.toolCallId);
26
+ return artifact === undefined ? [] : [[part.toolCallId, artifact] as const];
27
+ });
28
+ return entries.length > 0 ? Object.fromEntries(entries) : undefined;
29
+ };
30
+
31
+ export const collectToolInteractions = (
32
+ message: ThreadMessage,
33
+ toolInteractions:
34
+ | ReadonlyMap<string, Unstable_ToolInteractionLog>
35
+ | undefined,
36
+ ) => {
37
+ if (!toolInteractions) return undefined;
38
+ const entries = message.content.flatMap((part) => {
39
+ if (part.type !== "tool-call") return [];
40
+ const interactions = toolInteractions.get(part.toolCallId);
41
+ return interactions === undefined
42
+ ? []
43
+ : [[part.toolCallId, interactions] as const];
44
+ });
45
+ return entries.length > 0 ? Object.fromEntries(entries) : undefined;
46
+ };
47
+
48
+ export const collectToolApprovalResponses = (
49
+ message: ThreadMessage,
50
+ toolApprovalResponses:
51
+ | ReadonlyMap<string, RespondToToolApprovalOptions>
52
+ | undefined,
53
+ ) => {
54
+ if (!toolApprovalResponses) return undefined;
55
+ const entries = message.content.flatMap((part) => {
56
+ if (part.type !== "tool-call" || !part.approval) return [];
57
+ const response = toolApprovalResponses.get(part.approval.id);
58
+ if (!response) return [];
59
+ return [
60
+ [
61
+ part.approval.id,
62
+ {
63
+ approved: response.approved,
64
+ ...(response.optionId != null && { optionId: response.optionId }),
65
+ ...(response.text != null && { text: response.text }),
66
+ ...(response.answers != null && { answers: response.answers }),
67
+ ...(response.reason != null && { reason: response.reason }),
68
+ },
69
+ ] as const,
70
+ ];
71
+ });
72
+ return entries.length > 0 ? Object.fromEntries(entries) : undefined;
73
+ };
74
+
75
+ export const addToolData = <TMessage>(
76
+ message: TMessage,
77
+ toolArtifacts: Record<string, unknown> | undefined,
78
+ toolInteractions: Record<string, Unstable_ToolInteractionLog> | undefined,
79
+ toolApprovalResponses: Record<string, StoredToolApprovalResponse> | undefined,
80
+ ): TMessage => {
81
+ if (
82
+ (!toolArtifacts && !toolInteractions && !toolApprovalResponses) ||
83
+ !isRecord(message) ||
84
+ !Array.isArray(message.parts)
85
+ )
86
+ return message;
87
+ const toolCallIds = message.parts.flatMap((part) => {
88
+ if (!isRecord(part) || typeof part.toolCallId !== "string") return [];
89
+ return [part.toolCallId];
90
+ });
91
+ const artifacts = toolArtifacts
92
+ ? Object.fromEntries(
93
+ toolCallIds.flatMap((toolCallId) =>
94
+ Object.hasOwn(toolArtifacts, toolCallId)
95
+ ? [[toolCallId, toolArtifacts[toolCallId]] as const]
96
+ : [],
97
+ ),
98
+ )
99
+ : undefined;
100
+ const interactions = toolInteractions
101
+ ? Object.fromEntries(
102
+ toolCallIds.flatMap((toolCallId) =>
103
+ Object.hasOwn(toolInteractions, toolCallId)
104
+ ? [[toolCallId, toolInteractions[toolCallId]] as const]
105
+ : [],
106
+ ),
107
+ )
108
+ : undefined;
109
+ const approvalIds = message.parts.flatMap((part) => {
110
+ if (!isRecord(part) || !isRecord(part.approval)) return [];
111
+ const approvalId = part.approval.id;
112
+ return typeof approvalId === "string" ? [approvalId] : [];
113
+ });
114
+ const approvalResponses = toolApprovalResponses
115
+ ? Object.fromEntries(
116
+ approvalIds.flatMap((approvalId) =>
117
+ Object.hasOwn(toolApprovalResponses, approvalId)
118
+ ? [[approvalId, toolApprovalResponses[approvalId]] as const]
119
+ : [],
120
+ ),
121
+ )
122
+ : undefined;
123
+ const hasArtifacts = !!artifacts && Object.keys(artifacts).length > 0;
124
+ const hasInteractions =
125
+ !!interactions && Object.keys(interactions).length > 0;
126
+ const hasApprovalResponses =
127
+ !!approvalResponses && Object.keys(approvalResponses).length > 0;
128
+ if (!hasArtifacts && !hasInteractions && !hasApprovalResponses)
129
+ return message;
130
+ const metadata = isRecord(message.metadata) ? message.metadata : {};
131
+ return {
132
+ ...message,
133
+ metadata: {
134
+ ...metadata,
135
+ ...(hasArtifacts && { [TOOL_ARTIFACTS_METADATA_KEY]: artifacts }),
136
+ ...(hasInteractions && {
137
+ [TOOL_INTERACTIONS_METADATA_KEY]: interactions,
138
+ }),
139
+ ...(hasApprovalResponses && {
140
+ [TOOL_APPROVAL_RESPONSES_METADATA_KEY]: approvalResponses,
141
+ }),
142
+ },
143
+ } as TMessage;
144
+ };
145
+
146
+ export const restoreToolData = <TMessage>(
147
+ message: TMessage,
148
+ toolArtifacts: Map<string, unknown> | undefined,
149
+ toolInteractions: Map<string, Unstable_ToolInteractionLog> | undefined,
150
+ toolApprovalResponses: Map<string, RespondToToolApprovalOptions> | undefined,
151
+ ): TMessage => {
152
+ if (!isRecord(message) || !isRecord(message.metadata)) return message;
153
+ const metadata = message.metadata;
154
+ const hasArtifacts = Object.hasOwn(metadata, TOOL_ARTIFACTS_METADATA_KEY);
155
+ const hasInteractions = Object.hasOwn(
156
+ metadata,
157
+ TOOL_INTERACTIONS_METADATA_KEY,
158
+ );
159
+ const hasApprovalResponses = Object.hasOwn(
160
+ metadata,
161
+ TOOL_APPROVAL_RESPONSES_METADATA_KEY,
162
+ );
163
+ if (!hasArtifacts && !hasInteractions && !hasApprovalResponses)
164
+ return message;
165
+ const artifacts = metadata[TOOL_ARTIFACTS_METADATA_KEY];
166
+ if (toolArtifacts && isRecord(artifacts)) {
167
+ for (const [toolCallId, artifact] of Object.entries(artifacts)) {
168
+ toolArtifacts.set(toolCallId, artifact);
169
+ }
170
+ }
171
+ const interactions = metadata[TOOL_INTERACTIONS_METADATA_KEY];
172
+ if (toolInteractions && isRecord(interactions)) {
173
+ for (const [toolCallId, value] of Object.entries(interactions)) {
174
+ const log = readToolInteractionLog(value);
175
+ if (log) toolInteractions.set(toolCallId, log);
176
+ }
177
+ }
178
+ const approvalResponses = metadata[TOOL_APPROVAL_RESPONSES_METADATA_KEY];
179
+ if (toolApprovalResponses && isRecord(approvalResponses)) {
180
+ for (const [approvalId, value] of Object.entries(approvalResponses)) {
181
+ if (!isRecord(value) || typeof value.approved !== "boolean") continue;
182
+ const answers = normalizeToolApprovalAnswers(value.answers);
183
+ toolApprovalResponses.set(approvalId, {
184
+ approvalId,
185
+ approved: value.approved,
186
+ ...(typeof value.optionId === "string" && {
187
+ optionId: value.optionId,
188
+ }),
189
+ ...(typeof value.text === "string" && { text: value.text }),
190
+ ...(answers && Object.keys(answers).length > 0 && { answers }),
191
+ ...(typeof value.reason === "string" && { reason: value.reason }),
192
+ });
193
+ }
194
+ }
195
+ const {
196
+ [TOOL_ARTIFACTS_METADATA_KEY]: _,
197
+ [TOOL_INTERACTIONS_METADATA_KEY]: __,
198
+ [TOOL_APPROVAL_RESPONSES_METADATA_KEY]: ___,
199
+ ...restMetadata
200
+ } = metadata;
201
+ const { metadata: _metadata, ...restMessage } = message;
202
+ return (
203
+ Object.keys(restMetadata).length === 0
204
+ ? restMessage
205
+ : { ...restMessage, metadata: restMetadata }
206
+ ) as TMessage;
207
+ };
@@ -203,6 +203,18 @@ describe("useAISDKRuntime tool approvals", () => {
203
203
  });
204
204
  });
205
205
 
206
+ it("finds the pending approval past malformed parts", async () => {
207
+ const onRespondToToolApproval = vi.fn(async () => {});
208
+ const { respond, messages } = setupPendingApproval(onRespondToToolApproval);
209
+ (messages[0]!.parts as unknown[]).unshift(null, { text: "no type" });
210
+
211
+ await act(async () => {
212
+ await respond({ approvalId: "approval-1", approved: true });
213
+ });
214
+
215
+ expect(onRespondToToolApproval).toHaveBeenCalledOnce();
216
+ });
217
+
206
218
  it("does not store a request the handler hands back through the AI SDK", async () => {
207
219
  const { respond, addToolApprovalResponse } = setupPendingApproval(
208
220
  (_response, { respondViaAISDK }) => respondViaAISDK(),