소스 검색

secondt commit

clarknexus 1 개월 전
부모
커밋
f588ea2820

+ 752 - 0
001-agent/2.DeepSeek驱动的数据分析Agent实战.md

@@ -0,0 +1,752 @@
+# DeepSeek 驱动的数据分析 Agent 实战
+> 用大模型让数据"开口说话"——从架构设计到代码落地的完整指南
+>
+
+---
+
+## 一、为什么要做数据分析 Agent
+### 1.1 传统数据分析的痛点
+做过数据分析的同学都知道,日常工作中最折磨人的不是算法本身,而是**从"业务问题"到"拿到结果"这条路上的各种摩擦**:
+
+| 痛点 | 典型场景 | 真实代价 |
+| :--- | :--- | :--- |
+| SQL 门槛高 | 运营同学想知道"上个月华东区 GMV 多少",得找数据分析师写 SQL | 排队等半天,分析师还得反复确认口径 |
+| 代码成本大 | 画个趋势图要写一堆 matplotlib 样板代码 | 80% 时间在调样式,20% 在思考业务 |
+| 沟通损耗重 | "我要看环比" -> 分析师理解成"同比" -> 返工 | 需求理解偏差是最大的隐性成本 |
+| 重复劳动多 | 每周固定出几张报表,逻辑一模一样 | 机械重复,毫无技术含量 |
+
+
+核心矛盾很清晰:**业务人员懂业务但不会写代码,技术人员会写代码但不一定懂业务**。中间这条鸿沟,正是大模型 Agent 可以填平的地方。
+
+### 1.2 大模型能解决什么
+大语言模型(LLM)在数据分析场景中,核心能力可以拆成三层:
+
++ **自然语言理解**:把"帮我看看上个月各区域销售情况"解析成结构化意图——时间范围(上月)、分组维度(区域)、指标(销售额)。这一层纯靠 LLM 就能搞定。
++ **代码生成**:根据解析出的意图,自动生成 SQL 查询或 Python 分析代码。这也是 LLM 的强项,尤其是 DeepSeek-Coder 系列在代码生成上的表现相当能打。
++ **推理与规划**:面对复杂分析需求(比如"找出销量下滑的原因"),需要 Agent 自主拆解任务、规划执行步骤。这一层需要 LLM + 工具调用 + 业务上下文共同协作。
+
+一句话总结:**LLM 负责"理解"和"生成",Agent 框架负责"调度"和"执行"**。
+
+---
+
+## 二、系统架构总览
+### 2.1 整体架构
+整个系统由两个核心 Agent 协同工作,各司其职:
+
+<!-- 这是一张图片,ocr 内容为: -->
+![](系统架构图.png)
+
+> ▲ 系统架构:用户提问 → 编排器路由 → NL2SQL Agent 查数据 / 可视化 Agent 画图表
+>
+
+### 2.2 两个 Agent 的分工
+| 维度 | NL2SQL 查询 Agent | 数据可视化 Agent |
+| :--- | :--- | :--- |
+| **职责** | 把自然语言翻译成 SQL,查询数据库 | 把查询结果翻译成 Python 代码,生成图表 |
+| **核心工具** | LangChain SQLDatabaseToolkit | 自定义 Python 代码执行工具 |
+| **输入** | 用户的自然语言问题 | 查询返回的结构化数据 |
+| **输出** | 查询结果(表格数据) | 统计分析 + 可视化图表 |
+| **典型问题** | "技术部有多少人?" | "画一个各部门平均薪资的柱状图" |
+
+
+### 2.3 技术栈一览
+| 组件 | 技术选型 | 选型理由 |
+| :--- | :--- | :--- |
+| 大模型 | DeepSeek-Chat / DeepSeek-Coder | 国产头部,代码能力强,性价比高 |
+| Agent 框架 | LangChain 1.3.1 | 生态成熟,create_agent API 简洁可靠 |
+| 数据库 | MySQL 8.0+ | 通用性强,适合演示 |
+| 数据处理 | Pandas + NumPy | Python 数据分析标准库 |
+| 可视化 | Matplotlib + Seaborn | 灵活可控,适合 Agent 生成代码 |
+| 沙箱执行 | 自定义 exec 沙箱 | 轻量级,可控性强 |
+
+
+---
+
+## 三、技术选型与环境准备
+### 3.1 依赖安装
+```bash
+# 核心框架(锁定 langchain 1.3.1,API 稳定可靠)
+pip install langchain==1.3.1 langchain-openai langchain-community
+
+# 数据库驱动(MySQL)
+pip install mysql-connector-python pymysql
+
+# 数据处理与可视化
+pip install pandas numpy matplotlib seaborn
+```
+
+### 3.2 初始化大模型客户端
+```python
+import os
+from langchain_openai import ChatOpenAI
+
+# 通过环境变量读取 API Key,避免硬编码泄露风险
+llm = ChatOpenAI(
+    model="deepseek-v4-flash",                          # DeepSeek 的对话模型
+    api_key=os.getenv("DEEPSEEK_API_KEY"),          # 从环境变量加载
+    base_url="https://api.deepseek.com",            # DeepSeek API 地址
+    temperature=0,                                   # 设为 0 保证 SQL 生成的确定性
+)
+
+print("✅ DeepSeek 模型初始化完成")
+```
+
+> **为什么 **`temperature=0`**?** 生成 SQL 和分析代码属于确定性任务,不需要随机性。温度设为 0 可以让模型输出更稳定、可复现。
+>
+
+---
+
+## 四、数据层搭建:建表与灌数据
+### 4.1 表结构设计
+我们用一个**电商订单系统**作为演示场景,涉及三张核心表:
+
+<!-- 这是一张图片,ocr 内容为: -->
+![](ER关系图.png)
+
+### 4.2 建表与数据初始化
+```python
+import os
+import mysql.connector
+
+# ============================================================
+# 第一步:建立数据库连接
+# ============================================================
+# 使用环境变量管理敏感信息,这是生产级代码的基本规范
+conn = mysql.connector.connect(
+    host=os.getenv("DB_HOST", "127.0.0.1"),
+    port=int(os.getenv("DB_PORT", 3306)),
+    user=os.getenv("DB_USER", "root"),
+    password=os.getenv("DB_PASSWORD", ""),
+    database=os.getenv("DB_NAME", "analytics_demo"),
+)
+cursor = conn.cursor()
+
+# ============================================================
+# 第二步:创建表结构
+# ============================================================
+
+# 员工表:存储公司内部人员信息
+# 用于分析:部门人数分布、薪资结构、入职趋势等
+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;
+""")
+
+# ============================================================
+# 第三步:插入演示数据
+# ============================================================
+
+# 员工数据:5 个员工,分属 3 个部门
+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"),
+]
+
+# 产品数据:4 款商品,覆盖 2 个品类
+products_data = [
+    (1, "笔记本电脑", "电子产品", 6999.00, 500),
+    (2, "机械键盘",   "电子产品", 399.00,  1000),
+    (3, "办公椅",     "办公用品", 499.00,  300),
+    (4, "显示器",     "电子产品", 1200.00, 400),
+]
+
+# 订单数据:5 笔订单,模拟真实购买行为
+orders_data = [
+    (1, 1, 1, 2, "2024-01-15"),  # 张三买了 2 台笔记本
+    (2, 2, 2, 15, "2024-01-16"),  # 李四买了 15 个键盘
+    (3, 3, 1, 10, "2024-01-17"),  # 王五买了 10 台笔记本
+    (4, 5, 3, 6, "2024-01-18"),  # 钱七买了 6 把办公椅
+    (5, 2, 4, 5, "2024-01-19"),  # 李四买了 5 台显示器
+]
+
+# executemany 批量插入,比逐条 insert 高效得多
+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 条订单记录")
+```
+
+---
+
+## 五、NL2SQL 自然语言查询 Agent
+### 5.1 工作原理
+NL2SQL Agent 的核心思路是:**让 LLM 借助工具链自主完成"看表结构 → 写 SQL → 检查语法 → 执行查询 → 总结结果"的全流程**。
+
+<!-- 这是一张图片,ocr 内容为: -->
+![](NL2SQL工作流.png)
+
+> ▲ NL2SQL Agent 工作流:侦察表结构 → 生成 SQL → 检查语法 → 执行查询 → 总结结果
+>
+
+> **关键点**:Agent 不是一次性生成 SQL 就完事了,它会**先侦察(看表结构)、再行动(写 SQL)、后验证(检查语法)**,这套工作流比人肉写 SQL 还严谨。
+>
+
+### 5.2 搭建 SQL Agent
+```python
+import os
+from langchain_community.utilities import SQLDatabase
+from langchain_community.agent_toolkits import SQLDatabaseToolkit
+from langchain_openai import ChatOpenAI
+from langchain.agents import create_agent  # langchain 1.3.1 的新 API
+
+# ============================================================
+# 第一步:连接数据库
+# ============================================================
+# SQLDatabase.from_uri 接受标准的数据库连接 URI
+# LangChain 内部会用 SQLAlchemy 管理连接池
+db_uri = (
+    f"mysql+pymysql://{os.getenv('DB_USER')}:{os.getenv('DB_PASSWORD')}"
+    f"@{os.getenv('DB_HOST')}:{os.getenv('DB_PORT')}/{os.getenv('DB_NAME')}"
+)
+db = SQLDatabase.from_uri(db_uri)
+
+# 验证连接:打印可用的表名
+print(f"数据库连接成功")
+print(f"   可用表:{db.get_usable_table_names()}")
+
+# ============================================================
+# 第二步:初始化大模型
+# ============================================================
+llm = ChatOpenAI(
+    model="deepseek-v4-flash",
+    api_key=os.getenv("DEEPSEEK_API_KEY"),
+    base_url="https://api.deepseek.com",
+    temperature=0,  # SQL 生成需要确定性输出
+)
+
+# ============================================================
+# 第三步:创建 SQL 工具包
+# ============================================================
+# SQLDatabaseToolkit 会自动注册 4 个工具:
+#   1. sql_db_list_tables    — 列出数据库中所有表
+#   2. sql_db_schema         — 获取指定表的 DDL 结构
+#   3. sql_db_query_checker  — 检查 SQL 语法是否正确
+#   4. sql_db_query          — 执行 SQL 并返回结果
+toolkit = SQLDatabaseToolkit(db=db, llm=llm)
+tools = toolkit.get_tools()
+
+print(f"\nSQL 工具包已加载({len(tools)} 个工具):")
+for t in tools:
+    print(f"   - {t.name}")
+
+# ============================================================
+# 第四步:定义 System Prompt(Agent 的"岗位说明书")
+# ============================================================
+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(langchain 1.3.1 一行搞定)
+# ============================================================
+# create_agent 返回一个可直接调用的 Runnable,无需再套 AgentExecutor
+# 它内部自动处理工具调用、循环迭代、错误重试等逻辑
+sql_agent = create_agent(
+    model=llm,
+    tools=tools,
+    system_prompt=SQL_AGENT_PROMPT,
+)
+
+print("\nNL2SQL Agent 创建完成,可以开始提问了!")
+```
+
+### 5.3 查询示例
+#### 示例 1:简单计数查询
+```python
+# 简单问题:Agent 只需一条 COUNT SQL 就能搞定
+# create_agent 的 invoke 接口使用 messages 格式
+response = sql_agent.invoke({
+    "messages": [{"role": "user", "content": "公司一共有多少名员工?"}]
+})
+
+# 最终回答在最后一条消息里
+final_msg = response["messages"][-1]
+print(f"回答:{final_msg.content}")
+```
+
+Agent 内部执行过程(简化版):
+
+```latex
+调用 sql_db_list_tables -> 发现 employees, products, orders
+调用 sql_db_schema("employees") -> 获取字段:id, name, department, salary, hire_date
+生成 SQL: SELECT COUNT(*) FROM employees
+返回结果: 5
+总结: "公司共有 5 名员工"
+```
+
+#### 示例 2:带条件的聚合查询
+```python
+# 中等难度:需要 JOIN + WHERE + GROUP BY
+response = sql_agent.invoke({
+    "messages": [{"role": "user", "content": "列出技术部所有员工的姓名和薪资,按薪资从高到低排序"}]
+})
+final_msg = response["messages"][-1]
+print(f"回答:{final_msg.content}")
+```
+
+#### 示例 3:跨表关联查询
+```python
+# 高难度:需要 JOIN 三张表 + 聚合计算
+response = sql_agent.invoke({
+    "messages": [{"role": "user", "content": "销售部的员工总共下了多少订单?订单总金额是多少?"}]
+})
+final_msg = response["messages"][-1]
+print(f"回答:{final_msg.content}")
+```
+
+### 5.4 调试:查看 Agent 的思考过程
+在实际开发中,你经常需要看 Agent 到底干了什么——生成了什么 SQL、中间结果是什么。create_agent 返回的 messages 列表包含完整的推理链:
+
+```python
+# 遍历所有消息,还原 Agent 的完整推理链
+response = sql_agent.invoke({
+    "messages": [{"role": "user", "content": "每个部门各有多少人?"}]
+})
+
+for i, msg in enumerate(response["messages"]):
+    msg_type = msg.__class__.__name__
+
+    # AIMessage 中如果有 tool_calls,说明 Agent 调用了工具
+    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']}")
+            print(f"  参数:{tc.get('args', {})}")
+
+    # ToolMessage 是工具返回的结果
+    elif msg_type == "ToolMessage":
+        print(f"  工具返回:{msg.content[:200]}")
+
+    # AIMessage 的最终文本回答
+    elif msg.content:
+        print(f"\n最终回答:{msg.content}")
+```
+
+---
+
+## 六、数据可视化 Agent 与代码沙箱
+### 6.1 设计思路
+NL2SQL Agent 解决了"查数据"的问题,但用户往往还想"看图表"。我们需要一个能**自主编写 Python 代码、执行数据分析、生成可视化图表**的 Agent。
+
+核心挑战在于:如何让 Agent 安全地执行动态生成的代码?
+
+<!-- 这是一张图片,ocr 内容为: -->
+![](可视化Agent工作流.png)
+
+> ▲ 可视化 Agent 工作流:Agent 生成代码 → 沙箱隔离执行 → 结果回传 → 解读后返回用户
+>
+
+### 6.2 加载数据到内存
+```python
+import pandas as pd
+import matplotlib.pyplot as plt
+import seaborn as sns
+import numpy as np
+
+# ============================================================
+# 配置 matplotlib 中文显示
+# ============================================================
+# Windows 用 SimHei(黑体),macOS 用 PingFang SC
+# 如果还是乱码,试试安装 fonts-noto-cjk 并清除缓存
+plt.rcParams["font.sans-serif"] = ["SimHei", "PingFang SC", "DejaVu Sans"]
+plt.rcParams["axes.unicode_minus"] = False  # 解决负号显示为方块的问题
+
+# ============================================================
+# 从数据库加载数据到 Pandas DataFrame
+# ============================================================
+# 用 SQLAlchemy engine 复用连接,避免重复创建连接池
+from sqlalchemy import create_engine
+
+engine = create_engine(db_uri)
+
+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("✅ 数据加载完成")
+print(f"   员工表:{len(employees_df)} 行 × {len(employees_df.columns)} 列")
+print(f"   产品表:{len(products_df)} 行 × {len(products_df.columns)} 列")
+print(f"   订单表:{len(orders_df)} 行 × {len(orders_df.columns)} 列")
+print(f"\n📋 员工表示例:")
+print(employees_df.head(3).to_string(index=False))
+```
+
+### 6.3 构建 Python 代码执行沙箱
+这是整个可视化 Agent 的核心组件——一个**受控的代码执行环境**。
+
+```python
+import traceback
+from io import StringIO
+from contextlib import redirect_stdout
+from langchain.tools import tool
+
+# ============================================================
+# 定义沙箱的"白名单"——只有这些库和数据可以被代码访问
+# ============================================================
+# 这是一种简单的安全策略:不在白名单里的东西,代码碰不到
+SANDBOX_GLOBALS = {
+    # 数据:Agent 可以分析这三张表
+    "employees_df": employees_df,
+    "products_df":  products_df,
+    "orders_df":    orders_df,
+    # 工具库:Agent 可以用这些库做分析和画图
+    "pd":  pd,       # pandas — 数据处理
+    "plt": plt,      # matplotlib — 基础绑图
+    "sns": sns,      # seaborn — 统计可视化
+    "np":  np,       # numpy — 数值计算
+}
+
+
+@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)
+
+      # 画柱状图
+      plt.figure(figsize=(10, 6))
+      employees_df.groupby('department')['salary'].mean().plot(kind='bar')
+      plt.title('各部门平均薪资')
+      plt.show()
+    """
+    # 准备隔离的执行环境
+    # globals_dict 提供白名单变量,locals_dict 收集执行过程中产生的新变量
+    exec_globals = dict(SANDBOX_GLOBALS)
+    exec_locals = {}
+
+    # 用 StringIO 捕获 print() 的输出
+    # 这样 Agent 生成的代码里所有的 print 语句都会被收集
+    output_buffer = StringIO()
+
+    try:
+        # redirect_stdout 会把标准输出重定向到我们的 buffer
+        with redirect_stdout(output_buffer):
+            exec(code, exec_globals, exec_locals)
+
+        result = output_buffer.getvalue()
+
+        # 如果代码没有 print 任何东西,给个默认提示
+        if not result.strip():
+            result = "✅ 代码执行成功(无文本输出,可能已生成图表)"
+
+        return f"执行成功:\n{result}"
+
+    except Exception as e:
+        # 出错时返回完整的错误堆栈,方便 Agent 自我修正
+        error_detail = traceback.format_exc()
+        return f"❌ 执行出错:{e}\n\n{error_detail}"
+
+
+print("✅ Python 代码执行沙箱创建成功")
+```
+
+> **安全提示**:上面的 `exec` 沙箱是最简实现,适合 Demo 和内部工具。生产环境建议用 **Docker 容器沙箱** 或 **WebAssembly 隔离**,彻底杜绝恶意代码风险。详见第七章。
+>
+
+### 6.4 搭建可视化 Agent
+```python
+from langchain.agents import create_agent
+
+# ============================================================
+# 定义可视化 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"
+"""
+
+# ============================================================
+# 组装可视化 Agent(langchain 1.3.1 写法)
+# ============================================================
+viz_tools = [execute_python_code]
+
+# 同样用 create_agent 一行搞定,和 SQL Agent 的创建方式完全一致
+visualization_agent = create_agent(
+    model=llm,
+    tools=viz_tools,
+    system_prompt=VISUALIZATION_PROMPT,
+)
+
+print("数据可视化 Agent 创建完成")
+```
+
+### 6.5 可视化示例
+#### 示例 1:基础统计分析
+```python
+import warnings
+warnings.filterwarnings("ignore")  # 抑制 matplotlib 的字体警告
+
+response = visualization_agent.invoke({
+    "messages": [{"role": "user", "content": "分析一下员工薪资的分布情况,包括均值、中位数、最大最小值等统计量"}]
+})
+final_msg = response["messages"][-1]
+print(f"\n分析结果:\n{final_msg.content}")
+```
+
+Agent 内部生成的代码(大致):
+
+```python
+# 数据探索
+print("=== 薪资统计概览 ===")
+print(employees_df['salary'].describe())
+
+# 计算额外统计量
+median_salary = employees_df['salary'].median()
+print(f"\n中位数薪资:{median_salary:,.0f} 元")
+print(f"薪资范围:{employees_df['salary'].min():,.0f} ~ {employees_df['salary'].max():,.0f} 元")
+print(f"薪资标准差:{employees_df['salary'].std():,.0f} 元")
+```
+
+#### 示例 2:分组对比图表
+```python
+response = visualization_agent.invoke({
+    "messages": [{"role": "user", "content": "画一个柱状图,展示各部门的平均薪资对比"}]
+})
+final_msg = response["messages"][-1]
+print(f"\n分析结果:\n{final_msg.content}")
+```
+
+Agent 内部生成的代码(大致):
+
+```python
+# 设置中文字体
+plt.rcParams['font.sans-serif'] = ['SimHei', 'PingFang SC', 'DejaVu Sans']
+plt.rcParams['axes.unicode_minus'] = False
+
+# 按部门计算平均薪资
+dept_salary = employees_df.groupby('department')['salary'].mean().sort_values(ascending=False)
+
+# 画柱状图
+plt.figure(figsize=(10, 6))
+bars = plt.bar(dept_salary.index, dept_salary.values, color=['#4C78A8', '#F58518', '#E45756'])
+plt.title('Average Salary by Department', fontsize=16)
+plt.xlabel('Department', fontsize=12)
+plt.ylabel('Average Salary (CNY)', fontsize=12)
+
+# 在柱子上方标注具体数值
+for bar in bars:
+    height = bar.get_height()
+    plt.text(bar.get_x() + bar.get_width()/2., height,
+             f'{height:,.0f}', ha='center', va='bottom', fontsize=11)
+
+plt.tight_layout()
+plt.show()
+```
+
+#### 示例 3:跨表关联分析
+```python
+response = visualization_agent.invoke({
+    "messages": [{"role": "user", "content": "计算每个产品类别的库存总量,用水平条形图展示,并标注具体数值"}]
+})
+final_msg = response["messages"][-1]
+print(f"\n分析结果:\n{final_msg.content}")
+```
+
+### 6.6 调试:查看 Agent 生成的代码
+```python
+# 查看 Agent 的完整推理过程,重点是它生成了什么代码
+response = visualization_agent.invoke({
+    "messages": [{"role": "user", "content": "分析各产品类别的平均价格,画一个饼图"}]
+})
+
+for i, msg in enumerate(response["messages"]):
+    msg_type = msg.__class__.__name__
+
+    # AIMessage 中的 tool_calls 包含 Agent 生成的代码
+    if msg_type == "AIMessage" and hasattr(msg, "tool_calls") and msg.tool_calls:
+        print(f"\n{'='*60}")
+        for tc in msg.tool_calls:
+            if tc["name"] == "execute_python_code":
+                print(f"Agent 生成的第 {i+1} 段代码:")
+                print(f"{'='*60}")
+                print(tc["args"].get("code", ""))
+
+    # ToolMessage 是代码执行的结果
+    elif msg_type == "ToolMessage":
+        print(f"\n--- 执行结果 ---")
+        print(msg.content[:500])
+
+# 最终回答
+final_msg = response["messages"][-1]
+print(f"\n{'='*60}")
+print(f"最终回答:\n{final_msg.content}")
+```
+
+---
+
+## 七、系统整合与工程化建议
+### 7.1 双 Agent 协作方案
+在实际项目中,两个 Agent 需要协同工作。这里有两种主流方案:
+
+#### 方案 A:编排器模式(推荐)
+**编排器模式**的核心思路是:用一个路由层判断用户意图,把请求分发到合适的 Agent:
+
+```latex
+用户提问 → 编排器判断意图
+  ├─ 查询类 → NL2SQL Agent → 返回数据
+  ├─ 分析类 → 可视化 Agent → 返回图表
+  └─ 混合类 → 先查数据 → 再画图表 → 返回结果
+```
+
+```python
+# 编排器:根据用户意图路由到合适的 Agent
+def run_data_analysis(user_query: str) -> str:
+    """
+    统一入口:自动判断用户意图,路由到对应的 Agent。
+    """
+    # 用 LLM 做意图分类(也可以用规则匹配,看场景复杂度)
+    classification_prompt = f"""判断以下用户问题属于哪种类型:
+    - "query":需要查询数据库获取数据
+    - "visualize":需要画图或做统计分析
+    - "both":需要先查数据,再画图分析
+
+    用户问题:{user_query}
+
+    只回复一个词:query / visualize / both"""
+
+    intent = llm.invoke(classification_prompt).content.strip().lower()
+
+    if intent == "query":
+        # create_agent 返回的消息格式:取最后一条消息
+        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
+        # 把查询结果传给可视化 Agent
+        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
+```
+
+#### 方案 B:统一 Agent 模式(作业)
+把 SQL 工具和 Python 代码执行工具都挂到同一个 Agent 上,让 LLM 自己决定调用哪些工具。优点是简单,缺点是工具太多时 LLM 容易"选择困难"。
+
+### 7.2 安全最佳实践
+| 风险 | 缓解措施 |
+| :--- | :--- |
+| SQL 注入 | 使用参数化查询,Agent 生成的 SQL 只读权限 |
+| 恶意代码执行 | 生产环境用 Docker/WASM 沙箱隔离 |
+| API Key 泄露 | 环境变量管理,禁止硬编码 |
+| 数据泄露 | 按用户角色控制数据访问范围 |
+| 大量数据导出 | 限制查询返回行数(`LIMIT`) |
+
+
+### 7.3 生产环境部署清单
+```latex
+[ ] 数据库用户只授予 SELECT 权限(禁止 DROP/DELETE/UPDATE)
+[ ] 代码沙箱替换为 Docker 容器(限制 CPU/内存/网络)
+[ ] 接入日志系统,记录所有 Agent 调用和 SQL 执行
+[ ] 添加速率限制,防止恶意刷接口
+[ ] 实现会话管理,支持多轮对话上下文
+[ ] 接入监控告警,异常查询及时通知
+[ ] 前端加一层用户认证,不要裸奔
+[ ] 考虑缓存热点查询结果,降低数据库压力
+```
+
+### 7.4 进阶方向
+当你把这套系统跑通之后,可以继续探索:
+
++ **接入更多数据源**:Excel、CSV、API 接口、MongoDB 等
++ **支持多轮对话**:用 LangChain 的 Memory 模块维护上下文
++ **图表自动美化**:用 LLM 选择最合适的图表类型和配色方案
++ **报告自动生成**:定期跑批,自动生成周报/月报
++ **RAG 增强**:把数据字典、业务规则做成向量库,提升 Agent 的业务理解能力
+
+---
+
+> **写在最后**:这套系统的核心价值不在于技术多复杂,而在于它**打通了"人的语言"到"数据的语言"的最后一公里**。让不懂 SQL 的人也能做数据分析,这才是 Agent 真正该干的事。
+>
+

