@assistant-ui/react-langchain 0.0.14 → 0.0.17
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/README.md +1 -1
- package/dist/convertMessages.d.ts +37 -3
- package/dist/convertMessages.d.ts.map +1 -1
- package/dist/convertMessages.js +63 -11
- package/dist/convertMessages.js.map +1 -1
- package/dist/findForkCheckpointInHistory.d.ts +26 -0
- package/dist/findForkCheckpointInHistory.d.ts.map +1 -0
- package/dist/findForkCheckpointInHistory.js +34 -0
- package/dist/findForkCheckpointInHistory.js.map +1 -0
- package/dist/hooks.d.ts +88 -0
- package/dist/hooks.d.ts.map +1 -0
- package/dist/hooks.js +99 -0
- package/dist/hooks.js.map +1 -0
- package/dist/index.d.ts +6 -3
- package/dist/index.js +4 -2
- package/dist/resolveForkCheckpoint.d.ts +17 -0
- package/dist/resolveForkCheckpoint.d.ts.map +1 -0
- package/dist/resolveForkCheckpoint.js +26 -0
- package/dist/resolveForkCheckpoint.js.map +1 -0
- package/dist/runtimeExtras.d.ts +7 -0
- package/dist/runtimeExtras.d.ts.map +1 -0
- package/dist/runtimeExtras.js +7 -0
- package/dist/runtimeExtras.js.map +1 -0
- package/dist/streamingTiming.d.ts +16 -0
- package/dist/streamingTiming.d.ts.map +1 -0
- package/dist/streamingTiming.js +53 -0
- package/dist/streamingTiming.js.map +1 -0
- package/dist/types.d.ts +113 -5
- package/dist/types.d.ts.map +1 -1
- package/dist/uiMessages.d.ts +24 -0
- package/dist/uiMessages.d.ts.map +1 -0
- package/dist/uiMessages.js +68 -0
- package/dist/uiMessages.js.map +1 -0
- package/dist/useStreamRuntime.d.ts +15 -76
- package/dist/useStreamRuntime.d.ts.map +1 -1
- package/dist/useStreamRuntime.js +190 -111
- package/dist/useStreamRuntime.js.map +1 -1
- package/package.json +9 -9
- package/src/__tests__/langChainTestUtils.ts +22 -0
- package/src/convertMessages.test.ts +210 -0
- package/src/convertMessages.ts +111 -11
- package/src/findForkCheckpointInHistory.test.ts +263 -0
- package/src/findForkCheckpointInHistory.ts +68 -0
- package/src/groupUIMessagesByParent.test.ts +61 -0
- package/src/hooks.ts +156 -0
- package/src/index.ts +20 -3
- package/src/resolveForkCheckpoint.test.ts +202 -0
- package/src/resolveForkCheckpoint.ts +49 -0
- package/src/runConfigToSubmitOptions.test.ts +24 -0
- package/src/runtimeExtras.ts +5 -0
- package/src/streamingTiming.test.ts +118 -0
- package/src/streamingTiming.ts +85 -0
- package/src/types.ts +147 -3
- package/src/uiMessages.test.ts +190 -0
- package/src/uiMessages.ts +91 -0
- package/src/useLangChainError.test.tsx +48 -0
- package/src/useLangChainInterrupts.test.tsx +48 -0
- package/src/useLangChainRespond.test.tsx +49 -0
- package/src/useLangChainRespondAll.test.tsx +52 -0
- package/src/useLangChainState.test.tsx +14 -26
- package/src/useLangChainStream.test.tsx +42 -0
- package/src/useLangChainSubagents.test.tsx +48 -0
- package/src/useLangChainSubgraphs.test.tsx +48 -0
- package/src/useLangChainToolCalls.test.tsx +48 -0
- package/src/useStreamRuntime.test.tsx +232 -0
- package/src/useStreamRuntime.ts +311 -234
|
@@ -0,0 +1,210 @@
|
|
|
1
|
+
import { describe, expect, it } from "vitest";
|
|
2
|
+
import { convertLangChainBaseMessage } from "./convertMessages";
|
|
3
|
+
import type { LangChainBaseMessage, UIMessage } from "./types";
|
|
4
|
+
|
|
5
|
+
const humanMessage = (content: unknown): LangChainBaseMessage => ({
|
|
6
|
+
_getType: () => "human",
|
|
7
|
+
id: "msg-1",
|
|
8
|
+
content,
|
|
9
|
+
});
|
|
10
|
+
|
|
11
|
+
const aiMessage = (content: unknown): LangChainBaseMessage => ({
|
|
12
|
+
_getType: () => "ai",
|
|
13
|
+
id: "msg-2",
|
|
14
|
+
content,
|
|
15
|
+
});
|
|
16
|
+
|
|
17
|
+
const contentOf = (result: ReturnType<typeof convertLangChainBaseMessage>) => {
|
|
18
|
+
if (result.role === "tool")
|
|
19
|
+
throw new Error("expected a content-bearing message");
|
|
20
|
+
return result.content;
|
|
21
|
+
};
|
|
22
|
+
|
|
23
|
+
describe("convertLangChainBaseMessage file content parts", () => {
|
|
24
|
+
it("converts a base64 file block", () => {
|
|
25
|
+
const result = convertLangChainBaseMessage(
|
|
26
|
+
humanMessage([
|
|
27
|
+
{
|
|
28
|
+
type: "file",
|
|
29
|
+
data: "ZmFrZQ==",
|
|
30
|
+
mime_type: "application/pdf",
|
|
31
|
+
source_type: "base64",
|
|
32
|
+
metadata: { filename: "a.pdf" },
|
|
33
|
+
},
|
|
34
|
+
]),
|
|
35
|
+
{},
|
|
36
|
+
);
|
|
37
|
+
|
|
38
|
+
expect(contentOf(result)).toEqual([
|
|
39
|
+
{
|
|
40
|
+
type: "file",
|
|
41
|
+
filename: "a.pdf",
|
|
42
|
+
data: "ZmFrZQ==",
|
|
43
|
+
mimeType: "application/pdf",
|
|
44
|
+
},
|
|
45
|
+
]);
|
|
46
|
+
});
|
|
47
|
+
|
|
48
|
+
it("falls back to a default filename when metadata is absent", () => {
|
|
49
|
+
const result = convertLangChainBaseMessage(
|
|
50
|
+
humanMessage([
|
|
51
|
+
{ type: "file", data: "ZmFrZQ==", mime_type: "application/pdf" },
|
|
52
|
+
]),
|
|
53
|
+
{},
|
|
54
|
+
);
|
|
55
|
+
|
|
56
|
+
expect(contentOf(result)).toEqual([
|
|
57
|
+
{
|
|
58
|
+
type: "file",
|
|
59
|
+
filename: "file",
|
|
60
|
+
data: "ZmFrZQ==",
|
|
61
|
+
mimeType: "application/pdf",
|
|
62
|
+
},
|
|
63
|
+
]);
|
|
64
|
+
});
|
|
65
|
+
});
|
|
66
|
+
|
|
67
|
+
describe("convertLangChainBaseMessage reasoning content parts", () => {
|
|
68
|
+
it("joins summary parts into a single reasoning part", () => {
|
|
69
|
+
const result = convertLangChainBaseMessage(
|
|
70
|
+
aiMessage([
|
|
71
|
+
{
|
|
72
|
+
type: "reasoning",
|
|
73
|
+
summary: [
|
|
74
|
+
{ type: "summary_text", text: "first" },
|
|
75
|
+
{ type: "summary_text", text: "second" },
|
|
76
|
+
],
|
|
77
|
+
},
|
|
78
|
+
]),
|
|
79
|
+
{},
|
|
80
|
+
);
|
|
81
|
+
|
|
82
|
+
expect(contentOf(result)).toEqual([
|
|
83
|
+
{ type: "reasoning", text: "first\n\n\nsecond" },
|
|
84
|
+
]);
|
|
85
|
+
});
|
|
86
|
+
|
|
87
|
+
it("falls back to the reasoning string when summary is absent", () => {
|
|
88
|
+
const result = convertLangChainBaseMessage(
|
|
89
|
+
aiMessage([{ type: "reasoning", reasoning: "thinking out loud" }]),
|
|
90
|
+
{},
|
|
91
|
+
);
|
|
92
|
+
|
|
93
|
+
expect(contentOf(result)).toEqual([
|
|
94
|
+
{ type: "reasoning", text: "thinking out loud" },
|
|
95
|
+
]);
|
|
96
|
+
});
|
|
97
|
+
|
|
98
|
+
it("does not throw when a reasoning block omits both summary and reasoning", () => {
|
|
99
|
+
const result = convertLangChainBaseMessage(
|
|
100
|
+
aiMessage([{ type: "reasoning" }]),
|
|
101
|
+
{},
|
|
102
|
+
);
|
|
103
|
+
|
|
104
|
+
expect(contentOf(result)).toEqual([{ type: "reasoning", text: "" }]);
|
|
105
|
+
});
|
|
106
|
+
|
|
107
|
+
it("tolerates null entries inside the summary array", () => {
|
|
108
|
+
const result = convertLangChainBaseMessage(
|
|
109
|
+
aiMessage([
|
|
110
|
+
{
|
|
111
|
+
type: "reasoning",
|
|
112
|
+
summary: [null, { type: "summary_text", text: "kept" }],
|
|
113
|
+
},
|
|
114
|
+
]),
|
|
115
|
+
{},
|
|
116
|
+
);
|
|
117
|
+
|
|
118
|
+
expect(contentOf(result)).toEqual([
|
|
119
|
+
{ type: "reasoning", text: "\n\n\nkept" },
|
|
120
|
+
]);
|
|
121
|
+
});
|
|
122
|
+
});
|
|
123
|
+
|
|
124
|
+
describe("convertLangChainBaseMessage generative UI from graph state", () => {
|
|
125
|
+
const uiMessage = (messageId: string): UIMessage => ({
|
|
126
|
+
type: "ui",
|
|
127
|
+
id: "ui-1",
|
|
128
|
+
name: "chart",
|
|
129
|
+
props: { points: [1, 2, 3] },
|
|
130
|
+
metadata: { message_id: messageId },
|
|
131
|
+
});
|
|
132
|
+
|
|
133
|
+
it("appends a data part for UI attached to the assistant message", () => {
|
|
134
|
+
const result = convertLangChainBaseMessage(aiMessage("hello"), {
|
|
135
|
+
uiMessagesByParent: new Map([["msg-2", [uiMessage("msg-2")]]]),
|
|
136
|
+
});
|
|
137
|
+
|
|
138
|
+
expect(contentOf(result)).toEqual([
|
|
139
|
+
{ type: "text", text: "hello" },
|
|
140
|
+
{ type: "data", name: "chart", data: { points: [1, 2, 3] } },
|
|
141
|
+
]);
|
|
142
|
+
});
|
|
143
|
+
|
|
144
|
+
it("leaves the message unchanged when no UI targets it", () => {
|
|
145
|
+
const result = convertLangChainBaseMessage(aiMessage("hello"), {
|
|
146
|
+
uiMessagesByParent: new Map([["other-id", [uiMessage("other-id")]]]),
|
|
147
|
+
});
|
|
148
|
+
|
|
149
|
+
expect(contentOf(result)).toEqual([{ type: "text", text: "hello" }]);
|
|
150
|
+
});
|
|
151
|
+
|
|
152
|
+
it("appends a data part per UI when several target the same message", () => {
|
|
153
|
+
const result = convertLangChainBaseMessage(aiMessage("hello"), {
|
|
154
|
+
uiMessagesByParent: new Map([
|
|
155
|
+
[
|
|
156
|
+
"msg-2",
|
|
157
|
+
[
|
|
158
|
+
{ ...uiMessage("msg-2"), name: "chart", props: { a: 1 } },
|
|
159
|
+
{ ...uiMessage("msg-2"), name: "table", props: { b: 2 } },
|
|
160
|
+
],
|
|
161
|
+
],
|
|
162
|
+
]),
|
|
163
|
+
});
|
|
164
|
+
|
|
165
|
+
expect(contentOf(result)).toEqual([
|
|
166
|
+
{ type: "text", text: "hello" },
|
|
167
|
+
{ type: "data", name: "chart", data: { a: 1 } },
|
|
168
|
+
{ type: "data", name: "table", data: { b: 2 } },
|
|
169
|
+
]);
|
|
170
|
+
});
|
|
171
|
+
|
|
172
|
+
it("does not attach UI to an assistant message without an id", () => {
|
|
173
|
+
const result = convertLangChainBaseMessage(
|
|
174
|
+
{ ...aiMessage("hello"), id: undefined },
|
|
175
|
+
{ uiMessagesByParent: new Map([["", [uiMessage("")]]]) },
|
|
176
|
+
);
|
|
177
|
+
|
|
178
|
+
expect(contentOf(result)).toEqual([{ type: "text", text: "hello" }]);
|
|
179
|
+
});
|
|
180
|
+
|
|
181
|
+
it("leaves the message unchanged when no converter metadata is passed", () => {
|
|
182
|
+
const result = convertLangChainBaseMessage(aiMessage("hello"));
|
|
183
|
+
|
|
184
|
+
expect(contentOf(result)).toEqual([{ type: "text", text: "hello" }]);
|
|
185
|
+
});
|
|
186
|
+
});
|
|
187
|
+
|
|
188
|
+
describe("convertLangChainBaseMessage image content parts", () => {
|
|
189
|
+
it("reads the url from an image_url object", () => {
|
|
190
|
+
const result = convertLangChainBaseMessage(
|
|
191
|
+
humanMessage([
|
|
192
|
+
{ type: "image_url", image_url: { url: "https://example.com/a.png" } },
|
|
193
|
+
]),
|
|
194
|
+
{},
|
|
195
|
+
);
|
|
196
|
+
|
|
197
|
+
expect(contentOf(result)).toEqual([
|
|
198
|
+
{ type: "image", image: "https://example.com/a.png" },
|
|
199
|
+
]);
|
|
200
|
+
});
|
|
201
|
+
|
|
202
|
+
it("drops the image part when image_url is undefined", () => {
|
|
203
|
+
const result = convertLangChainBaseMessage(
|
|
204
|
+
humanMessage([{ type: "image_url" }]),
|
|
205
|
+
{},
|
|
206
|
+
);
|
|
207
|
+
|
|
208
|
+
expect(contentOf(result)).toEqual([]);
|
|
209
|
+
});
|
|
210
|
+
});
|
package/src/convertMessages.ts
CHANGED
|
@@ -1,8 +1,29 @@
|
|
|
1
1
|
"use client";
|
|
2
2
|
|
|
3
3
|
import type { useExternalMessageConverter } from "@assistant-ui/core/react";
|
|
4
|
+
import type {
|
|
5
|
+
AppendMessage,
|
|
6
|
+
DataMessagePart,
|
|
7
|
+
MessageTiming,
|
|
8
|
+
} from "@assistant-ui/core";
|
|
4
9
|
import type { ReadonlyJSONObject } from "assistant-stream/utils";
|
|
5
|
-
import type {
|
|
10
|
+
import type {
|
|
11
|
+
LangChainBaseMessage,
|
|
12
|
+
LangChainContentBlock,
|
|
13
|
+
UIMessage,
|
|
14
|
+
} from "./types";
|
|
15
|
+
|
|
16
|
+
type LangChainMessageConverterMetadata =
|
|
17
|
+
useExternalMessageConverter.Metadata & {
|
|
18
|
+
uiMessagesByParent?: Map<string, UIMessage[]>;
|
|
19
|
+
messageTiming?: Record<string, MessageTiming>;
|
|
20
|
+
};
|
|
21
|
+
|
|
22
|
+
const uiMessageToDataPart = (ui: UIMessage): DataMessagePart => ({
|
|
23
|
+
type: "data",
|
|
24
|
+
name: ui.name,
|
|
25
|
+
data: ui.props,
|
|
26
|
+
});
|
|
6
27
|
|
|
7
28
|
export const getMessageType = (message: LangChainBaseMessage): string => {
|
|
8
29
|
if (typeof message._getType === "function") return message._getType();
|
|
@@ -23,17 +44,30 @@ const contentToParts = (content: unknown) => {
|
|
|
23
44
|
case "text":
|
|
24
45
|
case "text_delta":
|
|
25
46
|
return { type: "text" as const, text: part.text };
|
|
26
|
-
case "image_url":
|
|
27
|
-
|
|
28
|
-
|
|
29
|
-
|
|
30
|
-
|
|
47
|
+
case "image_url": {
|
|
48
|
+
const image =
|
|
49
|
+
typeof part.image_url === "string"
|
|
50
|
+
? part.image_url
|
|
51
|
+
: part.image_url?.url;
|
|
52
|
+
if (!image) return null;
|
|
53
|
+
return { type: "image" as const, image };
|
|
54
|
+
}
|
|
55
|
+
case "file":
|
|
56
|
+
return {
|
|
57
|
+
type: "file" as const,
|
|
58
|
+
filename: part.metadata?.filename ?? "file",
|
|
59
|
+
data: part.data,
|
|
60
|
+
mimeType: part.mime_type,
|
|
61
|
+
};
|
|
31
62
|
case "thinking":
|
|
32
63
|
return { type: "reasoning" as const, text: part.thinking };
|
|
33
64
|
case "reasoning":
|
|
34
65
|
return {
|
|
35
66
|
type: "reasoning" as const,
|
|
36
|
-
text:
|
|
67
|
+
text:
|
|
68
|
+
part.summary?.map((s) => s?.text ?? "").join("\n\n\n") ??
|
|
69
|
+
part.reasoning ??
|
|
70
|
+
"",
|
|
37
71
|
};
|
|
38
72
|
case "tool_use":
|
|
39
73
|
case "input_json_delta":
|
|
@@ -59,9 +93,10 @@ const getStringContent = (content: unknown): string => {
|
|
|
59
93
|
.join("");
|
|
60
94
|
};
|
|
61
95
|
|
|
62
|
-
export const convertLangChainBaseMessage
|
|
63
|
-
LangChainBaseMessage
|
|
64
|
-
|
|
96
|
+
export const convertLangChainBaseMessage = (
|
|
97
|
+
message: LangChainBaseMessage,
|
|
98
|
+
metadata: LangChainMessageConverterMetadata = {},
|
|
99
|
+
): useExternalMessageConverter.Message => {
|
|
65
100
|
const type = getMessageType(message);
|
|
66
101
|
|
|
67
102
|
switch (type) {
|
|
@@ -98,12 +133,26 @@ export const convertLangChainBaseMessage: useExternalMessageConverter.Callback<
|
|
|
98
133
|
const assistantStatus =
|
|
99
134
|
typeof message.status === "object" ? message.status : undefined;
|
|
100
135
|
|
|
136
|
+
const uiDataParts =
|
|
137
|
+
(message.id
|
|
138
|
+
? metadata.uiMessagesByParent
|
|
139
|
+
?.get(message.id)
|
|
140
|
+
?.map(uiMessageToDataPart)
|
|
141
|
+
: undefined) ?? [];
|
|
142
|
+
|
|
143
|
+
const timing = metadata.messageTiming?.[message.id ?? ""];
|
|
144
|
+
|
|
101
145
|
return {
|
|
102
146
|
role: "assistant",
|
|
103
147
|
id: message.id,
|
|
104
|
-
content: [
|
|
148
|
+
content: [
|
|
149
|
+
...contentToParts(message.content),
|
|
150
|
+
...toolCallParts,
|
|
151
|
+
...uiDataParts,
|
|
152
|
+
],
|
|
105
153
|
metadata: {
|
|
106
154
|
custom: getCustomMetadata(message.additional_kwargs),
|
|
155
|
+
...(timing && { timing }),
|
|
107
156
|
},
|
|
108
157
|
...(assistantStatus && { status: assistantStatus }),
|
|
109
158
|
};
|
|
@@ -135,3 +184,54 @@ export const convertLangChainBaseMessage: useExternalMessageConverter.Callback<
|
|
|
135
184
|
};
|
|
136
185
|
}
|
|
137
186
|
};
|
|
187
|
+
|
|
188
|
+
export const getMessageContent = (msg: AppendMessage) => {
|
|
189
|
+
const allContent = [
|
|
190
|
+
...msg.content,
|
|
191
|
+
...(msg.attachments?.flatMap((a) => a.content) ?? []),
|
|
192
|
+
];
|
|
193
|
+
|
|
194
|
+
const hasNonText = allContent.some(
|
|
195
|
+
(part) => part.type === "file" || part.type === "image",
|
|
196
|
+
);
|
|
197
|
+
const hasText = allContent.some((part) => part.type === "text");
|
|
198
|
+
if (hasNonText && !hasText) {
|
|
199
|
+
allContent.unshift({ type: "text", text: " " });
|
|
200
|
+
}
|
|
201
|
+
|
|
202
|
+
const content = allContent.map((part) => {
|
|
203
|
+
const type = part.type;
|
|
204
|
+
switch (type) {
|
|
205
|
+
case "text":
|
|
206
|
+
return { type: "text" as const, text: part.text };
|
|
207
|
+
case "image":
|
|
208
|
+
return { type: "image_url" as const, image_url: { url: part.image } };
|
|
209
|
+
case "file":
|
|
210
|
+
return {
|
|
211
|
+
type: "file" as const,
|
|
212
|
+
data: part.data,
|
|
213
|
+
mime_type: part.mimeType,
|
|
214
|
+
metadata: { filename: part.filename ?? "file" },
|
|
215
|
+
source_type: "base64" as const,
|
|
216
|
+
};
|
|
217
|
+
case "tool-call":
|
|
218
|
+
throw new Error("Tool call appends are not supported.");
|
|
219
|
+
default: {
|
|
220
|
+
const _exhaustiveCheck:
|
|
221
|
+
| "reasoning"
|
|
222
|
+
| "source"
|
|
223
|
+
| "audio"
|
|
224
|
+
| "data"
|
|
225
|
+
| "generative-ui" = type;
|
|
226
|
+
throw new Error(
|
|
227
|
+
`Unsupported append message part type: ${_exhaustiveCheck}`,
|
|
228
|
+
);
|
|
229
|
+
}
|
|
230
|
+
}
|
|
231
|
+
});
|
|
232
|
+
|
|
233
|
+
if (content.length === 1 && content[0]?.type === "text") {
|
|
234
|
+
return content[0].text ?? "";
|
|
235
|
+
}
|
|
236
|
+
return content;
|
|
237
|
+
};
|
|
@@ -0,0 +1,263 @@
|
|
|
1
|
+
import { describe, expect, it, vi } from "vitest";
|
|
2
|
+
import type { LangChainBaseMessage } from "./types";
|
|
3
|
+
import { findForkCheckpointInHistory } from "./findForkCheckpointInHistory";
|
|
4
|
+
|
|
5
|
+
const msg = (id: string | undefined): LangChainBaseMessage => ({
|
|
6
|
+
_getType: () => "human",
|
|
7
|
+
content: "",
|
|
8
|
+
id,
|
|
9
|
+
});
|
|
10
|
+
|
|
11
|
+
const makeClient = (history: unknown[]) => ({
|
|
12
|
+
threads: {
|
|
13
|
+
getHistory: vi.fn(async () => history as never),
|
|
14
|
+
},
|
|
15
|
+
});
|
|
16
|
+
|
|
17
|
+
type HistoryState = {
|
|
18
|
+
values: Record<string, unknown>;
|
|
19
|
+
checkpoint: { checkpoint_id?: string };
|
|
20
|
+
};
|
|
21
|
+
|
|
22
|
+
const makePaginatedClient = (states: HistoryState[]) => ({
|
|
23
|
+
threads: {
|
|
24
|
+
getHistory: vi.fn(
|
|
25
|
+
async (
|
|
26
|
+
_threadId: string,
|
|
27
|
+
options?: {
|
|
28
|
+
limit?: number;
|
|
29
|
+
before?: { configurable: { checkpoint_id: string } };
|
|
30
|
+
},
|
|
31
|
+
) => {
|
|
32
|
+
const limit = options?.limit ?? 10;
|
|
33
|
+
const beforeId = options?.before?.configurable.checkpoint_id;
|
|
34
|
+
const start = beforeId
|
|
35
|
+
? states.findIndex((s) => s.checkpoint.checkpoint_id === beforeId) + 1
|
|
36
|
+
: 0;
|
|
37
|
+
return states.slice(start, start + limit) as never;
|
|
38
|
+
},
|
|
39
|
+
),
|
|
40
|
+
},
|
|
41
|
+
});
|
|
42
|
+
|
|
43
|
+
describe("findForkCheckpointInHistory", () => {
|
|
44
|
+
it("returns the checkpoint_id of the state whose messages match by id", async () => {
|
|
45
|
+
const client = makeClient([
|
|
46
|
+
{
|
|
47
|
+
values: { messages: [msg("a"), msg("b"), msg("c")] },
|
|
48
|
+
checkpoint: { checkpoint_id: "cp-too-long" },
|
|
49
|
+
},
|
|
50
|
+
{
|
|
51
|
+
values: { messages: [msg("a"), msg("b")] },
|
|
52
|
+
checkpoint: { checkpoint_id: "cp-match" },
|
|
53
|
+
},
|
|
54
|
+
{
|
|
55
|
+
values: { messages: [msg("a")] },
|
|
56
|
+
checkpoint: { checkpoint_id: "cp-too-short" },
|
|
57
|
+
},
|
|
58
|
+
]);
|
|
59
|
+
|
|
60
|
+
const result = await findForkCheckpointInHistory(
|
|
61
|
+
client,
|
|
62
|
+
"thread-1",
|
|
63
|
+
[msg("a"), msg("b")],
|
|
64
|
+
"messages",
|
|
65
|
+
);
|
|
66
|
+
|
|
67
|
+
expect(result).toBe("cp-match");
|
|
68
|
+
expect(client.threads.getHistory).toHaveBeenCalledWith("thread-1", {
|
|
69
|
+
limit: 100,
|
|
70
|
+
});
|
|
71
|
+
});
|
|
72
|
+
|
|
73
|
+
it("returns null when no state has the same message ids", async () => {
|
|
74
|
+
const client = makeClient([
|
|
75
|
+
{
|
|
76
|
+
values: { messages: [msg("a"), msg("x")] },
|
|
77
|
+
checkpoint: { checkpoint_id: "cp-1" },
|
|
78
|
+
},
|
|
79
|
+
]);
|
|
80
|
+
|
|
81
|
+
const result = await findForkCheckpointInHistory(
|
|
82
|
+
client,
|
|
83
|
+
"thread-1",
|
|
84
|
+
[msg("a"), msg("b")],
|
|
85
|
+
"messages",
|
|
86
|
+
);
|
|
87
|
+
|
|
88
|
+
expect(result).toBeNull();
|
|
89
|
+
});
|
|
90
|
+
|
|
91
|
+
it("returns null when message ids are unstable (missing)", async () => {
|
|
92
|
+
const client = makeClient([
|
|
93
|
+
{
|
|
94
|
+
values: { messages: [msg("a"), msg(undefined)] },
|
|
95
|
+
checkpoint: { checkpoint_id: "cp-1" },
|
|
96
|
+
},
|
|
97
|
+
]);
|
|
98
|
+
|
|
99
|
+
const result = await findForkCheckpointInHistory(
|
|
100
|
+
client,
|
|
101
|
+
"thread-1",
|
|
102
|
+
[msg("a"), msg(undefined)],
|
|
103
|
+
"messages",
|
|
104
|
+
);
|
|
105
|
+
|
|
106
|
+
expect(result).toBeNull();
|
|
107
|
+
});
|
|
108
|
+
|
|
109
|
+
it("reads messages from a custom messagesKey", async () => {
|
|
110
|
+
const client = makeClient([
|
|
111
|
+
{
|
|
112
|
+
values: { history: [msg("a")] },
|
|
113
|
+
checkpoint: { checkpoint_id: "cp-custom" },
|
|
114
|
+
},
|
|
115
|
+
]);
|
|
116
|
+
|
|
117
|
+
const result = await findForkCheckpointInHistory(
|
|
118
|
+
client,
|
|
119
|
+
"thread-1",
|
|
120
|
+
[msg("a")],
|
|
121
|
+
"history",
|
|
122
|
+
);
|
|
123
|
+
|
|
124
|
+
expect(result).toBe("cp-custom");
|
|
125
|
+
});
|
|
126
|
+
|
|
127
|
+
it("returns null when the matching checkpoint has no checkpoint_id", async () => {
|
|
128
|
+
const client = makeClient([
|
|
129
|
+
{
|
|
130
|
+
values: { messages: [msg("a")] },
|
|
131
|
+
checkpoint: {},
|
|
132
|
+
},
|
|
133
|
+
]);
|
|
134
|
+
|
|
135
|
+
const result = await findForkCheckpointInHistory(
|
|
136
|
+
client,
|
|
137
|
+
"thread-1",
|
|
138
|
+
[msg("a")],
|
|
139
|
+
"messages",
|
|
140
|
+
);
|
|
141
|
+
|
|
142
|
+
expect(result).toBeNull();
|
|
143
|
+
});
|
|
144
|
+
|
|
145
|
+
it("keeps scanning when a matching state has no checkpoint_id", async () => {
|
|
146
|
+
const client = makeClient([
|
|
147
|
+
{
|
|
148
|
+
values: { messages: [msg("a")] },
|
|
149
|
+
checkpoint: {},
|
|
150
|
+
},
|
|
151
|
+
{
|
|
152
|
+
values: { messages: [msg("a")] },
|
|
153
|
+
checkpoint: { checkpoint_id: "cp-second" },
|
|
154
|
+
},
|
|
155
|
+
]);
|
|
156
|
+
|
|
157
|
+
const result = await findForkCheckpointInHistory(
|
|
158
|
+
client,
|
|
159
|
+
"thread-1",
|
|
160
|
+
[msg("a")],
|
|
161
|
+
"messages",
|
|
162
|
+
);
|
|
163
|
+
|
|
164
|
+
expect(result).toBe("cp-second");
|
|
165
|
+
});
|
|
166
|
+
|
|
167
|
+
it("does not throw when a state's messages payload is not an array", async () => {
|
|
168
|
+
const client = makeClient([
|
|
169
|
+
{
|
|
170
|
+
values: { messages: "ab" },
|
|
171
|
+
checkpoint: { checkpoint_id: "cp-1" },
|
|
172
|
+
},
|
|
173
|
+
]);
|
|
174
|
+
|
|
175
|
+
const result = await findForkCheckpointInHistory(
|
|
176
|
+
client,
|
|
177
|
+
"thread-1",
|
|
178
|
+
[msg("a"), msg("b")],
|
|
179
|
+
"messages",
|
|
180
|
+
);
|
|
181
|
+
|
|
182
|
+
expect(result).toBeNull();
|
|
183
|
+
});
|
|
184
|
+
|
|
185
|
+
it("forwards a custom history limit to getHistory", async () => {
|
|
186
|
+
const client = makeClient([]);
|
|
187
|
+
|
|
188
|
+
await findForkCheckpointInHistory(
|
|
189
|
+
client,
|
|
190
|
+
"thread-1",
|
|
191
|
+
[msg("a")],
|
|
192
|
+
"messages",
|
|
193
|
+
25,
|
|
194
|
+
);
|
|
195
|
+
|
|
196
|
+
expect(client.threads.getHistory).toHaveBeenCalledWith("thread-1", {
|
|
197
|
+
limit: 25,
|
|
198
|
+
});
|
|
199
|
+
});
|
|
200
|
+
|
|
201
|
+
it("pages through history to find a match beyond the first page", async () => {
|
|
202
|
+
const client = makePaginatedClient([
|
|
203
|
+
{
|
|
204
|
+
values: { messages: [msg("a"), msg("b"), msg("c"), msg("d")] },
|
|
205
|
+
checkpoint: { checkpoint_id: "cp-4" },
|
|
206
|
+
},
|
|
207
|
+
{
|
|
208
|
+
values: { messages: [msg("a"), msg("b"), msg("c")] },
|
|
209
|
+
checkpoint: { checkpoint_id: "cp-3" },
|
|
210
|
+
},
|
|
211
|
+
{
|
|
212
|
+
values: { messages: [msg("a"), msg("b")] },
|
|
213
|
+
checkpoint: { checkpoint_id: "cp-2" },
|
|
214
|
+
},
|
|
215
|
+
{
|
|
216
|
+
values: { messages: [msg("a")] },
|
|
217
|
+
checkpoint: { checkpoint_id: "cp-1" },
|
|
218
|
+
},
|
|
219
|
+
]);
|
|
220
|
+
|
|
221
|
+
const result = await findForkCheckpointInHistory(
|
|
222
|
+
client,
|
|
223
|
+
"thread-1",
|
|
224
|
+
[msg("a")],
|
|
225
|
+
"messages",
|
|
226
|
+
2,
|
|
227
|
+
);
|
|
228
|
+
|
|
229
|
+
expect(result).toBe("cp-1");
|
|
230
|
+
expect(client.threads.getHistory).toHaveBeenCalledTimes(2);
|
|
231
|
+
expect(client.threads.getHistory).toHaveBeenLastCalledWith("thread-1", {
|
|
232
|
+
limit: 2,
|
|
233
|
+
before: { configurable: { checkpoint_id: "cp-3" } },
|
|
234
|
+
});
|
|
235
|
+
});
|
|
236
|
+
|
|
237
|
+
it("stops instead of looping when the cursor does not advance", async () => {
|
|
238
|
+
const page = [
|
|
239
|
+
{
|
|
240
|
+
values: { messages: [msg("x"), msg("y")] },
|
|
241
|
+
checkpoint: { checkpoint_id: "cp-1" },
|
|
242
|
+
},
|
|
243
|
+
{
|
|
244
|
+
values: { messages: [msg("x")] },
|
|
245
|
+
checkpoint: { checkpoint_id: "cp-2" },
|
|
246
|
+
},
|
|
247
|
+
];
|
|
248
|
+
const client = {
|
|
249
|
+
threads: { getHistory: vi.fn(async () => page as never) },
|
|
250
|
+
};
|
|
251
|
+
|
|
252
|
+
const result = await findForkCheckpointInHistory(
|
|
253
|
+
client,
|
|
254
|
+
"thread-1",
|
|
255
|
+
[msg("a")],
|
|
256
|
+
"messages",
|
|
257
|
+
2,
|
|
258
|
+
);
|
|
259
|
+
|
|
260
|
+
expect(result).toBeNull();
|
|
261
|
+
expect(client.threads.getHistory).toHaveBeenCalledTimes(2);
|
|
262
|
+
});
|
|
263
|
+
});
|
|
@@ -0,0 +1,68 @@
|
|
|
1
|
+
import type { LangChainBaseMessage } from "./types";
|
|
2
|
+
|
|
3
|
+
type ForkCheckpointHistoryState = {
|
|
4
|
+
values?: Record<string, unknown> | undefined;
|
|
5
|
+
checkpoint?: { checkpoint_id?: string | null | undefined } | undefined;
|
|
6
|
+
};
|
|
7
|
+
|
|
8
|
+
type ForkHistoryCursor = { configurable: { checkpoint_id: string } };
|
|
9
|
+
|
|
10
|
+
export type ForkCheckpointClient = {
|
|
11
|
+
threads: {
|
|
12
|
+
getHistory: (
|
|
13
|
+
threadId: string,
|
|
14
|
+
options?: { limit?: number; before?: ForkHistoryCursor },
|
|
15
|
+
) => Promise<readonly ForkCheckpointHistoryState[]>;
|
|
16
|
+
};
|
|
17
|
+
};
|
|
18
|
+
|
|
19
|
+
export const findForkCheckpointInHistory = async (
|
|
20
|
+
client: ForkCheckpointClient,
|
|
21
|
+
threadId: string,
|
|
22
|
+
messagesUpToParent: readonly LangChainBaseMessage[],
|
|
23
|
+
messagesKey: string,
|
|
24
|
+
pageSize = 100,
|
|
25
|
+
): Promise<string | null> => {
|
|
26
|
+
if (!messagesUpToParent.every((m) => typeof m.id === "string")) return null;
|
|
27
|
+
|
|
28
|
+
const scanned = new Set<string>();
|
|
29
|
+
let before: ForkHistoryCursor | undefined;
|
|
30
|
+
|
|
31
|
+
for (;;) {
|
|
32
|
+
const page = await client.threads.getHistory(
|
|
33
|
+
threadId,
|
|
34
|
+
before ? { limit: pageSize, before } : { limit: pageSize },
|
|
35
|
+
);
|
|
36
|
+
|
|
37
|
+
let oldestCheckpointId: string | undefined;
|
|
38
|
+
let advanced = false;
|
|
39
|
+
for (const state of page) {
|
|
40
|
+
const checkpointId = state.checkpoint?.checkpoint_id;
|
|
41
|
+
if (typeof checkpointId === "string") {
|
|
42
|
+
oldestCheckpointId = checkpointId;
|
|
43
|
+
if (!scanned.has(checkpointId)) {
|
|
44
|
+
scanned.add(checkpointId);
|
|
45
|
+
advanced = true;
|
|
46
|
+
}
|
|
47
|
+
}
|
|
48
|
+
const stateMessages = state.values?.[messagesKey] as
|
|
49
|
+
| readonly LangChainBaseMessage[]
|
|
50
|
+
| undefined;
|
|
51
|
+
if (
|
|
52
|
+
!Array.isArray(stateMessages) ||
|
|
53
|
+
stateMessages.length !== messagesUpToParent.length
|
|
54
|
+
)
|
|
55
|
+
continue;
|
|
56
|
+
if (!stateMessages.every((m) => typeof m.id === "string")) continue;
|
|
57
|
+
const isMatch = messagesUpToParent.every(
|
|
58
|
+
(m, i) => m.id === stateMessages[i]?.id,
|
|
59
|
+
);
|
|
60
|
+
if (isMatch && checkpointId) return checkpointId;
|
|
61
|
+
}
|
|
62
|
+
|
|
63
|
+
// `advanced` halts paging if the backend ignores `before` and repeats a page
|
|
64
|
+
if (page.length < pageSize || oldestCheckpointId === undefined || !advanced)
|
|
65
|
+
return null;
|
|
66
|
+
before = { configurable: { checkpoint_id: oldestCheckpointId } };
|
|
67
|
+
}
|
|
68
|
+
};
|