@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.
- package/dist/useStreamRuntime.d.ts.map +1 -1
- package/dist/useStreamRuntime.js +112 -25
- package/dist/useStreamRuntime.js.map +1 -1
- package/package.json +8 -8
- package/src/useStreamRuntime.test.tsx +593 -0
- package/src/useStreamRuntime.ts +201 -36
package/src/useStreamRuntime.ts
CHANGED
|
@@ -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
|
|
216
|
-
|
|
217
|
-
|
|
218
|
-
|
|
219
|
-
|
|
220
|
-
if (!
|
|
221
|
-
|
|
222
|
-
|
|
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:
|
|
276
|
-
|
|
277
|
-
|
|
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
|
-
|
|
313
|
-
|
|
314
|
-
|
|
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
|
-
|
|
325
|
-
|
|
326
|
-
|
|
327
|
-
|
|
328
|
-
|
|
329
|
-
|
|
330
|
-
|
|
331
|
-
|
|
332
|
-
|
|
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(
|
|
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,
|