BIN
001-agent/__pycache__/unified_agent.cpython-313.pyc


BIN
001-agent/department_salary_comparison.png


+ 252 - 0
001-agent/unified_agent.py

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

BIN
002-rag/docs/浦发上海浦东发展银行西安分行个金客户经理考核办法.pdf


+ 159 - 0
002-rag/rag_agent.py

@@ -0,0 +1,159 @@
+import os
+
+import chromadb
+from dotenv import load_dotenv
+from langchain_community.document_loaders import PyMuPDFLoader
+from langchain_community.embeddings import DashScopeEmbeddings
+from langchain_community.vectorstores import Chroma
+from langchain_core.output_parsers import StrOutputParser
+from langchain_core.prompts import ChatPromptTemplate
+from langchain_openai import ChatOpenAI
+from langchain_classic.chains.combine_documents import create_stuff_documents_chain
+from langchain_text_splitters import RecursiveCharacterTextSplitter
+
+# ── 配置 ──────────────────────────────────────────────────────────
+BASE_DIR = os.path.dirname(os.path.abspath(__file__))
+
+PDF_PATH = os.path.join(BASE_DIR, "docs", "浦发上海浦东发展银行西安分行个金客户经理考核办法.pdf")
+CHROMA_DB_PATH = os.path.join(BASE_DIR, "chroma")
+COLLECTION_NAME = "shanghai_bank_policy"
+
+PROMPT_TEMPLATE = """
+你是一个专业的知识库助手。请根据以下上下文回答问题。
+
+**规则:**
+- 只基于提供的上下文回答,不要编造
+- 如果上下文中没有相关信息,直接说「根据现有资料,我找不到这个问题的答案」
+- 回答要简洁直接,引用原文时用引号
+
+**上下文:**
+{context}
+
+**问题:**
+{question}
+"""
+
+
+def load_config() -> dict:
+    """加载 .env 中的配置项"""
+    load_dotenv()
+    config = {
+        "api_key": os.getenv("ALIYUN_API_KEY"),
+        "base_url": os.getenv("ALIYUN_BASE_URL"),
+        "chat_model": os.getenv("ALIYUN_CHAT_MODEL"),
+        "embedding_model": os.getenv("ALIYUN_EMBEDDING_MODEL", "text-embedding-v3"),
+    }
+
+    missing = [k for k, v in config.items() if not v and k != "embedding_model"]
+    if missing:
+        raise EnvironmentError(f"缺少环境变量: {missing},请检查 .env 文件")
+
+    return config
+
+
+def init_llm(config: dict) -> ChatOpenAI:
+    return ChatOpenAI(
+        base_url=config["base_url"],
+        api_key=config["api_key"],
+        model=config["chat_model"],
+        temperature=0,
+        timeout=10,
+        max_retries=1,
+    )
+
+
+def init_embeddings(config: dict) -> DashScopeEmbeddings:
+    return DashScopeEmbeddings(
+        model=config["embedding_model"],
+        dashscope_api_key=config["api_key"],
+    )
+
+
+def clean_pdf_text(text: str) -> str:
+    """清洗 PDF 解析出的文本,去除常见噪声"""
+    import re
+
+    # 删除非中文字符之间的换行符
+    text = re.sub(r'[^一](\n)[^一]',
+                  lambda m: m.group(0).replace('\n', ''), text)
+
+    # 删除项目符号和多余空格
+    text = text.replace('•', '').replace('  ', ' ')
+
+    # 删除连续的换行符(保留一个)
+    text = re.sub(r'\n{2,}', '\n', text)
+
+    return text.strip()
+
+def build_or_load_vectorstore(embeddings: DashScopeEmbeddings) -> Chroma:
+    """如果 collection 已存在则直接加载,否则解析 PDF 并建库"""
+    chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH)
+    existing_collections = [col.name for col in chroma_client.list_collections()]
+
+    if COLLECTION_NAME in existing_collections:
+        print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载")
+        return Chroma(
+            embedding_function=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+        )
+
+    print(f"ℹ️  未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库")
+
+    loader = PyMuPDFLoader(PDF_PATH)
+    pages = loader.load()
+
+    # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
+    clean_pages = [clean_pdf_text(page) for page in pages]
+
+
+    text_splitter = RecursiveCharacterTextSplitter(
+        chunk_size=500,
+        chunk_overlap=50,
+        separators=["\n\n", "\n", "。", ";", ",", " ", ""],
+    )
+    docs = text_splitter.split_documents(pages)
+
+    vectorstore = Chroma.from_documents(
+        embedding=embeddings,
+        collection_name=COLLECTION_NAME,
+        client=chroma_client,
+        documents=docs,
+    )
+    print(f"✅ 建库完成,共 {len(docs)} 个分块")
+    return vectorstore
+
+
+def ask(query: str, retriever, llm) -> str:
+    """检索 + 生成回答"""
+    relevant_docs = retriever.invoke(query)
+    context = "\n\n---\n\n".join(d.page_content for d in relevant_docs)
+
+    prompt = ChatPromptTemplate.from_template(PROMPT_TEMPLATE)
+    chain = prompt | llm | StrOutputParser()
+
+
+    return chain.invoke({"context": context, "question": query})
+
+
+def main():
+    config = load_config()
+
+    llm = init_llm(config)
+    embeddings = init_embeddings(config)
+    print("✅ 模型初始化完成")
+
+    vectorstore = build_or_load_vectorstore(embeddings)
+    retriever = vectorstore.as_retriever(search_kwargs={"k": 3})
+
+    query = "什么是工作质量考核标准?"
+    answer = ask(query, retriever, llm)
+
+    print("\n" + "=" * 60)
+    print(f"问题:{query}")
+    print("=" * 60)
+    print(answer)
+
+
+if __name__ == "__main__":
+    main()

