analysis_graph.py 3.4 KB

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