health_butler_graph.py 12 KB

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