BIN
003-rag-optimized/docs/浦发上海浦东发展银行西安分行个金客户经理考核办法.pdf


+ 334 - 0
003-rag-optimized/rag_agent_optimized-decomposition.py

@@ -0,0 +1,334 @@
+import os
+
+import chromadb
+from dotenv import load_dotenv
+from langchain_community.document_loaders import PyMuPDFLoader
+from langchain_community.embeddings import DashScopeEmbeddings
+from langchain_community.retrievers import BM25Retriever
+from schema import RAGWithQueryRewriting
+from schema import RAGWithDecomposition
+from langchain_community.vectorstores import Chroma
+from langchain_core.output_parsers import StrOutputParser
+from langchain_core.prompts import ChatPromptTemplate
+from langchain_openai import ChatOpenAI
+from langchain_community.retrievers import BM25Retriever
+from langchain_community.embeddings import DashScopeEmbeddings
+from langchain_classic.retrievers import EnsembleRetriever
+from langchain_classic.chains.combine_documents import create_stuff_documents_chain
+from langchain_text_splitters import RecursiveCharacterTextSplitter
+from langchain_core.prompts import PromptTemplate
+from langchain_community.chat_models import ChatTongyi
+from langchain_classic.retrievers.multi_query import MultiQueryRetriever
+from langchain_experimental.text_splitter import SemanticChunker
+# ── 配置 ──────────────────────────────────────────────────────────
+BASE_DIR = os.path.dirname(os.path.abspath(__file__))
+
+PDF_PATH = os.path.join(BASE_DIR, "docs", "浦发上海浦东发展银行西安分行个金客户经理考核办法.pdf")
+CHROMA_DB_PATH = os.path.join(BASE_DIR, "chroma")
+COLLECTION_NAME = "shanghai_bank_policy"
+
+PROMPT_TEMPLATE = """
+你是一个专业的知识库助手。请根据以下上下文回答问题。
+
+**规则:**
+- 只基于提供的上下文回答,不要编造
+- 如果上下文中没有相关信息,直接说「根据现有资料,我找不到这个问题的答案」
+- 回答要简洁直接,引用原文时用引号
+
+**上下文:**
+{context}
+
+**问题:**
+{question}
+"""
+
+
+def load_config() -> dict:
+    """加载 .env 中的配置项"""
+    load_dotenv()
+    config = {
+        "api_key": os.getenv("ALIYUN_API_KEY"),
+        "base_url": os.getenv("ALIYUN_BASE_URL"),
+        "chat_model": os.getenv("ALIYUN_CHAT_MODEL"),
+        "embedding_model": os.getenv("ALIYUN_EMBEDDING_MODEL", "text-embedding-v3"),
+    }
+
+    missing = [k for k, v in config.items() if not v and k != "embedding_model"]
+    if missing:
+        raise EnvironmentError(f"缺少环境变量: {missing},请检查 .env 文件")
+
+    return config
+
+
+def init_llm(config: dict) -> ChatOpenAI:
+    return ChatOpenAI(
+        base_url=config["base_url"],
+        api_key=config["api_key"],
+        model=config["chat_model"],
+        temperature=0,
+        timeout=60,
+        max_retries=1,
+    )
+
+
+def init_embeddings(config: dict) -> DashScopeEmbeddings:
+    return DashScopeEmbeddings(
+        model=config["embedding_model"],
+        dashscope_api_key=config["api_key"],
+    )
+
+
+def clean_pdf_text(text: str) -> str:
+    """清洗 PDF 解析出的文本,去除常见噪声"""
+    import re
+
+    # 删除非中文字符之间的换行符
+    text = re.sub(r'[^一](\n)[^一]',
+                  lambda m: m.group(0).replace('\n', ''), text)
+
+    # 删除项目符号和多余空格
+    text = text.replace('•', '').replace('  ', ' ')
+
+    # 删除连续的换行符(保留一个)
+    text = re.sub(r'\n{2,}', '\n', text)
+
+    return text.strip()
+
+def build_or_load_vectorstore(embeddings: DashScopeEmbeddings) -> Chroma:
+    """如果 collection 已存在则直接加载,否则解析 PDF 并建库"""
+    chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH)
+    existing_collections = [col.name for col in chroma_client.list_collections()]
+
+    if COLLECTION_NAME in existing_collections:
+        print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载")
+        return Chroma(
+            embedding_function=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+        )
+
+    print(f"ℹ️  未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库")
+
+    loader = PyMuPDFLoader(PDF_PATH)
+    pages = loader.load()
+
+    # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
+    pages = clean_documents(pages)
+
+
+    # 创建语义分块器
+    # breakpoint_threshold_type="percentile" 表示用百分位数法确定切分阈值
+    # breakpoint_threshold_amount=95 表示只有相似度排名后 5% 的位置才会被切开
+    text_splitter = SemanticChunker(
+        embeddings=embeddings,
+        breakpoint_threshold_type="percentile",
+        breakpoint_threshold_amount=95  # 值越大,切出来的块越少(越粗)
+    )
+
+
+    docs = text_splitter.split_documents(pages)
+
+    vectorstore = Chroma.from_documents(
+        embedding=embeddings,
+        collection_name=COLLECTION_NAME,
+        client=chroma_client,
+        documents=docs,
+    )
+    print(f"✅ 建库完成,共 {len(docs)} 个分块")
+    return vectorstore
+
+
+def build_ensemble_retriever(embeddings: DashScopeEmbeddings) -> EnsembleRetriever:
+    """如果 collection 已存在则直接加载,否则解析 PDF 并建库"""
+    chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH)
+    existing_collections = [col.name for col in chroma_client.list_collections()]
+
+    loader = PyMuPDFLoader(PDF_PATH)
+    pages = loader.load()
+
+    # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
+    pages = clean_documents(pages)
+
+    # 创建语义分块器
+    # breakpoint_threshold_type="percentile" 表示用百分位数法确定切分阈值
+    # breakpoint_threshold_amount=95 表示只有相似度排名后 5% 的位置才会被切开
+    text_splitter = SemanticChunker(
+        embeddings=embeddings,
+        breakpoint_threshold_type="percentile",
+        breakpoint_threshold_amount=95  # 值越大,切出来的块越少(越粗)
+    )
+
+    docs = text_splitter.split_documents(pages)
+
+    vectorstore = None
+
+    if COLLECTION_NAME in existing_collections:
+        print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载")
+        vectorstore =  Chroma(
+            embedding_function=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+        )
+    else:
+
+        print(f"ℹ️  未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库")
+
+        vectorstore = Chroma.from_documents(
+            embedding=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+            documents=docs,
+        )
+
+    # ---- 创建 BM25 检索器 ----
+    # BM25 不需要向量,直接基于文本的关键词匹配
+    bm25_retriever = BM25Retriever.from_documents(docs)
+    bm25_retriever.k = 10
+
+    vector_retriever = vectorstore.as_retriever(search_kwargs={"k": 10})
+
+    ensemble_retriever = EnsembleRetriever(
+        retrievers=[bm25_retriever, vector_retriever],
+        weights=[0.4, 0.6],
+        normalize_scores=True  # 将不同检索器的分数归一化到 [0,1],避免偏差
+    )
+
+    return ensemble_retriever
+
+
+def build_multi_query_retriever(embeddings: DashScopeEmbeddings, llm) -> MultiQueryRetriever:
+    # 自定义改写提示词(可选,不写则用默认的)
+    CUSTOM_PROMPT = PromptTemplate(
+        input_variables=["question"],
+        template="""你是一个专业的问题改写助手。请为下面的问题生成 4 个不同的改写版本,
+    每个版本应该:
+    - 从不同角度表达相同的意图
+    - 使用不同的关键词和表达方式
+    - 保持问题的核心含义
+
+    每个问题单独一行,不要编号。
+
+    原始问题: {question}
+
+    改写后的问题:"""
+    )
+
+    ensemble_retriever = build_ensemble_retriever(embeddings=embeddings)
+
+    # 创建多查询检索器
+    # 底层检索器用的是上面的混合检索器,这样每个改写问题都会走混合检索
+    multi_query_retriever = MultiQueryRetriever.from_llm(
+        retriever=ensemble_retriever,
+        llm=llm,
+        prompt=CUSTOM_PROMPT
+    )
+    return multi_query_retriever
+
+
+def ask(query: str, retriever, llm) -> str:
+    """检索 + 生成回答"""
+    relevant_docs = retriever.invoke(query)
+    context = "\n\n---\n\n".join(d.page_content for d in relevant_docs)
+
+    prompt = ChatPromptTemplate.from_template(PROMPT_TEMPLATE)
+    chain = prompt | llm | StrOutputParser()
+
+
+    return chain.invoke({"context": context, "question": query})
+
+
+
+def build_hyde_chain(llm) -> dict:
+    # HyDE Prompt:生成假设性文档
+    hyde_prompt = PromptTemplate(
+        input_variables=["question"],
+        template="""请根据以下问题,写一段可能包含答案的文档片段。
+    要求:
+    1. 像真实文档一样专业、详细
+    2. 包含具体的数据、步骤或事实
+    3. 长度在 100-200 字之间
+    4. 即使你不确定答案,也要根据问题合理推测,写出一段"看起来像真的"文档
+
+    问题: {question}
+
+    假设性文档:"""
+    )
+
+    hyde_chain = hyde_prompt | llm | StrOutputParser()
+    return hyde_chain
+
+
+
+def build_rewrite_chain(llm) -> dict:
+    # 查询重写 Prompt
+    rewrite_prompt = PromptTemplate(
+        input_variables=["query"],
+        template="""你是一个查询优化助手。请将用户的口语化问题改写为更适合信息检索的精确查询。
+
+    改写要求:
+    1. 补充隐含的上下文信息
+    2. 将口语化表达转为专业表述
+    3. 消除歧义,明确查询意图
+    4. 保持原意不变,不要添加原问题未提及的内容
+    5. 直接输出改写后的查询,不要解释
+
+    用户问题: {query}
+
+    改写后的查询:"""
+    )
+
+    # 构建重写链
+    rewrite_chain = rewrite_prompt | llm | StrOutputParser()
+    return rewrite_chain
+
+
+def clean_pdf_text(text: str) -> str:
+    """清洗 PDF 解析出的文本,去除常见噪声"""
+    import re
+
+    # 删除非中文字符之间的换行符
+    text = re.sub(r'[^一](\n)[^一]',
+                  lambda m: m.group(0).replace('\n', ''), text)
+
+    # 删除项目符号和多余空格
+    text = text.replace('•', '').replace('  ', ' ')
+
+    # 删除连续的换行符(保留一个)
+    text = re.sub(r'\n{2,}', '\n', text)
+
+    return text.strip()
+
+
+def clean_documents(pages: list) -> list:
+    """对 Document 列表逐个清洗 page_content,返回清洗后的新 Document 列表"""
+    for page in pages:
+        page.page_content = clean_pdf_text(page.page_content)
+    return pages
+
+
+
+def main():
+    config = load_config()
+
+    llm = init_llm(config)
+    embeddings = init_embeddings(config)
+    print("✅ 模型初始化完成")
+
+    rewrite_chain = build_rewrite_chain(llm)
+
+
+    multi_query_retriever = build_multi_query_retriever(embeddings, llm)
+    rag = RAGWithDecomposition(retriever = multi_query_retriever,llm = llm)
+
+    query = "什么是工作质量考核标准,还有聘任考核程序是什么?"
+    result = rag.invoke(query)
+
+
+    print("\n" + "=" * 60)
+    print(f"问题:{query}")
+    print("=" * 60)
+    print(result["answer"])
+
+
+if __name__ == "__main__":
+    main()

