health_coach_graph.py 7.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193
  1. """
  2. AI 健康教练 LangGraph - 专为儿童健康咨询设计的对话图
  3. Uses RAG retrieval from health knowledge base.
  4. """
  5. from typing import TypedDict
  6. from langgraph.graph import StateGraph, START, END
  7. from app.checkpointer import get_checkpointer
  8. from langchain_openai import ChatOpenAI
  9. from langchain_core.messages import SystemMessage, HumanMessage
  10. from app.rag.retriever import RagRetriever
  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. from app.tools.java_client import JavaClient
  16. import logging
  17. logger = logging.getLogger(__name__)
  18. class HealthCoachState(TypedDict):
  19. query: str
  20. user_id: int
  21. child_id: int | None
  22. conversation_id: str | None
  23. context: dict | None
  24. answer: str | None
  25. sources: list[dict]
  26. memory_messages: list | None
  27. DEFAULT_COACH_PROMPT = """你是一个儿童健康教练,喜欢用提问的方式引导家长发现问题和解决方案。
  28. 核心原则:
  29. 1. 先问后答:不要直接给答案,用问题引导家长思考
  30. - "孩子这种情况多久了?"、"平时饮食习惯是怎样的?"、"体重身高发育曲线有没有记录?"
  31. 2. 用开放式问题了解情况,再给针对性建议
  32. 3. 语气温暖,像一个耐心的教练在陪伴家长成长
  33. 你的职责:
  34. 1. 儿童体质调理(如挑食、瘦弱、易疲劳、反复生病)—— 用提问了解情况
  35. 2. 科学喂养和营养建议 —— 先问饮食现状再给建议
  36. 3. 日常保健(睡眠、运动、季节防护)—— 引导家长发现生活中的问题
  37. 4. 体质辨识基础咨询 —— 通过提问帮助家长理解孩子体质特点
  38. 5. 健康问题判断 —— 先收集症状信息再建议是否就医
  39. 回答结构(教练风格):
  40. 1. 提问:先问1-2个关键问题了解情况
  41. 2. 分析:简短总结你观察到的问题
  42. 3. 建议:给出2-3个可操作的建议
  43. 4. 跟进问句:留一个后续跟进的问题,体现教练陪伴感
  44. 原则:
  45. - 涉及严重症状(急性腹痛、高烧不退、严重外伤等)立刻建议就医,不提问
  46. - 不确定时明确说"建议咨询专业医生"
  47. - 始终用中文回答
  48. """
  49. def create_health_coach_graph():
  50. """创建AI健康教练 StateGraph"""
  51. # LLM
  52. llm = ChatOpenAI(
  53. model=settings.llm_model,
  54. api_key=settings.llm_api_key,
  55. base_url=settings.llm_base_url,
  56. temperature=settings.llm_temperature,
  57. )
  58. # 嵌入和检索
  59. from app.rag.embeddings import get_embeddings
  60. embeddings = get_embeddings()
  61. retriever = RagRetriever(collection_name="health_knowledge")
  62. memory_mgr = MemoryManager(embeddings)
  63. java = JavaClient()
  64. # 带工具的 LLM(健康知识库检索)
  65. async def health_retrieve(query: str, user_id: int) -> list[dict]:
  66. try:
  67. results = await retriever.retrieve(
  68. query,
  69. filters={"user_id": user_id},
  70. k=3,
  71. )
  72. return results
  73. except Exception as e:
  74. logger.warning("健康知识检索失败: %s", e)
  75. return []
  76. async def generate_answer(state: HealthCoachState) -> dict:
  77. """带健康知识检索的 LLM 生成"""
  78. # 人格路由:context.coach_id (xibao/fubao) -> 专属prompt,三级回退保证可用性
  79. _ctx = state.get("context") or {}
  80. _coach_id = _ctx.get("coach_id")
  81. if _coach_id not in ("xibao", "fubao"):
  82. _coach_id = None
  83. _persona_prompt = await get_prompt(f"health_coach_{_coach_id}") if _coach_id else None
  84. _base_prompt = await get_prompt("health_coach")
  85. messages = [SystemMessage(content=_persona_prompt or _base_prompt or DEFAULT_COACH_PROMPT)]
  86. # 画像注入
  87. member_id = _ctx.get("child_id")
  88. portrait_text = await get_user_portrait_text(java, state["user_id"], member_id)
  89. if portrait_text:
  90. messages.insert(1, SystemMessage(content=portrait_text))
  91. # 家庭上下文
  92. ctx = state.get("context") or {}
  93. if ctx:
  94. child_info = []
  95. if ctx.get("child_id"):
  96. child_info.append(f"孩子ID: {ctx['child_id']}")
  97. if ctx.get("family_id"):
  98. child_info.append(f"家庭ID: {ctx['family_id']}")
  99. if child_info:
  100. messages.append(SystemMessage(
  101. content="用户背景信息: " + ", ".join(child_info)
  102. ))
  103. # 健康知识检索
  104. retrieved = await health_retrieve(state["query"], state["user_id"])
  105. sources = []
  106. if retrieved:
  107. context_texts = []
  108. for i, r in enumerate(retrieved[:3]):
  109. src = r.get("metadata", {})
  110. title = src.get("title", "健康知识")
  111. content = r["content"][:300]
  112. context_texts.append(f"[{i+1}] {title}:\n{content}")
  113. sources.append({
  114. "type": "knowledge",
  115. "title": title,
  116. "content": content,
  117. })
  118. knowledge_context = "\n\n---\n".join(context_texts)
  119. messages.append(SystemMessage(
  120. content=f"以下是与问题相关的健康知识参考:\n{knowledge_context}"
  121. ))
  122. # 长期记忆
  123. try:
  124. memories = await memory_mgr.recall(state["user_id"], state["query"])
  125. if memories:
  126. mem_text = "\n".join([f"- {m}" for m in memories[:3]])
  127. messages.append(SystemMessage(
  128. content=f"该用户的历史健康咨询记录:\n{mem_text}"
  129. ))
  130. except Exception as e:
  131. logger.warning("召回健康记忆失败: %s", e)
  132. # 用户问题
  133. messages.append(HumanMessage(content=state["query"]))
  134. response = await llm.ainvoke(messages)
  135. answer = response.content
  136. return {
  137. "answer": answer,
  138. "sources": sources,
  139. "memory_messages": [
  140. {"role": "user", "content": state["query"]},
  141. {"role": "assistant", "content": answer},
  142. ],
  143. }
  144. async def save_memory(state: HealthCoachState) -> dict:
  145. """保存对话到长期记忆"""
  146. try:
  147. if state.get("memory_messages"):
  148. await memory_mgr.save_conversation(
  149. state["user_id"],
  150. state.get("conversation_id") or f"health_{state['user_id']}",
  151. state["memory_messages"],
  152. )
  153. except Exception as e:
  154. logger.warning("保存健康教练记忆失败: %s", e)
  155. return {}
  156. # ── 构建图 ──
  157. builder = StateGraph(HealthCoachState)
  158. builder.add_node("generate_answer", generate_answer)
  159. builder.add_node("save_memory", save_memory)
  160. builder.add_edge(START, "generate_answer")
  161. builder.add_edge("generate_answer", "save_memory")
  162. builder.add_edge("save_memory", END)
  163. checkpointer = get_checkpointer()
  164. graph = builder.compile(checkpointer=checkpointer)
  165. return graph