|
@@ -0,0 +1,110 @@
|
|
|
|
|
+import contextlib
|
|
|
|
|
+import io
|
|
|
|
|
+import re
|
|
|
|
|
+from pathlib import Path
|
|
|
|
|
+from uuid import uuid4
|
|
|
|
|
+
|
|
|
|
|
+import matplotlib
|
|
|
|
|
+
|
|
|
|
|
+matplotlib.use("Agg")
|
|
|
|
|
+
|
|
|
|
|
+import matplotlib.pyplot as plt
|
|
|
|
|
+import numpy as np
|
|
|
|
|
+import pandas as pd
|
|
|
|
|
+from langchain_core.tools import tool
|
|
|
|
|
+from sqlalchemy import create_engine
|
|
|
|
|
+
|
|
|
|
|
+from agent.config import load_config
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+OUTPUT_DIR = Path("outputs/charts")
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def _read_sql(sql: str) -> pd.DataFrame:
|
|
|
|
|
+ config = load_config()
|
|
|
|
|
+ engine = create_engine(config.database_url)
|
|
|
|
|
+ with engine.connect() as connection:
|
|
|
|
|
+ return pd.read_sql(sql, connection)
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def _is_safe_python_code(code: str) -> bool:
|
|
|
|
|
+ lowered_code = code.lower()
|
|
|
|
|
+ forbidden_patterns = [
|
|
|
|
|
+ r"\bimport\s+os\b",
|
|
|
|
|
+ r"\bimport\s+sys\b",
|
|
|
|
|
+ r"\bimport\s+subprocess\b",
|
|
|
|
|
+ r"\bfrom\s+os\b",
|
|
|
|
|
+ r"\bfrom\s+sys\b",
|
|
|
|
|
+ r"\bfrom\s+subprocess\b",
|
|
|
|
|
+ r"\bopen\s*\(",
|
|
|
|
|
+ r"\beval\s*\(",
|
|
|
|
|
+ r"\bexec\s*\(",
|
|
|
|
|
+ r"\bcompile\s*\(",
|
|
|
|
|
+ r"__",
|
|
|
|
|
+ r"\bdelete\b",
|
|
|
|
|
+ r"\bdrop\b",
|
|
|
|
|
+ r"\bupdate\b",
|
|
|
|
|
+ r"\binsert\b",
|
|
|
|
|
+ r"\balter\b",
|
|
|
|
|
+ r"\btruncate\b",
|
|
|
|
|
+ ]
|
|
|
|
|
+ return not any(re.search(pattern, lowered_code) for pattern in forbidden_patterns)
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+@tool
|
|
|
|
|
+def execute_python_code(code: str) -> str:
|
|
|
|
|
+ """
|
|
|
|
|
+ 执行一段用于数据分析和图表可视化的 Python 代码,并返回 print 输出和图表文件路径。
|
|
|
|
|
+
|
|
|
|
|
+ 当用户要求做数据探索、统计分析、用 Pandas 处理数据、绘制柱状图、折线图、
|
|
|
|
|
+ 饼图、散点图、直方图等可视化结果时,应该调用这个工具。
|
|
|
|
|
+
|
|
|
|
|
+ 参数 code 必须是一段完整 Python 代码。代码中可以直接使用 pd、np、plt、
|
|
|
|
|
+ read_sql(sql) 和 database_url。需要数据库数据时,先用 read_sql("SELECT ...")
|
|
|
|
|
+ 读取为 DataFrame。绘图时请使用 matplotlib,并用 print() 输出关键统计量。
|
|
|
|
|
+ 工具会自动把所有 matplotlib 图表保存为 PNG 文件,不需要调用 plt.show()。
|
|
|
|
|
+
|
|
|
|
|
+ 安全要求:只做数据读取、分析和可视化,不要读写本地文件,不要执行系统命令,
|
|
|
|
|
+ 不要生成 INSERT、UPDATE、DELETE、DROP、ALTER 等修改数据库的 SQL。
|
|
|
|
|
+ """
|
|
|
|
|
+ if not _is_safe_python_code(code):
|
|
|
|
|
+ return "拒绝执行:代码包含潜在危险操作,只允许数据分析和可视化代码。"
|
|
|
|
|
+
|
|
|
|
|
+ config = load_config()
|
|
|
|
|
+ OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
|
|
|
|
+ plt.close("all")
|
|
|
|
|
+
|
|
|
|
|
+ namespace = {
|
|
|
|
|
+ "pd": pd,
|
|
|
|
|
+ "np": np,
|
|
|
|
|
+ "plt": plt,
|
|
|
|
|
+ "read_sql": _read_sql,
|
|
|
|
|
+ "database_url": config.database_url,
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ stdout = io.StringIO()
|
|
|
|
|
+ try:
|
|
|
|
|
+ with contextlib.redirect_stdout(stdout):
|
|
|
|
|
+ exec(code, {"__builtins__": __builtins__}, namespace)
|
|
|
|
|
+ except Exception as exc:
|
|
|
|
|
+ plt.close("all")
|
|
|
|
|
+ return f"代码执行失败:{type(exc).__name__}: {exc}"
|
|
|
|
|
+
|
|
|
|
|
+ chart_paths = []
|
|
|
|
|
+ for figure_number in plt.get_fignums():
|
|
|
|
|
+ figure = plt.figure(figure_number)
|
|
|
|
|
+ chart_path = OUTPUT_DIR / f"chart_{uuid4().hex}.png"
|
|
|
|
|
+ figure.tight_layout()
|
|
|
|
|
+ figure.savefig(chart_path, dpi=150, bbox_inches="tight")
|
|
|
|
|
+ chart_paths.append(str(chart_path))
|
|
|
|
|
+
|
|
|
|
|
+ plt.close("all")
|
|
|
|
|
+
|
|
|
|
|
+ output = stdout.getvalue().strip()
|
|
|
|
|
+ result_parts = []
|
|
|
|
|
+ if output:
|
|
|
|
|
+ result_parts.append(f"代码输出:\n{output}")
|
|
|
|
|
+ if chart_paths:
|
|
|
|
|
+ result_parts.append("图表已保存:\n" + "\n".join(chart_paths))
|
|
|
|
|
+
|
|
|
|
|
+ return "\n\n".join(result_parts) if result_parts else "代码执行完成,但没有输出文本或图表。"
|