@assistant-ui/react-langchain 0.0.25 → 0.0.27

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.
@@ -1,7 +1,7 @@
1
1
  /// <reference types="@assistant-ui/core/store" />
2
2
  "use client";
3
3
 
4
- import { useEffect, useMemo, useRef, useState } from "react";
4
+ import { useCallback, useEffect, useMemo, useRef, useState } from "react";
5
5
  import type { AppendMessage, ToolExecutionStatus } from "@assistant-ui/core";
6
6
  import {
7
7
  generateId,
@@ -15,7 +15,7 @@ import {
15
15
  useExternalMessageConverter,
16
16
  useRemoteThreadListRuntime,
17
17
  } from "@assistant-ui/core/react";
18
- import { useAuiState } from "@assistant-ui/store";
18
+ import { useAui, useAuiState } from "@assistant-ui/store";
19
19
  import { STREAM_CONTROLLER, useChannel, useStream } from "@langchain/react";
20
20
  import type { Channel } from "@langchain/react";
21
21
  import type {
@@ -43,6 +43,10 @@ export const runConfigToSubmitOptions = (
43
43
  ? { config: { configurable: runConfig.custom } }
44
44
  : undefined;
45
45
 
46
+ type NormalizedRunConfigOptions = NonNullable<
47
+ ReturnType<typeof runConfigToSubmitOptions>
48
+ >;
49
+
46
50
  /**
47
51
  * Group the graph's accumulated `UIMessage`s by the assistant message they
48
52
  * belong to. Non-array state and entries without a parent link are dropped.
@@ -91,6 +95,26 @@ const toStagedHumanMessage = (
91
95
  content: getMessageContent(msg),
92
96
  });
93
97
 
98
+ const humanContentText = (content: LangChainBaseMessage["content"]) => {
99
+ if (typeof content === "string") return content;
100
+ if (!Array.isArray(content)) return "";
101
+ return content
102
+ .filter(
103
+ (part): part is { type: "text"; text: string } =>
104
+ typeof part === "object" &&
105
+ part !== null &&
106
+ part.type === "text" &&
107
+ typeof part.text === "string",
108
+ )
109
+ .map((part) => part.text)
110
+ .join("");
111
+ };
112
+
113
+ const hasSameMessageContent = (
114
+ a: LangChainBaseMessage,
115
+ b: LangChainBaseMessage,
116
+ ) => humanContentText(a.content) === humanContentText(b.content);
117
+
94
118
  const truncateLangChainBaseMessages = (
95
119
  threadMessages: readonly ThreadMessage[],
96
120
  parentId: string | null,
@@ -119,6 +143,7 @@ const useStreamThreadRuntime = (
119
143
  ) => {
120
144
  const { adapters, autoCancelPendingToolCalls, unstable_allowCancellation } =
121
145
  options;
146
+ const aui = useAui();
122
147
  const messagesKey = options.messagesKey ?? "messages";
123
148
  const uiStateKey = options.uiStateKey ?? "ui";
124
149
 
@@ -184,6 +209,67 @@ const useStreamThreadRuntime = (
184
209
  const streamRef = useRef(stream);
185
210
  streamRef.current = stream;
186
211
 
212
+ const activeRunConfigRef = useRef<
213
+ NormalizedRunConfigOptions["config"] | undefined
214
+ >(undefined);
215
+ const runConfigByMessageIdRef = useRef(
216
+ new Map<string, NormalizedRunConfigOptions["config"] | undefined>(),
217
+ );
218
+ const activeThreadIdRef = useRef(externalId);
219
+ const setActiveRunConfig = useCallback(
220
+ (runConfig: AppendMessage["runConfig"]) => {
221
+ activeRunConfigRef.current = runConfigToSubmitOptions(runConfig)?.config;
222
+ },
223
+ [],
224
+ );
225
+ const withActiveRunConfig = useCallback(
226
+ (submitOptions?: Record<string, unknown>) => {
227
+ if (submitOptions && "config" in submitOptions) return submitOptions;
228
+ if (activeRunConfigRef.current === undefined) return submitOptions;
229
+ return { ...submitOptions, config: activeRunConfigRef.current };
230
+ },
231
+ [],
232
+ );
233
+
234
+ useEffect(() => {
235
+ if (
236
+ activeThreadIdRef.current !== null &&
237
+ activeThreadIdRef.current !== externalId
238
+ ) {
239
+ activeRunConfigRef.current = undefined;
240
+ runConfigByMessageIdRef.current.clear();
241
+ }
242
+ activeThreadIdRef.current = externalId;
243
+ }, [externalId]);
244
+
245
+ useEffect(() => {
246
+ const messages = stream.messages as readonly LangChainBaseMessage[];
247
+ const owned = runConfigByMessageIdRef.current;
248
+ for (let i = messages.length - 1; i >= 0; i--) {
249
+ const message = messages.at(i);
250
+ if (
251
+ !message?.id ||
252
+ getMessageType(message) !== "ai" ||
253
+ !message.tool_calls?.length
254
+ ) {
255
+ continue;
256
+ }
257
+ if (owned.has(message.id)) return;
258
+ break;
259
+ }
260
+ for (const message of messages) {
261
+ if (
262
+ !message.id ||
263
+ getMessageType(message) !== "ai" ||
264
+ !message.tool_calls?.length ||
265
+ owned.has(message.id)
266
+ ) {
267
+ continue;
268
+ }
269
+ owned.set(message.id, activeRunConfigRef.current);
270
+ }
271
+ }, [stream.messages]);
272
+
187
273
  const visibleMessagesRef = useRef(visibleMessages);
188
274
  visibleMessagesRef.current = visibleMessages;
189
275
 
@@ -196,11 +282,12 @@ const useStreamThreadRuntime = (
196
282
  {
197
283
  message: LangChainBaseMessage & { id: string };
198
284
  runConfig: AppendMessage["runConfig"];
285
+ reconcileOnEcho: boolean;
286
+ baseMessageCount: number;
199
287
  }
200
288
  >(),
201
289
  );
202
290
  const stagedBaseMessagesRef = useRef<LangChainBaseMessage[] | null>(null);
203
-
204
291
  useEffect(() => {
205
292
  if (stagedMessagesRef.current.size === 0) return;
206
293
 
@@ -208,18 +295,32 @@ const useStreamThreadRuntime = (
208
295
  const baseMessages =
209
296
  stagedBaseMessagesRef.current ??
210
297
  (stream.messages as LangChainBaseMessage[]);
211
- const baseMessageIds = new Set(
212
- baseMessages.flatMap((message) => (message.id ? [message.id] : [])),
213
- );
214
298
  const remainingStagedMessages: LangChainBaseMessage[] = [];
215
- const seenStagedIds = new Set<string>();
216
- for (const message of visibleMessagesRef.current) {
217
- if (!message.id || seenStagedIds.has(message.id)) continue;
218
- if (baseMessageIds.has(message.id)) continue;
219
- const staged = stagedMessagesRef.current.get(message.id);
220
- if (!staged) continue;
221
- remainingStagedMessages.push(staged.message);
222
- seenStagedIds.add(message.id);
299
+ const matchedBaseMessageIndexes = new Set<number>();
300
+ const visibleStagedIds = new Set(
301
+ visibleMessagesRef.current.flatMap((m) => (m.id ? [m.id] : [])),
302
+ );
303
+ for (const [id, staged] of stagedMessagesRef.current) {
304
+ if (!visibleStagedIds.has(id)) continue;
305
+ const echoed = baseMessages.some((message, index) => {
306
+ if (matchedBaseMessageIndexes.has(index)) return false;
307
+ if (message.id === id) {
308
+ matchedBaseMessageIndexes.add(index);
309
+ return true;
310
+ }
311
+ if (
312
+ !staged.reconcileOnEcho ||
313
+ index < staged.baseMessageCount ||
314
+ getMessageType(message) !== "human" ||
315
+ !hasSameMessageContent(message, staged.message)
316
+ ) {
317
+ return false;
318
+ }
319
+ matchedBaseMessageIndexes.add(index);
320
+ return true;
321
+ });
322
+ if (echoed) stagedMessagesRef.current.delete(id);
323
+ else remainingStagedMessages.push(staged.message);
223
324
  }
224
325
 
225
326
  if (remainingStagedMessages.length === 0) {
@@ -251,15 +352,32 @@ const useStreamThreadRuntime = (
251
352
  };
252
353
  };
253
354
 
254
- const stageUserMessage = (msg: AppendMessage) => {
355
+ const stageUserMessage = (msg: AppendMessage, reconcileOnEcho = false) => {
255
356
  const stagedMessage = toStagedHumanMessage(msg);
256
357
  stagedMessagesRef.current.set(stagedMessage.id, {
257
358
  message: stagedMessage,
258
359
  runConfig: msg.runConfig,
360
+ reconcileOnEcho,
361
+ baseMessageCount: streamRef.current.messages.length,
259
362
  });
260
363
  const nextMessages = [...visibleMessagesRef.current, stagedMessage];
261
364
  visibleMessagesRef.current = nextMessages;
262
365
  setStagedMessages(nextMessages);
366
+ return stagedMessage;
367
+ };
368
+
369
+ const removeStagedMessage = (id: string) => {
370
+ if (!stagedMessagesRef.current.delete(id)) return;
371
+ const nextMessages = visibleMessagesRef.current.filter(
372
+ (message) => message.id !== id,
373
+ );
374
+ visibleMessagesRef.current = nextMessages;
375
+ if (stagedMessagesRef.current.size === 0) {
376
+ stagedBaseMessagesRef.current = null;
377
+ setStagedMessages(null);
378
+ } else {
379
+ setStagedMessages(nextMessages);
380
+ }
263
381
  };
264
382
 
265
383
  const extras = useMemo(
@@ -272,13 +390,21 @@ const useStreamThreadRuntime = (
272
390
  subgraphs: stream.subgraphs,
273
391
  stream,
274
392
  error: stream.error,
275
- submit: stream.submit,
276
- respond: stream.respond,
277
- respondAll: stream.respondAll,
393
+ submit: (values, submitOptions) => {
394
+ const isResume = values == null || submitOptions?.command != null;
395
+ return stream.submit(
396
+ values,
397
+ isResume ? withActiveRunConfig(submitOptions) : submitOptions,
398
+ );
399
+ },
400
+ respond: (response, respondOptions) =>
401
+ stream.respond(response, withActiveRunConfig(respondOptions)),
402
+ respondAll: (responsesById, respondOptions) =>
403
+ stream.respondAll(responsesById, withActiveRunConfig(respondOptions)),
278
404
  values: stream.values,
279
405
  messagesKey,
280
406
  }),
281
- [stream, messagesKey],
407
+ [stream, messagesKey, withActiveRunConfig],
282
408
  );
283
409
 
284
410
  const runtime = useExternalStoreRuntime({
@@ -296,6 +422,9 @@ const useStreamThreadRuntime = (
296
422
  return;
297
423
  }
298
424
 
425
+ const stagedMessage = stageUserMessage(msg, true);
426
+ const stagedMessageId = stagedMessage.id;
427
+ setActiveRunConfig(msg.runConfig);
299
428
  const content = getMessageContent(msg);
300
429
  const cancellations =
301
430
  autoCancelPendingToolCalls !== false
@@ -309,30 +438,59 @@ const useStreamThreadRuntime = (
309
438
  status: "error" as const,
310
439
  }))
311
440
  : [];
312
- await stream.submit(
313
- { [messagesKey]: [...cancellations, { type: "human", content }] },
314
- runConfigToSubmitOptions(msg.runConfig),
315
- );
441
+ // A null threadId is not a no-op for the SDK: it rebinds the controller
442
+ // away from its self-created thread and forces a fresh one, so the
443
+ // submit waits for initialization to produce an identity; core no
444
+ // longer holds appends on that barrier.
445
+ try {
446
+ const { externalId } = await aui.threadListItem.initialize();
447
+ await streamRef.current.submit(
448
+ {
449
+ [messagesKey]: [
450
+ ...cancellations,
451
+ {
452
+ id: stagedMessageId,
453
+ type: "human",
454
+ content,
455
+ },
456
+ ],
457
+ },
458
+ {
459
+ ...runConfigToSubmitOptions(msg.runConfig),
460
+ ...(externalId != null ? { threadId: externalId } : {}),
461
+ },
462
+ );
463
+ } catch (error) {
464
+ removeStagedMessage(stagedMessageId);
465
+ throw error;
466
+ }
316
467
  },
317
468
  onAddToolResult: async ({
469
+ messageId,
318
470
  toolCallId,
319
471
  toolName,
320
472
  result,
321
473
  isError,
322
474
  artifact,
323
475
  }) => {
324
- await stream.submit({
325
- [messagesKey]: [
326
- {
327
- type: "tool",
328
- name: toolName,
329
- tool_call_id: toolCallId,
330
- content: JSON.stringify(result),
331
- ...(artifact !== undefined && { artifact }),
332
- status: isError ? "error" : "success",
333
- },
334
- ],
335
- });
476
+ const runConfig = runConfigByMessageIdRef.current.has(messageId)
477
+ ? runConfigByMessageIdRef.current.get(messageId)
478
+ : activeRunConfigRef.current;
479
+ await stream.submit(
480
+ {
481
+ [messagesKey]: [
482
+ {
483
+ type: "tool",
484
+ name: toolName,
485
+ tool_call_id: toolCallId,
486
+ content: JSON.stringify(result),
487
+ ...(artifact !== undefined && { artifact }),
488
+ status: isError ? "error" : "success",
489
+ },
490
+ ],
491
+ },
492
+ runConfig === undefined ? undefined : { config: runConfig },
493
+ );
336
494
  },
337
495
  onReload: async (parentId, config) => {
338
496
  const stagedRun = getStagedRun(parentId);
@@ -353,6 +511,8 @@ const useStreamThreadRuntime = (
353
511
  } else {
354
512
  setStagedMessages(null);
355
513
  }
514
+ const runConfig = config.runConfig ?? stagedRun.runConfig;
515
+ setActiveRunConfig(runConfig);
356
516
  await stream.submit(
357
517
  {
358
518
  [messagesKey]: stagedRun.messages.map((message) => ({
@@ -361,7 +521,7 @@ const useStreamThreadRuntime = (
361
521
  content: message.content,
362
522
  })),
363
523
  },
364
- runConfigToSubmitOptions(config.runConfig ?? stagedRun.runConfig),
524
+ runConfigToSubmitOptions(runConfig),
365
525
  );
366
526
  return;
367
527
  }
@@ -379,6 +539,7 @@ const useStreamThreadRuntime = (
379
539
  messagesKey,
380
540
  );
381
541
  if (!checkpointId) return;
542
+ setActiveRunConfig(config.runConfig);
382
543
  await s.submit(null, {
383
544
  forkFrom: checkpointId,
384
545
  ...runConfigToSubmitOptions(config.runConfig),
@@ -394,6 +555,8 @@ const useStreamThreadRuntime = (
394
555
  stagedMessagesRef.current.set(stagedMessage.id, {
395
556
  message: stagedMessage,
396
557
  runConfig: message.runConfig,
558
+ reconcileOnEcho: false,
559
+ baseMessageCount: 0,
397
560
  });
398
561
  stagedBaseMessagesRef.current = truncated;
399
562
  const nextMessages = [...truncated, stagedMessage];
@@ -416,6 +579,7 @@ const useStreamThreadRuntime = (
416
579
  );
417
580
  if (!checkpointId) return;
418
581
  const content = getMessageContent(message);
582
+ setActiveRunConfig(message.runConfig);
419
583
  await s.submit(
420
584
  { [messagesKey]: [{ type: "human", content }] },
421
585
  {
@@ -427,6 +591,7 @@ const useStreamThreadRuntime = (
427
591
  onCancel:
428
592
  unstable_allowCancellation !== false
429
593
  ? async () => {
594
+ activeRunConfigRef.current = undefined;
430
595
  await stream.stop();
431
596
  }
432
597
  : undefined,