unified_agent.py 9.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252
  1. """
  2. 方案 B:统一 Agent 模式
  3. 把 SQL 查询工具与 Python 代码执行工具挂到同一个 Agent 上,
  4. 让 LLM 自己决定何时查数据库、何时写 Python 分析/画图。
  5. """
  6. import os
  7. import traceback
  8. import warnings
  9. from io import StringIO
  10. from contextlib import redirect_stdout
  11. import pandas as pd
  12. import numpy as np
  13. import matplotlib.pyplot as plt
  14. import seaborn as sns
  15. from sqlalchemy import create_engine
  16. from langchain_openai import ChatOpenAI
  17. from langchain_community.utilities import SQLDatabase
  18. from langchain_community.agent_toolkits import SQLDatabaseToolkit
  19. from langchain.tools import tool
  20. from langchain.agents import create_agent
  21. import os
  22. from dotenv import load_dotenv
  23. load_dotenv()
  24. warnings.filterwarnings("ignore")
  25. # ============================================================
  26. # 第一步:配置环境与大模型
  27. # ============================================================
  28. # ── Environment ────────────────────────────────────────────────────────────
  29. AliYunKey = os.getenv("ALIYUN_API_KEY")
  30. ALIYUN_BASE_URL = os.getenv("ALIYUN_BASE_URL")
  31. ALIYUN_CHAT_MODEL = os.getenv("ALIYUN_CHAT_MODEL")
  32. ALIYUN_EMBEDDING_MODEL = os.getenv("ALIYUN_EMBEDDING_MODEL")
  33. DB_HOST = os.getenv("DB_HOST", "127.0.0.1")
  34. DB_PORT = int(os.getenv("DB_PORT", "3306"))
  35. DB_USER = os.getenv("DB_USER", "root")
  36. DB_PASSWORD = os.getenv("DB_PASSWORD", "root")
  37. DB_NAME = os.getenv("DB_NAME", "dyxz")
  38. # 数据库连接 URI(SQLAlchemy 格式)
  39. DB_URI = f"mysql+pymysql://{DB_USER}:{DB_PASSWORD}@{DB_HOST}:{DB_PORT}/{DB_NAME}"
  40. llm = ChatOpenAI(
  41. base_url=ALIYUN_BASE_URL, api_key=AliYunKey,
  42. model=ALIYUN_CHAT_MODEL, temperature=0, timeout=10, max_retries=1,
  43. )
  44. print("✅ 模型初始化完成")
  45. # ============================================================
  46. # 第二步:连接数据库并加载数据到 Pandas
  47. # ============================================================
  48. db = SQLDatabase.from_uri(DB_URI)
  49. engine = create_engine(DB_URI)
  50. print(f"✅ 数据库连接成功")
  51. print(f" 可用表:{db.get_usable_table_names()}")
  52. employees_df = pd.read_sql("SELECT * FROM employees", engine)
  53. products_df = pd.read_sql("SELECT * FROM products", engine)
  54. orders_df = pd.read_sql("SELECT * FROM orders", engine)
  55. print(f"✅ 数据加载完成")
  56. print(f" employees:{len(employees_df)} 行")
  57. print(f" products :{len(products_df)} 行")
  58. print(f" orders :{len(orders_df)} 行")
  59. # 配置 matplotlib 中文显示
  60. plt.rcParams["font.sans-serif"] = ["SimHei", "PingFang SC", "DejaVu Sans"]
  61. plt.rcParams["axes.unicode_minus"] = False
  62. # ============================================================
  63. # 第三步:准备工具
  64. # ============================================================
  65. # 3.1 SQL 工具包
  66. sql_toolkit = SQLDatabaseToolkit(db=db, llm=llm)
  67. sql_tools = sql_toolkit.get_tools()
  68. print(f"\n✅ SQL 工具包已加载({len(sql_tools)} 个工具):")
  69. for t in sql_tools:
  70. print(f" - {t.name}")
  71. # 3.2 Python 代码执行沙箱
  72. SANDBOX_GLOBALS = {
  73. "employees_df": employees_df,
  74. "products_df": products_df,
  75. "orders_df": orders_df,
  76. "pd": pd,
  77. "plt": plt,
  78. "sns": sns,
  79. "np": np,
  80. }
  81. @tool
  82. def execute_python_code(code: str) -> str:
  83. """
  84. 执行 Python 代码进行数据分析和可视化。
  85. 可用变量:
  86. - employees_df: 员工表 DataFrame
  87. - products_df: 产品表 DataFrame
  88. - orders_df: 订单表 DataFrame
  89. - pd, plt, sns, np
  90. 使用示例:
  91. result = employees_df.groupby('department')['salary'].mean()
  92. print(result)
  93. plt.figure(figsize=(10, 6))
  94. employees_df.groupby('department')['salary'].mean().plot(kind='bar')
  95. plt.title('Average Salary by Department')
  96. plt.show()
  97. """
  98. exec_globals = dict(SANDBOX_GLOBALS)
  99. exec_locals = {}
  100. output_buffer = StringIO()
  101. try:
  102. with redirect_stdout(output_buffer):
  103. exec(code, exec_globals, exec_locals)
  104. result = output_buffer.getvalue()
  105. if not result.strip():
  106. result = "✅ 代码执行成功(无文本输出,可能已生成图表)"
  107. return f"执行成功:\n{result}"
  108. except Exception as e:
  109. error_detail = traceback.format_exc()
  110. return f"❌ 执行出错:{e}\n\n{error_detail}"
  111. print("\n✅ Python 代码执行沙箱创建成功")
  112. # 合并所有工具
  113. all_tools = sql_tools + [execute_python_code]
  114. print(f"\n✅ 统一 Agent 共挂载 {len(all_tools)} 个工具")
  115. # ============================================================
  116. # 第四步:定义统一 Agent 的 System Prompt
  117. # ============================================================
  118. UNIFIED_AGENT_PROMPT = """你是一名全能的数据分析 Agent,同时具备两种核心能力:
  119. 1. SQL 数据库查询:通过 SQL 工具查询 MySQL 数据库中的 employees、products、orders 表。
  120. 2. Python 数据分析与可视化:通过 execute_python_code 工具编写并执行 Python 代码。
  121. ## 数据库表结构
  122. - employees(员工表):id, name, department, salary, hire_date
  123. - products(产品表):id, product_name, category, price, stock
  124. - orders(订单表):id, employee_id, product_id, quantity, order_date
  125. ## 已加载到内存的 DataFrame
  126. - employees_df, products_df, orders_df(字段与数据库表一致)
  127. ## 可用工具
  128. """ + "\n".join([f"- {t.name}: {t.description.split(chr(10))[0] if t.description else 'No description'}" for t in all_tools]) + """
  129. ## 工作流程
  130. 1. 先判断用户问题更适合用 SQL 查询,还是更适合用 Python 分析/可视化。
  131. 2. 如果需要查数据:先用 sql_db_list_tables / sql_db_schema 了解表结构,再生成并执行 SQL。
  132. 3. 如果需要分析或画图:用 execute_python_code 编写 Python 代码,可先用 head/describe 探索数据。
  133. 4. 复杂问题可组合使用:先 SQL 查数,再 Python 分析/画图。
  134. 5. 最后用中文给出简洁、专业的业务洞察。
  135. ## 代码规范(使用 execute_python_code 时)
  136. - 绑图前设置中文字体:plt.rcParams['font.sans-serif'] = ['SimHei', 'PingFang SC', 'DejaVu Sans']
  137. - 设置 plt.rcParams['axes.unicode_minus'] = False
  138. - 图表尺寸统一用 plt.figure(figsize=(10, 6))
  139. - 图表标题用英文(避免渲染问题),但向用户解释时用中文
  140. - 用 print() 输出关键统计量
  141. ## 约束
  142. - 只使用数据库中实际存在的表和字段
  143. - SQL 单次查询结果限制在 50 条以内
  144. - 如果执行出错,分析原因后修正并重新执行
  145. - 回答简洁专业,不要啰嗦
  146. """
  147. # ============================================================
  148. # 第五步:组装统一 Agent
  149. # ============================================================
  150. unified_agent = create_agent(
  151. model=llm,
  152. tools=all_tools,
  153. system_prompt=UNIFIED_AGENT_PROMPT,
  154. )
  155. print("\n🚀 统一 Agent 创建完成,可以开始提问!")
  156. # ============================================================
  157. # 第六步:封装一个友好的调用接口
  158. # ============================================================
  159. def ask_agent(query: str, verbose: bool = True):
  160. """
  161. 向统一 Agent 提问,并返回最终回答。
  162. Args:
  163. query: 用户的自然语言问题
  164. verbose: 是否打印 Agent 的完整思考过程
  165. Returns:
  166. 最终回答文本
  167. """
  168. print(f"\n{'='*60}")
  169. print(f"用户问题:{query}")
  170. print(f"{'='*60}")
  171. response = unified_agent.invoke({
  172. "messages": [{"role": "user", "content": query}]
  173. })
  174. if verbose:
  175. print("\n--- Agent 思考过程 ---")
  176. for i, msg in enumerate(response["messages"]):
  177. msg_type = msg.__class__.__name__
  178. if hasattr(msg, "tool_calls") and msg.tool_calls:
  179. print(f"\n步骤 {i+1} [{msg_type}]:")
  180. for tc in msg.tool_calls:
  181. print(f" 调用工具:{tc['name']}")
  182. args = tc.get("args", {})
  183. if "code" in args:
  184. print(f" 参数(代码):\n{args['code'][:500]}{'...' if len(args['code']) > 500 else ''}")
  185. else:
  186. print(f" 参数:{args}")
  187. elif msg_type == "ToolMessage":
  188. content = msg.content
  189. print(f" 工具返回:{content[:300]}{'...' if len(content) > 300 else ''}")
  190. final_msg = response["messages"][-1]
  191. print(f"\n--- 最终回答 ---")
  192. print(final_msg.content)
  193. return final_msg.content
  194. # ============================================================
  195. # 第七步:示例提问
  196. # ============================================================
  197. if __name__ == "__main__":
  198. # 示例 1:纯 SQL 查询类问题
  199. ask_agent("公司一共有多少名员工?每个部门各有多少人?")
  200. # 示例 2:纯 Python 可视化类问题
  201. ask_agent("画一个柱状图,展示各部门的平均薪资对比,并在柱子上标注数值")
  202. # 示例 3:需要先 SQL 查数、再 Python 分析/可视化的混合问题
  203. ask_agent("查询每个产品类别的总销售额,并用水平条形图展示")