Browse Source

feat: implement stream buffering for chat sessions and add stream resume functionality

feat/ai-sdk-v6-upgrade
npmrun 10 hours ago
parent
commit
3e5c22335d
  1. 54
      app/composables/useAgentChat.ts
  2. BIN
      packages/drizzle-pkg/db.sqlite
  3. 48
      server/api/agent/chat/index.post.ts
  4. 89
      server/api/agent/chat/stream.get.ts
  5. 129
      server/service/agent/stream-buffer.ts

54
app/composables/useAgentChat.ts

@ -90,11 +90,65 @@ export function useAgentChat(options: UseAgentChatOptions) {
...m, ...m,
parts: m.parts ? (typeof m.parts === "string" ? JSON.parse(m.parts) : m.parts) : undefined, parts: m.parts ? (typeof m.parts === "string" ? JSON.parse(m.parts) : m.parts) : undefined,
})); }));
await tryResumeStream(sid);
} catch { } catch {
messages.value = []; 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() { async function cleanupStaleApprovals() {
const staleMessages: { msg: AgentMessage; parts: MessagePart[] }[] = []; const staleMessages: { msg: AgentMessage; parts: MessagePart[] }[] = [];
for (const msg of messages.value) { for (const msg of messages.value) {

BIN
packages/drizzle-pkg/db.sqlite

Binary file not shown.

48
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 { getModelWithProviderById, getModelWithProviderByIdAny } from "#server/service/llm";
import { generateSessionTitle } from "#server/service/agent/title"; import { generateSessionTitle } from "#server/service/agent/title";
import { type StoredPart } from "./types"; import { type StoredPart } from "./types";
import { createStreamBuffer, appendChunk, markBufferDone, removeStreamBuffer } from "#server/service/agent/stream-buffer";
import log4js from "logger"; import log4js from "logger";
const logger = log4js.getLogger("APP"); const logger = log4js.getLogger("APP");
@ -408,6 +409,11 @@ export default defineEventHandler(async (event) => {
isApprovalContinue ? "yes" : "no", isApprovalContinue ? "yes" : "no",
); );
const streamBuffer = createStreamBuffer(sessionId, {
modelId: model.id,
userMessageId: userMessage?.id ?? null,
});
const result = streamText({ const result = streamText({
model: languageModel, model: languageModel,
system: systemPrompt || undefined, system: systemPrompt || undefined,
@ -569,6 +575,33 @@ export default defineEventHandler(async (event) => {
headers: responseHeaders, headers: responseHeaders,
}); });
function wrapStreamWithBuffer(originalBody: ReadableStream<Uint8Array>): ReadableStream<Uint8Array> {
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) { if (continueAfterApproval && Object.keys(approvalToolResults).length > 0) {
const originalBody = response.body; const originalBody = response.body;
if (originalBody) { if (originalBody) {
@ -581,6 +614,7 @@ export default defineEventHandler(async (event) => {
})}\n\n`, })}\n\n`,
); );
const prefix = new TextEncoder().encode(prefixChunks.join("")); const prefix = new TextEncoder().encode(prefixChunks.join(""));
appendChunk(sessionId, prefix);
const transformed = new ReadableStream({ const transformed = new ReadableStream({
async start(controller) { async start(controller) {
@ -590,15 +624,20 @@ export default defineEventHandler(async (event) => {
while (true) { while (true) {
const { done, value } = await reader.read(); const { done, value } = await reader.read();
if (done) break; if (done) break;
if (value) {
appendChunk(sessionId, value);
controller.enqueue(value); controller.enqueue(value);
} }
}
} catch (err) { } catch (err) {
logger.error("[%s] [AGENT-CHAT] stream read error: %s", event.context.requestId ?? "-", err instanceof Error ? err.message : String(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( const errorChunk = new TextEncoder().encode(
`data: ${JSON.stringify({ type: "error", errorText: "流式响应中断" })}\n\n`, `data: ${JSON.stringify({ type: "error", errorText: "流式响应中断" })}\n\n`,
); );
appendChunk(sessionId, errorChunk);
controller.enqueue(errorChunk); controller.enqueue(errorChunk);
} finally { } finally {
markBufferDone(sessionId);
controller.close(); 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; return response;
}); });

89
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<Uint8Array>({
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 ?? "",
},
});
});

129
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<string, StreamBuffer>();
let cleanupTimer: ReturnType<typeof setInterval> | 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);
};
}
Loading…
Cancel
Save