Add rejoining of streams

This commit is contained in:
Tat Dat Duong
2025-05-22 17:45:51 +02:00
parent 42eff39cd0
commit 62e688bb37
+184 -63
View File
@@ -502,6 +502,9 @@ export interface UseStreamOptions<
* Callback that is called when the thread ID is updated (ie when a new thread is created).
*/
onThreadId?: (threadId: string) => void;
/** Will rejoin the stream on mount */
joinOnMount?: boolean;
}
export interface UseStream<
@@ -647,7 +650,7 @@ export function useStream<
| ErrorStreamEvent
| FeedbackStreamEvent;
let { assistantId, messagesKey, onError, onFinish } = options;
let { assistantId, messagesKey, onError, onFinish, joinOnMount } = options;
messagesKey ??= "messages";
const client = useMemo(
@@ -805,10 +808,10 @@ export function useStream<
abortRef.current = null;
}, []);
const submit = async (
values: UpdateType | null | undefined,
submitOptions?: SubmitOptions<StateType, ConfigurableType>,
) => {
const join = async (metadata: {
runId: string;
threadId?: string | undefined | null;
}) => {
try {
setIsLoading(true);
setStreamError(undefined);
@@ -816,65 +819,8 @@ export function useStream<
submittingRef.current = true;
abortRef.current = new AbortController();
// Unbranch things
const newPath = submitOptions?.checkpoint?.checkpoint_id
? branchByCheckpoint[submitOptions?.checkpoint?.checkpoint_id]?.branch
: undefined;
if (newPath != null) setBranch(newPath ?? "");
// Assumption: we're setting the initial value
// Used for instant feedback
setStreamValues(() => {
const values = { ...historyValues };
if (submitOptions?.optimisticValues != null) {
return {
...values,
...(typeof submitOptions.optimisticValues === "function"
? submitOptions.optimisticValues(values)
: submitOptions.optimisticValues),
};
}
return values;
});
let usableThreadId = threadId;
if (!usableThreadId) {
const thread = await client.threads.create();
onThreadId(thread.thread_id);
usableThreadId = thread.thread_id;
}
const streamMode = unique([
...(submitOptions?.streamMode ?? []),
...trackStreamModeRef.current,
...callbackStreamMode,
]);
const checkpoint =
submitOptions?.checkpoint ?? threadHead?.checkpoint ?? undefined;
// @ts-expect-error
if (checkpoint != null) delete checkpoint.thread_id;
const run = client.runs.stream(usableThreadId, assistantId, {
input: values as Record<string, unknown>,
config: submitOptions?.config,
command: submitOptions?.command,
interruptBefore: submitOptions?.interruptBefore,
interruptAfter: submitOptions?.interruptAfter,
metadata: submitOptions?.metadata,
multitaskStrategy: submitOptions?.multitaskStrategy,
onCompletion: submitOptions?.onCompletion,
onDisconnect: submitOptions?.onDisconnect ?? "cancel",
const run = client.runs.joinStream(metadata.threadId, metadata.runId, {
signal: abortRef.current.signal,
checkpoint,
streamMode,
streamSubgraphs: submitOptions?.streamSubgraphs,
}) as AsyncGenerator<EventStreamEvent>;
let streamError: StreamError | undefined;
@@ -929,6 +875,181 @@ export function useStream<
}
}
// TODO: stream created checkpoints to avoid an unnecessary network request
if (metadata.threadId) {
const result = await history.mutate(metadata.threadId);
setStreamValues(null);
if (streamError != null) throw streamError;
const lastHead = result.at(0);
if (lastHead) onFinish?.(lastHead);
}
} catch (error) {
if (
!(
error instanceof Error &&
(error.name === "AbortError" || error.name === "TimeoutError")
)
) {
console.error(error);
setStreamError(error);
onError?.(error);
}
} finally {
setIsLoading(false);
// Assumption: messages are already handled, we can clear the manager
messageManagerRef.current.clear();
submittingRef.current = false;
abortRef.current = null;
}
};
const joinRef = useRef<typeof join>(join);
joinRef.current = join;
const joinKey = useMemo(() => {
if (joinOnMount) return undefined;
return { runId: localStorage.get(`lg:rejoin:${threadId}`), threadId };
}, [joinOnMount, threadId]);
useEffect(() => {
if (joinKey) joinRef.current?.(joinKey);
}, [joinKey]);
const submit = async (
values: UpdateType | null | undefined,
submitOptions?: SubmitOptions<StateType, ConfigurableType>,
) => {
try {
setIsLoading(true);
setStreamError(undefined);
submittingRef.current = true;
abortRef.current = new AbortController();
// Unbranch things
const newPath = submitOptions?.checkpoint?.checkpoint_id
? branchByCheckpoint[submitOptions?.checkpoint?.checkpoint_id]?.branch
: undefined;
if (newPath != null) setBranch(newPath ?? "");
// Assumption: we're setting the initial value
// Used for instant feedback
setStreamValues(() => {
const values = { ...historyValues };
if (submitOptions?.optimisticValues != null) {
return {
...values,
...(typeof submitOptions.optimisticValues === "function"
? submitOptions.optimisticValues(values)
: submitOptions.optimisticValues),
};
}
return values;
});
let usableThreadId = threadId;
if (!usableThreadId) {
const thread = await client.threads.create();
onThreadId(thread.thread_id);
usableThreadId = thread.thread_id;
}
const streamMode = unique([
...(submitOptions?.streamMode ?? []),
...trackStreamModeRef.current,
...callbackStreamMode,
]);
const checkpoint =
submitOptions?.checkpoint ?? threadHead?.checkpoint ?? undefined;
// @ts-expect-error
if (checkpoint != null) delete checkpoint.thread_id;
let rejoinKey: string | undefined;
const run = client.runs.stream(usableThreadId, assistantId, {
input: values as Record<string, unknown>,
config: submitOptions?.config,
command: submitOptions?.command,
interruptBefore: submitOptions?.interruptBefore,
interruptAfter: submitOptions?.interruptAfter,
metadata: submitOptions?.metadata,
multitaskStrategy: submitOptions?.multitaskStrategy,
onCompletion: submitOptions?.onCompletion,
onDisconnect: submitOptions?.onDisconnect ?? "cancel",
signal: abortRef.current.signal,
checkpoint,
streamMode,
streamSubgraphs: submitOptions?.streamSubgraphs,
onRunCreated(params) {
// rejoin the stream if needed
rejoinKey = `lg:rejoin:${params.thread_id ?? "temporary"}`;
window.localStorage.setItem(rejoinKey, params.run_id);
},
}) as AsyncGenerator<EventStreamEvent>;
let streamError: StreamError | undefined;
for await (const { event, data } of run) {
if (event === "error") {
streamError = new StreamError(data);
break;
}
if (event === "updates") options.onUpdateEvent?.(data);
if (event === "custom")
options.onCustomEvent?.(data, {
mutate: (update) =>
setStreamValues((prev) => {
// should not happen
if (prev == null) return prev;
return {
...prev,
...(typeof update === "function" ? update(prev) : update),
};
}),
});
if (event === "metadata") options.onMetadataEvent?.(data);
if (event === "events") options.onLangChainEvent?.(data);
if (event === "debug") options.onDebugEvent?.(data);
if (event === "values") setStreamValues(data);
if (event === "messages") {
const [serialized] = data;
const messageId = messageManagerRef.current.add(serialized);
if (!messageId) {
console.warn(
"Failed to add message to manager, no message ID found",
);
continue;
}
setStreamValues((streamValues) => {
const values = { ...historyValues, ...streamValues };
// Assumption: we're concatenating the message
const messages = getMessages(values).slice();
const { chunk, index } =
messageManagerRef.current.get(messageId, messages.length) ?? {};
if (!chunk || index == null) return values;
messages[index] = toMessageDict(chunk);
return { ...values, [messagesKey!]: messages };
});
}
}
if (rejoinKey) window.localStorage.removeItem(rejoinKey);
// TODO: stream created checkpoints to avoid an unnecessary network request
const result = await history.mutate(usableThreadId);
setStreamValues(null);