@assistant-ui/ai-sdk 0.0.5 → 0.0.7
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/adapters/aiSDKFormatAdapter.d.ts +2 -8
- package/dist/adapters/aiSDKFormatAdapter.js +1 -25
- package/dist/adapters/vercelAttachmentAdapter.d.ts +1 -2
- package/dist/adapters/vercelAttachmentAdapter.d.ts.map +1 -1
- package/dist/aiSDKExtras.d.ts +2 -3
- package/dist/aiSDKExtras.d.ts.map +1 -1
- package/dist/converters/convertMessage.d.ts +5 -4
- package/dist/converters/convertMessage.d.ts.map +1 -1
- package/dist/converters/convertMessage.js +40 -3
- package/dist/converters/convertMessage.js.map +1 -1
- package/dist/converters/modelContentEnvelope.d.ts +4 -5
- package/dist/converters/modelContentEnvelope.d.ts.map +1 -1
- package/dist/converters/toCreateMessage.d.ts +1 -2
- package/dist/converters/toCreateMessage.d.ts.map +1 -1
- package/dist/converters/toolOutputConversion.d.ts +2 -3
- package/dist/converters/toolOutputConversion.d.ts.map +1 -1
- package/dist/hooks.d.ts +2 -3
- package/dist/hooks.d.ts.map +1 -1
- package/dist/model-context/injectInteractableContext.d.ts +1 -2
- package/dist/model-context/injectInteractableContext.d.ts.map +1 -1
- package/dist/model-context/injectQuoteContext.d.ts +1 -2
- package/dist/model-context/injectQuoteContext.d.ts.map +1 -1
- package/dist/runtime/AISDKChat.d.ts +2 -3
- package/dist/runtime/AISDKChat.d.ts.map +1 -1
- package/dist/runtime/AISDKThreads.d.ts +2 -3
- package/dist/runtime/AISDKThreads.d.ts.map +1 -1
- package/dist/runtime/AISDKThreads.js +87 -68
- package/dist/runtime/AISDKThreads.js.map +1 -1
- package/dist/runtime/sdkIdentity.d.ts +5 -0
- package/dist/runtime/sdkIdentity.d.ts.map +1 -0
- package/dist/runtime/sdkIdentity.js +9 -0
- package/dist/runtime/sdkIdentity.js.map +1 -0
- package/dist/runtime/useAISDKRuntime.d.ts +17 -5
- package/dist/runtime/useAISDKRuntime.d.ts.map +1 -1
- package/dist/runtime/useAISDKRuntime.js +103 -38
- package/dist/runtime/useAISDKRuntime.js.map +1 -1
- package/dist/runtime/useChatRuntime.d.ts +2 -3
- package/dist/runtime/useChatRuntime.d.ts.map +1 -1
- package/dist/runtime/useChatRuntime.js +5 -1
- package/dist/runtime/useChatRuntime.js.map +1 -1
- package/dist/runtime/useChatThread.d.ts +19 -5
- package/dist/runtime/useChatThread.d.ts.map +1 -1
- package/dist/runtime/useChatThread.js +31 -10
- package/dist/runtime/useChatThread.js.map +1 -1
- package/dist/runtime/useExternalHistory.d.ts +2 -3
- package/dist/runtime/useExternalHistory.d.ts.map +1 -1
- package/dist/runtime/useExternalHistory.js +31 -20
- package/dist/runtime/useExternalHistory.js.map +1 -1
- package/dist/runtime/useResourceCleanup.d.ts +1 -2
- package/dist/runtime/useResourceCleanup.d.ts.map +1 -1
- package/dist/runtime/useStreamingTiming.d.ts +2 -3
- package/dist/runtime/useStreamingTiming.d.ts.map +1 -1
- package/dist/tools/frontendTools.d.ts +3 -4
- package/dist/tools/frontendTools.d.ts.map +1 -1
- package/dist/tools/generativeTools.d.ts +5 -6
- package/dist/tools/generativeTools.d.ts.map +1 -1
- package/dist/tools/mcp-stdio.unsupported.d.ts +1 -2
- package/dist/tools/mcp-stdio.unsupported.d.ts.map +1 -1
- package/dist/transport/AssistantChatTransport.d.ts +3 -4
- package/dist/transport/AssistantChatTransport.d.ts.map +1 -1
- package/dist/transport/resumable.d.ts +4 -5
- package/dist/transport/resumable.d.ts.map +1 -1
- package/dist/usage.d.ts +4 -5
- package/dist/usage.d.ts.map +1 -1
- package/dist/utils/getVercelAIMessages.d.ts +1 -2
- package/dist/utils/getVercelAIMessages.d.ts.map +1 -1
- package/dist/utils/sliceMessagesUntil.d.ts +1 -2
- package/dist/utils/sliceMessagesUntil.d.ts.map +1 -1
- package/package.json +15 -14
- package/src/adapters/aiSDKFormatAdapter.ts +4 -41
- package/src/converters/convertMessage.test.ts +152 -0
- package/src/converters/convertMessage.ts +95 -5
- package/src/runtime/AISDKChat.test.ts +4 -5
- package/src/runtime/AISDKThreads.cloud.test.ts +12 -3
- package/src/runtime/AISDKThreads.test.ts +140 -13
- package/src/runtime/AISDKThreads.ts +23 -4
- package/src/runtime/sdkIdentity.ts +9 -0
- package/src/runtime/useAISDKRuntime.approval.integration.test.tsx +287 -0
- package/src/runtime/useAISDKRuntime.approval.test.tsx +257 -1
- package/src/runtime/useAISDKRuntime.test.ts +46 -4
- package/src/runtime/useAISDKRuntime.ts +164 -24
- package/src/runtime/useAISDKRuntime.voice.test.tsx +270 -0
- package/src/runtime/useChatRuntime.test.ts +72 -5
- package/src/runtime/useChatRuntime.ts +2 -1
- package/src/runtime/useChatThread.test.ts +74 -0
- package/src/runtime/useChatThread.ts +45 -6
- package/src/runtime/useExternalHistory.test.ts +75 -0
- package/src/runtime/useExternalHistory.ts +28 -11
- package/src/tools/generativeTools.test.ts +7 -1
- package/src/transport/AssistantChatTransport.test.ts +1 -9
- package/dist/adapters/aiSDKFormatAdapter.d.ts.map +0 -1
- package/dist/adapters/aiSDKFormatAdapter.js.map +0 -1
|
@@ -48,6 +48,20 @@ describe("AISDKMessageConverter", () => {
|
|
|
48
48
|
expect(converted[0]?.metadata).not.toHaveProperty("usage");
|
|
49
49
|
});
|
|
50
50
|
|
|
51
|
+
it("keeps modality metadata at the top level", () => {
|
|
52
|
+
const converted = AISDKMessageConverter.toThreadMessages([
|
|
53
|
+
{
|
|
54
|
+
id: "a1",
|
|
55
|
+
role: "assistant",
|
|
56
|
+
parts: [{ type: "text", text: "yo" }],
|
|
57
|
+
metadata: { modality: "voice" },
|
|
58
|
+
},
|
|
59
|
+
] as any);
|
|
60
|
+
|
|
61
|
+
expect(converted[0]?.metadata.modality).toBe("voice");
|
|
62
|
+
expect(converted[0]?.metadata.custom).not.toHaveProperty("modality");
|
|
63
|
+
});
|
|
64
|
+
|
|
51
65
|
it("does not flag messages when no optimistic id is provided", () => {
|
|
52
66
|
const converted = AISDKMessageConverter.toThreadMessages([
|
|
53
67
|
{ id: "a1", role: "assistant", parts: [{ type: "text", text: "yo" }] },
|
|
@@ -473,6 +487,144 @@ describe("AISDKMessageConverter", () => {
|
|
|
473
487
|
});
|
|
474
488
|
});
|
|
475
489
|
|
|
490
|
+
it("preserves rich approval fields for a custom response channel", () => {
|
|
491
|
+
const metadata: AISDKMessageConverterMetadata = {
|
|
492
|
+
supportsRichToolApprovalResponses: true,
|
|
493
|
+
};
|
|
494
|
+
const converted = AISDKMessageConverter.toThreadMessages(
|
|
495
|
+
[
|
|
496
|
+
{
|
|
497
|
+
id: "a1",
|
|
498
|
+
role: "assistant",
|
|
499
|
+
parts: [
|
|
500
|
+
{
|
|
501
|
+
type: "tool-deploy",
|
|
502
|
+
toolCallId: "tc-1",
|
|
503
|
+
state: "approval-responded",
|
|
504
|
+
input: {},
|
|
505
|
+
approval: {
|
|
506
|
+
id: "approval-1",
|
|
507
|
+
display: "select",
|
|
508
|
+
allowFreeform: true,
|
|
509
|
+
options: [
|
|
510
|
+
{
|
|
511
|
+
id: "once",
|
|
512
|
+
kind: "allow-once",
|
|
513
|
+
label: "Only once",
|
|
514
|
+
grants: ["repository", 42],
|
|
515
|
+
confirm: {
|
|
516
|
+
title: "Confirm access",
|
|
517
|
+
description: { invalid: true },
|
|
518
|
+
},
|
|
519
|
+
},
|
|
520
|
+
"invalid",
|
|
521
|
+
{ id: 1, kind: "allow-always" },
|
|
522
|
+
{ id: "always", kind: 2 },
|
|
523
|
+
],
|
|
524
|
+
optionId: "once",
|
|
525
|
+
text: "an answer",
|
|
526
|
+
},
|
|
527
|
+
},
|
|
528
|
+
],
|
|
529
|
+
} as any,
|
|
530
|
+
],
|
|
531
|
+
false,
|
|
532
|
+
metadata,
|
|
533
|
+
);
|
|
534
|
+
|
|
535
|
+
const toolCall = converted[0]?.content.find(
|
|
536
|
+
(part): part is any => part.type === "tool-call",
|
|
537
|
+
);
|
|
538
|
+
expect(toolCall?.approval).toEqual({
|
|
539
|
+
id: "approval-1",
|
|
540
|
+
display: "select",
|
|
541
|
+
allowFreeform: true,
|
|
542
|
+
options: [
|
|
543
|
+
{
|
|
544
|
+
id: "once",
|
|
545
|
+
kind: "allow-once",
|
|
546
|
+
label: "Only once",
|
|
547
|
+
grants: ["repository"],
|
|
548
|
+
confirm: { title: "Confirm access" },
|
|
549
|
+
},
|
|
550
|
+
],
|
|
551
|
+
optionId: "once",
|
|
552
|
+
text: "an answer",
|
|
553
|
+
});
|
|
554
|
+
});
|
|
555
|
+
|
|
556
|
+
it("applies a host answer to an approval the message has not recorded", () => {
|
|
557
|
+
const metadata: AISDKMessageConverterMetadata = {
|
|
558
|
+
supportsRichToolApprovalResponses: true,
|
|
559
|
+
toolApprovalResponses: new Map([
|
|
560
|
+
[
|
|
561
|
+
"approval-1",
|
|
562
|
+
{
|
|
563
|
+
approvalId: "approval-1",
|
|
564
|
+
approved: true,
|
|
565
|
+
optionId: "staging",
|
|
566
|
+
text: "only staging",
|
|
567
|
+
},
|
|
568
|
+
],
|
|
569
|
+
["approval-2", { approvalId: "approval-2", approved: true }],
|
|
570
|
+
["approval-3", { approvalId: "approval-3", approved: true }],
|
|
571
|
+
]),
|
|
572
|
+
};
|
|
573
|
+
const converted = AISDKMessageConverter.toThreadMessages(
|
|
574
|
+
[
|
|
575
|
+
{
|
|
576
|
+
id: "a1",
|
|
577
|
+
role: "assistant",
|
|
578
|
+
parts: [
|
|
579
|
+
{
|
|
580
|
+
type: "tool-deploy",
|
|
581
|
+
toolCallId: "tc-1",
|
|
582
|
+
state: "approval-requested",
|
|
583
|
+
input: {},
|
|
584
|
+
approval: {
|
|
585
|
+
id: "approval-1",
|
|
586
|
+
display: "select",
|
|
587
|
+
options: [{ id: "staging", kind: "_target" }],
|
|
588
|
+
},
|
|
589
|
+
},
|
|
590
|
+
{
|
|
591
|
+
type: "tool-deploy",
|
|
592
|
+
toolCallId: "tc-2",
|
|
593
|
+
state: "approval-responded",
|
|
594
|
+
input: {},
|
|
595
|
+
approval: { id: "approval-2", approved: false, reason: "no" },
|
|
596
|
+
},
|
|
597
|
+
{
|
|
598
|
+
type: "tool-deploy",
|
|
599
|
+
toolCallId: "tc-3",
|
|
600
|
+
state: "approval-requested",
|
|
601
|
+
input: {},
|
|
602
|
+
approval: { id: "approval-3", resolution: "expired" },
|
|
603
|
+
},
|
|
604
|
+
],
|
|
605
|
+
} as any,
|
|
606
|
+
],
|
|
607
|
+
false,
|
|
608
|
+
metadata,
|
|
609
|
+
);
|
|
610
|
+
|
|
611
|
+
const approvals = converted[0]?.content.map(
|
|
612
|
+
(part) => (part as { approval?: unknown }).approval,
|
|
613
|
+
);
|
|
614
|
+
expect(approvals).toEqual([
|
|
615
|
+
{
|
|
616
|
+
id: "approval-1",
|
|
617
|
+
display: "select",
|
|
618
|
+
options: [{ id: "staging", kind: "_target" }],
|
|
619
|
+
approved: true,
|
|
620
|
+
optionId: "staging",
|
|
621
|
+
text: "only staging",
|
|
622
|
+
},
|
|
623
|
+
{ id: "approval-2", approved: false, reason: "no" },
|
|
624
|
+
{ id: "approval-3", resolution: "expired" },
|
|
625
|
+
]);
|
|
626
|
+
});
|
|
627
|
+
|
|
476
628
|
it("drops a resolution the core contract does not declare", () => {
|
|
477
629
|
const converted = AISDKMessageConverter.toThreadMessages([
|
|
478
630
|
{
|
|
@@ -12,6 +12,7 @@ import {
|
|
|
12
12
|
import {
|
|
13
13
|
isMcpAppUri,
|
|
14
14
|
type ReasoningMessagePart,
|
|
15
|
+
type ToolApprovalOption,
|
|
15
16
|
type ToolCallMessagePart,
|
|
16
17
|
type TextMessagePart,
|
|
17
18
|
type DataMessagePart,
|
|
@@ -22,6 +23,7 @@ import {
|
|
|
22
23
|
type ThreadMessageLike,
|
|
23
24
|
type McpAppMetadata,
|
|
24
25
|
type MessagePartStreamStatus,
|
|
26
|
+
type RespondToToolApprovalOptions,
|
|
25
27
|
} from "@assistant-ui/core";
|
|
26
28
|
import { stableStringifyToolArgs } from "@assistant-ui/core/internal";
|
|
27
29
|
import {
|
|
@@ -40,6 +42,7 @@ const THREAD_METADATA_KEYS = new Set([
|
|
|
40
42
|
"timing",
|
|
41
43
|
"submittedFeedback",
|
|
42
44
|
"isOptimistic",
|
|
45
|
+
"modality",
|
|
43
46
|
"custom",
|
|
44
47
|
]);
|
|
45
48
|
|
|
@@ -60,6 +63,8 @@ export type AISDKMessageConverterMetadata =
|
|
|
60
63
|
toolArgsKeyOrderCache?: Map<string, Map<string, string[]>>;
|
|
61
64
|
toolLastInputCache?: Map<string, ReadonlyJSONObject>;
|
|
62
65
|
mcpAppMetadataCache?: Map<string, McpAppMetadata>;
|
|
66
|
+
supportsRichToolApprovalResponses?: boolean;
|
|
67
|
+
toolApprovalResponses?: ReadonlyMap<string, RespondToToolApprovalOptions>;
|
|
63
68
|
/** Id of the currently-streaming message, flagged optimistic (#4037). */
|
|
64
69
|
optimisticMessageId?: string | undefined;
|
|
65
70
|
};
|
|
@@ -152,19 +157,79 @@ function extractMcpAppMetadata(
|
|
|
152
157
|
return out;
|
|
153
158
|
}
|
|
154
159
|
|
|
160
|
+
const normalizeToolApprovalOptions = (
|
|
161
|
+
options: unknown,
|
|
162
|
+
): readonly ToolApprovalOption[] | undefined => {
|
|
163
|
+
if (!Array.isArray(options)) return undefined;
|
|
164
|
+
|
|
165
|
+
return options.flatMap<ToolApprovalOption>((value) => {
|
|
166
|
+
if (!value || typeof value !== "object" || Array.isArray(value)) return [];
|
|
167
|
+
const option = value as Record<string, unknown>;
|
|
168
|
+
if (typeof option.id !== "string" || typeof option.kind !== "string")
|
|
169
|
+
return [];
|
|
170
|
+
|
|
171
|
+
const confirm = option.confirm;
|
|
172
|
+
const confirmDetails =
|
|
173
|
+
confirm && typeof confirm === "object" && !Array.isArray(confirm)
|
|
174
|
+
? (confirm as Record<string, unknown>)
|
|
175
|
+
: undefined;
|
|
176
|
+
|
|
177
|
+
return [
|
|
178
|
+
{
|
|
179
|
+
id: option.id,
|
|
180
|
+
kind: option.kind,
|
|
181
|
+
...(typeof option.label === "string" && { label: option.label }),
|
|
182
|
+
...(typeof option.description === "string" && {
|
|
183
|
+
description: option.description,
|
|
184
|
+
}),
|
|
185
|
+
...(Array.isArray(option.grants) && {
|
|
186
|
+
grants: option.grants.filter(
|
|
187
|
+
(grant): grant is string => typeof grant === "string",
|
|
188
|
+
),
|
|
189
|
+
}),
|
|
190
|
+
...(typeof confirm === "boolean"
|
|
191
|
+
? { confirm }
|
|
192
|
+
: confirmDetails
|
|
193
|
+
? {
|
|
194
|
+
confirm: {
|
|
195
|
+
...(typeof confirmDetails.title === "string" && {
|
|
196
|
+
title: confirmDetails.title,
|
|
197
|
+
}),
|
|
198
|
+
...(typeof confirmDetails.description === "string" && {
|
|
199
|
+
description: confirmDetails.description,
|
|
200
|
+
}),
|
|
201
|
+
},
|
|
202
|
+
}
|
|
203
|
+
: {}),
|
|
204
|
+
},
|
|
205
|
+
];
|
|
206
|
+
});
|
|
207
|
+
};
|
|
208
|
+
|
|
155
209
|
function getToolApprovalAndInterrupt(
|
|
156
210
|
part: {
|
|
157
211
|
approval?: Record<string, unknown> | undefined;
|
|
158
212
|
},
|
|
159
213
|
toolStatus: { type: string; payload?: unknown } | undefined,
|
|
214
|
+
supportsRichToolApprovalResponses: boolean,
|
|
215
|
+
toolApprovalResponses:
|
|
216
|
+
| ReadonlyMap<string, RespondToToolApprovalOptions>
|
|
217
|
+
| undefined,
|
|
160
218
|
): {
|
|
161
219
|
approval?: NonNullable<ToolCallMessagePart["approval"]>;
|
|
162
220
|
interrupt?: NonNullable<ToolCallMessagePart["interrupt"]>;
|
|
163
221
|
} {
|
|
164
222
|
if (part.approval) {
|
|
165
|
-
|
|
166
|
-
|
|
167
|
-
|
|
223
|
+
const response =
|
|
224
|
+
typeof part.approval.id === "string" &&
|
|
225
|
+
part.approval.approved === undefined &&
|
|
226
|
+
part.approval.resolution !== "cancelled" &&
|
|
227
|
+
part.approval.resolution !== "expired"
|
|
228
|
+
? toolApprovalResponses?.get(part.approval.id)
|
|
229
|
+
: undefined;
|
|
230
|
+
// The built-in AI SDK channel sends only id, approved and reason back to
|
|
231
|
+
// the server, so a request shape promising any other answer would render
|
|
232
|
+
// controls whose response cannot travel.
|
|
168
233
|
const {
|
|
169
234
|
id,
|
|
170
235
|
prompt,
|
|
@@ -178,7 +243,18 @@ function getToolApprovalAndInterrupt(
|
|
|
178
243
|
optionId,
|
|
179
244
|
text,
|
|
180
245
|
...additionalApprovalFields
|
|
181
|
-
} =
|
|
246
|
+
} = response
|
|
247
|
+
? {
|
|
248
|
+
...part.approval,
|
|
249
|
+
approved: response.approved,
|
|
250
|
+
...(response.reason != null && { reason: response.reason }),
|
|
251
|
+
...(response.optionId != null && { optionId: response.optionId }),
|
|
252
|
+
...(response.text != null && { text: response.text }),
|
|
253
|
+
}
|
|
254
|
+
: part.approval;
|
|
255
|
+
const normalizedOptions = supportsRichToolApprovalResponses
|
|
256
|
+
? normalizeToolApprovalOptions(options)
|
|
257
|
+
: undefined;
|
|
182
258
|
const requestReason = additionalApprovalFields.requestReason;
|
|
183
259
|
if (typeof id === "string")
|
|
184
260
|
return {
|
|
@@ -193,6 +269,15 @@ function getToolApprovalAndInterrupt(
|
|
|
193
269
|
...(typeof approved === "boolean" && { approved }),
|
|
194
270
|
...(typeof reason === "string" && { reason }),
|
|
195
271
|
...(isAutomatic === true && { isAutomatic: true }),
|
|
272
|
+
...(supportsRichToolApprovalResponses && {
|
|
273
|
+
...((display === "decision" ||
|
|
274
|
+
display === "select" ||
|
|
275
|
+
display === "text") && { display }),
|
|
276
|
+
...(typeof allowFreeform === "boolean" && { allowFreeform }),
|
|
277
|
+
...(normalizedOptions && { options: normalizedOptions }),
|
|
278
|
+
...(typeof optionId === "string" && { optionId }),
|
|
279
|
+
...(typeof text === "string" && { text }),
|
|
280
|
+
}),
|
|
196
281
|
...((resolution === "cancelled" || resolution === "expired") && {
|
|
197
282
|
resolution,
|
|
198
283
|
}),
|
|
@@ -347,7 +432,12 @@ function convertParts(
|
|
|
347
432
|
part.callProviderMetadata as PartProviderMetadata,
|
|
348
433
|
}
|
|
349
434
|
: undefined),
|
|
350
|
-
...getToolApprovalAndInterrupt(
|
|
435
|
+
...getToolApprovalAndInterrupt(
|
|
436
|
+
part,
|
|
437
|
+
toolStatus,
|
|
438
|
+
metadata.supportsRichToolApprovalResponses === true,
|
|
439
|
+
metadata.toolApprovalResponses,
|
|
440
|
+
),
|
|
351
441
|
} satisfies ToolCallMessagePart;
|
|
352
442
|
}
|
|
353
443
|
|
|
@@ -153,7 +153,9 @@ describe("AISDKChat as a standalone client config entry", () => {
|
|
|
153
153
|
]
|
|
154
154
|
.map((chunk) => `data: ${JSON.stringify(chunk)}\n\n`)
|
|
155
155
|
.join("");
|
|
156
|
-
const fetchMock = vi.fn
|
|
156
|
+
const fetchMock = vi.fn<
|
|
157
|
+
(input: RequestInfo | URL, init: RequestInit) => Promise<Response>
|
|
158
|
+
>(
|
|
157
159
|
async () =>
|
|
158
160
|
new Response(sse, {
|
|
159
161
|
headers: { "content-type": "text/event-stream" },
|
|
@@ -178,10 +180,7 @@ describe("AISDKChat as a standalone client config entry", () => {
|
|
|
178
180
|
expect(state.messages).toHaveLength(2);
|
|
179
181
|
});
|
|
180
182
|
|
|
181
|
-
const [url, init] = fetchMock.mock.calls[0]
|
|
182
|
-
RequestInfo,
|
|
183
|
-
RequestInit,
|
|
184
|
-
];
|
|
183
|
+
const [url, init] = fetchMock.mock.calls[0]!;
|
|
185
184
|
expect(String(url)).toContain("/api/chat");
|
|
186
185
|
const body = JSON.parse(init.body as string);
|
|
187
186
|
expect(body.id).toBe("test-thread-1");
|
|
@@ -54,23 +54,28 @@ const mocks = vi.hoisted(() => {
|
|
|
54
54
|
return { history };
|
|
55
55
|
},
|
|
56
56
|
};
|
|
57
|
-
return {
|
|
57
|
+
return {
|
|
58
|
+
adapter,
|
|
59
|
+
useCloudThreadListAdapter: vi.fn(() => adapter),
|
|
60
|
+
};
|
|
58
61
|
});
|
|
59
62
|
|
|
60
63
|
vi.mock("@assistant-ui/core/react", async (importOriginal) => ({
|
|
61
64
|
...(await importOriginal<typeof import("@assistant-ui/core/react")>()),
|
|
62
|
-
useCloudThreadListAdapter:
|
|
65
|
+
useCloudThreadListAdapter: mocks.useCloudThreadListAdapter,
|
|
63
66
|
}));
|
|
64
67
|
|
|
65
68
|
import { AISDKThreads } from "./AISDKThreads";
|
|
66
69
|
import { createCancellableTransport } from "./__tests__/controlled-transport";
|
|
70
|
+
import { AI_SDK_SDK } from "./sdkIdentity";
|
|
67
71
|
|
|
68
72
|
describe("AISDKThreads cloud", () => {
|
|
69
73
|
it("reloads history when switching a keyed cloud thread", async () => {
|
|
74
|
+
const cloud = {} as AssistantCloud;
|
|
70
75
|
const handle = createAssistantClient(
|
|
71
76
|
AuiConfig({
|
|
72
77
|
threads: AISDKThreads({
|
|
73
|
-
cloud
|
|
78
|
+
cloud,
|
|
74
79
|
threadId: "t1",
|
|
75
80
|
}),
|
|
76
81
|
}),
|
|
@@ -84,6 +89,10 @@ describe("AISDKThreads cloud", () => {
|
|
|
84
89
|
await vi.waitFor(() => {
|
|
85
90
|
expect(load).toHaveBeenCalled();
|
|
86
91
|
});
|
|
92
|
+
expect(mocks.useCloudThreadListAdapter).toHaveBeenCalledWith({
|
|
93
|
+
cloud,
|
|
94
|
+
sdk: AI_SDK_SDK,
|
|
95
|
+
});
|
|
87
96
|
const afterFirst = load.mock.calls.length;
|
|
88
97
|
flushTapSync(() => aui.threads.switchToThread("t2"));
|
|
89
98
|
await vi.waitFor(() => {
|
|
@@ -251,17 +251,28 @@ describe("AISDKThreads", () => {
|
|
|
251
251
|
}
|
|
252
252
|
});
|
|
253
253
|
|
|
254
|
-
it("forwards ChatInit callbacks to each thread's chat", async () => {
|
|
254
|
+
it("forwards ChatInit callbacks to each thread's chat from the latest render", async () => {
|
|
255
255
|
const { transport, emit, close } = createControlledTransport();
|
|
256
|
-
const
|
|
257
|
-
const
|
|
258
|
-
|
|
259
|
-
|
|
260
|
-
|
|
261
|
-
|
|
256
|
+
const onFinishA = vi.fn();
|
|
257
|
+
const onFinishB = vi.fn();
|
|
258
|
+
let onFinish = onFinishA;
|
|
259
|
+
const listeners = new Set<() => void>();
|
|
260
|
+
const handle = createAssistantClient({
|
|
261
|
+
getConfig: () =>
|
|
262
|
+
AuiConfig({
|
|
263
|
+
threads: AISDKThreads({ transport: () => transport, onFinish }),
|
|
264
|
+
}),
|
|
265
|
+
subscribe: (listener) => {
|
|
266
|
+
listeners.add(listener);
|
|
267
|
+
return () => listeners.delete(listener);
|
|
268
|
+
},
|
|
269
|
+
});
|
|
262
270
|
handle.subscribe(() => {});
|
|
263
271
|
const aui = handle.getClient();
|
|
264
272
|
|
|
273
|
+
onFinish = onFinishB;
|
|
274
|
+
flushTapSync(() => listeners.forEach((listener) => listener()));
|
|
275
|
+
|
|
265
276
|
flushTapSync(() => aui.composer.setText("hi"));
|
|
266
277
|
flushTapSync(() => aui.composer.send());
|
|
267
278
|
await vi.waitFor(() => {
|
|
@@ -271,11 +282,123 @@ describe("AISDKThreads", () => {
|
|
|
271
282
|
});
|
|
272
283
|
emit(...textReply("done"));
|
|
273
284
|
close();
|
|
274
|
-
await vi.waitFor(() => expect(
|
|
285
|
+
await vi.waitFor(() => expect(onFinishB).toHaveBeenCalledTimes(1));
|
|
286
|
+
expect(onFinishA).not.toHaveBeenCalled();
|
|
287
|
+
|
|
288
|
+
handle.destroy();
|
|
289
|
+
});
|
|
290
|
+
|
|
291
|
+
it("forwards the latest callbacks to a switched-away thread still streaming in the background", async () => {
|
|
292
|
+
const { transport, emit, close } = createControlledTransport();
|
|
293
|
+
const onFinishA = vi.fn();
|
|
294
|
+
const onFinishB = vi.fn();
|
|
295
|
+
let onFinish = onFinishA;
|
|
296
|
+
const listeners = new Set<() => void>();
|
|
297
|
+
const handle = createAssistantClient({
|
|
298
|
+
getConfig: () =>
|
|
299
|
+
AuiConfig({
|
|
300
|
+
threads: AISDKThreads({ transport: () => transport, onFinish }),
|
|
301
|
+
}),
|
|
302
|
+
subscribe: (listener) => {
|
|
303
|
+
listeners.add(listener);
|
|
304
|
+
return () => listeners.delete(listener);
|
|
305
|
+
},
|
|
306
|
+
});
|
|
307
|
+
handle.subscribe(() => {});
|
|
308
|
+
const aui = handle.getClient();
|
|
309
|
+
|
|
310
|
+
flushTapSync(() => aui.composer.setText("stream me"));
|
|
311
|
+
flushTapSync(() => aui.composer.send());
|
|
312
|
+
await vi.waitFor(() => {
|
|
313
|
+
expect(
|
|
314
|
+
handle.getClient().thread.getState().messages.length,
|
|
315
|
+
).toBeGreaterThan(0);
|
|
316
|
+
});
|
|
317
|
+
emit(
|
|
318
|
+
{ type: "start" },
|
|
319
|
+
{ type: "text-start", id: "t1" },
|
|
320
|
+
{ type: "text-delta", id: "t1", delta: "partial" },
|
|
321
|
+
);
|
|
322
|
+
|
|
323
|
+
flushTapSync(() => aui.threads.switchToNewThread());
|
|
324
|
+
onFinish = onFinishB;
|
|
325
|
+
flushTapSync(() => listeners.forEach((listener) => listener()));
|
|
326
|
+
|
|
327
|
+
emit({ type: "text-end", id: "t1" }, { type: "finish" });
|
|
328
|
+
close();
|
|
329
|
+
await vi.waitFor(() => expect(onFinishB).toHaveBeenCalledTimes(1));
|
|
330
|
+
expect(onFinishA).not.toHaveBeenCalled();
|
|
275
331
|
|
|
276
332
|
handle.destroy();
|
|
277
333
|
});
|
|
278
334
|
|
|
335
|
+
it("forwards the latest callbacks to a cloud thread's chat", async () => {
|
|
336
|
+
const cloudThread = (id: string) => ({
|
|
337
|
+
id,
|
|
338
|
+
title: id,
|
|
339
|
+
is_archived: false,
|
|
340
|
+
last_message_at: null,
|
|
341
|
+
external_id: null,
|
|
342
|
+
metadata: null,
|
|
343
|
+
});
|
|
344
|
+
const cloud = {
|
|
345
|
+
threads: {
|
|
346
|
+
list: vi.fn(async () => ({ threads: [cloudThread("t1")] })),
|
|
347
|
+
create: vi.fn(),
|
|
348
|
+
update: vi.fn(),
|
|
349
|
+
delete: vi.fn(),
|
|
350
|
+
get: vi.fn(async (id: string) => cloudThread(id)),
|
|
351
|
+
messages: {
|
|
352
|
+
list: vi.fn(async () => ({ messages: [] })),
|
|
353
|
+
create: vi.fn(async () => ({ message_id: "remote-message-1" })),
|
|
354
|
+
update: vi.fn(),
|
|
355
|
+
},
|
|
356
|
+
},
|
|
357
|
+
runs: { stream: vi.fn(), report: vi.fn() },
|
|
358
|
+
telemetry: { enabled: false },
|
|
359
|
+
} as unknown as AssistantCloud;
|
|
360
|
+
const { transport, emit, close } = createControlledTransport();
|
|
361
|
+
const onFinishA = vi.fn();
|
|
362
|
+
const onFinishB = vi.fn();
|
|
363
|
+
let onFinish = onFinishA;
|
|
364
|
+
const listeners = new Set<() => void>();
|
|
365
|
+
const handle = createAssistantClient({
|
|
366
|
+
getConfig: () =>
|
|
367
|
+
AuiConfig({
|
|
368
|
+
threads: AISDKThreads({ cloud, threadId: "t1", transport, onFinish }),
|
|
369
|
+
}),
|
|
370
|
+
subscribe: (listener) => {
|
|
371
|
+
listeners.add(listener);
|
|
372
|
+
return () => listeners.delete(listener);
|
|
373
|
+
},
|
|
374
|
+
});
|
|
375
|
+
handle.subscribe(() => {});
|
|
376
|
+
try {
|
|
377
|
+
await handle.getClient().threads.getLoadThreadsPromise();
|
|
378
|
+
await vi.waitFor(() => {
|
|
379
|
+
expect(handle.getClient().threads.getState().mainThreadId).toBe("t1");
|
|
380
|
+
});
|
|
381
|
+
await vi.waitFor(() => {
|
|
382
|
+
expect(handle.getClient().thread.getState().isLoading).toBe(false);
|
|
383
|
+
});
|
|
384
|
+
|
|
385
|
+
onFinish = onFinishB;
|
|
386
|
+
flushTapSync(() => listeners.forEach((listener) => listener()));
|
|
387
|
+
|
|
388
|
+
flushTapSync(() => handle.getClient().composer.setText("hi"));
|
|
389
|
+
flushTapSync(() => handle.getClient().composer.send());
|
|
390
|
+
await vi.waitFor(() => {
|
|
391
|
+
expect(handle.getClient().thread.getState().isRunning).toBe(true);
|
|
392
|
+
});
|
|
393
|
+
emit(...textReply("done"));
|
|
394
|
+
close();
|
|
395
|
+
await vi.waitFor(() => expect(onFinishB).toHaveBeenCalledTimes(1));
|
|
396
|
+
expect(onFinishA).not.toHaveBeenCalled();
|
|
397
|
+
} finally {
|
|
398
|
+
handle.destroy();
|
|
399
|
+
}
|
|
400
|
+
});
|
|
401
|
+
|
|
279
402
|
it("posts each thread's own id as the chat id", async () => {
|
|
280
403
|
const bodies: unknown[] = [];
|
|
281
404
|
const fetchStub = vi.fn(async (_url: unknown, init?: RequestInit) => {
|
|
@@ -332,7 +455,9 @@ describe("AISDKThreads", () => {
|
|
|
332
455
|
const list = vi.fn(async () => ({
|
|
333
456
|
threads: [cloudThread("cloud-1"), cloudThread("cloud-2")],
|
|
334
457
|
}));
|
|
335
|
-
const create = vi.fn(async () => ({
|
|
458
|
+
const create = vi.fn<AssistantCloud["threads"]["create"]>(async () => ({
|
|
459
|
+
thread_id: "cloud-created",
|
|
460
|
+
}));
|
|
336
461
|
const deleteThread = vi.fn(async () => {});
|
|
337
462
|
const cloud = {
|
|
338
463
|
threads: {
|
|
@@ -398,7 +523,9 @@ describe("AISDKThreads", () => {
|
|
|
398
523
|
external_id: null,
|
|
399
524
|
metadata: null,
|
|
400
525
|
});
|
|
401
|
-
const create = vi.fn
|
|
526
|
+
const create = vi.fn<AssistantCloud["threads"]["messages"]["create"]>(
|
|
527
|
+
async () => ({ message_id: "remote-message-1" }),
|
|
528
|
+
);
|
|
402
529
|
const cloud = {
|
|
403
530
|
threads: {
|
|
404
531
|
list: vi.fn(async () => ({
|
|
@@ -479,9 +606,9 @@ describe("AISDKThreads", () => {
|
|
|
479
606
|
external_id: null,
|
|
480
607
|
metadata: null,
|
|
481
608
|
});
|
|
482
|
-
const createMessage = vi.fn
|
|
483
|
-
|
|
484
|
-
}));
|
|
609
|
+
const createMessage = vi.fn<
|
|
610
|
+
AssistantCloud["threads"]["messages"]["create"]
|
|
611
|
+
>(async () => ({ message_id: "remote-message-1" }));
|
|
485
612
|
const cloud = {
|
|
486
613
|
threads: {
|
|
487
614
|
list: vi.fn(async () => ({ threads: [cloudThread("t1")] })),
|
|
@@ -2,7 +2,7 @@
|
|
|
2
2
|
|
|
3
3
|
import { resource, useResource, withKey } from "@assistant-ui/tap";
|
|
4
4
|
import { useEffect, useMemo, useState } from "react";
|
|
5
|
-
import { Chat,
|
|
5
|
+
import type { Chat, UIMessage } from "@ai-sdk/react";
|
|
6
6
|
import type { ChatTransport } from "ai";
|
|
7
7
|
import type { AssistantCloud } from "assistant-cloud";
|
|
8
8
|
import {
|
|
@@ -20,12 +20,14 @@ import {
|
|
|
20
20
|
import { useAui } from "@assistant-ui/store";
|
|
21
21
|
import { AssistantChatTransport } from "../transport/AssistantChatTransport";
|
|
22
22
|
import {
|
|
23
|
+
createChat,
|
|
23
24
|
splitChatThreadOptions,
|
|
24
25
|
useChatThread,
|
|
25
26
|
type ChatThreadOptions,
|
|
26
27
|
} from "./useChatThread";
|
|
27
28
|
import { MessageRepository } from "@assistant-ui/core/internal";
|
|
28
29
|
import { useResourceCleanup } from "./useResourceCleanup";
|
|
30
|
+
import { AI_SDK_SDK } from "./sdkIdentity";
|
|
29
31
|
|
|
30
32
|
export type AISDKThreadsOptions<UI_MESSAGE extends UIMessage = UIMessage> =
|
|
31
33
|
Omit<ChatThreadOptions<UI_MESSAGE>, "id" | "transport" | "messages"> & {
|
|
@@ -62,16 +64,22 @@ type AISDKThreadChatOptions<UI_MESSAGE extends UIMessage = UIMessage> = Omit<
|
|
|
62
64
|
"cloud" | "threadId" | "onThreadIdChange"
|
|
63
65
|
>;
|
|
64
66
|
|
|
67
|
+
type ChatOptionsRef<UI_MESSAGE extends UIMessage> = {
|
|
68
|
+
current: AISDKThreadChatOptions<UI_MESSAGE> | undefined;
|
|
69
|
+
};
|
|
70
|
+
|
|
65
71
|
type ChatEntry<UI_MESSAGE extends UIMessage> = {
|
|
66
72
|
chat: Chat<UI_MESSAGE>;
|
|
67
73
|
transport: ChatTransport<UI_MESSAGE>;
|
|
68
74
|
repository: MessageRepository;
|
|
75
|
+
optionsRef: ChatOptionsRef<UI_MESSAGE>;
|
|
69
76
|
};
|
|
70
77
|
|
|
71
78
|
const createChatEntry = <UI_MESSAGE extends UIMessage>(
|
|
72
79
|
threadId: string,
|
|
73
80
|
options: AISDKThreadChatOptions<UI_MESSAGE> | undefined,
|
|
74
81
|
): ChatEntry<UI_MESSAGE> => {
|
|
82
|
+
const optionsRef: ChatOptionsRef<UI_MESSAGE> = { current: options };
|
|
75
83
|
const { chatInit } = splitChatThreadOptions(
|
|
76
84
|
options as ChatThreadOptions<UI_MESSAGE> | undefined,
|
|
77
85
|
);
|
|
@@ -84,9 +92,10 @@ const createChatEntry = <UI_MESSAGE extends UIMessage>(
|
|
|
84
92
|
? options.transport.__internal_clone()
|
|
85
93
|
: options.transport;
|
|
86
94
|
return {
|
|
87
|
-
chat:
|
|
95
|
+
chat: createChat({ ...chatInit, id: threadId, transport }, optionsRef),
|
|
88
96
|
transport,
|
|
89
97
|
repository: new MessageRepository(),
|
|
98
|
+
optionsRef,
|
|
90
99
|
};
|
|
91
100
|
};
|
|
92
101
|
|
|
@@ -116,9 +125,13 @@ const useAISDKChatThread = <UI_MESSAGE extends UIMessage = UIMessage>({
|
|
|
116
125
|
const [owned] = useState(() =>
|
|
117
126
|
cloud ? createChatEntry(threadId, options) : undefined,
|
|
118
127
|
);
|
|
119
|
-
const { chat, transport, repository } =
|
|
128
|
+
const { chat, transport, repository, optionsRef } =
|
|
120
129
|
owned ?? getOrCreateChatEntry(threadId, options, chats);
|
|
121
130
|
|
|
131
|
+
useEffect(() => {
|
|
132
|
+
if (cloud) optionsRef.current = options;
|
|
133
|
+
});
|
|
134
|
+
|
|
122
135
|
useEffect(() => {
|
|
123
136
|
if (!cloud) return undefined;
|
|
124
137
|
return () => {
|
|
@@ -173,13 +186,19 @@ const useAISDKThreads = <UI_MESSAGE extends UIMessage = UIMessage>(
|
|
|
173
186
|
const [chats] = useState(() => new Map<string, ChatEntry<UI_MESSAGE>>());
|
|
174
187
|
const bindCloud = cloud !== undefined;
|
|
175
188
|
|
|
189
|
+
useEffect(() => {
|
|
190
|
+
for (const { optionsRef } of chats.values()) {
|
|
191
|
+
optionsRef.current = threadOptions;
|
|
192
|
+
}
|
|
193
|
+
});
|
|
194
|
+
|
|
176
195
|
useResourceCleanup(true, () => {
|
|
177
196
|
for (const { chat } of chats.values()) {
|
|
178
197
|
void chat.stop().catch(() => {});
|
|
179
198
|
}
|
|
180
199
|
});
|
|
181
200
|
|
|
182
|
-
const cloudAdapter = useCloudThreadListAdapter({ cloud });
|
|
201
|
+
const cloudAdapter = useCloudThreadListAdapter({ cloud, sdk: AI_SDK_SDK });
|
|
183
202
|
const thread = (id: string) => {
|
|
184
203
|
const element = AISDKChatThread({
|
|
185
204
|
threadId: id,
|