@assistant-ui/react-langchain 0.0.27 → 0.0.29

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,76 +1,13 @@
1
1
  "use client";
2
2
 
3
3
  import type { MessageTiming } from "@assistant-ui/core";
4
- import {
5
- useStreamingTiming,
6
- type StreamingTimingAccessors,
7
- } from "@assistant-ui/core/react";
8
- import type { LangChainBaseMessage, LangChainContentBlock } from "./types";
4
+ import { useStreamingTiming } from "@assistant-ui/core/react";
5
+ import { createLangChainStreamingTimingAccessors } from "./converter";
9
6
  import { getMessageType } from "./convertMessages";
7
+ import type { LangChainBaseMessage } from "./types";
10
8
 
11
- const findAiMessage = (
12
- messages: readonly LangChainBaseMessage[],
13
- messageId: string,
14
- ): LangChainBaseMessage | undefined =>
15
- messages.find((m) => getMessageType(m) === "ai" && m.id === messageId);
16
-
17
- const reasoningTextLength = (part: {
18
- readonly summary?: ReadonlyArray<{ readonly text?: string }>;
19
- readonly reasoning?: string;
20
- }): number => {
21
- if (part.summary && part.summary.length > 0)
22
- return part.summary.map((s) => s?.text ?? "").join("\n\n\n").length;
23
- return part.reasoning?.length ?? 0;
24
- };
25
-
26
- const getTextLength = (
27
- messages: readonly LangChainBaseMessage[],
28
- messageId: string,
29
- ): number => {
30
- const m = findAiMessage(messages, messageId);
31
- if (!m) return 0;
32
- const content = m.content;
33
- if (typeof content === "string") return content.length;
34
- if (!Array.isArray(content)) return 0;
35
- let len = 0;
36
- for (const part of content as readonly LangChainContentBlock[]) {
37
- switch (part.type) {
38
- case "text":
39
- case "text_delta":
40
- if (typeof part.text === "string") len += part.text.length;
41
- break;
42
- case "thinking":
43
- if (typeof part.thinking === "string") len += part.thinking.length;
44
- break;
45
- case "reasoning":
46
- len += reasoningTextLength(part);
47
- break;
48
- }
49
- }
50
- return len;
51
- };
52
-
53
- const getToolCallCount = (
54
- messages: readonly LangChainBaseMessage[],
55
- messageId: string,
56
- ): number => findAiMessage(messages, messageId)?.tool_calls?.length ?? 0;
57
-
58
- const getAssistantMessageId = (
59
- messages: readonly LangChainBaseMessage[],
60
- ): string | undefined => {
61
- for (let i = messages.length - 1; i >= 0; i--) {
62
- const m = messages[i];
63
- if (m && getMessageType(m) === "ai" && m.id) return m.id;
64
- }
65
- return undefined;
66
- };
67
-
68
- export const langChainStreamingTimingAccessors: StreamingTimingAccessors<LangChainBaseMessage> =
69
- {
70
- getAssistantMessageId,
71
- getTextLength,
72
- getToolCallCount,
73
- };
9
+ export const langChainStreamingTimingAccessors =
10
+ createLangChainStreamingTimingAccessors<LangChainBaseMessage>(getMessageType);
74
11
 
75
12
  /**
76
13
  * Tracks per-message streaming timing for LangChain messages. Delegates to
package/src/types.ts CHANGED
@@ -17,45 +17,7 @@ import type {
17
17
  SubgraphDiscoverySnapshot,
18
18
  } from "@langchain/react";
19
19
 
20
- /** Known content block types from @langchain/core messages. */
21
- export type LangChainContentBlock =
22
- | { type: "text"; text: string }
23
- | { type: "text_delta"; text: string }
24
- | { type: "image_url"; image_url: string | { url?: string } }
25
- | { type: "thinking"; thinking: string }
26
- | {
27
- type: "reasoning";
28
- summary?: Array<{ type: "summary_text"; text?: string }>;
29
- reasoning?: string;
30
- }
31
- | {
32
- type: "file";
33
- data: string;
34
- mime_type: string;
35
- source_type?: "base64";
36
- metadata?: { filename?: string };
37
- }
38
- | {
39
- type: "file";
40
- url: string;
41
- mime_type?: string;
42
- source_type: "url";
43
- metadata?: { filename?: string };
44
- }
45
- | {
46
- type: "file";
47
- id: string;
48
- mime_type?: string;
49
- source_type: "id";
50
- metadata?: { filename?: string };
51
- }
52
- | {
53
- type: "audio";
54
- data: string;
55
- mime_type: string;
56
- source_type: "base64";
57
- }
58
- | { type: "tool_use" | "input_json_delta" };
20
+ export type { LangChainContentBlock } from "./converter";
59
21
 
