work_01.py 8.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201
  1. """
  2. 1. 加载配置
  3. 2. 创建大模型
  4. 3. 配置并获取db,db数据,db工具
  5. 4. 配置沙箱
  6. 5. 创建python代码画图工具
  7. 6. 创建智能体
  8. 7. 使用智能体
  9. """
  10. from dotenv import load_dotenv
  11. import os
  12. from langchain_openai import ChatOpenAI
  13. from langchain.agents import create_agent
  14. from langchain.tools import tool
  15. import traceback
  16. from io import StringIO
  17. from contextlib import redirect_stdout
  18. from langchain.tools import tool
  19. import pandas as pd
  20. import matplotlib.pyplot as plt
  21. import seaborn as sns
  22. import numpy as np
  23. from langchain_community.utilities import SQLDatabase
  24. from langchain_community.agent_toolkits import SQLDatabaseToolkit
  25. from sqlalchemy import create_engine
  26. load_dotenv(override=True)
  27. deepseek_base_url = os.getenv("DEEPSEEK_BASE_URL")
  28. deepseek_base_key = os.getenv("DEEPSEEK_BASE_KEY")
  29. deepseek_base_name = os.getenv("DEEPSEEK_BASE_NAME")
  30. print(deepseek_base_key)
  31. print(deepseek_base_name)
  32. print(deepseek_base_url)
  33. llm = ChatOpenAI(
  34. base_url = deepseek_base_url,
  35. api_key = deepseek_base_key,
  36. model= deepseek_base_name
  37. )
  38. db_uri = (
  39. f"mysql+pymysql://{os.getenv('DB_USER')}:{os.getenv('DB_PASSWORD')}"
  40. f"@{os.getenv('DB_HOST')}:{os.getenv('DB_PORT')}/{os.getenv('DB_NAME')}"
  41. )
  42. db = SQLDatabase.from_uri(db_uri)
  43. # ============================================================
  44. # 配置 matplotlib 中⽂显示
  45. # ============================================================
  46. # Windows ⽤ SimHei(⿊体),macOS ⽤ PingFang SC
  47. # 如果还是乱码,试试安装 fonts-noto-cjk 并清除缓存
  48. plt.rcParams["font.sans-serif"] = ["SimHei", "PingFang SC", "DejaVu Sans"]
  49. plt.rcParams["axes.unicode_minus"] = False # 解决负号显示为⽅块的问题
  50. # ============================================================
  51. # 从数据库加载数据到 Pandas DataFrame
  52. # ============================================================
  53. # ⽤ SQLAlchemy engine 复⽤连接,避免重复创建连接池
  54. engine = create_engine(db_uri)
  55. employees_df = pd.read_sql("SELECT * FROM employees", engine)
  56. products_df = pd.read_sql("SELECT * FROM products", engine)
  57. orders_df = pd.read_sql("SELECT * FROM orders", engine)
  58. SANDBOX_GLOBALS = {
  59. # 数据:Agent 可以分析这三张表
  60. "employees_df": employees_df,
  61. "products_df": products_df,
  62. "orders_df": orders_df,
  63. # ⼯具库:Agent 可以⽤这些库做分析和画图
  64. "pd": pd, # pandas — 数据处理
  65. "plt": plt, # matplotlib — 基础绑图
  66. "sns": sns, # seaborn — 统计可视化
  67. "np": np, # numpy — 数值计算
  68. }
  69. toolkit = SQLDatabaseToolkit(db=db, llm=llm)
  70. sql_tool = toolkit.get_tools()
  71. @tool
  72. def plot_tool(code):
  73. """
  74. 执⾏ Python 代码进⾏数据分析和可视化。
  75. 可⽤变量:
  76. - employees_df: 员⼯表 DataFrame(字段:id, name, department, salary,
  77. hire_date)
  78. - products_df: 产品表 DataFrame(字段:id, product_name, category, pr
  79. ice, stock)
  80. - orders_df: 订单表 DataFrame(字段:id, employee_id, product_id, q
  81. uantity, order_date)
  82. - pd: pandas 库
  83. - plt: matplotlib.pyplot
  84. - sns: seaborn
  85. - np: numpy
  86. 使⽤示例:
  87. # 统计各部⻔平均薪资
  88. result = employees_df.groupby('department')['salary'].mean()
  89. print(result)
  90. # 画柱状图
  91. plt.figure(figsize=(10, 6))
  92. employees_df.groupby('department')['salary'].mean().plot(kind='bar') plt.title('各部⻔平均薪资')
  93. plt.show()
  94. """
  95. # 准备隔离的执⾏环境
  96. # globals_dict 提供⽩名单变量,locals_dict 收集执⾏过程中产⽣的新变量
  97. exec_globals = dict(SANDBOX_GLOBALS)
  98. exec_locals = {}
  99. # ⽤ StringIO 捕获 print() 的输出
  100. # 这样 Agent ⽣成的代码⾥所有的 print 语句都会被收集
  101. output_buffer = StringIO()
  102. try:
  103. # redirect_stdout 会把标准输出重定向到我们的 buffer
  104. with redirect_stdout(output_buffer):
  105. exec(code, exec_globals, exec_locals)
  106. result = output_buffer.getvalue()
  107. # 如果代码没有 print 任何东⻄,给个默认提示
  108. if not result.strip():
  109. result = "✅ 代码执⾏成功(⽆⽂本输出,可能已⽣成图表)"
  110. return f"执⾏成功:\n{result}"
  111. except Exception as e:
  112. # 出错时返回完整的错误堆栈,⽅便 Agent ⾃我修正
  113. error_detail = traceback.format_exc()
  114. return f"❌ 执⾏出错:{e}\n\n{error_detail}"
  115. VISUALIZATION_PROMPT = """
  116. # 角色与目标
  117. 你是一个智能数据分析助手。你的核心职责是根据用户的自然语言描述,准确判断需要调用哪些工具(SQL查询工具、Python绘图工具)来完成用户的任务,并按照严格的格式输出调用指令。
  118. # 可用工具定义
  119. 1. **SQL查询工具 (`sql_tool`)**:用于执行SQL语句,从数据库中获取原始数据。
  120. 2. **Python绘图工具 (`plot_tool`)**:用于执行Python代码,根据数据生成柱状图、折线图、饼图等可视化图表。
  121. # 决策逻辑与判断规则
  122. 请严格按照以下优先级和条件进行判断,不要自行臆测用户未提及的需求:
  123. 1. **情况一:需要“查询数据”且“绘制图表”**
  124. - **触发条件**:用户描述中**同时包含**“获取/查询/查找数据”和“画图/图表/可视化/柱状/折线”等意图。
  125. - **执行动作**:**必须**先调用 `sql_tool`工具获取数据结果,再根据数据结果调用 `plot_tool`工具画图。
  126. 2. **情况二:仅需要“绘制图表”**
  127. - **触发条件**:用户描述中**只**包含“画图/图表/可视化”等意图,且**未提及**需要从数据库查询新数据(例如用户已提供了数据,或要求基于已有数据绘图)。
  128. - **执行动作**:**仅**调用 `plot_tool`工具画图。
  129. 3. **情况三:仅需要“查询数据”**
  130. - **触发条件**:用户描述中**只**包含“查询/查找/获取数据”等意图,且**未提及**任何画图或可视化的需求。
  131. - **执行动作**:**仅**调用 `sql_tool`工具获取数据结果。
  132. 4. **情况四:无需调用任何工具**
  133. - **触发条件**:用户描述与数据查询、图表绘制均无关(例如:闲聊、问好、咨询业务定义、询问天气等);或者用户的需求模糊不清,无法明确判断需要哪个工具。
  134. - **执行动作**:**不调用**任何工具,直接回复用户说明无法处理或进行澄清。
  135. # plot_tool工具提示词
  136. 你是⼀名资深数据分析师,精通 Python、Pandas 和 Matplo tlib 数据可视化。
  137. ## 可⽤数据
  138. 1. employees_df — 员⼯表(字段:id, name, department, salary, hire_date)
  139. 2. products_df — 产品表(字段:id, product_name, category, price, stock)
  140. 3. orders_df — 订单表(字段:id, employee_id, product_id, quantity, order
  141. _date)
  142. ## ⼯作流程
  143. 1. 理解⽤户的分析需求
  144. 2. ⽤ execute_python_code ⼯具编写并执⾏ Python 代码
  145. 3. 先做数据探索(head、describe、info),再做深⼊分析
  146. 4. ⽤中⽂解释分析结果,给出业务洞察
  147. ## 代码规范
  148. - 绑图前设置中⽂字体:plt.rcParams['font.sans-serif'] = ['SimHei', 'PingFang
  149. SC', 'DejaVu Sans']
  150. - 设置 plt.rcParams['axes.unicode_minus'] = False
  151. - 图表尺⼨统⼀⽤ plt.figure(figsize=(10, 6))
  152. - 必须添加标题、坐标轴标签,让图表⾃解释
  153. - ⽤ print() 输出关键统计量,不要只画图不说话
  154. - 图表标题⽤英⽂(避免渲染问题),但⽤中⽂向⽤户解释结果
  155. ## 注意事项
  156. - 每次只执⾏⼀段完整的代码,不要拆成多段
  157. - 先探索数据结构,再做分析——不要上来就画图
  158. - 结果要有业务洞察,不只是"最⼤值是 XXX"
  159. """
  160. sql_plot_tools = sql_tool + [plot_tool]
  161. sql_py_agent = create_agent(
  162. model=llm,
  163. tools=sql_plot_tools,
  164. system_prompt=VISUALIZATION_PROMPT,
  165. )
  166. print("==================正常开始================")
  167. def do_sql_py_agent(user_query):
  168. data_content = sql_py_agent.invoke({"messages": [{"role": "user", "content": user_query}]})
  169. return data_content["messages"][-1].content
  170. print(do_sql_py_agent("你是谁?"))
  171. print(do_sql_py_agent("帮我查一下一共多少名员工?"))
  172. print(do_sql_py_agent("帮我画一个员工销售分布图"))
  173. print("==================正常结束================")