main.py 1.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748
  1. """
  2. 数据分析入口:根据用户问题自动路由到 SQL Agent 或可视化 Agent。
  3. """
  4. from config import llm
  5. from agents import sql_agent, visualization_agent
  6. def classify_intent(user_query: str) -> str:
  7. """用 LLM 判断用户意图:query / visualize / both。"""
  8. prompt = f"""判断以下用户问题属于哪种类型:
  9. - "query":需要查询数据库获取数据
  10. - "visualize":需要画图或做统计分析
  11. - "both":需要先查数据,再画图分析
  12. 用户问题:{user_query}
  13. 只回复一个词:query / visualize / both"""
  14. return llm.invoke(prompt).content.strip().lower()
  15. def run_data_analysis(user_query: str) -> str:
  16. """统一入口:自动判断意图并路由到对应 Agent。"""
  17. intent = classify_intent(user_query)
  18. if intent == "query":
  19. result = sql_agent.invoke({"messages": [{"role": "user", "content": user_query}]})
  20. return result["messages"][-1].content
  21. elif intent == "visualize":
  22. result = visualization_agent.invoke({"messages": [{"role": "user", "content": user_query}]})
  23. return result["messages"][-1].content
  24. else: # both
  25. data_result = sql_agent.invoke({"messages": [{"role": "user", "content": user_query}]})
  26. data_content = data_result["messages"][-1].content
  27. viz_input = f"基于以下数据进行可视化分析:\n{data_content}\n\n原始问题:{user_query}"
  28. viz_result = visualization_agent.invoke({"messages": [{"role": "user", "content": viz_input}]})
  29. return viz_result["messages"][-1].content
  30. if __name__ == "__main__":
  31. user_query = """
  32. 列出所有2023年入职的员工,并列出他们的所属部门、姓名和薪资,按薪资从高到低排序?并画柱状图展示
  33. """
  34. result = run_data_analysis(user_query)
  35. print(result)