@assistant-ui/react-langchain 0.0.32 → 0.0.33

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 (55) hide show
  1. package/dist/attachSubagentTranscripts.d.ts +13 -14
  2. package/dist/attachSubagentTranscripts.d.ts.map +1 -1
  3. package/dist/convertMessages.d.ts +6 -9
  4. package/dist/convertMessages.d.ts.map +1 -1
  5. package/dist/convertMessages.js +53 -15
  6. package/dist/convertMessages.js.map +1 -1
  7. package/dist/converter.d.ts +128 -128
  8. package/dist/converter.d.ts.map +1 -1
  9. package/dist/converter.js +2 -1
  10. package/dist/converter.js.map +1 -1
  11. package/dist/findForkCheckpointInHistory.d.ts +15 -16
  12. package/dist/findForkCheckpointInHistory.d.ts.map +1 -1
  13. package/dist/hooks.d.ts +5 -7
  14. package/dist/hooks.d.ts.map +1 -1
  15. package/dist/index.d.ts +7 -8
  16. package/dist/index.d.ts.map +1 -0
  17. package/dist/resolveForkCheckpoint.d.ts +4 -5
  18. package/dist/resolveForkCheckpoint.d.ts.map +1 -1
  19. package/dist/runtimeExtras.d.ts +1 -3
  20. package/dist/runtimeExtras.d.ts.map +1 -1
  21. package/dist/sdkIdentity.d.ts +1 -3
  22. package/dist/sdkIdentity.d.ts.map +1 -1
  23. package/dist/sdkIdentity.js +1 -1
  24. package/dist/streamingTiming.d.ts +2 -4
  25. package/dist/streamingTiming.d.ts.map +1 -1
  26. package/dist/subagentMessagesProjection.d.ts +17 -0
  27. package/dist/subagentMessagesProjection.d.ts.map +1 -0
  28. package/dist/subagentMessagesProjection.js +40 -0
  29. package/dist/subagentMessagesProjection.js.map +1 -0
  30. package/dist/types.d.ts +100 -103
  31. package/dist/types.d.ts.map +1 -1
  32. package/dist/uiMessages.d.ts +18 -6
  33. package/dist/uiMessages.d.ts.map +1 -1
  34. package/dist/uiMessages.js +33 -1
  35. package/dist/uiMessages.js.map +1 -1
  36. package/dist/useStreamRuntime.d.ts +7 -10
  37. package/dist/useStreamRuntime.d.ts.map +1 -1
  38. package/dist/useStreamRuntime.js +99 -25
  39. package/dist/useStreamRuntime.js.map +1 -1
  40. package/dist/useSubagentTranscripts.d.ts +3 -5
  41. package/dist/useSubagentTranscripts.d.ts.map +1 -1
  42. package/dist/useSubagentTranscripts.js +3 -2
  43. package/dist/useSubagentTranscripts.js.map +1 -1
  44. package/package.json +9 -5
  45. package/src/convertMessages.test.ts +69 -0
  46. package/src/convertMessages.ts +46 -8
  47. package/src/converter.ts +6 -0
  48. package/src/subagentMessagesProjection.test.ts +175 -0
  49. package/src/subagentMessagesProjection.ts +50 -0
  50. package/src/uiMessages.test.ts +106 -0
  51. package/src/uiMessages.ts +43 -0
  52. package/src/useStreamRuntime.test.tsx +100 -0
  53. package/src/useStreamRuntime.ts +193 -44
  54. package/src/useStreamRuntime.voice.test.tsx +712 -0
  55. package/src/useSubagentTranscripts.ts +3 -6
@@ -1072,6 +1072,106 @@ describe("useStreamRuntime subagent transcripts", () => {
1072
1072
  expect(nestedTranscript()).toBe(rendered);
1073
1073
  });
1074
1074
 
