@assistant-ui/react-langchain 0.0.26 → 0.0.28

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.
@@ -5,6 +5,7 @@ import { describe, expect, it, vi } from "vitest";
5
5
  import { AssistantRuntimeProvider } from "@assistant-ui/core/react";
6
6
  import type {
7
7
  AssistantRuntime,
8
+ AppendMessage,
8
9
  RemoteThreadListAdapter,
9
10
  } from "@assistant-ui/core";
10
11
  import { useAui } from "@assistant-ui/store";
@@ -156,10 +157,12 @@ const makeThreadListAdapter = (): RemoteThreadListAdapter => ({
156
157
 
157
158
  const deferred = <T,>() => {
158
159
  let resolve!: (value: T) => void;
159
- const promise = new Promise<T>((res) => {
160
+ let reject!: (reason?: unknown) => void;
161
+ const promise = new Promise<T>((res, rej) => {
160
162
  resolve = res;
163
+ reject = rej;
161
164
  });
162
- return { promise, resolve };
165
+ return { promise, resolve, reject };
163
166
  };
164
167
 
165
168
  describe("useStreamRuntime thread options", () => {
@@ -202,7 +205,7 @@ describe("useStreamRuntime thread options", () => {
202
205
  view.unmount();
203
206
  });
204
207
 
205
- it("submits with the initialized thread id when initialization resolves before rerender", async () => {
208
+ it("renders before initialization and submits with the initialized thread id", async () => {
206
209
  const stream = createMockStream();
207
210
  mockUseStream.mockReturnValue(stream);
208
211
  const initialization = deferred<{
@@ -213,18 +216,29 @@ describe("useStreamRuntime thread options", () => {
213
216
  threadListAdapter.list = vi.fn(async () => ({ threads: [] }));
214
217
  threadListAdapter.initialize = vi.fn(() => initialization.promise);
215
218
 
216
- const capture: { runtime: AssistantRuntime | null } = { runtime: null };
219
+ const capture: {
220
+ runtime: AssistantRuntime | null;
221
+ aui?: ReturnType<typeof useAui>;
222
+ } = { runtime: null };
223
+ const Capture = () => {
224
+ capture.aui = useAui();
225
+ return null;
226
+ };
217
227
  const TestRuntime = () => {
218
228
  const runtime = useStreamRuntime({
219
229
  apiUrl: "/api",
220
230
  unstable_threadListAdapter: threadListAdapter,
221
231
  } as never);
222
232
  capture.runtime = runtime;
223
- return <AssistantRuntimeProvider runtime={runtime} />;
233
+ return (
234
+ <AssistantRuntimeProvider runtime={runtime}>
235
+ <Capture />
236
+ </AssistantRuntimeProvider>
237
+ );
224
238
  };
225
239
 
226
240
  const view = render(<TestRuntime />);
227
- await waitFor(() => expect(capture.runtime).not.toBeNull());
241
+ await waitFor(() => expect(capture.aui).toBeDefined());
228
242
 
229
243
  await act(async () => {
230
244
  capture.runtime!.thread.append({
@@ -235,6 +249,7 @@ describe("useStreamRuntime thread options", () => {
235
249
  });
236
250
 
237
251
  expect(stream.submit).not.toHaveBeenCalled();
252
+ expect(getText(capture.aui!)).toEqual(["hello"]);
238
253
 
239
254
  await act(async () => {
240
255
  initialization.resolve({ remoteId: "thread-b", externalId: "thread-b" });
@@ -242,10 +257,25 @@ describe("useStreamRuntime thread options", () => {
242
257
 
243
258
  await waitFor(() =>
244
259
  expect(stream.submit).toHaveBeenCalledWith(
245
- { messages: [{ type: "human", content: "hello" }] },
260
+ {
261
+ messages: [
262
+ expect.objectContaining({
263
+ id: expect.any(String),
264
+ type: "human",
265
+ content: "hello",
266
+ }),
267
+ ],
268
+ },
246
269
  { threadId: "thread-b" },
247
270
  ),
248
271
  );
272
+
273
+ stream.messages = [message("echo-hello", "human", "hello")];
274
+ view.rerender(<TestRuntime />);
275
+ await waitFor(() => {
276
+ expect(getText(capture.aui!)).toEqual(["hello"]);
277
+ expect(capture.aui!.thread.getState().messages).toHaveLength(1);
278
+ });
249
279
  view.unmount();
250
280
  });
251
281
 
@@ -280,6 +310,100 @@ describe("useStreamRuntime thread options", () => {
280
310
  }
281
311
  view.unmount();
282
312
  });
313
+
314
+ it.each(["initialization", "submit"] as const)(
315
+ "removes the staged message when %s fails",
316
+ async (failurePoint) => {
317
+ const stream = createMockStream();
318
+ mockUseStream.mockReturnValue(stream);
319
+ const initialization = deferred<{
320
+ remoteId: string;
321
+ externalId: string;
322
+ }>();
323
+ const threadListAdapter = makeThreadListAdapter();
324
+ threadListAdapter.list = vi.fn(async () => ({ threads: [] }));
325
+ threadListAdapter.initialize = vi.fn(() => initialization.promise);
326
+
327
+ const capture: {
328
+ runtime: AssistantRuntime | null;
329
+ aui?: ReturnType<typeof useAui>;
330
+ } = { runtime: null };
331
+ const Capture = () => {
332
+ capture.aui = useAui();
333
+ return null;
334
+ };
335
+ const TestRuntime = () => {
336
+ const runtime = useStreamRuntime({
337
+ apiUrl: "/api",
338
+ unstable_threadListAdapter: threadListAdapter,
339
+ } as never);
340
+ capture.runtime = runtime;
341
+ return (
342
+ <AssistantRuntimeProvider runtime={runtime}>
343
+ <Capture />
344
+ </AssistantRuntimeProvider>
345
+ );
346
+ };
347
+
348
+ const view = render(<TestRuntime />);
349
+ await waitFor(() => expect(capture.aui).toBeDefined());
350
+
351
+ if (failurePoint === "submit") {
352
+ stream.submit.mockRejectedValueOnce(new Error("submit failed"));
353
+ }
354
+
355
+ const core = (
356
+ capture.runtime!.thread as unknown as {
357
+ __internal_threadBinding: {
358
+ getState(): { append(message: AppendMessage): Promise<void> };
359
+ };
360
+ }
361
+ ).__internal_threadBinding.getState();
362
+ let appendPromise!: Promise<void>;
363
+ await act(async () => {
364
+ appendPromise = core.append({
365
+ role: "user",
366
+ content: [{ type: "text", text: "failed" }],
367
+ parentId: null,
368
+ sourceId: null,
369
+ runConfig: undefined,
370
+ attachments: [],
371
+ metadata: { custom: {} },
372
+ createdAt: new Date(0),
373
+ });
374
+ await Promise.resolve();
375
+ });
376
+ const appendResult = appendPromise.then(
377
+ () => undefined,
378
+ (error: unknown) => error,
379
+ );
380
+ expect(getText(capture.aui!)).toEqual(["failed"]);
381
+
382
+ await act(async () => {
383
+ if (failurePoint === "initialization") {
384
+ initialization.reject(new Error("initialize failed"));
385
+ } else {
386
+ initialization.resolve({
387
+ remoteId: "thread-failed",
388
+ externalId: "thread-failed",
389
+ });
390
+ }
391
+ });
392
+ await expect(appendResult).resolves.toMatchObject({
393
+ message:
394
+ failurePoint === "initialization"
395
+ ? "initialize failed"
396
+ : "submit failed",
397
+ });
398
+ await waitFor(() => expect(getText(capture.aui!)).toEqual([]));
399
+ if (failurePoint === "initialization") {
400
+ expect(stream.submit).not.toHaveBeenCalled();
401
+ } else {
402
+ expect(stream.submit).toHaveBeenCalledTimes(1);
403
+ }
404
+ view.unmount();
405
+ },
406
+ );
283
407
  });
284
408
 
285
409
  describe("useStreamRuntime run configuration", () => {
@@ -331,7 +455,9 @@ describe("useStreamRuntime run configuration", () => {
331
455
  expect(stream.submit).toHaveBeenNthCalledWith(
332
456
  1,
333
457
  {
334
- messages: [{ type: "human", content: "hello" }],
458
+ messages: [
459
+ expect.objectContaining({ type: "human", content: "hello" }),
460
+ ],
335
461
  },
336
462
  config,
337
463
  );
@@ -659,6 +785,55 @@ describe("useStreamRuntime staged messages", () => {
659
785
  expect(stream.submit).not.toHaveBeenCalled();
660
786
  });
661
787
 
788
+ it("does not resurrect a staged draft that an edit already truncated", async () => {
789
+ const stream = createMockStream([
790
+ message("u1", "human", "first"),
791
+ message("a1", "ai", "first answer"),
792
+ message("u2", "human", "second"),
793
+ ]);
794
+ const { auiResult, rerender } = renderAui(stream);
795
+
796
+ await act(async () => {
797
+ auiResult.current.thread.append({
798
+ role: "user",
799
+ content: [{ type: "text", text: "draft" }],
800
+ startRun: false,
801
+ });
802
+ });
803
+ await waitFor(() => {
804
+ expect(getText(auiResult.current)).toEqual([
805
+ "first",
806
+ "first answer",
807
+ "second",
808
+ "draft",
809
+ ]);
810
+ });
811
+
812
+ await act(async () => {
813
+ auiResult.current.thread.append({
814
+ role: "user",
815
+ parentId: "u1",
816
+ content: [{ type: "text", text: "edited" }],
817
+ startRun: false,
818
+ });
819
+ });
820
+ await waitFor(() => {
821
+ expect(getText(auiResult.current)).toEqual(["first", "edited"]);
822
+ });
823
+
824
+ stream.messages = [
825
+ message("u1", "human", "first"),
826
+ message("a1", "ai", "first answer from refresh"),
827
+ message("u2", "human", "second from refresh"),
828
+ ];
829
+ rerender();
830
+
831
+ await waitFor(() => {
832
+ expect(getText(auiResult.current)).toEqual(["first", "edited"]);
833
+ });
834
+ expect(stream.submit).not.toHaveBeenCalled();
835
+ });
836
+
662
837
  it("keeps later staged messages visible after promoting one staged parent", async () => {
663
838
  const stream = createMockStream([message("u1", "human", "earlier")]);
664
839
  const { auiResult, rerender } = renderAui(stream);
@@ -9,6 +9,11 @@ import {
9
9
  pickExternalStoreSharedOptions,
10
10
  } from "@assistant-ui/core";
11
11
  import type { ThreadMessage } from "@assistant-ui/core";
12
+ import {
13
+ createCloudThreadListAdapterCreateFallback,
14
+ createToolCallCancellationStub,
15
+ scanPendingToolCalls,
16
+ } from "@assistant-ui/core/internal";
12
17
  import {
13
18
  useCloudThreadListAdapter,
14
19
  useExternalStoreRuntime,
@@ -24,6 +29,8 @@ import type {
24
29
  UIMessage,
25
30
  UseStreamRuntimeOptions,
26
31
  } from "./types";
32
+ import { groupUIMessagesByParent } from "./converter";
33
+ export { groupUIMessagesByParent } from "./converter";
27
34
  import {
28
35
  convertLangChainBaseMessage,
29
36
  getMessageContent,
@@ -47,44 +54,21 @@ type NormalizedRunConfigOptions = NonNullable<
47
54
  ReturnType<typeof runConfigToSubmitOptions>
48
55
  >;
49
56
 
50
- /**
51
- * Group the graph's accumulated `UIMessage`s by the assistant message they
52
- * belong to. Non-array state and entries without a parent link are dropped.
53
- * The parent id comes from `metadata.message_id` (Python SDK) or
54
- * `metadata.id` (JS SDK).
55
- */
56
- export const groupUIMessagesByParent = (
57
- value: unknown,
58
- ): Map<string, UIMessage[]> => {
59
- const map = new Map<string, UIMessage[]>();
60
- if (!Array.isArray(value)) return map;
61
- for (const ui of value as UIMessage[]) {
62
- const parentId = ui.metadata?.message_id ?? ui.metadata?.id;
63
- if (!parentId) continue;
64
- const existing = map.get(parentId);
65
- if (existing) {
66
- existing.push(ui);
67
- } else {
68
- map.set(parentId, [ui]);
69
- }
70
- }
71
- return map;
72
- };
73
-
74
57
  const getPendingToolCalls = (
75
58
  messages: readonly LangChainBaseMessage[],
76
- ): LangChainToolCall[] => {
77
- const pending = new Map<string, LangChainToolCall>();
78
- for (const m of messages) {
79
- const type = getMessageType(m);
80
- if (type === "ai") {
81
- for (const tc of m.tool_calls ?? []) pending.set(tc.id, tc);
82
- } else if (type === "tool" && m.tool_call_id) {
83
- pending.delete(m.tool_call_id);
84
- }
85
- }
86
- return [...pending.values()];
87
- };
59
+ ): LangChainToolCall[] =>
60
+ scanPendingToolCalls(
61
+ messages,
62
+ (message) => {
63
+ const type = getMessageType(message);
64
+ if (type === "ai") return { toolCalls: message.tool_calls ?? [] };
65
+ if (type === "tool" && message.tool_call_id) {
66
+ return { toolCallId: message.tool_call_id };
67
+ }
68
+ return undefined;
69
+ },
70
+ (toolCall) => toolCall.id,
71
+ );
88
72
 
89
73
  const toStagedHumanMessage = (
90
74
  msg: AppendMessage,
@@ -95,6 +79,26 @@ const toStagedHumanMessage = (
95
79
  content: getMessageContent(msg),
96
80
  });
97
81
 
82
+ const humanContentText = (content: LangChainBaseMessage["content"]) => {
83
+ if (typeof content === "string") return content;
84
+ if (!Array.isArray(content)) return "";
85
+ return content
86
+ .filter(
87
+ (part): part is { type: "text"; text: string } =>
88
+ typeof part === "object" &&
89
+ part !== null &&
90
+ part.type === "text" &&
91
+ typeof part.text === "string",
92
+ )
93
+ .map((part) => part.text)
94
+ .join("");
95
+ };
96
+
97
+ const hasSameMessageContent = (
98
+ a: LangChainBaseMessage,
99
+ b: LangChainBaseMessage,
100
+ ) => humanContentText(a.content) === humanContentText(b.content);
101
+
98
102
  const truncateLangChainBaseMessages = (
99
103
  threadMessages: readonly ThreadMessage[],
100
104
  parentId: string | null,
@@ -171,7 +175,8 @@ const useStreamThreadRuntime = (
171
175
  const convertWithUI = useMemo<
172
176
  useExternalMessageConverter.Callback<LangChainBaseMessage>
173
177
  >(() => {
174
- const uiMessagesByParent = groupUIMessagesByParent(mergedUiMessages);
178
+ const uiMessagesByParent =
179
+ groupUIMessagesByParent<UIMessage>(mergedUiMessages);
175
180
  return (message, metadata) =>
176
181
  convertLangChainBaseMessage(message, {
177
182
  ...metadata,
@@ -262,11 +267,12 @@ const useStreamThreadRuntime = (
262
267
  {
263
268
  message: LangChainBaseMessage & { id: string };
264
269
  runConfig: AppendMessage["runConfig"];
270
+ reconcileOnEcho: boolean;
271
+ baseMessageCount: number;
265
272
  }
266
273
  >(),
267
274
  );
268
275
  const stagedBaseMessagesRef = useRef<LangChainBaseMessage[] | null>(null);
269
-
270
276
  useEffect(() => {
271
277
  if (stagedMessagesRef.current.size === 0) return;
272
278
 
@@ -274,18 +280,32 @@ const useStreamThreadRuntime = (
274
280
  const baseMessages =
275
281
  stagedBaseMessagesRef.current ??
276
282
  (stream.messages as LangChainBaseMessage[]);
277
- const baseMessageIds = new Set(
278
- baseMessages.flatMap((message) => (message.id ? [message.id] : [])),
279
- );
280
283
  const remainingStagedMessages: LangChainBaseMessage[] = [];
281
- const seenStagedIds = new Set<string>();
282
- for (const message of visibleMessagesRef.current) {
283
- if (!message.id || seenStagedIds.has(message.id)) continue;
284
- if (baseMessageIds.has(message.id)) continue;
285
- const staged = stagedMessagesRef.current.get(message.id);
286
- if (!staged) continue;
287
- remainingStagedMessages.push(staged.message);
288
- seenStagedIds.add(message.id);
284
+ const matchedBaseMessageIndexes = new Set<number>();
285
+ const visibleStagedIds = new Set(
286
+ visibleMessagesRef.current.flatMap((m) => (m.id ? [m.id] : [])),
287
+ );
288
+ for (const [id, staged] of stagedMessagesRef.current) {
289
+ if (!visibleStagedIds.has(id)) continue;
290
+ const echoed = baseMessages.some((message, index) => {
291
+ if (matchedBaseMessageIndexes.has(index)) return false;
292
+ if (message.id === id) {
293
+ matchedBaseMessageIndexes.add(index);
294
+ return true;
295
+ }
296
+ if (
297
+ !staged.reconcileOnEcho ||
298
+ index < staged.baseMessageCount ||
299
+ getMessageType(message) !== "human" ||
300
+ !hasSameMessageContent(message, staged.message)
301
+ ) {
302
+ return false;
303
+ }
304
+ matchedBaseMessageIndexes.add(index);
305
+ return true;
306
+ });
307
+ if (echoed) stagedMessagesRef.current.delete(id);
308
+ else remainingStagedMessages.push(staged.message);
289
309
  }
290
310
 
291
311
  if (remainingStagedMessages.length === 0) {
@@ -317,15 +337,32 @@ const useStreamThreadRuntime = (
317
337
  };
318
338
  };
319
339
 
320
- const stageUserMessage = (msg: AppendMessage) => {
340
+ const stageUserMessage = (msg: AppendMessage, reconcileOnEcho = false) => {
321
341
  const stagedMessage = toStagedHumanMessage(msg);
322
342
  stagedMessagesRef.current.set(stagedMessage.id, {
323
343
  message: stagedMessage,
324
344
  runConfig: msg.runConfig,
345
+ reconcileOnEcho,
346
+ baseMessageCount: streamRef.current.messages.length,
325
347
  });
326
348
  const nextMessages = [...visibleMessagesRef.current, stagedMessage];
327
349
  visibleMessagesRef.current = nextMessages;
328
350
  setStagedMessages(nextMessages);
351
+ return stagedMessage;
352
+ };
353
+
354
+ const removeStagedMessage = (id: string) => {
355
+ if (!stagedMessagesRef.current.delete(id)) return;
356
+ const nextMessages = visibleMessagesRef.current.filter(
357
+ (message) => message.id !== id,
358
+ );
359
+ visibleMessagesRef.current = nextMessages;
360
+ if (stagedMessagesRef.current.size === 0) {
361
+ stagedBaseMessagesRef.current = null;
362
+ setStagedMessages(null);
363
+ } else {
364
+ setStagedMessages(nextMessages);
365
+ }
329
366
  };
330
367
 
331
368
  const extras = useMemo(
@@ -357,7 +394,7 @@ const useStreamThreadRuntime = (
357
394
 
358
395
  const runtime = useExternalStoreRuntime({
359
396
  ...pickExternalStoreSharedOptions(options),
360
- isRunning: effectiveIsRunning,
397
+ isRunning: stream.isLoading,
361
398
  isLoading: stream.isThreadLoading,
362
399
  messages: threadMessages,
363
400
  adapters,
@@ -370,32 +407,42 @@ const useStreamThreadRuntime = (
370
407
  return;
371
408
  }
372
409
 
410
+ const stagedMessage = stageUserMessage(msg, true);
411
+ const stagedMessageId = stagedMessage.id;
373
412
  setActiveRunConfig(msg.runConfig);
374
413
  const content = getMessageContent(msg);
375
414
  const cancellations =
376
415
  autoCancelPendingToolCalls !== false
377
416
  ? getPendingToolCalls(
378
417
  streamRef.current.messages as readonly LangChainBaseMessage[],
379
- ).map((t) => ({
380
- type: "tool" as const,
381
- name: t.name,
382
- tool_call_id: t.id,
383
- content: JSON.stringify({ cancelled: true }),
384
- status: "error" as const,
385
- }))
418
+ ).map(createToolCallCancellationStub)
386
419
  : [];
387
420
  // A null threadId is not a no-op for the SDK: it rebinds the controller
388
421
  // away from its self-created thread and forces a fresh one, so the
389
422
  // submit waits for initialization to produce an identity; core no
390
423
  // longer holds appends on that barrier.
391
- const { externalId } = await aui.threadListItem.initialize();
392
- await streamRef.current.submit(
393
- { [messagesKey]: [...cancellations, { type: "human", content }] },
394
- {
395
- ...runConfigToSubmitOptions(msg.runConfig),
396
- ...(externalId != null ? { threadId: externalId } : {}),
397
- },
398
- );
424
+ try {
425
+ const { externalId } = await aui.threadListItem.initialize();
426
+ await streamRef.current.submit(
427
+ {
428
+ [messagesKey]: [
429
+ ...cancellations,
430
+ {
431
+ id: stagedMessageId,
432
+ type: "human",
433
+ content,
434
+ },
435
+ ],
436
+ },
437
+ {
438
+ ...runConfigToSubmitOptions(msg.runConfig),
439
+ ...(externalId != null ? { threadId: externalId } : {}),
440
+ },
441
+ );
442
+ } catch (error) {
443
+ removeStagedMessage(stagedMessageId);
444
+ throw error;
445
+ }
399
446
  },
400
447
  onAddToolResult: async ({
401
448
  messageId,
@@ -487,6 +534,8 @@ const useStreamThreadRuntime = (
487
534
  stagedMessagesRef.current.set(stagedMessage.id, {
488
535
  message: stagedMessage,
489
536
  runConfig: message.runConfig,
537
+ reconcileOnEcho: false,
538
+ baseMessageCount: 0,
490
539
  });
491
540
  stagedBaseMessagesRef.current = truncated;
492
541
  const nextMessages = [...truncated, stagedMessage];
@@ -567,9 +616,13 @@ export const useStreamRuntime = (rawOptions: UseStreamRuntimeOptions) => {
567
616
  const optionsRef = useRef(options);
568
617
  optionsRef.current = options;
569
618
 
619
+ const aui = useAui();
570
620
  const cloudAdapter = useCloudThreadListAdapter({
571
621
  cloud,
572
- create,
622
+ create: createCloudThreadListAdapterCreateFallback(
623
+ create,
624
+ aui.threadListItem,
625
+ ),
573
626
  delete: deleteFn,
574
627
  });
575
628
  const adapter = unstable_threadListAdapter ?? cloudAdapter;