Browse Source

feat: enhance chat session management with abort functionality and update stream handling

feat/ai-sdk-v6-upgrade
npmrun 6 hours ago
parent
commit
6140743352
  1. 4
      .codegraph/daemon.pid
  2. 80
      app/composables/useAgentChat.ts
  3. 11
      app/composables/useAgentChatStore.ts
  4. 2
      app/pages/index.vue
  5. 2
      packages/common/config/index.ts
  6. BIN
      packages/drizzle-pkg/db.sqlite
  7. 177
      server/api/agent/chat/index.post.ts
  8. 61
      server/api/agent/chat/stop.post.ts
  9. 2
      server/api/agent/chat/stream.get.ts
  10. 7
      server/service/agent-tool/index.ts

4
.codegraph/daemon.pid

@ -1,6 +1,6 @@
{
"pid": 18358,
"pid": 1727,
"version": "0.9.7",
"socketPath": "/home/dash/code/nuxt-app/.codegraph/daemon.sock",
"startedAt": 1786169176794
"startedAt": 1786201226096
}

80
app/composables/useAgentChat.ts

@ -98,6 +98,8 @@ export function useAgentChat(options: UseAgentChatOptions) {
}
async function tryResumeStream(sid: string) {
if (isStopped.value) return;
const lastMsg = messages.value[messages.value.length - 1];
if (!lastMsg) return;
@ -116,10 +118,21 @@ export function useAgentChat(options: UseAgentChatOptions) {
return;
}
if (!res.ok) return;
if (!res.ok) {
await loadMessages(sid);
return;
}
const isStreamResume = res.headers.get("X-Stream-Resume") === "true";
if (!isStreamResume) return;
if (!isStreamResume) {
try {
const body = await res.json();
if (body?.reason === "already-done" || body?.reason === "no-buffer") {
await loadMessages(sid);
}
} catch { /* not JSON */ }
return;
}
const modelIdHeader = res.headers.get("X-Stream-Model-Id");
const userMessageIdHeader = res.headers.get("X-Stream-User-Message-Id");
@ -562,6 +575,7 @@ export function useAgentChat(options: UseAgentChatOptions) {
);
isLoading.value = true;
isStopped.value = false;
abortController = new AbortController();
try {
@ -592,7 +606,13 @@ export function useAgentChat(options: UseAgentChatOptions) {
validateAssistantContent(messages.value.indexOf(assistantMsg), true);
onStreamComplete?.();
} catch (err: any) {
if (err.name !== "AbortError") {
if (err.name === "AbortError") {
isStopped.value = true;
const msg = messages.value[messages.value.indexOf(assistantMsg)];
if (msg) {
cleanupInProgressToolCalls(msg);
}
} else {
errorMessage.value = err.message || "请求失败";
}
} finally {
@ -627,6 +647,59 @@ export function useAgentChat(options: UseAgentChatOptions) {
}
}
async function stopAndSave() {
const sid = sessionId();
if (abortController) {
abortController.abort();
abortController = null;
}
isStopped.value = true;
if (sid) {
try {
await $fetch("/api/agent/chat/stop", {
method: "POST",
body: { sessionId: sid },
});
} catch {
}
await waitForAbortSave(sid, 3000);
await refreshMessagesNoResume(sid);
}
}
async function waitForAbortSave(sid: string, maxWait: number) {
const start = Date.now();
while (Date.now() - start < maxWait) {
try {
const res = await $fetch<{ code: number; data: { messages: AgentMessage[] } }>(
`/api/agent/sessions/${sid}/messages`,
{ method: "GET" },
);
const msgs = res.data?.messages ?? [];
const lastMsg = msgs[msgs.length - 1];
if (lastMsg && lastMsg.role === "assistant" && lastMsg.content) {
return;
}
} catch {
}
await new Promise((r) => setTimeout(r, 300));
}
}
async function refreshMessagesNoResume(sid: string) {
try {
const res = await $fetch<{ code: number; data: { messages: AgentMessage[] } }>(
`/api/agent/sessions/${sid}/messages`,
{ method: "GET" },
);
messages.value = (res.data?.messages ?? []).map((m) => ({
...m,
parts: m.parts ? (typeof m.parts === "string" ? JSON.parse(m.parts) : m.parts) : undefined,
}));
} catch {
}
}
function clear() {
messages.value = [];
errorMessage.value = "";
@ -642,6 +715,7 @@ export function useAgentChat(options: UseAgentChatOptions) {
rateLimit,
send,
stopGeneration,
stopAndSave,
clear,
loadMessages,
respondToApproval,

11
app/composables/useAgentChatStore.ts

@ -11,6 +11,7 @@ export interface AgentChatInstance {
rateLimit: ReturnType<typeof useAgentRateLimit>;
send: ReturnType<typeof useAgentChat>["send"];
stopGeneration: ReturnType<typeof useAgentChat>["stopGeneration"];
stopAndSave: ReturnType<typeof useAgentChat>["stopAndSave"];
clear: ReturnType<typeof useAgentChat>["clear"];
loadMessages: ReturnType<typeof useAgentChat>["loadMessages"];
respondToApproval: ReturnType<typeof useAgentChat>["respondToApproval"];
@ -119,7 +120,14 @@ export function useAgentChatStore(storeOptions: UseAgentChatStoreOptions) {
await inst.sendFeedback(messageId, feedback);
}
function stopGeneration() {
async function stopGeneration() {
const sid = currentSessionId.value;
if (!sid) return;
const inst = getInstance(sid);
await inst.stopAndSave();
}
function forceStop() {
const sid = currentSessionId.value;
if (!sid) return;
const inst = getInstance(sid);
@ -163,6 +171,7 @@ export function useAgentChatStore(storeOptions: UseAgentChatStoreOptions) {
respondToApproval,
sendFeedback,
stopGeneration,
forceStop,
stopAll,
clear,
setCurrentSessionId,

2
app/pages/index.vue

@ -191,7 +191,7 @@ async function handleApprove(toolCallId: string, approved: boolean) {
}
async function handleDeleteSession(id: string) {
chat.stopGeneration();
chat.forceStop();
chat.removeInstance(id);
await sessions.deleteSession(id);
}

2
packages/common/config/index.ts

@ -45,6 +45,8 @@ export const API_ALLOWLIST: RouteRule[] = [
{ path: "/api/agent/sessions/:id/config", methods: ["PUT"] },
{ path: "/api/agent/sessions/:id/messages", methods: ["GET"] },
{ path: "/api/agent/chat", methods: ["POST"] },
{ path: "/api/agent/chat/stream", methods: ["GET"] },
{ path: "/api/agent/chat/stop", methods: ["POST"] },
{ path: "/api/agent/chat/tool-approve", methods: ["POST"] },
{ path: "/api/agent/feedback", methods: ["POST"] },
{ path: "/api/agent-tools", methods: ["GET"] },

BIN
packages/drizzle-pkg/db.sqlite

Binary file not shown.

177
server/api/agent/chat/index.post.ts

@ -12,6 +12,7 @@ import { getModelWithProviderById, getModelWithProviderByIdAny } from "#server/s
import { generateSessionTitle } from "#server/service/agent/title";
import { type StoredPart } from "./types";
import { createStreamBuffer, appendChunk, markBufferDone, removeStreamBuffer } from "#server/service/agent/stream-buffer";
import { registerAbortController, unregisterAbortController } from "./stop.post";
import log4js from "logger";
const logger = log4js.getLogger("APP");
@ -414,6 +415,24 @@ export default defineEventHandler(async (event) => {
userMessageId: userMessage?.id ?? null,
});
const serverAbortController = new AbortController();
registerAbortController(sessionId, serverAbortController);
let abortedPartialContent = "";
let abortedPartialParts: Array<Record<string, unknown>> = [];
let streamedText = "";
let streamedReasoning = "";
let streamedToolCalls: Array<Record<string, unknown>> = [];
let chunkPartId = 0;
event.node.req.on("close", () => {
if (!serverAbortController.signal.aborted) {
logger.info("[%s] [AGENT-CHAT] client disconnected, aborting streamText", event.context.requestId ?? "-");
serverAbortController.abort();
}
});
let approvalPrefixChunk: Uint8Array | null = null;
if (isApprovalContinue && Object.keys(approvalToolResults).length > 0) {
const prefixChunks = Object.entries(approvalToolResults).map(
@ -433,6 +452,7 @@ export default defineEventHandler(async (event) => {
model: languageModel,
system: systemPrompt || undefined,
messages: modelMessages,
abortSignal: serverAbortController.signal,
maxOutputTokens: model.maxTokens && model.maxTokens >= 1 && model.maxTokens <= 393216 ? model.maxTokens : undefined,
...(Object.keys(effectiveTools).length > 0
? { tools: effectiveTools, stopWhen: stepCountIs(8) }
@ -450,7 +470,34 @@ export default defineEventHandler(async (event) => {
: String(errorData?.error ?? "未知错误");
logger.error("[%s] [AGENT-CHAT] streamText error: %s", event.context.requestId ?? "-", errMsg);
},
onChunk: ({ chunk }) => {
if (chunk.type === "text-delta") {
streamedText += chunk.text;
} else if (chunk.type === "reasoning-delta") {
streamedReasoning += chunk.text;
} else if (chunk.type === "tool-call") {
streamedToolCalls.push({
id: `p_${chunkPartId++}`,
type: "tool-call",
toolCallId: chunk.toolCallId,
toolName: chunk.toolName,
args: chunk.input,
state: "call",
});
} else if (chunk.type === "tool-result") {
const tcIdx = streamedToolCalls.findIndex((tc) => tc.toolCallId === chunk.toolCallId);
if (tcIdx >= 0) {
streamedToolCalls[tcIdx].result = chunk.output;
streamedToolCalls[tcIdx].state = "result";
}
}
},
onFinish: async ({ finishReason, usage, steps, text: assistantText }) => {
if (serverAbortController.signal.aborted) {
logger.info("[%s] [AGENT-CHAT] onFinish skipped due to abort", event.context.requestId ?? "-");
unregisterAbortController(sessionId);
return;
}
logger.info(
"[%s] [AGENT-CHAT] finished: reason=%s steps=%d inputTokens=%d outputTokens=%d",
event.context.requestId ?? "-",
@ -578,6 +625,96 @@ export default defineEventHandler(async (event) => {
if (shouldGenerateTitle && titleModelId && assistantText.trim().length > 0) {
generateSessionTitle(sessionId, content.trim(), assistantText, titleModelId, user?.id ?? null).catch((e) => { logger.error("[AGENT-CHAT] title gen failed: %s", e?.message ?? e); });
}
unregisterAbortController(sessionId);
},
onAbort: async () => {
try {
const partialText = streamedText;
logger.info("[%s] [AGENT-CHAT] aborted: streamedTextLen=%d streamedToolCalls=%d", event.context.requestId ?? "-", partialText.length, streamedToolCalls.length);
const parts: Array<Record<string, unknown>> = [];
if (streamedReasoning.trim()) {
parts.push({ id: `p_${parts.length}`, type: "reasoning", text: streamedReasoning });
}
if (partialText.trim()) {
parts.push({ id: `p_${parts.length}`, type: "text", text: partialText });
}
for (const tc of streamedToolCalls) {
parts.push({
id: `p_${parts.length}`,
type: "tool-call",
toolCallId: tc.toolCallId,
toolName: tc.toolName,
args: tc.args,
result: tc.result,
state: tc.result !== undefined ? "result" : "call",
});
}
abortedPartialContent = partialText;
abortedPartialParts = parts;
if (isApprovalContinue) {
const lastAssistant = historyMessages.filter((m) => m.role === "assistant").pop();
if (lastAssistant) {
let existingParts: Array<Record<string, unknown>> = [];
try {
existingParts = lastAssistant.parts ? JSON.parse(lastAssistant.parts) : [];
} catch {
existingParts = [];
}
for (const p of existingParts) {
if (p.type === "tool-call" && p.state === "approval-responded") {
if (p.approved === true && approvalToolResults[p.toolCallId] !== undefined) {
p.result = approvalToolResults[p.toolCallId];
p.state = "result";
} else if (p.approved === false) {
p.result = "工具执行被拒绝";
p.state = "result";
}
}
}
const newParts = parts.filter(
(np) => np.type !== "tool-call" || !existingParts.some((ep) => ep.toolCallId === np.toolCallId),
);
const idOffset = existingParts.length;
for (let i = 0; i < newParts.length; i++) {
newParts[i].id = `p_${idOffset + i}`;
}
const mergedParts = [...existingParts, ...newParts];
const mergedContent = (lastAssistant.content || "") + partialText;
await updateMessageContent(lastAssistant.id, mergedContent);
await updateMessageParts(lastAssistant.id, JSON.stringify(mergedParts));
}
} else {
const hasContent = partialText.trim().length > 0 || parts.length > 0;
if (hasContent) {
const assistantSortOrder = (await getMaxSortOrder(sessionId)) + 1;
await saveMessage({
sessionId,
role: "assistant",
content: partialText,
parts: parts.length > 0 ? JSON.stringify(parts) : null,
modelId: model.id,
sortOrder: assistantSortOrder,
});
}
}
await touchSession(sessionId);
if (!user) {
incrementRateLimit(sessionId, ip);
}
markBufferDone(sessionId);
unregisterAbortController(sessionId);
logger.info("[%s] [AGENT-CHAT] abort save complete: contentLen=%d partsCount=%d", event.context.requestId ?? "-", partialText.length, parts.length);
} catch (abortErr) {
logger.error("[%s] [AGENT-CHAT] onAbort error: %s", event.context.requestId ?? "-", abortErr instanceof Error ? abortErr.message : String(abortErr));
markBufferDone(sessionId);
unregisterAbortController(sessionId);
}
},
});
@ -594,22 +731,37 @@ export default defineEventHandler(async (event) => {
return new ReadableStream({
start(controller) {
const reader = originalBody.getReader();
let clientDisconnected = false;
function pump() {
reader.read().then(({ done, value }) => {
if (done) {
markBufferDone(sessionId);
controller.close();
if (!clientDisconnected) {
try { controller.close(); } catch { /* already closed */ }
}
return;
}
if (value) {
appendChunk(sessionId, value);
controller.enqueue(value);
if (!clientDisconnected) {
try {
controller.enqueue(value);
} catch {
clientDisconnected = true;
}
}
}
pump();
}).catch((err) => {
logger.error("[%s] [AGENT-CHAT] buffer stream read error: %s", event.context.requestId ?? "-", err instanceof Error ? err.message : String(err));
if (serverAbortController.signal.aborted) {
logger.info("[%s] [AGENT-CHAT] stream reader aborted by server abort", event.context.requestId ?? "-");
} else {
logger.error("[%s] [AGENT-CHAT] buffer stream read error: %s", event.context.requestId ?? "-", err instanceof Error ? err.message : String(err));
}
markBufferDone(sessionId);
controller.close();
if (!clientDisconnected) {
try { controller.close(); } catch { /* already closed */ }
}
});
}
pump();
@ -626,13 +778,20 @@ export default defineEventHandler(async (event) => {
async start(controller) {
controller.enqueue(prefix);
const reader = originalBody.getReader();
let clientDisconnected = false;
try {
while (true) {
const { done, value } = await reader.read();
if (done) break;
if (value) {
appendChunk(sessionId, value);
controller.enqueue(value);
if (!clientDisconnected) {
try {
controller.enqueue(value);
} catch {
clientDisconnected = true;
}
}
}
}
} catch (err) {
@ -641,10 +800,14 @@ export default defineEventHandler(async (event) => {
`data: ${JSON.stringify({ type: "error", errorText: "流式响应中断" })}\n\n`,
);
appendChunk(sessionId, errorChunk);
controller.enqueue(errorChunk);
if (!clientDisconnected) {
try { controller.enqueue(errorChunk); } catch { /* client gone */ }
}
} finally {
markBufferDone(sessionId);
controller.close();
if (!clientDisconnected) {
try { controller.close(); } catch { /* already closed */ }
}
}
},
});

61
server/api/agent/chat/stop.post.ts

@ -0,0 +1,61 @@
import { defineEventHandler, readBody } from "h3";
import { getCurrentUser } from "#server/utils/context";
import { getTempTokenFromCookie } from "#server/service/agent/temp-token";
import { getSessionByIdAndUser } from "#server/service/agent/session";
import { getStreamBuffer, markBufferDone } from "#server/service/agent/stream-buffer";
import { R } from "#server/utils/response";
import log4js from "logger";
const logger = log4js.getLogger("APP");
const activeControllers = new Map<string, AbortController>();
export function registerAbortController(sessionId: string, controller: AbortController) {
activeControllers.set(sessionId, controller);
}
export function unregisterAbortController(sessionId: string) {
activeControllers.delete(sessionId);
}
export function abortSessionStream(sessionId: string): boolean {
const controller = activeControllers.get(sessionId);
if (controller && !controller.signal.aborted) {
controller.abort();
logger.info("[AGENT-STOP] aborted streamText for sessionId=%s", sessionId);
return true;
}
return false;
}
export default defineEventHandler(async (event) => {
const user = await getCurrentUser(event);
const tempToken = getTempTokenFromCookie(event);
if (!user && !tempToken) {
throw createError({ statusCode: 401, statusMessage: "请先登录或创建临时会话" });
}
const body = await readBody(event);
const { sessionId } = body as { sessionId?: string };
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 aborted = abortSessionStream(sessionId);
if (!aborted) {
const buf = getStreamBuffer(sessionId);
if (buf && !buf.done) {
markBufferDone(sessionId);
}
}
return R.success({ aborted });
});

2
server/api/agent/chat/stream.get.ts

@ -34,7 +34,7 @@ export default defineEventHandler(async (event) => {
return R.success({ active: false, reason: "no-buffer" });
}
if (buf.done) {
if (buf.done && buf.chunks.length === 0) {
return R.success({ active: false, reason: "already-done", modelId: buf.modelId });
}

7
server/service/agent-tool/index.ts

@ -1,7 +1,7 @@
import { dbGlobal } from "drizzle-pkg/lib/db";
import { agentTools } from "drizzle-pkg/lib/schema/agent-tool";
import type { UserRole } from "drizzle-pkg/lib/schema/auth";
import { eq, asc, inArray } from "drizzle-orm";
import { eq, asc, inArray, and } from "drizzle-orm";
import { tool, jsonSchema } from "ai";
import { z } from "zod";
@ -139,7 +139,10 @@ export async function getAgentToolBySlug(slug: string): Promise<AgentToolRow | n
export async function getAgentToolsBySlugs(slugs: string[]): Promise<Map<string, AgentToolRow>> {
if (slugs.length === 0) return new Map();
const rows = await dbGlobal.select().from(agentTools).where(inArray(agentTools.slug, slugs));
const rows = await dbGlobal
.select()
.from(agentTools)
.where(and(inArray(agentTools.slug, slugs), eq(agentTools.enabled, 1)));
return new Map(rows.map((r) => [r.slug, r]));
}

Loading…
Cancel
Save