health_butler_graph.py 12 KB

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