diff --git a/.codegraph/daemon.pid b/.codegraph/daemon.pid index f81707b..a0e5091 100644 --- a/.codegraph/daemon.pid +++ b/.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 } diff --git a/app/composables/useAgentChat.ts b/app/composables/useAgentChat.ts index 883af7f..5d6cedd 100644 --- a/app/composables/useAgentChat.ts +++ b/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, diff --git a/app/composables/useAgentChatStore.ts b/app/composables/useAgentChatStore.ts index 554e264..fe218ec 100644 --- a/app/composables/useAgentChatStore.ts +++ b/app/composables/useAgentChatStore.ts @@ -11,6 +11,7 @@ export interface AgentChatInstance { rateLimit: ReturnType; send: ReturnType["send"]; stopGeneration: ReturnType["stopGeneration"]; + stopAndSave: ReturnType["stopAndSave"]; clear: ReturnType["clear"]; loadMessages: ReturnType["loadMessages"]; respondToApproval: ReturnType["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, diff --git a/app/pages/index.vue b/app/pages/index.vue index 8ef7da1..d5f355f 100644 --- a/app/pages/index.vue +++ b/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); } diff --git a/packages/common/config/index.ts b/packages/common/config/index.ts index 79322cd..fe88b1f 100644 --- a/packages/common/config/index.ts +++ b/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"] }, diff --git a/packages/drizzle-pkg/db.sqlite b/packages/drizzle-pkg/db.sqlite index af9fd14..2494383 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 eef2389..b52cb75 100644 --- a/server/api/agent/chat/index.post.ts +++ b/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> = []; + let streamedText = ""; + let streamedReasoning = ""; + let streamedToolCalls: Array> = []; + 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> = []; + 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> = []; + 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 */ } + } } }, }); diff --git a/server/api/agent/chat/stop.post.ts b/server/api/agent/chat/stop.post.ts new file mode 100644 index 0000000..f297f4a --- /dev/null +++ b/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(); + +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 }); +}); diff --git a/server/api/agent/chat/stream.get.ts b/server/api/agent/chat/stream.get.ts index bf050fb..f8ce8f9 100644 --- a/server/api/agent/chat/stream.get.ts +++ b/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 }); } diff --git a/server/service/agent-tool/index.ts b/server/service/agent-tool/index.ts index e7f29a4..51a4335 100644 --- a/server/service/agent-tool/index.ts +++ b/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> { 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])); }