+ 425 - 0
003-rag-optimized/rag_agent_optimized-hyde-search.py

@@ -0,0 +1,425 @@
+import os
+
+import chromadb
+from dotenv import load_dotenv
+from langchain_community.document_loaders import PyMuPDFLoader
+from langchain_community.embeddings import DashScopeEmbeddings
+from langchain_community.retrievers import BM25Retriever
+from schema import RAGWithQueryRewriting
+from schema import RAGWithDecomposition
+from langchain_community.vectorstores import Chroma
+from langchain_core.output_parsers import StrOutputParser
+from langchain_core.prompts import ChatPromptTemplate
+from langchain_openai import ChatOpenAI
+from langchain_community.retrievers import BM25Retriever
+from langchain_community.embeddings import DashScopeEmbeddings
+from langchain_classic.retrievers import EnsembleRetriever
+from langchain_classic.chains.combine_documents import create_stuff_documents_chain
+from langchain_text_splitters import RecursiveCharacterTextSplitter
+from langchain_core.prompts import PromptTemplate
+from langchain_community.chat_models import ChatTongyi
+from langchain_classic.retrievers.multi_query import MultiQueryRetriever
+from langchain_experimental.text_splitter import SemanticChunker
+# ── 配置 ──────────────────────────────────────────────────────────
+BASE_DIR = os.path.dirname(os.path.abspath(__file__))
+
+PDF_PATH = os.path.join(BASE_DIR, "docs", "浦发上海浦东发展银行西安分行个金客户经理考核办法.pdf")
+CHROMA_DB_PATH = os.path.join(BASE_DIR, "chroma")
+COLLECTION_NAME = "shanghai_bank_policy"
+
+PROMPT_TEMPLATE = """
+你是一个专业的知识库助手。请根据以下上下文回答问题。
+
+**规则:**
+- 只基于提供的上下文回答,不要编造
+- 如果上下文中没有相关信息,直接说「根据现有资料,我找不到这个问题的答案」
+- 回答要简洁直接,引用原文时用引号
+
+**上下文:**
+{context}
+
+**问题:**
+{question}
+"""
+
+
+def load_config() -> dict:
+    """加载 .env 中的配置项"""
+    load_dotenv()
+    config = {
+        "api_key": os.getenv("ALIYUN_API_KEY"),
+        "base_url": os.getenv("ALIYUN_BASE_URL"),
+        "chat_model": os.getenv("ALIYUN_CHAT_MODEL"),
+        "embedding_model": os.getenv("ALIYUN_EMBEDDING_MODEL", "text-embedding-v3"),
+    }
+
+    missing = [k for k, v in config.items() if not v and k != "embedding_model"]
+    if missing:
+        raise EnvironmentError(f"缺少环境变量: {missing},请检查 .env 文件")
+
+    return config
+
+
+def init_llm(config: dict) -> ChatOpenAI:
+    return ChatOpenAI(
+        base_url=config["base_url"],
+        api_key=config["api_key"],
+        model=config["chat_model"],
+        temperature=0,
+        timeout=120,
+        max_retries=1,
+    )
+
+
+def init_embeddings(config: dict) -> DashScopeEmbeddings:
+    return DashScopeEmbeddings(
+        model=config["embedding_model"],
+        dashscope_api_key=config["api_key"],
+    )
+
+
+def clean_pdf_text(text: str) -> str:
+    """清洗 PDF 解析出的文本,去除常见噪声"""
+    import re
+
+    # 删除非中文字符之间的换行符
+    text = re.sub(r'[^一](\n)[^一]',
+                  lambda m: m.group(0).replace('\n', ''), text)
+
+    # 删除项目符号和多余空格
+    text = text.replace('•', '').replace('  ', ' ')
+
+    # 删除连续的换行符(保留一个)
+    text = re.sub(r'\n{2,}', '\n', text)
+
+    return text.strip()
+
+def build_or_load_vectorstore(embeddings: DashScopeEmbeddings) -> Chroma:
+    """如果 collection 已存在则直接加载,否则解析 PDF 并建库"""
+    chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH)
+    existing_collections = [col.name for col in chroma_client.list_collections()]
+
+    if COLLECTION_NAME in existing_collections:
+        print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载")
+        return Chroma(
+            embedding_function=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+        )
+
+    print(f"ℹ️  未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库")
+
+    loader = PyMuPDFLoader(PDF_PATH)
+    pages = loader.load()
+
+    # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
+    pages = clean_documents(pages)
+
+
+    # 创建语义分块器
+    # breakpoint_threshold_type="percentile" 表示用百分位数法确定切分阈值
+    # breakpoint_threshold_amount=95 表示只有相似度排名后 5% 的位置才会被切开
+    text_splitter = SemanticChunker(
+        embeddings=embeddings,
+        breakpoint_threshold_type="percentile",
+        breakpoint_threshold_amount=95  # 值越大,切出来的块越少(越粗)
+    )
+
+
+    docs = text_splitter.split_documents(pages)
+
+    vectorstore = Chroma.from_documents(
+        embedding=embeddings,
+        collection_name=COLLECTION_NAME,
+        client=chroma_client,
+        documents=docs,
+    )
+    print(f"✅ 建库完成,共 {len(docs)} 个分块")
+    return vectorstore
+
+
+def build_ensemble_retriever(embeddings: DashScopeEmbeddings) -> EnsembleRetriever:
+    """如果 collection 已存在则直接加载,否则解析 PDF 并建库"""
+    chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH)
+    existing_collections = [col.name for col in chroma_client.list_collections()]
+
+    loader = PyMuPDFLoader(PDF_PATH)
+    pages = loader.load()
+
+    # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
+    pages = clean_documents(pages)
+
+    # 创建语义分块器
+    # breakpoint_threshold_type="percentile" 表示用百分位数法确定切分阈值
+    # breakpoint_threshold_amount=95 表示只有相似度排名后 5% 的位置才会被切开
+    text_splitter = SemanticChunker(
+        embeddings=embeddings,
+        breakpoint_threshold_type="percentile",
+        breakpoint_threshold_amount=95  # 值越大,切出来的块越少(越粗)
+    )
+
+    docs = text_splitter.split_documents(pages)
+
+    vectorstore = None
+
+    if COLLECTION_NAME in existing_collections:
+        print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载")
+        vectorstore =  Chroma(
+            embedding_function=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+        )
+    else:
+
+        print(f"ℹ️  未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库")
+
+        vectorstore = Chroma.from_documents(
+            embedding=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+            documents=docs,
+        )
+
+    # ---- 创建 BM25 检索器 ----
+    # BM25 不需要向量,直接基于文本的关键词匹配
+    bm25_retriever = BM25Retriever.from_documents(docs)
+    bm25_retriever.k = 10
+
+    vector_retriever = vectorstore.as_retriever(search_kwargs={"k": 10})
+
+    ensemble_retriever = EnsembleRetriever(
+        retrievers=[bm25_retriever, vector_retriever],
+        weights=[0.4, 0.6],
+        normalize_scores=True  # 将不同检索器的分数归一化到 [0,1],避免偏差
+    )
+
+    return ensemble_retriever
+
+
+def build_multi_query_retriever(embeddings: DashScopeEmbeddings, llm) -> MultiQueryRetriever:
+    # 自定义改写提示词(可选,不写则用默认的)
+    CUSTOM_PROMPT = PromptTemplate(
+        input_variables=["question"],
+        template="""你是一个专业的问题改写助手。请为下面的问题生成 4 个不同的改写版本,
+    每个版本应该:
+    - 从不同角度表达相同的意图
+    - 使用不同的关键词和表达方式
+    - 保持问题的核心含义
+
+    每个问题单独一行,不要编号。
+
+    原始问题: {question}
+
+    改写后的问题:"""
+    )
+
+    ensemble_retriever = build_ensemble_retriever(embeddings=embeddings)
+
+    # 创建多查询检索器
+    # 底层检索器用的是上面的混合检索器,这样每个改写问题都会走混合检索
+    multi_query_retriever = MultiQueryRetriever.from_llm(
+        retriever=ensemble_retriever,
+        llm=llm,
+        prompt=CUSTOM_PROMPT
+    )
+    return multi_query_retriever
+
+
+def ask(query: str, retriever, llm) -> str:
+    """检索 + 生成回答"""
+    relevant_docs = retriever.invoke(query)
+    context = "\n\n---\n\n".join(d.page_content for d in relevant_docs)
+
+    prompt = ChatPromptTemplate.from_template(PROMPT_TEMPLATE)
+    chain = prompt | llm | StrOutputParser()
+
+
+    return chain.invoke({"context": context, "question": query})
+
+
+
+def build_hyde_chain(llm) -> dict:
+    # HyDE Prompt:生成假设性文档
+    hyde_prompt = PromptTemplate(
+        input_variables=["question"],
+        template="""请根据以下问题,写一段可能包含答案的文档片段。
+    要求:
+    1. 像真实文档一样专业、详细
+    2. 包含具体的数据、步骤或事实
+    3. 长度在 100-200 字之间
+    4. 即使你不确定答案,也要根据问题合理推测,写出一段"看起来像真的"文档
+
+    问题: {question}
+
+    假设性文档:"""
+    )
+
+    hyde_chain = hyde_prompt | llm | StrOutputParser()
+    return hyde_chain
+
+
+
+def build_rewrite_chain(llm) -> dict:
+    # 查询重写 Prompt
+    rewrite_prompt = PromptTemplate(
+        input_variables=["query"],
+        template="""你是一个查询优化助手。请将用户的口语化问题改写为更适合信息检索的精确查询。
+
+    改写要求:
+    1. 补充隐含的上下文信息
+    2. 将口语化表达转为专业表述
+    3. 消除歧义,明确查询意图
+    4. 保持原意不变,不要添加原问题未提及的内容
+    5. 直接输出改写后的查询,不要解释
+
+    用户问题: {query}
+
+    改写后的查询:"""
+    )
+
+    # 构建重写链
+    rewrite_chain = rewrite_prompt | llm | StrOutputParser()
+    return rewrite_chain
+
+
+
+def hybrid_hyde_search(llm, embeddings: DashScopeEmbeddings, hyde_chain, question) -> dict:
+    """如果 collection 已存在则直接加载,否则解析 PDF 并建库"""
+    chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH)
+    existing_collections = [col.name for col in chroma_client.list_collections()]
+
+    loader = PyMuPDFLoader(PDF_PATH)
+    pages = loader.load()
+
+    # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
+    pages = clean_documents(pages)
+
+    # 创建语义分块器
+    # breakpoint_threshold_type="percentile" 表示用百分位数法确定切分阈值
+    # breakpoint_threshold_amount=95 表示只有相似度排名后 5% 的位置才会被切开
+    text_splitter = SemanticChunker(
+        embeddings=embeddings,
+        breakpoint_threshold_type="percentile",
+        breakpoint_threshold_amount=95  # 值越大,切出来的块越少(越粗)
+    )
+
+    docs = text_splitter.split_documents(pages)
+
+    vectorstore = None
+
+    if COLLECTION_NAME in existing_collections:
+        print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载")
+        vectorstore =  Chroma(
+            embedding_function=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+        )
+    else:
+
+        print(f"ℹ️  未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库")
+
+        vectorstore = Chroma.from_documents(
+            embedding=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+            documents=docs,
+        )
+
+    # ---- 创建 BM25 检索器 ----
+    # BM25 不需要向量,直接基于文本的关键词匹配
+    bm25_retriever = BM25Retriever.from_documents(docs)
+    bm25_retriever.k = 10
+
+    hypothetical_doc = hyde_chain.invoke({"question": question})
+
+    # 方案一:用假设性文档检索(偏语义)
+    hyde_docs = vectorstore.similarity_search(hypothetical_doc, k=10)
+
+
+    bm25_docs = bm25_retriever.invoke(hypothetical_doc, k=10)  # question must be a str here
+
+    # 方案二:用原始问题检索(偏精确)
+    original_docs = vectorstore.similarity_search(question, k=10)
+
+    # 方案二:用原始问题检索(偏精确)
+
+    # 合并去重
+    seen = set()
+    merged_docs = []
+    for doc in hyde_docs + bm25_docs:
+        if hash(doc.page_content) not in seen:
+            seen.add(hash(doc.page_content))
+            merged_docs.append(doc)
+
+    all_docs = merged_docs[:10]
+
+    # 第三步:用原始问题 + 检索结果生成答案
+    context = "\n\n".join([doc.page_content for doc in all_docs])
+    answer_prompt = f"""基于以下上下文回答用户问题。如果上下文中没有相关信息,请说明。
+
+    上下文:{context}
+
+    用户问题:{question}
+    答案:"""
+
+    answer = llm.invoke(answer_prompt).content
+
+    return {
+        "original_query": question,
+        "hypothetical_doc": hypothetical_doc,
+        "retrieved_docs": docs,
+        "answer": answer
+    }
+
+
+
+
+
+
+def clean_pdf_text(text: str) -> str:
+    """清洗 PDF 解析出的文本,去除常见噪声"""
+    import re
+
+    # 删除非中文字符之间的换行符
+    text = re.sub(r'[^一](\n)[^一]',
+                  lambda m: m.group(0).replace('\n', ''), text)
+
+    # 删除项目符号和多余空格
+    text = text.replace('•', '').replace('  ', ' ')
+
+    # 删除连续的换行符(保留一个)
+    text = re.sub(r'\n{2,}', '\n', text)
+
+    return text.strip()
+
+
+def clean_documents(pages: list) -> list:
+    """对 Document 列表逐个清洗 page_content,返回清洗后的新 Document 列表"""
+    for page in pages:
+        page.page_content = clean_pdf_text(page.page_content)
+    return pages
+
+
+
+def main():
+    config = load_config()
+
+    llm = init_llm(config)
+    embeddings = init_embeddings(config)
+    print("✅ 模型初始化完成")
+
+
+    query = "什么是工作质量考核标准,还有聘任考核程序是什么?"
+
+    hyde_chain = build_hyde_chain(llm)
+
+    result = hybrid_hyde_search(llm, embeddings, hyde_chain, query)
+
+    print("\n" + "=" * 60)
+    print(f"问题:{query}")
+    print("=" * 60)
+    print(result["answer"])
+
+
+if __name__ == "__main__":
+    main()

