@assistant-ui/react-langchain 0.0.27 → 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.
@@ -0,0 +1,372 @@
1
+ import type {
2
+ AppendMessage,
3
+ DataMessagePart,
4
+ ThreadAssistantMessage,
5
+ ThreadUserMessage,
6
+ } from "@assistant-ui/core";
7
+ import {
8
+ parseDataUrl,
9
+ resolveFilePartSource,
10
+ } from "@assistant-ui/core/internal";
11
+ import type { StreamingTimingAccessors } from "@assistant-ui/core/react";
12
+
13
+ /** Known content block types from @langchain/core messages. */
14
+ export type LangChainContentBlock =
15
+ | { type: "text"; text: string }
16
+ | { type: "text_delta"; text: string }
17
+ | { type: "image_url"; image_url: string | { url?: string } }
18
+ | { type: "thinking"; thinking: string }
19
+ | {
20
+ type: "reasoning";
21
+ summary?: Array<{ type: "summary_text"; text?: string }>;
22
+ reasoning?: string;
23
+ }
24
+ | {
25
+ type: "file";
26
+ data: string;
27
+ mime_type: string;
28
+ source_type?: "base64";
29
+ metadata?: { filename?: string };
30
+ }
31
+ | {
32
+ type: "file";
33
+ url: string;
34
+ mime_type?: string;
35
+ source_type: "url";
36
+ metadata?: { filename?: string };
37
+ }
38
+ | {
39
+ type: "file";
40
+ id: string;
41
+ mime_type?: string;
42
+ source_type: "id";
43
+ metadata?: { filename?: string };
44
+ }
45
+ | {
46
+ type: "audio";
47
+ data: string;
48
+ mime_type: string;
49
+ source_type: "base64";
50
+ }
51
+ | { type: "tool_use" | "input_json_delta" };
52
+
53
+ type ConvertedContentPart =
54
+ | ThreadUserMessage["content"][number]
55
+ | ThreadAssistantMessage["content"][number];
56
+
57
+ export const convertLangChainContentBlock = (
58
+ part: LangChainContentBlock,
59
+ ): ConvertedContentPart | null | undefined => {
60
+ const type = part.type;
61
+ switch (type) {
62
+ case "text":
63
+ case "text_delta":
64
+ return { type: "text" as const, text: part.text };
65
+ case "image_url": {
66
+ const image =
67
+ typeof part.image_url === "string"
68
+ ? part.image_url
69
+ : part.image_url?.url;
70
+ if (!image) return null;
71
+ return { type: "image" as const, image };
72
+ }
73
+ case "file":
74
+ return {
75
+ type: "file" as const,
76
+ filename: part.metadata?.filename ?? "file",
77
+ data:
78
+ part.source_type === "url"
79
+ ? part.url
80
+ : part.source_type === "id"
81
+ ? part.id
82
+ : part.data,
83
+ mimeType: part.mime_type ?? "application/octet-stream",
84
+ ...((part.source_type === "url" || part.source_type === "id") && {
85
+ sourceType: part.source_type,
86
+ }),
87
+ };
88
+ case "audio": {
89
+ const mimeType = part.mime_type ?? "application/octet-stream";
90
+ const subtype = mimeType.startsWith("audio/")
91
+ ? mimeType.slice("audio/".length)
92
+ : undefined;
93
+ return {
94
+ type: "file" as const,
95
+ filename: subtype ? `audio.${subtype}` : "audio",
96
+ data: part.data,
97
+ mimeType,
98
+ };
99
+ }
100
+ case "thinking":
101
+ return { type: "reasoning" as const, text: part.thinking };
102
+ case "reasoning":
103
+ return {
104
+ type: "reasoning" as const,
105
+ text:
106
+ part.summary && part.summary.length > 0
107
+ ? part.summary.map((s) => s?.text ?? "").join("\n\n\n")
108
+ : (part.reasoning ?? ""),
109
+ };
110
+ case "tool_use":
111
+ case "input_json_delta":
112
+ return null;
113
+ default:
114
+ return undefined;
115
+ }
116
+ };
117
+
118
+ const hasVisibleText = (text: unknown): boolean =>
119
+ typeof text === "string" && text.trim() !== "";
120
+
121
+ /**
122
+ * Audio output arrives outside the content array: providers leave `content`
123
+ * empty and put the spoken text in `additional_kwargs.audio.transcript`. The
124
+ * audio bytes stay behind because no provider reports their media type, and a
125
+ * streamed response carries raw PCM rather than a playable file.
126
+ */
127
+ export const withAudioTranscript = <T extends { type: string; text?: unknown }>(
128
+ parts: readonly T[],
129
+ additionalKwargs: Record<string, unknown> | undefined,
130
+ ): readonly (T | { type: "text"; text: string })[] => {
131
+ const audio = additionalKwargs?.audio as { transcript?: unknown } | undefined;
132
+ const transcript = audio?.transcript;
133
+ if (typeof transcript !== "string" || !hasVisibleText(transcript))
134
+ return parts;
135
+ if (parts.some((part) => part.type === "text" && hasVisibleText(part.text)))
136
+ return parts;
137
+ return [
138
+ ...parts.filter((part) => part.type !== "text"),
139
+ { type: "text" as const, text: transcript },
140
+ ];
141
+ };
142
+
143
+ export const getCustomMetadata = (
144
+ additionalKwargs: Record<string, unknown> | undefined,
145
+ ): Record<string, unknown> =>
146
+ (additionalKwargs?.metadata as Record<string, unknown>) ?? {};
147
+
148
+ export const uiMessageToDataPart = <
149
+ TUIMessage extends { name: string; props: Record<string, unknown> },
150
+ >(
151
+ ui: TUIMessage,
152
+ ): DataMessagePart => ({
153
+ type: "data",
154
+ name: ui.name,
155
+ data: ui.props,
156
+ });
157
+
158
+ /**
159
+ * Audio media types that reach a provider's audio input through the LangChain
160
+ * `audio` block. langchain-core derives OpenAI's `input_audio.format` by
161
+ * splitting `mime_type` on `/`, and that format is a wav-or-mp3 enum, so
162
+ * `audio/mpeg` passes the converter and is rejected at the provider.
163
+ */
164
+ const audioBlockMimeTypes = new Map<string, "audio/mp3" | "audio/wav">([
165
+ ["audio/mp3", "audio/mp3"],
166
+ ["audio/mpeg", "audio/mp3"],
167
+ ["audio/wav", "audio/wav"],
168
+ ["audio/wave", "audio/wav"],
169
+ ["audio/x-wav", "audio/wav"],
170
+ ]);
171
+
172
+ export const getMessageContent = (msg: AppendMessage) => {
173
+ const allContent = [
174
+ ...msg.content,
175
+ ...(msg.attachments?.flatMap((a) => a.content) ?? []),
176
+ ];
177
+
178
+ const hasNonText = allContent.some(
179
+ (part) =>
180
+ part.type === "file" || part.type === "image" || part.type === "audio",
181
+ );
182
+ const hasText = allContent.some((part) => part.type === "text");
183
+ if (hasNonText && !hasText) {
184
+ allContent.unshift({ type: "text", text: " " });
185
+ }
186
+
187
+ const content = allContent.flatMap((part) => {
188
+ const type = part.type;
189
+ switch (type) {
190
+ case "text":
191
+ return { type: "text" as const, text: part.text };
192
+ case "image":
193
+ return { type: "image_url" as const, image_url: { url: part.image } };
194
+ case "file": {
195
+ const metadata = { filename: part.filename ?? "file" };
196
+ if (part.sourceType === "id") {
197
+ return {
198
+ type: "file" as const,
199
+ id: part.data,
200
+ mime_type: part.mimeType,
201
+ filename: metadata.filename,
202
+ metadata,
203
+ source_type: "id" as const,
204
+ };
205
+ }
206
+ const source = resolveFilePartSource(part);
207
+ if (source.kind === "url") {
208
+ return {
209
+ type: "file" as const,
210
+ url: source.url,
211
+ mime_type: part.mimeType,
212
+ filename: metadata.filename,
213
+ metadata,
214
+ source_type: "url" as const,
215
+ };
216
+ }
217
+ const audioMimeType = audioBlockMimeTypes.get(
218
+ source.mimeType.toLowerCase(),
219
+ );
220
+ if (audioMimeType) {
221
+ return {
222
+ type: "audio" as const,
223
+ data: source.data,
224
+ mime_type: audioMimeType,
225
+ source_type: "base64" as const,
226
+ };
227
+ }
228
+ return {
229
+ type: "file" as const,
230
+ data: source.data,
231
+ mime_type: source.mimeType,
232
+ filename: metadata.filename,
233
+ metadata,
234
+ source_type: "base64" as const,
235
+ };
236
+ }
237
+ case "audio": {
238
+ const parsed = parseDataUrl(part.audio.data);
239
+ return {
240
+ type: "audio" as const,
241
+ data: parsed?.data ?? part.audio.data,
242
+ mime_type: `audio/${part.audio.format}`,
243
+ source_type: "base64" as const,
244
+ };
245
+ }
246
+ case "data":
247
+ return [];
248
+ case "tool-call":
249
+ throw new Error("Tool call appends are not supported.");
250
+ default: {
251
+ const _exhaustiveCheck: "reasoning" | "source" | "generative-ui" = type;
252
+ throw new Error(
253
+ `Unsupported append message part type: ${_exhaustiveCheck}`,
254
+ );
255
+ }
256
+ }
257
+ });
258
+
259
+ if (content.length === 1 && content[0]?.type === "text") {
260
+ return content[0].text ?? "";
261
+ }
262
+ return content;
263
+ };
264
+
265
+ const reasoningTextLength = (part: {
266
+ readonly summary?: ReadonlyArray<{ readonly text?: string }>;
267
+ readonly reasoning?: string;
268
+ }): number => {
269
+ if (part.summary && part.summary.length > 0)
270
+ return part.summary.map((s) => s?.text ?? "").join("\n\n\n").length;
271
+ return part.reasoning?.length ?? 0;
272
+ };
273
+
274
+ export const createLangChainStreamingTimingAccessors = <
275
+ TMessage extends {
276
+ id?: string | undefined;
277
+ content?: unknown;
278
+ tool_calls?: readonly unknown[] | undefined;
279
+ },
280
+ >(
281
+ getType: (message: TMessage) => string,
282
+ ): StreamingTimingAccessors<TMessage> => {
283
+ const findAiMessage = (
284
+ messages: readonly TMessage[],
285
+ messageId: string,
286
+ ): TMessage | undefined =>
287
+ messages.find(
288
+ (message) => getType(message) === "ai" && message.id === messageId,
289
+ );
290
+
291
+ const getTextLength = (
292
+ messages: readonly TMessage[],
293
+ messageId: string,
294
+ ): number => {
295
+ const message = findAiMessage(messages, messageId);
296
+ if (!message) return 0;
297
+ const content = message.content;
298
+ if (typeof content === "string") return content.length;
299
+ if (!Array.isArray(content)) return 0;
300
+ let len = 0;
301
+ for (const part of content as readonly LangChainContentBlock[]) {
302
+ switch (part.type) {
303
+ case "text":
304
+ case "text_delta":
305
+ if (typeof part.text === "string") len += part.text.length;
306
+ break;
307
+ case "thinking":
308
+ if (typeof part.thinking === "string") len += part.thinking.length;
309
+ break;
310
+ case "reasoning":
311
+ len += reasoningTextLength(part);
312
+ break;
313
+ }
314
+ }
315
+ return len;
316
+ };
317
+
318
+ const getToolCallCount = (
319
+ messages: readonly TMessage[],
320
+ messageId: string,
321
+ ): number => findAiMessage(messages, messageId)?.tool_calls?.length ?? 0;
322
+
323
+ const getAssistantMessageId = (
324
+ messages: readonly TMessage[],
325
+ ): string | undefined => {
326
+ for (let i = messages.length - 1; i >= 0; i--) {
327
+ const message = messages[i];
328
+ if (message && getType(message) === "ai" && message.id) return message.id;
329
+ }
330
+ return undefined;
331
+ };
332
+
333
+ return {
334
+ getAssistantMessageId,
335
+ getTextLength,
336
+ getToolCallCount,
337
+ };
338
+ };
339
+
340
+ /**
341
+ * Resolve the assistant message a `UIMessage` belongs to: the parent id comes
342
+ * from `metadata.message_id` (Python SDK) or `metadata.id` (JS SDK).
343
+ */
344
+ export const getUIMessageParentId = (ui: {
345
+ metadata?: { message_id?: string; id?: string } | undefined;
346
+ }): string | undefined => ui.metadata?.message_id ?? ui.metadata?.id;
347
+
348
+ /**
349
+ * Group the graph's accumulated `UIMessage`s by the assistant message they
350
+ * belong to. Non-array state and entries without a parent link are dropped.
351
+ */
352
+ export const groupUIMessagesByParent = <
353
+ T extends {
354
+ metadata?: { message_id?: string; id?: string } | undefined;
355
+ },
356
+ >(
357
+ value: unknown,
358
+ ): Map<string, T[]> => {
359
+ const map = new Map<string, T[]>();
360
+ if (!Array.isArray(value)) return map;
361
+ for (const ui of value as T[]) {
362
+ const parentId = getUIMessageParentId(ui);
363
+ if (!parentId) continue;
364
+ const existing = map.get(parentId);
365
+ if (existing) {
366
+ existing.push(ui);
367
+ } else {
368
+ map.set(parentId, [ui]);
369
+ }
370
+ }
371
+ return map;
372
+ };
@@ -0,0 +1,37 @@
1
+ // @vitest-environment jsdom
2
+
3
+ import { act, renderHook } from "@testing-library/react";
4
+ import { describe, expect, it } from "vitest";
5
+ import type { LangChainBaseMessage } from "./types";
6
+ import { useLangChainStreamingTiming } from "./streamingTiming";
7
+
8
+ describe("useLangChainStreamingTiming", () => {
9
+ it("counts the reasoning fallback when summary is empty", () => {
10
+ const messages: LangChainBaseMessage[] = [
11
+ {
12
+ id: "msg-1",
13
+ _getType: () => "ai",
14
+ content: [
15
+ {
16
+ type: "reasoning",
17
+ summary: [],
18
+ reasoning: "deduced",
19
+ },
20
+ ],
21
+ },
22
+ ];
23
+
24
+ const { result, rerender } = renderHook(
25
+ ({ msgs, running }) => useLangChainStreamingTiming(msgs, running),
26
+ { initialProps: { msgs: messages, running: true } },
27
+ );
28
+
29
+ act(() => {
30
+ rerender({ msgs: messages, running: false });
31
+ });
32
+
33
+ expect(result.current["msg-1"]?.tokenCount).toBe(
34
+ Math.ceil("deduced".length / 4),
35
+ );
36
+ });
37
+ });
@@ -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;
@@ -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,
@@ -191,7 +175,8 @@ const useStreamThreadRuntime = (
191
175
  const convertWithUI = useMemo<
192
176
  useExternalMessageConverter.Callback<LangChainBaseMessage>
193
177
  >(() => {
194
- const uiMessagesByParent = groupUIMessagesByParent(mergedUiMessages);
178
+ const uiMessagesByParent =
179
+ groupUIMessagesByParent<UIMessage>(mergedUiMessages);
195
180
  return (message, metadata) =>
196
181
  convertLangChainBaseMessage(message, {
197
182
  ...metadata,
@@ -409,7 +394,7 @@ const useStreamThreadRuntime = (
409
394
 
410
395
  const runtime = useExternalStoreRuntime({
411
396
  ...pickExternalStoreSharedOptions(options),
412
- isRunning: effectiveIsRunning,
397
+ isRunning: stream.isLoading,
413
398
  isLoading: stream.isThreadLoading,
414
399
  messages: threadMessages,
415
400
  adapters,
@@ -430,13 +415,7 @@ const useStreamThreadRuntime = (
430
415
  autoCancelPendingToolCalls !== false
431
416
  ? getPendingToolCalls(
432
417
  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
- }))
418
+ ).map(createToolCallCancellationStub)
440
419
  : [];
441
420
  // A null threadId is not a no-op for the SDK: it rebinds the controller
442
421
  // away from its self-created thread and forces a fresh one, so the
@@ -637,9 +616,13 @@ export const useStreamRuntime = (rawOptions: UseStreamRuntimeOptions) => {
637
616
  const optionsRef = useRef(options);
638
617
  optionsRef.current = options;
639
618
 
619
+ const aui = useAui();
640
620
  const cloudAdapter = useCloudThreadListAdapter({
641
621
  cloud,
642
- create,
622
+ create: createCloudThreadListAdapterCreateFallback(
623
+ create,
624
+ aui.threadListItem,
625
+ ),
643
626
  delete: deleteFn,
644
627
  });
645
628
  const adapter = unstable_threadListAdapter ?? cloudAdapter;