Explorar el Código

Initial commit: data analyst project

linruying hace 2 meses
commit
d7bb405479

+ 23 - 0
.gitignore

@@ -0,0 +1,23 @@
+# Python
+__pycache__/
+*.py[cod]
+*.egg-info/
+dist/
+build/
+.venv/
+venv/
+
+# Environment
+.env
+.env.local
+
+# IDE
+.vscode/
+.idea/
+
+# Jupyter
+.ipynb_checkpoints/
+
+# OS
+Thumbs.db
+.DS_Store

+ 2 - 0
01_dataanalyst/agents/__init__.py

@@ -0,0 +1,2 @@
+from .sql_agent import sql_agent
+from .viz_agent import visualization_agent

+ 50 - 0
01_dataanalyst/agents/sql_agent.py

@@ -0,0 +1,50 @@
+"""
+NL2SQL Agent:将自然语言转换为 SQL 查询,自动执行并返回分析结果。
+"""
+
+import os
+import sys
+
+# 确保父目录可导入(支持直接运行 python agents/sql_agent.py)
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+
+from langchain_community.agent_toolkits import SQLDatabaseToolkit
+from langchain.agents import create_agent
+from config import llm, db
+
+# ---- 创建 SQL 工具包 ----
+toolkit = SQLDatabaseToolkit(db=db, llm=llm)
+tools = toolkit.get_tools()
+
+# ---- System Prompt ----
+SQL_AGENT_PROMPT = """你是一名专业的 SQL 数据分析师。
+
+## 工作流程
+1. 先用 sql_db_list_tables 查看数据库中有哪些表
+2. 用 sql_db_schema 获取相关表的字段结构和类型
+3. 生成 SQL 之前,用 sql_db_query_checker 检查语法
+4. 确认无误后,用 sql_db_query 执行查询
+5. 用中文总结查询结果,给出简洁的业务洞察
+
+## 约束
+- 只使用数据库中实际存在的表和字段,不要凭空编造
+- 单次查询结果限制在 50 条以内
+- 如果查询出错,分析错误原因后重新生成 SQL
+- 回答要简洁专业,不要啰嗦
+"""
+
+# ---- 创建 Agent ----
+sql_agent = create_agent(
+    model=llm,
+    tools=tools,
+    system_prompt=SQL_AGENT_PROMPT,
+)
+
+
+if __name__ == "__main__":
+    print(f"📊 数据库:{db}")
+    print(f"   可用表:{db.get_usable_table_names()}")
+    print(f"\nSQL 工具包已加载({len(tools)} 个工具):")
+    for t in tools:
+        print(f"   - {t.name}")
+    print("\n✅ NL2SQL Agent 创建完成,可以开始提问了!")

+ 138 - 0
01_dataanalyst/agents/viz_agent.py