+ 315 - 0
003-rag-optimized/rag_agent_optimized-two.py

@@ -0,0 +1,315 @@
+import os
+
+import chromadb
+from dotenv import load_dotenv
+from langchain_community.document_loaders import PyMuPDFLoader
+from langchain_community.embeddings import DashScopeEmbeddings
+from langchain_community.retrievers import BM25Retriever
+from schema import RAGWithQueryRewriting
+from langchain_community.vectorstores import Chroma
+from langchain_core.output_parsers import StrOutputParser
+from langchain_core.prompts import ChatPromptTemplate
+from langchain_openai import ChatOpenAI
+from langchain_community.retrievers import BM25Retriever
+from langchain_community.embeddings import DashScopeEmbeddings
+from langchain_classic.retrievers import EnsembleRetriever
+from langchain_classic.chains.combine_documents import create_stuff_documents_chain
+from langchain_text_splitters import RecursiveCharacterTextSplitter
+from langchain_core.prompts import PromptTemplate
+from langchain_community.chat_models import ChatTongyi
+from langchain_classic.retrievers.multi_query import MultiQueryRetriever
+from langchain_experimental.text_splitter import SemanticChunker
+# ── 配置 ──────────────────────────────────────────────────────────
+BASE_DIR = os.path.dirname(os.path.abspath(__file__))
+
+PDF_PATH = os.path.join(BASE_DIR, "docs", "浦发上海浦东发展银行西安分行个金客户经理考核办法.pdf")
+CHROMA_DB_PATH = os.path.join(BASE_DIR, "chroma")
+COLLECTION_NAME = "shanghai_bank_policy"
+
+PROMPT_TEMPLATE = """
+你是一个专业的知识库助手。请根据以下上下文回答问题。
+
+**规则:**
+- 只基于提供的上下文回答,不要编造
+- 如果上下文中没有相关信息,直接说「根据现有资料,我找不到这个问题的答案」
+- 回答要简洁直接,引用原文时用引号
+
+**上下文:**
+{context}
+
+**问题:**
+{question}
+"""
+
+
+def load_config() -> dict:
+    """加载 .env 中的配置项"""
+    load_dotenv()
+    config = {
+        "api_key": os.getenv("ALIYUN_API_KEY"),
+        "base_url": os.getenv("ALIYUN_BASE_URL"),
+        "chat_model": os.getenv("ALIYUN_CHAT_MODEL"),
+        "embedding_model": os.getenv("ALIYUN_EMBEDDING_MODEL", "text-embedding-v3"),
+    }
+
+    missing = [k for k, v in config.items() if not v and k != "embedding_model"]
+    if missing:
+        raise EnvironmentError(f"缺少环境变量: {missing},请检查 .env 文件")
+
+    return config
+
+
+def init_llm(config: dict) -> ChatOpenAI:
+    return ChatOpenAI(
+        base_url=config["base_url"],
+        api_key=config["api_key"],
+        model=config["chat_model"],
+        temperature=0,
+        timeout=10,
+        max_retries=1,
+    )
+
+
+def init_embeddings(config: dict) -> DashScopeEmbeddings:
+    return DashScopeEmbeddings(
+        model=config["embedding_model"],
+        dashscope_api_key=config["api_key"],
+    )
+
+
+def clean_pdf_text(text: str) -> str:
+    """清洗 PDF 解析出的文本,去除常见噪声"""
+    import re
+
+    # 删除非中文字符之间的换行符
+    text = re.sub(r'[^一](\n)[^一]',
+                  lambda m: m.group(0).replace('\n', ''), text)
+
+    # 删除项目符号和多余空格
+    text = text.replace('•', '').replace('  ', ' ')
+
+    # 删除连续的换行符(保留一个)
+    text = re.sub(r'\n{2,}', '\n', text)
+
+    return text.strip()
+
+def build_or_load_vectorstore(embeddings: DashScopeEmbeddings) -> Chroma:
+    """如果 collection 已存在则直接加载,否则解析 PDF 并建库"""
+    chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH)
+    existing_collections = [col.name for col in chroma_client.list_collections()]
+
+    if COLLECTION_NAME in existing_collections:
+        print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载")
+        return Chroma(
+            embedding_function=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+        )
+
+    print(f"ℹ️  未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库")
+
+    loader = PyMuPDFLoader(PDF_PATH)
+    pages = loader.load()
+
+    # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
+    pages = clean_documents(pages)
+
+
+    # 创建语义分块器
+    # breakpoint_threshold_type="percentile" 表示用百分位数法确定切分阈值
+    # breakpoint_threshold_amount=95 表示只有相似度排名后 5% 的位置才会被切开
+    text_splitter = SemanticChunker(
+        embeddings=embeddings,
+        breakpoint_threshold_type="percentile",
+        breakpoint_threshold_amount=95  # 值越大,切出来的块越少(越粗)
+    )
+
+
+    docs = text_splitter.split_documents(pages)
+
+    vectorstore = Chroma.from_documents(
+        embedding=embeddings,
+        collection_name=COLLECTION_NAME,
+        client=chroma_client,
+        documents=docs,
+    )
+    print(f"✅ 建库完成,共 {len(docs)} 个分块")
+    return vectorstore
+
+
+def build_ensemble_retriever(embeddings: DashScopeEmbeddings) -> EnsembleRetriever:
+    """如果 collection 已存在则直接加载,否则解析 PDF 并建库"""
+    chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH)
+    existing_collections = [col.name for col in chroma_client.list_collections()]
+
+    loader = PyMuPDFLoader(PDF_PATH)
+    pages = loader.load()
+
+    # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
+    pages = clean_documents(pages)
+
+    # 创建语义分块器
+    # breakpoint_threshold_type="percentile" 表示用百分位数法确定切分阈值
+    # breakpoint_threshold_amount=95 表示只有相似度排名后 5% 的位置才会被切开
+    text_splitter = SemanticChunker(
+        embeddings=embeddings,
+        breakpoint_threshold_type="percentile",
+        breakpoint_threshold_amount=95  # 值越大,切出来的块越少(越粗)
+    )
+
+    docs = text_splitter.split_documents(pages)
+
+    vectorstore = None
+
+    if COLLECTION_NAME in existing_collections:
+        print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载")
+        vectorstore =  Chroma(
+            embedding_function=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+        )
+    else:
+
+        print(f"ℹ️  未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库")
+
+        vectorstore = Chroma.from_documents(
+            embedding=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+            documents=docs,
+        )
+
+    # ---- 创建 BM25 检索器 ----
+    # BM25 不需要向量,直接基于文本的关键词匹配
+    bm25_retriever = BM25Retriever.from_documents(docs)
+    bm25_retriever.k = 10
+
+    vector_retriever = vectorstore.as_retriever(search_kwargs={"k": 10})
+
+    ensemble_retriever = EnsembleRetriever(
+        retrievers=[bm25_retriever, vector_retriever],
+        weights=[0.4, 0.6],
+        normalize_scores=True  # 将不同检索器的分数归一化到 [0,1],避免偏差
+    )
+
+    return ensemble_retriever
+
+
+def build_multi_query_retriever(embeddings: DashScopeEmbeddings, llm) -> MultiQueryRetriever:
+    # 自定义改写提示词(可选,不写则用默认的)
+    CUSTOM_PROMPT = PromptTemplate(
+        input_variables=["question"],
+        template="""你是一个专业的问题改写助手。请为下面的问题生成 4 个不同的改写版本,
+    每个版本应该:
+    - 从不同角度表达相同的意图
+    - 使用不同的关键词和表达方式
+    - 保持问题的核心含义
+
+    每个问题单独一行,不要编号。
+
+    原始问题: {question}
+
+    改写后的问题:"""
+    )
+
+    ensemble_retriever = build_ensemble_retriever(embeddings=embeddings)
+
+    # 创建多查询检索器
+    # 底层检索器用的是上面的混合检索器,这样每个改写问题都会走混合检索
+    multi_query_retriever = MultiQueryRetriever.from_llm(
+        retriever=ensemble_retriever,
+        llm=llm,
+        prompt=CUSTOM_PROMPT
+    )
+    return multi_query_retriever
+
+
+def ask(query: str, retriever, llm) -> str:
+    """检索 + 生成回答"""
+    relevant_docs = retriever.invoke(query)
+    context = "\n\n---\n\n".join(d.page_content for d in relevant_docs)
+
+    prompt = ChatPromptTemplate.from_template(PROMPT_TEMPLATE)
+    chain = prompt | llm | StrOutputParser()
+
+
+    return chain.invoke({"context": context, "question": query})
+
+
+def build_rewrite_chain(llm) -> dict:
+    # 查询重写 Prompt
+    rewrite_prompt = PromptTemplate(
+        input_variables=["query"],
+        template="""你是一个查询优化助手。请将用户的口语化问题改写为更适合信息检索的精确查询。
+
+    改写要求:
+    1. 补充隐含的上下文信息
+    2. 将口语化表达转为专业表述
+    3. 消除歧义,明确查询意图
+    4. 保持原意不变,不要添加原问题未提及的内容
+    5. 直接输出改写后的查询,不要解释
+
+    用户问题: {query}
+
+    改写后的查询:"""
+    )
+
+    # 构建重写链
+    rewrite_chain = rewrite_prompt | llm | StrOutputParser()
+    return rewrite_chain
+
+
+def clean_pdf_text(text: str) -> str:
+    """清洗 PDF 解析出的文本,去除常见噪声"""
+    import re
+
+    # 删除非中文字符之间的换行符
+    text = re.sub(r'[^一](\n)[^一]',
+                  lambda m: m.group(0).replace('\n', ''), text)
+
+    # 删除项目符号和多余空格
+    text = text.replace('•', '').replace('  ', ' ')
+
+    # 删除连续的换行符(保留一个)
+    text = re.sub(r'\n{2,}', '\n', text)
+
+    return text.strip()
+
+
+def clean_documents(pages: list) -> list:
+    """对 Document 列表逐个清洗 page_content,返回清洗后的新 Document 列表"""
+    for page in pages:
+        page.page_content = clean_pdf_text(page.page_content)
+    return pages
+
+
+
+def main():
+    config = load_config()
+
+    llm = init_llm(config)
+    embeddings = init_embeddings(config)
+    print("✅ 模型初始化完成")
+
+    rewrite_chain = build_rewrite_chain(llm)
+
+
+    multi_query_retriever = build_multi_query_retriever(embeddings, llm)
+    rag = RAGWithQueryRewriting(
+        retriever=multi_query_retriever,
+        llm=llm,
+        rewrite_chain=rewrite_chain
+    )
+
+    query = "什么是工作质量考核标准?"
+    result = rag.invoke(query)
+
+
+    print("\n" + "=" * 60)
+    print(f"问题:{query}")
+    print("=" * 60)
+    print(result["answer"])
+
+
+if __name__ == "__main__":
+    main()

