sse_core.py 9.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206
  1. # -*- coding: utf-8 -*-
  2. """
  3. SSE 核心工具函数 —— 流式事件生成器、文本提取、响应清理等通用辅助
  4. """
  5. import json
  6. import re
  7. import time
  8. import traceback
  9. from typing import AsyncGenerator
  10. from tools.tool_names_cn import TOOL_NAME_CN as _TOOL_CN
  11. from db.chat_store import save_message, get_messages
  12. # ── 工具名 → 中文描述映射 ──
  13. def _cn_tool_desc(tool_name: str) -> str:
  14. """将工具函数名映射为中文描述短语(用于 SSE message 字段)。
  15. 映射表统一维护在 tools/tool_names_cn.py 中,新增工具只需改那一处。
  16. """
  17. return _TOOL_CN.get(tool_name, tool_name)
  18. # ── 通用 SSE 流式事件生成器 ──
  19. async def sse_event_generator(
  20. agent,
  21. user_message: str,
  22. thread_id: str,
  23. session_id: str,
  24. config: dict,
  25. agent_cn_name: str = "智能助手",
  26. save_to_db: bool = True,
  27. ) -> AsyncGenerator[str, None]:
  28. """
  29. 通用 SSE 流式事件生成器。
  30. 使用 stream_mode=["updates", "messages"] 获取完整执行图景,
  31. 前端可根据 event type 区分:thinking / executing / tool_call / token / done / error。
  32. Args:
  33. save_to_db: 是否将消息保存到数据库。点选解读等非对话场景应设为 False。
  34. """
  35. # 保存用户消息(仅对话场景写库)
  36. if save_to_db:
  37. save_message(session_id, "user", user_message)
  38. # 用于收集完整回复
  39. full_response = ""
  40. start_time = time.time()
  41. try:
  42. # ── 发送 agent_start 事件,前端据此创建 Agent 卡片 ──
  43. yield f"data: {json.dumps({'type': 'agent_start', 'agent': agent_cn_name, 'cn_agent': agent_cn_name, 'message': f'「{agent_cn_name}」开始处理...'}, ensure_ascii=False)}\n\n"
  44. # 构建消息列表(包含历史消息 + 当前消息)
  45. if save_to_db:
  46. messages = [{"role": "user", "content": msg["content"]}
  47. for msg in get_messages(session_id, limit=50)]
  48. else:
  49. messages = []
  50. messages.append({"role": "user", "content": user_message})
  51. # 使用 agent.astream() 异步流式调用,多模式获取完整图景
  52. async for chunk in agent.astream(
  53. {"messages": messages},
  54. stream_mode=["updates", "messages"],
  55. config=config,
  56. version="v2",
  57. ):
  58. # 判断事件来源:主代理 vs 子代理
  59. is_subagent = any(s.startswith("tools:") for s in chunk.get("ns", []))
  60. source = "subagent" if is_subagent else "main"
  61. # ── updates 模式:步骤级事件(thinking / executing)──
  62. if chunk["type"] == "updates":
  63. for node_name in chunk["data"]:
  64. if node_name == "model_request":
  65. # 代理正在思考/推理
  66. yield f"data: {json.dumps({'type': 'thinking', 'source': source, 'node': node_name, 'cn_agent': agent_cn_name, 'message': '正在分析您的问题...'}, ensure_ascii=False)}\n\n"
  67. elif node_name == "tools":
  68. # 代理正在执行工具调用,提取工具名称
  69. tools_data = chunk["data"].get(node_name, {})
  70. tool_names = []
  71. for msg in tools_data.get("messages", []):
  72. if hasattr(msg, "name"):
  73. tool_names.append(msg.name)
  74. elif isinstance(msg, dict) and msg.get("name"):
  75. tool_names.append(msg["name"])
  76. # 生成中文描述
  77. cn_names = [_cn_tool_desc(t) for t in tool_names] if tool_names else ["工具调用"]
  78. msg_text = f"正在{'、'.join(cn_names)}..."
  79. yield f"data: {json.dumps({'type': 'executing', 'source': source, 'cn_agent': agent_cn_name, 'tools': tool_names, 'cn_tools': cn_names, 'message': msg_text}, ensure_ascii=False)}\n\n"
  80. # ── messages 模式:token 级事件(token / tool_call / tool_result / updated_todo_list)──
  81. elif chunk["type"] == "messages":
  82. token_data = chunk["data"]
  83. if isinstance(token_data, (list, tuple)) and len(token_data) >= 1:
  84. msg_obj, _metadata = token_data[0], token_data[1] if len(token_data) > 1 else {}
  85. else:
  86. msg_obj = token_data
  87. # 检测工具调用(tool_call_chunks 在流式传输中逐步到达)
  88. if hasattr(msg_obj, "tool_call_chunks") and msg_obj.tool_call_chunks:
  89. for tc in msg_obj.tool_call_chunks:
  90. if tc.get("name"):
  91. cn = _cn_tool_desc(tc["name"])
  92. yield f"data: {json.dumps({'type': 'tool_call', 'source': source, 'cn_agent': agent_cn_name, 'tool': tc['name'], 'cn_tool': cn, 'message': f'调用工具:{cn}'}, ensure_ascii=False)}\n\n"
  93. # 检测工具结果
  94. is_tool_msg = hasattr(msg_obj, "type") and msg_obj.type == "tool"
  95. if is_tool_msg:
  96. tool_name = getattr(msg_obj, "name", "unknown")
  97. if tool_name == "write_todos":
  98. # write_todos 特殊处理:提取 JSON 并发送 updated_todo_list 事件
  99. tool_content = _extract_text_content(msg_obj) or ""
  100. todo_list = _parse_todo_list(tool_content)
  101. if todo_list:
  102. yield f"data: {json.dumps({'type': 'updated_todo_list', 'cn_agent': agent_cn_name, 'source': source, 'todos': todo_list, 'message': '任务进度已更新'}, ensure_ascii=False)}\n\n"
  103. else:
  104. # 解析失败时也至少发一个 tool_result 事件,并打印诊断日志
  105. print(f"[write_todos] 解析失败,原始内容: {tool_content[:300]}")
  106. yield f"data: {json.dumps({'type': 'tool_result', 'source': source, 'cn_agent': agent_cn_name, 'tool': tool_name, 'cn_tool': _cn_tool_desc(tool_name), 'message': '任务进度已更新'}, ensure_ascii=False)}\n\n"
  107. else:
  108. cn = _cn_tool_desc(tool_name)
  109. yield f"data: {json.dumps({'type': 'tool_result', 'source': source, 'cn_agent': agent_cn_name, 'tool': tool_name, 'cn_tool': cn, 'message': f'{cn} 完成'}, ensure_ascii=False)}\n\n"
  110. # 提取并流式输出文本内容(工具结果不当作 token 输出,仅输出 AI 生成的文本)
  111. if not is_tool_msg:
  112. content = _extract_text_content(msg_obj)
  113. if content and isinstance(content, str):
  114. full_response += content
  115. yield f"data: {json.dumps({'type': 'token', 'source': source, 'content': content}, ensure_ascii=False)}\n\n"
  116. # 保存助手回复(仅对话场景写库)
  117. if save_to_db and full_response.strip():
  118. cleaned = _clean_response(full_response)
  119. save_message(session_id, "assistant", cleaned)
  120. print(full_response)
  121. # 发送 agent_done + done 事件
  122. duration_ms = int((time.time() - start_time) * 1000)
  123. yield f"data: {json.dumps({'type': 'agent_done', 'agent': agent_cn_name, 'cn_agent': agent_cn_name, 'duration_ms': duration_ms, 'message': f'「{agent_cn_name}」完成({duration_ms}ms)'}, ensure_ascii=False)}\n\n"
  124. yield f"data: {json.dumps({'type': 'done', 'thread_id': thread_id, 'session_id': session_id, 'duration_ms': duration_ms, 'message': f'回答完成({duration_ms}ms)'}, ensure_ascii=False)}\n\n"
  125. except Exception as e:
  126. traceback.print_exc()
  127. yield f"data: {json.dumps({'type': 'error', 'message': str(e)}, ensure_ascii=False)}\n\n"
  128. # ── 文本提取 / 解析 / 清理 ──
  129. def _extract_text_content(msg_obj) -> str | None:
  130. """从 LangChain 消息对象中提取文本内容(兼容 dict 和对象两种形式)。"""
  131. if isinstance(msg_obj, dict):
  132. return msg_obj.get("content")
  133. elif hasattr(msg_obj, "content"):
  134. raw = getattr(msg_obj, "content", None)
  135. if isinstance(raw, str):
  136. return raw
  137. elif isinstance(raw, list):
  138. # 多模态内容块:合并所有 text 类型的块
  139. parts = []
  140. for block in raw:
  141. if isinstance(block, dict) and block.get("type") == "text":
  142. parts.append(block.get("text", ""))
  143. elif hasattr(block, "type") and getattr(block, "type", "") == "text":
  144. parts.append(getattr(block, "text", ""))
  145. return "".join(parts) if parts else None
  146. return None
  147. def _parse_todo_list(text: str) -> list | None:
  148. """从 write_todos 输出中提取 todo 列表。
  149. write_todos 输出格式: "Updated todo list to [{'content': '...', 'status': '...'}, ...]"
  150. 返回 JSON-serializable list of dicts,失败返回 None。
  151. """
  152. import ast
  153. # 贪婪匹配最外层 [...](内容中的 [ ] 在字符串字面量内,ast.literal_eval 可正确处理)
  154. match = re.search(r"\[.*\]", text)
  155. if not match:
  156. return None
  157. try:
  158. python_list = ast.literal_eval(match.group())
  159. if isinstance(python_list, list):
  160. return python_list
  161. except (ValueError, SyntaxError, TypeError):
  162. pass
  163. return None
  164. def _clean_response(text: str) -> str:
  165. """清理 Agent 回复中的格式噪音"""
  166. # 移除 ANSI 转义序列
  167. text = re.sub(r'\x1b\[[0-9;]*m', '', text)
  168. text = re.sub(r'\x1b\[[0-9;]*[a-zA-Z]', '', text)
  169. text = re.sub(r'\x1b\][^\x07]*\x07', '', text)
  170. # 移除工具调用残留(大括号 JSON 块如果独立成行则移除)
  171. text = re.sub(r'^\s*\{[^}]*\}\s*$', '', text, flags=re.MULTILINE)
  172. return text.strip()