analysis_graph.py 3.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596
  1. from typing import TypedDict, Optional
  2. from langgraph.graph import StateGraph, START, END
  3. from langgraph.checkpoint import MemorySaver
  4. from langchain_openai import ChatOpenAI
  5. from langchain_core.messages import SystemMessage, HumanMessage
  6. from app.tools.report_tools import get_report_detail, get_survey_data, get_dimension_scores
  7. from app.config import settings
  8. import logging
  9. logger = logging.getLogger(__name__)
  10. class AnalysisState(TypedDict):
  11. report_id: int
  12. user_id: int
  13. focus: Optional[str] # nutrition / gut / chronic / overall
  14. report_data: Optional[dict]
  15. survey_data: Optional[dict]
  16. dimension_scores: Optional[dict]
  17. analysis: Optional[str]
  18. recommendations: list[str]
  19. ANALYSIS_SYSTEM_PROMPT = """你是一个儿童健康报告解读专家。根据健康报告数据和调研问卷, 提供专业、易懂的分析。
  20. 分析原则:
  21. 1. 用通俗语言解释各项指标含义
  22. 2. 关注异常指标, 给出改善建议
  23. 3. 结合问卷数据提供个性化分析
  24. 4. 五维能量(身/心/智/行/富)角度解读整体状况
  25. 5. 输出格式: 总体评估 → 分项分析 → 改善建议
  26. """
  27. def create_analysis_graph():
  28. """创建报告解读 StateGraph"""
  29. llm = ChatOpenAI(
  30. model=settings.llm_model,
  31. api_key=settings.llm_api_key,
  32. base_url=settings.llm_base_url,
  33. temperature=0.3,
  34. )
  35. llm_with_tools = llm.bind_tools([
  36. get_report_detail,
  37. get_survey_data,
  38. get_dimension_scores,
  39. ])
  40. builder = StateGraph(AnalysisState)
  41. async def gather_data(state: AnalysisState) -> dict:
  42. """收集报告 + 问卷 + 维度数据"""
  43. result = {}
  44. if state.get("report_id"):
  45. try:
  46. detail = await get_report_detail.ainvoke({"report_id": state["report_id"]})
  47. result["report_data"] = detail
  48. except Exception as e:
  49. logger.warning("获取报告详情失败: %s", e)
  50. try:
  51. survey = await get_survey_data.ainvoke({"report_id": state["report_id"]})
  52. result["survey_data"] = survey
  53. except Exception as e:
  54. logger.warning("获取问卷数据失败: %s", e)
  55. return result
  56. async def analyze(state: AnalysisState) -> dict:
  57. """LLM 分析"""
  58. context_parts = []
  59. if state.get("report_data"):
  60. context_parts.append(f"报告数据: {state['report_data']}")
  61. if state.get("survey_data"):
  62. context_parts.append(f"问卷数据: {state['survey_data']}")
  63. if state.get("focus"):
  64. context_parts.append(f"重点关注: {state['focus']}")
  65. context_text = "\n".join(context_parts) if context_parts else "暂无数据"
  66. messages = [
  67. SystemMessage(content=ANALYSIS_SYSTEM_PROMPT),
  68. SystemMessage(content=f"待分析数据:\n{context_text}"),
  69. HumanMessage(content="请分析以上健康数据, 给出评估和建议。"),
  70. ]
  71. response = await llm_with_tools.ainvoke(messages)
  72. return {"analysis": response.content}
  73. builder.add_node("gather_data", gather_data)
  74. builder.add_node("analyze", analyze)
  75. builder.add_edge(START, "gather_data")
  76. builder.add_edge("gather_data", "analyze")
  77. builder.add_edge("analyze", END)
  78. checkpointer = MemorySaver()
  79. return builder.compile(checkpointer=checkpointer)