60
22
  export type LangChainToolCall = {
61
23
  id: string;
@@ -10,7 +10,7 @@ import type {
10
10
  } from "@assistant-ui/core";
11
11
  import { useAui } from "@assistant-ui/store";
12
12
  import type { LangChainBaseMessage } from "./types";
13
- import type { ReactNode } from "react";
13
+ import { startTransition, Suspense, type ReactNode } from "react";
14
14
  import {
15
15
  useLangChainRespond,
16
16
  useLangChainRespondAll,
@@ -898,3 +898,76 @@ describe("useStreamRuntime staged messages", () => {
898
898
  });
899
899
  });
900
900
  });
901
+
902
+ describe("useStreamRuntime committed refs", () => {
903
+ it("submits through the committed stream after an abandoned render", async () => {
904
+ const streamA = createMockStream();
905
+ const streamB = createMockStream();
906
+ const adapter = makeThreadListAdapter();
907
+ mockUseStream.mockImplementation((options: { apiUrl: string }) =>
908
+ options.apiUrl === "/api/b" ? streamB : streamA,
909
+ );
910
+ const host = renderHook(() =>
911
+ useStreamRuntime({
912
+ apiUrl: "/api/a",
913
+ unstable_threadListAdapter: adapter,
914
+ } as never),
915
+ );
916
+
917
+ const pending = new Promise<never>(() => {});
918
+ let blocked = false;
919
+ const interruptedRender = vi.fn();
920
+ const Blocker = () => {
921
+ if (blocked) {
922
+ interruptedRender();
923
+ throw pending;
924
+ }
925
+ return null;
926
+ };
927
+
928
+ const capture: { runtime: AssistantRuntime | null } = { runtime: null };
929
+ const Nested = ({ apiUrl }: { apiUrl: string }) => {
930
+ capture.runtime = useStreamRuntime({
931
+ apiUrl,
932
+ unstable_threadListAdapter: adapter,
933
+ } as never);
934
+ return null;
935
+ };
936
+ const Tree = ({ apiUrl }: { apiUrl: string }) => (
937
+ <AssistantRuntimeProvider runtime={host.result.current}>
938
+ <Suspense fallback={null}>
939
+ <Nested apiUrl={apiUrl} />
940
+ <Blocker />
941
+ </Suspense>
942
+ </AssistantRuntimeProvider>
943
+ );
944
+
945
+ const view = render(<Tree apiUrl="/api/a" />);
946
+ expect(capture.runtime).not.toBeNull();
947
+
948
+ act(() => {
949
+ blocked = true;
950
+ startTransition(() => view.rerender(<Tree apiUrl="/api/b" />));
951
+ });
952
+ expect(interruptedRender).toHaveBeenCalled();
953
+
954
+ await act(async () => {
955
+ await capture.runtime!.thread.append("hello");
956
+ });
957
+
958
+ expect(streamA.submit).toHaveBeenCalledOnce();
959
+ expect(streamB.submit).not.toHaveBeenCalled();
960
+
961
+ await act(async () => {
962
+ blocked = false;
963
+ view.rerender(<Tree apiUrl="/api/b" />);
964
+ });
965
+ await act(async () => {
966
+ await capture.runtime!.thread.append("second");
967
+ });
968
+
969
+ expect(streamB.submit).toHaveBeenCalledOnce();
970
+ view.unmount();
971
+ host.unmount();
972
+ });
973
+ });
@@ -1,7 +1,14 @@
1
1
  /// <reference types="@assistant-ui/core/store" />
2
2
  "use client";
3
3
 
4
- import { useCallback, useEffect, useMemo, useRef, useState } from "react";
4
+ import {
5
+ useCallback,
6
+ useEffect,
7
+ useInsertionEffect,
8
+ useMemo,
9
+ useRef,
10
+ useState,
11
+ } from "react";
5
12
  import type { AppendMessage, ToolExecutionStatus } from "@assistant-ui/core";