@@ -0,0 +1,138 @@
+"""
+数据可视化 Agent:用 Python 代码对 DataFrame 做统计分析和绑图。
+"""
+
+import os
+import sys
+import traceback
+import warnings
+from io import StringIO
+from contextlib import redirect_stdout
+
+# 确保父目录可导入(支持直接运行 python agents/viz_agent.py)
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+
+import pandas as pd
+import matplotlib.pyplot as plt
+import seaborn as sns
+import numpy as np
+from langchain.tools import tool
+from langchain.agents import create_agent
+
+from config import llm, engine
+
+warnings.filterwarnings("ignore")
+
+# ============================================================
+# 1. 中文字体配置
+# ============================================================
+plt.rcParams["font.sans-serif"] = ["SimHei", "PingFang SC", "DejaVu Sans"]
+plt.rcParams["axes.unicode_minus"] = False
+
+# ============================================================
+# 2. 延迟加载 DataFrame(避免 import 时就连库)
+# ============================================================
+_dataframes = None
+
+
+def load_dataframes():
+    """从数据库加载 DataFrame。首次调用后缓存,后续直接返回缓存。"""
+    global _dataframes
+    if _dataframes is None:
+        _dataframes = {
+            "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),
+        }
+    return _dataframes
+
+
+# ============================================================
+# 3. 代码执行沙箱
+# ============================================================
+@tool
+def execute_python_code(code: str) -> str:
+    """
+    执行 Python 代码进行数据分析和可视化。
+
+    可用变量:
+      - employees_df: 员工表 DataFrame(字段:id, name, department, salary, hire_date)
+      - products_df:  产品表 DataFrame(字段:id, product_name, category, price, stock)
+      - orders_df:    订单表 DataFrame(字段:id, employee_id, product_id, quantity, order_date)
+      - pd:   pandas 库
+      - plt:  matplotlib.pyplot
+      - sns:  seaborn
+      - np:   numpy
+
+    使用示例:
+      result = employees_df.groupby('department')['salary'].mean()
+      print(result)
+    """
+    dfs = load_dataframes()
+    sandbox = {
+        **dfs,
+        "pd":  pd,
+        "plt": plt,
+        "sns": sns,
+        "np":  np,
+    }
+    exec_locals = {}
+    output_buffer = StringIO()
+
+    try:
+        with redirect_stdout(output_buffer):
+            exec(code, sandbox, exec_locals)
+        result = output_buffer.getvalue()
+        if not result.strip():
+            result = "✅ 代码执行成功(无文本输出,可能已生成图表)"
+        return f"执行成功:\n{result}"
+    except Exception as e:
+        return f"❌ 执行出错:{e}\n\n{traceback.format_exc()}"
+
+
+# ============================================================
+# 4. Agent System Prompt
+# ============================================================
+VISUALIZATION_PROMPT = """你是一名资深数据分析师,精通 Python、Pandas 和 Matplotlib 数据可视化。
+
+## 可用数据
+1. employees_df — 员工表(字段:id, name, department, salary, hire_date)
+2. products_df  — 产品表(字段:id, product_name, category, price, stock)
+3. orders_df    — 订单表(字段:id, employee_id, product_id, quantity, order_date)
+
+## 工作流程
+1. 理解用户的分析需求
+2. 用 execute_python_code 工具编写并执行 Python 代码
+3. 先做数据探索(head、describe、info),再做深入分析
+4. 用中文解释分析结果,给出业务洞察
+
+## 代码规范
+- 绑图前设置中文字体:plt.rcParams['font.sans-serif'] = ['SimHei', 'PingFang SC', 'DejaVu Sans']
+- 设置 plt.rcParams['axes.unicode_minus'] = False
+- 图表尺寸统一用 plt.figure(figsize=(10, 6))
+- 必须添加标题、坐标轴标签,让图表自解释
+- 用 print() 输出关键统计量,不要只画图不说话
+- 图表标题用英文(避免渲染问题),但用中文向用户解释结果
+
+## 注意事项
+- 每次只执行一段完整的代码,不要拆成多段
+- 先探索数据结构,再做分析——不要上来就画图
+- 结果要有业务洞察,不只是"最大值是 XXX"
+"""
+
+# ============================================================
+# 5. 创建 Agent
+# ============================================================
+visualization_agent = create_agent(
+    model=llm,
+    tools=[execute_python_code],
+    system_prompt=VISUALIZATION_PROMPT,
+)
+
+
+if __name__ == "__main__":
+    dfs = load_dataframes()
+    print("✅ 数据加载完成")
+    for name, df in dfs.items():
+        print(f"   {name}:{len(df)} 行 × {len(df.columns)} 列")
+    print("\n✅ 数据可视化 Agent 创建完成")

+ 49 - 0
01_dataanalyst/config.py

