From 735a76a16c0573209a66f1ecc64f1f317998689e Mon Sep 17 00:00:00 2001 From: Tat Dat Duong Date: Mon, 10 Feb 2025 17:53:22 -0800 Subject: [PATCH] Clean up the API --- libs/sdk-js/src/react/index.ts | 2 +- libs/sdk-js/src/react/stream.tsx | 291 +++++++++++++++++-------------- 2 files changed, 159 insertions(+), 134 deletions(-) diff --git a/libs/sdk-js/src/react/index.ts b/libs/sdk-js/src/react/index.ts index 7fc550ada..97a58f3d1 100644 --- a/libs/sdk-js/src/react/index.ts +++ b/libs/sdk-js/src/react/index.ts @@ -1 +1 @@ -export { useStream, LangGraphConfig } from "./stream.js"; +export { useStream, type MessageMetadata } from "./stream.js"; diff --git a/libs/sdk-js/src/react/stream.tsx b/libs/sdk-js/src/react/stream.tsx index d3941ce1e..aca6bf098 100644 --- a/libs/sdk-js/src/react/stream.tsx +++ b/libs/sdk-js/src/react/stream.tsx @@ -4,23 +4,23 @@ import { Client } from "../client.js"; import type { Command } from "../types.js"; import type { Message } from "../types.messages.js"; -import type { Config, ThreadState } from "../schema.js"; +import type { Checkpoint, Config, ThreadState } from "../schema.js"; import type { CustomStreamEvent, DebugStreamEvent, + ErrorStreamEvent, EventsStreamEvent, MessagesStreamEvent, MessagesTupleStreamEvent, + MetadataStreamEvent, + StreamMode, UpdatesStreamEvent, ValuesStreamEvent, } from "../types.stream.js"; import { type MutableRefObject, - type ReactNode, - createContext, useCallback, - useContext, useEffect, useMemo, useRef, @@ -33,6 +33,12 @@ import { convertToChunk, } from "@langchain/core/messages"; +class StreamError extends Error { + constructor(data: { error: string; message: string }) { + super([data.error, data.message].filter(Boolean).join(": ")); + } +} + class MessageTupleManager { chunks: Record = {}; @@ -106,12 +112,12 @@ interface ValidSequence { items: [Node, ...(Node | ValidFork)[]]; } -// forks -export type CheckpointBranchPath = string[]; +export type MessageMetadata> = { + messageId: string; + firstSeenState: ThreadState | undefined; -export type MessageBranch = { - current: CheckpointBranchPath; - options: CheckpointBranchPath[]; + branch: string | undefined; + branchOptions: string[] | undefined; }; function fetchHistory>( @@ -129,15 +135,18 @@ function useThreadHistory>( const [history, setHistory] = useState[]>([]); const fetcher = useCallback( - (threadId: string | undefined | null): Promise => { + ( + threadId: string | undefined | null, + ): Promise[]> => { if (threadId != null) { - return fetchHistory(client, threadId).then((history) => - setHistory(history), - ); + return fetchHistory(client, threadId).then((history) => { + setHistory(history); + return history; + }); } setHistory([]); - return Promise.resolve(); + return Promise.resolve([]); }, [], ); @@ -153,33 +162,18 @@ function useThreadHistory>( }; } -interface LangGraphConfig { - withMessages?: string; - onError?: (error: unknown) => void; - client?: Client; -} - -const ConfigProvider = createContext(null!); -export const LangGraphConfig = (props: { - config: LangGraphConfig; - children?: ReactNode; -}) => { - return ( - - {props.children} - - ); -}; - export function useStream< StateType extends Record = Record, UpdateType extends Record = Partial, CustomType = unknown, ->(options?: { - client?: Client; +>(options: { + assistantId: string; + client: Client; withMessages?: string; + onError?: (error: unknown) => void; + onFinish?: (state: ThreadState) => void; // TODO: can we make threadId uncontrollable / controllable? threadId?: string | null; @@ -192,14 +186,12 @@ export function useStream< | DebugStreamEvent | MessagesStreamEvent | MessagesTupleStreamEvent - | EventsStreamEvent; + | EventsStreamEvent + | MetadataStreamEvent + | ErrorStreamEvent; - const contextConfig = useContext(ConfigProvider); - const { withMessages, onError, threadId, client } = Object.assign( - {}, - contextConfig, - options, - ); + const { assistantId, threadId, client, withMessages, onError, onFinish } = + options; if (client == null) { throw new Error( @@ -207,22 +199,29 @@ export function useStream< ); } - const [branchPath, setBranchPath] = useState([]); - + const [branchPath, setBranchPath] = useState([]); + const [isLoading, setIsLoading] = useState(false); const [error, setError] = useState(undefined); const [events, setEvents] = useState([]); const [streamValues, setStreamValues] = useState(null); - const [streamMode, setStreamMode] = useState< + const messageManagerRef = useRef(new MessageTupleManager()); + const submittingRef = useRef(false); + const trackStreamModeRef = useRef< Array<"values" | "updates" | "events" | "custom" | "messages-tuple"> >(["values", "messages-tuple"]); - const manager = useRef(new MessageTupleManager()); - const submittingRef = useRef(false); + const trackStreamMode = useCallback( + (mode: Exclude) => { + if (!trackStreamModeRef.current.includes(mode)) + trackStreamModeRef.current.push(mode); + }, + [], + ); - // TODO: is this a responsibility of SWR / React Query? - // Maybe we should use that instead to allow rehydration? + // TODO: this should be done on the server to avoid pagination + // TODO: should we permit adapter? SWR / React Query? const history = useThreadHistory(threadId, client, submittingRef); const getMessages = useMemo(() => { @@ -297,11 +296,14 @@ export function useStream< return [rootSequence as ValidSequence, pathMap]; })(); - const [flatValues, checkpointPathMap] = (() => { + const [flatValues, flatPaths] = (() => { const result: ThreadState[] = []; // TODO: this is kinda ugly - const checkpointPathMap: Record = {}; + const flatPaths: Record< + string, + { current: string[] | undefined; branches: string[][] | undefined } + > = {}; const forkStack = branchPath.slice(); const queue: (Node | Fork)[] = [...sequence.items]; @@ -311,9 +313,9 @@ export function useStream< if (item.type === "node") { result.push(item.value); - checkpointPathMap[item.value.checkpoint.checkpoint_id!] = { + flatPaths[item.value.checkpoint.checkpoint_id!] = { current: item.path, - options: + branches: item.path.length > 0 ? pathMap[item.path.at(-2) ?? "$"] ?? [] : [], }; } @@ -333,62 +335,68 @@ export function useStream< } } - return [result, checkpointPathMap]; + return [result, flatPaths]; })(); - const lastSeenValue = flatValues.at(-1); - const historyValues = lastSeenValue?.values ?? ({} as StateType); + const threadHead: ThreadState | undefined = flatValues.at(-1); + const historyValues = threadHead?.values ?? ({} as StateType); - const messageMeta = (() => { + const messageMetadata = (() => { if (getMessages == null) return undefined; const alreadyShown = new Set(); - return getMessages(historyValues).map((message, idx) => { - const messageId = message.id ?? idx; - const firstSeenIdx = findLastIndex(history.data, (state) => - getMessages(state.values) - .map((m, idx) => m.id ?? idx) - .includes(messageId), - ); + return getMessages(historyValues).map( + (message, idx): MessageMetadata => { + const messageId = message.id ?? idx; + const firstSeenIdx = findLastIndex(history.data, (state) => + getMessages(state.values) + .map((m, idx) => m.id ?? idx) + .includes(messageId), + ); - const firstState = history.data[firstSeenIdx] as - | ThreadState - | undefined; + const firstSeen = history.data[firstSeenIdx] as + | ThreadState + | undefined; - let branch = firstState - ? checkpointPathMap[firstState.checkpoint.checkpoint_id!] - : undefined; - if (!branch?.current.length) branch = undefined; + let branch = firstSeen + ? flatPaths[firstSeen.checkpoint.checkpoint_id!] + : undefined; - const optionsShown = branch?.options?.flat(2).join(","); - if (optionsShown) { - if (alreadyShown.has(optionsShown)) branch = undefined; - alreadyShown.add(optionsShown); - } + if (!branch?.current?.length) branch = undefined; - return { messageId, firstState, branch }; - }); + // serialize branches + const optionsShown = branch?.branches?.flat(2).join(","); + if (optionsShown) { + if (alreadyShown.has(optionsShown)) branch = undefined; + alreadyShown.add(optionsShown); + } + + return { + messageId: messageId.toString(), + firstSeenState: firstSeen, + branch: branch?.current?.join(">"), + branchOptions: branch?.branches?.map((b) => b.join(">")), + }; + }, + ); })(); const handleSubmit = async ( values: UpdateType | undefined, submitOptions?: { config?: Config; + checkpoint?: Omit | null; command?: Command; + streamMode?: Array; optimisticValues?: | Partial | ((prev: StateType) => Partial); }, ) => { try { - // TODO: have loading state as well + setIsLoading(true); submittingRef.current = true; - // This is used to reset the path to make sure we always fetch the - // latest generation / edit. - // TODO: make sure it's actually aware of the config passed to handleSubmit - setBranchPath((path) => path.slice(0, -1)); - let usableThreadId = threadId; if (!usableThreadId) { const thread = await client.threads.create(); @@ -396,36 +404,51 @@ export function useStream< usableThreadId = thread.thread_id; } + const streamMode = unique([ + ...(submitOptions?.streamMode ?? []), + ...trackStreamModeRef.current, + ]); + + const checkpoint = + submitOptions?.checkpoint ?? threadHead?.checkpoint ?? undefined; + // @ts-expect-error + if (checkpoint != null) delete checkpoint.thread_id; + // TODO: why non-existent assistant ID does not throw an error here? - const run = (await client.runs.stream(usableThreadId, "agent", { + const run = (await client.runs.stream(usableThreadId, assistantId, { input: values as Record, - config: { - ...submitOptions?.config, - configurable: { - ...lastSeenValue?.checkpoint, - ...submitOptions?.config?.configurable, - }, - }, + config: submitOptions?.config, + checkpoint, streamMode, })) as AsyncGenerator; // Assumption: we're setting the initial value // Used for instant feedback - if (submitOptions?.optimisticValues != null) { - setStreamValues((streamValues) => { - const values = { ...historyValues, ...streamValues }; + setStreamValues(() => { + const values = { ...historyValues }; + + if (submitOptions?.optimisticValues != null) { return { ...values, ...(typeof submitOptions.optimisticValues === "function" ? submitOptions.optimisticValues(values) : submitOptions.optimisticValues), }; - }); - } + } + + return values; + }); for await (const { event, data } of run) { setEvents((events) => [...events, { event, data } as EventStreamEvent]); + if (event === "error") { + const error = new StreamError(data); + setError(error); + onError?.(error); + break; + } + if (event === "values") { setStreamValues(data); } else if (event === "messages") { @@ -433,7 +456,7 @@ export function useStream< const [serialized] = data; - const messageId = manager.current.add(serialized); + const messageId = messageManagerRef.current.add(serialized); if (!messageId) { console.warn( "Failed to add message to manager, no message ID found", @@ -447,7 +470,7 @@ export function useStream< // Assumption: we're concating the message const messages = getMessages(values).slice(); const { chunk, index } = - manager.current.get(messageId, messages.length) ?? {}; + messageManagerRef.current.get(messageId, messages.length) ?? {}; if (!chunk || index == null) return values; messages[index] = toMessageDict(chunk); @@ -457,15 +480,20 @@ export function useStream< } } - // TODO: add a "checkpoint" stream mode to get the branches directly - await history.mutate(usableThreadId); + // TODO: stream created checkpoints to avoid an unnecessary network request + const result = await history.mutate(usableThreadId); setStreamValues(null); + + const lastHead = result.at(0); + if (lastHead) onFinish?.(lastHead); } catch (error) { setError(error); onError?.(error); } finally { + setIsLoading(false); + // Assumption: messages are already handled, we can clear the manager - manager.current.clear(); + messageManagerRef.current.clear(); submittingRef.current = false; } }; @@ -473,9 +501,7 @@ export function useStream< const values = streamValues ?? historyValues; const stream = { get custom() { - if (!streamMode.includes("custom")) { - setStreamMode((mode) => unique([...mode, "custom"])); - } + trackStreamMode("custom"); return events .filter((item) => item.event === "custom") @@ -483,17 +509,12 @@ export function useStream< }, get events() { - if (!streamMode.includes("events")) { - setStreamMode((mode) => unique([...mode, "events"])); - } - + trackStreamMode("events"); return events; }, - get updates() { - if (!streamMode.includes("updates")) { - setStreamMode((mode) => unique([...mode, "updates"])); - } + get updates() { + trackStreamMode("updates"); return events .filter( (item): item is UpdatesStreamEvent => @@ -503,18 +524,27 @@ export function useStream< }, }; - return { - error, - handleSubmit, - setBranchPath, + const setBranch = useCallback( + (path: string) => setBranchPath(path.split(">")), + [setBranchPath], + ); + + return { + get values() { + trackStreamMode("values"); + return values; + }, + + error, + isLoading, + + handleSubmit, + setBranch, - sequence, stream, get messages() { - if (!streamMode.includes("messages-tuple")) { - setStreamMode((mode) => unique([...mode, "messages-tuple"])); - } + trackStreamMode("messages-tuple"); if (getMessages == null) { throw new Error( @@ -525,10 +555,11 @@ export function useStream< return getMessages(values); }, - getMessagesMeta(message: Message, index?: number) { - if (!streamMode.includes("messages-tuple")) { - setStreamMode((mode) => unique([...mode, "messages-tuple"])); - } + getMessagesMetadata( + message: Message, + index?: number, + ): MessageMetadata | undefined { + trackStreamMode("messages-tuple"); if (getMessages == null) { throw new Error( @@ -536,15 +567,9 @@ export function useStream< ); } - return messageMeta?.find((m) => m.messageId === (message.id ?? index)); - }, - - get values() { - if (!streamMode.includes("values")) { - setStreamMode((mode) => unique([...mode, "values"])); - } - - return values; + return messageMetadata?.find( + (m) => m.messageId === (message.id ?? index), + ); }, }; }