6
13
  import {
7
14
  generateId,
@@ -9,6 +16,11 @@ import {
9
16
  pickExternalStoreSharedOptions,
10
17
  } from "@assistant-ui/core";
11
18
  import type { ThreadMessage } from "@assistant-ui/core";
19
+ import {
20
+ createCloudThreadListAdapterCreateFallback,
21
+ createToolCallCancellationStub,
22
+ scanPendingToolCalls,
23
+ } from "@assistant-ui/core/internal";
12
24
  import {
13
25
  useCloudThreadListAdapter,
14
26
  useExternalStoreRuntime,
@@ -24,6 +36,8 @@ import type {
24
36
  UIMessage,
25
37
  UseStreamRuntimeOptions,
26
38
  } from "./types";
39
+ import { groupUIMessagesByParent } from "./converter";
40
+ export { groupUIMessagesByParent } from "./converter";
27
41
  import {
28
42
  convertLangChainBaseMessage,
29
43
  getMessageContent,
@@ -47,44 +61,21 @@ type NormalizedRunConfigOptions = NonNullable<
47
61
  ReturnType<typeof runConfigToSubmitOptions>
48
62
  >;
49
63
 
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
64
  const getPendingToolCalls = (
75
65
  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
- };
66
+ ): LangChainToolCall[] =>
67
+ scanPendingToolCalls(
68
+ messages,
69
+ (message) => {
70
+ const type = getMessageType(message);
71
+ if (type === "ai") return { toolCalls: message.tool_calls ?? [] };
72
+ if (type === "tool" && message.tool_call_id) {
73
+ return { toolCallId: message.tool_call_id };
74
+ }
75
+ return undefined;
76
+ },
77
+ (toolCall) => toolCall.id,
78
+ );
88
79
 
89
80
  const toStagedHumanMessage = (
90
81
  msg: AppendMessage,
@@ -191,7 +182,8 @@ const useStreamThreadRuntime = (
191
182
  const convertWithUI = useMemo<
192
183
  useExternalMessageConverter.Callback<LangChainBaseMessage>
193
184
  >(() => {
194
- const uiMessagesByParent = groupUIMessagesByParent(mergedUiMessages);
185
+ const uiMessagesByParent =
186
+ groupUIMessagesByParent<UIMessage>(mergedUiMessages);
195
187
  return (message, metadata) =>
196
188
  convertLangChainBaseMessage(message, {
197
189
  ...metadata,
@@ -207,7 +199,9 @@ const useStreamThreadRuntime = (
207
199
  });
208
200
 
209
201
  const streamRef = useRef(stream);
210
- streamRef.current = stream;
202
+ useInsertionEffect(() => {
203
+ streamRef.current = stream;
204
+ }, [stream]);
211
205
 
212
206
  const activeRunConfigRef = useRef<
213
207
  NormalizedRunConfigOptions["config"] | undefined
@@ -271,10 +265,14 @@ const useStreamThreadRuntime = (
271
265
  }, [stream.messages]);
272
266
 
273
267
  const visibleMessagesRef = useRef(visibleMessages);
274
- visibleMessagesRef.current = visibleMessages;
268
+ useInsertionEffect(() => {
269
+ visibleMessagesRef.current = visibleMessages;
270
+ }, [visibleMessages]);
275
271
 
276
272
  const threadMessagesRef = useRef(threadMessages);
277
- threadMessagesRef.current = threadMessages;
273
+ useInsertionEffect(() => {
274
+ threadMessagesRef.current = threadMessages;
275
+ }, [threadMessages]);
278
276
 
279
277
  const stagedMessagesRef = useRef(
280
278
  new Map<
@@ -409,7 +407,7 @@ const useStreamThreadRuntime = (
409
407
 
410
408
  const runtime = useExternalStoreRuntime({
411
409
  ...pickExternalStoreSharedOptions(options),
412
- isRunning: effectiveIsRunning,
410
+ isRunning: stream.isLoading,
413
411
  isLoading: stream.isThreadLoading,
414
412
  messages: threadMessages,
415
413
  adapters,
@@ -430,13 +428,7 @@ const useStreamThreadRuntime = (
430
428
  autoCancelPendingToolCalls !== false
431
429
  ? getPendingToolCalls(
432
430
  streamRef.current.messages as readonly LangChainBaseMessage[],
433
- ).map((t) => ({
434
- type: "tool" as const,
435
- name: t.name,
436
- tool_call_id: t.id,
437
- content: JSON.stringify({ cancelled: true }),
438
- status: "error" as const,
439
- }))
431
+ ).map(createToolCallCancellationStub)
440
432
  : [];
441
433
  // A null threadId is not a no-op for the SDK: it rebinds the controller
442
434
  // away from its self-created thread and forces a fresh one, so the
@@ -634,19 +626,20 @@ export const useStreamRuntime = (rawOptions: UseStreamRuntimeOptions) => {
634
626
  ...options
635
627
  } = rawOptions;
636
628
 
637
- const optionsRef = useRef(options);
638
- optionsRef.current = options;
639
-
629
+ const aui = useAui();
640
630
  const cloudAdapter = useCloudThreadListAdapter({
641
631
  cloud,
642
- create,
632
+ create: createCloudThreadListAdapterCreateFallback(
633
+ create,
634
+ aui.threadListItem,
635
+ ),
643
636
  delete: deleteFn,
644
637
  });
645
638
  const adapter = unstable_threadListAdapter ?? cloudAdapter;
646
639
 
647
640
  return useRemoteThreadListRuntime({
648
641
  runtimeHook: function RuntimeHook() {
649
- return useStreamThreadRuntime(optionsRef.current);
642
+ return useStreamThreadRuntime(options);
650
643
  },
651
644
  adapter,
652
645
  allowNesting: true,