@@ -0,0 +1,49 @@
+"""
+统一配置模块:环境变量、数据库连接、LLM 实例。
+
+其他模块直接 `from config import llm, db, engine` 即可,
+无需重复写 load_dotenv 和 db_uri 构造逻辑。
+"""
+
+import os
+from dotenv import load_dotenv
+from langchain_openai import ChatOpenAI
+from langchain_community.utilities import SQLDatabase
+from sqlalchemy import create_engine
+
+# ============================================================
+# 1. 加载 .env(整个项目只调用一次)
+# ============================================================
+load_dotenv(os.path.join(os.path.dirname(__file__), "..", ".env"))
+
+# ============================================================
+# 2. 数据库配置
+# ============================================================
+DB_HOST = os.getenv("DB_HOST", "127.0.0.1")
+DB_PORT = os.getenv("DB_PORT", "3306")
+DB_USER = os.getenv("DB_USER", "root")
+DB_PASSWORD = os.getenv("DB_PASSWORD", "")
+DB_NAME = os.getenv("DB_NAME", "analytics_demo")
+
+DB_URI = f"mysql+pymysql://{DB_USER}:{DB_PASSWORD}@{DB_HOST}:{DB_PORT}/{DB_NAME}"
+
+# SQLAlchemy 引擎(DataFrame 加载用)
+engine = create_engine(DB_URI)
+
+# LangChain SQLDatabase 实例(Agent 工具包用)
+db = SQLDatabase.from_uri(DB_URI)
+
+# ============================================================
+# 3. LLM 配置
+# ============================================================
+DEEPSEEK_API_KEY = os.getenv("DEEPSEEK_API_KEY")
+DEEPSEEK_BASE_URL = os.getenv("DEEPSEEK_BASE_URL")
+
+llm = ChatOpenAI(
+    model="deepseek-v4-flash",
+    api_key=DEEPSEEK_API_KEY,
+    base_url=DEEPSEEK_BASE_URL,
+    temperature=0,
+)
+
+print("✅ 配置加载完成(LLM + 数据库连接)")

+ 95 - 0
01_dataanalyst/database_build.py

@@ -0,0 +1,95 @@
+"""
+数据库初始化脚本:创建 analytics_demo 库和表结构,插入演示数据。
+
+使用方式:python database_build.py
+"""
+
+import mysql.connector
+from config import DB_HOST, DB_PORT, DB_USER, DB_PASSWORD, DB_NAME
+
+
+def init_database():
+    """创建数据库、表结构并插入演示数据。"""
+    # 先连接 MySQL(不指定数据库),确保目标库存在
+    conn = mysql.connector.connect(
+        host=DB_HOST,
+        port=int(DB_PORT),
+        user=DB_USER,
+        password=DB_PASSWORD,
+    )
+    cursor = conn.cursor()
+    cursor.execute(f"CREATE DATABASE IF NOT EXISTS {DB_NAME} DEFAULT CHARSET utf8mb4")
+    conn.database = DB_NAME
+
+    # ---- 建表 ----
+    cursor.execute("""
+        CREATE TABLE IF NOT EXISTS employees (
+            id          INT PRIMARY KEY AUTO_INCREMENT,
+            name        VARCHAR(50)  NOT NULL,
+            department  VARCHAR(50)  NOT NULL,
+            salary      DECIMAL(10,2) NOT NULL,
+            hire_date   DATE
+        ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4;
+    """)
+
+    cursor.execute("""
+        CREATE TABLE IF NOT EXISTS products (
+            id           INT PRIMARY KEY AUTO_INCREMENT,
+            product_name VARCHAR(100) NOT NULL,
+            category     VARCHAR(50)  NOT NULL,
+            price        DECIMAL(10,2) NOT NULL,
+            stock        INT DEFAULT 0
+        ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4;
+    """)
+
+    cursor.execute("""
+        CREATE TABLE IF NOT EXISTS orders (
+            id           INT PRIMARY KEY AUTO_INCREMENT,
+            employee_id  INT NOT NULL,
+            product_id   INT NOT NULL,
+            quantity     INT NOT NULL,
+            order_date   DATE NOT NULL,
+            FOREIGN KEY (employee_id) REFERENCES employees(id),
+            FOREIGN KEY (product_id)  REFERENCES products(id)
+        ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4;
+    """)
+
+    # ---- 插入演示数据 ----
+    employees_data = [
+        (1, "张三", "技术部",   20000.00, "2023-01-15"),
+        (2, "李四", "销售部",   11000.00, "2023-02-20"),
+        (3, "王五", "技术部",   16000.00, "2022-11-10"),
+        (4, "赵六", "人力资源", 5000.00, "2023-03-01"),
+        (5, "钱七", "销售部",   17000.00, "2022-12-05"),
+    ]
+
+    products_data = [
+        (1, "笔记本电脑", "电子产品", 6999.00, 500),
+        (2, "机械键盘",   "电子产品", 399.00,  1000),
+        (3, "办公椅",     "办公用品", 499.00,  300),
+        (4, "显示器",     "电子产品", 1200.00, 400),
+    ]
+
+    orders_data = [
+        (1, 1, 1, 2,  "2024-01-15"),
+        (2, 2, 2, 15, "2024-01-16"),
+        (3, 3, 1, 10, "2024-01-17"),
+        (4, 5, 3, 6,  "2024-01-18"),
+        (5, 2, 4, 5,  "2024-01-19"),
+    ]
+
+    cursor.executemany("INSERT IGNORE INTO employees VALUES (%s,%s,%s,%s,%s)", employees_data)
+    cursor.executemany("INSERT IGNORE INTO products  VALUES (%s,%s,%s,%s,%s)", products_data)
+    cursor.executemany("INSERT IGNORE INTO orders    VALUES (%s,%s,%s,%s,%s)", orders_data)
+
+    conn.commit()
+    conn.close()
+
+    print("✅ 数据库初始化完成")
+    print("   - employees 表:5 条员工记录")
+    print("   - products  表:4 条产品记录")
+    print("   - orders    表:5 条订单记录")
+
+
+if __name__ == "__main__":
+    init_database()

