From 9f60450ffa9853fa9310587d880a8da1df71140e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=89=AB=E5=9C=B0=E5=83=A7?= <357099073@qq.com> Date: Thu, 10 Sep 2026 21:51:00 +0800 Subject: [PATCH] =?UTF-8?q?=E6=9B=B4=E6=96=B0ai=E5=B9=BF=E5=9C=BAbug?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/components.d.ts | 1 + .../views/apps/ai/components/chatWindow.vue | 261 +++++++++++++++--- .../apps/oa/document/documentgraph/index.vue | 101 ++++--- go/controllers/backend_ai_chat.go | 239 ++++++++++++---- .../alter_backend_ai_chat_message_stats.sql | 6 + go/models/backend_ai_chat_message.go | 6 + 6 files changed, 482 insertions(+), 132 deletions(-) create mode 100644 go/docs/sql/alter_backend_ai_chat_message_stats.sql diff --git a/backend/components.d.ts b/backend/components.d.ts index 2483e63..2f3d0d6 100644 --- a/backend/components.d.ts +++ b/backend/components.d.ts @@ -54,6 +54,7 @@ declare module 'vue' { ElMenuItem: typeof import('element-plus/es')['ElMenuItem'] ElOption: typeof import('element-plus/es')['ElOption'] ElPagination: typeof import('element-plus/es')['ElPagination'] + ElPopover: typeof import('element-plus/es')['ElPopover'] ElProgress: typeof import('element-plus/es')['ElProgress'] ElRadio: typeof import('element-plus/es')['ElRadio'] ElRadioButton: typeof import('element-plus/es')['ElRadioButton'] diff --git a/backend/src/views/apps/ai/components/chatWindow.vue b/backend/src/views/apps/ai/components/chatWindow.vue index fd42b9d..c00f12b 100644 --- a/backend/src/views/apps/ai/components/chatWindow.vue +++ b/backend/src/views/apps/ai/components/chatWindow.vue @@ -42,13 +42,37 @@ :value="m" /> - - - - MCP - {{ activeMcpTools.length }} - - + + +
+
+ MCP 工具 + +
+
{{ mcpTip }}
+ +
+
@@ -92,13 +116,16 @@ class="message-bubble" :class="{ typing: msg.role === 'assistant' && !msg.content && sending, - 'has-tools': msg.role === 'assistant' && msg.toolCalls && msg.toolCalls.length + 'has-tools': msg.role === 'assistant' && msg.toolCalls && msg.toolCalls.length, + generating: msg.role === 'assistant' && sending, + error: msg.role === 'assistant' && msg.error }" > +
+ +
+
+ {{ formatDuration(msg.durationMs) }} + · + {{ msg.tokens }} tokens + · + {{ msg.model }} +
@@ -196,6 +233,10 @@ const sending = ref(false) const inputText = ref('') const messages = ref([]) const messagesRef = ref() +// 流式回复进行中时 sessionId 变化的暂存值(等待本轮结束后再加载历史) +const pendingSessionId = ref(null) +// 标记「本次 sessionId 变化是本轮发送自己创建的会话」,避免拉取历史覆盖正在流式追加的消息 +const selfSessionSwitch = ref(false) // 消息编辑 const editingId = ref(-1) @@ -213,14 +254,56 @@ const defaultPreset = ref(null) // MCP 接入状态 const activeMcpTools = ref([]) +// 本轮对话是否启用 MCP 工具(本地记忆上次选择) +const savedMcp = loadMcpSetting() +const mcpEnabled = ref(savedMcp.enabled) +// 参与本轮对话的 MCP 服务器(默认全部已启用的服务) +const selectedMcpServers = ref(savedMcp.serverIds) +const mcpServerList = ref([]) +// 开关打开且至少勾选了一个服务才算真正启用 +const mcpOn = computed(() => mcpEnabled.value && selectedMcpServers.value.length > 0) const mcpTip = computed(() => { - if (activeMcpTools.value.length === 0) { - return '当前未启用 MCP 工具,可在侧边栏 CPU 按钮中配置' + if (!mcpEnabled.value) { + return 'MCP 工具已关闭,本轮对话不会调用任何 MCP 工具' } - const servers = [...new Set(activeMcpTools.value.map((t) => t.server_name).filter(Boolean))] - return `已接入 ${activeMcpTools.value.length} 个 MCP 工具(${servers.join('、')}),会话中 AI 会自动调用` + if (selectedMcpServers.value.length === 0) { + return '未勾选任何 MCP 服务,本轮对话不会调用工具' + } + const names = mcpServerList.value + .filter((s) => selectedMcpServers.value.includes(s.id)) + .map((s) => s.name) + return `本轮将调用:${names.join('、')}` }) +function loadMcpSetting() { + const setting = { enabled: true, serverIds: [] } + try { + const e = localStorage.getItem('ai_mcp_enabled') + if (e !== null) setting.enabled = e === '1' + const s = localStorage.getItem('ai_mcp_servers') + if (s) setting.serverIds = JSON.parse(s) || [] + } catch (e) { + // 忽略 + } + return setting +} + +function persistMcpSetting() { + try { + localStorage.setItem('ai_mcp_enabled', mcpEnabled.value ? '1' : '0') + localStorage.setItem('ai_mcp_servers', JSON.stringify(selectedMcpServers.value)) + } catch (e) { + // 忽略 + } +} + +// 将毫秒格式化为可读耗时 +function formatDuration(ms) { + if (!ms || ms <= 0) return '0ms' + if (ms < 1000) return `${ms}ms` + return `${(ms / 1000).toFixed(1)}s` +} + const currentModels = computed(() => { const p = providerList.value.find(item => item.id === selectedProviderId.value) return p && p.models ? p.models : [] @@ -237,11 +320,10 @@ async function fetchMcpStatus() { const res = await getMcpServerList() if (res.code === 200 && res.data && res.data.list) { const enabled = res.data.list.filter((s) => s.enabled === 1) - const tools = [] - enabled.forEach((s) => { - tools.push({ key: '', server_id: s.id, server_name: s.name, tool_name: '' }) - }) - activeMcpTools.value = tools + mcpServerList.value = enabled.map((s) => ({ id: s.id, name: s.name })) + // 过滤掉已删除/已停用的服务;没有历史选择时默认全选 + const valid = selectedMcpServers.value.filter((id) => mcpServerList.value.some((s) => s.id === id)) + selectedMcpServers.value = valid.length > 0 ? valid : mcpServerList.value.map((s) => s.id) } } catch (e) { // 静默失败 @@ -295,6 +377,15 @@ const quickTips = [ watch( () => props.sessionId, (val) => { + if (sending.value) { + // 流式回复期间不能重新拉取历史:会把本地正在追加的 assistant 消息整体替换掉,导致回复不显示 + if (selfSessionSwitch.value) { + selfSessionSwitch.value = false + return + } + pendingSessionId.value = val + return + } if (val) { fetchMessages(val) } else { @@ -452,9 +543,12 @@ async function doSend(content) { // 添加用户消息 messages.value.push({ role: 'user', content }) + // 注意:必须取数组内的响应式代理对象来修改,直接持有 push 进去的原始对象不会触发视图更新 + const userMsg = messages.value[messages.value.length - 1] // 添加空的AI消息,用于流式更新 - const aiMsg = { id: null, role: 'assistant', content: '', toolCalls: [] } - messages.value.push(aiMsg) + messages.value.push({ id: null, role: 'assistant', content: '', toolCalls: [], tokens: 0, durationMs: 0, model: '', error: false }) + const aiIndex = messages.value.length - 1 + const aiMsg = messages.value[aiIndex] sending.value = true scrollToBottom() @@ -471,7 +565,9 @@ async function doSend(content) { session_id: props.sessionId || 0, content, provider_id: selectedProviderId.value, - model: selectedModel.value + model: selectedModel.value, + mcp_enabled: mcpOn.value, + mcp_server_ids: mcpOn.value ? selectedMcpServers.value : [] }) }) @@ -485,14 +581,19 @@ async function doSend(content) { let sessionCreated = false let currentEvent = '' let currentData = '' + let streamError = null const handleSSEEvent = (eventType, data) => { if (!data) return try { const parsed = JSON.parse(data) if (eventType === 'session') { + if (parsed.user_message_id) { + userMsg.id = parsed.user_message_id + } if (!props.sessionId && parsed.session_id && !sessionCreated) { sessionCreated = true + selfSessionSwitch.value = true emit('session-created', { id: parsed.session_id, title: content.length > 20 ? content.slice(0, 20) + '...' : content @@ -532,11 +633,13 @@ async function doSend(content) { scrollToBottom() } else if (eventType === 'done') { aiMsg.id = parsed.message_id + if (parsed.tokens) aiMsg.tokens = parsed.tokens + if (parsed.duration_ms) aiMsg.durationMs = parsed.duration_ms + if (parsed.model) aiMsg.model = parsed.model } else if (eventType === 'error') { - throw new Error(parsed.error || 'AI调用失败') + streamError = new Error(parsed.error || 'AI调用失败') } } catch (e) { - if (e.message === 'AI调用失败') throw e // JSON解析错误忽略 } } @@ -569,15 +672,28 @@ async function doSend(content) { if (currentData) { handleSSEEvent(currentEvent, currentData) } + if (streamError) { + throw streamError + } } catch (e) { - ElMessage.error(e.message || '发送失败') - // 移除空的AI消息 - if (aiMsg.content === '') { - const idx = messages.value.indexOf(aiMsg) - if (idx >= 0) messages.value.splice(idx, 1) + // 将错误以对话形式展示,不弹出顶部提示 + const errText = e.message || '发送失败,请稍后重试' + if (messages.value[aiIndex] === aiMsg) { + if (!aiMsg.content) { + aiMsg.content = errText + } else { + aiMsg.content += '\n\n⚠️ ' + errText + } + aiMsg.error = true } } finally { sending.value = false + // 流式期间被挂起的会话切换,此时再加载历史 + if (pendingSessionId.value) { + const id = pendingSessionId.value + pendingSessionId.value = null + await fetchMessages(id) + } scrollToBottom() } } @@ -733,6 +849,46 @@ defineExpose({ } } + // 助手回复下方的响应信息(耗时 / tokens / 模型) + .message-meta { + display: flex; + align-items: center; + gap: 6px; + margin-top: 6px; + font-size: 12px; + color: var(--el-text-color-secondary); + + .meta-dot { + color: var(--el-border-color); + } + + .meta-item { + white-space: nowrap; + } + } + + // 模型运行中的加载指示(流式输出期间显示) + .generating-indicator { + display: inline-flex; + align-items: center; + margin-bottom: 4px; + color: var(--el-color-primary); + font-size: 16px; + } + + .thinking-text { + color: var(--el-text-color-secondary); + font-size: 13px; + } + + // 助手消息报错(对话内错误展示) + .message-bubble.error { + background: var(--el-color-danger-light-9); + border: 1px solid var(--el-color-danger-light-5); + color: var(--el-color-danger); + white-space: pre-wrap; + } + .edit-area { display: inline-block; text-align: left; @@ -850,8 +1006,14 @@ defineExpose({ border: 1px solid var(--el-border-color); color: var(--el-text-color-secondary); font-size: 12px; - cursor: default; + cursor: pointer; height: 32px; + user-select: none; + + &:hover { + border-color: var(--el-color-primary); + color: var(--el-color-primary); + } &.has-mcp { background: var(--el-color-success-light-9); @@ -871,6 +1033,43 @@ defineExpose({ } } +// MCP 开关面板 +.mcp-panel { + .mcp-panel-head { + display: flex; + align-items: center; + justify-content: space-between; + margin-bottom: 6px; + } + + .mcp-panel-title { + font-size: 14px; + font-weight: 600; + } + + .mcp-panel-tip { + font-size: 12px; + color: var(--el-text-color-secondary); + line-height: 1.6; + margin-bottom: 10px; + word-break: break-word; + } + + .mcp-panel-empty { + font-size: 12px; + color: var(--el-text-color-placeholder); + padding: 6px 0; + } + + .mcp-check-list { + display: flex; + flex-direction: column; + gap: 6px; + max-height: 220px; + overflow-y: auto; + } +} + // 工具调用卡片 .tool-call-list { display: flex; diff --git a/backend/src/views/apps/oa/document/documentgraph/index.vue b/backend/src/views/apps/oa/document/documentgraph/index.vue index b8125c0..2281cda 100644 --- a/backend/src/views/apps/oa/document/documentgraph/index.vue +++ b/backend/src/views/apps/oa/document/documentgraph/index.vue @@ -4,7 +4,7 @@

文档图谱

- 默认展示顶级往下 3 级分类脉络,点击带 ▸ 的分类节点可继续延展 3 级;滚轮缩放,按住空格 + 左键拖拽平移 + 默认展示顶级往下 3 级分类脉络,点击带 ▸ 的分类节点可继续延展 3 级;滚轮缩放,按住右键拖拽平移画布

@@ -66,7 +66,12 @@
-
+
分类层级 滚轮缩放 · 空格+左键拖拽平移 · 点击 ▸ 分类展开下级(更深层级循环使用这 3 色)滚轮缩放 · 右键拖拽平移 · 左键拖拽节点 · 点击 ▸ 分类展开下级(更深层级循环使用这 3 + 色)
@@ -210,8 +216,8 @@ const detailId = ref(0); // 分层展开:默认展示顶级往下 3 级,点击分支节点后在该分支继续延展 3 级 const INITIAL_DEPTH = 3; const expandedIds = ref(new Set()); -// 按住空格 + 鼠标左键拖拽平移 -const spaceHeld = ref(false); +// 按住右键拖拽平移画布(ECharts 原生 roam 只响应左键,右键需自行接管) +const panning = ref(false); const filters = reactive({ keyword: "", @@ -455,10 +461,10 @@ function buildOption() { { type: "graph", layout: filters.layout, - // roam: 滚轮缩放 + 拖拽平移;按住空格时禁用节点拖拽,让左键拖拽整体平移 + // roam: 滚轮缩放 + 左键在空白处拖拽平移;节点左键可拖拽,右键拖拽整体平移 roam: true, scaleLimit: { min: 0.15, max: 6 }, - draggable: !spaceHeld.value, + draggable: true, data, links: edgeList, categories: [ @@ -540,39 +546,50 @@ function openNode(node) { } } -// ---------------- 空格 + 鼠标左键拖拽平移 ---------------- -function isTypingTarget(el) { - return ( - !!el && - (el.tagName === "INPUT" || - el.tagName === "TEXTAREA" || - el.isContentEditable === true) - ); -} +// ---------------- 按住右键拖拽平移 ---------------- +// ECharts 的 roam 会忽略右键(isMiddleOrRightButtonOnMouseUpDown), +// 这里自行接管右键拖拽,通过 graphRoam 动作平移,不改动节点数据与布局。 +let panLastX = 0; +let panLastY = 0; -function applyPanMode(on) { - if (chartRef.value) chartRef.value.style.cursor = on ? "grab" : ""; - try { - // 空格按下时禁用节点拖拽,使左键拖拽变为整体平移 - chart?.setOption({ series: [{ draggable: !on }] }); - } catch { - /* ignore */ - } -} - -function onSpaceDown(e) { - if (e.code !== "Space") return; - if (isTypingTarget(e.target)) return; +function onChartMouseDown(e) { + if (e.button !== 2) return; // 仅响应右键 e.preventDefault(); - if (spaceHeld.value) return; - spaceHeld.value = true; - applyPanMode(true); + panning.value = true; + panLastX = e.clientX; + panLastY = e.clientY; + chartRef.value?.classList.add("is-panning"); + window.addEventListener("mousemove", onPanMouseMove); + window.addEventListener("mouseup", onPanMouseUp); } -function onSpaceUp(e) { - if (e.code !== "Space") return; - spaceHeld.value = false; - applyPanMode(false); +function onPanMouseMove(e) { + if (!panning.value) return; + // 右键在画布外松开时 mouseup 不会触发,用 buttons 兜底结束拖拽 + if ((e.buttons & 2) === 0) { + stopPan(); + return; + } + const dx = e.clientX - panLastX; + const dy = e.clientY - panLastY; + panLastX = e.clientX; + panLastY = e.clientY; + if (!dx && !dy) return; + e.preventDefault(); + chart?.dispatchAction({ type: "graphRoam", seriesIndex: 0, dx, dy }); +} + +function onPanMouseUp(e) { + if (e.button !== 2) return; + stopPan(); +} + +function stopPan() { + if (!panning.value) return; + panning.value = false; + chartRef.value?.classList.remove("is-panning"); + window.removeEventListener("mousemove", onPanMouseMove); + window.removeEventListener("mouseup", onPanMouseUp); } async function loadCategories() { @@ -677,14 +694,10 @@ onMounted(async () => { attributes: true, attributeFilter: ["class"] }); - // 空格 + 左键拖拽平移 - window.addEventListener("keydown", onSpaceDown); - window.addEventListener("keyup", onSpaceUp); }); onBeforeUnmount(() => { - window.removeEventListener("keydown", onSpaceDown); - window.removeEventListener("keyup", onSpaceUp); + stopPan(); resizeObserver?.disconnect(); themeObserver?.disconnect(); chart?.dispose(); @@ -752,6 +765,12 @@ onBeforeUnmount(() => { .chart { width: 100%; height: 100%; + + // 右键拖拽平移时强制抓手光标(zrender 会改写 canvas 内联 cursor) + &.is-panning, + &.is-panning * { + cursor: grabbing !important; + } } .chart-empty { diff --git a/go/controllers/backend_ai_chat.go b/go/controllers/backend_ai_chat.go index 0f3f7dd..aa80e2b 100644 --- a/go/controllers/backend_ai_chat.go +++ b/go/controllers/backend_ai_chat.go @@ -54,10 +54,20 @@ func (c *BackendAiChatController) chatOk(data interface{}) { } type chatSendPayload struct { - SessionID uint64 `json:"session_id"` - Content string `json:"content"` + SessionID uint64 `json:"session_id"` + Content string `json:"content"` + // ProviderID/Model 由前端选择 ProviderID uint64 `json:"provider_id"` Model string `json:"model"` + // McpEnabled 是否为本轮对话注入 MCP 工具;不传(nil)视为开启,兼容旧客户端 + McpEnabled *bool `json:"mcp_enabled"` + // McpServerIDs 限定参与本轮对话的 MCP 服务器;为空表示使用全部已启用的服务器 + McpServerIDs []uint64 `json:"mcp_server_ids"` +} + +// mcpEnabled 解析本轮是否启用 MCP(未传默认开启) +func (p *chatSendPayload) mcpEnabled() bool { + return p.McpEnabled == nil || *p.McpEnabled } // ============ 通用消息/工具结构 ============ @@ -124,6 +134,28 @@ type mcpToolSummary struct { const maxToolRounds = 6 +// aiUsage 一次 AI 调用的 token 用量 +type aiUsage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` +} + +// add 累加用量 +func (u *aiUsage) add(o aiUsage) { + u.PromptTokens += o.PromptTokens + u.CompletionTokens += o.CompletionTokens + u.TotalTokens += o.TotalTokens +} + +// total 返回本次调用的总 token 数(部分厂商不返回 total,则自行求和) +func (u aiUsage) total() int { + if u.TotalTokens > 0 { + return u.TotalTokens + } + return u.PromptTokens + u.CompletionTokens +} + // sanitizeToolName 工具名规范化为 OpenAI 允许的字符集 func sanitizeToolName(s string) string { var sb strings.Builder @@ -150,16 +182,23 @@ func mcpToolKey(serverID uint64, name string) string { return key } -// collectMcpTools 汇总当前用户「已启用」的 MCP 服务器工具,转换为 LLM 工具列表 -func collectMcpTools(claims *jwtutil.Claims) ([]openaiTool, map[string]toolRef, []mcpToolSummary, error) { +// collectMcpTools 汇总「已启用」的 MCP 服务器工具,转换为 LLM 工具列表 +// enabled=false 时不注入任何工具;serverIDs 非空时只取指定的服务器 +func collectMcpTools(claims *jwtutil.Claims, enabled bool, serverIDs []uint64) ([]openaiTool, map[string]toolRef, []mcpToolSummary, error) { + if !enabled { + return nil, map[string]toolRef{}, nil, nil + } + var servers []models.BackendMcpServer - _, err := models.Orm.QueryTable(new(models.BackendMcpServer)). + qs := models.Orm.QueryTable(new(models.BackendMcpServer)). Filter("tenant_id", fmt.Sprintf("%d", claims.TenantId)). Filter("user_id", uint64(claims.UserID)). Filter("enabled", 1). - Filter("delete_time__isnull", true). - OrderBy("id"). - All(&servers) + Filter("delete_time__isnull", true) + if len(serverIDs) > 0 { + qs = qs.Filter("id__in", serverIDs) + } + _, err := qs.OrderBy("id").All(&servers) if err != nil { return nil, nil, nil, err } @@ -265,12 +304,13 @@ func buildMcpAssistantMessage(text string, pending []pendingToolCall) openaiMess // runToolLoop 运行带已启用 MCP 工具的完整对话循环(非流式),返回最终文本。 // 供聊天非流式接口与智能生成等模块复用。 func runToolLoop(claims *jwtutil.Claims, provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage) (string, error) { - llmTools, refMap, _, _ := collectMcpTools(claims) + // 智能生成等内部调用:默认注入全部已启用的 MCP 工具 + llmTools, refMap, _, _ := collectMcpTools(claims, true, nil) current := messages rounds := 0 for { rounds++ - text, pending, err := callAITools(provider, model, systemPrompt, current, llmTools) + text, pending, _, err := callAITools(provider, model, systemPrompt, current, llmTools) if err != nil { return "", err } @@ -542,19 +582,23 @@ func (c *BackendAiChatController) Send() { } _, _ = models.Orm.Insert(&userMsg) - // 汇总启用的 MCP 工具 - llmTools, refMap, _, _ := collectMcpTools(claims) + // 汇总启用的 MCP 工具(受前端开关与所选服务控制) + llmTools, refMap, _, _ := collectMcpTools(claims, p.mcpEnabled(), p.McpServerIDs) current := messages reply := "" rounds := 0 + startTime := time.Now() + var totalUsage aiUsage for { rounds++ - text, pending, callErr := callAITools(provider, useModel, systemPrompt, current, llmTools) + var roundUsage aiUsage + text, pending, roundUsage, callErr := callAITools(provider, useModel, systemPrompt, current, llmTools) if callErr != nil { c.chatJsonErr(500, 500, "AI调用失败: "+callErr.Error()) return } + totalUsage.add(roundUsage) if len(pending) == 0 { reply = text break @@ -568,12 +612,16 @@ func (c *BackendAiChatController) Send() { current = append(current, toolMsgs...) } + durationMs := int(time.Since(startTime).Milliseconds()) assistantMsg := models.BackendAiChatMessage{ - TenantID: fmt.Sprintf("%d", claims.TenantId), - SessionID: session.ID, - Role: "assistant", - Content: reply, - CreateTime: time.Now(), + TenantID: fmt.Sprintf("%d", claims.TenantId), + SessionID: session.ID, + Role: "assistant", + Content: reply, + Tokens: totalUsage.total(), + DurationMs: durationMs, + Model: useModel, + CreateTime: time.Now(), } _, _ = models.Orm.Insert(&assistantMsg) @@ -581,9 +629,12 @@ func (c *BackendAiChatController) Send() { _, _ = models.Orm.Update(&session, "update_time") c.chatOk(map[string]interface{}{ - "session_id": session.ID, - "reply": reply, - "message_id": assistantMsg.ID, + "session_id": session.ID, + "reply": reply, + "message_id": assistantMsg.ID, + "tokens": totalUsage.total(), + "duration_ms": durationMs, + "model": useModel, }) } @@ -645,8 +696,8 @@ func (c *BackendAiChatController) SendStream() { } _, _ = models.Orm.Insert(&userMsg) - // 汇总启用的 MCP 工具 - llmTools, refMap, toolSummaries, _ := collectMcpTools(claims) + // 汇总启用的 MCP 工具(受前端开关与所选服务控制) + llmTools, refMap, toolSummaries, _ := collectMcpTools(claims, p.mcpEnabled(), p.McpServerIDs) // 设置SSE响应头 rw := c.Ctx.ResponseWriter.ResponseWriter @@ -672,8 +723,11 @@ func (c *BackendAiChatController) SendStream() { writeSSE(event, string(b)) } - // session 事件 - writeJSON("session", map[string]interface{}{"session_id": session.ID}) + // session 事件(带上用户消息ID,便于前端后续编辑/删除该条消息) + writeJSON("session", map[string]interface{}{ + "session_id": session.ID, + "user_message_id": userMsg.ID, + }) // 当前启用的工具列表(前端展示) if len(toolSummaries) > 0 { @@ -684,17 +738,21 @@ func (c *BackendAiChatController) SendStream() { rounds := 0 fullReply := "" streamErr := error(nil) + startTime := time.Now() + var totalUsage aiUsage for { rounds++ var text string var pending []pendingToolCall - text, pending, streamErr = callAIStreamTools(provider, useModel, systemPrompt, current, llmTools, func(chunk string) { + var roundUsage aiUsage + text, pending, roundUsage, streamErr = callAIStreamTools(provider, useModel, systemPrompt, current, llmTools, func(chunk string) { writeJSON("content", map[string]string{"content": chunk}) }) if streamErr != nil { break } + totalUsage.add(roundUsage) if len(pending) == 0 { fullReply += text break @@ -736,12 +794,16 @@ func (c *BackendAiChatController) SendStream() { } // 保存AI回复 + durationMs := int(time.Since(startTime).Milliseconds()) assistantMsg := models.BackendAiChatMessage{ - TenantID: fmt.Sprintf("%d", claims.TenantId), - SessionID: session.ID, - Role: "assistant", - Content: fullReply, - CreateTime: time.Now(), + TenantID: fmt.Sprintf("%d", claims.TenantId), + SessionID: session.ID, + Role: "assistant", + Content: fullReply, + Tokens: totalUsage.total(), + DurationMs: durationMs, + Model: useModel, + CreateTime: time.Now(), } _, _ = models.Orm.Insert(&assistantMsg) @@ -749,8 +811,11 @@ func (c *BackendAiChatController) SendStream() { _, _ = models.Orm.Update(&session, "update_time") writeJSON("done", map[string]interface{}{ - "session_id": session.ID, - "message_id": assistantMsg.ID, + "session_id": session.ID, + "message_id": assistantMsg.ID, + "tokens": totalUsage.total(), + "duration_ms": durationMs, + "model": useModel, }) } @@ -758,7 +823,7 @@ func (c *BackendAiChatController) SendStream() { // callAI 兼容入口(无 MCP 工具),供智能生成等模块使用 func callAI(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage) (string, error) { - text, pending, err := callAITools(provider, model, systemPrompt, messages, nil) + text, pending, _, err := callAITools(provider, model, systemPrompt, messages, nil) if err != nil { return "", err } @@ -767,7 +832,7 @@ func callAI(provider models.BackendAiProvider, model string, systemPrompt string return text, nil } -func callAITools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool) (string, []pendingToolCall, error) { +func callAITools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool) (string, []pendingToolCall, aiUsage, error) { client := &http.Client{Timeout: 180 * time.Second} if provider.ProviderType == "openai" { @@ -790,20 +855,20 @@ func callAITools(provider models.BackendAiProvider, model string, systemPrompt s } req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData)) if err != nil { - return "", nil, err + return "", nil, aiUsage{}, err } req.Header.Set("Content-Type", "application/json") req.Header.Set("Authorization", "Bearer "+provider.ApiKey) resp, err := client.Do(req) if err != nil { - return "", nil, err + return "", nil, aiUsage{}, err } defer resp.Body.Close() bodyBytes, _ := io.ReadAll(resp.Body) if resp.StatusCode != 200 { - return "", nil, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(bodyBytes)) + return "", nil, aiUsage{}, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(bodyBytes)) } var result struct { Choices []struct { @@ -812,25 +877,38 @@ func callAITools(provider models.BackendAiProvider, model string, systemPrompt s ToolCalls []openaiToolCall `json:"tool_calls"` } `json:"message"` } `json:"choices"` + Usage *struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` + } `json:"usage"` Error *struct { Message string `json:"message"` } `json:"error"` } if err := json.Unmarshal(bodyBytes, &result); err != nil { - return "", nil, fmt.Errorf("解析响应失败: %s", string(bodyBytes)) + return "", nil, aiUsage{}, fmt.Errorf("解析响应失败: %s", string(bodyBytes)) } if result.Error != nil { - return "", nil, fmt.Errorf(result.Error.Message) + return "", nil, aiUsage{}, fmt.Errorf(result.Error.Message) } if len(result.Choices) == 0 { - return "", nil, fmt.Errorf("AI未返回内容") + return "", nil, aiUsage{}, fmt.Errorf("AI未返回内容") + } + var usage aiUsage + if result.Usage != nil { + usage = aiUsage{ + PromptTokens: result.Usage.PromptTokens, + CompletionTokens: result.Usage.CompletionTokens, + TotalTokens: result.Usage.TotalTokens, + } } content := result.Choices[0].Message.Content var pending []pendingToolCall for _, tc := range result.Choices[0].Message.ToolCalls { pending = append(pending, parseOpenAIToolCall(tc)) } - return content, pending, nil + return content, pending, usage, nil } // Anthropic 非流式 @@ -853,7 +931,7 @@ func callAITools(provider models.BackendAiProvider, model string, systemPrompt s } req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData)) if err != nil { - return "", nil, err + return "", nil, aiUsage{}, err } req.Header.Set("Content-Type", "application/json") req.Header.Set("x-api-key", provider.ApiKey) @@ -861,13 +939,13 @@ func callAITools(provider models.BackendAiProvider, model string, systemPrompt s resp, err := client.Do(req) if err != nil { - return "", nil, err + return "", nil, aiUsage{}, err } defer resp.Body.Close() bodyBytes, _ := io.ReadAll(resp.Body) if resp.StatusCode != 200 { - return "", nil, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(bodyBytes)) + return "", nil, aiUsage{}, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(bodyBytes)) } var result struct { Content []struct { @@ -877,15 +955,27 @@ func callAITools(provider models.BackendAiProvider, model string, systemPrompt s Name string `json:"name"` Input map[string]interface{} `json:"input"` } `json:"content"` + Usage *struct { + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` + } `json:"usage"` Error *struct { Message string `json:"message"` } `json:"error"` } if err := json.Unmarshal(bodyBytes, &result); err != nil { - return "", nil, fmt.Errorf("解析响应失败: %s", string(bodyBytes)) + return "", nil, aiUsage{}, fmt.Errorf("解析响应失败: %s", string(bodyBytes)) } if result.Error != nil { - return "", nil, fmt.Errorf(result.Error.Message) + return "", nil, aiUsage{}, fmt.Errorf(result.Error.Message) + } + var usage aiUsage + if result.Usage != nil { + usage = aiUsage{ + PromptTokens: result.Usage.InputTokens, + CompletionTokens: result.Usage.OutputTokens, + TotalTokens: result.Usage.InputTokens + result.Usage.OutputTokens, + } } var text strings.Builder var pending []pendingToolCall @@ -897,7 +987,7 @@ func callAITools(provider models.BackendAiProvider, model string, systemPrompt s pending = append(pending, pendingToolCall{ID: block.ID, Name: block.Name, Args: block.Input}) } } - return text.String(), pending, nil + return text.String(), pending, usage, nil } // parseOpenAIToolCall 解析非流式 OpenAI 工具调用 @@ -917,7 +1007,7 @@ func parseOpenAIToolCall(tc openaiToolCall) pendingToolCall { // ============ AI 调用(流式,含工具) ============ -func callAIStreamTools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool, onChunk func(string)) (string, []pendingToolCall, error) { +func callAIStreamTools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool, onChunk func(string)) (string, []pendingToolCall, aiUsage, error) { if provider.ProviderType == "openai" { return callOpenAIStreamTools(provider, model, systemPrompt, messages, tools, onChunk) } @@ -925,7 +1015,7 @@ func callAIStreamTools(provider models.BackendAiProvider, model string, systemPr } // callOpenAIStreamTools OpenAI 兼容流式工具调用 -func callOpenAIStreamTools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool, onChunk func(string)) (string, []pendingToolCall, error) { +func callOpenAIStreamTools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool, onChunk func(string)) (string, []pendingToolCall, aiUsage, error) { client := &http.Client{Timeout: 180 * time.Second} reqMessages := messages @@ -940,6 +1030,8 @@ func callOpenAIStreamTools(provider models.BackendAiProvider, model string, syst if len(tools) > 0 { reqBody["tools"] = tools } + // 请求在最后一个 chunk 返回 token 用量统计 + reqBody["stream_options"] = map[string]interface{}{"include_usage": true} jsonData, _ := json.Marshal(reqBody) url := strings.TrimRight(provider.ApiBase, "/") @@ -948,23 +1040,24 @@ func callOpenAIStreamTools(provider models.BackendAiProvider, model string, syst } req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData)) if err != nil { - return "", nil, err + return "", nil, aiUsage{}, err } req.Header.Set("Content-Type", "application/json") req.Header.Set("Authorization", "Bearer "+provider.ApiKey) resp, err := client.Do(req) if err != nil { - return "", nil, err + return "", nil, aiUsage{}, err } defer resp.Body.Close() if resp.StatusCode != 200 { body, _ := io.ReadAll(resp.Body) - return "", nil, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(body)) + return "", nil, aiUsage{}, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(body)) } var text strings.Builder + var usage aiUsage acc := make(map[int]*struct { ID string Name string @@ -991,6 +1084,11 @@ func callOpenAIStreamTools(provider models.BackendAiProvider, model string, syst } `json:"delta"` FinishReason string `json:"finish_reason"` } `json:"choices"` + Usage *struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` + } `json:"usage"` Error *struct { Message string `json:"message"` } `json:"error"` @@ -999,7 +1097,14 @@ func callOpenAIStreamTools(provider models.BackendAiProvider, model string, syst continue } if chunk.Error != nil { - return "", nil, fmt.Errorf(chunk.Error.Message) + return "", nil, aiUsage{}, fmt.Errorf(chunk.Error.Message) + } + if chunk.Usage != nil { + usage = aiUsage{ + PromptTokens: chunk.Usage.PromptTokens, + CompletionTokens: chunk.Usage.CompletionTokens, + TotalTokens: chunk.Usage.TotalTokens, + } } if len(chunk.Choices) == 0 { continue @@ -1052,11 +1157,11 @@ func callOpenAIStreamTools(provider models.BackendAiProvider, model string, syst } pending = append(pending, pc) } - return text.String(), pending, nil + return text.String(), pending, usage, nil } // callAnthropicStreamTools Anthropic 流式工具调用 -func callAnthropicStreamTools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool, onChunk func(string)) (string, []pendingToolCall, error) { +func callAnthropicStreamTools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool, onChunk func(string)) (string, []pendingToolCall, aiUsage, error) { client := &http.Client{Timeout: 180 * time.Second} reqBody := map[string]interface{}{ @@ -1079,7 +1184,7 @@ func callAnthropicStreamTools(provider models.BackendAiProvider, model string, s } req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData)) if err != nil { - return "", nil, err + return "", nil, aiUsage{}, err } req.Header.Set("Content-Type", "application/json") req.Header.Set("x-api-key", provider.ApiKey) @@ -1087,16 +1192,17 @@ func callAnthropicStreamTools(provider models.BackendAiProvider, model string, s resp, err := client.Do(req) if err != nil { - return "", nil, err + return "", nil, aiUsage{}, err } defer resp.Body.Close() if resp.StatusCode != 200 { body, _ := io.ReadAll(resp.Body) - return "", nil, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(body)) + return "", nil, aiUsage{}, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(body)) } var text strings.Builder + var usage aiUsage type toolUseAcc struct { ID string Name string @@ -1127,10 +1233,16 @@ func callAnthropicStreamTools(provider models.BackendAiProvider, model string, s ID string `json:"id"` Name string `json:"name"` } `json:"content_block"` + Usage *struct { + InputTokens int `json:"input_tokens"` + } `json:"usage"` } if err := json.Unmarshal([]byte(data), &block); err != nil { continue } + if block.Usage != nil { + usage.PromptTokens = block.Usage.InputTokens + } if block.ContentBlock.Type == "tool_use" { if _, ok := acc[block.Index]; !ok { acc[block.Index] = &toolUseAcc{} @@ -1169,9 +1281,15 @@ func callAnthropicStreamTools(provider models.BackendAiProvider, model string, s Delta struct { StopReason string `json:"stop_reason"` } `json:"delta"` + Usage *struct { + OutputTokens int `json:"output_tokens"` + } `json:"usage"` } if err := json.Unmarshal([]byte(data), &md); err == nil { _ = md.Delta.StopReason + if md.Usage != nil { + usage.CompletionTokens = md.Usage.OutputTokens + } } case "message_stop": // 结束 @@ -1198,7 +1316,8 @@ func callAnthropicStreamTools(provider models.BackendAiProvider, model string, s } pending = append(pending, pc) } - return text.String(), pending, nil + usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens + return text.String(), pending, usage, nil } // ============ Anthropic 消息/工具格式转换 ============ diff --git a/go/docs/sql/alter_backend_ai_chat_message_stats.sql b/go/docs/sql/alter_backend_ai_chat_message_stats.sql new file mode 100644 index 0000000..5e7d87d --- /dev/null +++ b/go/docs/sql/alter_backend_ai_chat_message_stats.sql @@ -0,0 +1,6 @@ +-- 为 AI 聊天消息表增加 token 用量与响应耗时字段 +-- 用途:在对话界面展示每条助手回复的响应时间与消耗的 tokens +ALTER TABLE yz_backend_ai_chat_message + ADD COLUMN tokens INT NOT NULL DEFAULT 0 COMMENT '本次响应消耗的 token 总数', + ADD COLUMN duration_ms INT NOT NULL DEFAULT 0 COMMENT 'AI 响应耗时(毫秒)', + ADD COLUMN model VARCHAR(64) NOT NULL DEFAULT '' COMMENT '生成该消息使用的模型'; diff --git a/go/models/backend_ai_chat_message.go b/go/models/backend_ai_chat_message.go index f2e26ba..f07c8a9 100644 --- a/go/models/backend_ai_chat_message.go +++ b/go/models/backend_ai_chat_message.go @@ -9,6 +9,12 @@ type BackendAiChatMessage struct { SessionID uint64 `orm:"column(session_id)" json:"session_id"` Role string `orm:"column(role);size(20)" json:"role"` // user / assistant Content string `orm:"column(content);type(text)" json:"content"` + // Tokens 本次响应消耗的 token 总数(0 表示未统计到) + Tokens int `orm:"column(tokens)" json:"tokens"` + // DurationMs AI 响应耗时(毫秒,含工具调用轮次) + DurationMs int `orm:"column(duration_ms)" json:"duration_ms"` + // Model 生成该消息使用的模型 + Model string `orm:"column(model);size(64)" json:"model"` CreateTime time.Time `orm:"column(create_time);auto_now_add;type(datetime)" json:"create_time"` }