+ 286 - 0
003-rag-optimized/rag_agent_optimized.py

@@ -0,0 +1,286 @@
+import os
+
+import chromadb
+from dotenv import load_dotenv
+from langchain_community.document_loaders import PyMuPDFLoader
+from langchain_community.embeddings import DashScopeEmbeddings
+from langchain_community.retrievers import BM25Retriever
+from langchain_community.vectorstores import Chroma
+from langchain_core.output_parsers import StrOutputParser
+from langchain_core.prompts import ChatPromptTemplate
+from langchain_openai import ChatOpenAI
+from langchain_community.retrievers import BM25Retriever
+from langchain_community.embeddings import DashScopeEmbeddings
+from langchain_classic.retrievers import EnsembleRetriever
+from langchain_classic.chains.combine_documents import create_stuff_documents_chain
+from langchain_text_splitters import RecursiveCharacterTextSplitter
+from langchain_core.prompts import PromptTemplate
+from langchain_community.chat_models import ChatTongyi
+from langchain_classic.retrievers.multi_query import MultiQueryRetriever
+from langchain_experimental.text_splitter import SemanticChunker
+# ── 配置 ──────────────────────────────────────────────────────────
+BASE_DIR = os.path.dirname(os.path.abspath(__file__))
+
+PDF_PATH = os.path.join(BASE_DIR, "docs", "浦发上海浦东发展银行西安分行个金客户经理考核办法.pdf")
+CHROMA_DB_PATH = os.path.join(BASE_DIR, "chroma")
+COLLECTION_NAME = "shanghai_bank_policy"
+
+PROMPT_TEMPLATE = """
+你是一个专业的知识库助手。请根据以下上下文回答问题。
+
+**规则:**
+- 只基于提供的上下文回答,不要编造
+- 如果上下文中没有相关信息,直接说「根据现有资料,我找不到这个问题的答案」
+- 回答要简洁直接,引用原文时用引号
+
+**上下文:**
+{context}
+
+**问题:**
+{question}
+"""
+
+
+def load_config() -> dict:
+    """加载 .env 中的配置项"""
+    load_dotenv()
+    config = {
+        "api_key": os.getenv("ALIYUN_API_KEY"),
+        "base_url": os.getenv("ALIYUN_BASE_URL"),
+        "chat_model": os.getenv("ALIYUN_CHAT_MODEL"),
+        "embedding_model": os.getenv("ALIYUN_EMBEDDING_MODEL", "text-embedding-v3"),
+    }
+
+    missing = [k for k, v in config.items() if not v and k != "embedding_model"]
+    if missing:
+        raise EnvironmentError(f"缺少环境变量: {missing},请检查 .env 文件")
+
+    return config
+
+
+def init_llm(config: dict) -> ChatOpenAI:
+    return ChatOpenAI(
+        base_url=config["base_url"],
+        api_key=config["api_key"],
+        model=config["chat_model"],
+        temperature=0,
+        timeout=10,
+        max_retries=1,
+    )
+
+
+def init_embeddings(config: dict) -> DashScopeEmbeddings:
+    return DashScopeEmbeddings(
+        model=config["embedding_model"],
+        dashscope_api_key=config["api_key"],
+    )
+
+
+def clean_pdf_text(text: str) -> str:
+    """清洗 PDF 解析出的文本,去除常见噪声"""
+    import re
+
+    # 删除非中文字符之间的换行符
+    text = re.sub(r'[^一](\n)[^一]',
+                  lambda m: m.group(0).replace('\n', ''), text)
+
+    # 删除项目符号和多余空格
+    text = text.replace('•', '').replace('  ', ' ')
+
+    # 删除连续的换行符(保留一个)
+    text = re.sub(r'\n{2,}', '\n', text)
+
+    return text.strip()
+
+def build_or_load_vectorstore(embeddings: DashScopeEmbeddings) -> Chroma:
+    """如果 collection 已存在则直接加载,否则解析 PDF 并建库"""
+    chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH)
+    existing_collections = [col.name for col in chroma_client.list_collections()]
+
+    if COLLECTION_NAME in existing_collections:
+        print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载")
+        return Chroma(
+            embedding_function=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+        )
+
+    print(f"ℹ️  未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库")
+
+    loader = PyMuPDFLoader(PDF_PATH)
+    pages = loader.load()
+
+    # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
+    pages = clean_documents(pages)
+
+
+    # 创建语义分块器
+    # breakpoint_threshold_type="percentile" 表示用百分位数法确定切分阈值
+    # breakpoint_threshold_amount=95 表示只有相似度排名后 5% 的位置才会被切开
+    text_splitter = SemanticChunker(
+        embeddings=embeddings,
+        breakpoint_threshold_type="percentile",
+        breakpoint_threshold_amount=95  # 值越大,切出来的块越少(越粗)
+    )
+
+
+    docs = text_splitter.split_documents(pages)
+
+    vectorstore = Chroma.from_documents(
+        embedding=embeddings,
+        collection_name=COLLECTION_NAME,
+        client=chroma_client,
+        documents=docs,
+    )
+    print(f"✅ 建库完成,共 {len(docs)} 个分块")
+    return vectorstore
+
+
+def build_ensemble_retriever(embeddings: DashScopeEmbeddings) -> EnsembleRetriever:
+    """如果 collection 已存在则直接加载,否则解析 PDF 并建库"""
+    chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH)
+    existing_collections = [col.name for col in chroma_client.list_collections()]
+
+    loader = PyMuPDFLoader(PDF_PATH)
+    pages = loader.load()
+
+    # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
+    pages = clean_documents(pages)
+
+    # 创建语义分块器
+    # breakpoint_threshold_type="percentile" 表示用百分位数法确定切分阈值
+    # breakpoint_threshold_amount=95 表示只有相似度排名后 5% 的位置才会被切开
+    text_splitter = SemanticChunker(
+        embeddings=embeddings,
+        breakpoint_threshold_type="percentile",
+        breakpoint_threshold_amount=95  # 值越大,切出来的块越少(越粗)
+    )
+
+    docs = text_splitter.split_documents(pages)
+
+    vectorstore = None
+
+    if COLLECTION_NAME in existing_collections:
+        print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载")
+        vectorstore =  Chroma(
+            embedding_function=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+        )
+    else:
+
+        print(f"ℹ️  未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库")
+
+        vectorstore = Chroma.from_documents(
+            embedding=embeddings,
+            collection_name=COLLECTION_NAME,
+            client=chroma_client,
+            documents=docs,
+        )
+
+    # ---- 创建 BM25 检索器 ----
+    # BM25 不需要向量,直接基于文本的关键词匹配
+    bm25_retriever = BM25Retriever.from_documents(docs)
+    bm25_retriever.k = 10
+
+    vector_retriever = vectorstore.as_retriever(search_kwargs={"k": 10})
+
+    ensemble_retriever = EnsembleRetriever(
+        retrievers=[bm25_retriever, vector_retriever],
+        weights=[0.4, 0.6],
+        normalize_scores=True  # 将不同检索器的分数归一化到 [0,1],避免偏差
+    )
+
+    return ensemble_retriever
+
+
+def build_multi_query_retriever(embeddings: DashScopeEmbeddings, llm) -> MultiQueryRetriever:
+    # 自定义改写提示词(可选,不写则用默认的)
+    CUSTOM_PROMPT = PromptTemplate(
+        input_variables=["question"],
+        template="""你是一个专业的问题改写助手。请为下面的问题生成 4 个不同的改写版本,
+    每个版本应该:
+    - 从不同角度表达相同的意图
+    - 使用不同的关键词和表达方式
+    - 保持问题的核心含义
+
+    每个问题单独一行,不要编号。
+
+    原始问题: {question}
+
+    改写后的问题:"""
+    )
+
+    ensemble_retriever = build_ensemble_retriever(embeddings=embeddings)
+
+    # 创建多查询检索器
+    # 底层检索器用的是上面的混合检索器,这样每个改写问题都会走混合检索
+    multi_query_retriever = MultiQueryRetriever.from_llm(
+        retriever=ensemble_retriever,
+        llm=llm,
+        prompt=CUSTOM_PROMPT
+    )
+    return multi_query_retriever
+
+
+def ask(query: str, retriever, llm) -> str:
+    """检索 + 生成回答"""
+    relevant_docs = retriever.invoke(query)
+    context = "\n\n---\n\n".join(d.page_content for d in relevant_docs)
+
+    prompt = ChatPromptTemplate.from_template(PROMPT_TEMPLATE)
+    chain = prompt | llm | StrOutputParser()
+
+
+    return chain.invoke({"context": context, "question": query})
+
+
+def clean_pdf_text(text: str) -> str:
+    """清洗 PDF 解析出的文本,去除常见噪声"""
+    import re
+
+    # 删除非中文字符之间的换行符
+    text = re.sub(r'[^一](\n)[^一]',
+                  lambda m: m.group(0).replace('\n', ''), text)
+
+    # 删除项目符号和多余空格
+    text = text.replace('•', '').replace('  ', ' ')
+
+    # 删除连续的换行符(保留一个)
+    text = re.sub(r'\n{2,}', '\n', text)
+
+    return text.strip()
+
+
+def clean_documents(pages: list) -> list:
+    """对 Document 列表逐个清洗 page_content,返回清洗后的新 Document 列表"""
+    for page in pages:
+        page.page_content = clean_pdf_text(page.page_content)
+    return pages
+
+
+def main():
+    config = load_config()
+
+    llm = init_llm(config)
+    embeddings = init_embeddings(config)
+    print("✅ 模型初始化完成")
+
+    vectorstore = build_or_load_vectorstore(embeddings)
+    # retriever = vectorstore.as_retriever(search_kwargs={"k": 3})
+
+    # ensemble_retriever = build_ensemble_retriever(embeddings)
+
+    multi_query_retriever = build_multi_query_retriever(embeddings, llm)
+
+    query = "什么是工作质量考核标准?"
+    answer = ask(query, multi_query_retriever, llm)
+
+    print("\n" + "=" * 60)
+    print(f"问题:{query}")
+    print("=" * 60)
+    print(answer)
+
+
+if __name__ == "__main__":
+    main()

