From 62e688bb3737c5831ed3d2279ca458c02ca304ab Mon Sep 17 00:00:00 2001 From: Tat Dat Duong Date: Thu, 15 May 2025 16:20:38 -0700 Subject: [PATCH] Add rejoining of streams --- libs/sdk-js/src/react/stream.tsx | 247 +++++++++++++++++++++++-------- 1 file changed, 184 insertions(+), 63 deletions(-) diff --git a/libs/sdk-js/src/react/stream.tsx b/libs/sdk-js/src/react/stream.tsx index 69c4a27a2..84f4bcf92 100644 --- a/libs/sdk-js/src/react/stream.tsx +++ b/libs/sdk-js/src/react/stream.tsx @@ -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, - ) => { + 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, - 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; 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(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, + ) => { + 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, + 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; + + 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);