From bfd1cbac56cba2cda4c12060e0adbbe0152aa49d Mon Sep 17 00:00:00 2001 From: matevip Date: Sat, 11 Apr 2026 08:46:21 +0800 Subject: [PATCH] feat(chat): ChatGPT tool calling + fix cross-turn message pollution --- .../mate/llm/chatgpt/ChatGPTChatModel.java | 110 ++++++++-- .../llm/chatgpt/ChatGPTResponsesClient.java | 199 ++++++++++++++---- .../src/main/resources/db/data-mysql-zh.sql | 4 +- .../src/main/resources/db/data-zh.sql | 4 +- .../src/components/chat/MessageBubble.vue | 12 ++ mateclaw-ui/src/composables/chat/useChat.ts | 116 ++++++---- .../src/composables/chat/useMessages.ts | 12 +- mateclaw-ui/src/utils/messageReconcile.ts | 8 +- mateclaw-ui/src/views/ChatConsole.vue | 32 +-- text.txt | 77 +++++++ 10 files changed, 450 insertions(+), 124 deletions(-) create mode 100644 text.txt diff --git a/mateclaw-server/src/main/java/vip/mate/llm/chatgpt/ChatGPTChatModel.java b/mateclaw-server/src/main/java/vip/mate/llm/chatgpt/ChatGPTChatModel.java index b0a1609d..ca3a5fb5 100644 --- a/mateclaw-server/src/main/java/vip/mate/llm/chatgpt/ChatGPTChatModel.java +++ b/mateclaw-server/src/main/java/vip/mate/llm/chatgpt/ChatGPTChatModel.java @@ -10,13 +10,17 @@ import org.springframework.ai.chat.model.ChatResponse; import org.springframework.ai.chat.model.Generation; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.model.tool.ToolCallingChatOptions; +import org.springframework.ai.tool.ToolCallback; +import org.springframework.ai.tool.definition.ToolDefinition; import reactor.core.publisher.Flux; -import java.util.List; +import java.util.*; /** - * ChatGPT 会员模型 — 实现 Spring AI ChatModel 接口, - * 内部通过 ChatGPTResponsesClient 调用 chatgpt.com/backend-api。 + * ChatGPT 会员模型 — 实现 Spring AI ChatModel 接口。 + * 支持 tool calling:从 Prompt 的 ChatOptions 中提取 toolCallbacks, + * 传递给 ChatGPTResponsesClient,并将响应中的 function_call 转换为 ToolCall。 */ @Slf4j public class ChatGPTChatModel implements ChatModel { @@ -36,9 +40,10 @@ public class ChatGPTChatModel implements ChatModel { List messages = prompt.getInstructions(); String model = resolveModel(prompt); Double temp = resolveTemperature(prompt); + List toolDefs = extractToolDefinitions(prompt); - log.debug("ChatGPT call: model={}, messages={}", model, messages.size()); - String content = client.call(model, messages, temp); + log.debug("[ChatGPT] call: model={}, messages={}, tools={}", model, messages.size(), toolDefs.size()); + String content = client.call(model, messages, temp, toolDefs); Generation generation = new Generation(new AssistantMessage(content), ChatGenerationMetadata.builder().finishReason("stop").build()); @@ -51,14 +56,65 @@ public class ChatGPTChatModel implements ChatModel { List messages = prompt.getInstructions(); String model = resolveModel(prompt); Double temp = resolveTemperature(prompt); + List toolDefs = extractToolDefinitions(prompt); - log.debug("ChatGPT stream: model={}, messages={}", model, messages.size()); - return client.stream(model, messages, temp) - .map(delta -> { - Generation generation = new Generation(new AssistantMessage(delta), - ChatGenerationMetadata.builder().finishReason(null).build()); - return new ChatResponse(List.of(generation), - ChatResponseMetadata.builder().model(model).build()); + log.debug("[ChatGPT] stream: model={}, messages={}, tools={}", model, messages.size(), toolDefs.size()); + + // 状态:累积 tool call arguments + Map toolCallNames = new LinkedHashMap<>(); + Map toolCallArgs = new LinkedHashMap<>(); + + return client.streamEvents(model, messages, temp, toolDefs) + .mapNotNull(event -> { + switch (event.type()) { + case "text" -> { + Generation gen = new Generation(new AssistantMessage(event.content()), + ChatGenerationMetadata.builder().finishReason(null).build()); + return new ChatResponse(List.of(gen), + ChatResponseMetadata.builder().model(model).build()); + } + case "tool_call_start" -> { + // 创建初始 ToolCall(空 arguments),让 NodeStreamingChatHelper 创建 accumulator + toolCallNames.put(event.toolCallId(), event.toolName()); + toolCallArgs.put(event.toolCallId(), new StringBuilder()); + List startCalls = List.of( + new AssistantMessage.ToolCall(event.toolCallId(), "function", event.toolName(), "") + ); + AssistantMessage startMsg = AssistantMessage.builder() + .content("") + .toolCalls(startCalls) + .build(); + Generation startGen = new Generation(startMsg, + ChatGenerationMetadata.builder().finishReason(null).build()); + return new ChatResponse(List.of(startGen), + ChatResponseMetadata.builder().model(model).build()); + } + case "tool_call_args_delta" -> { + // 增量追加 arguments(通过空 id 的 ToolCall 让 accumulator 追加) + StringBuilder sb = toolCallArgs.get(event.toolCallId()); + if (sb != null) sb.append(event.toolArgsDelta()); + List deltaCalls = List.of( + new AssistantMessage.ToolCall("", "function", "", event.toolArgsDelta()) + ); + AssistantMessage deltaMsg = AssistantMessage.builder() + .content("") + .toolCalls(deltaCalls) + .build(); + Generation deltaGen = new Generation(deltaMsg, + ChatGenerationMetadata.builder().finishReason(null).build()); + return new ChatResponse(List.of(deltaGen), + ChatResponseMetadata.builder().model(model).build()); + } + case "tool_call_done" -> { + // 不需要再发一次完整的 — accumulator 已经有了 + return null; + } + case "done" -> { + // 流结束,如果没有任何 tool call 产生过 done event 但有未完成的,忽略 + return null; + } + default -> { return null; } + } }); } @@ -70,6 +126,36 @@ public class ChatGPTChatModel implements ChatModel { .build(); } + /** + * 从 Prompt 的 ChatOptions 中提取 ToolDefinition 列表 + */ + private List extractToolDefinitions(Prompt prompt) { + ChatOptions options = prompt.getOptions(); + if (options == null) return List.of(); + + // ToolCallingChatOptions 或 OpenAiChatOptions 都可能包含 toolCallbacks + List callbacks = null; + if (options instanceof ToolCallingChatOptions tcOpts) { + callbacks = tcOpts.getToolCallbacks(); + } else { + // 尝试反射获取(Spring AI 的 OpenAiChatOptions 也有 toolCallbacks) + try { + var method = options.getClass().getMethod("getToolCallbacks"); + @SuppressWarnings("unchecked") + var result = (List) method.invoke(options); + callbacks = result; + } catch (Exception ignored) { + // 不支持 tool callbacks + } + } + + if (callbacks == null || callbacks.isEmpty()) return List.of(); + + return callbacks.stream() + .map(ToolCallback::getToolDefinition) + .toList(); + } + private String resolveModel(Prompt prompt) { if (prompt.getOptions() != null && prompt.getOptions().getModel() != null) { return prompt.getOptions().getModel(); diff --git a/mateclaw-server/src/main/java/vip/mate/llm/chatgpt/ChatGPTResponsesClient.java b/mateclaw-server/src/main/java/vip/mate/llm/chatgpt/ChatGPTResponsesClient.java index 9bebc8ce..007e1a89 100644 --- a/mateclaw-server/src/main/java/vip/mate/llm/chatgpt/ChatGPTResponsesClient.java +++ b/mateclaw-server/src/main/java/vip/mate/llm/chatgpt/ChatGPTResponsesClient.java @@ -7,6 +7,7 @@ import com.fasterxml.jackson.databind.node.ObjectNode; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; import org.springframework.ai.chat.messages.*; +import org.springframework.ai.tool.definition.ToolDefinition; import org.springframework.http.HttpHeaders; import org.springframework.http.MediaType; import org.springframework.stereotype.Component; @@ -15,10 +16,11 @@ import reactor.core.publisher.Flux; import vip.mate.exception.MateClawException; import vip.mate.llm.oauth.OpenAIOAuthService; -import java.util.List; +import java.util.*; /** * ChatGPT Backend API 客户端 — 调用 chatgpt.com/backend-api/codex/responses(Responses API 格式) + * 支持 tool calling(function_call)。 */ @Slf4j @Component @@ -33,25 +35,41 @@ public class ChatGPTResponsesClient { private final WebClient webClient = WebClient.create(); /** - * 同步调用 — ChatGPT Backend API 强制要求 stream=true, - * 所以实际仍走 SSE,只是收集完整响应后再返回。 + * 流式调用结果 — 包含文本增量和 tool call 事件 */ - public String call(String model, List messages, Double temperature) { - return stream(model, messages, temperature) + public record StreamEvent(String type, String content, String toolCallId, String toolName, String toolArgsDelta) { + public static StreamEvent text(String delta) { return new StreamEvent("text", delta, null, null, null); } + public static StreamEvent toolCallStart(String callId, String name) { return new StreamEvent("tool_call_start", null, callId, name, null); } + public static StreamEvent toolCallArgsDelta(String callId, String delta) { return new StreamEvent("tool_call_args_delta", null, callId, null, delta); } + public static StreamEvent toolCallDone(String callId, String args) { return new StreamEvent("tool_call_done", null, callId, null, args); } + public static StreamEvent done() { return new StreamEvent("done", null, null, null, null); } + } + + /** + * 同步调用 — 收集完整响应(仅文本部分) + */ + public String call(String model, List messages, Double temperature, List tools) { + return streamEvents(model, messages, temperature, tools) + .filter(e -> "text".equals(e.type())) + .map(StreamEvent::content) .collectList() .map(chunks -> String.join("", chunks)) .block(); } /** - * 流式调用 Responses API (SSE) + * 流式调用 — 返回结构化事件(文本 + tool_call) */ - public Flux stream(String model, List messages, Double temperature) { + public Flux streamEvents(String model, List messages, Double temperature, + List tools) { String accessToken = oauthService.ensureValidAccessToken(); String accountId = oauthService.getAccountId(); - ObjectNode requestBody = buildRequestBody(model, messages, temperature); + ObjectNode requestBody = buildRequestBody(model, messages, temperature, tools); String bodyJson = requestBody.toString(); - log.info("ChatGPT request body: {}", bodyJson); + log.info("[ChatGPT] Request: model={}, messages={}, tools={}", model, messages.size(), + tools != null ? tools.size() : 0); + log.debug("[ChatGPT] Request body: {}", bodyJson.length() > 2000 + ? bodyJson.substring(0, 2000) + "..." : bodyJson); return webClient.post() .uri(BASE_URL + RESPONSES_PATH) @@ -63,33 +81,42 @@ public class ChatGPTResponsesClient { .onStatus(status -> status.is4xxClientError() || status.is5xxServerError(), response -> response.bodyToMono(String.class) .map(errorBody -> { - log.error("ChatGPT API error {}: {}", response.statusCode(), errorBody); + log.error("[ChatGPT] API error {}: {}", response.statusCode(), errorBody); return new MateClawException("ChatGPT API " + response.statusCode() + ": " + errorBody); })) .bodyToFlux(String.class) - .doOnNext(raw -> log.debug("ChatGPT SSE raw: {}", raw.length() > 200 ? raw.substring(0, 200) + "..." : raw)) + .doOnNext(raw -> log.debug("[ChatGPT] SSE raw: {}", raw.length() > 200 ? raw.substring(0, 200) + "..." : raw)) .filter(line -> !line.isBlank() && !line.equals("[DONE]")) + .filter(line -> line.startsWith("data:")) .map(line -> { - // SSE 格式:每行以 "data: " 开头,需要去掉前缀 if (line.startsWith("data: ")) return line.substring(6); - if (line.startsWith("data:")) return line.substring(5); - return line; + return line.substring(5); }) .filter(line -> !line.isBlank() && !line.equals("[DONE]")) - .mapNotNull(this::extractDeltaContent) + .mapNotNull(this::parseSSEEvent) .onErrorMap(e -> e instanceof MateClawException ? e : new MateClawException("ChatGPT 流式调用失败: " + e.getMessage())); } + /** + * 纯文本流(向后兼容) + */ + public Flux stream(String model, List messages, Double temperature) { + return streamEvents(model, messages, temperature, null) + .filter(e -> "text".equals(e.type())) + .map(StreamEvent::content); + } + // ==================== 请求构建 ==================== - ObjectNode buildRequestBody(String model, List messages, Double temperature) { + ObjectNode buildRequestBody(String model, List messages, Double temperature, + List tools) { ObjectNode body = objectMapper.createObjectNode(); body.put("model", model); - body.put("stream", true); // ChatGPT Backend API 强制要求 stream=true + body.put("stream", true); body.put("store", false); - // 从 messages 中提取 system prompt → instructions + // system prompt → instructions String systemPrompt = null; for (Message msg : messages) { if (msg.getMessageType() == MessageType.SYSTEM) { @@ -101,14 +128,13 @@ public class ChatGPTResponsesClient { body.put("instructions", systemPrompt); } - // 非 system 消息 → input 数组(Responses API 格式) + // 消息 → input 数组(Responses API 格式) ArrayNode input = objectMapper.createArrayNode(); int msgIndex = 0; for (Message msg : messages) { if (msg.getMessageType() == MessageType.SYSTEM) continue; if (msg.getMessageType() == MessageType.USER) { - // User: content 必须是 [{ type: "input_text", text: "..." }] 格式 ObjectNode item = objectMapper.createObjectNode(); item.put("role", "user"); ArrayNode contentArr = objectMapper.createArrayNode(); @@ -119,35 +145,92 @@ public class ChatGPTResponsesClient { item.set("content", contentArr); input.add(item); } else if (msg.getMessageType() == MessageType.ASSISTANT) { - // Assistant: 转为 output message item - ObjectNode item = objectMapper.createObjectNode(); - item.put("type", "message"); - item.put("role", "assistant"); - item.put("id", "msg_" + msgIndex); - ArrayNode contentArr = objectMapper.createArrayNode(); - ObjectNode textPart = objectMapper.createObjectNode(); - textPart.put("type", "output_text"); - textPart.put("text", msg.getText() != null ? msg.getText() : ""); - contentArr.add(textPart); - item.set("content", contentArr); - input.add(item); + AssistantMessage assistantMsg = (AssistantMessage) msg; + + // 如果 assistant 消息包含 tool calls,需要输出 function_call items + if (assistantMsg.hasToolCalls()) { + // 先输出文本部分(如果有) + String text = assistantMsg.getText(); + if (text != null && !text.isBlank()) { + ObjectNode textItem = objectMapper.createObjectNode(); + textItem.put("type", "message"); + textItem.put("role", "assistant"); + textItem.put("id", "msg_" + msgIndex); + ArrayNode contentArr = objectMapper.createArrayNode(); + ObjectNode textPart = objectMapper.createObjectNode(); + textPart.put("type", "output_text"); + textPart.put("text", text); + contentArr.add(textPart); + textItem.set("content", contentArr); + input.add(textItem); + } + // 输出 function_call items + for (AssistantMessage.ToolCall tc : assistantMsg.getToolCalls()) { + ObjectNode fcItem = objectMapper.createObjectNode(); + fcItem.put("type", "function_call"); + fcItem.put("call_id", tc.id()); + fcItem.put("name", tc.name()); + fcItem.put("arguments", tc.arguments()); + input.add(fcItem); + } + } else { + ObjectNode item = objectMapper.createObjectNode(); + item.put("type", "message"); + item.put("role", "assistant"); + item.put("id", "msg_" + msgIndex); + ArrayNode contentArr = objectMapper.createArrayNode(); + ObjectNode textPart = objectMapper.createObjectNode(); + textPart.put("type", "output_text"); + textPart.put("text", msg.getText() != null ? msg.getText() : ""); + contentArr.add(textPart); + item.set("content", contentArr); + input.add(item); + } + } else if (msg.getMessageType() == MessageType.TOOL) { + // Tool result → function_call_output + ToolResponseMessage toolMsg = (ToolResponseMessage) msg; + for (ToolResponseMessage.ToolResponse response : toolMsg.getResponses()) { + ObjectNode fcoItem = objectMapper.createObjectNode(); + fcoItem.put("type", "function_call_output"); + fcoItem.put("call_id", response.id()); + fcoItem.put("output", response.responseData()); + input.add(fcoItem); + } } msgIndex++; } body.set("input", input); - // 注意:ChatGPT Backend API 的部分模型(如 gpt-5.4)不支持 temperature, - // 仅对非推理类旧模型(如 gpt-4o)传递此参数 + // temperature(推理类模型不支持) if (temperature != null && !model.startsWith("gpt-5") && !model.startsWith("o")) { body.put("temperature", temperature); } + // tools 数组(Responses API flat format) + if (tools != null && !tools.isEmpty()) { + ArrayNode toolsArr = objectMapper.createArrayNode(); + for (ToolDefinition tool : tools) { + ObjectNode toolNode = objectMapper.createObjectNode(); + toolNode.put("type", "function"); + toolNode.put("name", tool.name()); + toolNode.put("description", tool.description()); + try { + JsonNode params = objectMapper.readTree(tool.inputSchema()); + toolNode.set("parameters", params); + } catch (Exception e) { + log.warn("[ChatGPT] Failed to parse tool schema for {}: {}", tool.name(), e.getMessage()); + } + toolsArr.add(toolNode); + } + body.set("tools", toolsArr); + body.put("tool_choice", "auto"); + } + // Responses API 特有参数 ObjectNode text = objectMapper.createObjectNode(); text.put("verbosity", "medium"); body.set("text", text); - // include reasoning(OpenClaw 的标准参数) ArrayNode include = objectMapper.createArrayNode(); include.add("reasoning.encrypted_content"); body.set("include", include); @@ -158,28 +241,56 @@ public class ChatGPTResponsesClient { // ==================== 响应解析 ==================== /** - * 从 SSE delta 事件中提取增量文本 + * 解析 SSE 事件 — 支持文本增量和 function_call 事件 */ - private String extractDeltaContent(String eventData) { + private StreamEvent parseSSEEvent(String eventData) { try { JsonNode node = objectMapper.readTree(eventData); String type = node.path("type").asText(""); - // response.output_text.delta — 文本增量 + // 文本增量 if ("response.output_text.delta".equals(type)) { - return node.path("delta").asText(null); + String delta = node.path("delta").asText(null); + return delta != null ? StreamEvent.text(delta) : null; } - // response.completed / response.done — 结束信号 + // function_call 开始(response.output_item.added with type=function_call) + if ("response.output_item.added".equals(type)) { + JsonNode item = node.path("item"); + if ("function_call".equals(item.path("type").asText(""))) { + String callId = item.path("call_id").asText(""); + String name = item.path("name").asText(""); + log.info("[ChatGPT] Tool call started: name={}, callId={}", name, callId); + return StreamEvent.toolCallStart(callId, name); + } + } + + // function_call arguments 增量 + if ("response.function_call_arguments.delta".equals(type)) { + String callId = node.path("call_id").asText(""); + String delta = node.path("delta").asText(""); + return StreamEvent.toolCallArgsDelta(callId, delta); + } + + // function_call arguments 完成 + if ("response.function_call_arguments.done".equals(type)) { + String callId = node.path("call_id").asText(""); + String args = node.path("arguments").asText("{}"); + log.info("[ChatGPT] Tool call done: callId={}, args={}", callId, + args.length() > 200 ? args.substring(0, 200) + "..." : args); + return StreamEvent.toolCallDone(callId, args); + } + + // 完成/结束 if (type.startsWith("response.completed") || type.startsWith("response.done")) { - return null; + return StreamEvent.done(); } - // response.failed — 错误 + // 错误 if ("response.failed".equals(type)) { String error = node.path("response").path("error").path("message").asText("Unknown error"); - log.error("ChatGPT Responses API 返回错误: {}", error); - throw new MateClawException("ChatGPT ���回错误: " + error); + log.error("[ChatGPT] Responses API error: {}", error); + throw new MateClawException("ChatGPT 返回错误: " + error); } return null; diff --git a/mateclaw-server/src/main/resources/db/data-mysql-zh.sql b/mateclaw-server/src/main/resources/db/data-mysql-zh.sql index 78dd43d5..812bba94 100644 --- a/mateclaw-server/src/main/resources/db/data-mysql-zh.sql +++ b/mateclaw-server/src/main/resources/db/data-mysql-zh.sql @@ -1568,7 +1568,7 @@ VALUES ( ## 边界 - 私密的保持私密。 -- 写文件和执行命令需要用户确认。 +- 需要执行文件操作或命令时,直接调用对应的工具(如 execute_shell_command、read_file 等),不要用文本描述你要做什么。系统会自动对危险操作弹出审批确认。 - 拿不准就先问。 ## 风格 @@ -1715,7 +1715,7 @@ VALUES ( ## 边界 - 私密的保持私密。 -- 写文件和执行命令需要用户确认。 +- 需要执行文件操作或命令时,直接调用对应的工具(如 execute_shell_command、read_file 等),不要用文本描述你要做什么。系统会自动对危险操作弹出审批确认。 - 拿不准就先问。 ## 风格 diff --git a/mateclaw-server/src/main/resources/db/data-zh.sql b/mateclaw-server/src/main/resources/db/data-zh.sql index 68b2dbed..ca2c93f2 100644 --- a/mateclaw-server/src/main/resources/db/data-zh.sql +++ b/mateclaw-server/src/main/resources/db/data-zh.sql @@ -1591,7 +1591,7 @@ VALUES ( ## 边界 - 私密的保持私密。 -- 写文件和执行命令需要用户确认。 +- 需要执行文件操作或命令时,直接调用对应的工具(如 execute_shell_command、read_file 等),不要用文本描述你要做什么。系统会自动对危险操作弹出审批确认。 - 拿不准就先问。 ## 风格 @@ -1738,7 +1738,7 @@ VALUES ( ## 边界 - 私密的保持私密。 -- 写文件和执行命令需要用户确认。 +- 需要执行文件操作或命令时,直接调用对应的工具(如 execute_shell_command、read_file 等),不要用文本描述你要做什么。系统会自动对危险操作弹出审批确认。 - 拿不准就先问。 ## 风格 diff --git a/mateclaw-ui/src/components/chat/MessageBubble.vue b/mateclaw-ui/src/components/chat/MessageBubble.vue index b7a6cf35..97bd6ae1 100644 --- a/mateclaw-ui/src/components/chat/MessageBubble.vue +++ b/mateclaw-ui/src/components/chat/MessageBubble.vue @@ -608,6 +608,18 @@ const segments = computed(() => { } } + // 去重:相同 toolName + toolArgs 的 tool_call segment 只保留第一个 + const seenToolCalls = new Set() + const deduped = segs.filter(seg => { + if (seg.type !== 'tool_call') return true + const key = `${seg.toolName}::${seg.toolArgs || ''}` + if (seenToolCalls.has(key)) return false + seenToolCalls.add(key) + return true + }) + segs.length = 0 + segs.push(...deduped) + // 修复历史消息顺序:如果 thinking 被落在 content 后面,提到首个 content 前 // 只处理单个 thinking 段的常见场景,避免破坏复杂交错时间线 const thinkingIndices = segs diff --git a/mateclaw-ui/src/composables/chat/useChat.ts b/mateclaw-ui/src/composables/chat/useChat.ts index 7f2c998c..9361c7d1 100644 --- a/mateclaw-ui/src/composables/chat/useChat.ts +++ b/mateclaw-ui/src/composables/chat/useChat.ts @@ -73,6 +73,8 @@ export interface UseChatReturn { clearMessages: () => void /** 重连到运行中的流 */ reconnectStream: (conversationId: string) => Promise + /** 彻底重置流上下文 — 切换/新建会话时调用 */ + resetForNewConversation: () => void } export interface SendMessageOptions { @@ -115,11 +117,23 @@ export function useChat(options: UseChatOptions): UseChatReturn { const segIdCounter = { value: 0 } const genSegId = () => `seg-${Date.now()}-${segIdCounter.value++}` + /** 当前 turn 的唯一标识 — 确保 flushSegmentsToMessage 不会把旧 turn 的 segments 写到新消息 */ + let activeTurnId = '' + + /** 重置当前 turn 的流式状态 — 必须在每次创建新 assistant placeholder 之前调用 */ + function resetCurrentTurnState() { + currentSegments.value = [] + segIdCounter.value = 0 + activeTurnId = `turn-${Date.now()}-${Math.random().toString(36).slice(2, 6)}` + } + /** 将当前 segments 同步到助手消息的 metadata 中(实时渲染用) */ const flushSegmentsToMessage = () => { if (!currentAssistantId.value || currentSegments.value.length === 0) return const msg = getMessage(currentAssistantId.value) if (!msg) return + // 保护:只写入当前 turn 创建的消息,避免旧 turn segments 污染新消息 + if ((msg as any)._turnId && (msg as any)._turnId !== activeTurnId) return const metadata = parseMetadata((msg as any).metadata) updateMessage(currentAssistantId.value, { ...msg, @@ -257,14 +271,10 @@ export function useChat(options: UseChatOptions): UseChatReturn { return } - // 重置分段列表 - currentSegments.value = [] - segIdCounter.value = 0 - - const assistantMessage = createAssistantMessage('') - if (streamConversationId) { - assistantMessage.conversationId = streamConversationId - } + // 没有 placeholder 时才创建(正常路径 placeholder 已在 sendMessage 中创建) + resetCurrentTurnState() + const assistantMessage = createAssistantMessage('', streamConversationId) + ;(assistantMessage as any)._turnId = activeTurnId currentAssistantId.value = assistantMessage.id as string }) @@ -627,7 +637,9 @@ export function useChat(options: UseChatOptions): UseChatReturn { } } if (!targetId) { - const placeholder = createAssistantMessage('') + resetCurrentTurnState() + const placeholder = createAssistantMessage('', streamConversationId) + ;(placeholder as any)._turnId = activeTurnId targetId = placeholder.id as string currentAssistantId.value = targetId } @@ -766,12 +778,14 @@ export function useChat(options: UseChatOptions): UseChatReturn { const queued = messageQueue.dequeue() const messageContent = data.message || queued?.content || '' if (messageContent) { - const userMessage = createUserMessage(messageContent, queued?.contentParts) - userMessage.conversationId = data.conversationId || streamConversationId + const convId = data.conversationId || streamConversationId + createUserMessage(messageContent, queued?.contentParts, convId) } // 2. 再创建 assistant 占位消息 - const assistantMessage = createAssistantMessage('') - assistantMessage.conversationId = data.conversationId || streamConversationId + resetCurrentTurnState() + const convId2 = data.conversationId || streamConversationId + const assistantMessage = createAssistantMessage('', convId2) + ;(assistantMessage as any)._turnId = activeTurnId currentAssistantId.value = assistantMessage.id as string streamPhase.value = 'thinking' phaseInfo.value = null @@ -893,12 +907,12 @@ export function useChat(options: UseChatOptions): UseChatReturn { try { if (!isApprovalCommand) { - const userMessage = createUserMessage(content, contentParts) - userMessage.conversationId = conversationId + createUserMessage(content, contentParts, conversationId) } - const assistantMessage = createAssistantMessage('') - assistantMessage.conversationId = conversationId + resetCurrentTurnState() + const assistantMessage = createAssistantMessage('', conversationId) + ;(assistantMessage as any)._turnId = activeTurnId currentAssistantId.value = assistantMessage.id as string // contentParts 已由 buildOutgoingParts 包含 file entries,不要重复合并 attachments @@ -949,10 +963,10 @@ export function useChat(options: UseChatOptions): UseChatReturn { } else { // 没有活跃的流,直接发送 messageQueue.clear() - const userMessage = createUserMessage(content, options.contentParts) - userMessage.conversationId = conversationId - const assistantMessage = createAssistantMessage('') - assistantMessage.conversationId = conversationId + createUserMessage(content, options.contentParts, conversationId) + resetCurrentTurnState() + const assistantMessage = createAssistantMessage('', conversationId) + ;(assistantMessage as any)._turnId = activeTurnId currentAssistantId.value = assistantMessage.id as string streamPhase.value = 'thinking' phaseInfo.value = null @@ -969,8 +983,7 @@ export function useChat(options: UseChatOptions): UseChatReturn { // 回退为本地可见消息 + 清队列,避免消息静默丢失。 const failedQueued = messageQueue.dequeue() if (failedQueued) { - const userMessage = createUserMessage(failedQueued.content, failedQueued.contentParts) - userMessage.conversationId = conversationId + createUserMessage(failedQueued.content, failedQueued.contentParts, conversationId) } error.value = new Error('Failed to queue message, please resend') } @@ -983,6 +996,10 @@ export function useChat(options: UseChatOptions): UseChatReturn { // 这样 done 事件能正常到达,onStreamEnd 被触发,消息状态和会话列表都能正确更新。 // 加一个 fallback timeout(3 秒),防止 done 事件因网络问题永远不到达。 const stopGeneration = async () => { + // 在任何 await 之前冻结标识符 + 安装 fallback timer,防止 resetForNewConversation 并发清空后丢失上下文 + const convId = streamConversationId + const assistantId = currentAssistantId.value + // 先取消排队消息 messageQueue.clear() @@ -990,30 +1007,19 @@ export function useChat(options: UseChatOptions): UseChatReturn { streamPhase.value = 'stopped' phaseInfo.value = null - if (streamConversationId) { - try { - await fetchWithAuth(`${baseUrl}/api/v1/chat/${streamConversationId}/stop`, { - method: 'POST', - }) - } catch (e) { - console.warn('[useChat] Stop API failed:', e) - } - } - - // 不立即 disconnect —— 等 done 事件自然到达(后端 doOnCancel 会广播 done) - // 设置 fallback timeout:如果 3 秒内 done 事件没到达,强制清理 - const convId = streamConversationId - const assistantId = currentAssistantId.value + // 在 await 之前安装 fallback timer,确保即使 resetForNewConversation 并发执行也不会遗漏 if (stopFallbackTimer) clearTimeout(stopFallbackTimer) stopFallbackTimer = setTimeout(() => { stopFallbackTimer = null console.warn('[useChat] Stop fallback: done event not received within 3s, force cleanup') - stream.disconnect() + // 只有当 stream 仍属于旧会话时才 disconnect,防止误杀新会话的流 + if (streamConversationId === convId || !streamConversationId) { + stream.disconnect() + } if (currentAssistantId.value === assistantId && assistantId) { setMessageStatus(assistantId, 'stopped') currentAssistantId.value = null } - // 强制触发 onStreamEnd 以刷新会话列表 onStreamEnd?.({ conversationId: convId, reason: 'stopped', @@ -1029,6 +1035,15 @@ export function useChat(options: UseChatOptions): UseChatReturn { if (stopFallbackTimer) { clearTimeout(stopFallbackTimer); stopFallbackTimer = null } unsubscribeError() }) + + // 发送后端 stop 请求(fire-and-forget,不阻塞 resetForNewConversation) + if (convId) { + fetchWithAuth(`${baseUrl}/api/v1/chat/${convId}/stop`, { + method: 'POST', + }).catch(e => { + console.warn('[useChat] Stop API failed:', e) + }) + } } // 取消排队消息 @@ -1059,8 +1074,9 @@ export function useChat(options: UseChatOptions): UseChatReturn { phaseInfo.value = null // 创建 assistant 占位消息用于接收重连后的流数据 - const assistantMessage = createAssistantMessage('') - assistantMessage.conversationId = conversationId + resetCurrentTurnState() + const assistantMessage = createAssistantMessage('', conversationId) + ;(assistantMessage as any)._turnId = activeTurnId currentAssistantId.value = assistantMessage.id as string try { @@ -1111,6 +1127,23 @@ export function useChat(options: UseChatOptions): UseChatReturn { }) } + /** 彻底重置流上下文 — 切换/新建会话时调用,确保旧流状态不污染新会话 */ + const resetForNewConversation = () => { + stream.disconnect() + streamConversationId = '' + currentAssistantId.value = null + currentSegments.value = [] + segIdCounter.value = 0 + streamPhase.value = 'idle' + phaseInfo.value = null + error.value = null + messageQueue.clear() + if (stopFallbackTimer) { + clearTimeout(stopFallbackTimer) + stopFallbackTimer = null + } + } + return { messages, isGenerating, @@ -1128,6 +1161,7 @@ export function useChat(options: UseChatOptions): UseChatReturn { addMessage, clearMessages, reconnectStream, + resetForNewConversation, } } diff --git a/mateclaw-ui/src/composables/chat/useMessages.ts b/mateclaw-ui/src/composables/chat/useMessages.ts index 317ed462..147bb452 100644 --- a/mateclaw-ui/src/composables/chat/useMessages.ts +++ b/mateclaw-ui/src/composables/chat/useMessages.ts @@ -46,9 +46,9 @@ export interface UseMessagesReturn { /** 获取消息 */ getMessage: (id: string | number) => Message | undefined /** 创建用户消息 */ - createUserMessage: (content: string, contentParts?: MessageContentPart[]) => Message + createUserMessage: (content: string, contentParts?: MessageContentPart[], conversationId?: string) => Message /** 创建助手消息 */ - createAssistantMessage: (content?: string) => Message + createAssistantMessage: (content?: string, conversationId?: string) => Message /** 在消息列表头部插入更早的消息(分页加载) */ prependMessages: (olderMessages: Message[]) => void /** 设置 hasMore 状态 */ @@ -191,14 +191,14 @@ export function useMessages(options: UseMessagesOptions = {}): UseMessagesReturn } // 创建用户消息 - const createUserMessage = (content: string, contentParts?: MessageContentPart[]): Message => { + const createUserMessage = (content: string, contentParts?: MessageContentPart[], conversationId?: string): Message => { const parts: MessageContentPart[] = contentParts || [ { type: 'text', text: content }, ] return addMessage({ role: 'user', - conversationId: '', // 由调用方设置 + conversationId: conversationId || '', content, contentParts: parts, status: 'completed', @@ -206,10 +206,10 @@ export function useMessages(options: UseMessagesOptions = {}): UseMessagesReturn } // 创建助手消息 - const createAssistantMessage = (content: string = ''): Message => { + const createAssistantMessage = (content: string = '', conversationId?: string): Message => { return addMessage({ role: 'assistant', - conversationId: '', // 由调用方设置 + conversationId: conversationId || '', content, contentParts: content ? [{ type: 'text', text: content, visibleLength: 0 }] : [], status: 'generating', diff --git a/mateclaw-ui/src/utils/messageReconcile.ts b/mateclaw-ui/src/utils/messageReconcile.ts index 1536a9a6..a9de68a9 100644 --- a/mateclaw-ui/src/utils/messageReconcile.ts +++ b/mateclaw-ui/src/utils/messageReconcile.ts @@ -187,19 +187,17 @@ export function reconcileMessages(local: Message[], fetched: Message[]): Message } } - // 保留 fetched 中不存在的本地 assistant 消息(防止 lagging snapshot 丢弃刚完成的消息) - // 推断当前对话 ID:取 fetched 中第一条消息的 conversationId + // 保留 fetched 中不存在的本地消息(user + assistant),防止 lagging snapshot 丢弃刚发送的消息 const fetchedConversationId = fetched.length > 0 ? (fetched[0] as any).conversationId : '' for (const lm of local) { const lid = String(lm.id) - if (!matchedLocalIds.has(lid) && lm.role === 'assistant') { + if (!matchedLocalIds.has(lid)) { // 跳过不属于当前对话的本地消息,防止跨对话污染 - // 无 conversationId 的 orphan 消息也不保留 const lmConvId = (lm as any).conversationId if (!lmConvId || (fetchedConversationId && lmConvId !== fetchedConversationId)) { continue } - // 检查是否是 fetched 末尾之后的消息(刚完成,DB 还没返回) + // 只保留在 fetched 末尾之后的消息(刚发送/刚完成,DB 还没返回) const lastFetchedTime = result.length > 0 ? result[result.length - 1].createTime : '' if (!lastFetchedTime || (lm.createTime && lm.createTime >= lastFetchedTime)) { result.push(lm) diff --git a/mateclaw-ui/src/views/ChatConsole.vue b/mateclaw-ui/src/views/ChatConsole.vue index f4dae3c7..0efc006f 100644 --- a/mateclaw-ui/src/views/ChatConsole.vue +++ b/mateclaw-ui/src/views/ChatConsole.vue @@ -563,6 +563,7 @@ const { stopGeneration: stopChatGeneration, cancelQueued, reconnectStream: reconnectChatStream, + resetForNewConversation, } = useChat({ baseUrl: '', onStreamEnd: async (meta) => { @@ -770,18 +771,19 @@ async function loadConversations() { async function refreshCurrentConversationMessages(conversationId: string) { if (!conversationId) return - // 如果已经在生成新消息(用户在 stop 后又快速发了新消息),不覆盖本地状态 if (isGenerating.value) return - // 审批挂起时本地状态比 DB 更丰富(含 thinking + text),不替换 if (streamPhase.value === 'awaiting_approval') return try { const res: any = await conversationApi.listMessages(conversationId) + // Stale guard:await 返回后确认仍是当前会话 + if (currentConversationId.value !== conversationId) return + // 二次 isGenerating 检查:如果 await 期间用户已发新消息,不覆盖本地状态 + if (isGenerating.value) return const fetched = extractMessages(res).messages.map((msg: Message) => normalizeMessage(msg)) - // 过滤掉不属于当前对话的本地消息,防止跨对话污染 + // 严格过滤:只保留 conversationId 完全匹配的本地消息,orphan(空 conversationId)直接丢弃 const currentMessages = messages.value.filter( - (m: any) => !m.conversationId || m.conversationId === conversationId + (m: any) => m.conversationId === conversationId ) - // 逐条 reconcile:只接受更丰富的版本,防止 poorer DB 快照覆盖 local rich message messages.value = reconcileMessages(currentMessages, fetched) } catch (e) { console.warn('[ChatConsole] Failed to refresh current conversation messages:', e) @@ -806,13 +808,15 @@ async function hydrateStateFromRoute() { messages.value = [] try { const res: any = await conversationApi.listMessages(conversationId) + if (currentConversationId.value !== conversationId) return messages.value = extractMessages(res).messages.map((msg: Message) => normalizeMessage(msg)) } catch { // 消息加载失败,保持空 } try { + if (currentConversationId.value !== conversationId) return const statusRes: any = await conversationApi.getStatus(conversationId) - if (statusRes.data?.streamStatus === 'running') { + if (currentConversationId.value === conversationId && statusRes.data?.streamStatus === 'running') { await reconnectStream(conversationId) } } catch { @@ -839,16 +843,19 @@ async function selectConversation(conv: Conversation) { resetStreamingState() currentConversationId.value = conv.conversationId selectedAgentId.value = conv.agentId || selectedAgentId.value + const requestedConvId = conv.conversationId try { - const res: any = await conversationApi.listMessages(conv.conversationId) + const res: any = await conversationApi.listMessages(requestedConvId) + // Stale guard:await 返回后确认仍是当前会话,否则丢弃 + if (currentConversationId.value !== requestedConvId) return messages.value = extractMessages(res).messages.map((msg: Message) => normalizeMessage(msg)) // Hydrate pending approvals:恢复刷新后丢失的审批卡片 try { - const approvalRes: any = await chatApi.getPendingApprovals(conv.conversationId) + const approvalRes: any = await chatApi.getPendingApprovals(requestedConvId) + if (currentConversationId.value !== requestedConvId) return const pendingApprovals = approvalRes.data || [] if (pendingApprovals.length > 0) { - // 将 pending approvals 绑定到最近的 assistant 消息 const assistantMessages = messages.value.filter(m => m.role === 'assistant') const lastAssistant = assistantMessages[assistantMessages.length - 1] if (lastAssistant) { @@ -862,7 +869,6 @@ async function selectConversation(conv: Conversation) { arguments: pa.toolArguments, reason: pa.reason, status: 'pending_approval', - // 增强字段(Phase 6: 结构化风险信息) findings: pa.findingsJson ? JSON.parse(pa.findingsJson) : undefined, maxSeverity: pa.maxSeverity || undefined, summary: pa.summary || undefined, @@ -875,8 +881,8 @@ async function selectConversation(conv: Conversation) { // hydration 失败不影响正常使用 } - if (conv.streamStatus === 'running') { - await reconnectStream(conv.conversationId) + if (currentConversationId.value === requestedConvId && conv.streamStatus === 'running') { + await reconnectStream(requestedConvId) } } catch (e) { ElMessage.error(t('chat.loadMessagesFailed')) @@ -1126,7 +1132,9 @@ function handleCancelQueued() { // 简化版重置函数 function resetStreamingState() { + // 先通知后端停止旧流(fire-and-forget),再彻底清理前端状态 stopChatGeneration() + resetForNewConversation() } // ============ 附件处理 ============ diff --git a/text.txt b/text.txt new file mode 100644 index 00000000..84398fb8 --- /dev/null +++ b/text.txt @@ -0,0 +1,77 @@ +请基于以下已确认事实,分析并修复前端 Chat 会话串线问题。不要泛泛而谈,直接围绕时序、状态隔离、消息归属和可验证修复方案展开。 + +问题背景: + +用户连续两次问了同一句话“你有记忆里有啥”,系统创建了两个不同的 conversationId,并且后端日志显示这两个会话都是独立、正常完成的: + +1. 第一次会话: +- conversationId: conv_1775859743291_z9pav7 +- SSE chat 建立时间:2026-04-11 06:22:32 +- user message 已落库 +- assistant message 已落库 +- done 已发送 +- stream fully completed + +2. 第二次会话: +- conversationId: conv_1775859773188_335bjk +- SSE chat 建立时间:2026-04-11 06:22:55 +- user message 已落库 +- assistant message 已落库 +- done 已发送 +- stream fully completed + +关键信号: + +- 两次请求是两个不同 conversationId。 +- 服务端日志没有显示 approval / awaiting_approval / interrupt / queued_input 相关链路。 +- 第一个会话已经完成后,前端仍然发了一次 stop,请求日志为 `stopped=false`,这说明 stop 到达时旧流已经结束,不是服务端还在跑旧流。 +- 因此,这更像是前端本地状态污染、会话切换时序竞争、或 reconcile 逻辑把旧本地消息错误带入新会话,而不是后端把旧会话内容串到了新会话。 + +当前高优先级怀疑点: + +1. 切换会话 / 新建会话时,只调用了 stopChatGeneration(),但没有等待旧 SSE 流和本地状态完全清理。 +- 这会导致旧流晚到的事件(delta / done / error)在新会话已经创建 assistant 占位消息后,继续命中新会话的共享状态。 +- useChat 内部当前使用共享的 `currentAssistantId`、`streamConversationId`、`messages`,如果不做 conversation 级别隔离,就有天然串线风险。 + +2. 审批占位 assistant message 在某些路径下创建后没有 conversationId。 +- 当前 refresh / reconcile 的过滤逻辑对 `!conversationId` 的本地消息仍可能放行。 +- 这类 orphan message 可能被错误并入后续任意会话。 + +3. reconcileMessages() 当前有“保留 fetched 中不存在的本地 assistant 消息”的策略。 +- 这个策略本来是为了防止 lagging snapshot 丢刚完成的 rich message。 +- 但如果 local 里混入了旧会话消息、orphan message、或者未彻底清理的占位消息,就会把错误消息保留下来。 + +你的任务: + +1. 先明确判断: +- 根因是否主要在前端状态管理,而非后端 conversation/message 落库。 +- 哪一条最可能导致“上一轮消息出现在新会话”。 + +2. 给出修复方案,要求具体到代码层面: +- 会话切换 / 新建会话时,如何确保旧流彻底解绑。 +- 如何避免旧流事件写入当前会话。 +- `currentAssistantId` / `streamConversationId` 是否应该按 conversation 隔离,还是至少在事件处理时校验 conversationId。 +- 所有本地新建 message 是否必须强制携带 conversationId。 +- reconcileMessages() 是否应该完全禁止保留非当前 conversation 的本地消息。 +- 对 `conversationId` 为空的本地消息,应该如何处理。 + +3. 给出建议的防御性约束: +- 每个 SSE 事件落地前必须校验所属 conversationId。 +- onStreamEnd / reconnect / refreshCurrentConversationMessages 只能作用于当前会话。 +- 新会话开始前,旧会话的 placeholder / generating message 必须被清理或隔离。 + +4. 输出格式要求: +- 先给“根因判断”。 +- 再给“最小修复方案”。 +- 再给“更稳妥的长期方案”。 +- 最后给“如何验证修复有效”,至少覆盖: + - 连续快速新建会话并发送相同问题 + - 旧会话刚结束时立刻切新会话 + - refreshCurrentConversationMessages 在流结束后执行 + - orphan assistant message / 空 conversationId message 不得污染新会话 + +补充要求: + +- 不要只说“加锁”或“避免 race condition”,要明确到状态变量、事件处理器、过滤条件和消息生命周期。 +- 如果你认为某个现有修补不够,请直接指出为什么不够。 +- 如果需要改 reconcileMessages,请说明保留本地 assistant message 的边界条件。