health_coach_graph.py 6.3 KB

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