health_butler_graph.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320
  1. from typing import TypedDict, Optional, Literal
  2. from langgraph.graph import StateGraph, START, END
  3. from langgraph.checkpoint.memory import MemorySaver
  4. from langchain_openai import ChatOpenAI
  5. from langchain_core.messages import SystemMessage, HumanMessage
  6. from app.tools.java_client import JavaClient
  7. from app.tools.product_tools import (
  8. search_product_by_keyword,
  9. search_article_by_keyword,
  10. )
  11. from app.memory.store import MemoryManager
  12. from app.config import settings
  13. from app.prompt_service import get_prompt
  14. import logging
  15. import json
  16. logger = logging.getLogger(__name__)
  17. class HealthButlerState(TypedDict):
  18. query: str
  19. user_id: int
  20. conversation_id: str
  21. family_id: int | None
  22. child_id: int | None
  23. report_id: int | None
  24. focus: str | None # gut / nutrition / overall / chronic
  25. kb_context: dict | None
  26. knowledge_results: list[dict]
  27. answer: str | None
  28. tasks: list[dict]
  29. sources: list[dict]
  30. messages: list | None
  31. DEFAULT_BUTLER_PROMPT = """你是一个家庭AI健康管家, 为用户提供专业、温暖的个性化健康建议。
  32. 你的能力:
  33. 1. 基于健康知识库(菌属功能、指标解读、营养素、食物)提供科学建议
  34. 2. 解读健康报告数据, 结合知识库给出改善方案
  35. 3. 根据菌群检测结果推荐精准营养方案
  36. 4. 提供可执行的健康微行动
  37. 回答原则:
  38. 1. 用中文, 像家庭医生一样亲切但专业
  39. 2. 引用知识库中的菌属/指标/食物信息作为证据
  40. 3. 优先使用知识库数据, 其次使用通用健康知识
  41. 4. 需要产品推荐时使用搜索工具
  42. 5. 生成 [TASK: {"title": "任务名", "dimension": "身", "points": 10, "frequency": "每天"}] 标记创建健康行动任务
  43. 6. 超出知识库范围的慢性病/急症问题, 建议咨询专业医生
  44. 7. 不要编造医学建议, 不夸大效果
  45. """
  46. def create_health_butler_graph():
  47. """创建 AI 健康管家 StateGraph — 知识检索 + LLM 建议 + 记忆"""
  48. llm = ChatOpenAI(
  49. model=settings.llm_model,
  50. api_key=settings.llm_api_key,
  51. base_url=settings.llm_base_url,
  52. temperature=settings.llm_temperature,
  53. )
  54. from app.rag.embeddings import get_embeddings
  55. embeddings = get_embeddings()
  56. memory_mgr = MemoryManager(embeddings)
  57. java = JavaClient()
  58. llm_with_tools = llm.bind_tools([
  59. search_product_by_keyword,
  60. search_article_by_keyword,
  61. ])
  62. builder = StateGraph(HealthButlerState)
  63. async def load_knowledge(state: HealthButlerState) -> dict:
  64. """阶段1: 从 Java 后端加载健康知识库上下文"""
  65. context = {}
  66. try:
  67. client = await java._get_client()
  68. resp = await client.post(
  69. "/api/health/knowledge/list",
  70. json={},
  71. )
  72. data = resp.json()
  73. if data.get("code") == 200:
  74. items = data.get("data", [])
  75. # 按 itemType 分组
  76. by_type = {}
  77. for item in items:
  78. t = item.get("itemType", "unknown")
  79. if t not in by_type:
  80. by_type[t] = []
  81. by_type[t].append(item)
  82. context["kb_items"] = len(items)
  83. context["kb_by_type"] = {k: len(v) for k, v in by_type.items()}
  84. context["kb_loaded"] = True
  85. else:
  86. context["kb_loaded"] = False
  87. except Exception as e:
  88. logger.warning("加载健康知识库失败: %s", e)
  89. context["kb_loaded"] = False
  90. return {"kb_context": context}
  91. async def search_knowledge(state: HealthButlerState) -> dict:
  92. """阶段2: 根据 query 意图精准搜索知识库"""
  93. query = state["query"]
  94. results = []
  95. # 提取 query 中可能的菌属名/指标名/食物名关键词
  96. # 策略: 对 query 做简单的关键词匹配, 调用 Java 知识库逐个查询
  97. keywords = _extract_health_keywords(query)
  98. if keywords:
  99. client = await java._get_client()
  100. for kw in keywords[:5]: # 最多查询5个关键词
  101. # 优先查 bacteria 类型
  102. for item_type in ["bacteria", "indicator", "food", "nutrient"]:
  103. try:
  104. resp = await client.post(
  105. "/api/health/knowledge/query",
  106. json={"itemType": item_type, "itemName": kw},
  107. )
  108. data = resp.json()
  109. if data.get("code") == 200 and data.get("data"):
  110. results.append({
  111. "type": item_type,
  112. "keyword": kw,
  113. "data": data["data"],
  114. })
  115. break # 命中即跳出
  116. except Exception:
  117. continue
  118. return {"knowledge_results": results}
  119. async def llm_answer(state: HealthButlerState) -> dict:
  120. """阶段3: LLM 生成个性化健康建议"""
  121. messages = [SystemMessage(content=await get_prompt("health_butler") or DEFAULT_BUTLER_PROMPT)]
  122. # 注入健康知识库上下文
  123. kb_ctx = state.get("kb_context", {})
  124. knowledge_results = state.get("knowledge_results", [])
  125. if knowledge_results:
  126. kb_text_parts = []
  127. for r in knowledge_results:
  128. d = r.get("data", {})
  129. name = d.get("itemName", r.get("keyword", ""))
  130. cat = d.get("category", "")
  131. desc = d.get("description", "")
  132. suggestion = d.get("suggestion", "")
  133. normal_range = d.get("normalRange", "")
  134. parts = [f"【{r['type']}】{name}"]
  135. if cat:
  136. parts.append(f"(分类: {cat})")
  137. if normal_range:
  138. parts.append(f"正常范围: {normal_range}")
  139. if desc:
  140. parts.append(f"{desc}")
  141. if suggestion:
  142. parts.append(f"建议: {suggestion}")
  143. kb_text_parts.append(" ".join(parts))
  144. kb_text = "\n---\n".join(kb_text_parts)
  145. messages.append(SystemMessage(content=f"健康知识库参考:\n{kb_text}"))
  146. # 注入家庭上下文
  147. family_id = state.get("family_id")
  148. if family_id:
  149. try:
  150. client = await java._get_client()
  151. resp = await client.post(
  152. "/api/ai/context",
  153. json={
  154. "user_id": str(state["user_id"]),
  155. "params": {
  156. "familyId": family_id,
  157. **({"childId": state["child_id"]} if state.get("child_id") else {}),
  158. },
  159. },
  160. )
  161. data = resp.json()
  162. if data.get("code") == 200:
  163. ctx = data.get("data", {})
  164. if ctx:
  165. messages.append(SystemMessage(content=f"家庭上下文:\n{json.dumps(ctx, ensure_ascii=False, indent=2)}"))
  166. except Exception as e:
  167. logger.debug("加载家庭上下文跳过: %s", e)
  168. # 注入长期记忆
  169. try:
  170. memories = await memory_mgr.recall(state["user_id"], state["query"])
  171. if memories:
  172. mem_text = "\n".join([f"- {m}" for m in memories])
  173. messages.append(SystemMessage(content=f"相关历史对话:\n{mem_text}"))
  174. except Exception as e:
  175. logger.warning("召回记忆失败: %s", e)
  176. # 用户消息
  177. focus = state.get("focus")
  178. query = state["query"]
  179. if focus:
  180. query = f"[重点关注: {focus}] {query}"
  181. messages.append(HumanMessage(content=query))
  182. response = await llm_with_tools.ainvoke(messages)
  183. answer = response.content
  184. # 提取任务
  185. import re
  186. tasks = []
  187. pattern = r'\[TASK:\s*\{([^}]+)\}\]'
  188. for match in re.findall(pattern, answer):
  189. task_info = {}
  190. for kv in match.split(","):
  191. if ":" in kv:
  192. k, v = kv.split(":", 1)
  193. task_info[k.strip()] = v.strip().strip('"').strip("'")
  194. if task_info.get("title"):
  195. tasks.append(task_info)
  196. sources = []
  197. for r in knowledge_results:
  198. d = r.get("data", {})
  199. sources.append({
  200. "type": "health_kb",
  201. "name": d.get("itemName", r.get("keyword", "")),
  202. "category": d.get("category", ""),
  203. })
  204. if response.response_metadata.get("tool_calls"):
  205. for tc in response.response_metadata["tool_calls"]:
  206. sources.append({
  207. "type": "tool",
  208. "name": tc.get("name", ""),
  209. })
  210. return {
  211. "answer": answer,
  212. "tasks": tasks,
  213. "sources": sources,
  214. "messages": [
  215. {"role": "user", "content": state["query"]},
  216. {"role": "assistant", "content": answer},
  217. ],
  218. }
  219. async def save_memory(state: HealthButlerState) -> dict:
  220. try:
  221. if state.get("messages"):
  222. await memory_mgr.save_conversation(
  223. state["user_id"],
  224. state["conversation_id"],
  225. state["messages"],
  226. )
  227. except Exception as e:
  228. logger.warning("保存记忆失败: %s", e)
  229. return {}
  230. # ── 路由 ──
  231. def route_after_knowledge(state: HealthButlerState) -> Literal["search_knowledge", "llm_answer"]:
  232. kb_ctx = state.get("kb_context", {})
  233. if kb_ctx.get("kb_loaded"):
  234. return "search_knowledge"
  235. return "llm_answer"
  236. # ── 构建图 ──
  237. builder.add_node("load_knowledge", load_knowledge)
  238. builder.add_node("search_knowledge", search_knowledge)
  239. builder.add_node("llm_answer", llm_answer)
  240. builder.add_node("save_memory", save_memory)
  241. builder.add_edge(START, "load_knowledge")
  242. builder.add_conditional_edges("load_knowledge", route_after_knowledge)
  243. builder.add_edge("search_knowledge", "llm_answer")
  244. builder.add_edge("llm_answer", "save_memory")
  245. builder.add_edge("save_memory", END)
  246. checkpointer = MemorySaver()
  247. graph = builder.compile(checkpointer=checkpointer)
  248. return graph
  249. def _extract_health_keywords(query: str) -> list[str]:
  250. """从 query 中提取健康相关关键词(菌属、指标、食物等名词)"""
  251. # 简单策略: 提取2字以上的中文词作为候选
  252. import re
  253. # 提取中文字符组成的连续词段
  254. chinese_words = re.findall(r'[\u4e00-\u9fff]{2,}', query)
  255. # 提取英文/拉丁名词组
  256. english_words = re.findall(r'[A-Za-z][A-Za-z\s]{2,20}[A-Za-z]', query)
  257. keywords = []
  258. # 过滤常见通用词
  259. stopwords = {"请问", "帮我", "能不能", "可不可以", "怎么样", "为什么",
  260. "是什么", "怎么办", "什么意思", "好不好", "是否", "你好",
  261. "谢谢", "请", "我想", "需要", "关于", "一个", "这个",
  262. "那个", "哪些", "什么", "怎么", "如何", "多少", "是不是"}
  263. for w in chinese_words:
  264. if w not in stopwords and len(w) >= 2:
  265. keywords.append(w)
  266. for w in english_words:
  267. w = w.strip()
  268. if len(w) > 2:
  269. keywords.append(w)
  270. # 去重 + 限数量
  271. seen = set()
  272. deduped = []
  273. for kw in keywords:
  274. if kw not in seen:
  275. seen.add(kw)
  276. deduped.append(kw)
  277. return deduped[:10]