health_butler_graph.py 12 KB

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