ZhouYI 2 mesiacov pred
commit
9ff4e42bae

+ 1 - 0
01_dataAnalysis/agent/__init__.py

@@ -0,0 +1 @@
+

+ 49 - 0
01_dataAnalysis/agent/chain.py

@@ -0,0 +1,49 @@
+from langchain.agents import create_agent
+
+from agent.tools.python_tool import execute_python_code
+from agent.tools.sql_tool import ask_database, get_table_schema, list_tables
+from agent.tools.time_tool import get_current_time
+from agent.tools.wheather_tool import get_weather
+
+from .config import AppConfig
+from .llm import create_llm
+
+
+def create_chat_agent(config: AppConfig):
+    llm = create_llm(config)
+
+    return create_agent(
+        model=llm,
+        tools=[
+            get_current_time,
+            get_weather,
+            list_tables,
+            get_table_schema,
+            ask_database,
+            execute_python_code,
+        ],
+        system_prompt="""
+你是一个耐心、清晰的 Python 和 Agent 开发学习助手。
+
+你可以使用工具获取当前时间、实时天气、查询数据库,也可以执行 Python
+代码做数据分析和图表可视化。
+
+当用户询问数据库、表结构、SQL 查询、数据统计或数据分析时:
+1. 先调用 list_tables 查看可用表;
+2. 再调用 get_table_schema 确认相关表的字段和类型;
+3. 然后生成只读 SQL,不要凭空编造表名和字段名;
+4. 简单查询或验证 SQL 时,可以调用 ask_database。
+
+当用户要求基于数据库生成图表、统计图、可视化分析或统计表时:
+1. 先用 list_tables 和 get_table_schema 确认表结构;
+2. 再调用 execute_python_code;
+3. 在 Python 代码里用 read_sql("SELECT ...") 把数据库数据读取成 DataFrame;
+4. 使用 pandas 做统计分析,使用 print() 输出关键统计量或统计表;
+5. 使用 matplotlib 绘图,图表标题建议使用英文;
+6. 最终回答中用中文解释分析结果,并给出图表保存路径。
+
+代码和 SQL 都应以只读分析为主,不要修改数据库或本地文件。
+如果用户的问题缺少必要信息,先追问,不要随便猜测。
+回答使用中文,表达简洁清楚。
+""",
+    )

+ 23 - 0
01_dataAnalysis/agent/cli.py

@@ -0,0 +1,23 @@
+from .chain import create_chat_agent
+from .config import load_config
+
+
+def run_chat() -> None:
+    config = load_config()
+    agent = create_chat_agent(config)
+    run_config = {"configurable": {"thread_id": config.default_session_id}}
+
+    while True:
+        question = input("请输入: ").strip()
+        if not question:
+            continue
+
+        if question.lower() in ["exit", "quit", "q"]:
+            break
+
+        result = agent.invoke(
+            {"messages": [{"role": "user", "content": question}]},
+            config=run_config,
+        )
+        answer = result["messages"][-1].content
+        print(f"AI: {answer}")

+ 30 - 0
01_dataAnalysis/agent/config.py

@@ -0,0 +1,30 @@
+import os
+from dataclasses import dataclass
+
+
+@dataclass(frozen=True)
+class AppConfig:
+    api_key: str
+    api_url: str
+    model: str
+    database_url: str
+    history_table_name: str
+    default_session_id: str
+
+
+def load_config() -> AppConfig:
+    api_key = os.getenv("API_KEY")
+    if not api_key:
+        raise ValueError("API_KEY is not set in environment variables.")
+
+    return AppConfig(
+        api_key=api_key,
+        api_url=os.getenv(
+            "API_URL",
+            "https://dashscope.aliyuncs.com/compatible-mode/v1",
+        ),
+        model=os.getenv("MODEL", "qwen3.6-plus"),
+        database_url=os.getenv("DATABASE_URL", "mysql+pymysql://root:1234@localhost:3306/mydb"),
+        history_table_name=os.getenv("HISTORY_TABLE_NAME", "chat_history"),
+        default_session_id=os.getenv("SESSION_ID", "user_001"),
+    )

+ 11 - 0
01_dataAnalysis/agent/history.py

@@ -0,0 +1,11 @@
+from langchain_community.chat_message_histories import SQLChatMessageHistory
+
+from .config import AppConfig
+
+
+def create_message_history(config: AppConfig, session_id: str) -> SQLChatMessageHistory:
+    return SQLChatMessageHistory(
+        session_id=session_id,
+        connection=config.database_url,
+        table_name=config.history_table_name,
+    )

+ 12 - 0
01_dataAnalysis/agent/llm.py

@@ -0,0 +1,12 @@
+from langchain_openai import ChatOpenAI
+
+from .config import AppConfig
+
+
+def create_llm(config: AppConfig) -> ChatOpenAI:
+    return ChatOpenAI(
+        model_name=config.model,
+        api_key=config.api_key,
+        base_url=config.api_url,
+        temperature=0.7,
+    )

