diff --git a/packages/react-headless/src/store/__tests__/createChatStore.test.ts b/packages/react-headless/src/store/__tests__/createChatStore.test.ts index 1d28dfbab..862bf3021 100644 --- a/packages/react-headless/src/store/__tests__/createChatStore.test.ts +++ b/packages/react-headless/src/store/__tests__/createChatStore.test.ts @@ -1,4 +1,5 @@ import { beforeEach, describe, expect, it, vi } from "vitest"; +import { EventType } from "../../types"; import type { Message, Thread, UserMessage } from "../types"; import { makeStore } from "./__helpers/makeStore"; @@ -424,5 +425,80 @@ describe("createChatStore", () => { expect(store.getState().messages).toEqual(newMessages); expect(store.getState().isLoadingMessages).toBe(false); }); + + it("does not write aborted stream events into the newly selected thread", async () => { + let release!: () => void; + const held = new Promise((resolve) => { + release = resolve; + }); + + const store = makeStore({ + send: vi.fn().mockResolvedValue(new Response("", { status: 200 })), + getMessages: vi.fn().mockResolvedValue([makeMessage("t2-m1")]), + streamProtocol: { + parse: async function* () { + await held; + yield { + type: EventType.TEXT_MESSAGE_START, + messageId: "asst-1", + role: "assistant", + }; + yield { + type: EventType.TEXT_MESSAGE_CONTENT, + messageId: "asst-1", + delta: "leaked", + }; + }, + }, + }); + store.setState({ selectedThreadId: "t1" }); + + const run = store.getState().processMessage({ role: "user", content: "hello" }); + await flushPromises(); + + store.getState().selectThread("t2"); + release(); + await run; + await flushPromises(); + + expect(store.getState().messages).toEqual([makeMessage("t2-m1")]); + expect(store.getState().messages.some((m) => m.role === "assistant")).toBe(false); + }); + + it("does not clear isRunning when a newer run replaced the aborted one", async () => { + let resolveFirstSend!: (response: Response) => void; + const firstSend = new Promise((resolve) => { + resolveFirstSend = resolve; + }); + let sendCount = 0; + const send = vi.fn().mockImplementation(() => { + sendCount += 1; + if (sendCount === 1) return firstSend; + return new Promise(() => {}); + }); + + const store = makeStore({ + send, + getMessages: vi.fn().mockResolvedValue([]), + streamProtocol: { parse: async function* () {} }, + }); + store.setState({ selectedThreadId: "t1" }); + + const first = store.getState().processMessage({ role: "user", content: "one" }); + await flushPromises(); + + store.getState().selectThread("t2"); + await flushPromises(); + + store.getState().processMessage({ role: "user", content: "two" }); + await flushPromises(); + expect(store.getState().isRunning).toBe(true); + + resolveFirstSend(new Response("", { status: 200 })); + await first; + await flushPromises(); + + expect(store.getState().isRunning).toBe(true); + }); }); }); diff --git a/packages/react-headless/src/store/createChatStore.ts b/packages/react-headless/src/store/createChatStore.ts index eed9fd96c..a1b721876 100644 --- a/packages/react-headless/src/store/createChatStore.ts +++ b/packages/react-headless/src/store/createChatStore.ts @@ -84,6 +84,7 @@ export const createChatStore = (configRef: React.RefObject(), + isRunning: false, }); }, @@ -103,6 +104,7 @@ export const createChatStore = (configRef: React.RefObject(), + isRunning: false, }); threadStorage .getMessages(threadId) @@ -162,7 +164,11 @@ export const createChatStore = (configRef: React.RefObject ({ messages: [...s.messages, optimisticMessage] })); abortController.signal.addEventListener("abort", () => { - set({ _abortController: null, isRunning: false }); + set({ + _abortController: null, + isRunning: false, + executingToolCallIds: new Set(), + }); }); try { @@ -214,28 +220,39 @@ export const createChatStore = (configRef: React.RefObject get()._abortController === abortController; + await processStreamedMessage({ response, - createMessage: (msg) => set((s) => ({ messages: [...s.messages, msg] })), - updateMessage: (msg) => + createMessage: (msg) => { + if (!belongsToThisRun()) return; + set((s) => ({ messages: [...s.messages, msg] })); + }, + updateMessage: (msg) => { + if (!belongsToThisRun()) return; set((s) => ({ messages: s.messages.map((m) => (m.id === msg.id ? msg : m)), - })), + })); + }, // A tool's args have closed (TOOL_CALL_END) → it is now executing. - markToolExecuting: (id) => + markToolExecuting: (id) => { + if (!belongsToThisRun()) return; set((s) => s.executingToolCallIds.has(id) ? s : { executingToolCallIds: new Set(s.executingToolCallIds).add(id) }, - ), + ); + }, // Its result landed (or it errored) → no longer executing. - clearToolExecuting: (id) => + clearToolExecuting: (id) => { + if (!belongsToThisRun()) return; set((s) => { if (!s.executingToolCallIds.has(id)) return s; const next = new Set(s.executingToolCallIds); next.delete(id); return { executingToolCallIds: next }; - }), + }); + }, adapter: configRef.current.llm.streamProtocol, }); } catch (e) { @@ -243,6 +260,9 @@ export const createChatStore = (configRef: React.RefObject { expect(settled[1]?.detail["errors"]).toEqual([queryError]); expect(settled[1]?.detail["response"]).toBe(response); }); + + it("settles when the Renderer unmounts while still streaming", () => { + const events: ObservabilityEvent[] = []; + const removeListener = observability.listenAll((event) => { + if (event.detail.kind === "react-lang:stream") events.push(event); + }); + + act(() => root.render(createElement(Harness, { isStreaming: true, queryErrors: [] }))); + expect(events.map((event) => event.detail["phase"])).toEqual(["streaming"]); + + act(() => root.unmount()); + expect(events.map((event) => event.detail["phase"])).toEqual(["streaming", "settled"]); + expect(events[0]?.detail.id).toBe(events[1]?.detail.id); + removeListener(); + + // afterEach also unmounts — give it a live root. + root = createRoot(container); + }); }); diff --git a/packages/react-lang/src/hooks/useStreamingObservability.ts b/packages/react-lang/src/hooks/useStreamingObservability.ts index b06c2e109..3bd3c6b50 100644 --- a/packages/react-lang/src/hooks/useStreamingObservability.ts +++ b/packages/react-lang/src/hooks/useStreamingObservability.ts @@ -124,9 +124,39 @@ function parserMetadata(result: ParseResult | null) { }; } +function emitSettled( + state: StreamingObservabilityState, + update: StreamingObservabilityUpdate, + response: string | null, + result: ParseResult | null, + errors: OpenUIError[], + libraryIdFields: { __libraryId?: string }, +) { + observability(errors.length > 0 ? "error" : "info", { + id: update.id, + kind: STREAM_EVENT_KIND, + phase: STREAM_PHASE_SETTLED, + updateIndex: update.updateIndex, + response, + responseLength: response?.length ?? 0, + parser: parserMetadata(result), + errors, + errorCount: errors.length, + ...captureStreamTiming(state), + ...libraryIdFields, + message: + errors.length > 0 + ? `OpenUI Lang settled with ${errors.length} error${errors.length === 1 ? "" : "s"}` + : "OpenUI Lang settled", + } satisfies SettledStreamEventDetail); +} + /** * Publishes the incremental OpenUI Lang stream lifecycle. The stable id is * created only after this Renderer instance has actually entered streaming. + * + * Switching chats unmounts the Renderer while `isStreaming` is still true, so + * unmount also settles — otherwise DevTools stays on "Streaming" forever. */ export function useStreamingObservability({ response, @@ -138,6 +168,15 @@ export function useStreamingObservability({ __libraryId, }: UseStreamingObservabilityOptions): void { const streamRef = useRef(createStreamingObservabilityState()); + const latestRef = useRef({ + publish, + isStreaming, + response, + result, + errorsRef, + __libraryId, + }); + latestRef.current = { publish, isStreaming, response, result, errorsRef, __libraryId }; useEffect(() => { if (!publish) return; @@ -170,23 +209,27 @@ export function useStreamingObservability({ } if (update?.phase === STREAM_PHASE_SETTLED) { - observability(errors.length > 0 ? "error" : "info", { - id: update.id, - kind: STREAM_EVENT_KIND, - phase: STREAM_PHASE_SETTLED, - updateIndex: update.updateIndex, - response, - responseLength: response?.length ?? 0, - parser: parserMetadata(result), - errors, - errorCount: errors.length, - ...captureStreamTiming(streamRef.current), - ...libraryIdFields, - message: - errors.length > 0 - ? `OpenUI Lang settled with ${errors.length} error${errors.length === 1 ? "" : "s"}` - : "OpenUI Lang settled", - } satisfies SettledStreamEventDetail); + emitSettled(streamRef.current, update, response, result, errors, libraryIdFields); } }, [publish, isStreaming, response, result, errorsRef, errorRevision, __libraryId]); + + useEffect(() => { + return () => { + const latest = latestRef.current; + if (!latest.publish || !latest.isStreaming) return; + const state = streamRef.current; + if (!state.id || state.settled) return; + const errors = latest.errorsRef.current; + const update = advanceStreamingObservability( + state, + false, + latest.response, + JSON.stringify(errors), + ); + if (update?.phase !== STREAM_PHASE_SETTLED) return; + const libraryIdFields = + latest.__libraryId !== undefined ? { __libraryId: latest.__libraryId } : {}; + emitSettled(state, update, latest.response, latest.result, errors, libraryIdFields); + }; + }, []); } diff --git a/packages/react-ui/src/components/AgentInterface/AgentInterface.tsx b/packages/react-ui/src/components/AgentInterface/AgentInterface.tsx index 07101e968..6a2c8abf3 100644 --- a/packages/react-ui/src/components/AgentInterface/AgentInterface.tsx +++ b/packages/react-ui/src/components/AgentInterface/AgentInterface.tsx @@ -215,8 +215,18 @@ export const AgentInterface: AgentInterfaceComponent = ((props: AgentInterfacePr const resolvedAssistantMessage = useMemo(() => { if (components?.AssistantMessage) return components.AssistantMessage; if (componentLibrary) { - const Cmp = ({ message }: { message: AssistantMessage }) => ( - + const Cmp = ({ + message, + isStreaming, + }: { + message: AssistantMessage; + isStreaming: boolean; + }) => ( + ); return Cmp; } diff --git a/packages/react-ui/src/components/OpenUIChat/GenUIAssistantMessage.tsx b/packages/react-ui/src/components/OpenUIChat/GenUIAssistantMessage.tsx index f36907ac6..ba4f40648 100644 --- a/packages/react-ui/src/components/OpenUIChat/GenUIAssistantMessage.tsx +++ b/packages/react-ui/src/components/OpenUIChat/GenUIAssistantMessage.tsx @@ -18,9 +18,12 @@ import { AssistantMessageContainer } from "./AssistantMessageContainer"; export const GenUIAssistantMessage = ({ message, library, + isStreaming: isStreamingProp, }: { message: AssistantMessage; library: Library; + /** When omitted, derived from the active thread run. */ + isStreaming?: boolean; }) => { const messages = useThread((s) => s.messages); const isRunning = useThread((s) => s.isRunning); @@ -28,7 +31,7 @@ export const GenUIAssistantMessage = ({ const updateMessage = useThread((s) => s.updateMessage); const lastAssistantId = useMemo(() => getLastAssistantMessageId(messages), [messages]); - const isStreaming = isRunning && lastAssistantId === message.id; + const isStreaming = isStreamingProp ?? (isRunning && lastAssistantId === message.id); // Strip the inline sentinels and separate any persisted form-state. const { content, contextString, contentHeader } = useMemo( diff --git a/packages/react-ui/src/components/OpenUIChat/withChatProvider.tsx b/packages/react-ui/src/components/OpenUIChat/withChatProvider.tsx index 5b1490627..3a08bba44 100644 --- a/packages/react-ui/src/components/OpenUIChat/withChatProvider.tsx +++ b/packages/react-ui/src/components/OpenUIChat/withChatProvider.tsx @@ -39,8 +39,12 @@ export function withChatProvider(WrappedComponent: React.Compon const genUIAssistantMessage = useMemo(() => { if (customAssistantMessage || !componentLibrary) return undefined; - return ({ message }: { message: AssistantMessage }) => ( - + return ({ message, isStreaming }: { message: AssistantMessage; isStreaming: boolean }) => ( + ); }, [customAssistantMessage, componentLibrary]);