sse_core.py 22 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488
  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 langgraph.types import Command # 程序化 interrupt() 恢复时需启用
  11. from tools.tool_names_cn import TOOL_NAME_CN as _TOOL_CN
  12. from db.chat_store import save_message, get_messages
  13. from tools.context_tracker import capture_context_usage
  14. # ── 工具名 → 中文描述映射 ──
  15. def _cn_tool_desc(tool_name: str) -> str:
  16. """将工具函数名映射为中文描述短语(用于 SSE message 字段)。
  17. 映射表统一维护在 tools/tool_names_cn.py 中,新增工具只需改那一处。
  18. """
  19. return _TOOL_CN.get(tool_name, tool_name)
  20. # ── 通用 SSE 流式事件生成器 ──
  21. async def sse_event_generator(
  22. agent,
  23. user_message: str,
  24. thread_id: str,
  25. session_id: str,
  26. config: dict,
  27. agent_cn_name: str = "智能助手",
  28. save_to_db: bool = True,
  29. original_user_message: str | None = None,
  30. interrupt_before: list | None = None,
  31. interrupt_after: list | None = None,
  32. ) -> AsyncGenerator[str, None]:
  33. """
  34. 通用 SSE 流式事件生成器。
  35. 使用 stream_mode=["updates", "messages"] 获取完整执行图景,
  36. 前端可根据 event type 区分:thinking / executing / tool_call / token / done / error。
  37. Args:
  38. save_to_db: 是否将消息保存到数据库。点选解读等非对话场景应设为 False。
  39. original_user_message: 用户的原始消息(不含系统注入的前缀)。
  40. 若提供,则 DB 保存此原始消息;user_message 仍作为 Agent 的输入。
  41. """
  42. # 保存用户消息(仅对话场景写库),优先使用原始消息
  43. if save_to_db:
  44. save_message(session_id, "user", original_user_message or user_message)
  45. # 用于收集完整回复
  46. full_response = ""
  47. start_time = time.time()
  48. _token_usage = {"prompt": 0, "completion": 0} # 从流中捕获的实际 token 用量
  49. try:
  50. # ── 发送 agent_start 事件,前端据此创建 Agent 卡片 ──
  51. 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"
  52. # 构建消息列表(包含历史消息 + 当前消息)
  53. if save_to_db:
  54. messages = [{"role": "user", "content": msg["content"]}
  55. for msg in get_messages(session_id, limit=50)]
  56. else:
  57. messages = []
  58. messages.append({"role": "user", "content": user_message})
  59. # 使用 agent.astream() 异步流式调用,多模式获取完整图景
  60. async for chunk in agent.astream(
  61. {"messages": messages},
  62. stream_mode=["updates", "messages"],
  63. config=config,
  64. interrupt_before=interrupt_before,
  65. interrupt_after=interrupt_after,
  66. version="v2",
  67. ):
  68. # 判断事件来源:主代理 vs 子代理
  69. is_subagent = any(s.startswith("tools:") for s in chunk.get("ns", []))
  70. source = "subagent" if is_subagent else "main"
  71. # ── updates 模式:步骤级事件(thinking / executing)──
  72. if chunk["type"] == "updates":
  73. for node_name in chunk["data"]:
  74. if node_name == "model":
  75. # 代理正在思考/推理
  76. yield f"data: {json.dumps({'type': 'thinking', 'source': source, 'node': node_name, 'cn_agent': agent_cn_name, 'message': '正在分析您的问题...'}, ensure_ascii=False)}\n\n"
  77. elif node_name == "tools":
  78. # 代理正在执行工具调用,提取工具名称
  79. tools_data = chunk["data"].get(node_name, {})
  80. tool_names = []
  81. for msg in tools_data.get("messages", []):
  82. if hasattr(msg, "name"):
  83. tool_names.append(msg.name)
  84. elif isinstance(msg, dict) and msg.get("name"):
  85. tool_names.append(msg["name"])
  86. # 生成中文描述
  87. cn_names = [_cn_tool_desc(t) for t in tool_names] if tool_names else ["工具调用"]
  88. msg_text = f"正在{'、'.join(cn_names)}..."
  89. 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"
  90. # ── messages 模式:token 级事件(token / tool_call / tool_result / updated_todo_list)──
  91. elif chunk["type"] == "messages":
  92. token_data = chunk["data"]
  93. if isinstance(token_data, (list, tuple)) and len(token_data) >= 1:
  94. msg_obj, _metadata = token_data[0], token_data[1] if len(token_data) > 1 else {}
  95. else:
  96. msg_obj = token_data
  97. # 检测工具调用(tool_call_chunks 在流式传输中逐步到达)
  98. if hasattr(msg_obj, "tool_call_chunks") and msg_obj.tool_call_chunks:
  99. for tc in msg_obj.tool_call_chunks:
  100. if tc.get("name"):
  101. cn = _cn_tool_desc(tc["name"])
  102. 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"
  103. # 检测工具结果
  104. is_tool_msg = hasattr(msg_obj, "type") and msg_obj.type == "tool"
  105. if is_tool_msg:
  106. tool_name = getattr(msg_obj, "name", "unknown")
  107. if tool_name == "write_todos":
  108. # write_todos 特殊处理:提取 JSON 并发送 updated_todo_list 事件
  109. tool_content = _extract_text_content(msg_obj) or ""
  110. todo_list = _parse_todo_list(tool_content)
  111. if todo_list:
  112. 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"
  113. else:
  114. # 解析失败时也至少发一个 tool_result 事件,并打印诊断日志
  115. print(f"[write_todos] 解析失败,原始内容: {tool_content[:300]}")
  116. 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"
  117. else:
  118. cn = _cn_tool_desc(tool_name)
  119. 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"
  120. # 提取并流式输出文本内容(工具结果不当作 token 输出,仅输出 AI 生成的文本)
  121. if not is_tool_msg:
  122. # 先检查推理/思考内容(DeepSeek 等模型的 reasoning_content)
  123. reasoning = _extract_reasoning_content(msg_obj)
  124. if reasoning:
  125. yield f"data: {json.dumps({'type': 'thinking_token', 'source': source, 'content': reasoning}, ensure_ascii=False)}\n\n"
  126. # 再提取常规文本内容
  127. content = _extract_text_content(msg_obj)
  128. if content and isinstance(content, str):
  129. full_response += content
  130. yield f"data: {json.dumps({'type': 'token', 'source': source, 'content': content}, ensure_ascii=False)}\n\n"
  131. # ── 提取 token 用量元数据(最后一个消息块通常携带 usage)──
  132. _capture_token_meta(msg_obj, _token_usage)
  133. # 保存助手回复(仅对话场景写库)
  134. if save_to_db and full_response.strip():
  135. cleaned = _clean_response(full_response)
  136. save_message(session_id, "assistant", cleaned)
  137. print(full_response)
  138. # ── 更新上下文用量 ──
  139. if save_to_db and _token_usage["prompt"] > 0:
  140. await capture_context_usage(
  141. session_id=session_id,
  142. prompt_tokens=_token_usage["prompt"],
  143. completion_tokens=_token_usage["completion"],
  144. current_message=user_message,
  145. )
  146. # 发送 agent_done + done 事件
  147. duration_ms = int((time.time() - start_time) * 1000)
  148. 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"
  149. # 检查 LangGraph 中断状态(Human-in-the-Loop)
  150. interrupt_info = _check_interrupt(agent, config)
  151. if interrupt_info:
  152. yield f"data: {json.dumps({'type': 'interrupt', **interrupt_info}, ensure_ascii=False)}\n\n"
  153. # 中断时不发送 done 事件,等待用户审批后恢复
  154. return
  155. 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"
  156. except Exception as e:
  157. traceback.print_exc()
  158. yield f"data: {json.dumps({'type': 'error', 'message': str(e)}, ensure_ascii=False)}\n\n"
  159. # ── 文本提取 / 解析 / 清理 ──
  160. def _extract_text_content(msg_obj) -> str | None:
  161. """从 LangChain 消息对象中提取文本内容(兼容 dict 和对象两种形式)。"""
  162. if isinstance(msg_obj, dict):
  163. return msg_obj.get("content")
  164. elif hasattr(msg_obj, "content"):
  165. raw = getattr(msg_obj, "content", None)
  166. if isinstance(raw, str):
  167. return raw
  168. elif isinstance(raw, list):
  169. # 多模态内容块:合并所有 text 类型的块
  170. parts = []
  171. for block in raw:
  172. if isinstance(block, dict) and block.get("type") == "text":
  173. parts.append(block.get("text", ""))
  174. elif hasattr(block, "type") and getattr(block, "type", "") == "text":
  175. parts.append(getattr(block, "text", ""))
  176. return "".join(parts) if parts else None
  177. return None
  178. # ── 调试开关:设为 True 时打印推理内容摘要(排查完毕后关闭)──
  179. _DEBUG_REASONING = False
  180. def _extract_reasoning_content(msg_obj) -> str | None:
  181. """从 LangChain 消息对象中提取推理/思考内容(DeepSeek 等模型的 reasoning_content)。
  182. 支持三种来源(按优先级):
  183. 1. msg_obj.additional_kwargs["reasoning_content"](OpenAI 兼容流式 delta)
  184. 2. msg_obj.reasoning_content(LangChain 直接属性)
  185. 3. content 中的 thinking 类型块
  186. Returns:
  187. 推理文本字符串,无推理内容时返回 None
  188. """
  189. # 方式 1:additional_kwargs 中的 reasoning_content(最常见)
  190. if hasattr(msg_obj, "additional_kwargs") and isinstance(msg_obj.additional_kwargs, dict):
  191. reasoning = msg_obj.additional_kwargs.get("reasoning_content", "")
  192. if reasoning:
  193. return reasoning
  194. # 方式 2:直接属性 reasoning_content
  195. reasoning = getattr(msg_obj, "reasoning_content", None)
  196. if reasoning:
  197. return reasoning
  198. # 方式 3:content 为 list 时,提取 thinking 类型的块
  199. if hasattr(msg_obj, "content"):
  200. raw = getattr(msg_obj, "content", None)
  201. if isinstance(raw, list):
  202. parts = []
  203. for block in raw:
  204. if isinstance(block, dict) and block.get("type") == "thinking":
  205. parts.append(block.get("thinking", ""))
  206. elif hasattr(block, "type") and getattr(block, "type", "") == "thinking":
  207. parts.append(getattr(block, "thinking", ""))
  208. if parts:
  209. return "".join(parts)
  210. return None
  211. def _capture_token_meta(msg_obj, usage_ref: dict):
  212. """从 LangChain 消息对象中提取 LLM token 用量元数据。
  213. 优先读取 usage_metadata(langchain ≥0.3),
  214. 其次读取 response_metadata.usage(OpenAI 兼容格式)。
  215. 结果写入 usage_ref dict(原地修改)。
  216. """
  217. # 方式 1:usage_metadata(langchain 标准字段)
  218. um = getattr(msg_obj, "usage_metadata", None)
  219. if um and isinstance(um, dict):
  220. inp = um.get("input_tokens", 0)
  221. out = um.get("output_tokens", 0)
  222. if inp or out:
  223. usage_ref["prompt"] = inp
  224. usage_ref["completion"] = out
  225. return
  226. # 方式 2:response_metadata.usage(OpenAI / DeepSeek 兼容)
  227. rm = getattr(msg_obj, "response_metadata", None)
  228. if rm and isinstance(rm, dict):
  229. usage = rm.get("usage", {}) or rm.get("token_usage", {})
  230. if isinstance(usage, dict):
  231. inp = usage.get("prompt_tokens", 0)
  232. out = usage.get("completion_tokens", 0)
  233. if inp or out:
  234. usage_ref["prompt"] = inp
  235. usage_ref["completion"] = out
  236. return
  237. def _parse_todo_list(text: str) -> list | None:
  238. """从 write_todos 输出中提取 todo 列表。
  239. write_todos 输出格式: "Updated todo list to [{'content': '...', 'status': '...'}, ...]"
  240. 返回 JSON-serializable list of dicts,失败返回 None。
  241. """
  242. import ast
  243. # 贪婪匹配最外层 [...](内容中的 [ ] 在字符串字面量内,ast.literal_eval 可正确处理)
  244. match = re.search(r"\[.*\]", text)
  245. if not match:
  246. return None
  247. try:
  248. python_list = ast.literal_eval(match.group())
  249. if isinstance(python_list, list):
  250. return python_list
  251. except (ValueError, SyntaxError, TypeError):
  252. pass
  253. return None
  254. def _clean_response(text: str) -> str:
  255. """清理 Agent 回复中的格式噪音"""
  256. # 移除 ANSI 转义序列
  257. text = re.sub(r'\x1b\[[0-9;]*m', '', text)
  258. text = re.sub(r'\x1b\[[0-9;]*[a-zA-Z]', '', text)
  259. text = re.sub(r'\x1b\][^\x07]*\x07', '', text)
  260. # 移除工具调用残留(大括号 JSON 块如果独立成行则移除)
  261. text = re.sub(r'^\s*\{[^}]*\}\s*$', '', text, flags=re.MULTILINE)
  262. return text.strip()
  263. # ── LangGraph Human-in-the-Loop 中断处理 ──
  264. def _check_interrupt(agent, config: dict) -> dict | None:
  265. """
  266. 检查 LangGraph 状态是否被中断。
  267. 支持两种中断检测:
  268. 1. 程序化中断(节点内调用 interrupt())→ state.interrupts 非空
  269. 2. interrupt_before / interrupt_after 暂停 → state.next 非空
  270. 当中断发生时返回中断信息 dict,否则返回 None。
  271. 返回格式:
  272. {"node": "tools", "message": "智能体准备执行工具,等待审批...",
  273. "interrupts": [...], "plan": "计划文本(程序化中断时)"}
  274. """
  275. try:
  276. state = agent.get_state(config)
  277. except Exception:
  278. return None
  279. if state is None:
  280. return None
  281. # 1. 检测程序化中断(节点内调用 interrupt())
  282. interrupts = getattr(state, "interrupts", None)
  283. if not interrupts:
  284. values = getattr(state, "values", {}) or {}
  285. interrupts = values.get("__interrupt__", [])
  286. # 2. 检测 interrupt_before / interrupt_after 暂停
  287. # 图暂停时 state.next 非空(有待执行节点);完成时为空
  288. next_nodes = getattr(state, "next", None)
  289. is_paused_by_interrupt_config = bool(next_nodes) if next_nodes is not None else False
  290. if not interrupts and not is_paused_by_interrupt_config:
  291. return None
  292. # 提取中断节点名
  293. if interrupts:
  294. interrupt_data = interrupts[0] if interrupts else {}
  295. # 处理 interrupt() 返回值可能是 Interrupt 对象或 dict
  296. if hasattr(interrupt_data, "value"):
  297. interrupt_data = interrupt_data.value
  298. if isinstance(interrupt_data, dict):
  299. node = interrupt_data.get("type", "tools")
  300. plan = interrupt_data.get("plan", "")
  301. message = interrupt_data.get("message", "智能体暂停执行,等待您的审批...")
  302. else:
  303. node = str(next_nodes or "tools")
  304. plan = ""
  305. message = "智能体暂停执行,等待您的审批..."
  306. else:
  307. node = next_nodes
  308. plan = ""
  309. message = "智能体暂停执行,等待您的审批..."
  310. if isinstance(node, (list, tuple)):
  311. node = node[0] if node else "unknown"
  312. result = {
  313. "node": str(node),
  314. "message": message,
  315. "interrupts": [str(i) for i in interrupts] if interrupts else [],
  316. }
  317. if plan:
  318. result["plan"] = plan
  319. return result
  320. async def resume_stream(
  321. agent,
  322. session_id: str,
  323. thread_id: str,
  324. config: dict,
  325. agent_cn_name: str = "智能助手",
  326. action: str = "approve",
  327. interrupt_before: list | None = None,
  328. interrupt_after: list | None = None,
  329. ):
  330. """
  331. 恢复被中断的 LangGraph 流式执行。
  332. 参数:
  333. action: "approve" 或 "reject"
  334. Yields:
  335. SSE 事件字符串
  336. """
  337. if action == "reject":
  338. yield f"data: {json.dumps({'type': 'error', 'message': '用户拒绝了工具执行'}, ensure_ascii=False)}\n\n"
  339. yield f"data: {json.dumps({'type': 'done', 'thread_id': thread_id, 'session_id': session_id, 'message': '已取消执行'}, ensure_ascii=False)}\n\n"
  340. return
  341. full_response = ""
  342. start_time = time.time()
  343. try:
  344. # 检测中断类型:程序化中断需用 Command(resume=...),配置式中断传 None
  345. state = agent.get_state(config)
  346. has_programmatic = bool(getattr(state, "interrupts", None)) if state else False
  347. stream_input = Command(resume={"action": action}) if has_programmatic else None
  348. # 恢复执行
  349. async for chunk in agent.astream(
  350. stream_input,
  351. stream_mode=["updates", "messages"],
  352. config=config,
  353. interrupt_before=interrupt_before,
  354. interrupt_after=interrupt_after,
  355. version="v2",
  356. ):
  357. is_subagent = any(s.startswith("tools:") for s in chunk.get("ns", []))
  358. if chunk["type"] == "updates":
  359. for node_name in chunk["data"]:
  360. if node_name == "model":
  361. yield f"data: {json.dumps({'type': 'thinking', 'source': 'main', 'cn_agent': agent_cn_name, 'message': '继续执行...'}, ensure_ascii=False)}\n\n"
  362. elif node_name == "tools":
  363. tools_data = chunk["data"].get(node_name, {})
  364. tool_names = []
  365. for msg in tools_data.get("messages", []):
  366. if hasattr(msg, "name"):
  367. tool_names.append(msg.name)
  368. elif isinstance(msg, dict) and msg.get("name"):
  369. tool_names.append(msg["name"])
  370. cn_names = [_cn_tool_desc(t) for t in tool_names] if tool_names else ["工具调用"]
  371. yield f"data: {json.dumps({'type': 'executing', 'source': 'main', 'cn_agent': agent_cn_name, 'tools': tool_names, 'cn_tools': cn_names, 'message': f"正在{'、'.join(cn_names)}..."}, ensure_ascii=False)}\n\n"
  372. elif chunk["type"] == "messages":
  373. token_data = chunk["data"]
  374. if isinstance(token_data, (list, tuple)) and len(token_data) >= 1:
  375. msg_obj = token_data[0]
  376. else:
  377. msg_obj = token_data
  378. # 工具结果
  379. is_tool_msg = hasattr(msg_obj, "type") and msg_obj.type == "tool"
  380. if is_tool_msg:
  381. tool_name = getattr(msg_obj, "name", "unknown")
  382. if tool_name != "write_todos":
  383. yield f"data: {json.dumps({'type': 'tool_result', 'source': 'main', 'cn_agent': agent_cn_name, 'tool': tool_name, 'cn_tool': _cn_tool_desc(tool_name), 'message': f'{_cn_tool_desc(tool_name)} 完成'}, ensure_ascii=False)}\n\n"
  384. # AI 文本
  385. if not is_tool_msg:
  386. # 先检查推理/思考内容(DeepSeek 等模型的 reasoning_content)
  387. reasoning = _extract_reasoning_content(msg_obj)
  388. if reasoning:
  389. yield f"data: {json.dumps({'type': 'thinking_token', 'source': 'main', 'content': reasoning}, ensure_ascii=False)}\n\n"
  390. # 再提取常规文本内容
  391. content = _extract_text_content(msg_obj)
  392. if content and isinstance(content, str):
  393. full_response += content
  394. yield f"data: {json.dumps({'type': 'token', 'source': 'main', 'content': content}, ensure_ascii=False)}\n\n"
  395. # 保存助手回复
  396. if full_response.strip():
  397. cleaned = _clean_response(full_response)
  398. save_message(session_id, "assistant", cleaned)
  399. # 再次检查中断(可能有多次中断)
  400. interrupt_info = _check_interrupt(agent, config)
  401. if interrupt_info:
  402. yield f"data: {json.dumps({'type': 'interrupt', **interrupt_info}, ensure_ascii=False)}\n\n"
  403. return
  404. duration_ms = int((time.time() - start_time) * 1000)
  405. 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"
  406. except Exception as e:
  407. traceback.print_exc()
  408. yield f"data: {json.dumps({'type': 'error', 'message': str(e)}, ensure_ascii=False)}\n\n"