+ 139 - 0
003-rag-optimized/schema.py

@@ -0,0 +1,139 @@
+from concurrent.futures import ThreadPoolExecutor
+from langchain_core.output_parsers import StrOutputParser
+from langchain_core.prompts import PromptTemplate
+import json
+
+class RAGWithQueryRewriting:
+    """集成查询重写的 RAG 系统"""
+
+    def __init__(self, retriever, llm, rewrite_chain):
+        self.retriever = retriever
+        self.llm = llm
+        self.rewrite_chain = rewrite_chain
+
+    def invoke(self, question):
+        # 第一步:重写查询
+        rewritten_query = self.rewrite_chain.invoke({"query": question})
+        print(f"原始问题: {question}")
+        print(f"重写后:   {rewritten_query}")
+
+        # 第二步:用重写后的查询进行检索
+        docs = self.retriever.invoke(rewritten_query)
+
+        # 第三步:用原始问题 + 检索结果生成答案
+        context = "\n\n".join([doc.page_content for doc in docs])
+        answer_prompt = f"""基于以下上下文回答用户问题。如果上下文中没有相关信息,请说明。
+
+上下文:{context}
+
+用户问题:{question}
+答案:"""
+        answer = self.llm.invoke(answer_prompt).content
+
+        return {
+            "original_query": question,
+            "rewritten_query": rewritten_query,
+            "retrieved_docs": docs,
+            "answer": answer
+        }
+
+
+class RAGWithDecomposition:
+    """集成查询分解的 RAG 系统"""
+
+    def __init__(self, retriever, llm):
+        self.retriever = retriever
+        self.llm = llm
+
+    def invoke(self, question):
+
+        decompose_prompt = PromptTemplate(
+            input_variables=["question"],
+            template="""你是一个问题分解助手。请将用户的复杂问题分解为 2-4 个独立的子问题,
+        每个子问题应该能独立检索和回答。
+
+        要求:
+        1. 子问题之间互不依赖,可以并行检索
+        2. 子问题覆盖原始问题的所有方面
+        3. 每个子问题简洁明确
+        4. 以 JSON 数组格式输出
+
+        用户问题: {question}
+
+        输出格式示例: ["子问题1", "子问题2", "子问题3"]
+
+        子问题列表:"""
+        )
+
+        decompose_chain = decompose_prompt | self.llm | StrOutputParser()
+
+
+
+
+        # 第一步:分解问题
+        sub_queries = self.decompose_query(question, decompose_chain)
+        print(f"分解为 {len(sub_queries)} 个子问题")
+
+        # 第二步:并行检索并合并
+        all_docs = self.parallel_retrieve_and_merge(sub_queries, self.retriever)
+        print(f"合并后共 {len(all_docs)} 个文档片段")
+
+        # 第三步:用所有上下文生成综合答案
+        context = "\n\n".join([doc.page_content for doc in all_docs])
+        answer_prompt = f"""基于以下上下文,全面回答用户的问题。请综合所有相关信息,给出完整、有条理的答案。上下文:{context}
+
+用户问题:{question}
+答案:"""
+        answer = self.llm.invoke(answer_prompt).content
+
+        return {
+            "question": question,
+            "sub_queries": sub_queries,
+            "doc_count": len(all_docs),
+            "answer": answer
+        }
+
+
+    # 使用 ThreadPoolExecutor 创建线程池,并行执行检索
+    # 每个子问题独立检索,互不干扰
+    # 按内容哈希值去重(避免重复文档)
+
+    def parallel_retrieve_and_merge(self, sub_queries, retriever, k_per_query=5):
+        """
+        对每个子问题并行检索,然后合并去重
+        """
+        all_docs = []
+        seen_contents = set()
+
+        def retrieve_one(query):
+            return retriever.invoke(query)
+
+        # 创建线程池,线程数等于子问题数量
+        with ThreadPoolExecutor(max_workers=len(sub_queries)) as executor:
+            # 提交所有检索任务
+            futures = [executor.submit(retrieve_one, q) for q in sub_queries]
+            ## 收集结果
+            for future in futures:
+                docs = future.result()
+                for doc in docs:
+                    # 按内容去重
+                    content_hash = hash(doc.page_content)
+                    if content_hash not in seen_contents:
+                        seen_contents.add(content_hash)
+                        all_docs.append(doc)
+
+        return all_docs
+
+    def decompose_query(self, question: str, decompose_chain) -> list[str]:
+        """将复杂问题分解为子问题"""
+        result = decompose_chain.invoke({"question": question})
+        # 解析 JSON 数组
+        try:
+            sub_queries = json.loads(result.strip())
+            return sub_queries
+        except json.JSONDecodeError:
+            # 兜底:按行分割
+            return [line.strip() for line in result.strip().split("\n") if line.strip()]
+
+
+