+ 48 - 0
01_dataanalyst/main.py

@@ -0,0 +1,48 @@
+"""
+数据分析入口:根据用户问题自动路由到 SQL Agent 或可视化 Agent。
+"""
+
+from config import llm
+from agents import sql_agent, visualization_agent
+
+
+def classify_intent(user_query: str) -> str:
+    """用 LLM 判断用户意图:query / visualize / both。"""
+    prompt = f"""判断以下用户问题属于哪种类型:
+    - "query":需要查询数据库获取数据
+    - "visualize":需要画图或做统计分析
+    - "both":需要先查数据,再画图分析
+
+    用户问题:{user_query}
+
+    只回复一个词:query / visualize / both"""
+
+    return llm.invoke(prompt).content.strip().lower()
+
+
+def run_data_analysis(user_query: str) -> str:
+    """统一入口:自动判断意图并路由到对应 Agent。"""
+    intent = classify_intent(user_query)
+
+    if intent == "query":
+        result = sql_agent.invoke({"messages": [{"role": "user", "content": user_query}]})
+        return result["messages"][-1].content
+
+    elif intent == "visualize":
+        result = visualization_agent.invoke({"messages": [{"role": "user", "content": user_query}]})
+        return result["messages"][-1].content
+
+    else:  # both
+        data_result = sql_agent.invoke({"messages": [{"role": "user", "content": user_query}]})
+        data_content = data_result["messages"][-1].content
+        viz_input = f"基于以下数据进行可视化分析:\n{data_content}\n\n原始问题:{user_query}"
+        viz_result = visualization_agent.invoke({"messages": [{"role": "user", "content": viz_input}]})
+        return viz_result["messages"][-1].content
+
+
+if __name__ == "__main__":
+    user_query = """
+    列出所有2023年入职的员工,并列出他们的所属部门、姓名和薪资,按薪资从高到低排序?并画柱状图展示
+    """
+    result = run_data_analysis(user_query)
+    print(result)

+ 219 - 0
test.ipynb

