nutrition_graph.py 5.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159
  1. """AI 营养助手 LangGraph — 基于健康报告的个性化营养建议"""
  2. from typing import TypedDict
  3. from langgraph.graph import StateGraph, START, END
  4. from langgraph.checkpoint.memory import MemorySaver
  5. from langchain_openai import ChatOpenAI
  6. from langchain_core.messages import SystemMessage, HumanMessage
  7. from app.tools.java_client import JavaClient
  8. from app.tools.product_tools import (
  9. search_product_by_keyword,
  10. search_article_by_keyword,
  11. )
  12. from app.memory.store import MemoryManager
  13. from app.config import settings
  14. from app.prompt_service import get_prompt
  15. from app.portrait_service import get_user_portrait_text
  16. import logging
  17. logger = logging.getLogger(__name__)
  18. class NutritionState(TypedDict):
  19. query: str
  20. user_id: int
  21. conversation_id: str
  22. child_id: int | None
  23. report_id: int | None
  24. context: dict | None
  25. answer: str | None
  26. tasks: list[dict]
  27. sources: list[dict]
  28. messages: list | None
  29. DEFAULT_NUTRITION_PROMPT = """你是一个儿童营养健康顾问,专注于基于检测报告提供精准的营养改善建议。
  30. 你的能力:
  31. 1. 解读儿童健康报告(菌群、营养指标、体检数据)
  32. 2. 基于报告数据推荐合适的营养产品和膳食方案
  33. 3. 结合孩子体质给出个性化的饮食建议
  34. 4. 推荐相关的健康活动和科普文章
  35. 回答原则:
  36. 1. 用中文,语气温暖专业,像营养师一样亲切
  37. 2. 必须基于报告数据给出建议,不凭空猜测
  38. 3. 引用知识库中的营养素/食物信息作为证据
  39. 4. 需要产品推荐时使用搜索工具,说明推荐理由
  40. 5. 生成 [TASK: {"title": "任务名", "dimension": "身", "points": 10, "frequency": "每天"}] 标记创建营养行动任务
  41. 6. 严重健康问题建议咨询专业医生或营养师
  42. 7. 不要编造医学建议,不夸大效果
  43. """
  44. def create_nutrition_graph():
  45. """创建营养助手 StateGraph"""
  46. llm = ChatOpenAI(
  47. model=settings.llm_model,
  48. api_key=settings.llm_api_key,
  49. base_url=settings.llm_base_url,
  50. temperature=settings.llm_temperature,
  51. )
  52. from app.rag.embeddings import get_embeddings
  53. embeddings = get_embeddings()
  54. memory_mgr = MemoryManager(embeddings)
  55. java = JavaClient()
  56. llm_with_tools = llm.bind_tools([
  57. search_product_by_keyword,
  58. search_article_by_keyword,
  59. ])
  60. builder = StateGraph(NutritionState)
  61. async def generate_answer(state: NutritionState) -> dict:
  62. """核心 LLM 调用 + 画像注入"""
  63. # 加载营养助手 prompt
  64. nutrition_prompt = await get_prompt("nutrition_assistant") or DEFAULT_NUTRITION_PROMPT
  65. messages = [SystemMessage(content=nutrition_prompt)]
  66. # 画像注入(在人格 prompt 之后)
  67. member_id = state.get("child_id")
  68. portrait_text = await get_user_portrait_text(java, state["user_id"], member_id)
  69. if portrait_text:
  70. messages.insert(1, SystemMessage(content=portrait_text))
  71. # 家庭上下文
  72. ctx = state.get("context") or {}
  73. if ctx:
  74. ctx_parts = []
  75. if ctx.get("child_id"):
  76. ctx_parts.append(f"孩子ID: {ctx['child_id']}")
  77. if ctx.get("family_id"):
  78. ctx_parts.append(f"家庭ID: {ctx['family_id']}")
  79. if ctx.get("report_id"):
  80. ctx_parts.append(f"报告ID: {ctx['report_id']}")
  81. if ctx_parts:
  82. messages.append(SystemMessage(content="用户背景信息: " + ", ".join(ctx_parts)))
  83. # 长期记忆
  84. try:
  85. memories = await memory_mgr.recall(state["user_id"], state["query"])
  86. if memories:
  87. mem_text = "\n".join([f"- {m}" for m in memories[:3]])
  88. messages.append(SystemMessage(content=f"该用户的历史咨询记录:\n{mem_text}"))
  89. except Exception as e:
  90. logger.warning("召回营养记忆失败: %s", e)
  91. # 用户消息
  92. messages.append(HumanMessage(content=state["query"]))
  93. response = await llm_with_tools.ainvoke(messages)
  94. answer = response.content
  95. # 提取任务和来源
  96. tasks = []
  97. sources = []
  98. if response.response_metadata.get("tool_calls"):
  99. for tc in response.response_metadata["tool_calls"]:
  100. sources.append({
  101. "type": "tool",
  102. "name": tc.get("name", ""),
  103. "input": tc.get("args", {}),
  104. })
  105. return {
  106. "answer": answer,
  107. "tasks": tasks,
  108. "sources": sources,
  109. "messages": [
  110. {"role": "user", "content": state["query"]},
  111. {"role": "assistant", "content": answer},
  112. ],
  113. }
  114. async def save_memory(state: NutritionState) -> dict:
  115. """保存对话到长期记忆"""
  116. try:
  117. if state.get("messages"):
  118. await memory_mgr.save_conversation(
  119. state["user_id"],
  120. state.get("conversation_id") or f"nutrition_{state['user_id']}",
  121. state["messages"],
  122. )
  123. except Exception as e:
  124. logger.warning("保存营养记忆失败: %s", e)
  125. return {}
  126. # 构建图
  127. builder.add_node("generate_answer", generate_answer)
  128. builder.add_node("save_memory", save_memory)
  129. builder.add_edge(START, "generate_answer")
  130. builder.add_edge("generate_answer", "save_memory")
  131. builder.add_edge("save_memory", END)
  132. checkpointer = MemorySaver()
  133. graph = builder.compile(checkpointer=checkpointer)
  134. return graph