1075
+ it("keeps messages and transcripts across equal copies of the UI state", async () => {
1076
+ const stream = createMockStream([
1077
+ message("human-1", "human", "delegate"),
1078
+ {
1079
+ id: "root-ai",
1080
+ _getType: () => "ai",
1081
+ content: "",
1082
+ tool_calls: [{ id: "task-one", name: "task", args: {} }],
1083
+ },
1084
+ ]);
1085
+ const transcript = [message("nested-ai", "ai", "nested answer")];
1086
+ stream.subagents = new Map([
1087
+ [
1088
+ "task-one",
1089
+ {
1090
+ id: "task-one",
1091
+ namespace: ["tools:task-one"],
1092
+ status: "running",
1093
+ parentId: null,
1094
+ depth: 1,
1095
+ startedAt: new Date(1_000),
1096
+ completedAt: null,
1097
+ },
1098
+ ],
1099
+ ]);
1100
+ stream[streamController]!.registry.acquire.mockReturnValue({
1101
+ store: { getSnapshot: () => transcript, subscribe: () => () => {} },
1102
+ release: vi.fn(),
1103
+ });
1104
+ const uiState = (points: number[]) => [
1105
+ {
1106
+ type: "ui",
1107
+ id: "ui-root",
1108
+ name: "chart",
1109
+ props: { points },
1110
+ metadata: { message_id: "root-ai" },
1111
+ },
1112
+ {
1113
+ type: "ui",
1114
+ id: "ui-nested",
1115
+ name: "chart",
1116
+ props: { points },
1117
+ metadata: { message_id: "nested-ai" },
1118
+ },
1119
+ ];
1120
+ stream.values = { ui: uiState([1, 2]) };
1121
+ const { auiResult, rerender } = renderAui(stream);
1122
+ const nestedTranscript = () => {
1123
+ const { messages } = auiResult.current.thread.getState();
1124
+ for (const threadMessage of messages) {
1125
+ for (const part of threadMessage.content) {
1126
+ if (part.type === "tool-call" && part.toolCallId === "task-one")
1127
+ return part.messages;
1128
+ }
1129
+ }
1130
+ return undefined;
1131
+ };
1132
+
1133
+ await waitFor(() =>
1134
+ expect(nestedTranscript()?.[0]?.content).toMatchObject([
1135
+ { type: "text", text: "nested answer" },
1136
+ { type: "data", name: "chart", data: { points: [1, 2] } },
1137
+ ]),
1138
+ );
1139
+ const [human, ai] = auiResult.current.thread.getState().messages;
1140
+ expect(ai?.content).toMatchObject([
1141
+ { type: "tool-call", toolCallId: "task-one" },
1142
+ { type: "data", name: "chart", data: { points: [1, 2] } },
1143
+ ]);
1144
+ const rendered = nestedTranscript();
1145
+
1146
+ for (let i = 0; i < 3; i++) {
1147
+ stream.values = { ui: uiState([1, 2]) };
1148
+ await act(async () => {
1149
+ rerender();
1150
+ });
1151
+ }
1152
+
1153
+ const messages = auiResult.current.thread.getState().messages;
1154
+ expect(messages[0]).toBe(human);
1155
+ expect(messages[1]).toBe(ai);
1156
+ expect(nestedTranscript()).toBe(rendered);
1157
+
1158
+ stream.values = { ui: uiState([1, 2, 3]) };
1159
+ await act(async () => {
1160
+ rerender();
1161
+ });
1162
+
1163
+ expect(
1164
+ auiResult.current.thread.getState().messages[1]?.content,
1165
+ ).toMatchObject([
1166
+ { type: "tool-call", toolCallId: "task-one" },
1167
+ { type: "data", name: "chart", data: { points: [1, 2, 3] } },
1168
+ ]);
1169
+ expect(nestedTranscript()?.[0]?.content).toMatchObject([
1170
+ { type: "text", text: "nested answer" },
1171
+ { type: "data", name: "chart", data: { points: [1, 2, 3] } },
1172
+ ]);
1173
+ });
1174
+
1075
1175
  it("keeps messages and transcripts when custom events carry no UI update", async () => {
1076
1176
  const stream = createMockStream([
1077
1177
  message("human-1", "human", "delegate"),
@@ -1,4 +1,4 @@
1
- /// <reference types="@assistant-ui/core/store" />
1
+ /// <reference types="@assistant-ui/core/store" preserve="true" />
2
2
  "use client";
3
3
 
4
4
  import {
@@ -19,6 +19,7 @@ import type { ThreadMessage } from "@assistant-ui/core";
19
19
  import {
20
20
  createCloudThreadListAdapterCreateFallback,
21
21
  createToolCallCancellationStub,
22
+ getThreadMessageText,
22
23
  scanPendingToolCalls,
23
24
  } from "@assistant-ui/core/internal";
24
25
  import {
@@ -35,7 +36,7 @@ import type {
35
36
  UIMessage,
36
37
  UseStreamRuntimeOptions,
37
38
  } from "./types";
38
- import { groupUIMessagesByParent } from "./converter";
39
+ import { getMessageModality, groupUIMessagesByParent } from "./converter";
39
40
  export { groupUIMessagesByParent } from "./converter";
40
41
  import {
41
42
  convertLangChainBaseMessage,
@@ -49,8 +50,10 @@ import {
49
50
  import { useSubagentTranscripts } from "./useSubagentTranscripts";
50
51
  import {
51
52
  createUIFoldMemo,
53
+ createUISnapshotMemo,
52
54
  foldUIUpdates,
53
55
  mergeUIMessages,
56
+ reconcileUISnapshot,
54
57
  UI_CUSTOM_CHANNELS,
55
58
  } from "./uiMessages";
56
59
  import { langChainExtras } from "./runtimeExtras";
@@ -94,6 +97,15 @@ const toStagedHumanMessage = (
94
97
  content: getMessageContent(msg),
95
98
  });
96
99
 
100
+ const toStagedMessageInput = (message: LangChainBaseMessage) => ({
101
+ id: message.id,
102
+ type: getMessageType(message) === "ai" ? ("ai" as const) : ("human" as const),
103
+ content: message.content,
104
+ ...(message.additional_kwargs && {
105
+ additional_kwargs: message.additional_kwargs,
106
+ }),
107
+ });
108
+
97
109
  const humanContentText = (content: LangChainBaseMessage["content"]) => {
98
110
  if (typeof content === "string") return content;
99
111
  if (!Array.isArray(content)) return "";
@@ -166,7 +178,11 @@ const useStreamThreadRuntime = (
166
178
  );
167
179
  const effectiveIsRunning = stream.isLoading || hasExecutingTools;
168
180
 
169
- const uiStateValue = stream.values[uiStateKey];
181
+ const [uiSnapshotMemo] = useState(createUISnapshotMemo);
182
+ const uiStateValue = reconcileUISnapshot(
183
+ stream.values[uiStateKey],
184
+ uiSnapshotMemo,
185
+ );
170
186
 
171
187
  const customEvents = useChannel(stream, UI_CUSTOM_CHANNELS);
172
188
  const [uiFoldMemo] = useState(createUIFoldMemo);
@@ -305,6 +321,7 @@ const useStreamThreadRuntime = (
305
321
  runConfig: AppendMessage["runConfig"];
306
322
  reconcileOnEcho: boolean;
307
323
  baseMessageCount: number;
324
+ transcriptStatus?: "unsent" | "sent";
308
325
  }
309
326
  >(),
310
327
  );
@@ -360,19 +377,121 @@ const useStreamThreadRuntime = (
360
377
  }, [stream.messages]);
361
378
 
362
379
  const getStagedRun = (parentId: string | null) => {
363
- if (!parentId || !stagedMessagesRef.current.has(parentId)) return null;
380
+ const parent = parentId
381
+ ? stagedMessagesRef.current.get(parentId)
382
+ : undefined;
383
+ if (!parent || parent.transcriptStatus === "sent") return null;
364
384
 
365
385
  const staged: LangChainBaseMessage[] = [];
366
386
  for (const message of visibleMessagesRef.current) {
367
- if (message.id && stagedMessagesRef.current.has(message.id)) {
368
- staged.push(stagedMessagesRef.current.get(message.id)!.message);
387
+ const entry = message.id
388
+ ? stagedMessagesRef.current.get(message.id)
389
+ : undefined;
390
+ if (entry && entry.transcriptStatus !== "sent") {
391
+ staged.push(entry.message);
369
392
  }
370
393
  if (message.id === parentId) break;
371
394
  }
372
395
 
373
396
  return {
374
397
  messages: staged,
375
- runConfig: stagedMessagesRef.current.get(parentId)!.runConfig,
398
+ runConfig: parent.runConfig,
399
+ };
400
+ };
401
+
402
+ const appendVoiceTranscript = (message: ThreadMessage) => {
403
+ const transcript = {
404
+ id: message.id,
405
+ _getType: () => (message.role === "assistant" ? "ai" : "human"),
406
+ content: getThreadMessageText(message),
407
+ ...(message.metadata.modality && {
408
+ additional_kwargs: { modality: message.metadata.modality },
409
+ }),
410
+ };
411
+ stagedMessagesRef.current.set(transcript.id, {
412
+ message: transcript,
413
+ runConfig: undefined,
414
+ reconcileOnEcho: false,
415
+ baseMessageCount: streamRef.current.messages.length,
416
+ transcriptStatus: "unsent",
417
+ });
418
+ const nextMessages = [...visibleMessagesRef.current, transcript];
419
+ visibleMessagesRef.current = nextMessages;
420
+ setStagedMessages(nextMessages);
421
+ };
422
+
423
+ const getUnsentTranscripts = () =>
424
+ visibleMessagesRef.current.filter(
425
+ (message) =>
426
+ message.id !== undefined &&
427
+ stagedMessagesRef.current.get(message.id)?.transcriptStatus ===
428
+ "unsent",
429
+ );
430
+
431
+ const setTranscriptStatus = (
432
+ messages: readonly LangChainBaseMessage[],
433
+ status: "unsent" | "sent",
434
+ ) => {
435
+ for (const message of messages) {
436
+ const staged = message.id
437
+ ? stagedMessagesRef.current.get(message.id)
438
+ : undefined;
439
+ if (staged?.transcriptStatus) staged.transcriptStatus = status;
440
+ }
441
+ };
442
+
443
+ // Reserving before the submit keeps an overlapping submit from carrying the
444
+ // same transcript; a failed submit hands it back to the next run.
445
+ const submitCarryingTranscripts = async (
446
+ transcripts: readonly LangChainBaseMessage[],
447
+ submit: () => Promise<void>,
448
+ ) => {
449
+ setTranscriptStatus(transcripts, "sent");
450
+ try {
451
+ await submit();
452
+ } catch (error) {
453
+ setTranscriptStatus(transcripts, "unsent");
454
+ throw error;
455
+ }
456
+ };
457
+
458
+ const dropTranscripts = (messages: readonly LangChainBaseMessage[]) => {
459
+ for (const message of messages) {
460
+ if (
461
+ message.id &&
462
+ stagedMessagesRef.current.get(message.id)?.transcriptStatus
463
+ )
464
+ removeStagedMessage(message.id);
465
+ }
466
+ };
467
+
468
+ const isTranscriptMessage = (message: LangChainBaseMessage) => {
469
+ const staged = message.id
470
+ ? stagedMessagesRef.current.get(message.id)
471
+ : undefined;
472
+ if (staged?.transcriptStatus === "unsent") return true;
473
+ const type = getMessageType(message);
474
+ return (
475
+ (type === "human" || type === "ai") &&
476
+ getMessageModality(message.additional_kwargs) !== undefined
477
+ );
478
+ };
479
+
480
+ // A transcript reaches the graph in the same input as the message after it, so
481
+ // no checkpoint ends at one; a fork starts before the trailing transcripts.
482
+ const planForkTranscripts = (parentId: string | null) => {
483
+ const visible = visibleMessagesRef.current;
484
+ const parentIndex =
485
+ parentId == null ? -1 : visible.findIndex((m) => m.id === parentId);
486
+ if (parentId != null && parentIndex === -1)
487
+ return { forkParentId: parentId, transcripts: [], truncated: [] };
488
+
489
+ let start = parentIndex + 1;
490
+ while (start > 0 && isTranscriptMessage(visible[start - 1]!)) start--;
491
+ return {
492
+ forkParentId: start > 0 ? (visible[start - 1]!.id ?? null) : null,
493
+ transcripts: visible.slice(start, parentIndex + 1),
494
+ truncated: visible.slice(parentIndex + 1),
376
495
  };
377
496
  };
378
497
 
@@ -462,27 +581,32 @@ const useStreamThreadRuntime = (
462
581
  // longer holds appends on that barrier.
463
582
  try {
464
583
  const { externalId } = await aui.threadListItem.initialize();
465
- await streamRef.current.submit(
466
- {
467
- [messagesKey]: [
468
- ...cancellations,
469
- {
470
- id: stagedMessageId,
471
- type: "human",
472
- content,
473
- },
474
- ],
475
- },
476
- {
477
- ...runConfigToSubmitOptions(msg.runConfig),
478
- ...(externalId != null ? { threadId: externalId } : {}),
479
- },
584
+ const transcripts = getUnsentTranscripts();
585
+ await submitCarryingTranscripts(transcripts, () =>
586
+ streamRef.current.submit(
587
+ {
588
+ [messagesKey]: [
589
+ ...cancellations,
590
+ ...transcripts.map(toStagedMessageInput),
591
+ {
592
+ id: stagedMessageId,
593
+ type: "human",
594
+ content,
595
+ },
596
+ ],
597
+ },
598
+ {
599
+ ...runConfigToSubmitOptions(msg.runConfig),
600
+ ...(externalId != null ? { threadId: externalId } : {}),
601
+ },
602
+ ),
480
603
  );
481
604
  } catch (error) {
482
605
  removeStagedMessage(stagedMessageId);
483
606
  throw error;
484
607
  }
485
608
  },
609
+ onVoiceTranscript: appendVoiceTranscript,
486
610
  onAddToolResult: async ({
487
611
  messageId,
488
612
  toolCallId,
@@ -513,9 +637,18 @@ const useStreamThreadRuntime = (
513
637
  onReload: async (parentId, config) => {
514
638
  const stagedRun = getStagedRun(parentId);
515
639
  if (stagedRun) {
640
+ if (
641
+ config.sourceId &&
642
+ stagedMessagesRef.current.get(config.sourceId)?.transcriptStatus
643
+ )
644
+ removeStagedMessage(config.sourceId);
516
645
  const promotedIds = new Set<string>();
517
646
  for (const message of stagedRun.messages) {
518
- if (!message.id) continue;
647
+ if (
648
+ !message.id ||
649
+ stagedMessagesRef.current.get(message.id)?.transcriptStatus
650
+ )
651
+ continue;
519
652
  promotedIds.add(message.id);
520
653
  stagedMessagesRef.current.delete(message.id);
521
654
  }
@@ -531,15 +664,13 @@ const useStreamThreadRuntime = (
531
664
  }
532
665
  const runConfig = config.runConfig ?? stagedRun.runConfig;
533
666
  setActiveRunConfig(runConfig);
534
- await stream.submit(
535
- {
536
- [messagesKey]: stagedRun.messages.map((message) => ({
537
- id: message.id,
538
- type: "human",
539
- content: message.content,
540
- })),
541
- },
542
- runConfigToSubmitOptions(runConfig),
667
+ await submitCarryingTranscripts(stagedRun.messages, () =>
668
+ stream.submit(
669
+ {
670
+ [messagesKey]: stagedRun.messages.map(toStagedMessageInput),
671
+ },
672
+ runConfigToSubmitOptions(runConfig),
673
+ ),
543
674
  );
544
675
  return;
545
676
  }
@@ -547,21 +678,30 @@ const useStreamThreadRuntime = (
547
678
  const threadId = externalId;
548
679
  if (!threadId || parentId == null) return;
549
680
  const s = streamRef.current;
681
+ const fork = planForkTranscripts(parentId);
550
682
  const checkpointId = await resolveForkCheckpoint(
551
683
  s.client,
552
684
  threadId,
553
685
  s.messages as readonly LangChainBaseMessage[],
554
- parentId,
686
+ fork.forkParentId,
555
687
  config.sourceId,
556
688
  s[STREAM_CONTROLLER]?.messageMetadataStore?.getSnapshot?.(),
557
689
  messagesKey,
558
690
  );
559
691
  if (!checkpointId) return;
692
+ dropTranscripts(fork.truncated);
560
693
  setActiveRunConfig(config.runConfig);
561
- await s.submit(null, {
562
- forkFrom: checkpointId,
563
- ...runConfigToSubmitOptions(config.runConfig),
564
- });
694
+ await submitCarryingTranscripts(fork.transcripts, () =>
695
+ s.submit(
696
+ fork.transcripts.length > 0
697
+ ? { [messagesKey]: fork.transcripts.map(toStagedMessageInput) }
698
+ : null,
699
+ {
700
+ forkFrom: checkpointId,
701
+ ...runConfigToSubmitOptions(config.runConfig),
702
+ },
703
+ ),
704
+ );
565
705
  },
566
706
  onEdit: async (message) => {
567
707
  if (!(message.startRun ?? message.role === "user")) {
@@ -586,24 +726,33 @@ const useStreamThreadRuntime = (
586
726
  const threadId = externalId;
587
727
  if (!threadId) return;
588
728
  const s = streamRef.current;
729
+ const fork = planForkTranscripts(message.parentId);
589
730
  const checkpointId = await resolveForkCheckpoint(
590
731
  s.client,
591
732
  threadId,
592
733
  s.messages as readonly LangChainBaseMessage[],
593
- message.parentId,
734
+ fork.forkParentId,
594
735
  message.sourceId,
595
736
  s[STREAM_CONTROLLER]?.messageMetadataStore?.getSnapshot?.(),
596
737
  messagesKey,
597
738
  );
598
739
  if (!checkpointId) return;
740
+ dropTranscripts(fork.truncated);
599
741
  const content = getMessageContent(message);
600
742
  setActiveRunConfig(message.runConfig);
601
- await s.submit(
602
- { [messagesKey]: [{ type: "human", content }] },
603
- {
604
- forkFrom: checkpointId,
605
- ...runConfigToSubmitOptions(message.runConfig),
606
- },
743
+ await submitCarryingTranscripts(fork.transcripts, () =>
744
+ s.submit(
745
+ {
746
+ [messagesKey]: [
747
+ ...fork.transcripts.map(toStagedMessageInput),
748
+ { type: "human", content },
749
+ ],
750
+ },
751
+ {
752
+ forkFrom: checkpointId,
753
+ ...runConfigToSubmitOptions(message.runConfig),
754
+ },
755
+ ),
607
756
  );
608
757
  },
609
758
  onCancel: