feat(sdk-js): expose branches

This commit is contained in:
Tat Dat Duong
2025-02-13 10:17:42 -08:00
parent f3403eab48
commit 2c66ac869d
+139 -120
View File
@@ -135,6 +135,123 @@ export type MessageMetadata<StateType extends Record<string, unknown>> = {
branchOptions: string[] | undefined;
};
function getBranchSequence<StateType extends Record<string, unknown>>(
history: ThreadState<StateType>[],
) {
const childrenMap: Record<string, ThreadState<StateType>[]> = {};
// First pass - collect nodes for each checkpoint
history.forEach((state) => {
const checkpointId = state.parent_checkpoint?.checkpoint_id ?? "$";
childrenMap[checkpointId] ??= [];
childrenMap[checkpointId].push(state);
});
// Second pass - create a tree of sequences
type Task = { id: string; sequence: Sequence; path: string[] };
const rootSequence: Sequence = { type: "sequence", items: [] };
const queue: Task[] = [{ id: "$", sequence: rootSequence, path: [] }];
const paths: string[][] = [];
const visited = new Set<string>();
while (queue.length > 0) {
const task = queue.shift()!;
if (visited.has(task.id)) continue;
visited.add(task.id);
const children = childrenMap[task.id];
if (children == null || children.length === 0) continue;
// If we've encountered a fork (2+ children), push the fork
// to the sequence and add a new sequence for each child
let fork: Fork | undefined;
if (children.length > 1) {
fork = { type: "fork", items: [] };
task.sequence.items.push(fork);
}
for (const value of children) {
const id = value.checkpoint.checkpoint_id!;
let sequence = task.sequence;
let path = task.path;
if (fork != null) {
sequence = { type: "sequence", items: [] };
fork.items.unshift(sequence);
path = path.slice();
path.push(id);
paths.push(path);
}
sequence.items.push({ type: "node", value, path });
queue.push({ id, sequence, path });
}
}
return { rootSequence, paths };
}
const PATH_SEP = ">";
const ROOT_ID = "$";
// Get flat view
function getBranchView<StateType extends Record<string, unknown>>(
sequence: Sequence<StateType>,
paths: string[][],
branch: string,
) {
const path = branch.split(PATH_SEP);
const pathMap: Record<string, string[][]> = {};
for (const path of paths) {
const parent = path.at(-2) ?? ROOT_ID;
pathMap[parent] ??= [];
pathMap[parent].unshift(path);
}
const history: ThreadState<StateType>[] = [];
const branchByCheckpoint: Record<
string,
{ branch: string | undefined; branchOptions: string[] | undefined }
> = {};
const forkStack = path.slice();
const queue: (Node<StateType> | Fork<StateType>)[] = [...sequence.items];
while (queue.length > 0) {
const item = queue.shift()!;
if (item.type === "node") {
history.push(item.value);
branchByCheckpoint[item.value.checkpoint.checkpoint_id!] = {
branch: item.path.join(PATH_SEP),
branchOptions: (item.path.length > 0
? pathMap[item.path.at(-2) ?? ROOT_ID] ?? []
: []
).map((p) => p.join(PATH_SEP)),
};
}
if (item.type === "fork") {
const forkId = forkStack.shift();
const index =
forkId != null
? item.items.findIndex((value) => {
const firstItem = value.items.at(0);
if (!firstItem || firstItem.type !== "node") return false;
return firstItem.value.checkpoint.checkpoint_id === forkId;
})
: -1;
const nextItems = item.items.at(index)?.items ?? [];
queue.push(...nextItems);
}
}
return { history, branchByCheckpoint };
}
function fetchHistory<StateType extends Record<string, unknown>>(
client: Client,
threadId: string,
@@ -274,9 +391,8 @@ export function useStream<
);
const [threadId, onThreadId] = useControllableThreadId(options);
const [branchPath, setBranchPath] = useState<string[]>([]);
const [branch, setBranch] = useState<string>("");
const [isLoading, setIsLoading] = useState(false);
const [_, setEvents] = useState<EventStreamEvent[]>([]);
const [streamError, setStreamError] = useState<unknown>(undefined);
const [streamValues, setStreamValues] = useState<StateType | null>(null);
@@ -330,111 +446,14 @@ export function useStream<
: [];
}, [messagesKey]);
const [sequence, pathMap] = (() => {
const childrenMap: Record<string, ThreadState<StateType>[]> = {};
const { rootSequence, paths } = getBranchSequence(history.data);
const { history: flatHistory, branchByCheckpoint } = getBranchView(
rootSequence,
paths,
branch,
);
// First pass - collect nodes for each checkpoint
history.data.forEach((state) => {
const checkpointId = state.parent_checkpoint?.checkpoint_id ?? "$";
childrenMap[checkpointId] ??= [];
childrenMap[checkpointId].push(state);
});
// Second pass - create a tree of sequences
type Task = { id: string; sequence: Sequence; path: string[] };
const rootSequence: Sequence = { type: "sequence", items: [] };
const queue: Task[] = [{ id: "$", sequence: rootSequence, path: [] }];
const paths: string[][] = [];
const visited = new Set<string>();
while (queue.length > 0) {
const task = queue.shift()!;
if (visited.has(task.id)) continue;
visited.add(task.id);
const children = childrenMap[task.id];
if (children == null || children.length === 0) continue;
// If we've encountered a fork (2+ children), push the fork
// to the sequence and add a new sequence for each child
let fork: Fork | undefined;
if (children.length > 1) {
fork = { type: "fork", items: [] };
task.sequence.items.push(fork);
}
for (const value of children) {
const id = value.checkpoint.checkpoint_id!;
let sequence = task.sequence;
let path = task.path;
if (fork != null) {
sequence = { type: "sequence", items: [] };
fork.items.unshift(sequence);
path = path.slice();
path.push(id);
paths.push(path);
}
sequence.items.push({ type: "node", value, path });
queue.push({ id, sequence, path });
}
}
// Third pass, create a map for available forks
const pathMap: Record<string, string[][]> = {};
for (const path of paths) {
const parent = path.at(-2) ?? "$";
pathMap[parent] ??= [];
pathMap[parent].unshift(path);
}
return [rootSequence as ValidSequence, pathMap];
})();
const [flatValues, flatPaths] = (() => {
const result: ThreadState<StateType>[] = [];
const flatPaths: Record<
string,
{ current: string[] | undefined; branches: string[][] | undefined }
> = {};
const forkStack = branchPath.slice();
const queue: (Node<StateType> | Fork<StateType>)[] = [...sequence.items];
while (queue.length > 0) {
const item = queue.shift()!;
if (item.type === "node") {
result.push(item.value);
flatPaths[item.value.checkpoint.checkpoint_id!] = {
current: item.path,
branches:
item.path.length > 0 ? pathMap[item.path.at(-2) ?? "$"] ?? [] : [],
};
}
if (item.type === "fork") {
const forkId = forkStack.shift();
const index =
forkId != null
? item.items.findIndex((value) => {
const firstItem = value.items.at(0);
if (!firstItem || firstItem.type !== "node") return false;
return firstItem.value.checkpoint.checkpoint_id === forkId;
})
: -1;
const nextItems = item.items.at(index)?.items ?? [];
queue.push(...nextItems);
}
}
return [result, flatPaths];
})();
const threadHead: ThreadState<StateType> | undefined = flatValues.at(-1);
const threadHead: ThreadState<StateType> | undefined = flatHistory.at(-1);
const historyValues = threadHead?.values ?? ({} as StateType);
const historyError = (() => {
const error = threadHead?.tasks?.at(-1)?.error;
@@ -470,13 +489,13 @@ export function useStream<
| undefined;
let branch = firstSeen
? flatPaths[firstSeen.checkpoint.checkpoint_id!]
? branchByCheckpoint[firstSeen.checkpoint.checkpoint_id!]
: undefined;
if (!branch?.current?.length) branch = undefined;
if (!branch?.branch?.length) branch = undefined;
// serialize branches
const optionsShown = branch?.branches?.flat(2).join(",");
const optionsShown = branch?.branchOptions?.flat(2).join(",");
if (optionsShown) {
if (alreadyShown.has(optionsShown)) branch = undefined;
alreadyShown.add(optionsShown);
@@ -485,8 +504,9 @@ export function useStream<
return {
messageId: messageId.toString(),
firstSeenState: firstSeen,
branch: branch?.current?.join(">"),
branchOptions: branch?.branches?.map((b) => b.join(">")),
branch: branch?.branch,
branchOptions: branch?.branchOptions,
};
},
);
@@ -561,9 +581,10 @@ export function useStream<
// Unbranch things
const newPath = submitOptions?.checkpoint?.checkpoint_id
? flatPaths[submitOptions?.checkpoint?.checkpoint_id]?.current
? branchByCheckpoint[submitOptions?.checkpoint?.checkpoint_id]?.branch
: undefined;
if (newPath != null) setBranchPath(newPath ?? []);
if (newPath != null) setBranch(newPath ?? "");
// Assumption: we're setting the initial value
// Used for instant feedback
@@ -584,8 +605,6 @@ export function useStream<
let streamError: StreamError | undefined;
for await (const { event, data } of run) {
setEvents((events) => [...events, { event, data } as EventStreamEvent]);
if (event === "error") {
streamError = new StreamError(data);
break;
@@ -656,11 +675,6 @@ export function useStream<
const error = isLoading ? streamError : historyError;
const values = streamValues ?? historyValues;
const setBranch = useCallback(
(path: string) => setBranchPath(path.split(">")),
[setBranchPath],
);
return {
get values() {
trackStreamMode("values");
@@ -672,8 +686,13 @@ export function useStream<
stop,
submit,
branch,
setBranch,
history: flatHistory,
experimental_branchTree: rootSequence,
get messages() {
trackStreamMode("messages-tuple");