diff --git a/app/composables/useAgentChat.ts b/app/composables/useAgentChat.ts index c099d09..5824418 100644 --- a/app/composables/useAgentChat.ts +++ b/app/composables/useAgentChat.ts @@ -90,11 +90,65 @@ export function useAgentChat(options: UseAgentChatOptions) { ...m, parts: m.parts ? (typeof m.parts === "string" ? JSON.parse(m.parts) : m.parts) : undefined, })); + + await tryResumeStream(sid); } catch { messages.value = []; } } + async function tryResumeStream(sid: string) { + const lastMsg = messages.value[messages.value.length - 1]; + if (!lastMsg || lastMsg.role !== "user") return; + + let res: Response; + try { + res = await fetch(`/api/agent/chat/stream?sessionId=${encodeURIComponent(sid)}`, { + method: "GET", + }); + } catch { + return; + } + + if (!res.ok) return; + + const isStreamResume = res.headers.get("X-Stream-Resume") === "true"; + if (!isStreamResume) return; + + const modelIdHeader = res.headers.get("X-Stream-Model-Id"); + const userMessageIdHeader = res.headers.get("X-Stream-User-Message-Id"); + + const assistantMsg: AgentMessage = { + id: generateId(), + role: "assistant", + content: "", + parts: [], + modelId: modelIdHeader ? Number(modelIdHeader) : null, + }; + messages.value.push(assistantMsg); + const assistantIdx = messages.value.length - 1; + + if (userMessageIdHeader && lastMsg) { + lastMsg.id = userMessageIdHeader; + } + + isLoading.value = true; + abortController = new AbortController(); + + try { + await processStream(res, assistantIdx); + validateAssistantContent(assistantIdx); + onStreamComplete?.(); + } catch (err: any) { + if (err.name !== "AbortError") { + errorMessage.value = err.message || "流式恢复失败"; + } + } finally { + isLoading.value = false; + abortController = null; + } + } + async function cleanupStaleApprovals() { const staleMessages: { msg: AgentMessage; parts: MessagePart[] }[] = []; for (const msg of messages.value) { diff --git a/packages/drizzle-pkg/db.sqlite b/packages/drizzle-pkg/db.sqlite index e0f01ed..31f118b 100644 Binary files a/packages/drizzle-pkg/db.sqlite and b/packages/drizzle-pkg/db.sqlite differ diff --git a/server/api/agent/chat/index.post.ts b/server/api/agent/chat/index.post.ts index ba5bd8c..8032432 100644 --- a/server/api/agent/chat/index.post.ts +++ b/server/api/agent/chat/index.post.ts @@ -11,6 +11,7 @@ import { checkRateLimit, incrementRateLimit } from "#server/service/agent/rate-l import { getModelWithProviderById, getModelWithProviderByIdAny } from "#server/service/llm"; import { generateSessionTitle } from "#server/service/agent/title"; import { type StoredPart } from "./types"; +import { createStreamBuffer, appendChunk, markBufferDone, removeStreamBuffer } from "#server/service/agent/stream-buffer"; import log4js from "logger"; const logger = log4js.getLogger("APP"); @@ -408,6 +409,11 @@ export default defineEventHandler(async (event) => { isApprovalContinue ? "yes" : "no", ); + const streamBuffer = createStreamBuffer(sessionId, { + modelId: model.id, + userMessageId: userMessage?.id ?? null, + }); + const result = streamText({ model: languageModel, system: systemPrompt || undefined, @@ -569,6 +575,33 @@ export default defineEventHandler(async (event) => { headers: responseHeaders, }); + function wrapStreamWithBuffer(originalBody: ReadableStream): ReadableStream { + return new ReadableStream({ + start(controller) { + const reader = originalBody.getReader(); + function pump() { + reader.read().then(({ done, value }) => { + if (done) { + markBufferDone(sessionId); + controller.close(); + return; + } + if (value) { + appendChunk(sessionId, value); + controller.enqueue(value); + } + pump(); + }).catch((err) => { + logger.error("[%s] [AGENT-CHAT] buffer stream read error: %s", event.context.requestId ?? "-", err instanceof Error ? err.message : String(err)); + markBufferDone(sessionId); + controller.close(); + }); + } + pump(); + }, + }); + } + if (continueAfterApproval && Object.keys(approvalToolResults).length > 0) { const originalBody = response.body; if (originalBody) { @@ -581,6 +614,7 @@ export default defineEventHandler(async (event) => { })}\n\n`, ); const prefix = new TextEncoder().encode(prefixChunks.join("")); + appendChunk(sessionId, prefix); const transformed = new ReadableStream({ async start(controller) { @@ -590,15 +624,20 @@ export default defineEventHandler(async (event) => { while (true) { const { done, value } = await reader.read(); if (done) break; - controller.enqueue(value); + if (value) { + appendChunk(sessionId, value); + controller.enqueue(value); + } } } catch (err) { logger.error("[%s] [AGENT-CHAT] stream read error: %s", event.context.requestId ?? "-", err instanceof Error ? err.message : String(err)); const errorChunk = new TextEncoder().encode( `data: ${JSON.stringify({ type: "error", errorText: "流式响应中断" })}\n\n`, ); + appendChunk(sessionId, errorChunk); controller.enqueue(errorChunk); } finally { + markBufferDone(sessionId); controller.close(); } }, @@ -612,5 +651,14 @@ export default defineEventHandler(async (event) => { } } + if (response.body) { + const bufferedBody = wrapStreamWithBuffer(response.body); + return new Response(bufferedBody, { + status: response.status, + statusText: response.statusText, + headers: response.headers, + }); + } + return response; }); diff --git a/server/api/agent/chat/stream.get.ts b/server/api/agent/chat/stream.get.ts new file mode 100644 index 0000000..bf050fb --- /dev/null +++ b/server/api/agent/chat/stream.get.ts @@ -0,0 +1,89 @@ +import { defineEventHandler, getQuery } from "h3"; +import { R } from "#server/utils/response"; +import { getCurrentUser } from "#server/utils/context"; +import { getTempTokenFromCookie } from "#server/service/agent/temp-token"; +import { getSessionByIdAndUser } from "#server/service/agent/session"; +import { getStreamBuffer, subscribeToBuffer } from "#server/service/agent/stream-buffer"; +import log4js from "logger"; + +const logger = log4js.getLogger("APP"); + +export default defineEventHandler(async (event) => { + const user = await getCurrentUser(event); + const tempToken = getTempTokenFromCookie(event); + + if (!user && !tempToken) { + throw createError({ statusCode: 401, statusMessage: "请先登录或创建临时会话" }); + } + + const query = getQuery(event); + const sessionId = query.sessionId as string | undefined; + + if (!sessionId) { + throw createError({ statusCode: 400, statusMessage: "缺少 sessionId 参数" }); + } + + const session = await getSessionByIdAndUser(sessionId, user?.id ?? null, tempToken); + if (!session) { + throw createError({ statusCode: 404, statusMessage: "会话不存在" }); + } + + const buf = getStreamBuffer(sessionId); + + if (!buf) { + return R.success({ active: false, reason: "no-buffer" }); + } + + if (buf.done) { + return R.success({ active: false, reason: "already-done", modelId: buf.modelId }); + } + + const stream = new ReadableStream({ + start(controller) { + for (const chunk of buf.chunks) { + controller.enqueue(chunk); + } + + if (buf.done) { + controller.close(); + return; + } + + const unsubscribe = subscribeToBuffer( + sessionId, + (chunk) => { + try { + controller.enqueue(chunk); + } catch (e) { + logger.error("[STREAM-RESUME] enqueue error: %s", e instanceof Error ? e.message : String(e)); + } + }, + () => { + try { + controller.close(); + } catch { + } + }, + ); + + event.node.req.on("close", () => { + unsubscribe(); + try { + controller.close(); + } catch { + } + }); + }, + }); + + return new Response(stream, { + headers: { + "Content-Type": "text/event-stream; charset=utf-8", + "Cache-Control": "no-cache", + "Connection": "keep-alive", + "X-Stream-Resume": "true", + "X-Stream-Model-Id": String(buf.modelId ?? ""), + "X-Stream-User-Message-Id": buf.userMessageId ?? "", + }, + }); +}); diff --git a/server/service/agent/stream-buffer.ts b/server/service/agent/stream-buffer.ts new file mode 100644 index 0000000..0b20479 --- /dev/null +++ b/server/service/agent/stream-buffer.ts @@ -0,0 +1,129 @@ +import log4js from "logger"; + +const logger = log4js.getLogger("APP"); + +const BUFFER_TTL_MS = 5 * 60 * 1000; +const CLEANUP_INTERVAL_MS = 60 * 1000; + +export interface StreamBuffer { + chunks: Uint8Array[]; + done: boolean; + createdAt: number; + doneAt: number | null; + modelId: number | null; + userMessageId: string | null; + subscribers: Array<(chunk: Uint8Array) => void>; + doneSubscribers: Array<() => void>; +} + +const buffers = new Map(); + +let cleanupTimer: ReturnType | null = null; + +function ensureCleanupTimer() { + if (cleanupTimer) return; + cleanupTimer = setInterval(() => { + const now = Date.now(); + for (const [sid, buf] of buffers) { + const expireAt = buf.doneAt ?? buf.createdAt + BUFFER_TTL_MS; + if (now > expireAt) { + buffers.delete(sid); + logger.info("[STREAM-BUFFER] expired and removed sessionId=%s", sid); + } + } + if (buffers.size === 0 && cleanupTimer) { + clearInterval(cleanupTimer); + cleanupTimer = null; + } + }, CLEANUP_INTERVAL_MS); + cleanupTimer.unref?.(); +} + +export function createStreamBuffer( + sessionId: string, + meta: { modelId?: number | null; userMessageId?: string | null }, +): StreamBuffer { + const buf: StreamBuffer = { + chunks: [], + done: false, + createdAt: Date.now(), + doneAt: null, + modelId: meta.modelId ?? null, + userMessageId: meta.userMessageId ?? null, + subscribers: [], + doneSubscribers: [], + }; + buffers.set(sessionId, buf); + ensureCleanupTimer(); + logger.info("[STREAM-BUFFER] created sessionId=%s", sessionId); + return buf; +} + +export function appendChunk(sessionId: string, chunk: Uint8Array): void { + const buf = buffers.get(sessionId); + if (!buf || buf.done) return; + buf.chunks.push(chunk); + for (const sub of buf.subscribers) { + try { + sub(chunk); + } catch (e) { + logger.error("[STREAM-BUFFER] subscriber error: %s", e instanceof Error ? e.message : String(e)); + } + } +} + +export function markBufferDone(sessionId: string): void { + const buf = buffers.get(sessionId); + if (!buf) return; + buf.done = true; + buf.doneAt = Date.now(); + const doneSubs = buf.doneSubscribers; + buf.doneSubscribers = []; + buf.subscribers = []; + for (const sub of doneSubs) { + try { + sub(); + } catch (e) { + logger.error("[STREAM-BUFFER] done subscriber error: %s", e instanceof Error ? e.message : String(e)); + } + } + logger.info("[STREAM-BUFFER] done sessionId=%s chunks=%d", sessionId, buf.chunks.length); +} + +export function getStreamBuffer(sessionId: string): StreamBuffer | undefined { + return buffers.get(sessionId); +} + +export function removeStreamBuffer(sessionId: string): void { + buffers.delete(sessionId); +} + +export function hasActiveStream(sessionId: string): boolean { + const buf = buffers.get(sessionId); + return !!buf && !buf.done; +} + +export function subscribeToBuffer( + sessionId: string, + onChunk: (chunk: Uint8Array) => void, + onDone: () => void, +): () => void { + const buf = buffers.get(sessionId); + if (!buf) { + onDone(); + return () => {}; + } + + if (buf.done) { + onDone(); + return () => {}; + } + + buf.subscribers.push(onChunk); + buf.doneSubscribers.push(onDone); + + return () => { + buf.subscribers = buf.subscribers.filter((s) => s !== onChunk); + buf.doneSubscribers = buf.doneSubscribers.filter((s) => s !== onDone); + }; +}