@assistant-ui/react-langchain 0.0.15 → 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/dist/convertMessages.d.ts +37 -3
- package/dist/convertMessages.d.ts.map +1 -1
- package/dist/convertMessages.js +57 -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 +105 -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 +10 -85
- package/dist/useStreamRuntime.d.ts.map +1 -1
- package/dist/useStreamRuntime.js +186 -135
- package/dist/useStreamRuntime.js.map +1 -1
- package/package.json +9 -9
- package/src/__tests__/langChainTestUtils.ts +3 -3
- package/src/convertMessages.test.ts +160 -3
- package/src/convertMessages.ts +104 -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 +18 -3
- package/src/resolveForkCheckpoint.test.ts +202 -0
- package/src/resolveForkCheckpoint.ts +49 -0
- package/src/runtimeExtras.ts +5 -0
- package/src/streamingTiming.test.ts +118 -0
- package/src/streamingTiming.ts +85 -0
- package/src/types.ts +140 -3
- package/src/uiMessages.test.ts +190 -0
- package/src/uiMessages.ts +91 -0
- package/src/useLangChainError.test.tsx +1 -1
- 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 +1 -1
- 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 +1 -1
- package/src/useStreamRuntime.test.tsx +232 -0
- package/src/useStreamRuntime.ts +301 -268
|
@@ -0,0 +1,61 @@
|
|
|
1
|
+
import { describe, expect, it } from "vitest";
|
|
2
|
+
import { groupUIMessagesByParent } from "./useStreamRuntime";
|
|
3
|
+
import type { UIMessage } from "./types";
|
|
4
|
+
|
|
5
|
+
const uiMessage = (
|
|
6
|
+
metadata: UIMessage["metadata"],
|
|
7
|
+
name = "chart",
|
|
8
|
+
): UIMessage => ({
|
|
9
|
+
type: "ui",
|
|
10
|
+
id: "ui-1",
|
|
11
|
+
name,
|
|
12
|
+
props: { points: [1, 2, 3] },
|
|
13
|
+
...(metadata !== undefined && { metadata }),
|
|
14
|
+
});
|
|
15
|
+
|
|
16
|
+
describe("groupUIMessagesByParent", () => {
|
|
17
|
+
it("returns an empty map for non-array state", () => {
|
|
18
|
+
expect(groupUIMessagesByParent(undefined).size).toBe(0);
|
|
19
|
+
expect(groupUIMessagesByParent(null).size).toBe(0);
|
|
20
|
+
expect(groupUIMessagesByParent({ ui: [] }).size).toBe(0);
|
|
21
|
+
});
|
|
22
|
+
|
|
23
|
+
it("groups by metadata.message_id (Python SDK)", () => {
|
|
24
|
+
const ui = uiMessage({ message_id: "msg-2" });
|
|
25
|
+
const map = groupUIMessagesByParent([ui]);
|
|
26
|
+
|
|
27
|
+
expect(map.get("msg-2")).toEqual([ui]);
|
|
28
|
+
});
|
|
29
|
+
|
|
30
|
+
it("falls back to metadata.id when message_id is absent (JS SDK)", () => {
|
|
31
|
+
const ui = uiMessage({ id: "msg-2" });
|
|
32
|
+
const map = groupUIMessagesByParent([ui]);
|
|
33
|
+
|
|
34
|
+
expect(map.get("msg-2")).toEqual([ui]);
|
|
35
|
+
});
|
|
36
|
+
|
|
37
|
+
it("prefers message_id over id when both are present", () => {
|
|
38
|
+
const ui = uiMessage({ message_id: "msg-2", id: "msg-9" });
|
|
39
|
+
const map = groupUIMessagesByParent([ui]);
|
|
40
|
+
|
|
41
|
+
expect(map.get("msg-2")).toEqual([ui]);
|
|
42
|
+
expect(map.has("msg-9")).toBe(false);
|
|
43
|
+
});
|
|
44
|
+
|
|
45
|
+
it("skips entries with no parent link", () => {
|
|
46
|
+
const map = groupUIMessagesByParent([
|
|
47
|
+
uiMessage(undefined),
|
|
48
|
+
uiMessage({ message_id: "" }),
|
|
49
|
+
]);
|
|
50
|
+
|
|
51
|
+
expect(map.size).toBe(0);
|
|
52
|
+
});
|
|
53
|
+
|
|
54
|
+
it("collects several UIs under the same parent in order", () => {
|
|
55
|
+
const a = uiMessage({ message_id: "msg-2" }, "chart");
|
|
56
|
+
const b = uiMessage({ message_id: "msg-2" }, "table");
|
|
57
|
+
const map = groupUIMessagesByParent([a, b]);
|
|
58
|
+
|
|
59
|
+
expect(map.get("msg-2")).toEqual([a, b]);
|
|
60
|
+
});
|
|
61
|
+
});
|
package/src/hooks.ts
ADDED
|
@@ -0,0 +1,156 @@
|
|
|
1
|
+
"use client";
|
|
2
|
+
|
|
3
|
+
import { useAui } from "@assistant-ui/store";
|
|
4
|
+
import type {
|
|
5
|
+
AssembledToolCall,
|
|
6
|
+
SubagentDiscoverySnapshot,
|
|
7
|
+
SubgraphDiscoverySnapshot,
|
|
8
|
+
} from "@langchain/react";
|
|
9
|
+
import { langChainExtras } from "./runtimeExtras";
|
|
10
|
+
import type { LangChainBaseMessage } from "./types";
|
|
11
|
+
|
|
12
|
+
const EMPTY_TOOL_CALLS: readonly AssembledToolCall[] = [];
|
|
13
|
+
|
|
14
|
+
const EMPTY_INTERRUPTS: readonly { id?: string; value?: unknown }[] = [];
|
|
15
|
+
|
|
16
|
+
const EMPTY_SUBAGENTS: ReadonlyMap<string, SubagentDiscoverySnapshot> =
|
|
17
|
+
new Map();
|
|
18
|
+
|
|
19
|
+
const EMPTY_SUBGRAPHS: ReadonlyMap<string, SubgraphDiscoverySnapshot> =
|
|
20
|
+
new Map();
|
|
21
|
+
|
|
22
|
+
/**
|
|
23
|
+
* Read the current LangGraph interrupt state from the runtime extras.
|
|
24
|
+
*/
|
|
25
|
+
export const useLangChainInterruptState = () =>
|
|
26
|
+
langChainExtras.use((e) => e.interrupt, undefined);
|
|
27
|
+
|
|
28
|
+
/**
|
|
29
|
+
* Read every interrupt pending at the current checkpoint, each with `id`
|
|
30
|
+
* and `value`. Defaults to an empty array. Pair with `useLangChainRespondAll`
|
|
31
|
+
* (keyed by interrupt id) to resolve several at once.
|
|
32
|
+
*/
|
|
33
|
+
export const useLangChainInterrupts = () =>
|
|
34
|
+
langChainExtras.use(
|
|
35
|
+
(e) => e.interrupts ?? EMPTY_INTERRUPTS,
|
|
36
|
+
EMPTY_INTERRUPTS,
|
|
37
|
+
);
|
|
38
|
+
|
|
39
|
+
/** Read the last run/hydration error from the runtime extras. */
|
|
40
|
+
export const useLangChainError = () =>
|
|
41
|
+
langChainExtras.use((e) => e.error, undefined);
|
|
42
|
+
|
|
43
|
+
/**
|
|
44
|
+
* Read the root tool calls assembled by `useStream` from the `tools`
|
|
45
|
+
* channel. Defaults to an empty array, so consumers can `.map` without
|
|
46
|
+
* a guard. Useful for rendering pending/streamed tool calls and
|
|
47
|
+
* approval UIs.
|
|
48
|
+
*/
|
|
49
|
+
export const useLangChainToolCalls = () =>
|
|
50
|
+
langChainExtras.use((e) => e.toolCalls ?? EMPTY_TOOL_CALLS, EMPTY_TOOL_CALLS);
|
|
51
|
+
|
|
52
|
+
/** Subagents discovered on the current run (multi-agent graphs). */
|
|
53
|
+
export const useLangChainSubagents = () =>
|
|
54
|
+
langChainExtras.use((e) => e.subagents ?? EMPTY_SUBAGENTS, EMPTY_SUBAGENTS);
|
|
55
|
+
|
|
56
|
+
/** Subgraphs discovered on the current run. */
|
|
57
|
+
export const useLangChainSubgraphs = () =>
|
|
58
|
+
langChainExtras.use((e) => e.subgraphs ?? EMPTY_SUBGRAPHS, EMPTY_SUBGRAPHS);
|
|
59
|
+
|
|
60
|
+
/**
|
|
61
|
+
* The underlying `useStream` handle, for advanced views — pass it to v1's
|
|
62
|
+
* `useMessages(stream, target)` / `useToolCalls(stream, target)` with a
|
|
63
|
+
* subagent/subgraph from `useLangChainSubagents` / `useLangChainSubgraphs`.
|
|
64
|
+
* `undefined` outside the runtime provider.
|
|
65
|
+
*/
|
|
66
|
+
export const useLangChainStream = () =>
|
|
67
|
+
langChainExtras.use((e) => e.stream, undefined);
|
|
68
|
+
|
|
69
|
+
/**
|
|
70
|
+
* Returns a function to submit raw state updates to the LangGraph agent,
|
|
71
|
+
* bypassing the normal message flow. Useful for sending interrupt resume
|
|
72
|
+
* commands.
|
|
73
|
+
*/
|
|
74
|
+
export const useLangChainSubmit = () => {
|
|
75
|
+
const aui = useAui();
|
|
76
|
+
return (
|
|
77
|
+
values: Record<string, unknown> | null | undefined,
|
|
78
|
+
options?: Record<string, unknown>,
|
|
79
|
+
) => langChainExtras.get(aui).submit(values, options);
|
|
80
|
+
};
|
|
81
|
+
|
|
82
|
+
/**
|
|
83
|
+
* Resume a LangGraph interrupt with a response payload via
|
|
84
|
+
* `useStream().respond`. Preferred over `useLangChainSendCommand`; it
|
|
85
|
+
* carries the response cleanly and handles interrupt namespaces.
|
|
86
|
+
*/
|
|
87
|
+
export const useLangChainRespond = () => {
|
|
88
|
+
const aui = useAui();
|
|
89
|
+
return (response: unknown, options?: Record<string, unknown>) =>
|
|
90
|
+
langChainExtras.get(aui).respond(response, options);
|
|
91
|
+
};
|
|
92
|
+
|
|
93
|
+
/**
|
|
94
|
+
* Resume several LangGraph interrupts pending at the same checkpoint in
|
|
95
|
+
* one run via `useStream().respondAll`. Use when a run pauses on multiple
|
|
96
|
+
* interrupts at once; sequential `useLangChainRespond` calls can't service
|
|
97
|
+
* them (the first resume starts a run, stranding the rest).
|
|
98
|
+
*/
|
|
99
|
+
export const useLangChainRespondAll = () => {
|
|
100
|
+
const aui = useAui();
|
|
101
|
+
return (
|
|
102
|
+
responsesById: Record<string, unknown>,
|
|
103
|
+
options?: Record<string, unknown>,
|
|
104
|
+
) => langChainExtras.get(aui).respondAll(responsesById, options);
|
|
105
|
+
};
|
|
106
|
+
|
|
107
|
+
/**
|
|
108
|
+
* Submit a list of LangChain-shaped messages on the current thread.
|
|
109
|
+
* Parity helper for migrating from `useLangGraphSend`. Routes to
|
|
110
|
+
* `useStream().submit({ [messagesKey]: messages }, options)`.
|
|
111
|
+
*/
|
|
112
|
+
export const useLangChainSend = () => {
|
|
113
|
+
const aui = useAui();
|
|
114
|
+
return (
|
|
115
|
+
messages: readonly LangChainBaseMessage[],
|
|
116
|
+
options?: Record<string, unknown>,
|
|
117
|
+
) => {
|
|
118
|
+
const extras = langChainExtras.get(aui);
|
|
119
|
+
return extras.submit({ [extras.messagesKey]: messages }, options);
|
|
120
|
+
};
|
|
121
|
+
};
|
|
122
|
+
|
|
123
|
+
/**
|
|
124
|
+
* Submit a `useStream` command (e.g. interrupt resume). Parity helper
|
|
125
|
+
* for migrating from `useLangGraphSendCommand`. Note that v1's command
|
|
126
|
+
* shape (`{ resume?, goto?, update? }`) differs from the legacy
|
|
127
|
+
* `{ resume: string }` form; to carry a payload, use the input or
|
|
128
|
+
* `stream.respond` instead.
|
|
129
|
+
*/
|
|
130
|
+
export const useLangChainSendCommand = () => {
|
|
131
|
+
const submit = useLangChainSubmit();
|
|
132
|
+
return (command: Record<string, unknown>) => submit(null, { command });
|
|
133
|
+
};
|
|
134
|
+
|
|
135
|
+
/**
|
|
136
|
+
* Read a custom LangGraph state key from the current thread. Mirrors
|
|
137
|
+
* `useStream().values[key]` from `@langchain/react` and updates when the
|
|
138
|
+
* stream emits new state.
|
|
139
|
+
*
|
|
140
|
+
* @example
|
|
141
|
+
* ```tsx
|
|
142
|
+
* const todos = useLangChainState<Todo[]>("todos");
|
|
143
|
+
* const files = useLangChainState<Record<string, string>>("files", {});
|
|
144
|
+
* ```
|
|
145
|
+
*/
|
|
146
|
+
export function useLangChainState<T>(key: string): T | undefined;
|
|
147
|
+
export function useLangChainState<T>(key: string, defaultValue: T): T;
|
|
148
|
+
export function useLangChainState<T>(
|
|
149
|
+
key: string,
|
|
150
|
+
defaultValue?: T,
|
|
151
|
+
): T | undefined {
|
|
152
|
+
return langChainExtras.use((e) => {
|
|
153
|
+
const value = e.values[key] as T | undefined;
|
|
154
|
+
return value !== undefined ? value : defaultValue;
|
|
155
|
+
}, defaultValue);
|
|
156
|
+
}
|
package/src/index.ts
CHANGED
|
@@ -1,19 +1,34 @@
|
|
|
1
|
+
export { useStreamRuntime } from "./useStreamRuntime";
|
|
2
|
+
|
|
1
3
|
export {
|
|
2
|
-
useStreamRuntime,
|
|
3
4
|
useLangChainError,
|
|
5
|
+
useLangChainInterrupts,
|
|
4
6
|
useLangChainInterruptState,
|
|
7
|
+
useLangChainRespond,
|
|
8
|
+
useLangChainRespondAll,
|
|
5
9
|
useLangChainSend,
|
|
6
10
|
useLangChainSendCommand,
|
|
7
11
|
useLangChainState,
|
|
12
|
+
useLangChainStream,
|
|
13
|
+
useLangChainSubagents,
|
|
14
|
+
useLangChainSubgraphs,
|
|
8
15
|
useLangChainSubmit,
|
|
9
16
|
useLangChainToolCalls,
|
|
10
|
-
} from "./
|
|
11
|
-
export type { UseStreamRuntimeOptions } from "./useStreamRuntime";
|
|
17
|
+
} from "./hooks";
|
|
12
18
|
|
|
13
19
|
export { convertLangChainBaseMessage } from "./convertMessages";
|
|
20
|
+
export { useLangChainStreamingTiming } from "./streamingTiming";
|
|
21
|
+
|
|
22
|
+
export type {
|
|
23
|
+
SubagentDiscoverySnapshot,
|
|
24
|
+
SubgraphDiscoverySnapshot,
|
|
25
|
+
} from "@langchain/react";
|
|
14
26
|
|
|
15
27
|
export type {
|
|
16
28
|
LangChainBaseMessage,
|
|
17
29
|
LangChainContentBlock,
|
|
18
30
|
LangChainToolCall,
|
|
31
|
+
RemoveUIMessage,
|
|
32
|
+
UIMessage,
|
|
33
|
+
UseStreamRuntimeOptions,
|
|
19
34
|
} from "./types";
|
|
@@ -0,0 +1,202 @@
|
|
|
1
|
+
import { describe, expect, it, vi } from "vitest";
|
|
2
|
+
import type { LangChainBaseMessage } from "./types";
|
|
3
|
+
import type { ForkCheckpointClient } from "./findForkCheckpointInHistory";
|
|
4
|
+
import { resolveForkCheckpoint } from "./resolveForkCheckpoint";
|
|
5
|
+
|
|
6
|
+
const msg = (id: string): LangChainBaseMessage => ({
|
|
7
|
+
_getType: () => "human",
|
|
8
|
+
content: "",
|
|
9
|
+
id,
|
|
10
|
+
});
|
|
11
|
+
|
|
12
|
+
type HistoryState = {
|
|
13
|
+
values: Record<string, unknown>;
|
|
14
|
+
checkpoint: { checkpoint_id?: string };
|
|
15
|
+
};
|
|
16
|
+
|
|
17
|
+
const makeClient = (history: HistoryState[]): ForkCheckpointClient => ({
|
|
18
|
+
threads: {
|
|
19
|
+
getHistory: vi.fn(async () => history as never),
|
|
20
|
+
},
|
|
21
|
+
});
|
|
22
|
+
|
|
23
|
+
const metadata = (entries: Record<string, string | undefined>) =>
|
|
24
|
+
new Map(
|
|
25
|
+
Object.entries(entries).map(([id, parentCheckpointId]) => [
|
|
26
|
+
id,
|
|
27
|
+
{ parentCheckpointId },
|
|
28
|
+
]),
|
|
29
|
+
);
|
|
30
|
+
|
|
31
|
+
describe("resolveForkCheckpoint", () => {
|
|
32
|
+
it("uses the head metadata fast-path without hitting getHistory", async () => {
|
|
33
|
+
const client = makeClient([]);
|
|
34
|
+
|
|
35
|
+
const result = await resolveForkCheckpoint(
|
|
36
|
+
client,
|
|
37
|
+
"thread-1",
|
|
38
|
+
[msg("h1"), msg("a1")],
|
|
39
|
+
"h1",
|
|
40
|
+
"a1",
|
|
41
|
+
metadata({ a1: "cp-head" }),
|
|
42
|
+
"messages",
|
|
43
|
+
);
|
|
44
|
+
|
|
45
|
+
expect(result).toBe("cp-head");
|
|
46
|
+
expect(client.threads.getHistory).not.toHaveBeenCalled();
|
|
47
|
+
});
|
|
48
|
+
|
|
49
|
+
it("falls back to history search when sourceId is not the head message", async () => {
|
|
50
|
+
const client = makeClient([
|
|
51
|
+
{
|
|
52
|
+
values: { messages: [msg("h1")] },
|
|
53
|
+
checkpoint: { checkpoint_id: "cp-h1" },
|
|
54
|
+
},
|
|
55
|
+
]);
|
|
56
|
+
|
|
57
|
+
const result = await resolveForkCheckpoint(
|
|
58
|
+
client,
|
|
59
|
+
"thread-1",
|
|
60
|
+
[msg("h1"), msg("a1"), msg("h2")],
|
|
61
|
+
"h1",
|
|
62
|
+
"h1",
|
|
63
|
+
metadata({ h2: "cp-head" }),
|
|
64
|
+
"messages",
|
|
65
|
+
);
|
|
66
|
+
|
|
67
|
+
expect(result).toBe("cp-h1");
|
|
68
|
+
expect(client.threads.getHistory).toHaveBeenCalled();
|
|
69
|
+
});
|
|
70
|
+
|
|
71
|
+
it("falls back to history search when the head metadata has no checkpoint", async () => {
|
|
72
|
+
const client = makeClient([
|
|
73
|
+
{
|
|
74
|
+
values: { messages: [msg("h1")] },
|
|
75
|
+
checkpoint: { checkpoint_id: "cp-h1" },
|
|
76
|
+
},
|
|
77
|
+
]);
|
|
78
|
+
|
|
79
|
+
const result = await resolveForkCheckpoint(
|
|
80
|
+
client,
|
|
81
|
+
"thread-1",
|
|
82
|
+
[msg("h1"), msg("a1")],
|
|
83
|
+
"h1",
|
|
84
|
+
"a1",
|
|
85
|
+
metadata({}),
|
|
86
|
+
"messages",
|
|
87
|
+
);
|
|
88
|
+
|
|
89
|
+
expect(result).toBe("cp-h1");
|
|
90
|
+
expect(client.threads.getHistory).toHaveBeenCalled();
|
|
91
|
+
});
|
|
92
|
+
|
|
93
|
+
it("matches the parent prefix through history search", async () => {
|
|
94
|
+
const client = makeClient([
|
|
95
|
+
{
|
|
96
|
+
values: { messages: [msg("h1"), msg("a1")] },
|
|
97
|
+
checkpoint: { checkpoint_id: "cp-full" },
|
|
98
|
+
},
|
|
99
|
+
{
|
|
100
|
+
values: { messages: [msg("h1")] },
|
|
101
|
+
checkpoint: { checkpoint_id: "cp-parent" },
|
|
102
|
+
},
|
|
103
|
+
]);
|
|
104
|
+
|
|
105
|
+
const result = await resolveForkCheckpoint(
|
|
106
|
+
client,
|
|
107
|
+
"thread-1",
|
|
108
|
+
[msg("h1"), msg("a1")],
|
|
109
|
+
"h1",
|
|
110
|
+
undefined,
|
|
111
|
+
undefined,
|
|
112
|
+
"messages",
|
|
113
|
+
);
|
|
114
|
+
|
|
115
|
+
expect(result).toBe("cp-parent");
|
|
116
|
+
});
|
|
117
|
+
|
|
118
|
+
it("returns null when a non-null parentId is not in the messages", async () => {
|
|
119
|
+
const client = makeClient([]);
|
|
120
|
+
|
|
121
|
+
const result = await resolveForkCheckpoint(
|
|
122
|
+
client,
|
|
123
|
+
"thread-1",
|
|
124
|
+
[msg("h1"), msg("a1")],
|
|
125
|
+
"missing",
|
|
126
|
+
undefined,
|
|
127
|
+
undefined,
|
|
128
|
+
"messages",
|
|
129
|
+
);
|
|
130
|
+
|
|
131
|
+
expect(result).toBeNull();
|
|
132
|
+
expect(client.threads.getHistory).not.toHaveBeenCalled();
|
|
133
|
+
});
|
|
134
|
+
|
|
135
|
+
it("forks the first message from the initial empty-message checkpoint", async () => {
|
|
136
|
+
const client = makeClient([
|
|
137
|
+
{
|
|
138
|
+
values: { messages: [msg("h1"), msg("a1")] },
|
|
139
|
+
checkpoint: { checkpoint_id: "cp-full" },
|
|
140
|
+
},
|
|
141
|
+
{
|
|
142
|
+
values: { messages: [] },
|
|
143
|
+
checkpoint: { checkpoint_id: "cp-initial" },
|
|
144
|
+
},
|
|
145
|
+
]);
|
|
146
|
+
|
|
147
|
+
const result = await resolveForkCheckpoint(
|
|
148
|
+
client,
|
|
149
|
+
"thread-1",
|
|
150
|
+
[msg("h1"), msg("a1")],
|
|
151
|
+
null,
|
|
152
|
+
"h1",
|
|
153
|
+
metadata({}),
|
|
154
|
+
"messages",
|
|
155
|
+
);
|
|
156
|
+
|
|
157
|
+
expect(result).toBe("cp-initial");
|
|
158
|
+
});
|
|
159
|
+
|
|
160
|
+
it("returns null when no checkpoint resolves", async () => {
|
|
161
|
+
const client = makeClient([
|
|
162
|
+
{
|
|
163
|
+
values: { messages: [msg("other")] },
|
|
164
|
+
checkpoint: { checkpoint_id: "cp-other" },
|
|
165
|
+
},
|
|
166
|
+
]);
|
|
167
|
+
|
|
168
|
+
const result = await resolveForkCheckpoint(
|
|
169
|
+
client,
|
|
170
|
+
"thread-1",
|
|
171
|
+
[msg("h1")],
|
|
172
|
+
"h1",
|
|
173
|
+
undefined,
|
|
174
|
+
undefined,
|
|
175
|
+
"messages",
|
|
176
|
+
);
|
|
177
|
+
|
|
178
|
+
expect(result).toBeNull();
|
|
179
|
+
});
|
|
180
|
+
|
|
181
|
+
it("swallows history search errors and returns null", async () => {
|
|
182
|
+
const client: ForkCheckpointClient = {
|
|
183
|
+
threads: {
|
|
184
|
+
getHistory: vi.fn(async () => {
|
|
185
|
+
throw new Error("network");
|
|
186
|
+
}),
|
|
187
|
+
},
|
|
188
|
+
};
|
|
189
|
+
|
|
190
|
+
const result = await resolveForkCheckpoint(
|
|
191
|
+
client,
|
|
192
|
+
"thread-1",
|
|
193
|
+
[msg("h1"), msg("a1")],
|
|
194
|
+
"h1",
|
|
195
|
+
undefined,
|
|
196
|
+
undefined,
|
|
197
|
+
"messages",
|
|
198
|
+
);
|
|
199
|
+
|
|
200
|
+
expect(result).toBeNull();
|
|
201
|
+
});
|
|
202
|
+
});
|
|
@@ -0,0 +1,49 @@
|
|
|
1
|
+
import type { LangChainBaseMessage } from "./types";
|
|
2
|
+
import {
|
|
3
|
+
findForkCheckpointInHistory,
|
|
4
|
+
type ForkCheckpointClient,
|
|
5
|
+
} from "./findForkCheckpointInHistory";
|
|
6
|
+
|
|
7
|
+
type MessageMetadataSnapshot =
|
|
8
|
+
| ReadonlyMap<string, { readonly parentCheckpointId: string | undefined }>
|
|
9
|
+
| undefined;
|
|
10
|
+
|
|
11
|
+
/**
|
|
12
|
+
* Hydration seeds every message with the head's parent checkpoint, so only the
|
|
13
|
+
* head's recorded fork checkpoint is reliable; older turns are matched against
|
|
14
|
+
* server history by message id. A `null` `parentId` edits the first human
|
|
15
|
+
* message and forks from the thread's initial (empty-message) checkpoint.
|
|
16
|
+
*/
|
|
17
|
+
export const resolveForkCheckpoint = async (
|
|
18
|
+
client: ForkCheckpointClient,
|
|
19
|
+
threadId: string,
|
|
20
|
+
messages: readonly LangChainBaseMessage[],
|
|
21
|
+
parentId: string | null,
|
|
22
|
+
sourceId: string | null | undefined,
|
|
23
|
+
metadata: MessageMetadataSnapshot,
|
|
24
|
+
messagesKey: string,
|
|
25
|
+
): Promise<string | null> => {
|
|
26
|
+
const lastMessage = messages[messages.length - 1];
|
|
27
|
+
let checkpointId: string | null =
|
|
28
|
+
sourceId != null && lastMessage?.id === sourceId
|
|
29
|
+
? (metadata?.get(sourceId)?.parentCheckpointId ?? null)
|
|
30
|
+
: null;
|
|
31
|
+
|
|
32
|
+
if (!checkpointId) {
|
|
33
|
+
const parentIndex =
|
|
34
|
+
parentId == null ? -1 : messages.findIndex((m) => m.id === parentId);
|
|
35
|
+
if (parentId != null && parentIndex === -1) return null;
|
|
36
|
+
try {
|
|
37
|
+
checkpointId = await findForkCheckpointInHistory(
|
|
38
|
+
client,
|
|
39
|
+
threadId,
|
|
40
|
+
messages.slice(0, parentIndex + 1),
|
|
41
|
+
messagesKey,
|
|
42
|
+
);
|
|
43
|
+
} catch {
|
|
44
|
+
return null;
|
|
45
|
+
}
|
|
46
|
+
}
|
|
47
|
+
|
|
48
|
+
return checkpointId;
|
|
49
|
+
};
|
|
@@ -0,0 +1,118 @@
|
|
|
1
|
+
import { describe, expect, it } from "vitest";
|
|
2
|
+
import type { LangChainBaseMessage } from "./types";
|
|
3
|
+
import { langChainStreamingTimingAccessors } from "./streamingTiming";
|
|
4
|
+
|
|
5
|
+
const ai = (
|
|
6
|
+
fields: Partial<LangChainBaseMessage> & { id: string; content: unknown },
|
|
7
|
+
): LangChainBaseMessage => ({
|
|
8
|
+
_getType: () => "ai",
|
|
9
|
+
...fields,
|
|
10
|
+
});
|
|
11
|
+
|
|
12
|
+
const human = (id: string, content: unknown): LangChainBaseMessage => ({
|
|
13
|
+
_getType: () => "human",
|
|
14
|
+
id,
|
|
15
|
+
content,
|
|
16
|
+
});
|
|
17
|
+
|
|
18
|
+
const { getAssistantMessageId, getTextLength, getToolCallCount } =
|
|
19
|
+
langChainStreamingTimingAccessors;
|
|
20
|
+
|
|
21
|
+
describe("langChainStreamingTimingAccessors", () => {
|
|
22
|
+
it("resolves the last assistant message id via _getType()", () => {
|
|
23
|
+
expect(
|
|
24
|
+
getAssistantMessageId([human("u1", "hi"), ai({ id: "a1", content: "" })]),
|
|
25
|
+
).toBe("a1");
|
|
26
|
+
expect(
|
|
27
|
+
getAssistantMessageId([
|
|
28
|
+
ai({ id: "a1", content: "" }),
|
|
29
|
+
human("u2", "hi"),
|
|
30
|
+
ai({ id: "a2", content: "" }),
|
|
31
|
+
]),
|
|
32
|
+
).toBe("a2");
|
|
33
|
+
});
|
|
34
|
+
|
|
35
|
+
it("returns undefined when there is no assistant message", () => {
|
|
36
|
+
expect(getAssistantMessageId([human("u1", "hi")])).toBeUndefined();
|
|
37
|
+
});
|
|
38
|
+
|
|
39
|
+
it("measures string content length", () => {
|
|
40
|
+
expect(getTextLength([ai({ id: "a1", content: "hello" })], "a1")).toBe(5);
|
|
41
|
+
});
|
|
42
|
+
|
|
43
|
+
it("sums text and thinking lengths across content blocks", () => {
|
|
44
|
+
expect(
|
|
45
|
+
getTextLength(
|
|
46
|
+
[
|
|
47
|
+
ai({
|
|
48
|
+
id: "a1",
|
|
49
|
+
content: [
|
|
50
|
+
{ type: "thinking", thinking: "hmm" },
|
|
51
|
+
{ type: "text", text: "answer" },
|
|
52
|
+
],
|
|
53
|
+
}),
|
|
54
|
+
],
|
|
55
|
+
"a1",
|
|
56
|
+
),
|
|
57
|
+
).toBe("hmm".length + "answer".length);
|
|
58
|
+
});
|
|
59
|
+
|
|
60
|
+
it("measures reasoning blocks (summary and reasoning fields)", () => {
|
|
61
|
+
// reasoning via summary_text entries, joined like contentToParts does.
|
|
62
|
+
expect(
|
|
63
|
+
getTextLength(
|
|
64
|
+
[
|
|
65
|
+
ai({
|
|
66
|
+
id: "a1",
|
|
67
|
+
content: [
|
|
68
|
+
{
|
|
69
|
+
type: "reasoning",
|
|
70
|
+
summary: [
|
|
71
|
+
{ type: "summary_text", text: "step one" },
|
|
72
|
+
{ type: "summary_text", text: "step two" },
|
|
73
|
+
],
|
|
74
|
+
},
|
|
75
|
+
],
|
|
76
|
+
}),
|
|
77
|
+
],
|
|
78
|
+
"a1",
|
|
79
|
+
),
|
|
80
|
+
).toBe("step one\n\n\nstep two".length);
|
|
81
|
+
|
|
82
|
+
// reasoning via the bare `reasoning` field when no summary is present.
|
|
83
|
+
expect(
|
|
84
|
+
getTextLength(
|
|
85
|
+
[
|
|
86
|
+
ai({
|
|
87
|
+
id: "a2",
|
|
88
|
+
content: [{ type: "reasoning", reasoning: "deduced" }],
|
|
89
|
+
}),
|
|
90
|
+
],
|
|
91
|
+
"a2",
|
|
92
|
+
),
|
|
93
|
+
).toBe("deduced".length);
|
|
94
|
+
});
|
|
95
|
+
|
|
96
|
+
it("counts tool_calls on the assistant message", () => {
|
|
97
|
+
expect(
|
|
98
|
+
getToolCallCount(
|
|
99
|
+
[
|
|
100
|
+
ai({
|
|
101
|
+
id: "a1",
|
|
102
|
+
content: "",
|
|
103
|
+
tool_calls: [
|
|
104
|
+
{ id: "t1", name: "search", args: {} },
|
|
105
|
+
{ id: "t2", name: "fetch", args: {} },
|
|
106
|
+
],
|
|
107
|
+
}),
|
|
108
|
+
],
|
|
109
|
+
"a1",
|
|
110
|
+
),
|
|
111
|
+
).toBe(2);
|
|
112
|
+
});
|
|
113
|
+
|
|
114
|
+
it("returns 0 when the message is missing or not assistant", () => {
|
|
115
|
+
expect(getTextLength([human("u1", "hi")], "a1")).toBe(0);
|
|
116
|
+
expect(getToolCallCount([human("u1", "hi")], "a1")).toBe(0);
|
|
117
|
+
});
|
|
118
|
+
});
|