+ 110 - 0
01_dataAnalysis/agent/tools/python_tool.py

@@ -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 "代码执行完成,但没有输出文本或图表。"

+ 84 - 0
01_dataAnalysis/agent/tools/sql_tool.py

@@ -0,0 +1,84 @@
+from functools import lru_cache
+
+from langchain_community.utilities import SQLDatabase
+from langchain_core.tools import tool
+
+from agent.config import load_config
+
+
+@lru_cache(maxsize=1)
+def get_database() -> SQLDatabase:
+    config = load_config()
+    return SQLDatabase.from_uri(config.database_url)
+
+
+def is_readonly_sql(sql: str) -> bool:
+    cleaned_sql = sql.strip().rstrip(";").strip()
+    lowered_sql = cleaned_sql.lower()
+
+    readonly_prefixes = ("select", "show", "describe", "desc", "explain")
+    forbidden_keywords = (
+        "insert",
+        "update",
+        "delete",
+        "drop",
+        "alter",
+        "create",
+        "truncate",
+        "replace",
+        "grant",
+        "revoke",
+    )
+
+    if not lowered_sql.startswith(readonly_prefixes):
+        return False
+
+    if ";" in cleaned_sql:
+        return False
+
+    return not any(f" {keyword} " in f" {lowered_sql} " for keyword in forbidden_keywords)
+
+
+@tool
+def list_tables() -> str:
+    """
+    列出当前数据库中可以使用的所有表名。
+
+    当用户询问数据库里有哪些表、需要做数据分析、需要查询某个业务数据,
+    或者你不确定应该查询哪张表时,应该先调用这个工具。
+    这个工具不需要参数,只返回表名列表,不会查询表里的具体数据。
+    """
+    db = get_database()
+    table_names = db.get_usable_table_names()
+    return "\n".join(table_names) if table_names else "当前数据库中没有可用的数据表。"
+
+
+@tool
+def get_table_schema(table_name: str) -> str:
+    """
+    查看指定数据表的字段结构、字段类型和部分样例信息。
+
+    当你准备生成 SQL 之前,应该先调用这个工具确认表结构,
+    不要凭空猜测字段名。参数 table_name 必须是数据库中真实存在的表名,
+    例如:chat_history。一次只传入一个表名。
+    """
+    db = get_database()
+    return db.get_table_info([table_name])
+
+
+@tool
+def ask_database(sql: str) -> str:
+    """
+    执行只读 SQL 查询并返回查询结果。
+
+    当你已经知道要查询的表和字段后,调用这个工具执行 SQL。
+    参数 sql 必须是一条完整的只读 SQL 语句,只允许 SELECT、SHOW、
+    DESCRIBE、DESC、EXPLAIN。不要传入 INSERT、UPDATE、DELETE、DROP、
+    ALTER、CREATE 等会修改数据库的语句。查询数据时建议使用 LIMIT 50
+    限制返回行数,避免一次返回太多数据。
+    """
+    if not is_readonly_sql(sql):
+        return "拒绝执行:只允许单条只读 SQL,例如 SELECT、SHOW、DESCRIBE、EXPLAIN。"
+
+    db = get_database()
+    return db.run(sql)

+ 14 - 0
01_dataAnalysis/agent/tools/time_tool.py

@@ -0,0 +1,14 @@
+from datetime import datetime
+
+from langchain_core.tools import tool
+
+
+@tool
+def get_current_time() -> str:
+    """
+    获取当前本地日期和时间。
+
+    当用户询问现在几点、今天几号、当前时间、当前日期等问题时,
+    应该调用这个工具。这个工具不需要参数。
+    """
+    return datetime.now().strftime("%Y-%m-%d %H:%M:%S")

+ 23 - 0
01_dataAnalysis/agent/tools/wheather_tool.py

@@ -0,0 +1,23 @@
+import requests
+from langchain_core.tools import tool
+
+
+@tool
+def get_weather(location: str) -> str:
+    """
+    查询指定城市或地区的实时天气。
+
+    当用户询问某个城市、地区或国家的天气、温度、气候状况时,
+    应该调用这个工具。参数 location 是地点名称,例如:北京、上海、
+    成都、Tokyo、New York。如果用户没有提供地点,应该先追问地点。
+    """
+    url = f"https://wttr.in/{location}?format=j1&lang=zh"
+    response = requests.get(url, timeout=10)
+    response.raise_for_status()
+
+    data = response.json()
+    current = data["current_condition"][0]
+    weather_desc = current["weatherDesc"][0]["value"]
+    temp = current["temp_C"]
+
+    return f"{location} 当前天气:{weather_desc},温度:{temp}°C"

+ 4 - 0
01_dataAnalysis/main.py

@@ -0,0 +1,4 @@
+from agent.cli import run_chat
+
+if __name__ == "__main__":
+    run_chat()