|
@@ -0,0 +1,201 @@
|
|
|
|
|
+"""
|
|
|
|
|
+1. 加载配置
|
|
|
|
|
+2. 创建大模型
|
|
|
|
|
+3. 配置并获取db,db数据,db工具
|
|
|
|
|
+4. 配置沙箱
|
|
|
|
|
+5. 创建python代码画图工具
|
|
|
|
|
+6. 创建智能体
|
|
|
|
|
+7. 使用智能体
|
|
|
|
|
+"""
|
|
|
|
|
+
|
|
|
|
|
+from dotenv import load_dotenv
|
|
|
|
|
+import os
|
|
|
|
|
+from langchain_openai import ChatOpenAI
|
|
|
|
|
+from langchain.agents import create_agent
|
|
|
|
|
+from langchain.tools import tool
|
|
|
|
|
+import traceback
|
|
|
|
|
+from io import StringIO
|
|
|
|
|
+from contextlib import redirect_stdout
|
|
|
|
|
+from langchain.tools import tool
|
|
|
|
|
+import pandas as pd
|
|
|
|
|
+import matplotlib.pyplot as plt
|
|
|
|
|
+import seaborn as sns
|
|
|
|
|
+import numpy as np
|
|
|
|
|
+from langchain_community.utilities import SQLDatabase
|
|
|
|
|
+from langchain_community.agent_toolkits import SQLDatabaseToolkit
|
|
|
|
|
+from sqlalchemy import create_engine
|
|
|
|
|
+
|
|
|
|
|
+load_dotenv(override=True)
|
|
|
|
|
+
|
|
|
|
|
+deepseek_base_url = os.getenv("DEEPSEEK_BASE_URL")
|
|
|
|
|
+deepseek_base_key = os.getenv("DEEPSEEK_BASE_KEY")
|
|
|
|
|
+deepseek_base_name = os.getenv("DEEPSEEK_BASE_NAME")
|
|
|
|
|
+print(deepseek_base_key)
|
|
|
|
|
+print(deepseek_base_name)
|
|
|
|
|
+print(deepseek_base_url)
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+llm = ChatOpenAI(
|
|
|
|
|
+ base_url = deepseek_base_url,
|
|
|
|
|
+ api_key = deepseek_base_key,
|
|
|
|
|
+ model= deepseek_base_name
|
|
|
|
|
+)
|
|
|
|
|
+
|
|
|
|
|
+db_uri = (
|
|
|
|
|
+ f"mysql+pymysql://{os.getenv('DB_USER')}:{os.getenv('DB_PASSWORD')}"
|
|
|
|
|
+ f"@{os.getenv('DB_HOST')}:{os.getenv('DB_PORT')}/{os.getenv('DB_NAME')}"
|
|
|
|
|
+)
|
|
|
|
|
+db = SQLDatabase.from_uri(db_uri)
|
|
|
|
|
+
|
|
|
|
|
+# ============================================================
|
|
|
|
|
+# 配置 matplotlib 中⽂显示
|
|
|
|
|
+# ============================================================
|
|
|
|
|
+# Windows ⽤ SimHei(⿊体),macOS ⽤ PingFang SC
|
|
|
|
|
+# 如果还是乱码,试试安装 fonts-noto-cjk 并清除缓存
|
|
|
|
|
+plt.rcParams["font.sans-serif"] = ["SimHei", "PingFang SC", "DejaVu Sans"]
|
|
|
|
|
+plt.rcParams["axes.unicode_minus"] = False # 解决负号显示为⽅块的问题
|
|
|
|
|
+# ============================================================
|
|
|
|
|
+# 从数据库加载数据到 Pandas DataFrame
|
|
|
|
|
+# ============================================================
|
|
|
|
|
+# ⽤ SQLAlchemy engine 复⽤连接,避免重复创建连接池
|
|
|
|
|
+
|
|
|
|
|
+engine = create_engine(db_uri)
|
|
|
|
|
+employees_df = pd.read_sql("SELECT * FROM employees", engine)
|
|
|
|
|
+products_df = pd.read_sql("SELECT * FROM products", engine)
|
|
|
|
|
+orders_df = pd.read_sql("SELECT * FROM orders", engine)
|
|
|
|
|
+
|
|
|
|
|
+SANDBOX_GLOBALS = {
|
|
|
|
|
+ # 数据:Agent 可以分析这三张表
|
|
|
|
|
+ "employees_df": employees_df,
|
|
|
|
|
+ "products_df": products_df,
|
|
|
|
|
+ "orders_df": orders_df,
|
|
|
|
|
+ # ⼯具库:Agent 可以⽤这些库做分析和画图
|
|
|
|
|
+ "pd": pd, # pandas — 数据处理
|
|
|
|
|
+ "plt": plt, # matplotlib — 基础绑图
|
|
|
|
|
+ "sns": sns, # seaborn — 统计可视化
|
|
|
|
|
+ "np": np, # numpy — 数值计算
|
|
|
|
|
+}
|
|
|
|
|
+
|
|
|
|
|
+toolkit = SQLDatabaseToolkit(db=db, llm=llm)
|
|
|
|
|
+sql_tool = toolkit.get_tools()
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+@tool
|
|
|
|
|
+def plot_tool(code):
|
|
|
|
|
+ """
|
|
|
|
|
+ 执⾏ Python 代码进⾏数据分析和可视化。
|
|
|
|
|
+ 可⽤变量:
|
|
|
|
|
+ - employees_df: 员⼯表 DataFrame(字段:id, name, department, salary,
|
|
|
|
|
+ hire_date)
|
|
|
|
|
+ - products_df: 产品表 DataFrame(字段:id, product_name, category, pr
|
|
|
|
|
+ ice, stock)
|
|
|
|
|
+ - orders_df: 订单表 DataFrame(字段:id, employee_id, product_id, q
|
|
|
|
|
+ uantity, order_date)
|
|
|
|
|
+ - pd: pandas 库
|
|
|
|
|
+ - plt: matplotlib.pyplot
|
|
|
|
|
+ - sns: seaborn
|
|
|
|
|
+ - np: numpy
|
|
|
|
|
+ 使⽤示例:
|
|
|
|
|
+ # 统计各部⻔平均薪资
|
|
|
|
|
+ result = employees_df.groupby('department')['salary'].mean()
|
|
|
|
|
+ print(result)
|
|
|
|
|
+ # 画柱状图
|
|
|
|
|
+ plt.figure(figsize=(10, 6))
|
|
|
|
|
+ employees_df.groupby('department')['salary'].mean().plot(kind='bar') plt.title('各部⻔平均薪资')
|
|
|
|
|
+ plt.show()
|
|
|
|
|
+ """
|
|
|
|
|
+ # 准备隔离的执⾏环境
|
|
|
|
|
+ # globals_dict 提供⽩名单变量,locals_dict 收集执⾏过程中产⽣的新变量
|
|
|
|
|
+ exec_globals = dict(SANDBOX_GLOBALS)
|
|
|
|
|
+ exec_locals = {}
|
|
|
|
|
+ # ⽤ StringIO 捕获 print() 的输出
|
|
|
|
|
+ # 这样 Agent ⽣成的代码⾥所有的 print 语句都会被收集
|
|
|
|
|
+ output_buffer = StringIO()
|
|
|
|
|
+ try:
|
|
|
|
|
+ # redirect_stdout 会把标准输出重定向到我们的 buffer
|
|
|
|
|
+ with redirect_stdout(output_buffer):
|
|
|
|
|
+ exec(code, exec_globals, exec_locals)
|
|
|
|
|
+ result = output_buffer.getvalue()
|
|
|
|
|
+ # 如果代码没有 print 任何东⻄,给个默认提示
|
|
|
|
|
+ if not result.strip():
|
|
|
|
|
+ result = "✅ 代码执⾏成功(⽆⽂本输出,可能已⽣成图表)"
|
|
|
|
|
+ return f"执⾏成功:\n{result}"
|
|
|
|
|
+ except Exception as e:
|
|
|
|
|
+ # 出错时返回完整的错误堆栈,⽅便 Agent ⾃我修正
|
|
|
|
|
+ error_detail = traceback.format_exc()
|
|
|
|
|
+ return f"❌ 执⾏出错:{e}\n\n{error_detail}"
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+VISUALIZATION_PROMPT = """
|
|
|
|
|
+ # 角色与目标
|
|
|
|
|
+ 你是一个智能数据分析助手。你的核心职责是根据用户的自然语言描述,准确判断需要调用哪些工具(SQL查询工具、Python绘图工具)来完成用户的任务,并按照严格的格式输出调用指令。
|
|
|
|
|
+
|
|
|
|
|
+ # 可用工具定义
|
|
|
|
|
+ 1. **SQL查询工具 (`sql_tool`)**:用于执行SQL语句,从数据库中获取原始数据。
|
|
|
|
|
+ 2. **Python绘图工具 (`plot_tool`)**:用于执行Python代码,根据数据生成柱状图、折线图、饼图等可视化图表。
|
|
|
|
|
+
|
|
|
|
|
+ # 决策逻辑与判断规则
|
|
|
|
|
+ 请严格按照以下优先级和条件进行判断,不要自行臆测用户未提及的需求:
|
|
|
|
|
+
|
|
|
|
|
+ 1. **情况一:需要“查询数据”且“绘制图表”**
|
|
|
|
|
+ - **触发条件**:用户描述中**同时包含**“获取/查询/查找数据”和“画图/图表/可视化/柱状/折线”等意图。
|
|
|
|
|
+ - **执行动作**:**必须**先调用 `sql_tool`工具获取数据结果,再根据数据结果调用 `plot_tool`工具画图。
|
|
|
|
|
+
|
|
|
|
|
+ 2. **情况二:仅需要“绘制图表”**
|
|
|
|
|
+ - **触发条件**:用户描述中**只**包含“画图/图表/可视化”等意图,且**未提及**需要从数据库查询新数据(例如用户已提供了数据,或要求基于已有数据绘图)。
|
|
|
|
|
+ - **执行动作**:**仅**调用 `plot_tool`工具画图。
|
|
|
|
|
+
|
|
|
|
|
+ 3. **情况三:仅需要“查询数据”**
|
|
|
|
|
+ - **触发条件**:用户描述中**只**包含“查询/查找/获取数据”等意图,且**未提及**任何画图或可视化的需求。
|
|
|
|
|
+ - **执行动作**:**仅**调用 `sql_tool`工具获取数据结果。
|
|
|
|
|
+
|
|
|
|
|
+ 4. **情况四:无需调用任何工具**
|
|
|
|
|
+ - **触发条件**:用户描述与数据查询、图表绘制均无关(例如:闲聊、问好、咨询业务定义、询问天气等);或者用户的需求模糊不清,无法明确判断需要哪个工具。
|
|
|
|
|
+ - **执行动作**:**不调用**任何工具,直接回复用户说明无法处理或进行澄清。
|
|
|
|
|
+
|
|
|
|
|
+ # plot_tool工具提示词
|
|
|
|
|
+ 你是⼀名资深数据分析师,精通 Python、Pandas 和 Matplo tlib 数据可视化。
|
|
|
|
|
+ ## 可⽤数据
|
|
|
|
|
+ 1. employees_df — 员⼯表(字段:id, name, department, salary, hire_date)
|
|
|
|
|
+ 2. products_df — 产品表(字段:id, product_name, category, price, stock)
|
|
|
|
|
+ 3. orders_df — 订单表(字段:id, employee_id, product_id, quantity, order
|
|
|
|
|
+ _date)
|
|
|
|
|
+ ## ⼯作流程
|
|
|
|
|
+ 1. 理解⽤户的分析需求
|
|
|
|
|
+ 2. ⽤ execute_python_code ⼯具编写并执⾏ Python 代码
|
|
|
|
|
+ 3. 先做数据探索(head、describe、info),再做深⼊分析
|
|
|
|
|
+ 4. ⽤中⽂解释分析结果,给出业务洞察
|
|
|
|
|
+ ## 代码规范
|
|
|
|
|
+ - 绑图前设置中⽂字体:plt.rcParams['font.sans-serif'] = ['SimHei', 'PingFang
|
|
|
|
|
+ SC', 'DejaVu Sans']
|
|
|
|
|
+ - 设置 plt.rcParams['axes.unicode_minus'] = False
|
|
|
|
|
+ - 图表尺⼨统⼀⽤ plt.figure(figsize=(10, 6))
|
|
|
|
|
+ - 必须添加标题、坐标轴标签,让图表⾃解释
|
|
|
|
|
+ - ⽤ print() 输出关键统计量,不要只画图不说话
|
|
|
|
|
+ - 图表标题⽤英⽂(避免渲染问题),但⽤中⽂向⽤户解释结果
|
|
|
|
|
+ ## 注意事项
|
|
|
|
|
+ - 每次只执⾏⼀段完整的代码,不要拆成多段
|
|
|
|
|
+ - 先探索数据结构,再做分析——不要上来就画图
|
|
|
|
|
+ - 结果要有业务洞察,不只是"最⼤值是 XXX"
|
|
|
|
|
+"""
|
|
|
|
|
+
|
|
|
|
|
+sql_plot_tools = sql_tool + [plot_tool]
|
|
|
|
|
+
|
|
|
|
|
+sql_py_agent = create_agent(
|
|
|
|
|
+ model=llm,
|
|
|
|
|
+ tools=sql_plot_tools,
|
|
|
|
|
+ system_prompt=VISUALIZATION_PROMPT,
|
|
|
|
|
+)
|
|
|
|
|
+print("==================正常开始================")
|
|
|
|
|
+def do_sql_py_agent(user_query):
|
|
|
|
|
+ data_content = sql_py_agent.invoke({"messages": [{"role": "user", "content": user_query}]})
|
|
|
|
|
+ return data_content["messages"][-1].content
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+print(do_sql_py_agent("你是谁?"))
|
|
|
|
|
+print(do_sql_py_agent("帮我查一下一共多少名员工?"))
|
|
|
|
|
+print(do_sql_py_agent("帮我画一个员工销售分布图"))
|
|
|
|
|
+print("==================正常结束================")
|
|
|
|
|
+
|