| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252 |
- """
- 方案 B:统一 Agent 模式
- 把 SQL 查询工具与 Python 代码执行工具挂到同一个 Agent 上,
- 让 LLM 自己决定何时查数据库、何时写 Python 分析/画图。
- """
- import os
- import traceback
- import warnings
- from io import StringIO
- from contextlib import redirect_stdout
- import pandas as pd
- import numpy as np
- import matplotlib.pyplot as plt
- import seaborn as sns
- from sqlalchemy import create_engine
- from langchain_openai import ChatOpenAI
- from langchain_community.utilities import SQLDatabase
- from langchain_community.agent_toolkits import SQLDatabaseToolkit
- from langchain.tools import tool
- from langchain.agents import create_agent
- import os
- from dotenv import load_dotenv
- load_dotenv()
- warnings.filterwarnings("ignore")
- # ============================================================
- # 第一步:配置环境与大模型
- # ============================================================
- # ── Environment ────────────────────────────────────────────────────────────
- AliYunKey = os.getenv("ALIYUN_API_KEY")
- ALIYUN_BASE_URL = os.getenv("ALIYUN_BASE_URL")
- ALIYUN_CHAT_MODEL = os.getenv("ALIYUN_CHAT_MODEL")
- ALIYUN_EMBEDDING_MODEL = os.getenv("ALIYUN_EMBEDDING_MODEL")
- DB_HOST = os.getenv("DB_HOST", "127.0.0.1")
- DB_PORT = int(os.getenv("DB_PORT", "3306"))
- DB_USER = os.getenv("DB_USER", "root")
- DB_PASSWORD = os.getenv("DB_PASSWORD", "root")
- DB_NAME = os.getenv("DB_NAME", "dyxz")
- # 数据库连接 URI(SQLAlchemy 格式)
- DB_URI = f"mysql+pymysql://{DB_USER}:{DB_PASSWORD}@{DB_HOST}:{DB_PORT}/{DB_NAME}"
- llm = ChatOpenAI(
- base_url=ALIYUN_BASE_URL, api_key=AliYunKey,
- model=ALIYUN_CHAT_MODEL, temperature=0, timeout=10, max_retries=1,
- )
- print("✅ 模型初始化完成")
- # ============================================================
- # 第二步:连接数据库并加载数据到 Pandas
- # ============================================================
- db = SQLDatabase.from_uri(DB_URI)
- engine = create_engine(DB_URI)
- print(f"✅ 数据库连接成功")
- print(f" 可用表:{db.get_usable_table_names()}")
- 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)
- print(f"✅ 数据加载完成")
- print(f" employees:{len(employees_df)} 行")
- print(f" products :{len(products_df)} 行")
- print(f" orders :{len(orders_df)} 行")
- # 配置 matplotlib 中文显示
- plt.rcParams["font.sans-serif"] = ["SimHei", "PingFang SC", "DejaVu Sans"]
- plt.rcParams["axes.unicode_minus"] = False
- # ============================================================
- # 第三步:准备工具
- # ============================================================
- # 3.1 SQL 工具包
- sql_toolkit = SQLDatabaseToolkit(db=db, llm=llm)
- sql_tools = sql_toolkit.get_tools()
- print(f"\n✅ SQL 工具包已加载({len(sql_tools)} 个工具):")
- for t in sql_tools:
- print(f" - {t.name}")
- # 3.2 Python 代码执行沙箱
- SANDBOX_GLOBALS = {
- "employees_df": employees_df,
- "products_df": products_df,
- "orders_df": orders_df,
- "pd": pd,
- "plt": plt,
- "sns": sns,
- "np": np,
- }
- @tool
- def execute_python_code(code: str) -> str:
- """
- 执行 Python 代码进行数据分析和可视化。
- 可用变量:
- - employees_df: 员工表 DataFrame
- - products_df: 产品表 DataFrame
- - orders_df: 订单表 DataFrame
- - pd, plt, sns, np
- 使用示例:
- 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('Average Salary by Department')
- plt.show()
- """
- exec_globals = dict(SANDBOX_GLOBALS)
- exec_locals = {}
- output_buffer = StringIO()
- try:
- with redirect_stdout(output_buffer):
- exec(code, exec_globals, exec_locals)
- result = output_buffer.getvalue()
- if not result.strip():
- result = "✅ 代码执行成功(无文本输出,可能已生成图表)"
- return f"执行成功:\n{result}"
- except Exception as e:
- error_detail = traceback.format_exc()
- return f"❌ 执行出错:{e}\n\n{error_detail}"
- print("\n✅ Python 代码执行沙箱创建成功")
- # 合并所有工具
- all_tools = sql_tools + [execute_python_code]
- print(f"\n✅ 统一 Agent 共挂载 {len(all_tools)} 个工具")
- # ============================================================
- # 第四步:定义统一 Agent 的 System Prompt
- # ============================================================
- UNIFIED_AGENT_PROMPT = """你是一名全能的数据分析 Agent,同时具备两种核心能力:
- 1. SQL 数据库查询:通过 SQL 工具查询 MySQL 数据库中的 employees、products、orders 表。
- 2. Python 数据分析与可视化:通过 execute_python_code 工具编写并执行 Python 代码。
- ## 数据库表结构
- - employees(员工表):id, name, department, salary, hire_date
- - products(产品表):id, product_name, category, price, stock
- - orders(订单表):id, employee_id, product_id, quantity, order_date
- ## 已加载到内存的 DataFrame
- - employees_df, products_df, orders_df(字段与数据库表一致)
- ## 可用工具
- """ + "\n".join([f"- {t.name}: {t.description.split(chr(10))[0] if t.description else 'No description'}" for t in all_tools]) + """
- ## 工作流程
- 1. 先判断用户问题更适合用 SQL 查询,还是更适合用 Python 分析/可视化。
- 2. 如果需要查数据:先用 sql_db_list_tables / sql_db_schema 了解表结构,再生成并执行 SQL。
- 3. 如果需要分析或画图:用 execute_python_code 编写 Python 代码,可先用 head/describe 探索数据。
- 4. 复杂问题可组合使用:先 SQL 查数,再 Python 分析/画图。
- 5. 最后用中文给出简洁、专业的业务洞察。
- ## 代码规范(使用 execute_python_code 时)
- - 绑图前设置中文字体:plt.rcParams['font.sans-serif'] = ['SimHei', 'PingFang SC', 'DejaVu Sans']
- - 设置 plt.rcParams['axes.unicode_minus'] = False
- - 图表尺寸统一用 plt.figure(figsize=(10, 6))
- - 图表标题用英文(避免渲染问题),但向用户解释时用中文
- - 用 print() 输出关键统计量
- ## 约束
- - 只使用数据库中实际存在的表和字段
- - SQL 单次查询结果限制在 50 条以内
- - 如果执行出错,分析原因后修正并重新执行
- - 回答简洁专业,不要啰嗦
- """
- # ============================================================
- # 第五步:组装统一 Agent
- # ============================================================
- unified_agent = create_agent(
- model=llm,
- tools=all_tools,
- system_prompt=UNIFIED_AGENT_PROMPT,
- )
- print("\n🚀 统一 Agent 创建完成,可以开始提问!")
- # ============================================================
- # 第六步:封装一个友好的调用接口
- # ============================================================
- def ask_agent(query: str, verbose: bool = True):
- """
- 向统一 Agent 提问,并返回最终回答。
- Args:
- query: 用户的自然语言问题
- verbose: 是否打印 Agent 的完整思考过程
- Returns:
- 最终回答文本
- """
- print(f"\n{'='*60}")
- print(f"用户问题:{query}")
- print(f"{'='*60}")
- response = unified_agent.invoke({
- "messages": [{"role": "user", "content": query}]
- })
- if verbose:
- print("\n--- Agent 思考过程 ---")
- for i, msg in enumerate(response["messages"]):
- msg_type = msg.__class__.__name__
- if hasattr(msg, "tool_calls") and msg.tool_calls:
- print(f"\n步骤 {i+1} [{msg_type}]:")
- for tc in msg.tool_calls:
- print(f" 调用工具:{tc['name']}")
- args = tc.get("args", {})
- if "code" in args:
- print(f" 参数(代码):\n{args['code'][:500]}{'...' if len(args['code']) > 500 else ''}")
- else:
- print(f" 参数:{args}")
- elif msg_type == "ToolMessage":
- content = msg.content
- print(f" 工具返回:{content[:300]}{'...' if len(content) > 300 else ''}")
- final_msg = response["messages"][-1]
- print(f"\n--- 最终回答 ---")
- print(final_msg.content)
- return final_msg.content
- # ============================================================
- # 第七步:示例提问
- # ============================================================
- if __name__ == "__main__":
- # 示例 1:纯 SQL 查询类问题
- ask_agent("公司一共有多少名员工?每个部门各有多少人?")
- # 示例 2:纯 Python 可视化类问题
- ask_agent("画一个柱状图,展示各部门的平均薪资对比,并在柱子上标注数值")
- # 示例 3:需要先 SQL 查数、再 Python 分析/可视化的混合问题
- ask_agent("查询每个产品类别的总销售额,并用水平条形图展示")
|