@@ -0,0 +1,219 @@
+{
+ "cells": [
+  {
+   "cell_type": "code",
+   "execution_count": 11,
+   "id": "3a522fca",
+   "metadata": {},
+   "outputs": [
+    {
+     "name": "stdout",
+     "output_type": "stream",
+     "text": [
+      "你问的这句话“chromaDB返回相似度分数的计算机制”,意思是:**ChromaDB(一个向量数据库)在返回查询结果时,里面附带的那个“相似度分数”到底是怎么算出来的?** 换句话说,是问背后的数学原理或算法。\n",
+      "\n",
+      "### 它在做什么?\n",
+      "\n",
+      "让我拆开来说:\n",
+      "\n",
+      "1. **ChromaDB 是什么?**  \n",
+      "   它是一个专门存储和搜索“向量”的数据库。向量可以理解为一段数字列表(比如 `[0.1, 0.5, -0.2]`),用来表示文字、图片、音频等内容的“语义特征”。\n",
+      "\n",
+      "2. **返回相似度分数**  \n",
+      "   当你向 ChromaDB 提交一个查询(比如一段文本或一个向量),它会在库中找出最接近的几个向量,并给每个结果附上一个分数,比如 `0.92`。这个分数表示查询向量和库中某个向量有多“像”。\n",
+      "\n",
+      "3. **计算机制**  \n",
+      "   就是指这个分数是通过**哪种距离或相似度度量**算出来的。常用的方法有:\n",
+      "   - **余弦相似度**:看两个向量的方向是否一致(取值范围 -1 ~ 1,越大越相似)。  \n",
+      "   - **欧几里得距离**:看两个向量在多维空间中的直线距离(越小越相似)。  \n",
+      "   - **点积**:简单的内积(适用于某些已归一化的场景)。\n",
+      "\n",
+      "   ChromaDB 默认使用 **余弦相似度**,但你也可以手动指定其他距离函数。\n",
+      "\n",
+      "### 总结一句话\n",
+      "\n",
+      "这句话是在询问:**当你用 ChromaDB 做相似度检索时,它内部用的是什么公式/算法,来计算那个“相似度分数”的?** 换句话说,就是搞清楚结果里的数字(比如 0.95)到底代表什么物理意义(是角度相似还是距离远近)。\n"
+     ]
+    }
+   ],
+   "source": [
+    "from langchain_openai import ChatOpenAI\n",
+    "\n",
+    "# 创建 DeepSeek 聊天模型实例\n",
+    "# base_url 指向 DeepSeek 的兼容端点,而非 OpenAI 官方地址\n",
+    "llm = ChatOpenAI(\n",
+    "    model_name=\"deepseek-v4-flash\",              # DeepSeek 的对话模型\n",
+    "    api_key=\"sk-8fd808112bf7488ab27dacf0cad09315\",                  # 在 platform.deepseek.com 获取\n",
+    "    base_url=\"https://api.deepseek.com\"      # DeepSeek API 地址\n",
+    ")\n",
+    "\n",
+    "query_stm = \"\"\"\n",
+    "chromaDB返回相似度分数的计算机制\n",
+    "\"\"\"\n",
+    "# invoke 是 LangChain 统一的调用方法,返回 AIMessage 对象\n",
+    "response = llm.invoke(f\"解释{query_stm}是什么意思,在做什么?\")\n",
+    "print(response.content)  # .content 拿到纯文本"
+   ]
+  },
+  {
+   "cell_type": "code",
+   "execution_count": null,
+   "id": "0f77c487",
+   "metadata": {},
+   "outputs": [],
+   "source": []
+  },
+  {
+   "cell_type": "code",
+   "execution_count": 7,
+   "id": "237fc3ca",
+   "metadata": {},
+   "outputs": [
+    {
+     "name": "stdout",
+     "output_type": "stream",
+     "text": [
+      "片名:绿皮书\n",
+      "年份:2018\n",
+      "导演:彼得·法雷里\n",
+      "评分:8.2\n",
+      "主角名字:托尼·利普\n"
+     ]
+    }
+   ],
+   "source": [
+    "from pydantic import BaseModel, Field\n",
+    "from langchain.chat_models import init_chat_model\n",
+    "\n",
+    "# 定义你期望的输出结构(Pydantic 模型)\n",
+    "class MovieInfo(BaseModel):\n",
+    "    \"\"\"电影信息\"\"\"\n",
+    "    title: str = Field(description=\"电影名称\")\n",
+    "    year: int = Field(description=\"上映年份\")\n",
+    "    director: str = Field(description=\"导演\")\n",
+    "    rating: float = Field(description=\"评分(10分制)\")\n",
+    "    character_name: str = Field(description=\"主角名字\")\n",
+    "\n",
+    "llm = init_chat_model(\n",
+    "    model=\"qwen-plus\",\n",
+    "    model_provider=\"openai\",\n",
+    "    api_key=\"sk-1283efb80348448180e87b3c56fb6f3d\",\n",
+    "    base_url=\"https://dashscope.aliyuncs.com/compatible-mode/v1\"\n",
+    ")\n",
+    "\n",
+    "# .with_structured_output() 会自动把 Schema 注入 prompt,\n",
+    "# 并在底层做 JSON 解析和类型校验\n",
+    "structured_llm = llm.with_structured_output(MovieInfo)\n",
+    "\n",
+    "result = structured_llm.invoke(\"介绍一下电影《绿皮书》\")\n",
+    "\n",
+    "# 返回的是 MovieInfo 对象,可以直接用属性访问\n",
+    "print(f\"片名:{result.title}\")\n",
+    "print(f\"年份:{result.year}\")\n",
+    "print(f\"导演:{result.director}\")\n",
+    "print(f\"评分:{result.rating}\")\n",
+    "print(f\"主角名字:{result.character_name}\")"
+   ]
+  },
+  {
+   "cell_type": "code",
+   "execution_count": 8,
+   "id": "cc05cc04",
+   "metadata": {},
+   "outputs": [
+    {
+     "name": "stdout",
+     "output_type": "stream",
+     "text": [
+      "雪花扇\n"
+     ]
+    }
+   ],
+   "source": [
+    "from typing_extensions import TypedDict\n",
+    "\n",
+    "# 方式二:TypedDict(更轻量,无运行时校验)\n",
+    "class MovieTypedDict(TypedDict):\n",
+    "    title: str\n",
+    "    year: int\n",
+    "    director: str\n",
+    "    rating: float\n",
+    "\n",
+    "structured_llm = llm.with_structured_output(MovieTypedDict)\n",
+    "result = structured_llm.invoke(\"介绍一下电影《雪花扇》\")\n",
+    "# 返回的是普通 dict\n",
+    "print(result[\"title\"])"
+   ]
+  },
+  {
+   "cell_type": "code",
+   "execution_count": 6,
+   "id": "3af58865",
+   "metadata": {
+    "vscode": {
+     "languageId": "javascript"
+    }
+   },
+   "outputs": [
+    {
+     "name": "stdout",
+     "output_type": "stream",
+     "text": [
+      "{'director': '饺子', 'rating': 8.4, 'title': '哪吒之魔童降世', 'year': 2019}\n"
+     ]
+    }
+   ],
+   "source": [
+    "json_schema = {\n",
+    "    \"title\": \"MovieInfo\",\n",
+    "    \"description\": \"电影信息对象\",\n",
+    "    \"type\": \"object\",\n",
+    "    \"properties\": {\n",
+    "        \"title\":    {\"type\": \"string\",  \"description\": \"电影名称\"},\n",
+    "        \"year\":     {\"type\": \"integer\", \"description\": \"上映年份\"},\n",
+    "        \"director\": {\"type\": \"string\",  \"description\": \"导演\"},\n",
+    "        \"rating\":   {\"type\": \"number\",  \"description\": \"评分(10分制)\"}\n",
+    "    },\n",
+    "    \"required\": [\"title\", \"year\", \"director\", \"rating\"]\n",
+    "}\n",
+    "\n",
+    "structured_llm = llm.with_structured_output(json_schema)\n",
+    "result = structured_llm.invoke(\"介绍一下电影《哪吒之魔童降世》\")\n",
+    "print(result)  # 返回的是 dict"
+   ]
+  },
+  {
+   "cell_type": "code",
+   "execution_count": null,
+   "id": "09fdc151",
+   "metadata": {
+    "vscode": {
+     "languageId": "javascript"
+    }
+   },
+   "outputs": [],
+   "source": []
+  }
+ ],
+ "metadata": {
+  "kernelspec": {
+   "display_name": "01_langchain (3.14.3)",
+   "language": "python",
+   "name": "python3"
+  },
+  "language_info": {
+   "codemirror_mode": {
+    "name": "ipython",
+    "version": 3
+   },
+   "file_extension": ".py",
+   "mimetype": "text/x-python",
+   "name": "python",
+   "nbconvert_exporter": "python",
+   "pygments_lexer": "ipython3",
+   "version": "3.14.3"
+  }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 5
+}