feat(FN-3843): add chat stream reattach support
Adds the changeset for FN-3843, completing the chat stream reattach feature for the `@runfusion/fusion` CLI package. Fusion-Task-Id: FN-3843
This commit is contained in:
5
.changeset/fn-3843-chat-stream-reattach.md
Normal file
5
.changeset/fn-3843-chat-stream-reattach.md
Normal file
@@ -0,0 +1,5 @@
|
|||||||
|
---
|
||||||
|
"@runfusion/fusion": patch
|
||||||
|
---
|
||||||
|
|
||||||
|
Reattach to in-flight chat stream after reload so streaming responses keep rendering instead of disappearing.
|
||||||
@@ -1,5 +1,5 @@
|
|||||||
import { afterEach, describe, expect, it, vi } from "vitest";
|
import { afterEach, describe, expect, it, vi } from "vitest";
|
||||||
import { streamChatResponse } from "../legacy";
|
import { attachChatStream, streamChatResponse } from "../legacy";
|
||||||
|
|
||||||
function createChunkedStream(chunks: string[]): ReadableStream<Uint8Array> {
|
function createChunkedStream(chunks: string[]): ReadableStream<Uint8Array> {
|
||||||
const encoder = new TextEncoder();
|
const encoder = new TextEncoder();
|
||||||
@@ -140,3 +140,81 @@ describe("streamChatResponse SSE parser", () => {
|
|||||||
vi.useRealTimers();
|
vi.useRealTimers();
|
||||||
});
|
});
|
||||||
});
|
});
|
||||||
|
|
||||||
|
describe("attachChatStream", () => {
|
||||||
|
afterEach(() => {
|
||||||
|
vi.restoreAllMocks();
|
||||||
|
});
|
||||||
|
|
||||||
|
it("replays buffered events and done", async () => {
|
||||||
|
vi.spyOn(globalThis, "fetch").mockResolvedValue(
|
||||||
|
new Response(
|
||||||
|
createChunkedStream([
|
||||||
|
"event: text\n",
|
||||||
|
"data: \"Hello\"\n\n",
|
||||||
|
"event: done\n",
|
||||||
|
"data: {\"messageId\":\"m-1\"}\n\n",
|
||||||
|
]),
|
||||||
|
{ status: 200 },
|
||||||
|
),
|
||||||
|
);
|
||||||
|
|
||||||
|
const textChunks: string[] = [];
|
||||||
|
const donePayloads: Array<{ messageId: string }> = [];
|
||||||
|
|
||||||
|
attachChatStream("s-1", {
|
||||||
|
onText: (data) => textChunks.push(data),
|
||||||
|
onDone: (data) => donePayloads.push(data),
|
||||||
|
onError: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
await vi.waitFor(() => {
|
||||||
|
expect(textChunks).toEqual(["Hello"]);
|
||||||
|
expect(donePayloads).toEqual([{ messageId: "m-1" }]);
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
it("delivers live events after replay", async () => {
|
||||||
|
vi.spyOn(globalThis, "fetch").mockResolvedValue(
|
||||||
|
new Response(
|
||||||
|
createChunkedStream([
|
||||||
|
"event: text\n",
|
||||||
|
"data: \"A\"\n\n",
|
||||||
|
"event: text\n",
|
||||||
|
"data: \"B\"\n\n",
|
||||||
|
]),
|
||||||
|
{ status: 200 },
|
||||||
|
),
|
||||||
|
);
|
||||||
|
|
||||||
|
const textChunks: string[] = [];
|
||||||
|
|
||||||
|
attachChatStream("s-1", {
|
||||||
|
onText: (data) => textChunks.push(data),
|
||||||
|
onError: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
await vi.waitFor(() => {
|
||||||
|
expect(textChunks).toEqual(["A", "B"]);
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
it("aborts fetch when close is called", async () => {
|
||||||
|
let signal: AbortSignal | undefined;
|
||||||
|
vi.spyOn(globalThis, "fetch").mockImplementation((_input, init) => {
|
||||||
|
signal = init?.signal;
|
||||||
|
return new Promise<Response>(() => {
|
||||||
|
// keep open until aborted
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
const stream = attachChatStream("s-1", { onError: vi.fn() });
|
||||||
|
await vi.waitFor(() => {
|
||||||
|
expect(signal).toBeDefined();
|
||||||
|
});
|
||||||
|
|
||||||
|
stream.close();
|
||||||
|
expect(signal?.aborted).toBe(true);
|
||||||
|
expect(stream.isConnected()).toBe(false);
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|||||||
@@ -8321,19 +8321,21 @@ export function cancelChatResponse(
|
|||||||
* When attachments are provided, the request body is sent as multipart form data;
|
* When attachments are provided, the request body is sent as multipart form data;
|
||||||
* otherwise it uses the existing JSON payload path.
|
* otherwise it uses the existing JSON payload path.
|
||||||
*/
|
*/
|
||||||
|
export interface ChatStreamHandlers {
|
||||||
|
onThinking?: (data: string) => void;
|
||||||
|
onText?: (data: string) => void;
|
||||||
|
onToolStart?: (data: { toolName: string; args?: Record<string, unknown> }) => void;
|
||||||
|
onToolEnd?: (data: { toolName: string; isError: boolean; result?: unknown }) => void;
|
||||||
|
onFallback?: (data: { primaryModel: string; fallbackModel: string; triggerPoint: "session-creation" | "prompt-time" }) => void;
|
||||||
|
onDone?: (data: { messageId: string; message?: ChatMessage }) => void;
|
||||||
|
onError?: (data: string) => void;
|
||||||
|
onConnectionStateChange?: (state: StreamConnectionState) => void;
|
||||||
|
}
|
||||||
|
|
||||||
export function streamChatResponse(
|
export function streamChatResponse(
|
||||||
sessionId: string,
|
sessionId: string,
|
||||||
content: string,
|
content: string,
|
||||||
handlers: {
|
handlers: ChatStreamHandlers,
|
||||||
onThinking?: (data: string) => void;
|
|
||||||
onText?: (data: string) => void;
|
|
||||||
onToolStart?: (data: { toolName: string; args?: Record<string, unknown> }) => void;
|
|
||||||
onToolEnd?: (data: { toolName: string; isError: boolean; result?: unknown }) => void;
|
|
||||||
onFallback?: (data: { primaryModel: string; fallbackModel: string; triggerPoint: "session-creation" | "prompt-time" }) => void;
|
|
||||||
onDone?: (data: { messageId: string; message?: ChatMessage }) => void;
|
|
||||||
onError?: (data: string) => void;
|
|
||||||
onConnectionStateChange?: (state: StreamConnectionState) => void;
|
|
||||||
},
|
|
||||||
attachments?: File[],
|
attachments?: File[],
|
||||||
projectId?: string,
|
projectId?: string,
|
||||||
options?: { maxReconnectAttempts?: number; firstEventTimeoutMs?: number },
|
options?: { maxReconnectAttempts?: number; firstEventTimeoutMs?: number },
|
||||||
@@ -8573,6 +8575,192 @@ export function streamChatResponse(
|
|||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export function attachChatStream(
|
||||||
|
sessionId: string,
|
||||||
|
handlers: ChatStreamHandlers,
|
||||||
|
projectId?: string,
|
||||||
|
options?: { lastEventId?: number },
|
||||||
|
): { close: () => void; isConnected: () => boolean } {
|
||||||
|
const url = buildApiUrl(withProjectId(`/chat/sessions/${encodeURIComponent(sessionId)}/stream`, projectId));
|
||||||
|
const abortController = new AbortController();
|
||||||
|
let closedByUser = false;
|
||||||
|
let terminated = false;
|
||||||
|
|
||||||
|
const dispatchEvent = (eventName: string, rawData: string): void => {
|
||||||
|
if (!eventName) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
switch (eventName) {
|
||||||
|
case "thinking":
|
||||||
|
try {
|
||||||
|
handlers.onThinking?.(JSON.parse(rawData));
|
||||||
|
} catch {
|
||||||
|
handlers.onThinking?.(rawData);
|
||||||
|
}
|
||||||
|
break;
|
||||||
|
case "text":
|
||||||
|
try {
|
||||||
|
handlers.onText?.(JSON.parse(rawData));
|
||||||
|
} catch {
|
||||||
|
handlers.onText?.(rawData);
|
||||||
|
}
|
||||||
|
break;
|
||||||
|
case "tool_start":
|
||||||
|
try {
|
||||||
|
handlers.onToolStart?.(JSON.parse(rawData));
|
||||||
|
} catch {
|
||||||
|
// skip malformed event
|
||||||
|
}
|
||||||
|
break;
|
||||||
|
case "tool_end":
|
||||||
|
try {
|
||||||
|
handlers.onToolEnd?.(JSON.parse(rawData));
|
||||||
|
} catch {
|
||||||
|
// skip malformed event
|
||||||
|
}
|
||||||
|
break;
|
||||||
|
case "fallback":
|
||||||
|
try {
|
||||||
|
handlers.onFallback?.(JSON.parse(rawData));
|
||||||
|
} catch {
|
||||||
|
// skip malformed event
|
||||||
|
}
|
||||||
|
break;
|
||||||
|
case "done":
|
||||||
|
terminated = true;
|
||||||
|
try {
|
||||||
|
const parsed = JSON.parse(rawData) as { messageId?: unknown; message?: unknown };
|
||||||
|
handlers.onDone?.({
|
||||||
|
messageId: typeof parsed.messageId === "string" ? parsed.messageId : "",
|
||||||
|
...(parsed.message && typeof parsed.message === "object" ? { message: parsed.message as ChatMessage } : {}),
|
||||||
|
});
|
||||||
|
} catch {
|
||||||
|
handlers.onDone?.({ messageId: "" });
|
||||||
|
}
|
||||||
|
break;
|
||||||
|
case "error":
|
||||||
|
terminated = true;
|
||||||
|
try {
|
||||||
|
const parsed = JSON.parse(rawData);
|
||||||
|
handlers.onError?.(parsed.message || parsed);
|
||||||
|
} catch {
|
||||||
|
handlers.onError?.(rawData || "Stream error");
|
||||||
|
}
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
(async () => {
|
||||||
|
try {
|
||||||
|
const requestHeaders = withTokenHeader();
|
||||||
|
if (typeof options?.lastEventId === "number") {
|
||||||
|
requestHeaders["Last-Event-ID"] = String(options.lastEventId);
|
||||||
|
}
|
||||||
|
|
||||||
|
const res = await fetch(url, {
|
||||||
|
method: "GET",
|
||||||
|
headers: requestHeaders,
|
||||||
|
signal: abortController.signal,
|
||||||
|
});
|
||||||
|
|
||||||
|
if (!res.ok) {
|
||||||
|
const errorBody = await res.text();
|
||||||
|
let errorMsg = `Request failed: ${res.status}`;
|
||||||
|
try {
|
||||||
|
const parsed = JSON.parse(errorBody);
|
||||||
|
errorMsg = parsed.error || errorMsg;
|
||||||
|
} catch { /* use default */ }
|
||||||
|
handlers.onError?.(errorMsg);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
if (!res.body) {
|
||||||
|
handlers.onError?.("No response body");
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
handlers.onConnectionStateChange?.("connected");
|
||||||
|
|
||||||
|
const reader = res.body.getReader();
|
||||||
|
const decoder = new TextDecoder();
|
||||||
|
let buffer = "";
|
||||||
|
let currentEvent = "";
|
||||||
|
let currentDataLines: string[] = [];
|
||||||
|
|
||||||
|
const processLines = (chunk: string, flushPendingEvent = false): void => {
|
||||||
|
buffer += chunk;
|
||||||
|
const lines = buffer.split("\n");
|
||||||
|
buffer = lines.pop() || "";
|
||||||
|
|
||||||
|
if (flushPendingEvent && buffer.length > 0) {
|
||||||
|
lines.push(buffer);
|
||||||
|
buffer = "";
|
||||||
|
}
|
||||||
|
|
||||||
|
for (const rawLine of lines) {
|
||||||
|
const line = rawLine.endsWith("\r") ? rawLine.slice(0, -1) : rawLine;
|
||||||
|
|
||||||
|
if (line.startsWith("event:")) {
|
||||||
|
currentEvent = line.slice(6).trim();
|
||||||
|
} else if (line.startsWith("data:")) {
|
||||||
|
const value = line.slice(5);
|
||||||
|
currentDataLines.push(value.startsWith(" ") ? value.slice(1) : value);
|
||||||
|
} else if (line === "") {
|
||||||
|
const currentData = currentDataLines.join("\n");
|
||||||
|
dispatchEvent(currentEvent, currentData);
|
||||||
|
currentEvent = "";
|
||||||
|
currentDataLines = [];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if (flushPendingEvent && currentEvent && currentDataLines.length > 0) {
|
||||||
|
const trailingData = currentDataLines.join("\n");
|
||||||
|
dispatchEvent(currentEvent, trailingData);
|
||||||
|
currentEvent = "";
|
||||||
|
currentDataLines = [];
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
while (true) {
|
||||||
|
const { done, value } = await reader.read();
|
||||||
|
if (done) {
|
||||||
|
processLines(decoder.decode(), true);
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
processLines(decoder.decode(value, { stream: true }));
|
||||||
|
}
|
||||||
|
|
||||||
|
const hasUndispatchedTrailingFragment =
|
||||||
|
buffer.length > 0 || currentEvent.length > 0 || currentDataLines.length > 0;
|
||||||
|
|
||||||
|
if (!terminated && !closedByUser && !hasUndispatchedTrailingFragment) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
} catch (err: unknown) {
|
||||||
|
if (err instanceof DOMException && err.name === "AbortError") {
|
||||||
|
if (!closedByUser && !terminated) {
|
||||||
|
handlers.onError?.("Connection aborted");
|
||||||
|
}
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
if (closedByUser) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
handlers.onError?.(err instanceof Error ? err.message : "Connection error");
|
||||||
|
}
|
||||||
|
})();
|
||||||
|
|
||||||
|
return {
|
||||||
|
close: () => {
|
||||||
|
closedByUser = true;
|
||||||
|
abortController.abort();
|
||||||
|
},
|
||||||
|
isConnected: () => !closedByUser,
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
// ── Insights API ─────────────────────────────────────────────────────────────
|
// ── Insights API ─────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ vi.mock("../../api", () => ({
|
|||||||
updateChatSession: vi.fn(),
|
updateChatSession: vi.fn(),
|
||||||
deleteChatSession: vi.fn(),
|
deleteChatSession: vi.fn(),
|
||||||
streamChatResponse: vi.fn(),
|
streamChatResponse: vi.fn(),
|
||||||
|
attachChatStream: vi.fn(),
|
||||||
cancelChatResponse: vi.fn(),
|
cancelChatResponse: vi.fn(),
|
||||||
fetchAgents: vi.fn().mockResolvedValue([
|
fetchAgents: vi.fn().mockResolvedValue([
|
||||||
{ id: "agent-001", name: "Alpha", role: "executor", state: "idle", icon: undefined, createdAt: "2026-04-08T00:00:00.000Z", updatedAt: "2026-04-08T00:00:00.000Z", metadata: {} },
|
{ id: "agent-001", name: "Alpha", role: "executor", state: "idle", icon: undefined, createdAt: "2026-04-08T00:00:00.000Z", updatedAt: "2026-04-08T00:00:00.000Z", metadata: {} },
|
||||||
@@ -50,6 +51,7 @@ const mockFetchChatMessages = vi.mocked(apiModule.fetchChatMessages);
|
|||||||
const mockUpdateChatSession = vi.mocked(apiModule.updateChatSession);
|
const mockUpdateChatSession = vi.mocked(apiModule.updateChatSession);
|
||||||
const mockDeleteChatSession = vi.mocked(apiModule.deleteChatSession);
|
const mockDeleteChatSession = vi.mocked(apiModule.deleteChatSession);
|
||||||
const mockStreamChatResponse = vi.mocked(apiModule.streamChatResponse);
|
const mockStreamChatResponse = vi.mocked(apiModule.streamChatResponse);
|
||||||
|
const mockAttachChatStream = vi.mocked(apiModule.attachChatStream);
|
||||||
const mockCancelChatResponse = vi.mocked(apiModule.cancelChatResponse);
|
const mockCancelChatResponse = vi.mocked(apiModule.cancelChatResponse);
|
||||||
const mockFetchAgents = vi.mocked(apiModule.fetchAgents);
|
const mockFetchAgents = vi.mocked(apiModule.fetchAgents);
|
||||||
|
|
||||||
@@ -100,6 +102,7 @@ describe("useChat", () => {
|
|||||||
});
|
});
|
||||||
mockDeleteChatSession.mockResolvedValue({ success: true });
|
mockDeleteChatSession.mockResolvedValue({ success: true });
|
||||||
mockStreamChatResponse.mockReturnValue({ close: vi.fn(), isConnected: () => true });
|
mockStreamChatResponse.mockReturnValue({ close: vi.fn(), isConnected: () => true });
|
||||||
|
mockAttachChatStream.mockReturnValue({ close: vi.fn(), isConnected: () => true });
|
||||||
mockCancelChatResponse.mockResolvedValue({ success: true });
|
mockCancelChatResponse.mockResolvedValue({ success: true });
|
||||||
});
|
});
|
||||||
|
|
||||||
@@ -2066,18 +2069,23 @@ describe("useChat", () => {
|
|||||||
});
|
});
|
||||||
});
|
});
|
||||||
|
|
||||||
it("clears recovery streaming state when SSE delivers assistant message", async () => {
|
it("clears recovery streaming state when attach stream completes", async () => {
|
||||||
let subscribeHandler: Record<string, (event: MessageEvent) => void> = {};
|
|
||||||
mockSubscribeSse.mockImplementation((_url, options) => {
|
|
||||||
if (options?.events) {
|
|
||||||
subscribeHandler = options.events as typeof subscribeHandler;
|
|
||||||
}
|
|
||||||
return () => {};
|
|
||||||
});
|
|
||||||
|
|
||||||
const session = { ...makeSession({ id: "session-001", agentId: "agent-001" }), isGenerating: true };
|
const session = { ...makeSession({ id: "session-001", agentId: "agent-001" }), isGenerating: true };
|
||||||
mockFetchChatSessions.mockResolvedValueOnce({ sessions: [session] });
|
mockFetchChatSessions.mockResolvedValueOnce({ sessions: [session] });
|
||||||
mockFetchChatMessages.mockResolvedValue({ messages: [] });
|
mockFetchChatMessages.mockResolvedValue({
|
||||||
|
messages: [
|
||||||
|
makeMessage({
|
||||||
|
id: "msg-assistant-001",
|
||||||
|
sessionId: "session-001",
|
||||||
|
role: "assistant",
|
||||||
|
content: "Generated response",
|
||||||
|
}),
|
||||||
|
],
|
||||||
|
});
|
||||||
|
mockAttachChatStream.mockImplementation((_sessionId, handlers) => {
|
||||||
|
setTimeout(() => handlers.onDone?.({ messageId: "msg-assistant-001" }), 0);
|
||||||
|
return { close: vi.fn(), isConnected: () => true };
|
||||||
|
});
|
||||||
|
|
||||||
const { result } = renderHook(() => useChat("proj-123"));
|
const { result } = renderHook(() => useChat("proj-123"));
|
||||||
|
|
||||||
@@ -2089,24 +2097,6 @@ describe("useChat", () => {
|
|||||||
result.current.selectSession("session-001");
|
result.current.selectSession("session-001");
|
||||||
});
|
});
|
||||||
|
|
||||||
await waitFor(() => {
|
|
||||||
expect(result.current.isStreaming).toBe(true);
|
|
||||||
});
|
|
||||||
|
|
||||||
// Simulate SSE delivering the completed assistant message
|
|
||||||
const assistantMessage = makeMessage({
|
|
||||||
id: "msg-assistant-001",
|
|
||||||
sessionId: "session-001",
|
|
||||||
role: "assistant",
|
|
||||||
content: "Generated response",
|
|
||||||
});
|
|
||||||
|
|
||||||
act(() => {
|
|
||||||
subscribeHandler["chat:message:added"](
|
|
||||||
new MessageEvent("chat:message:added", { data: JSON.stringify(assistantMessage) }),
|
|
||||||
);
|
|
||||||
});
|
|
||||||
|
|
||||||
await waitFor(() => {
|
await waitFor(() => {
|
||||||
expect(result.current.isStreaming).toBe(false);
|
expect(result.current.isStreaming).toBe(false);
|
||||||
expect(result.current.streamingText).toBe("");
|
expect(result.current.streamingText).toBe("");
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ vi.mock("../../api", () => ({
|
|||||||
createChatSession: vi.fn(),
|
createChatSession: vi.fn(),
|
||||||
fetchChatMessages: vi.fn(),
|
fetchChatMessages: vi.fn(),
|
||||||
streamChatResponse: vi.fn(),
|
streamChatResponse: vi.fn(),
|
||||||
|
attachChatStream: vi.fn(),
|
||||||
cancelChatResponse: vi.fn(),
|
cancelChatResponse: vi.fn(),
|
||||||
}));
|
}));
|
||||||
|
|
||||||
@@ -20,6 +21,7 @@ const mockFetchChatSession = vi.mocked(apiModule.fetchChatSession);
|
|||||||
const mockCreateChatSession = vi.mocked(apiModule.createChatSession);
|
const mockCreateChatSession = vi.mocked(apiModule.createChatSession);
|
||||||
const mockFetchChatMessages = vi.mocked(apiModule.fetchChatMessages);
|
const mockFetchChatMessages = vi.mocked(apiModule.fetchChatMessages);
|
||||||
const mockStreamChatResponse = vi.mocked(apiModule.streamChatResponse);
|
const mockStreamChatResponse = vi.mocked(apiModule.streamChatResponse);
|
||||||
|
const mockAttachChatStream = vi.mocked(apiModule.attachChatStream);
|
||||||
const mockCancelChatResponse = vi.mocked(apiModule.cancelChatResponse);
|
const mockCancelChatResponse = vi.mocked(apiModule.cancelChatResponse);
|
||||||
|
|
||||||
function makeSession(overrides: Partial<ChatSession> & Pick<ChatSession, "id" | "agentId">): ChatSession {
|
function makeSession(overrides: Partial<ChatSession> & Pick<ChatSession, "id" | "agentId">): ChatSession {
|
||||||
@@ -57,6 +59,7 @@ describe("useQuickChat", () => {
|
|||||||
session: { ...makeSession({ id: "session-001", agentId: "agent-001" }), isGenerating: false },
|
session: { ...makeSession({ id: "session-001", agentId: "agent-001" }), isGenerating: false },
|
||||||
});
|
});
|
||||||
mockStreamChatResponse.mockReturnValue({ close: vi.fn(), isConnected: () => true });
|
mockStreamChatResponse.mockReturnValue({ close: vi.fn(), isConnected: () => true });
|
||||||
|
mockAttachChatStream.mockReturnValue({ close: vi.fn(), isConnected: () => true });
|
||||||
mockCancelChatResponse.mockResolvedValue({ success: true });
|
mockCancelChatResponse.mockResolvedValue({ success: true });
|
||||||
});
|
});
|
||||||
|
|
||||||
@@ -1305,22 +1308,18 @@ describe("useQuickChat", () => {
|
|||||||
});
|
});
|
||||||
});
|
});
|
||||||
|
|
||||||
it("clears recovery streaming state when polling detects generation complete", async () => {
|
it("clears recovery streaming state when attach stream completes", async () => {
|
||||||
vi.useFakeTimers({ shouldAdvanceTime: true });
|
|
||||||
|
|
||||||
const session = { ...makeSession({ id: "session-001", agentId: "agent-001" }), isGenerating: true };
|
const session = { ...makeSession({ id: "session-001", agentId: "agent-001" }), isGenerating: true };
|
||||||
mockFetchResumeChatSession.mockResolvedValue({ session });
|
mockFetchResumeChatSession.mockResolvedValue({ session });
|
||||||
mockFetchChatMessages.mockResolvedValue({ messages: [] });
|
|
||||||
|
|
||||||
// After first poll, server reports generation is done and has a new assistant message
|
|
||||||
mockFetchChatSession.mockResolvedValue({
|
|
||||||
session: { ...makeSession({ id: "session-001", agentId: "agent-001" }), isGenerating: false },
|
|
||||||
});
|
|
||||||
mockFetchChatMessages.mockResolvedValue({
|
mockFetchChatMessages.mockResolvedValue({
|
||||||
messages: [
|
messages: [
|
||||||
{ id: "msg-1", sessionId: "session-001", role: "assistant", content: "Done", thinkingOutput: null, metadata: null, createdAt: new Date().toISOString() },
|
{ id: "msg-1", sessionId: "session-001", role: "assistant", content: "Done", thinkingOutput: null, metadata: null, createdAt: new Date().toISOString() },
|
||||||
],
|
],
|
||||||
});
|
});
|
||||||
|
mockAttachChatStream.mockImplementation((_sessionId, handlers) => {
|
||||||
|
setTimeout(() => handlers.onDone?.({ messageId: "msg-1" }), 0);
|
||||||
|
return { close: vi.fn(), isConnected: () => true };
|
||||||
|
});
|
||||||
|
|
||||||
const { result } = renderHook(() => useQuickChat("proj-123"));
|
const { result } = renderHook(() => useQuickChat("proj-123"));
|
||||||
|
|
||||||
@@ -1328,22 +1327,11 @@ describe("useQuickChat", () => {
|
|||||||
await result.current.switchSession("agent-001");
|
await result.current.switchSession("agent-001");
|
||||||
});
|
});
|
||||||
|
|
||||||
await waitFor(() => {
|
|
||||||
expect(result.current.isStreaming).toBe(true);
|
|
||||||
});
|
|
||||||
|
|
||||||
// Advance time to trigger the polling interval (3s)
|
|
||||||
await act(async () => {
|
|
||||||
vi.advanceTimersByTime(3500);
|
|
||||||
});
|
|
||||||
|
|
||||||
await waitFor(() => {
|
await waitFor(() => {
|
||||||
expect(result.current.isStreaming).toBe(false);
|
expect(result.current.isStreaming).toBe(false);
|
||||||
expect(result.current.streamingText).toBe("");
|
expect(result.current.streamingText).toBe("");
|
||||||
expect(result.current.messages.some((m) => m.id === "msg-1")).toBe(true);
|
expect(result.current.messages.some((m) => m.id === "msg-1")).toBe(true);
|
||||||
});
|
});
|
||||||
|
|
||||||
vi.useRealTimers();
|
|
||||||
});
|
});
|
||||||
});
|
});
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import {
|
|||||||
fetchChatMessages,
|
fetchChatMessages,
|
||||||
updateChatSession,
|
updateChatSession,
|
||||||
deleteChatSession,
|
deleteChatSession,
|
||||||
|
attachChatStream,
|
||||||
streamChatResponse,
|
streamChatResponse,
|
||||||
cancelChatResponse,
|
cancelChatResponse,
|
||||||
fetchAgents,
|
fetchAgents,
|
||||||
@@ -329,6 +330,56 @@ export function useChat(
|
|||||||
setIsStreaming(false);
|
setIsStreaming(false);
|
||||||
}, []);
|
}, []);
|
||||||
|
|
||||||
|
const attachIfGenerating = useCallback((sessionId: string) => {
|
||||||
|
if (streamRef.current || !sessionId) {
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
|
||||||
|
cancelledByUserRef.current = false;
|
||||||
|
setIsStreaming(true);
|
||||||
|
|
||||||
|
const { handlers } = createChatStreamHandlers({
|
||||||
|
sessionId,
|
||||||
|
tempUserMessageId: "",
|
||||||
|
setStreamingText,
|
||||||
|
setStreamingThinking,
|
||||||
|
setStreamingToolCalls,
|
||||||
|
cancelStreamingFlushesRef,
|
||||||
|
addToast,
|
||||||
|
onFallbackSession: (data, fallbackSessionId) => {
|
||||||
|
const nextModel = parseModelDescriptor(data.fallbackModel);
|
||||||
|
setSessions((prev) => prev.map((session) =>
|
||||||
|
session.id === fallbackSessionId ? { ...session, ...nextModel } : session,
|
||||||
|
));
|
||||||
|
setActiveSession((prev) => prev && prev.id === fallbackSessionId ? { ...prev, ...nextModel } : prev);
|
||||||
|
},
|
||||||
|
onDone: () => {
|
||||||
|
setStreamingText("");
|
||||||
|
setStreamingThinking("");
|
||||||
|
setStreamingToolCalls([]);
|
||||||
|
setIsStreaming(false);
|
||||||
|
isStreamingRef.current = false;
|
||||||
|
streamRef.current = null;
|
||||||
|
void loadMessages(sessionId);
|
||||||
|
},
|
||||||
|
onError: (data) => {
|
||||||
|
setStreamingText("");
|
||||||
|
setStreamingThinking("");
|
||||||
|
setStreamingToolCalls([]);
|
||||||
|
setIsStreaming(false);
|
||||||
|
isStreamingRef.current = false;
|
||||||
|
streamRef.current = null;
|
||||||
|
const errorMessage = typeof data === "string" && data.trim() ? data : "Failed to get response";
|
||||||
|
addToast?.(errorMessage, "error");
|
||||||
|
void loadMessages(sessionId);
|
||||||
|
},
|
||||||
|
});
|
||||||
|
|
||||||
|
const stream = attachChatStream(sessionId, handlers, projectId);
|
||||||
|
streamRef.current = stream;
|
||||||
|
return true;
|
||||||
|
}, [addToast, loadMessages, projectId]);
|
||||||
|
|
||||||
// Select a session
|
// Select a session
|
||||||
const selectSession = useCallback(
|
const selectSession = useCallback(
|
||||||
(id: string, sessionOverride?: ChatSessionInfo) => {
|
(id: string, sessionOverride?: ChatSessionInfo) => {
|
||||||
@@ -363,8 +414,8 @@ export function useChat(
|
|||||||
// all streaming state. Showing "Connecting…" immediately tells the
|
// all streaming state. Showing "Connecting…" immediately tells the
|
||||||
// user the AI is still working.
|
// user the AI is still working.
|
||||||
if (session?.isGenerating) {
|
if (session?.isGenerating) {
|
||||||
setIsStreaming(true);
|
|
||||||
setStreamingText("");
|
setStreamingText("");
|
||||||
|
attachIfGenerating(session.id);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Persist active session to localStorage
|
// Persist active session to localStorage
|
||||||
@@ -374,7 +425,7 @@ export function useChat(
|
|||||||
removeScopedItem(ACTIVE_SESSION_STORAGE_KEY, projectId);
|
removeScopedItem(ACTIVE_SESSION_STORAGE_KEY, projectId);
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
[sessions, loadMessages, projectId, resetTransientComposerState],
|
[attachIfGenerating, sessions, loadMessages, projectId, resetTransientComposerState],
|
||||||
);
|
);
|
||||||
|
|
||||||
// Update the ref to point to the actual selectSession function
|
// Update the ref to point to the actual selectSession function
|
||||||
@@ -637,7 +688,13 @@ export function useChat(
|
|||||||
// Recovery mode polling: if reloaded mid-generation, keep waiting state alive
|
// Recovery mode polling: if reloaded mid-generation, keep waiting state alive
|
||||||
// until generation finishes and messages can be reloaded.
|
// until generation finishes and messages can be reloaded.
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
if (!isStreaming || streamRef.current || !activeSessionRef.current) return;
|
if (!activeSessionRef.current?.isGenerating) return;
|
||||||
|
|
||||||
|
if (!streamRef.current) {
|
||||||
|
attachIfGenerating(activeSessionRef.current.id);
|
||||||
|
}
|
||||||
|
|
||||||
|
if (!isStreamingRef.current || streamRef.current || !activeSessionRef.current) return;
|
||||||
|
|
||||||
const interval = setInterval(async () => {
|
const interval = setInterval(async () => {
|
||||||
if (!isStreamingRef.current || streamRef.current || !activeSessionRef.current) {
|
if (!isStreamingRef.current || streamRef.current || !activeSessionRef.current) {
|
||||||
@@ -662,7 +719,7 @@ export function useChat(
|
|||||||
}, 3000);
|
}, 3000);
|
||||||
|
|
||||||
return () => clearInterval(interval);
|
return () => clearInterval(interval);
|
||||||
}, [isStreaming, loadMessages, projectId]);
|
}, [attachIfGenerating, loadMessages, projectId, activeSession]);
|
||||||
|
|
||||||
// SSE real-time updates
|
// SSE real-time updates
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
@@ -692,6 +749,10 @@ export function useChat(
|
|||||||
// If this is the active session, update it too
|
// If this is the active session, update it too
|
||||||
if (activeSessionRef.current?.id === updatedSession.id) {
|
if (activeSessionRef.current?.id === updatedSession.id) {
|
||||||
setActiveSession(updatedSession);
|
setActiveSession(updatedSession);
|
||||||
|
if (updatedSession.isGenerating && !streamRef.current) {
|
||||||
|
setStreamingText("");
|
||||||
|
attachIfGenerating(updatedSession.id);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -782,7 +843,7 @@ export function useChat(
|
|||||||
});
|
});
|
||||||
|
|
||||||
return unsubscribe;
|
return unsubscribe;
|
||||||
}, [projectId]);
|
}, [attachIfGenerating, projectId]);
|
||||||
|
|
||||||
// Cleanup on unmount
|
// Cleanup on unmount
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import {
|
|||||||
fetchChatSession,
|
fetchChatSession,
|
||||||
createChatSession,
|
createChatSession,
|
||||||
fetchChatMessages,
|
fetchChatMessages,
|
||||||
|
attachChatStream,
|
||||||
streamChatResponse,
|
streamChatResponse,
|
||||||
cancelChatResponse,
|
cancelChatResponse,
|
||||||
} from "../api";
|
} from "../api";
|
||||||
@@ -268,6 +269,59 @@ export function useQuickChat(
|
|||||||
[projectId],
|
[projectId],
|
||||||
);
|
);
|
||||||
|
|
||||||
|
const attachIfGenerating = useCallback((sessionId: string) => {
|
||||||
|
if (streamRef.current || !sessionId) {
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
|
||||||
|
cancelledByUserRef.current = false;
|
||||||
|
setIsStreaming(true);
|
||||||
|
|
||||||
|
const { handlers } = createChatStreamHandlers({
|
||||||
|
sessionId,
|
||||||
|
tempUserMessageId: "",
|
||||||
|
setStreamingText,
|
||||||
|
setStreamingThinking,
|
||||||
|
setStreamingToolCalls,
|
||||||
|
cancelStreamingFlushesRef,
|
||||||
|
addToast,
|
||||||
|
onFallbackSession: (data, fallbackSessionId) => {
|
||||||
|
const nextModel = parseModelDescriptor(data.fallbackModel);
|
||||||
|
setSessions((prev) => prev.map((session) =>
|
||||||
|
session.id === fallbackSessionId ? { ...session, ...nextModel } : session,
|
||||||
|
));
|
||||||
|
setActiveSession((prev) => prev && prev.id === fallbackSessionId ? { ...prev, ...nextModel } : prev);
|
||||||
|
},
|
||||||
|
onDone: () => {
|
||||||
|
setStreamingText("");
|
||||||
|
setStreamingThinking("");
|
||||||
|
setStreamingToolCalls([]);
|
||||||
|
setIsStreaming(false);
|
||||||
|
isStreamingRef.current = false;
|
||||||
|
streamRef.current = null;
|
||||||
|
void fetchChatMessages(sessionId, { limit: 50 }, projectId).then((data) => {
|
||||||
|
setMessages(data.messages.map(mapChatMessageToInfo));
|
||||||
|
}).catch(() => {});
|
||||||
|
},
|
||||||
|
onError: (data) => {
|
||||||
|
setStreamingText("");
|
||||||
|
setStreamingThinking("");
|
||||||
|
setStreamingToolCalls([]);
|
||||||
|
setIsStreaming(false);
|
||||||
|
isStreamingRef.current = false;
|
||||||
|
streamRef.current = null;
|
||||||
|
const errorMessage = typeof data === "string" && data.trim() ? data : "Failed to get response";
|
||||||
|
addToast?.(errorMessage, "error");
|
||||||
|
void fetchChatMessages(sessionId, { limit: 50 }, projectId).then((resp) => {
|
||||||
|
setMessages(resp.messages.map(mapChatMessageToInfo));
|
||||||
|
}).catch(() => {});
|
||||||
|
},
|
||||||
|
});
|
||||||
|
|
||||||
|
streamRef.current = attachChatStream(sessionId, handlers, projectId);
|
||||||
|
return true;
|
||||||
|
}, [addToast, projectId]);
|
||||||
|
|
||||||
// Fetch existing sessions and find/create one for the given target
|
// Fetch existing sessions and find/create one for the given target
|
||||||
const initializeSession = useCallback(
|
const initializeSession = useCallback(
|
||||||
async (agentId: string, modelProvider?: string, modelId?: string) => {
|
async (agentId: string, modelProvider?: string, modelId?: string) => {
|
||||||
@@ -295,8 +349,8 @@ export function useQuickChat(
|
|||||||
// After a reload/HMR, the server keeps generating but the UI loses
|
// After a reload/HMR, the server keeps generating but the UI loses
|
||||||
// all streaming state. Show the "Connecting…" indicator immediately.
|
// all streaming state. Show the "Connecting…" indicator immediately.
|
||||||
if (existingSession.isGenerating) {
|
if (existingSession.isGenerating) {
|
||||||
setIsStreaming(true);
|
|
||||||
setStreamingText("");
|
setStreamingText("");
|
||||||
|
attachIfGenerating(existingSession.id);
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
const newSession = await createSessionForTarget(target);
|
const newSession = await createSessionForTarget(target);
|
||||||
@@ -319,7 +373,7 @@ export function useQuickChat(
|
|||||||
setSessionsLoading(false);
|
setSessionsLoading(false);
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
[projectId, addToast, createSessionForTarget],
|
[attachIfGenerating, projectId, addToast, createSessionForTarget],
|
||||||
);
|
);
|
||||||
|
|
||||||
// Load messages for the active session
|
// Load messages for the active session
|
||||||
@@ -352,7 +406,13 @@ export function useQuickChat(
|
|||||||
// Poll every 3s until the server reports isGenerating=false, then reload messages
|
// Poll every 3s until the server reports isGenerating=false, then reload messages
|
||||||
// and clear streaming state.
|
// and clear streaming state.
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
if (!isStreaming || streamRef.current || !activeSession) return;
|
if (!activeSession?.isGenerating) return;
|
||||||
|
|
||||||
|
if (!streamRef.current) {
|
||||||
|
attachIfGenerating(activeSession.id);
|
||||||
|
}
|
||||||
|
|
||||||
|
if (!isStreamingRef.current || streamRef.current || !activeSession) return;
|
||||||
|
|
||||||
const interval = setInterval(async () => {
|
const interval = setInterval(async () => {
|
||||||
// Re-check conditions inside the callback (state may have changed)
|
// Re-check conditions inside the callback (state may have changed)
|
||||||
@@ -379,7 +439,7 @@ export function useQuickChat(
|
|||||||
}, 3000);
|
}, 3000);
|
||||||
|
|
||||||
return () => clearInterval(interval);
|
return () => clearInterval(interval);
|
||||||
}, [isStreaming, activeSession, projectId]);
|
}, [activeSession, attachIfGenerating, projectId]);
|
||||||
|
|
||||||
// Reload messages from server (for same-session revisit)
|
// Reload messages from server (for same-session revisit)
|
||||||
const reloadMessages = useCallback(async () => {
|
const reloadMessages = useCallback(async () => {
|
||||||
@@ -463,7 +523,11 @@ export function useQuickChat(
|
|||||||
|
|
||||||
resetTransientComposerState();
|
resetTransientComposerState();
|
||||||
setActiveSession(session);
|
setActiveSession(session);
|
||||||
}, [resetTransientComposerState]);
|
if (session.isGenerating) {
|
||||||
|
setStreamingText("");
|
||||||
|
attachIfGenerating(session.id);
|
||||||
|
}
|
||||||
|
}, [attachIfGenerating, resetTransientComposerState]);
|
||||||
|
|
||||||
const startModelChat = useCallback(
|
const startModelChat = useCallback(
|
||||||
async (modelProvider: string, modelId: string) => {
|
async (modelProvider: string, modelId: string) => {
|
||||||
|
|||||||
@@ -48,25 +48,37 @@ function createSSEResponse(): {
|
|||||||
const mockInit = vi.fn().mockResolvedValue(undefined);
|
const mockInit = vi.fn().mockResolvedValue(undefined);
|
||||||
|
|
||||||
// Create mock functions before vi.mock
|
// Create mock functions before vi.mock
|
||||||
const { mockCreateFnAgent, mockChatStreamManager, mockSendMessage, mockCancelGeneration, mockBeginGeneration } = vi.hoisted(() => {
|
const {
|
||||||
|
mockCreateFnAgent,
|
||||||
|
mockChatStreamManager,
|
||||||
|
mockSendMessage,
|
||||||
|
mockCancelGeneration,
|
||||||
|
mockBeginGeneration,
|
||||||
|
mockIsGenerating,
|
||||||
|
mockGetActiveGenerationId,
|
||||||
|
} = vi.hoisted(() => {
|
||||||
// Store subscribers per session for broadcast simulation
|
// Store subscribers per session for broadcast simulation
|
||||||
const subscribers = new Map<string, Set<(event: any, eventId?: number) => void>>();
|
const subscribers = new Map<string, Set<{ callback: (event: any, eventId?: number) => void; generationId?: number }>>();
|
||||||
|
|
||||||
const chatStreamManager = {
|
const chatStreamManager = {
|
||||||
subscribe: vi.fn((sessionId: string, callback: (event: any, eventId?: number) => void) => {
|
subscribe: vi.fn((sessionId: string, callback: (event: any, eventId?: number) => void, options?: { generationId?: number }) => {
|
||||||
if (!subscribers.has(sessionId)) {
|
if (!subscribers.has(sessionId)) {
|
||||||
subscribers.set(sessionId, new Set());
|
subscribers.set(sessionId, new Set());
|
||||||
}
|
}
|
||||||
subscribers.get(sessionId)!.add(callback);
|
const entry = { callback, generationId: options?.generationId };
|
||||||
|
subscribers.get(sessionId)!.add(entry);
|
||||||
return () => {
|
return () => {
|
||||||
subscribers.get(sessionId)?.delete(callback);
|
subscribers.get(sessionId)?.delete(entry);
|
||||||
};
|
};
|
||||||
}),
|
}),
|
||||||
broadcast: vi.fn((sessionId: string, event: any) => {
|
broadcast: vi.fn((sessionId: string, event: any, options?: { generationId?: number }) => {
|
||||||
const callbacks = subscribers.get(sessionId);
|
const callbacks = subscribers.get(sessionId);
|
||||||
if (callbacks) {
|
if (callbacks) {
|
||||||
let eventId = 1;
|
let eventId = 1;
|
||||||
for (const callback of callbacks) {
|
for (const { callback, generationId } of callbacks) {
|
||||||
|
if (options?.generationId !== undefined && generationId !== undefined && options.generationId !== generationId) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
callback(event, eventId++);
|
callback(event, eventId++);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -88,7 +100,7 @@ const { mockCreateFnAgent, mockChatStreamManager, mockSendMessage, mockCancelGen
|
|||||||
__triggerDone: (sessionId: string, messageId: string) => {
|
__triggerDone: (sessionId: string, messageId: string) => {
|
||||||
const callbacks = subscribers.get(sessionId);
|
const callbacks = subscribers.get(sessionId);
|
||||||
if (callbacks) {
|
if (callbacks) {
|
||||||
for (const callback of callbacks) {
|
for (const { callback } of callbacks) {
|
||||||
callback({ type: "done", data: { messageId } }, 1);
|
callback({ type: "done", data: { messageId } }, 1);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -96,7 +108,7 @@ const { mockCreateFnAgent, mockChatStreamManager, mockSendMessage, mockCancelGen
|
|||||||
__triggerError: (sessionId: string, error: string) => {
|
__triggerError: (sessionId: string, error: string) => {
|
||||||
const callbacks = subscribers.get(sessionId);
|
const callbacks = subscribers.get(sessionId);
|
||||||
if (callbacks) {
|
if (callbacks) {
|
||||||
for (const callback of callbacks) {
|
for (const { callback } of callbacks) {
|
||||||
callback({ type: "error", data: error }, 1);
|
callback({ type: "error", data: error }, 1);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -108,6 +120,8 @@ const { mockCreateFnAgent, mockChatStreamManager, mockSendMessage, mockCancelGen
|
|||||||
mockSendMessage: vi.fn(),
|
mockSendMessage: vi.fn(),
|
||||||
mockCancelGeneration: vi.fn(),
|
mockCancelGeneration: vi.fn(),
|
||||||
mockBeginGeneration: vi.fn(() => ({ generationId: 1, abortController: new AbortController() })),
|
mockBeginGeneration: vi.fn(() => ({ generationId: 1, abortController: new AbortController() })),
|
||||||
|
mockIsGenerating: vi.fn(() => false),
|
||||||
|
mockGetActiveGenerationId: vi.fn(() => undefined),
|
||||||
mockChatStreamManager: chatStreamManager,
|
mockChatStreamManager: chatStreamManager,
|
||||||
};
|
};
|
||||||
});
|
});
|
||||||
@@ -164,6 +178,8 @@ vi.mock("../chat.js", () => {
|
|||||||
sendMessage = mockSendMessage;
|
sendMessage = mockSendMessage;
|
||||||
cancelGeneration = mockCancelGeneration;
|
cancelGeneration = mockCancelGeneration;
|
||||||
beginGeneration = mockBeginGeneration;
|
beginGeneration = mockBeginGeneration;
|
||||||
|
isGenerating = mockIsGenerating;
|
||||||
|
getActiveGenerationId = mockGetActiveGenerationId;
|
||||||
},
|
},
|
||||||
chatStreamManager: mockChatStreamManager,
|
chatStreamManager: mockChatStreamManager,
|
||||||
checkRateLimit: vi.fn().mockReturnValue(true),
|
checkRateLimit: vi.fn().mockReturnValue(true),
|
||||||
@@ -266,6 +282,8 @@ function createMockChatManager() {
|
|||||||
sendMessage: mockSendMessage,
|
sendMessage: mockSendMessage,
|
||||||
cancelGeneration: mockCancelGeneration,
|
cancelGeneration: mockCancelGeneration,
|
||||||
beginGeneration: mockBeginGeneration,
|
beginGeneration: mockBeginGeneration,
|
||||||
|
isGenerating: mockIsGenerating,
|
||||||
|
getActiveGenerationId: mockGetActiveGenerationId,
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -314,6 +332,8 @@ describe("Chat API Routes", () => {
|
|||||||
mockDeleteMessage.mockReset();
|
mockDeleteMessage.mockReset();
|
||||||
mockSendMessage.mockReset();
|
mockSendMessage.mockReset();
|
||||||
mockCancelGeneration.mockReset();
|
mockCancelGeneration.mockReset();
|
||||||
|
mockIsGenerating.mockReset();
|
||||||
|
mockGetActiveGenerationId.mockReset();
|
||||||
mockAgentStoreInit.mockResolvedValue(undefined);
|
mockAgentStoreInit.mockResolvedValue(undefined);
|
||||||
mockAgentStoreGetAgent.mockReset();
|
mockAgentStoreGetAgent.mockReset();
|
||||||
mockGetOrCreateProjectStore.mockReset();
|
mockGetOrCreateProjectStore.mockReset();
|
||||||
@@ -323,6 +343,8 @@ describe("Chat API Routes", () => {
|
|||||||
mockGetMessages.mockReturnValue([]);
|
mockGetMessages.mockReturnValue([]);
|
||||||
mockGetLastMessageForSessions.mockReturnValue(new Map());
|
mockGetLastMessageForSessions.mockReturnValue(new Map());
|
||||||
mockCancelGeneration.mockReturnValue(false);
|
mockCancelGeneration.mockReturnValue(false);
|
||||||
|
mockIsGenerating.mockReturnValue(false);
|
||||||
|
mockGetActiveGenerationId.mockReturnValue(undefined);
|
||||||
|
|
||||||
// Default agent mock - agent with model config
|
// Default agent mock - agent with model config
|
||||||
mockAgentStoreGetAgent.mockResolvedValue({
|
mockAgentStoreGetAgent.mockResolvedValue({
|
||||||
@@ -1015,30 +1037,27 @@ describe("Chat API Routes", () => {
|
|||||||
store: any,
|
store: any,
|
||||||
chatStore: any,
|
chatStore: any,
|
||||||
chatManager: any,
|
chatManager: any,
|
||||||
|
routePath = "/chat/sessions/:id/messages",
|
||||||
|
method: "get" | "post" = "post",
|
||||||
): Promise<void> {
|
): Promise<void> {
|
||||||
// Dynamically import to get the current module state (with mocks applied)
|
|
||||||
const { createApiRoutes } = await import("../routes.js");
|
const { createApiRoutes } = await import("../routes.js");
|
||||||
const router = createApiRoutes(store, {
|
const router = createApiRoutes(store, {
|
||||||
chatStore,
|
chatStore,
|
||||||
chatManager,
|
chatManager,
|
||||||
});
|
});
|
||||||
|
|
||||||
// Find the SSE route handler
|
|
||||||
const stack = router.stack || [];
|
const stack = router.stack || [];
|
||||||
const handler = stack.find(
|
const handler = stack.find(
|
||||||
(layer: any) =>
|
(layer: any) =>
|
||||||
layer.route?.path === "/chat/sessions/:id/messages" &&
|
layer.route?.path === routePath &&
|
||||||
layer.route?.methods?.post,
|
layer.route?.methods?.[method],
|
||||||
);
|
);
|
||||||
|
|
||||||
if (!handler) {
|
if (!handler) {
|
||||||
throw new Error(`SSE route handler not found. Stack has ${stack.length} layers.`);
|
throw new Error(`SSE route handler not found (${method.toUpperCase()} ${routePath}). Stack has ${stack.length} layers.`);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get the actual handler function from the layer
|
|
||||||
const routeHandler = handler.route.stack[handler.route.stack.length - 1].handle;
|
const routeHandler = handler.route.stack[handler.route.stack.length - 1].handle;
|
||||||
|
|
||||||
// The handler is wrapped in middleware (rateLimit), so we need to call next
|
|
||||||
const next = vi.fn();
|
const next = vi.fn();
|
||||||
await routeHandler(req, res, next);
|
await routeHandler(req, res, next);
|
||||||
}
|
}
|
||||||
@@ -1104,6 +1123,117 @@ describe("Chat API Routes", () => {
|
|||||||
});
|
});
|
||||||
|
|
||||||
|
|
||||||
|
it("attach stream returns 404 for unknown session", async () => {
|
||||||
|
mockGetSession.mockReturnValue(null);
|
||||||
|
|
||||||
|
const response = await request(
|
||||||
|
app,
|
||||||
|
"GET",
|
||||||
|
"/api/chat/sessions/chat-missing/stream",
|
||||||
|
);
|
||||||
|
|
||||||
|
expect(response.status).toBe(404);
|
||||||
|
});
|
||||||
|
|
||||||
|
it("attach stream replays buffered events and ends when not generating", async () => {
|
||||||
|
mockGetSession.mockReturnValue(sampleSession);
|
||||||
|
mockIsGenerating.mockReturnValue(false);
|
||||||
|
mockChatStreamManager.getBufferedEvents.mockReturnValue([
|
||||||
|
{ id: 2, event: "text", data: JSON.stringify("hello") },
|
||||||
|
{ id: 3, event: "done", data: JSON.stringify({ messageId: "msg-1" }) },
|
||||||
|
]);
|
||||||
|
|
||||||
|
const req = createSSERequest();
|
||||||
|
const { res, chunks } = createSSEResponse();
|
||||||
|
req.params = { id: "chat-abc123" };
|
||||||
|
req.query = {} as any;
|
||||||
|
req.headers = {} as any;
|
||||||
|
|
||||||
|
await invokeSSEHandler(req, res, store, mockChatStore, mockChatManager, "/chat/sessions/:id/stream", "get");
|
||||||
|
|
||||||
|
const output = chunks.join("");
|
||||||
|
expect(output).toContain("id: 2");
|
||||||
|
expect(output).toContain("event: text");
|
||||||
|
expect(output).toContain("id: 3");
|
||||||
|
expect(output).toContain("event: done");
|
||||||
|
expect(res.end).toHaveBeenCalled();
|
||||||
|
expect(mockChatStreamManager.subscribe).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it("attach stream replays buffered events and receives live generation events", async () => {
|
||||||
|
mockGetSession.mockReturnValue(sampleSession);
|
||||||
|
mockIsGenerating.mockReturnValue(true);
|
||||||
|
mockGetActiveGenerationId.mockReturnValue(42);
|
||||||
|
mockChatStreamManager.getBufferedEvents.mockReturnValue([
|
||||||
|
{ id: 5, event: "text", data: JSON.stringify("buffer") },
|
||||||
|
]);
|
||||||
|
|
||||||
|
const req = createSSERequest();
|
||||||
|
const { res, chunks } = createSSEResponse();
|
||||||
|
req.params = { id: "chat-abc123" };
|
||||||
|
req.query = {} as any;
|
||||||
|
req.headers = {} as any;
|
||||||
|
|
||||||
|
await invokeSSEHandler(req, res, store, mockChatStore, mockChatManager, "/chat/sessions/:id/stream", "get");
|
||||||
|
expect(mockChatStreamManager.subscribe).toHaveBeenCalledWith(
|
||||||
|
"chat-abc123",
|
||||||
|
expect.any(Function),
|
||||||
|
{ generationId: 42 },
|
||||||
|
);
|
||||||
|
|
||||||
|
const subscriber = mockChatStreamManager.subscribe.mock.calls.at(-1)?.[1] as ((event: any, id?: number) => void);
|
||||||
|
subscriber({ type: "text", data: "live" }, 6);
|
||||||
|
subscriber({ type: "done", data: { messageId: "msg-2" } }, 7);
|
||||||
|
|
||||||
|
const output = chunks.join("");
|
||||||
|
expect(output).toContain("id: 5");
|
||||||
|
expect(output).toContain("data: \"buffer\"");
|
||||||
|
expect(output).toContain("data: \"live\"");
|
||||||
|
});
|
||||||
|
|
||||||
|
it("attach stream honors Last-Event-ID replay cutoff", async () => {
|
||||||
|
mockGetSession.mockReturnValue(sampleSession);
|
||||||
|
mockIsGenerating.mockReturnValue(false);
|
||||||
|
|
||||||
|
const req = createSSERequest();
|
||||||
|
const { res } = createSSEResponse();
|
||||||
|
req.params = { id: "chat-abc123" };
|
||||||
|
req.query = {} as any;
|
||||||
|
req.headers = { "last-event-id": "9" } as any;
|
||||||
|
|
||||||
|
await invokeSSEHandler(req, res, store, mockChatStore, mockChatManager, "/chat/sessions/:id/stream", "get");
|
||||||
|
|
||||||
|
expect(mockChatStreamManager.getBufferedEvents).toHaveBeenCalledWith("chat-abc123", 9);
|
||||||
|
});
|
||||||
|
|
||||||
|
it("attach stream filters out events from a different generation", async () => {
|
||||||
|
mockGetSession.mockReturnValue(sampleSession);
|
||||||
|
mockIsGenerating.mockReturnValue(true);
|
||||||
|
mockGetActiveGenerationId.mockReturnValue(7);
|
||||||
|
mockChatStreamManager.getBufferedEvents.mockReturnValue([]);
|
||||||
|
|
||||||
|
const req = createSSERequest();
|
||||||
|
const { res, chunks } = createSSEResponse();
|
||||||
|
req.params = { id: "chat-abc123" };
|
||||||
|
req.query = {} as any;
|
||||||
|
req.headers = {} as any;
|
||||||
|
|
||||||
|
await invokeSSEHandler(req, res, store, mockChatStore, mockChatManager, "/chat/sessions/:id/stream", "get");
|
||||||
|
mockChatStreamManager.broadcast("chat-abc123", { type: "text", data: "wrong" }, { generationId: 8 });
|
||||||
|
const subscriber = mockChatStreamManager.subscribe.mock.calls.at(-1)?.[1] as ((event: any, id?: number) => void);
|
||||||
|
subscriber({ type: "text", data: "right" }, 3);
|
||||||
|
subscriber({ type: "done", data: { messageId: "msg-3" } }, 4);
|
||||||
|
|
||||||
|
const output = chunks.join("");
|
||||||
|
expect(mockChatStreamManager.subscribe).toHaveBeenCalledWith(
|
||||||
|
"chat-abc123",
|
||||||
|
expect.any(Function),
|
||||||
|
{ generationId: 7 },
|
||||||
|
);
|
||||||
|
expect(output).toContain("right");
|
||||||
|
expect(output).not.toContain("wrong");
|
||||||
|
});
|
||||||
|
|
||||||
it("SSE route passes through tool_start and tool_end events", async () => {
|
it("SSE route passes through tool_start and tool_end events", async () => {
|
||||||
mockGetSession.mockReturnValue(sampleSession);
|
mockGetSession.mockReturnValue(sampleSession);
|
||||||
|
|
||||||
|
|||||||
@@ -1225,6 +1225,13 @@ export class ChatManager {
|
|||||||
return this.activeGenerations.has(sessionId);
|
return this.activeGenerations.has(sessionId);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Return the active generation ID for a session, if any.
|
||||||
|
*/
|
||||||
|
getActiveGenerationId(sessionId: string): number | undefined {
|
||||||
|
return this.activeGenerations.get(sessionId)?.generationId;
|
||||||
|
}
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* Return all session IDs that currently have an active generation.
|
* Return all session IDs that currently have an active generation.
|
||||||
* Useful for batch-enriching session lists without N+1 lookups.
|
* Useful for batch-enriching session lists without N+1 lookups.
|
||||||
|
|||||||
@@ -976,6 +976,7 @@ export function createApiRoutes(store: TaskStore, options?: ServerOptions): Rout
|
|||||||
});
|
});
|
||||||
registerChatRoutes(routeContext, {
|
registerChatRoutes(routeContext, {
|
||||||
parseLastEventId,
|
parseLastEventId,
|
||||||
|
replayBufferedSSE,
|
||||||
validateOptionalModelField,
|
validateOptionalModelField,
|
||||||
upload,
|
upload,
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -5,11 +5,12 @@ import { basename, join, resolve } from "node:path";
|
|||||||
import type { EnrichedChatSession, ChatAttachment } from "@fusion/core";
|
import type { EnrichedChatSession, ChatAttachment } from "@fusion/core";
|
||||||
import { ApiError, badRequest, internalError, notFound } from "../api-error.js";
|
import { ApiError, badRequest, internalError, notFound } from "../api-error.js";
|
||||||
import { rateLimit, RATE_LIMITS } from "../rate-limit.js";
|
import { rateLimit, RATE_LIMITS } from "../rate-limit.js";
|
||||||
import { writeSSEEvent } from "../sse-buffer.js";
|
import { writeSSEEvent, type SessionBufferedEvent } from "../sse-buffer.js";
|
||||||
import type { ApiRoutesContext } from "./types.js";
|
import type { ApiRoutesContext } from "./types.js";
|
||||||
|
|
||||||
interface ChatRouteDeps {
|
interface ChatRouteDeps {
|
||||||
parseLastEventId: (req: import("express").Request) => number | undefined;
|
parseLastEventId: (req: import("express").Request) => number | undefined;
|
||||||
|
replayBufferedSSE: (res: import("express").Response, bufferedEvents: SessionBufferedEvent[]) => boolean;
|
||||||
validateOptionalModelField: (value: unknown, fieldName: string) => string | undefined;
|
validateOptionalModelField: (value: unknown, fieldName: string) => string | undefined;
|
||||||
upload: import("multer").Multer;
|
upload: import("multer").Multer;
|
||||||
}
|
}
|
||||||
@@ -41,7 +42,7 @@ function resolveAttachmentPath(rootDir: string, sessionId: string, filename: str
|
|||||||
|
|
||||||
export function registerChatRoutes(ctx: ApiRoutesContext, deps: ChatRouteDeps): void {
|
export function registerChatRoutes(ctx: ApiRoutesContext, deps: ChatRouteDeps): void {
|
||||||
const { router, options, getProjectContext, chatLogger, rethrowAsApiError } = ctx;
|
const { router, options, getProjectContext, chatLogger, rethrowAsApiError } = ctx;
|
||||||
const { parseLastEventId, validateOptionalModelField, upload } = deps;
|
const { parseLastEventId, replayBufferedSSE, validateOptionalModelField, upload } = deps;
|
||||||
|
|
||||||
const uploadChatAttachment: import("express").RequestHandler = (req, res, next) => {
|
const uploadChatAttachment: import("express").RequestHandler = (req, res, next) => {
|
||||||
upload.single("file")(req, res, (err?: unknown) => {
|
upload.single("file")(req, res, (err?: unknown) => {
|
||||||
@@ -463,6 +464,88 @@ export function registerChatRoutes(ctx: ApiRoutesContext, deps: ChatRouteDeps):
|
|||||||
}
|
}
|
||||||
});
|
});
|
||||||
|
|
||||||
|
/**
|
||||||
|
* GET /api/chat/sessions/:id/stream
|
||||||
|
* Attach to an in-flight generation stream for an existing session.
|
||||||
|
*/
|
||||||
|
router.get("/chat/sessions/:id/stream", rateLimit(RATE_LIMITS.sse), async (req, res) => {
|
||||||
|
try {
|
||||||
|
const chatStore = options?.chatStore;
|
||||||
|
const chatManager = options?.chatManager;
|
||||||
|
if (!chatStore || !chatManager) {
|
||||||
|
throw internalError("Chat store or manager not available");
|
||||||
|
}
|
||||||
|
|
||||||
|
const sessionId = String(req.params.id);
|
||||||
|
const session = chatStore.getSession(sessionId);
|
||||||
|
if (!session) {
|
||||||
|
throw notFound(`Chat session ${sessionId} not found`);
|
||||||
|
}
|
||||||
|
|
||||||
|
const { projectId } = await getProjectContext(req);
|
||||||
|
if (projectId !== undefined && session.projectId !== projectId) {
|
||||||
|
throw notFound(`Chat session ${sessionId} not found`);
|
||||||
|
}
|
||||||
|
|
||||||
|
res.setHeader("Content-Type", "text/event-stream");
|
||||||
|
res.setHeader("Cache-Control", "no-cache");
|
||||||
|
res.setHeader("Connection", "keep-alive");
|
||||||
|
res.setHeader("X-Accel-Buffering", "no");
|
||||||
|
res.flushHeaders();
|
||||||
|
res.write(": connected\n\n");
|
||||||
|
|
||||||
|
const { chatStreamManager } = await import("../chat.js");
|
||||||
|
const lastEventId = parseLastEventId(req);
|
||||||
|
const buffered = chatStreamManager.getBufferedEvents(sessionId, lastEventId ?? 0);
|
||||||
|
if (!replayBufferedSSE(res, buffered)) {
|
||||||
|
res.end();
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
if (!chatManager.isGenerating(sessionId)) {
|
||||||
|
res.end();
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
const generationId = chatManager.getActiveGenerationId(sessionId);
|
||||||
|
if (generationId === undefined) {
|
||||||
|
res.end();
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
const unsubscribe = chatStreamManager.subscribe(sessionId, (event, eventId) => {
|
||||||
|
const data = (event as { data?: unknown }).data;
|
||||||
|
if (!writeSSEEvent(res, event.type, JSON.stringify(data ?? {}), eventId)) {
|
||||||
|
unsubscribe();
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
if (event.type === "done" || event.type === "error") {
|
||||||
|
unsubscribe();
|
||||||
|
res.end();
|
||||||
|
}
|
||||||
|
}, { generationId });
|
||||||
|
|
||||||
|
const heartbeat = setInterval(() => {
|
||||||
|
if (res.writableEnded) {
|
||||||
|
clearInterval(heartbeat);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
res.write(": heartbeat\n\n");
|
||||||
|
}, 30_000);
|
||||||
|
|
||||||
|
req.on("close", () => {
|
||||||
|
clearInterval(heartbeat);
|
||||||
|
unsubscribe();
|
||||||
|
});
|
||||||
|
} catch (err: unknown) {
|
||||||
|
if (err instanceof ApiError) {
|
||||||
|
throw err;
|
||||||
|
}
|
||||||
|
rethrowAsApiError(err, "Failed to attach chat stream");
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* POST /api/chat/sessions/:id/messages
|
* POST /api/chat/sessions/:id/messages
|
||||||
* Send a message and stream AI response via SSE.
|
* Send a message and stream AI response via SSE.
|
||||||
@@ -687,6 +770,7 @@ export function registerChatRoutes(ctx: ApiRoutesContext, deps: ChatRouteDeps):
|
|||||||
"POST /chat/sessions/:id/attachments",
|
"POST /chat/sessions/:id/attachments",
|
||||||
"GET /chat/sessions/:id/attachments/:filename",
|
"GET /chat/sessions/:id/attachments/:filename",
|
||||||
"DELETE /chat/sessions/:id/attachments/:filename",
|
"DELETE /chat/sessions/:id/attachments/:filename",
|
||||||
|
"GET /chat/sessions/:id/stream",
|
||||||
"POST /chat/sessions/:id/messages",
|
"POST /chat/sessions/:id/messages",
|
||||||
"POST /chat/sessions/:id/cancel",
|
"POST /chat/sessions/:id/cancel",
|
||||||
"DELETE /chat/sessions/:id/messages/:messageId",
|
"DELETE /chat/sessions/:id/messages/:messageId",
|
||||||
|
|||||||
Reference in New Issue
Block a user