""" 方案 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("查询每个产品类别的总销售额,并用水平条形图展示")