health_plan_graph.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308
  1. """
  2. 健康方案生成 LangGraph - 分步骤结构化方案生成
  3. 工作流:
  4. 1. 数据收集: 家庭成员信息 + 健康指标 + 知识库检索
  5. 2. LLM 生成: 总览概述 + 营养/饮食/运动 三个 section
  6. 3. 返回结构化 JSON
  7. 每次调用 LLM 时遵循统一的数据组装方式:
  8. - 用户指标 (来自 Java context API)
  9. - 知识库参考 (RAG 检索)
  10. - 系统提示词 (带输出格式模板)
  11. """
  12. from typing import TypedDict, Literal
  13. from langgraph.graph import StateGraph, START, END
  14. from langchain_openai import ChatOpenAI
  15. from langchain_core.messages import SystemMessage, HumanMessage
  16. from app.rag.retriever import RagRetriever
  17. from app.config import settings
  18. from app.tools.java_client import JavaClient
  19. from app.prompt_service import get_prompt
  20. import logging
  21. logger = logging.getLogger(__name__)
  22. SECTION_KEYS = ["nutrition", "diet", "exercise"]
  23. DEFAULT_PROMPT = """你是一个专业的家庭健康方案规划师。根据用户的健康数据和目标,生成结构化的改善方案。
  24. ## 输出格式(必须严格遵守 JSON)
  25. ```json
  26. {
  27. "overview": "总体概述(100字以内,说明方案目标和核心策略)",
  28. "sections": [
  29. {
  30. "key": "nutrition",
  31. "title": "营养补充建议",
  32. "content": "Markdown 格式的详细内容,包含具体产品推荐、用量、服用时间",
  33. "items": [
  34. {"name": "产品名", "dosage": "用量", "timing": "服用时间", "reason": "推荐理由"}
  35. ]
  36. },
  37. {
  38. "key": "diet",
  39. "title": "饮食建议",
  40. "content": "Markdown 格式的餐饮建议,包含早餐/午餐/晚餐建议",
  41. "items": [{"meal": "餐型", "food": "食物建议", "notes": "注意事项"}]
  42. },
  43. {
  44. "key": "exercise",
  45. "title": "运动计划",
  46. "content": "Markdown 格式的运动建议,包含频率、时长、类型",
  47. "items": [{"type": "运动类型", "duration": "时长", "frequency": "频率", "notes": "注意事项"}]
  48. }
  49. ]
  50. }
  51. ```
  52. ## 原则
  53. 1. 基于实际数据给出建议,不编造
  54. 2. 引用知识库时标注来源
  55. 3. 建议要具体可执行,避免空泛
  56. 4. 营养补充部分要具体到品牌/产品类型和用量
  57. 5. 严重健康问题建议咨询医生
  58. """
  59. REGENERATE_SECTION_PROMPT = """你是一个家庭健康方案规划师。根据用户反馈重新生成指定部分的内容。
  60. ## 当前方案内容
  61. {existing_section_content}
  62. ## 用户反馈
  63. {feedback}
  64. ## 相关背景数据
  65. {context_summary}
  66. 请重新生成该部分内容,保持与原格式一致。只输出新的 content 字段值(Markdown 格式),不需要输出 JSON 结构。"""
  67. class PlanState(TypedDict):
  68. member_ids: str
  69. dimensions: str
  70. goal: str
  71. family_id: int
  72. members_info: dict
  73. indicators: list
  74. abnormal_indicators: list
  75. kb_results: list
  76. overview: str
  77. nutrition_section: str
  78. diet_section: str
  79. exercise_section: str
  80. full_response: dict
  81. error: str
  82. async def collect_data(state: PlanState) -> dict:
  83. """Step 1: 收集用户数据 + 知识库检索"""
  84. java = JavaClient()
  85. retriever = RagRetriever(collection_name="cfc_knowledge")
  86. member_ids_str = state.get("member_ids", "")
  87. member_ids = [m.strip() for m in member_ids_str.split(",") if m.strip()]
  88. goal = state.get("goal", "")
  89. dimensions = state.get("dimensions", "") or ""
  90. # 1a. 获取家庭成员信息
  91. members_info = []
  92. for uid_str in member_ids:
  93. ctx = await java.get_family_context(int(uid_str), "child_info")
  94. children = ctx.get("children", []) if isinstance(ctx, dict) else []
  95. for child in children:
  96. cid = str(child.get("用户ID", ""))
  97. if cid == uid_str:
  98. members_info.append({
  99. "id": cid,
  100. "name": child.get("姓名", f"成员{cid}"),
  101. "age": child.get("年龄", "未知"),
  102. "energy": child.get("能量", 0),
  103. })
  104. break
  105. if not any(m["id"] == uid_str for m in members_info):
  106. members_info.append({"id": uid_str, "name": f"成员{uid_str}", "age": "未知", "energy": 0})
  107. # 1b. 获取健康指标
  108. all_indicators = []
  109. abnormal_list = []
  110. for member in members_info:
  111. reports = await java.get_member_reports(int(member["id"]))
  112. if not reports:
  113. continue
  114. latest = max(reports, key=lambda r: r.get("reportDate", ""))
  115. report_id = latest.get("id")
  116. indicators = await java.get_report_indicators(report_id)
  117. for ind in indicators:
  118. ind["_member_id"] = member["id"]
  119. ind["_member_name"] = member["name"]
  120. all_indicators.append(ind)
  121. # 1c. 识别异常指标
  122. known_indicators = {}
  123. for ind in all_indicators:
  124. name = ind.get("indicatorName", "").strip()
  125. if not name or name in known_indicators:
  126. continue
  127. for itype in ["indicator", "bacteria", "nutrient"]:
  128. kb = await java.query_health_knowledge(itype, name)
  129. if kb:
  130. known_indicators[name] = kb
  131. break
  132. for ind in all_indicators:
  133. name = ind.get("indicatorName", "")
  134. status = ind.get("status", "")
  135. kb = known_indicators.get(name, {})
  136. entry = {
  137. "member": ind.get("_member_name", ""),
  138. "indicator": name,
  139. "value": ind.get("indicatorValue", ""),
  140. "unit": ind.get("unit", ""),
  141. "ref_range": kb.get("normalRange", ind.get("refRange", "")),
  142. "description": kb.get("description", ""),
  143. "suggestion": kb.get("suggestion", ""),
  144. }
  145. if status in ("abnormal", "high", "low", "偏高", "偏低"):
  146. abnormal_list.append(entry)
  147. # 1d. 知识库检索
  148. kb_results = []
  149. queries = set()
  150. for ind in abnormal_list:
  151. queries.add(ind["indicator"])
  152. queries.add(goal)
  153. if dimensions:
  154. queries.add(dimensions)
  155. for q in list(queries)[:8]:
  156. results = await retriever.retrieve(q, k=3)
  157. kb_results.extend(results)
  158. return {
  159. "members_info": members_info,
  160. "indicators": all_indicators,
  161. "abnormal_indicators": abnormal_list,
  162. "kb_results": kb_results,
  163. }
  164. async def build_prompt(state: PlanState) -> str:
  165. """组装 LLM prompt"""
  166. base_prompt = await get_prompt("health_plan") or DEFAULT_PROMPT
  167. parts = [base_prompt]
  168. parts.append(f"\n## 用户目标\n{state['goal']}")
  169. if state.get("dimensions"):
  170. parts.append(f"\n## 重点关注维度\n{state['dimensions']}")
  171. parts.append("\n## 家庭成员")
  172. for m in state["members_info"]:
  173. parts.append(f"- {m['name']} (年龄: {m['age']})")
  174. if state["abnormal_indicators"]:
  175. parts.append("\n## 异常指标")
  176. for ind in state["abnormal_indicators"][:8]:
  177. parts.append(
  178. f"- {ind['member']} - {ind['indicator']}: {ind['value']}{ind.get('unit','')} "
  179. f"(参考: {ind['ref_range']})"
  180. )
  181. if ind.get("description"):
  182. parts.append(f" 说明: {ind['description']}")
  183. if state["kb_results"]:
  184. parts.append("\n## 知识库参考")
  185. for r in state["kb_results"][:6]:
  186. title = r.get("metadata", {}).get("title", "")
  187. content = r.get("content", "")[:200]
  188. parts.append(f"---\n{title}\n{content}")
  189. return "\n".join(parts)
  190. async def generate_plan(state: PlanState) -> dict:
  191. """Step 2: 调用 LLM 生成结构化方案"""
  192. llm = ChatOpenAI(
  193. model=settings.llm_model,
  194. api_key=settings.llm_api_key,
  195. base_url=settings.llm_base_url,
  196. temperature=0.3,
  197. )
  198. prompt = await build_prompt(state)
  199. messages = [SystemMessage(content=prompt)]
  200. try:
  201. response = await llm.ainvoke(messages)
  202. answer = response.content
  203. # 解析 JSON
  204. import json
  205. try:
  206. # 提取 JSON 块
  207. start = answer.find("{")
  208. end = answer.rfind("}") + 1
  209. if start >= 0 and end > start:
  210. json_str = answer[start:end]
  211. parsed = json.loads(json_str)
  212. return {"full_response": parsed, "overview": parsed.get("overview", "")}
  213. except (json.JSONDecodeError, Exception) as e:
  214. logger.warning("解析方案 JSON 失败,使用原始文本: %s", e)
  215. return {"full_response": {"raw": answer}, "overview": answer[:200]}
  216. except Exception as e:
  217. logger.error("LLM 生成方案失败: %s", e)
  218. return {"error": str(e)}
  219. async def regenerate_section(state: PlanState) -> dict:
  220. """重新生成指定 section"""
  221. section_key = state.get("section", "nutrition")
  222. feedback = state.get("feedback", "")
  223. existing_content = state.get("existing_section_content", "")
  224. llm = ChatOpenAI(
  225. model=settings.llm_model,
  226. api_key=settings.llm_api_key,
  227. base_url=settings.llm_base_url,
  228. temperature=0.3,
  229. )
  230. # 构建上下文摘要
  231. ctx_parts = []
  232. for m in state.get("members_info", []):
  233. ctx_parts.append(f"- {m['name']} (年龄: {m['age']})")
  234. if state.get("goal"):
  235. ctx_parts.append(f"目标: {state['goal']}")
  236. if state.get("abnormal_indicators"):
  237. for ind in state["abnormal_indicators"][:5]:
  238. ctx_parts.append(f"- {ind['member']}: {ind['indicator']}={ind['value']}")
  239. regenerate_template = await get_prompt("health_plan_regenerate") or REGENERATE_SECTION_PROMPT
  240. try:
  241. prompt = regenerate_template.format(
  242. existing_section_content=existing_content[:500],
  243. feedback=feedback,
  244. context_summary="\n".join(ctx_parts),
  245. )
  246. except (KeyError, IndexError, ValueError):
  247. logger.warning("Java 配置的 health_plan_regenerate 模板缺少占位符,回退本地模板")
  248. prompt = REGENERATE_SECTION_PROMPT.format(
  249. existing_section_content=existing_content[:500],
  250. feedback=feedback,
  251. context_summary="\n".join(ctx_parts),
  252. )
  253. try:
  254. response = await llm.ainvoke([SystemMessage(content=prompt)])
  255. return {"regenerated_content": response.content}
  256. except Exception as e:
  257. logger.error("重新生成方案 section 失败: %s", e)
  258. return {"error": str(e)}
  259. def create_health_plan_graph():
  260. builder = StateGraph(PlanState)
  261. builder.add_node("collect_data", collect_data)
  262. builder.add_node("generate_plan", generate_plan)
  263. builder.add_node("regenerate_section", regenerate_section)
  264. builder.add_edge(START, "collect_data")
  265. builder.add_edge("collect_data", "generate_plan")
  266. builder.add_edge("generate_plan", END)
  267. # regenerate_section 从外部直接调用,不走图
  268. graph = builder.compile()
  269. return graph