1. 参考 #

1.1 安装依赖 #

# AI/LLM 相关
uv add "openai>=2.53.0" "chromadb>=1.5.9" "langchain-text-splitters>=1.1.2" "sentence-transformers>=5.7.0"

# 文档处理
uv add "pymupdf>=1.28.0" "python-docx>=1.2.0" "python-pptx>=1.0.2" "openpyxl>=3.1.5"

# 工具库
uv add "beautifulsoup4>=4.15.0" "rich>=15.0.0"

依赖说明

依赖包 版本要求 分类 作用说明
openai >=2.53.0 AI/LLM 调用 OpenAI 兼容接口,完成大语言模型问答生成
chromadb >=1.5.9 AI/LLM 本地向量数据库,持久化存储文档块嵌入并支持相似度检索
langchain-text-splitters >=1.1.2 AI/LLM 按字符/语义规则切分长文档,生成适合检索的文本块
sentence-transformers >=5.7.0 AI/LLM 加载多语言嵌入模型,将文本编码为向量
pymupdf >=1.28.0 文档处理 解析 PDF,提取正文文本用于入库
python-docx >=1.2.0 文档处理 解析 Word(.docx)文档内容
python-pptx >=1.0.2 文档处理 解析 PowerPoint(.pptx)幻灯片文本
openpyxl >=3.1.5 文档处理 读取 Excel(.xlsx)工作表数据
beautifulsoup4 >=4.15.0 工具库 解析 HTML/XML,清洗网页类文档中的标签与噪声
rich >=15.0.0 工具库 增强终端输出,支持彩色、表格与富文本展示

1.2 测试命令 #

1.2.1 入库(ingest) #

uv run cli.py ingest --path handbook.md

1.2.2 问答(query) #

uv run cli.py query --question 公司年假怎么申请?

建议先执行 ingest 完成入库,再执行 query 验证检索与生成是否正常。

1.3 参考链接 #

2. 命令行入口(cli.py) #

本节介绍 RAG 项目的命令行入口:如何用 argparse 定义子命令、解析参数,并在 main 中统一调度与异常处理。当前实现已支持 ingest(入库)子命令,后续可按相同模式扩展 ask 等问答能力。

整体流程:用户在终端输入命令 → main 构建解析器并解析参数 → 按子命令调用对应处理函数 → 返回退出码。

sequenceDiagram autonumber actor User as 用户 participant Main as main() participant Parser as build_parser() participant Ingest as cmd_ingest() User->>Main: python cli.py ingest --path handbook.md Main->>Parser: 构建 ArgumentParser / 子命令 Parser-->>Main: parser Main->>Parser: parse_args() Parser-->>Main: args(command=ingest, path=..., func=cmd_ingest) Main->>Main: print(args) Main->>Ingest: args.func(args) Ingest->>Ingest: json.dumps({"file": path}) Ingest-->>Main: 打印 JSON 结果 alt 正常结束 Main-->>User: 退出码 0 else KeyboardInterrupt Main-->>User: 提示「已中断」,退出码 130 else 其他异常 Main->>Main: logger.exception(...) Main-->>User: 打印错误,退出码 1 end

2.1 cli.py #

# 导入日志模块,用于记录程序运行信息
import logging
# 导入系统模块,用于访问标准错误输出等系统功能
import sys
# 导入命令行参数解析模块
import argparse
# 导入 JSON 模块,用于序列化输出数据
import json
# 从 rich 库导入增强版 print,支持彩色与富文本输出
from rich import print

# 配置全局日志的基本格式与级别
logging.basicConfig(
    # 设置日志级别为 INFO,输出信息及以上级别日志
    level=logging.INFO,
    # 定义日志输出格式:时间、级别、记录器名称与消息内容
    format="%(asctime)s [%(levelname)s] %(name)s - %(message)s",
    # 定义时间戳的显示格式为年-月-日 时:分:秒
    datefmt="%Y-%m-%d %H:%M:%S",
# 结束 logging.basicConfig 的参数配置
)
# 获取当前模块对应的日志记录器实例
logger = logging.getLogger(__name__)


# 定义 ingest 子命令的处理函数,接收解析后的参数对象
def cmd_ingest(args):
    # 将文件路径封装为字典并以格式化 JSON 打印输出
    print(json.dumps({"file": args.path}, ensure_ascii=False, indent=2))


# 构建并返回命令行参数解析器
def build_parser():
    # 创建顶层 ArgumentParser,设置程序描述与帮助格式化器
    parser = argparse.ArgumentParser(
        # 设置程序用途说明:RAG 文件入库与问答
        description="RAG:文件入库 + 问答",
        # 使用原始描述帮助格式化器,保留描述中的换行与缩进
        formatter_class=argparse.RawDescriptionHelpFormatter,
    # 结束 ArgumentParser 的参数配置
    )
    # 添加子命令解析器,结果写入 args.command,且必须指定子命令
    sub = parser.add_subparsers(dest="command", required=True)

    # 添加名为 ingest 的子命令,用于文件入库
    p_ingest = sub.add_parser("ingest", help="入库文件")
    # 为 ingest 子命令添加必填的 --path 参数,表示待入库文件路径
    p_ingest.add_argument("--path", required=True, help="文件路径")
    # 将 ingest 子命令的默认处理函数绑定为 cmd_ingest
    p_ingest.set_defaults(func=cmd_ingest)

    # 返回配置完成的参数解析器
    return parser


# 定义程序主入口函数
def main():
    # 调用 build_parser 构建命令行参数解析器
    parser = build_parser()
    # 解析命令行参数,得到命名空间对象 args
    args = parser.parse_args()
    # 打印解析得到的参数对象,便于调试查看
    print(args)
    # 尝试执行子命令对应的处理函数
    try:
        # 调用当前子命令绑定的处理函数并传入参数
        args.func(args)
        # 执行成功时返回退出码 0
        return 0
    # 捕获用户按下 Ctrl+C 产生的键盘中断异常
    except KeyboardInterrupt:
        # 打印已中断提示信息
        print("\n已中断")
        # 返回标准的键盘中断退出码 130
        return 130
    # 捕获其余所有异常(忽略过宽捕获的静态检查警告)
    except Exception as exc:  # noqa: BLE001
        # 记录完整异常堆栈,提示执行失败
        logger.exception("执行失败")
        # 将错误信息输出到标准错误流
        print(f"[错误] {exc}", file=sys.stderr)
        # 返回表示失败的退出码 1
        return 1


# 判断当前模块是否作为脚本直接运行
if __name__ == "__main__":
    # 调用主函数启动程序
    main()

3. 文档加载与入库编排(loader + pipeline) #

本节在第 1 节 CLI 骨架之上,补齐「读文件 → 解析文本」的核心能力:DocumentLoader 按扩展名分派到各类解析器并做空白清洗;RAGPipeline 作为编排层校验路径、调用加载器并输出结果;cli.pycmd_ingest 改为委托 RAGPipeline.ingest_file,从而把命令行与具体解析逻辑解耦。

整体流程:ingest 子命令 → RAGPipeline.ingest_fileDocumentLoader.load(校验 / 选解析器 / 规范化)→ 打印文本;空文本时返回 0

sequenceDiagram autonumber actor User as 用户 participant CLI as cmd_ingest() participant Pipe as RAGPipeline participant Loader as DocumentLoader participant Ext as 对应 _load_* 解析器 User->>CLI: ingest --path handbook.md CLI->>Pipe: RAGPipeline() / ingest_file(path) Pipe->>Pipe: Path.resolve(),校验 is_file alt 路径不是文件 Pipe-->>CLI: raise ValueError else 路径合法 Pipe->>Loader: load(path) Loader->>Loader: 校验 exists / is_file,取 suffix alt 不支持的扩展名 Loader-->>Pipe: raise ValueError else 支持的类型 Loader->>Ext: extractors[ext](path) Ext-->>Loader: 原始文本 Loader->>Loader: _normalize(text) Loader-->>Pipe: 清洗后文本 Pipe->>Pipe: print(text) alt 文本为空 Pipe-->>CLI: return 0 else 文本非空 Pipe-->>CLI: (当前未继续分块/入库) end end end CLI-->>User: 打印 {"file": path} JSON

3.1 loader.py #

loader.py

# 导入 csv 模块,用于解析 CSV 文件内容
import csv

# 导入 json 模块,用于解析与序列化 JSON 数据
import json

# 导入 logging 模块,用于记录文档解析过程中的日志信息
import logging

# 从 pathlib 导入 Path,用于处理与校验文件路径
from pathlib import Path

# 获取名为 rag 的日志记录器实例
logger = logging.getLogger("rag")

# 定义当前加载器支持的文件扩展名集合
SUPPORTED_EXTENSIONS = {
    # PDF 文档扩展名
    ".pdf",
    # Word 新版文档扩展名
    ".docx",
    # Word 旧版文档扩展名
    ".doc",
    # Excel 新版表格扩展名
    ".xlsx",
    # Excel 旧版表格扩展名
    ".xls",
    # PowerPoint 新版演示文稿扩展名
    ".pptx",
    # PowerPoint 旧版演示文稿扩展名
    ".ppt",
    # HTML 网页扩展名
    ".html",
    # HTML 网页简写扩展名
    ".htm",
    # XML 文档扩展名
    ".xml",
    # CSV 表格扩展名
    ".csv",
    # JSON 数据扩展名
    ".json",
    # Markdown 文档扩展名
    ".md",
    # 纯文本扩展名
    ".txt",
    # JSON Lines 文本扩展名
    ".jsonl",
}


# 定义多格式非结构化文档加载器类
class DocumentLoader:
    # 类文档字符串:说明该类用于加载多种格式的非结构化文档
    """多格式非结构化文档加载器。"""

    # 定义文档加载入口方法,接收文件路径并返回清洗后的文本
    def load(self, file_path):
        # 将输入路径转换为 Path 对象,便于后续路径操作
        path = Path(file_path)
        # 判断路径是否真实存在
        if not path.exists():
            # 路径不存在时抛出文件未找到异常
            raise FileNotFoundError(f"文件不存在: {path}")
        # 判断路径是否指向普通文件
        if not path.is_file():
            # 路径不是文件时抛出参数错误
            raise ValueError(f"路径不是文件: {path}")

        # 获取文件扩展名并转换为小写,统一后续匹配逻辑
        ext = path.suffix.lower()
        # 记录当前正在解析的文档名称与类型
        logger.info("解析文档: %s (type=%s)", path.name, ext or "unknown")

        # 构建扩展名到对应解析方法的映射表
        extractors = {
            # PDF 文件使用 PDF 解析方法
            ".pdf": self._load_pdf,
            # DOCX 文件使用 Word 解析方法
            ".docx": self._load_docx,
            # DOC 文件同样使用 Word 解析方法
            ".doc": self._load_docx,
            # XLSX 文件使用 Excel 解析方法
            ".xlsx": self._load_excel,
            # XLS 文件同样使用 Excel 解析方法
            ".xls": self._load_excel,
            # PPTX 文件使用 PowerPoint 解析方法
            ".pptx": self._load_pptx,
            # PPT 文件同样使用 PowerPoint 解析方法
            ".ppt": self._load_pptx,
            # HTML 文件使用 HTML 解析方法
            ".html": self._load_html,
            # HTM 文件同样使用 HTML 解析方法
            ".htm": self._load_html,
            # XML 文件使用 XML 解析方法
            ".xml": self._load_xml,
            # CSV 文件使用 CSV 解析方法
            ".csv": self._load_csv,
            # JSON 文件使用 JSON 解析方法
            ".json": self._load_json,
            # Markdown 文件使用纯文本解析方法
            ".md": self._load_plain,
            # TXT 文件使用纯文本解析方法
            ".txt": self._load_plain,
            # JSONL 文件使用纯文本解析方法
            ".jsonl": self._load_plain,
        }
        # 根据扩展名查找对应的解析函数
        extractor = extractors.get(ext)
        # 判断是否找到可用的解析函数
        if extractor is None:
            # 未找到时抛出不支持的文件类型错误,并列出全部支持扩展名
            raise ValueError(
                f"不支持的文件类型: {ext},支持: {', '.join(sorted(SUPPORTED_EXTENSIONS))}"
            )

        # 调用对应解析函数提取文档原始文本
        text = extractor(path)
        # 对提取结果做空白清洗与规范化处理
        text = self._normalize(text)
        # 判断清洗后的文本是否为空
        if not text:
            # 文本为空时记录警告日志
            logger.warning("文档内容为空: %s", path)
        # 文本非空时进入成功分支
        else:
            # 记录解析完成日志,并输出文档字符数
            logger.info("解析完成: %s,字符数=%d", path.name, len(text))
        # 返回最终可用于后续分块与入库的文本
        return text

    # 将方法声明为静态方法,不依赖实例状态
    @staticmethod
    # 定义文本规范化方法,用于清洗多余空白
    def _normalize(text):
        # 方法文档字符串:说明清洗目的是降低噪声对分块与检索的影响
        """清洗多余空白,降低噪声对分块与检索的影响。"""
        # 按行拆分文本,并对每一行去除首尾空白
        lines = [line.strip() for line in text.splitlines()]
        # 过滤掉清洗后变成空字符串的行
        lines = [line for line in lines if line]
        # 用换行符重新拼接有效行并返回
        return "\n".join(lines)

    # 将方法声明为静态方法,不依赖实例状态
    @staticmethod
    # 定义 PDF 文档解析方法
    def _load_pdf(path):
        # 延迟导入 PyMuPDF,避免未使用 PDF 时增加启动开销
        import fitz  # PyMuPDF

        # 打开 PDF 文件并确保使用完后自动关闭
        with fitz.open(path) as pdf:
            # 逐页提取纯文本,并用换行符拼接后返回
            return "\n".join(page.get_text("text") for page in pdf)

    # 将方法声明为静态方法,不依赖实例状态
    @staticmethod
    # 定义 Word 文档解析方法
    def _load_docx(path):
        # 延迟导入 python-docx 的 Document 类
        from docx import Document

        # 根据路径打开 Word 文档对象
        doc = Document(str(path))
        # 提取所有非空段落文本,作为正文主体
        parts = [p.text for p in doc.paragraphs if p.text.strip()]
        # 遍历文档中的全部表格
        for table in doc.tables:
            # 遍历表格中的每一行
            for row in table.rows:
                # 提取当前行中非空单元格文本,并去除首尾空白
                cells = [cell.text.strip() for cell in row.cells if cell.text.strip()]
                # 判断当前行是否存在有效单元格内容
                if cells:
                    # 将单元格内容用制表符拼接后追加到结果列表
                    parts.append("\t".join(cells))
        # 用换行符拼接段落与表格内容后返回
        return "\n".join(parts)

    # 将方法声明为静态方法,不依赖实例状态
    @staticmethod
    # 定义 Excel 表格解析方法
    def _load_excel(path):
        # 延迟导入 openpyxl,用于读取 Excel 工作簿
        import openpyxl

        # 以只读计算值模式加载工作簿
        wb = openpyxl.load_workbook(str(path), data_only=True)
        # 使用 try/finally 确保工作簿最终会被关闭
        try:
            # 初始化用于收集各工作表文本行的列表
            rows = []
            # 遍历工作簿中的每一个工作表
            for sheet in wb.worksheets:
                # 先追加当前工作表名称作为分隔标记
                rows.append(f"[Sheet: {sheet.title}]")
                # 按行迭代工作表中的全部单元格值
                for row in sheet.iter_rows(values_only=True):
                    # 将单元格值转为字符串,空值统一转为空串
                    cells = [str(c) if c is not None else "" for c in row]
                    # 判断当前行是否存在任意非空内容
                    if any(cells):
                        # 将有效行用制表符拼接后追加到结果列表
                        rows.append("\t".join(cells))
            # 用换行符拼接全部工作表内容后返回
            return "\n".join(rows)
        # 无论解析是否成功,都执行收尾逻辑
        finally:
            # 关闭工作簿,释放文件句柄与相关资源
            wb.close()

    # 将方法声明为静态方法,不依赖实例状态
    @staticmethod
    # 定义 PowerPoint 演示文稿解析方法
    def _load_pptx(path):
        # 延迟导入 python-pptx 的 Presentation 类
        from pptx import Presentation

        # 根据路径打开演示文稿对象
        ppt = Presentation(str(path))
        # 初始化用于收集全部幻灯片文本的列表
        texts = []
        # 从 1 开始枚举每一页幻灯片
        for i, slide in enumerate(ppt.slides, start=1):
            # 提取当前幻灯片中所有带文本且非空的形状内容
            slide_texts = [
                # 去除形状文本首尾空白
                shape.text.strip()
                # 遍历当前幻灯片中的全部形状
                for shape in slide.shapes
                # 仅保留具备 text 属性且文本非空的形状
                if hasattr(shape, "text") and shape.text and shape.text.strip()
            ]
            # 判断当前幻灯片是否提取到有效文本
            if slide_texts:
                # 追加幻灯片序号标记,便于区分来源页
                texts.append(f"[Slide {i}]")
                # 将当前页提取到的文本逐条追加到总结果中
                texts.extend(slide_texts)
        # 用换行符拼接全部幻灯片文本后返回
        return "\n".join(texts)

    # 将方法声明为静态方法,不依赖实例状态
    @staticmethod
    # 定义 HTML 网页解析方法
    def _load_html(path):
        # 延迟导入 BeautifulSoup,用于解析 HTML 结构
        from bs4 import BeautifulSoup

        # 以 UTF-8 读取 HTML 文件内容,忽略无法解码的字符
        html = path.read_text(encoding="utf-8", errors="ignore")
        # 使用 lxml 解析器构建 BeautifulSoup 文档树
        soup = BeautifulSoup(html, "lxml")
        # 遍历并移除脚本、样式与 noscript 等噪声标签
        for tag in soup(["script", "style", "noscript"]):
            # 从文档树中彻底删除当前噪声标签
            tag.decompose()
        # 提取纯文本,行间用换行分隔,并自动去除多余空白
        return soup.get_text(separator="\n", strip=True)

    # 将方法声明为静态方法,不依赖实例状态
    @staticmethod
    # 定义 XML 文档解析方法
    def _load_xml(path):
        # 延迟导入 lxml.etree,用于解析 XML 树结构
        from lxml import etree

        # 解析 XML 文件并获取根节点
        root = etree.parse(str(path)).getroot()
        # 遍历全部文本节点,清洗后用空格拼接成一段文本返回
        return " ".join(t.strip() for t in root.itertext() if t and t.strip())

    # 将方法声明为静态方法,不依赖实例状态
    @staticmethod
    # 定义 CSV 文件解析方法
    def _load_csv(path):
        # 以只读方式打开 CSV 文件,忽略解码错误并禁用换行转换
        with path.open("r", encoding="utf-8", errors="ignore", newline="") as f:
            # 创建 CSV 读取器,按行解析字段
            reader = csv.reader(f)
            # 将每一行字段用逗号加空格拼接,再按行拼接成完整文本返回
            return "\n".join(", ".join(row) for row in reader)

    # 将方法声明为静态方法,不依赖实例状态
    @staticmethod
    # 定义 JSON 文件解析方法
    def _load_json(path):
        # 读取文件内容并反序列化为 Python 对象
        data = json.loads(path.read_text(encoding="utf-8"))
        # 将对象重新格式化为缩进 JSON 字符串后返回,保留非 ASCII 字符
        return json.dumps(data, ensure_ascii=False, indent=2)

    # 将方法声明为静态方法,不依赖实例状态
    @staticmethod
    # 定义纯文本类文件解析方法,适用于 md/txt/jsonl 等
    def _load_plain(path):
        # 以 UTF-8 读取文件全文,忽略无法解码的字符后直接返回
        return path.read_text(encoding="utf-8", errors="ignore")

3.2 pipeline.py #

pipeline.py

# 从 pathlib 导入 Path,用于处理与解析文件路径
from pathlib import Path

# 从 loader 模块导入 DocumentLoader,用于加载多格式非结构化文档
from loader import DocumentLoader

# 从 rich 库导入增强版 print,支持彩色与富文本输出
from rich import print


# 定义 RAG 主流程编排器类
class RAGPipeline:
    # 类文档字符串:说明该类负责编排 RAG 主流程
    """RAG 主流程编排器。"""

    # 初始化方法,可选接收配置参数
    def __init__(self, config=None):
        # 创建文档加载器实例并保存为实例属性
        self.loader = DocumentLoader()

    # 定义文件入库方法,接收待入库的文件路径
    def ingest_file(self, file_path):
        # 将输入路径转换为绝对路径对象
        path = Path(file_path).resolve()
        # 校验路径是否指向一个真实存在的文件
        if not path.is_file():
            # 若不是文件则抛出参数错误,提示需指定单个文件路径
            raise ValueError(f"请指定单个文件路径: {path}")
        # 调用文档加载器解析文件内容为文本
        text = self.loader.load(path)
        # 打印解析得到的文档文本内容
        print(text)
        # 判断清洗后的文本是否为空
        if not text.strip():
            # 文本为空时返回 0,表示未生成有效内容
            return 0

3.3 cli.py #

cli.py

# 导入日志模块,用于记录程序运行信息
import logging

# 导入系统模块,用于访问标准错误输出等系统功能
import sys

# 导入命令行参数解析模块
import argparse

# 导入 JSON 模块,用于序列化输出数据
import json

# 从 rich 库导入增强版 print,支持彩色与富文本输出
from rich import print

# 导入 RAG 主流程编排器
+from pipeline import RAGPipeline

# 配置全局日志的基本格式与级别
logging.basicConfig(
    # 设置日志级别为 INFO,输出信息及以上级别日志
    level=logging.INFO,
    # 定义日志输出格式:时间、级别、记录器名称与消息内容
    format="%(asctime)s [%(levelname)s] %(name)s - %(message)s",
    # 定义时间戳的显示格式为年-月-日 时:分:秒
    datefmt="%Y-%m-%d %H:%M:%S",
    # 结束 logging.basicConfig 的参数配置
)
# 获取当前模块对应的日志记录器实例
logger = logging.getLogger(__name__)


# 定义 ingest 子命令的处理函数,接收解析后的参数对象
def cmd_ingest(args):
    # 创建 RAG 主流程编排器实例
+   pipeline = RAGPipeline()
    # 调用 ingest_file 方法入库文件
+   chunks = pipeline.ingest_file(args.path)
    # 将文件路径封装为字典并以格式化 JSON 打印输出
+   print(print(json.dumps({"file": args.path}, ensure_ascii=False, indent=2)))


# 构建并返回命令行参数解析器
def build_parser():
    # 创建顶层 ArgumentParser,设置程序描述与帮助格式化器
    parser = argparse.ArgumentParser(
        # 设置程序用途说明:RAG 文件入库与问答
        description="RAG:文件入库 + 问答",
        # 使用原始描述帮助格式化器,保留描述中的换行与缩进
        formatter_class=argparse.RawDescriptionHelpFormatter,
        # 结束 ArgumentParser 的参数配置
    )
    # 添加子命令解析器,结果写入 args.command,且必须指定子命令
    sub = parser.add_subparsers(dest="command", required=True)

    # 添加名为 ingest 的子命令,用于文件入库
    p_ingest = sub.add_parser("ingest", help="入库文件")
    # 为 ingest 子命令添加必填的 --path 参数,表示待入库文件路径
    p_ingest.add_argument("--path", required=True, help="文件路径")
    # 将 ingest 子命令的默认处理函数绑定为 cmd_ingest
    p_ingest.set_defaults(func=cmd_ingest)

    # 返回配置完成的参数解析器
    return parser


# 定义程序主入口函数
def main():
    # 调用 build_parser 构建命令行参数解析器
    parser = build_parser()
    # 解析命令行参数,得到命名空间对象 args
    args = parser.parse_args()
    # 打印解析得到的参数对象,便于调试查看
    print(args)
    # 尝试执行子命令对应的处理函数
    try:
        # 调用当前子命令绑定的处理函数并传入参数
        args.func(args)
        # 执行成功时返回退出码 0
        return 0
    # 捕获用户按下 Ctrl+C 产生的键盘中断异常
    except KeyboardInterrupt:
        # 打印已中断提示信息
        print("\n已中断")
        # 返回标准的键盘中断退出码 130
        return 130
    # 捕获其余所有异常(忽略过宽捕获的静态检查警告)
    except Exception as exc:  # noqa: BLE001
        # 记录完整异常堆栈,提示执行失败
        logger.exception("执行失败")
        # 将错误信息输出到标准错误流
        print(f"[错误] {exc}", file=sys.stderr)
        # 返回表示失败的退出码 1
        return 1


# 判断当前模块是否作为脚本直接运行
if __name__ == "__main__":
    # 调用主函数启动程序
    main()

4. 向量化与向量库 #

本节在「能读文件」的基础上,引入可替换的运行时配置与向量基础设施:.env / RAGConfig 统一管理库路径、集合名与嵌入模型;EmbeddingService 封装本地 SentenceTransformer(本节先落模型名);VectorStore 封装 ChromaDB,并提供 reset 清空重建集合。RAGPipeline 初始化时组装 config → embedder → store,入库在加载到非空文本后调用 store.reset(),为后续写入向量腾出干净集合。

整体流程:创建 Pipeline → 从环境变量加载配置 → 构造嵌入服务与向量库 → ingest_file 加载文本 → 非空则重置集合。

sequenceDiagram autonumber actor User as 用户 participant CLI as cmd_ingest() participant Pipe as RAGPipeline participant Cfg as RAGConfig participant Emb as EmbeddingService participant Store as VectorStore participant Chroma as ChromaDB Client User->>CLI: ingest --path handbook.md CLI->>Pipe: RAGPipeline() Pipe->>Cfg: from_env()(RAG_DB_PATH / COLLECTION / EMBEDDING_MODEL) Cfg-->>Pipe: config Pipe->>Emb: EmbeddingService(model_name) Emb-->>Pipe: embedder Pipe->>Store: VectorStore(db_path, collection, embedder) Store-->>Pipe: store CLI->>Pipe: ingest_file(path) Pipe->>Pipe: loader.load(path) → text alt 文本为空 Pipe-->>CLI: return 0 else 文本非空 Pipe->>Store: reset() Store->>Chroma: delete_collection(name)(失败则忽略) Store->>Chroma: get_or_create_collection(cosine / metadata) Chroma-->>Store: 空集合 Store-->>Pipe: 重建完成 end CLI-->>User: 打印 {"file": path} JSON

4.1 .env #

.env

# 向量数据库存储路径,默认值为当前目录下的 chroma_db
RAG_DB_PATH="./chroma_db"
# 向量集合名称,默认值为 rag
RAG_COLLECTION="rag"
# 嵌入模型名称,默认值为多语言 MiniLM 模型
RAG_EMBEDDING_MODEL="paraphrase-multilingual-MiniLM-L12-v2"

4.2 config.py #

config.py

# 导入操作系统模块,用于读取环境变量
import os

# 从 dataclasses 模块导入 dataclass,用于定义数据类
from dataclasses import dataclass


# 使用 dataclass 装饰器,并将实例设为不可变(frozen=True)
@dataclass(frozen=True)
# 定义 RAG 系统运行时配置类
class RAGConfig:
    # 类文档字符串:说明该类用于保存 RAG 系统运行时配置
    """RAG 系统运行时配置。"""

    # 向量数据库存储路径,默认值为当前目录下的 chroma_db
    db_path: str = "./chroma_db"
    # 向量集合名称,默认值为 rag
    collection_name: str = "rag"
    # 嵌入模型名称,默认值为多语言 MiniLM 模型
    embedding_model: str = "paraphrase-multilingual-MiniLM-L12-v2"

    # 将 from_env 声明为类方法
    @classmethod
    # 定义从环境变量加载配置的类方法
    def from_env(cls):
        # 使用环境变量(若不存在则回退到类默认值)创建并返回配置实例
        return cls(
            # 从环境变量 RAG_DB_PATH 读取数据库路径,缺省使用类默认值
            db_path=os.getenv("RAG_DB_PATH", cls.db_path),
            # 从环境变量 RAG_COLLECTION 读取集合名称,缺省使用类默认值
            collection_name=os.getenv("RAG_COLLECTION", cls.collection_name),
            # 从环境变量 RAG_EMBEDDING_MODEL 读取嵌入模型名称,缺省使用类默认值
            embedding_model=os.getenv("RAG_EMBEDDING_MODEL", cls.embedding_model),
        )

4.3 embeddings.py #

embeddings.py

# 导入日志模块,用于记录程序运行信息
import logging

# 从 sentence_transformers 库导入 SentenceTransformer 模型类
from sentence_transformers import SentenceTransformer

# 获取名为 "rag" 的日志记录器实例
logger = logging.getLogger("rag")


# 定义 EmbeddingService 类,封装本地向量化能力
class EmbeddingService:
    # 类说明文档字符串:本地 SentenceTransformer 向量化服务
    """本地 SentenceTransformer 向量化服务。"""

    # 定义初始化方法,接收模型名称参数
    def __init__(self, model_name):
        # 将传入的模型名称保存为实例属性,供后续加载使用
        self.model_name = model_name

4.4 store.py #

store.py

# 导入日志模块,用于记录程序运行信息
import logging

# 从 datetime 模块导入 datetime 与 timezone,用于生成带时区的时间戳
from datetime import datetime, timezone

# 从 pathlib 模块导入 Path,用于处理文件路径
from pathlib import Path

# 获取当前模块对应的日志记录器实例
logger = logging.getLogger(__name__)


# 定义向量存储封装类
class VectorStore:
    # 类文档字符串:说明该类用于封装 ChromaDB 向量库
    """ChromaDB 向量库封装。"""

    # 定义重置集合的实例方法
    def reset(self):
        # 方法文档字符串:说明该方法用于清空并重建集合
        """清空并重建集合。"""
        # 尝试删除已有集合,忽略删除失败的情况
        try:
            # 调用客户端删除当前集合名称对应的集合
            self._client.delete_collection(self.collection_name)
        # 捕获所有异常(忽略过宽捕获的静态检查警告)
        except Exception:  # noqa: BLE001
            # 删除失败时忽略异常,继续后续重建流程
            pass
        # 获取或创建集合,并将结果赋值给实例的集合属性
        self._collection = self._client.get_or_create_collection(
            # 指定集合名称为当前配置的集合名
            name=self.collection_name,
            # 设置集合元数据字典
            metadata={
                # 指定 HNSW 索引使用余弦相似度空间
                "hnsw:space": "cosine",
                # 设置集合描述为 RAG 知识库
                "description": "RAG 知识库",
                # 记录集合创建时间为当前 UTC 时间的 ISO 格式字符串
                "created_at": datetime.now(timezone.utc).isoformat(),
                # 结束 metadata 字典
            },
            # 结束 get_or_create_collection 调用
        )
        # 记录已重建空集合的信息日志,并输出集合名称
        logger.info("已重建空集合: %s", self.collection_name)

4.5 cli.py #

cli.py

# 导入日志模块,用于记录程序运行信息
import logging

# 导入系统模块,用于访问标准错误输出等系统功能
import sys

# 导入命令行参数解析模块
import argparse

# 导入 JSON 模块,用于序列化输出数据
import json

# 从 rich 库导入增强版 print,支持彩色与富文本输出
from rich import print

# 导入 RAG 主流程编排器
from pipeline import RAGPipeline

# 配置全局日志的基本格式与级别
logging.basicConfig(
    # 设置日志级别为 INFO,输出信息及以上级别日志
    level=logging.INFO,
    # 定义日志输出格式:时间、级别、记录器名称与消息内容
    format="%(asctime)s [%(levelname)s] %(name)s - %(message)s",
    # 定义时间戳的显示格式为年-月-日 时:分:秒
    datefmt="%Y-%m-%d %H:%M:%S",
    # 结束 logging.basicConfig 的参数配置
)
# 获取当前模块对应的日志记录器实例
logger = logging.getLogger(__name__)


# 定义 ingest 子命令的处理函数,接收解析后的参数对象
def cmd_ingest(args):
    # 创建 RAG 主流程编排器实例
    pipeline = RAGPipeline()
    # 调用 ingest_file 方法入库文件
+   pipeline.ingest_file(args.path)
    # 将文件路径封装为字典并以格式化 JSON 打印输出
    print(print(json.dumps({"file": args.path}, ensure_ascii=False, indent=2)))


# 构建并返回命令行参数解析器
def build_parser():
    # 创建顶层 ArgumentParser,设置程序描述与帮助格式化器
    parser = argparse.ArgumentParser(
        # 设置程序用途说明:RAG 文件入库与问答
        description="RAG:文件入库 + 问答",
        # 使用原始描述帮助格式化器,保留描述中的换行与缩进
        formatter_class=argparse.RawDescriptionHelpFormatter,
        # 结束 ArgumentParser 的参数配置
    )
    # 添加子命令解析器,结果写入 args.command,且必须指定子命令
    sub = parser.add_subparsers(dest="command", required=True)

    # 添加名为 ingest 的子命令,用于文件入库
    p_ingest = sub.add_parser("ingest", help="入库文件")
    # 为 ingest 子命令添加必填的 --path 参数,表示待入库文件路径
    p_ingest.add_argument("--path", required=True, help="文件路径")
    # 将 ingest 子命令的默认处理函数绑定为 cmd_ingest
    p_ingest.set_defaults(func=cmd_ingest)

    # 返回配置完成的参数解析器
    return parser


# 定义程序主入口函数
def main():
    # 调用 build_parser 构建命令行参数解析器
    parser = build_parser()
    # 解析命令行参数,得到命名空间对象 args
    args = parser.parse_args()
    # 打印解析得到的参数对象,便于调试查看
    print(args)
    # 尝试执行子命令对应的处理函数
    try:
        # 调用当前子命令绑定的处理函数并传入参数
        args.func(args)
        # 执行成功时返回退出码 0
        return 0
    # 捕获用户按下 Ctrl+C 产生的键盘中断异常
    except KeyboardInterrupt:
        # 打印已中断提示信息
        print("\n已中断")
        # 返回标准的键盘中断退出码 130
        return 130
    # 捕获其余所有异常(忽略过宽捕获的静态检查警告)
    except Exception as exc:  # noqa: BLE001
        # 记录完整异常堆栈,提示执行失败
        logger.exception("执行失败")
        # 将错误信息输出到标准错误流
        print(f"[错误] {exc}", file=sys.stderr)
        # 返回表示失败的退出码 1
        return 1


# 判断当前模块是否作为脚本直接运行
if __name__ == "__main__":
    # 调用主函数启动程序
    main()

4.6 pipeline.py #

pipeline.py

# 从 pathlib 导入 Path,用于处理与解析文件路径
from pathlib import Path

# 从 loader 模块导入 DocumentLoader,用于加载多格式非结构化文档
from loader import DocumentLoader

# 从 rich 库导入增强版 print,支持彩色与富文本输出
from rich import print

# 从 store 模块导入 VectorStore,用于向量存储与检索
+from store import VectorStore

# 从 config 模块导入 RAGConfig,用于配置 RAG 系统
+from config import RAGConfig

# 从 embeddings 模块导入 EmbeddingService,用于向量化服务
+from embeddings import EmbeddingService


# 定义 RAG 主流程编排器类
class RAGPipeline:
    # 类文档字符串:说明该类负责编排 RAG 主流程
    """RAG 主流程编排器。"""

    # 初始化方法,可选接收配置参数
    def __init__(self, config=None):
+       self.config = config or RAGConfig.from_env()
        # 创建文档加载器实例并保存为实例属性
        self.loader = DocumentLoader()
        # 创建向量化服务实例并保存为实例属性
+       self.embedder = EmbeddingService(self.config.embedding_model)
        # 创建向量存储实例并保存为实例属性
+       self.store = VectorStore(
+           db_path=self.config.db_path,  # 数据库路径
+           collection_name=self.config.collection_name,  # 集合名称
+           embedding_service=self.embedder,  # 向量化服务
+       )

    # 定义文件入库方法,接收待入库的文件路径
    def ingest_file(self, file_path):
        # 将输入路径转换为绝对路径对象
        path = Path(file_path).resolve()
        # 校验路径是否指向一个真实存在的文件
        if not path.is_file():
            # 若不是文件则抛出参数错误,提示需指定单个文件路径
            raise ValueError(f"请指定单个文件路径: {path}")
        # 调用文档加载器解析文件内容为文本
        text = self.loader.load(path)
        # 打印解析得到的文档文本内容
        print(text)
        # 判断清洗后的文本是否为空
        if not text.strip():
            # 文本为空时返回 0,表示未生成有效内容
            return 0
        # 重置向量存储
+       self.store.reset()

5. 分块、向量化与入库 #

本节把第 3 节搭好的配置与空集合,接成可真正写入知识库的入库闭环:DocumentChunk 定义块数据模型;TextChunker 用递归字符切分(可配 chunk_size / chunk_overlap)生成带稳定 SHA256 ID 的块;EmbeddingService.embed 懒加载 SentenceTransformer 并批量归一化编码;VectorStore 补齐持久化客户端与 upsert_chunks(按批 embed → Chroma upsert)。RAGPipeline.ingest_file 串联「加载 → reset → 分块 → 写入」,CLI 输出写入块数。

整体流程:加载全文 → 清空集合 → 分块 → 分批向量化并 upsert → 返回写入数量。

sequenceDiagram autonumber actor User as 用户 participant CLI as cmd_ingest() participant Pipe as RAGPipeline participant Chunker as TextChunker participant Store as VectorStore participant Emb as EmbeddingService participant Chroma as ChromaDB User->>CLI: ingest --path handbook.md CLI->>Pipe: ingest_file(path) Pipe->>Pipe: loader.load(path) → text alt 文本为空 Pipe-->>CLI: return 0 else 文本非空 Pipe->>Store: reset()(清空并重建集合) Pipe->>Chunker: split(text, source=文件名) Chunker->>Chunker: RecursiveCharacterTextSplitter → DocumentChunk[](稳定 ID) Chunker-->>Pipe: chunks Pipe->>Store: upsert_chunks(chunks) loop 每批 batch_size=64 Store->>Emb: embed([content...]) Emb->>Emb: 懒加载模型 / encode(normalize) Emb-->>Store: vectors Store->>Chroma: upsert(ids, documents, embeddings, metadatas) end Store-->>Pipe: written = n Pipe-->>CLI: return n end CLI-->>User: JSON {"file": path, "chunks": n}

5.1 chunker.py #

chunker.py

# 导入 hashlib 模块,用于生成稳定的哈希 ID
import hashlib

# 导入 logging 模块,用于记录分块过程日志
import logging

# 从 langchain_text_splitters 导入 RecursiveCharacterTextSplitter,用于递归字符文本切分
from langchain_text_splitters import RecursiveCharacterTextSplitter

# 从 models 模块导入 DocumentChunk,用于表示入库后的文本块
from models import DocumentChunk

# 获取当前模块对应的日志记录器实例
logger = logging.getLogger(__name__)


# 定义文本分块器类
class TextChunker:
    # 类文档字符串:说明该类基于 RecursiveCharacterTextSplitter 的默认分块策略
    """基于 RecursiveCharacterTextSplitter 的默认分块策略。"""

    # 初始化方法,设置分块大小与重叠长度
    def __init__(self, chunk_size=500, chunk_overlap=80):
        # 创建递归字符文本切分器实例并保存为实例属性
        self.splitter = RecursiveCharacterTextSplitter(
            # 设置每个文本块的目标字符长度
            chunk_size=chunk_size,
            # 设置相邻文本块之间的重叠字符数
            chunk_overlap=chunk_overlap,
            # 优先按段落 → 行 → 句子 → 字符递归切分,尽量保住语义完整
            separators=[
                # 优先按空行(段落)切分
                "\n\n",
                # 其次按换行切分
                "\n",
                # 按中文句号切分
                "。",
                # 按中文感叹号切分
                "!",
                # 按中文问号切分
                "?",
                # 按中文分号切分
                ";",
                # 按英文句号加空格切分
                ". ",
                # 按英文感叹号加空格切分
                "! ",
                # 按英文问号加空格切分
                "? ",
                # 按空格切分
                " ",
                # 最后按字符强制切分
                "",
                # 结束分隔符列表
            ],
            # 使用内置 len 函数计算文本长度
            length_function=len,
            # 指定分隔符按普通字符串匹配,而非正则表达式
            is_separator_regex=False,
            # 结束 RecursiveCharacterTextSplitter 构造调用
        )

    # 定义文本切分方法,接收原文与来源标识
    def split(self, text, source):
        # 调用切分器将原文切分为原始文本块列表
        raw_chunks = self.splitter.split_text(text)
        # 初始化用于存放 DocumentChunk 对象的结果列表
        chunks = []
        # 遍历每个原始文本块及其序号索引
        for idx, content in enumerate(raw_chunks):
            # 判断当前文本块去除空白后是否为空
            if not content.strip():
                # 跳过空文本块,不纳入结果
                continue
            # 基于来源、索引与内容生成稳定的块 ID
            chunk_id = self._stable_id(source, idx, content)
            # 将构造好的 DocumentChunk 追加到结果列表
            chunks.append(
                # 创建入库后的文本块数据模型实例
                DocumentChunk(
                    # 设置文本块唯一标识符
                    id=chunk_id,
                    # 设置去除首尾空白后的正文内容
                    content=content.strip(),
                    # 设置文本块来源(如文件名)
                    source=source,
                    # 设置文本块在原文中的序号索引
                    chunk_index=idx,
                    # 设置附加元数据字典
                    metadata={
                        # 记录来源信息
                        "source": source,
                        # 记录块序号
                        "chunk_index": idx,
                        # 记录原始内容字符长度
                        "char_len": len(content),
                        # 结束元数据字典
                    },
                    # 结束 DocumentChunk 构造调用
                )
                # 结束 append 调用
            )
        # 记录分块完成日志,包含来源与块数量
        logger.info("分块完成: source=%s, chunks=%d", source, len(chunks))
        # 返回构造完成的文本块列表
        return chunks

    # 将 _stable_id 声明为静态方法
    @staticmethod
    # 定义生成稳定块 ID 的方法
    def _stable_id(source, index, content):
        # 方法文档字符串:说明稳定 ID 用于同内容重复入库时幂等跳过
        """稳定 ID:同内容重复入库可幂等跳过(不依赖 PYTHONHASHSEED)。"""
        # 对来源、索引与内容拼接后做 SHA256 哈希并取十六进制摘要
        digest = hashlib.sha256(
            # 将来源、索引与内容拼接为 UTF-8 字节串
            f"{source}::{index}::{content}".encode("utf-8")
            # 结束 sha256 调用并转换为十六进制摘要字符串
        ).hexdigest()
        # 返回摘要前 32 个字符作为稳定 ID
        return digest[:32]

5.2 models.py #

models.py

# 从 dataclasses 模块导入 asdict、dataclass 与 field,用于定义数据类与字段默认值
from dataclasses import asdict, dataclass, field


# 使用 dataclass 装饰器,将类自动转换为数据类
@dataclass
# 定义入库后的文本块数据模型类
class DocumentChunk:
    # 类文档字符串:说明该类表示入库后的文本块
    """入库后的文本块。"""

    # 文本块的唯一标识符
    id: str
    # 文本块的正文内容
    content: str
    # 文本块来源(如文件名)
    source: str
    # 文本块在原文中的序号索引
    chunk_index: int
    # 附加元数据字典,默认创建空字典
    metadata: dict = field(default_factory=dict)

5.3 .env #

.env

# 向量数据库存储路径,默认值为当前目录下的 chroma_db
RAG_DB_PATH="./chroma_db"
# 向量集合名称,默认值为 rag
RAG_COLLECTION="rag"
# 嵌入模型名称,默认值为多语言 MiniLM 模型
RAG_EMBEDDING_MODEL="C:/Users/83687/.cache/modelscope/models/Liudef--paraphrase-multilingual-MiniLM-L12-v2/snapshots/master"

5.4 cli.py #

cli.py

# 导入日志模块,用于记录程序运行信息
import logging

# 导入系统模块,用于访问标准错误输出等系统功能
import sys

# 导入命令行参数解析模块
import argparse

# 导入 JSON 模块,用于序列化输出数据
import json

# 从 rich 库导入增强版 print,支持彩色与富文本输出
from rich import print

# 导入 RAG 主流程编排器
from pipeline import RAGPipeline

# 配置全局日志的基本格式与级别
logging.basicConfig(
    # 设置日志级别为 INFO,输出信息及以上级别日志
    level=logging.INFO,
    # 定义日志输出格式:时间、级别、记录器名称与消息内容
    format="%(asctime)s [%(levelname)s] %(name)s - %(message)s",
    # 定义时间戳的显示格式为年-月-日 时:分:秒
    datefmt="%Y-%m-%d %H:%M:%S",
    # 结束 logging.basicConfig 的参数配置
)
# 获取当前模块对应的日志记录器实例
logger = logging.getLogger(__name__)


# 定义 ingest 子命令的处理函数,接收解析后的参数对象
def cmd_ingest(args):
    # 创建 RAG 主流程编排器实例
    pipeline = RAGPipeline()
    # 调用 ingest_file 方法入库文件
+   chunks = pipeline.ingest_file(args.path)
    # 将文件路径封装为字典并以格式化 JSON 打印输出
+   print(
+       print(
+           json.dumps(
+               {"file": args.path, "chunks": chunks}, ensure_ascii=False, indent=2
+           )
+       )
+   )


# 构建并返回命令行参数解析器
def build_parser():
    # 创建顶层 ArgumentParser,设置程序描述与帮助格式化器
    parser = argparse.ArgumentParser(
        # 设置程序用途说明:RAG 文件入库与问答
        description="RAG:文件入库 + 问答",
        # 使用原始描述帮助格式化器,保留描述中的换行与缩进
        formatter_class=argparse.RawDescriptionHelpFormatter,
        # 结束 ArgumentParser 的参数配置
    )
    # 添加子命令解析器,结果写入 args.command,且必须指定子命令
    sub = parser.add_subparsers(dest="command", required=True)

    # 添加名为 ingest 的子命令,用于文件入库
    p_ingest = sub.add_parser("ingest", help="入库文件")
    # 为 ingest 子命令添加必填的 --path 参数,表示待入库文件路径
    p_ingest.add_argument("--path", required=True, help="文件路径")
    # 将 ingest 子命令的默认处理函数绑定为 cmd_ingest
    p_ingest.set_defaults(func=cmd_ingest)

    # 返回配置完成的参数解析器
    return parser


# 定义程序主入口函数
def main():
    # 调用 build_parser 构建命令行参数解析器
    parser = build_parser()
    # 解析命令行参数,得到命名空间对象 args
    args = parser.parse_args()
    # 打印解析得到的参数对象,便于调试查看
    print(args)
    # 尝试执行子命令对应的处理函数
    try:
        # 调用当前子命令绑定的处理函数并传入参数
        args.func(args)
        # 执行成功时返回退出码 0
        return 0
    # 捕获用户按下 Ctrl+C 产生的键盘中断异常
    except KeyboardInterrupt:
        # 打印已中断提示信息
        print("\n已中断")
        # 返回标准的键盘中断退出码 130
        return 130
    # 捕获其余所有异常(忽略过宽捕获的静态检查警告)
    except Exception as exc:  # noqa: BLE001
        # 记录完整异常堆栈,提示执行失败
        logger.exception("执行失败")
        # 将错误信息输出到标准错误流
        print(f"[错误] {exc}", file=sys.stderr)
        # 返回表示失败的退出码 1
        return 1


# 判断当前模块是否作为脚本直接运行
if __name__ == "__main__":
    # 调用主函数启动程序
    main()

5.5 config.py #

config.py

# 导入操作系统模块,用于读取环境变量
import os

# 从 dataclasses 模块导入 dataclass,用于定义数据类
from dataclasses import dataclass


# 使用 dataclass 装饰器,并将实例设为不可变(frozen=True)
@dataclass(frozen=True)
# 定义 RAG 系统运行时配置类
class RAGConfig:
    # 类文档字符串:说明该类用于保存 RAG 系统运行时配置
    """RAG 系统运行时配置。"""

    # 向量数据库存储路径,默认值为当前目录下的 chroma_db
    db_path: str = "./chroma_db"
    # 向量集合名称,默认值为 rag
    collection_name: str = "rag"
    # 嵌入模型名称,默认值为多语言 MiniLM 模型
+   embedding_model: str = (
+       "C:/Users/83687/.cache/modelscope/models/Liudef--paraphrase-multilingual-MiniLM-L12-v2/snapshots/master"
+   )
    # 分块大小,默认值为 500
+   chunk_size: int = 500
    # 分块重叠长度,默认值为 80
+   chunk_overlap: int = 80

    # 将 from_env 声明为类方法
    @classmethod
    # 定义从环境变量加载配置的类方法
    def from_env(cls):
        # 使用环境变量(若不存在则回退到类默认值)创建并返回配置实例
        return cls(
            # 从环境变量 RAG_DB_PATH 读取数据库路径,缺省使用类默认值
            db_path=os.getenv("RAG_DB_PATH", cls.db_path),
            # 从环境变量 RAG_COLLECTION 读取集合名称,缺省使用类默认值
            collection_name=os.getenv("RAG_COLLECTION", cls.collection_name),
            # 从环境变量 RAG_EMBEDDING_MODEL 读取嵌入模型名称,缺省使用类默认值
            embedding_model=os.getenv("RAG_EMBEDDING_MODEL", cls.embedding_model),
            # 从环境变量 RAG_CHUNK_SIZE 读取分块大小,缺省使用类默认值
+           chunk_size=int(os.getenv("RAG_CHUNK_SIZE", str(cls.chunk_size))),
            # 从环境变量 RAG_CHUNK_OVERLAP 读取分块重叠长度,缺省使用类默认值
+           chunk_overlap=int(os.getenv("RAG_CHUNK_OVERLAP", str(cls.chunk_overlap))),
        )

5.6 embeddings.py #

embeddings.py

# 导入日志模块,用于记录程序运行信息
import logging

# 从 sentence_transformers 库导入 SentenceTransformer 模型类
from sentence_transformers import SentenceTransformer

# 获取名为 "rag" 的日志记录器实例
logger = logging.getLogger("rag")


# 定义 EmbeddingService 类,封装本地向量化能力
class EmbeddingService:
    # 类说明文档字符串:本地 SentenceTransformer 向量化服务
    """本地 SentenceTransformer 向量化服务。"""

    # 定义初始化方法,接收模型名称参数
    def __init__(self, model_name):
        # 将传入的模型名称保存为实例属性,供后续加载使用
        self.model_name = model_name
        # 初始化模型实例属性为 None,表示模型尚未加载
+       self._model = None

    # 定义懒加载获取嵌入模型的私有方法
+   def _get_model(self):
        # 判断模型是否尚未加载到实例属性中
+       if self._model is None:
            # 记录正在加载 Embedding 模型的信息日志,并输出模型名称
+           logger.info("加载 Embedding 模型: %s", self.model_name)
            # 使用指定模型名称创建 SentenceTransformer 实例并缓存
+           self._model = SentenceTransformer(self.model_name)
            # 记录 Embedding 模型已就绪的信息日志
+           logger.info("Embedding 模型就绪")
        # 返回已加载的模型实例
+       return self._model

    # 定义文本向量化方法,接收单个字符串或字符串列表
+   def embed(self, texts):
        # 通过懒加载方式获取已就绪的嵌入模型
+       model = self._get_model()
        # 若输入为单个字符串则包装为列表,否则直接使用原列表
+       batch = [texts] if isinstance(texts, str) else texts
        # 调用模型对批次文本进行编码,生成归一化向量
+       vectors = model.encode(
            # 传入文本批次;批次大于 16 时显示进度条,并对向量做归一化
+           batch,
+           show_progress_bar=len(batch) > 16,
+           normalize_embeddings=True,
            # 结束 model.encode 调用
+       )
        # 将各向量转换为 Python 列表并返回
+       return [v.tolist() for v in vectors]

5.7 pipeline.py #

pipeline.py

# 从 pathlib 导入 Path,用于处理与解析文件路径
from pathlib import Path
+import logging

# 从 loader 模块导入 DocumentLoader,用于加载多格式非结构化文档
from loader import DocumentLoader

# 从 rich 库导入增强版 print,支持彩色与富文本输出
from rich import print

# 从 store 模块导入 VectorStore,用于向量存储与检索
from store import VectorStore

# 从 config 模块导入 RAGConfig,用于配置 RAG 系统
from config import RAGConfig

# 从 embeddings 模块导入 EmbeddingService,用于向量化服务
from embeddings import EmbeddingService

# 从 chunker 模块导入 TextChunker,用于文本分块
+from chunker import TextChunker

+logger = logging.getLogger(__name__)


# 定义 RAG 主流程编排器类
class RAGPipeline:
    # 类文档字符串:说明该类负责编排 RAG 主流程
    """RAG 主流程编排器。"""

    # 初始化方法,可选接收配置参数
    def __init__(self, config=None):
        self.config = config or RAGConfig.from_env()
        # 创建文档加载器实例并保存为实例属性
        self.loader = DocumentLoader()
        # 创建向量化服务实例并保存为实例属性
        self.embedder = EmbeddingService(self.config.embedding_model)
        # 创建向量存储实例并保存为实例属性
        self.store = VectorStore(
            db_path=self.config.db_path,  # 数据库路径
            collection_name=self.config.collection_name,  # 集合名称
            embedding_service=self.embedder,  # 向量化服务
        )
        # 创建文本分块器实例并保存为实例属性
+       self.chunker = TextChunker(
+           chunk_size=self.config.chunk_size,  # 分块大小
+           chunk_overlap=self.config.chunk_overlap,  # 分块重叠长度
+       )

    # 定义文件入库方法,接收待入库的文件路径
    def ingest_file(self, file_path):
        # 将输入路径转换为绝对路径对象
        path = Path(file_path).resolve()
        # 校验路径是否指向一个真实存在的文件
        if not path.is_file():
            # 若不是文件则抛出参数错误,提示需指定单个文件路径
            raise ValueError(f"请指定单个文件路径: {path}")
        # 调用文档加载器解析文件内容为文本
        text = self.loader.load(path)
        # 判断清洗后的文本是否为空
        if not text.strip():
            # 文本为空时返回 0,表示未生成有效内容
            return 0
        # 重置向量存储
        self.store.reset()
        # 调用文本分块器将文本切分为文本块列表
+       chunks = self.chunker.split(text, source=str(path.name))
        # 调用向量存储实例将文本块列表插入到向量存储中
+       n = self.store.upsert_chunks(chunks)
        # 记录入库成功日志,包含文件名与文本块数量
+       logger.info("入库成功: %s → %d chunks", path.name, n)
+       return n

5.8 store.py #

store.py

# 导入日志模块,用于记录程序运行信息
import logging

# 从 datetime 模块导入 datetime 与 timezone,用于生成带时区的时间戳
from datetime import datetime, timezone

# 从 pathlib 模块导入 Path,用于处理文件路径
from pathlib import Path
+import chromadb
+from chromadb.config import Settings

# 获取当前模块对应的日志记录器实例
logger = logging.getLogger(__name__)


# 定义向量存储封装类
class VectorStore:
    # 类文档字符串:说明该类用于封装 ChromaDB 向量库
    """ChromaDB 向量库封装。"""

    # 初始化方法,接收数据库路径、集合名称与向量化服务
+   def __init__(self, db_path, collection_name, embedding_service):
        # 将数据库路径保存为实例属性
+       self.db_path = db_path
        # 将集合名称保存为实例属性
+       self.collection_name = collection_name
        # 将向量化服务保存为实例属性
+       self.embedding_service = embedding_service
        # 确保数据库目录存在,若不存在则递归创建
+       Path(db_path).mkdir(parents=True, exist_ok=True)

        # 创建 ChromaDB 持久化客户端实例并保存为私有属性
+       self._client = chromadb.PersistentClient(
            # 指定向量数据库的持久化存储路径
+           path=db_path,
            # 配置客户端设置,关闭匿名遥测上报
+           settings=Settings(anonymized_telemetry=False),
            # 结束 PersistentClient 构造调用
+       )
        # 获取或创建向量集合,并将结果赋值给实例的集合属性
+       self._collection = self._client.get_or_create_collection(
            # 指定集合名称为传入的集合名
+           name=collection_name,
            # 设置集合元数据字典
+           metadata={
                # 指定 HNSW 索引使用余弦相似度空间
+               "hnsw:space": "cosine",
                # 设置集合描述为RAG 知识库
+               "description": "RAG knowledge base",
                # 记录集合创建时间为当前 UTC 时间的 ISO 格式字符串
+               "created_at": datetime.now(timezone.utc).isoformat(),
                # 结束 metadata 字典
+           },
            # 结束 get_or_create_collection 调用
+       )
        # 记录向量库就绪的信息日志,包含路径、集合名与当前文档数量
+       logger.info(
            # 日志消息模板:路径、集合名称与文档计数
+           "向量库就绪: path=%s, collection=%s, count=%d",
            # 传入数据库路径参数
+           db_path,
            # 传入集合名称参数
+           collection_name,
            # 传入当前集合中的文档数量
+           self._collection.count(),
            # 结束 logger.info 调用
+       )

    # 将 count 声明为只读属性
+   @property
    # 定义获取集合文档数量的属性方法
+   def count(self):
        # 返回当前向量集合中的文档数量
+       return self._collection.count()

    # 定义重置集合的实例方法
    def reset(self):
        # 方法文档字符串:说明该方法用于清空并重建集合
        """清空并重建集合。"""
        # 尝试删除已有集合,忽略删除失败的情况
        try:
            # 调用客户端删除当前集合名称对应的集合
            self._client.delete_collection(self.collection_name)
        # 捕获所有异常(忽略过宽捕获的静态检查警告)
        except Exception:  # noqa: BLE001
            # 删除失败时忽略异常,继续后续重建流程
            pass
        # 获取或创建集合,并将结果赋值给实例的集合属性
        self._collection = self._client.get_or_create_collection(
            # 指定集合名称为当前配置的集合名
            name=self.collection_name,
            # 设置集合元数据字典
            metadata={
                # 指定 HNSW 索引使用余弦相似度空间
                "hnsw:space": "cosine",
                # 设置集合描述为 RAG 知识库
                "description": "RAG 知识库",
                # 记录集合创建时间为当前 UTC 时间的 ISO 格式字符串
                "created_at": datetime.now(timezone.utc).isoformat(),
                # 结束 metadata 字典
            },
            # 结束 get_or_create_collection 调用
        )
        # 记录已重建空集合的信息日志,并输出集合名称
        logger.info("已重建空集合: %s", self.collection_name)

    # 定义批量向量化并写入文本块的实例方法,默认批大小为 64
+   def upsert_chunks(self, chunks, batch_size=64):
        # 方法文档字符串:说明批量向量化并写入;已存在 ID 则 upsert 覆盖
+       """批量向量化并写入;已存在 ID 则 upsert 覆盖。"""
        # 判断文本块列表是否为空
+       if not chunks:
            # 若为空则直接返回 0,表示未写入任何内容
+           return 0

        # 初始化已写入文本块数量计数器
+       written = 0
        # 按批大小步进遍历文本块列表的起始索引
+       for start in range(0, len(chunks), batch_size):
            # 截取当前批次的文本块子列表
+           batch = chunks[start : start + batch_size]
            # 调用嵌入服务对当前批次各文本块内容进行向量化
+           embeddings = self.embedding_service.embed([c.content for c in batch])
            # 调用集合的 upsert 方法批量写入或覆盖当前批次数据
+           self._collection.upsert(
                # 传入当前批次各文本块的唯一标识符列表
+               ids=[c.id for c in batch],
                # 传入当前批次各文本块的正文内容列表
+               documents=[c.content for c in batch],
                # 传入当前批次对应的向量嵌入列表
+               embeddings=embeddings,
                # 传入当前批次各文本块的元数据列表
+               metadatas=[
                    # 为每个文本块构造元数据字典
+                   {
                        # 记录文本块来源(如文件名)
+                       "source": c.source,
                        # 记录文本块在原文中的序号索引
+                       "chunk_index": c.chunk_index,
                        # 记录文本块正文的字符长度
+                       "char_len": len(c.content),
                        # 结束元数据字典
+                   }
                    # 遍历当前批次中的每个文本块以生成元数据
+                   for c in batch
                    # 结束元数据列表推导式
+               ],
                # 结束 upsert 调用
+           )
            # 累加当前批次写入的文本块数量
+           written += len(batch)
            # 记录当前批次写入进度的信息日志
+           logger.info(
                # 日志格式字符串:已写入数量、总数量与集合当前总量
+               "写入进度: %d/%d (collection total=%d)",
                # 传入当前已处理到的文本块数量
+               start + len(batch),
                # 传入待写入文本块的总数量
+               len(chunks),
                # 传入集合中当前文档总量
+               self.count,
                # 结束 logger.info 调用
+           )
        # 返回本次成功写入的文本块总数量
+       return written

6. 检索问答 #

本节在入库闭环之上补齐「提问 → 相似检索 → 返回命中」:配置增加 top_kscore_thresholdRetrievalHit 描述单条命中;embed_query 将问题编码为查询向量;VectorStore.similarity_search 在 Chroma 中查询,把 cosine 距离转为相似度(1 - distance)并按阈值过滤、按分数排序。CLI 新增 query -q,经 askretrieve 打印命中列表。

整体流程:校验问题非空 → 问题向量化 → 向量库 top-k 检索 → 阈值过滤与排序 → 输出 RetrievalHit[]

sequenceDiagram autonumber actor User as 用户 participant CLI as cmd_query() participant Pipe as RAGPipeline.ask / retrieve participant Store as VectorStore participant Emb as EmbeddingService participant Chroma as ChromaDB User->>CLI: query -q "年假有多少天?" CLI->>Pipe: ask(question) alt 问题为空 Pipe-->>CLI: raise ValueError else 问题合法 Pipe->>Pipe: retrieve(question, top_k) Pipe->>Store: similarity_search(query, top_k, score_threshold) alt 知识库为空 Store-->>Pipe: raise ValueError(请先 ingest) else 有文档 Store->>Emb: embed_query(query) Emb-->>Store: query_vec Store->>Chroma: query(embeddings, n_results=min(top_k, count)) Chroma-->>Store: documents / metadatas / distances / ids Store->>Store: score = 1 - distance,过滤 score < threshold Store->>Store: 构造 RetrievalHit[] 并按 score 降序 opt 全部低于阈值 Store->>Store: warning(提示可能混入旧 Embedding 数据) end Store-->>Pipe: hits Pipe->>Pipe: print(hits) end end

6.1 .env #

.env

# 向量数据库存储路径,默认值为当前目录下的 chroma_db
RAG_DB_PATH="./chroma_db"
# 向量集合名称,默认值为 rag
RAG_COLLECTION="rag"
# 嵌入模型名称,默认值为多语言 MiniLM 模型
RAG_EMBEDDING_MODEL="C:/Users/83687/.cache/modelscope/models/Liudef--paraphrase-multilingual-MiniLM-L12-v2/snapshots/master"
# 检索返回的相似文档数量,默认值为 5
+RAG_TOP_K=5
# 相似度阈值,默认值为 0.2
+RAG_SCORE_THRESHOLD=0.2

6.2 cli.py #

cli.py

# 导入日志模块,用于记录程序运行信息
import logging

# 导入系统模块,用于访问标准错误输出等系统功能
import sys

# 导入命令行参数解析模块
import argparse

# 导入 JSON 模块,用于序列化输出数据
import json

# 从 rich 库导入增强版 print,支持彩色与富文本输出
from rich import print

# 导入 RAG 主流程编排器
from pipeline import RAGPipeline

# 导入 RAG 配置
+from config import RAGConfig


# 配置全局日志的基本格式与级别
logging.basicConfig(
    # 设置日志级别为 INFO,输出信息及以上级别日志
    level=logging.INFO,
    # 定义日志输出格式:时间、级别、记录器名称与消息内容
    format="%(asctime)s [%(levelname)s] %(name)s - %(message)s",
    # 定义时间戳的显示格式为年-月-日 时:分:秒
    datefmt="%Y-%m-%d %H:%M:%S",
    # 结束 logging.basicConfig 的参数配置
)
# 获取当前模块对应的日志记录器实例
logger = logging.getLogger(__name__)


# 定义 ingest 子命令的处理函数,接收解析后的参数对象
def cmd_ingest(args):
    # 创建 RAG 主流程编排器实例
    pipeline = RAGPipeline()
    # 调用 ingest_file 方法入库文件
    chunks = pipeline.ingest_file(args.path)
    # 将文件路径封装为字典并以格式化 JSON 打印输出
    print(
        print(
            json.dumps(
                {"file": args.path, "chunks": chunks}, ensure_ascii=False, indent=2
            )
        )
    )


# 定义 query 子命令的处理函数,接收解析后的参数对象
+def cmd_query(args):
    # 从环境变量加载配置并创建 RAG 主流程编排器实例
+   pipeline = RAGPipeline(RAGConfig.from_env())
    # 调用 ask 方法对用户问题进行检索增强问答
+   pipeline.ask(args.question)


# 构建并返回命令行参数解析器
def build_parser():
    # 创建顶层 ArgumentParser,设置程序描述与帮助格式化器
    parser = argparse.ArgumentParser(
        # 设置程序用途说明:RAG 文件入库与问答
        description="RAG:文件入库 + 问答",
        # 使用原始描述帮助格式化器,保留描述中的换行与缩进
        formatter_class=argparse.RawDescriptionHelpFormatter,
        # 结束 ArgumentParser 的参数配置
    )
    # 添加子命令解析器,结果写入 args.command,且必须指定子命令
    sub = parser.add_subparsers(dest="command", required=True)

    # 添加名为 ingest 的子命令,用于文件入库
    p_ingest = sub.add_parser("ingest", help="入库文件")
    # 为 ingest 子命令添加必填的 --path 参数,表示待入库文件路径
    p_ingest.add_argument("--path", required=True, help="文件路径")
    # 将 ingest 子命令的默认处理函数绑定为 cmd_ingest
    p_ingest.set_defaults(func=cmd_ingest)

    # 添加名为 query 的子命令,用于问答
+   p_query = sub.add_parser("query", help="问答")
    # 为 query 子命令添加必填的 --question/-q 参数,表示用户问题
+   p_query.add_argument("--question", "-q", required=True, help="用户问题")
    # 将 query 子命令的默认处理函数绑定为 cmd_query
+   p_query.set_defaults(func=cmd_query)
    # 返回配置完成的参数解析器
    return parser


# 定义程序主入口函数
def main():
    # 调用 build_parser 构建命令行参数解析器
    parser = build_parser()
    # 解析命令行参数,得到命名空间对象 args
    args = parser.parse_args()
    # 打印解析得到的参数对象,便于调试查看
    print(args)
    # 尝试执行子命令对应的处理函数
    try:
        # 调用当前子命令绑定的处理函数并传入参数
        args.func(args)
        # 执行成功时返回退出码 0
        return 0
    # 捕获用户按下 Ctrl+C 产生的键盘中断异常
    except KeyboardInterrupt:
        # 打印已中断提示信息
        print("\n已中断")
        # 返回标准的键盘中断退出码 130
        return 130
    # 捕获其余所有异常(忽略过宽捕获的静态检查警告)
    except Exception as exc:  # noqa: BLE001
        # 记录完整异常堆栈,提示执行失败
        logger.exception("执行失败")
        # 将错误信息输出到标准错误流
        print(f"[错误] {exc}", file=sys.stderr)
        # 返回表示失败的退出码 1
        return 1


# 判断当前模块是否作为脚本直接运行
if __name__ == "__main__":
    # 调用主函数启动程序
    main()

6.3 config.py #

config.py

# 导入操作系统模块,用于读取环境变量
import os

# 从 dataclasses 模块导入 dataclass,用于定义数据类
from dataclasses import dataclass


# 使用 dataclass 装饰器,并将实例设为不可变(frozen=True)
@dataclass(frozen=True)
# 定义 RAG 系统运行时配置类
class RAGConfig:
    # 类文档字符串:说明该类用于保存 RAG 系统运行时配置
    """RAG 系统运行时配置。"""

    # 向量数据库存储路径,默认值为当前目录下的 chroma_db
    db_path: str = "./chroma_db"
    # 向量集合名称,默认值为 rag
    collection_name: str = "rag"
    # 嵌入模型名称,默认值为多语言 MiniLM 模型
    embedding_model: str = (
        "C:/Users/83687/.cache/modelscope/models/Liudef--paraphrase-multilingual-MiniLM-L12-v2/snapshots/master"
    )
    # 分块大小,默认值为 500
    chunk_size: int = 500
    # 分块重叠长度,默认值为 80
    chunk_overlap: int = 80
    # 检索返回的相似文档数量,默认值为 5
+   top_k: int = 5
    # 相似度阈值,默认值为 0.2
    # cosine 距离越小越相似;此处用 1 - distance 作为相似度
+   score_threshold: float = 0.2

    # 将 from_env 声明为类方法
    @classmethod
    # 定义从环境变量加载配置的类方法
    def from_env(cls):
        # 使用环境变量(若不存在则回退到类默认值)创建并返回配置实例
        return cls(
            # 从环境变量 RAG_DB_PATH 读取数据库路径,缺省使用类默认值
            db_path=os.getenv("RAG_DB_PATH", cls.db_path),
            # 从环境变量 RAG_COLLECTION 读取集合名称,缺省使用类默认值
            collection_name=os.getenv("RAG_COLLECTION", cls.collection_name),
            # 从环境变量 RAG_EMBEDDING_MODEL 读取嵌入模型名称,缺省使用类默认值
            embedding_model=os.getenv("RAG_EMBEDDING_MODEL", cls.embedding_model),
            # 从环境变量 RAG_CHUNK_SIZE 读取分块大小,缺省使用类默认值
            chunk_size=int(os.getenv("RAG_CHUNK_SIZE", str(cls.chunk_size))),
            # 从环境变量 RAG_CHUNK_OVERLAP 读取分块重叠长度,缺省使用类默认值
            chunk_overlap=int(os.getenv("RAG_CHUNK_OVERLAP", str(cls.chunk_overlap))),
            # 从环境变量 RAG_TOP_K 读取检索返回的相似文档数量,缺省使用类默认值
+           top_k=int(os.getenv("RAG_TOP_K", str(cls.top_k))),
            # 从环境变量 RAG_SCORE_THRESHOLD 读取相似度阈值,缺省使用类默认值
+           score_threshold=float(
                # 读取环境变量 RAG_SCORE_THRESHOLD,若不存在则回退到类默认值并转为字符串
+               os.getenv("RAG_SCORE_THRESHOLD", str(cls.score_threshold))
                # 结束 float 转换与 score_threshold 参数赋值
+           ),
        )

6.4 embeddings.py #

embeddings.py

# 导入日志模块,用于记录程序运行信息
import logging

# 从 sentence_transformers 库导入 SentenceTransformer 模型类
from sentence_transformers import SentenceTransformer

# 获取名为 "rag" 的日志记录器实例
logger = logging.getLogger("rag")


# 定义 EmbeddingService 类,封装本地向量化能力
class EmbeddingService:
    # 类说明文档字符串:本地 SentenceTransformer 向量化服务
    """本地 SentenceTransformer 向量化服务。"""

    # 定义初始化方法,接收模型名称参数
    def __init__(self, model_name):
        # 将传入的模型名称保存为实例属性,供后续加载使用
        self.model_name = model_name
        # 初始化模型实例属性为 None,表示模型尚未加载
        self._model = None

    # 定义懒加载获取嵌入模型的私有方法
    def _get_model(self):
        # 判断模型是否尚未加载到实例属性中
        if self._model is None:
            # 记录正在加载 Embedding 模型的信息日志,并输出模型名称
            logger.info("加载 Embedding 模型: %s", self.model_name)
            # 使用指定模型名称创建 SentenceTransformer 实例并缓存
            self._model = SentenceTransformer(self.model_name)
            # 记录 Embedding 模型已就绪的信息日志
            logger.info("Embedding 模型就绪")
        # 返回已加载的模型实例
        return self._model

    # 定义文本向量化方法,接收单个字符串或字符串列表
    def embed(self, texts):
        # 通过懒加载方式获取已就绪的嵌入模型
        model = self._get_model()
        # 若输入为单个字符串则包装为列表,否则直接使用原列表
        batch = [texts] if isinstance(texts, str) else texts
        # 调用模型对批次文本进行编码,生成归一化向量
        vectors = model.encode(
            # 传入文本批次;批次大于 16 时显示进度条,并对向量做归一化
            batch,
            show_progress_bar=len(batch) > 16,
            normalize_embeddings=True,
            # 结束 model.encode 调用
        )
        # 将各向量转换为 Python 列表并返回
        return [v.tolist() for v in vectors]

    # 定义单个查询文本向量化方法,接收单个字符串
+   def embed_query(self, query):
        # 调用文本向量化方法对单个查询文本进行编码,生成归一化向量
+       vectors = self.embed(query)
        # 返回第一个向量(单个查询文本只有一个向量)
+       return vectors[0]

6.5 models.py #

models.py

# 从 dataclasses 模块导入 asdict、dataclass 与 field,用于定义数据类与字段默认值
from dataclasses import asdict, dataclass, field


# 使用 dataclass 装饰器,将类自动转换为数据类
@dataclass
# 定义入库后的文本块数据模型类
class DocumentChunk:
    # 类文档字符串:说明该类表示入库后的文本块
    """入库后的文本块。"""

    # 文本块的唯一标识符
    id: str
    # 文本块的正文内容
    content: str
    # 文本块来源(如文件名)
    source: str
    # 文本块在原文中的序号索引
    chunk_index: int
    # 附加元数据字典,默认创建空字典
    metadata: dict = field(default_factory=dict)


# 使用 dataclass 装饰器,将类自动转换为数据类
+@dataclass
# 定义单条检索命中数据模型类
+class RetrievalHit:
    # 类文档字符串:说明该类表示单条检索命中
+   """单条检索命中。"""

    # 检索命中的正文内容
+   content: str
    # 检索命中的来源(如文件名)
+   source: str
    # 检索命中对应文本块的唯一标识符
+   chunk_id: str
    # 检索相似度得分
+   score: float
    # 附加元数据字典,默认创建空字典
+   metadata: dict = field(default_factory=dict)

6.6 pipeline.py #

pipeline.py

# 从 pathlib 导入 Path,用于处理与解析文件路径
from pathlib import Path
import logging

# 从 loader 模块导入 DocumentLoader,用于加载多格式非结构化文档
from loader import DocumentLoader

# 从 rich 库导入增强版 print,支持彩色与富文本输出
from rich import print

# 从 store 模块导入 VectorStore,用于向量存储与检索
from store import VectorStore

# 从 config 模块导入 RAGConfig,用于配置 RAG 系统
from config import RAGConfig

# 从 embeddings 模块导入 EmbeddingService,用于向量化服务
from embeddings import EmbeddingService

# 从 chunker 模块导入 TextChunker,用于文本分块
from chunker import TextChunker

logger = logging.getLogger(__name__)


# 定义 RAG 主流程编排器类
class RAGPipeline:
    # 类文档字符串:说明该类负责编排 RAG 主流程
    """RAG 主流程编排器。"""

    # 初始化方法,可选接收配置参数
    def __init__(self, config=None):
        self.config = config or RAGConfig.from_env()
        # 创建文档加载器实例并保存为实例属性
        self.loader = DocumentLoader()
        # 创建向量化服务实例并保存为实例属性
        self.embedder = EmbeddingService(self.config.embedding_model)
        # 创建向量存储实例并保存为实例属性
        self.store = VectorStore(
            db_path=self.config.db_path,  # 数据库路径
            collection_name=self.config.collection_name,  # 集合名称
            embedding_service=self.embedder,  # 向量化服务
        )
        # 创建文本分块器实例并保存为实例属性
        self.chunker = TextChunker(
            chunk_size=self.config.chunk_size,  # 分块大小
            chunk_overlap=self.config.chunk_overlap,  # 分块重叠长度
        )

    # 定义文件入库方法,接收待入库的文件路径
    def ingest_file(self, file_path):
        # 将输入路径转换为绝对路径对象
        path = Path(file_path).resolve()
        # 校验路径是否指向一个真实存在的文件
        if not path.is_file():
            # 若不是文件则抛出参数错误,提示需指定单个文件路径
            raise ValueError(f"请指定单个文件路径: {path}")
        # 调用文档加载器解析文件内容为文本
        text = self.loader.load(path)
        # 判断清洗后的文本是否为空
        if not text.strip():
            # 文本为空时返回 0,表示未生成有效内容
            return 0
        # 重置向量存储
        self.store.reset()
        # 调用文本分块器将文本切分为文本块列表
        chunks = self.chunker.split(text, source=str(path.name))
        # 调用向量存储实例将文本块列表插入到向量存储中
        n = self.store.upsert_chunks(chunks)
        # 记录入库成功日志,包含文件名与文本块数量
        logger.info("入库成功: %s → %d chunks", path.name, n)
        return n

    # 定义检索方法,接收问题文本与可选的返回数量
+   def retrieve(self, question, top_k=None):
        # 调用向量存储的相似度检索,并返回命中结果
+       return self.store.similarity_search(
            # 传入查询问题文本
+           query=question,
            # 传入返回数量,未指定时回退为配置中的默认 top_k
+           top_k=top_k or self.config.top_k,
            # 传入配置中的相似度分数阈值
+           score_threshold=self.config.score_threshold,
            # 结束 similarity_search 调用
+       )

    # 定义问答方法,接收问题文本与可选的返回数量
+   def ask(self, question, top_k=None):
        # 去除问题文本首尾空白字符
+       question = question.strip()
        # 判断清洗后的问题是否为空
+       if not question:
            # 问题为空时抛出参数错误
+           raise ValueError("问题不能为空")

        # 记录开始 RAG 问答的信息日志,包含问题内容
+       logger.info("开始 RAG 问答: %s", question)
        # 调用检索方法获取与问题相关的命中结果
+       hits = self.retrieve(question, top_k=top_k)
        # 将检索命中结果打印到标准输出
+       print(hits)

6.7 store.py #

store.py

# 导入日志模块,用于记录程序运行信息
import logging

# 从 datetime 模块导入 datetime 与 timezone,用于生成带时区的时间戳
from datetime import datetime, timezone

# 从 pathlib 模块导入 Path,用于处理文件路径
from pathlib import Path
import chromadb

# 从 chromadb.config 模块导入 Settings,用于配置 ChromaDB 客户端设置
from chromadb.config import Settings

# 从 models 模块导入 RetrievalHit,用于表示单条检索命中
+from models import RetrievalHit

# 获取当前模块对应的日志记录器实例
logger = logging.getLogger(__name__)


# 定义向量存储封装类
class VectorStore:
    # 类文档字符串:说明该类用于封装 ChromaDB 向量库
    """ChromaDB 向量库封装。"""

    # 初始化方法,接收数据库路径、集合名称与向量化服务
    def __init__(self, db_path, collection_name, embedding_service):
        # 将数据库路径保存为实例属性
        self.db_path = db_path
        # 将集合名称保存为实例属性
        self.collection_name = collection_name
        # 将向量化服务保存为实例属性
        self.embedding_service = embedding_service
        # 确保数据库目录存在,若不存在则递归创建
        Path(db_path).mkdir(parents=True, exist_ok=True)

        # 创建 ChromaDB 持久化客户端实例并保存为私有属性
        self._client = chromadb.PersistentClient(
            # 指定向量数据库的持久化存储路径
            path=db_path,
            # 配置客户端设置,关闭匿名遥测上报
            settings=Settings(anonymized_telemetry=False),
            # 结束 PersistentClient 构造调用
        )
        # 获取或创建向量集合,并将结果赋值给实例的集合属性
        self._collection = self._client.get_or_create_collection(
            # 指定集合名称为传入的集合名
            name=collection_name,
            # 设置集合元数据字典
            metadata={
                # 指定 HNSW 索引使用余弦相似度空间
                "hnsw:space": "cosine",
                # 设置集合描述为RAG 知识库
                "description": "RAG knowledge base",
                # 记录集合创建时间为当前 UTC 时间的 ISO 格式字符串
                "created_at": datetime.now(timezone.utc).isoformat(),
                # 结束 metadata 字典
            },
            # 结束 get_or_create_collection 调用
        )
        # 记录向量库就绪的信息日志,包含路径、集合名与当前文档数量
        logger.info(
            # 日志消息模板:路径、集合名称与文档计数
            "向量库就绪: path=%s, collection=%s, count=%d",
            # 传入数据库路径参数
            db_path,
            # 传入集合名称参数
            collection_name,
            # 传入当前集合中的文档数量
            self._collection.count(),
            # 结束 logger.info 调用
        )

    # 将 count 声明为只读属性
    @property
    # 定义获取集合文档数量的属性方法
    def count(self):
        # 返回当前向量集合中的文档数量
        return self._collection.count()

    # 定义重置集合的实例方法
    def reset(self):
        # 方法文档字符串:说明该方法用于清空并重建集合
        """清空并重建集合。"""
        # 尝试删除已有集合,忽略删除失败的情况
        try:
            # 调用客户端删除当前集合名称对应的集合
            self._client.delete_collection(self.collection_name)
        # 捕获所有异常(忽略过宽捕获的静态检查警告)
        except Exception:  # noqa: BLE001
            # 删除失败时忽略异常,继续后续重建流程
            pass
        # 获取或创建集合,并将结果赋值给实例的集合属性
        self._collection = self._client.get_or_create_collection(
            # 指定集合名称为当前配置的集合名
            name=self.collection_name,
            # 设置集合元数据字典
            metadata={
                # 指定 HNSW 索引使用余弦相似度空间
                "hnsw:space": "cosine",
                # 设置集合描述为 RAG 知识库
                "description": "RAG 知识库",
                # 记录集合创建时间为当前 UTC 时间的 ISO 格式字符串
                "created_at": datetime.now(timezone.utc).isoformat(),
                # 结束 metadata 字典
            },
            # 结束 get_or_create_collection 调用
        )
        # 记录已重建空集合的信息日志,并输出集合名称
        logger.info("已重建空集合: %s", self.collection_name)

    # 定义批量向量化并写入文本块的实例方法,默认批大小为 64
    def upsert_chunks(self, chunks, batch_size=64):
        # 方法文档字符串:说明批量向量化并写入;已存在 ID 则 upsert 覆盖
        """批量向量化并写入;已存在 ID 则 upsert 覆盖。"""
        # 判断文本块列表是否为空
        if not chunks:
            # 若为空则直接返回 0,表示未写入任何内容
            return 0

        # 初始化已写入文本块数量计数器
        written = 0
        # 按批大小步进遍历文本块列表的起始索引
        for start in range(0, len(chunks), batch_size):
            # 截取当前批次的文本块子列表
            batch = chunks[start : start + batch_size]
            # 调用嵌入服务对当前批次各文本块内容进行向量化
            embeddings = self.embedding_service.embed([c.content for c in batch])
            # 调用集合的 upsert 方法批量写入或覆盖当前批次数据
            self._collection.upsert(
                # 传入当前批次各文本块的唯一标识符列表
                ids=[c.id for c in batch],
                # 传入当前批次各文本块的正文内容列表
                documents=[c.content for c in batch],
                # 传入当前批次对应的向量嵌入列表
                embeddings=embeddings,
                # 传入当前批次各文本块的元数据列表
                metadatas=[
                    # 为每个文本块构造元数据字典
                    {
                        # 记录文本块来源(如文件名)
                        "source": c.source,
                        # 记录文本块在原文中的序号索引
                        "chunk_index": c.chunk_index,
                        # 记录文本块正文的字符长度
                        "char_len": len(c.content),
                        # 结束元数据字典
                    }
                    # 遍历当前批次中的每个文本块以生成元数据
                    for c in batch
                    # 结束元数据列表推导式
                ],
                # 结束 upsert 调用
            )
            # 累加当前批次写入的文本块数量
            written += len(batch)
            # 记录当前批次写入进度的信息日志
            logger.info(
                # 日志格式字符串:已写入数量、总数量与集合当前总量
                "写入进度: %d/%d (collection total=%d)",
                # 传入当前已处理到的文本块数量
                start + len(batch),
                # 传入待写入文本块的总数量
                len(chunks),
                # 传入集合中当前文档总量
                self.count,
                # 结束 logger.info 调用
            )
        # 返回本次成功写入的文本块总数量
        return written

    # 定义相似度检索方法,接收查询文本、返回数量、分数阈值与可选过滤条件
+   def similarity_search(self, query, top_k=5, score_threshold=0.0, where=None):
        # 判断向量库中是否没有任何文档
+       if self.count == 0:
            # 知识库为空时抛出错误,提示需先执行入库
+           raise ValueError("知识库为空,请先执行入库(ingest)")

        # 调用嵌入服务将查询文本转换为查询向量
+       query_vec = self.embedding_service.embed_query(query)
        # 取请求数量与库内文档总数的较小值,避免超额查询
+       n_results = min(top_k, self.count)
        # 构造传给集合查询接口的关键字参数字典
+       kwargs = {
            # 传入查询向量列表(单条查询)
+           "query_embeddings": [query_vec],
            # 传入实际要返回的结果数量
+           "n_results": n_results,
            # 指定返回文档正文、元数据与距离字段
+           "include": ["documents", "metadatas", "distances"],
            # 结束 kwargs 字典
+       }
        # 判断是否传入了元数据过滤条件
+       if where:
            # 将过滤条件写入查询参数
+           kwargs["where"] = where

        # 调用集合的 query 方法执行向量相似度检索
+       raw = self._collection.query(**kwargs)
        # 初始化检索命中结果列表
+       hits = []

        # 取出第一组查询对应的文档正文列表,缺失时回退为空列表
+       documents = (raw.get("documents") or [[]])[0]
        # 取出第一组查询对应的元数据列表,缺失时回退为空列表
+       metadatas = (raw.get("metadatas") or [[]])[0]
        # 取出第一组查询对应的距离列表,缺失时回退为空列表
+       distances = (raw.get("distances") or [[]])[0]
        # 取出第一组查询对应的文本块 ID 列表,缺失时回退为空列表
+       ids = (raw.get("ids") or [[]])[0]

        # 初始化带分数的候选结果列表
+       scored = []
        # 并行遍历文档、元数据、距离与 ID,组装候选结果
+       for doc, meta, dist, chunk_id in zip(documents, metadatas, distances, ids):
            # 余弦空间下距离属于 [0, 2],相似度近似为 1 减去距离
            # cosine space: distance ∈ [0, 2],相似度 ≈ 1 - distance
            # 将距离转换为相似度分数
+           score = 1.0 - float(dist)
            # 将分数与对应字段打包后加入候选列表
+           scored.append((score, doc, meta, chunk_id))

        # 遍历所有带分数的候选结果,按阈值过滤并构造命中对象
+       for score, doc, meta, chunk_id in scored:
            # 判断当前分数是否低于设定阈值
+           if score < score_threshold:
                # 低于阈值则跳过该条结果
+               continue
            # 元数据为空时回退为空字典,避免后续取值报错
+           meta = meta or {}
            # 将通过阈值的结果追加到命中列表
+           hits.append(
                # 构造单条检索命中数据对象
+               RetrievalHit(
                    # 写入命中正文,空值时回退为空字符串
+                   content=doc or "",
                    # 写入来源字段,缺失时标记为 unknown
+                   source=str(meta.get("source", "unknown")),
                    # 写入文本块唯一标识符
+                   chunk_id=str(chunk_id),
                    # 写入四舍五入到四位小数的相似度分数
+                   score=round(score, 4),
                    # 写入完整元数据副本
+                   metadata=dict(meta),
                    # 结束 RetrievalHit 构造
+               )
                # 结束 hits.append 调用
+           )

        # 按相似度分数从高到低对命中结果排序
+       hits.sort(key=lambda h: h.score, reverse=True)
        # 判断是否没有任何命中,但存在未过阈值的候选结果
+       if not hits and scored:
            # 取出分数最高的前 5 个候选分数,格式化为逗号分隔字符串
+           top = ", ".join(f"{s:.3f}" for s, *_ in sorted(scored, reverse=True)[:5])
            # 记录警告日志,提示结果均低于阈值及可能的原因
+           logger.warning(
                # 警告消息模板:阈值与最高分提示
+               "检索结果均低于阈值 %.2f,最高分: %s(可能混入了其他 Embedding 模型的旧数据,请重新 ingest)",
                # 传入相似度阈值参数
+               score_threshold,
                # 传入最高分摘要字符串
+               top,
                # 结束 logger.warning 调用
+           )
        # 记录检索完成的信息日志,包含查询摘要与命中数量
+       logger.info("检索完成: query=%r, hits=%d/%d", query[:40], len(hits), n_results)
        # 返回过滤并排序后的检索命中列表
+       return hits

7. LLM 生成与完整 RAG 问答 #

本节在第 5 节「只检索、不生成」之上接入大模型,形成完整 RAG:LLMClient 通过 OpenAI 兼容接口(默认 DeepSeek)懒加载客户端,用固定 SYSTEM_PROMPT 约束「仅依据参考资料作答」;配置增加 base_url / API Key / 模型名 / temperature / max_context_chars,并由 dotenv 加载 .envbuild_prompt 把命中块拼成带来源与相似度的上下文(受字符上限截断);ask 在检索后:无命中直接返回固定提示,有命中则 llm.generate 产出最终答案。

整体流程:检索 hits → 组装 user prompt →(有命中)调用聊天补全 → 打印答案。

sequenceDiagram autonumber actor User as 用户 participant CLI as cmd_query() participant Pipe as RAGPipeline.ask participant Store as retrieve / similarity_search participant LLM as LLMClient participant API as OpenAI 兼容 API(DeepSeek) User->>CLI: query -q "年假有多少天?" CLI->>Pipe: ask(question) Pipe->>Store: retrieve(question) Store-->>Pipe: hits[] Pipe->>Pipe: build_prompt(question, hits) Note over Pipe: 按 max_context_chars 拼接<br/>[i] 来源 | 相似度 + 正文 alt 无命中 Pipe->>Pipe: answer = 固定「无法确定」提示 else 有命中 Pipe->>LLM: generate(prompt) LLM->>LLM: _get_client()(校验 API Key,懒加载) LLM->>API: chat.completions.create<br/>system=SYSTEM_PROMPT + user=prompt API-->>LLM: message.content LLM-->>Pipe: answer end Pipe-->>User: print(answer)

7.1 llm.py #

llm.py

import logging

from openai import OpenAI

logger = logging.getLogger("rag")

SYSTEM_PROMPT = """你是知识库助手。请严格根据「检索到的参考资料」回答用户问题。

回答要求:
1. 仅依据参考资料作答;资料不足时明确说明「根据现有知识库无法确定」,不要编造。
2. 关键结论尽量引用资料来源(如文件名),便于业务同事核对。
3. 表述专业、简洁,适合企业内部沟通。
4. 若资料之间存在冲突,请指出差异并说明依据。
"""


class LLMClient:
    """大模型调用封装。"""

    def __init__(self, config):
        self.config = config
        self._client = None

    def _get_client(self):
        if self._client is None:
            if not self.config.openai_api_key:
                raise ValueError(
                    "未配置 DEEPSEEK_API_KEY / OPENAI_API_KEY。"
                    "请在 .env 中设置后再进行问答。"
                )
            self._client = OpenAI(
                api_key=self.config.openai_api_key,
                base_url=self.config.openai_base_url,
            )
            logger.info(
                "DeepSeek 客户端已初始化: %s | model=%s",
                self.config.openai_base_url,
                self.config.openai_model,
            )
        return self._client

    def generate(self, user_prompt):
        client = self._get_client()
        response = client.chat.completions.create(
            model=self.config.openai_model,
            temperature=self.config.temperature,
            messages=[
                {"role": "system", "content": SYSTEM_PROMPT},
                {"role": "user", "content": user_prompt},
            ],
        )
        content = response.choices[0].message.content
        return (content or "").strip()

7.2 .env #

.env

# 向量数据库存储路径,默认值为当前目录下的 chroma_db
RAG_DB_PATH="./chroma_db"
# 向量集合名称,默认值为 rag
RAG_COLLECTION="rag"
# 嵌入模型名称,默认值为多语言 MiniLM 模型
RAG_EMBEDDING_MODEL="C:/Users/83687/.cache/modelscope/models/Liudef--paraphrase-multilingual-MiniLM-L12-v2/snapshots/master"
# 检索返回的相似文档数量,默认值为 5
RAG_TOP_K=5
# 相似度阈值,默认值为 0.2
RAG_SCORE_THRESHOLD=0.2
# OpenAI 兼容接口的基础 URL,默认指向 DeepSeek API
+OPENAI_API_BASE=https://api.deepseek.com
# OpenAI 兼容接口的 API 密钥,默认值为空字符串
+DEEPSEEK_API_KEY=sk-0b24393e7dbb4b7d854f59c6e8be527e
# 大语言模型名称,默认值为 deepseek-v4-flash
+OPENAI_MODEL_NAME="deepseek-v4-flash"
# 生成温度参数,默认值为 0.2,数值越低输出越稳定
+RAG_TEMPERATURE=0.2
# 最大上下文字符数,默认值为 6000
+RAG_MAX_CONTEXT_CHARS=6000

7.3 config.py #

config.py

# 导入操作系统模块,用于读取环境变量
import os

# 从 dataclasses 模块导入 dataclass,用于定义数据类
from dataclasses import dataclass

+import dotenv

+dotenv.load_dotenv(override=True)


# 使用 dataclass 装饰器,并将实例设为不可变(frozen=True)
@dataclass(frozen=True)
# 定义 RAG 系统运行时配置类
class RAGConfig:
    # 类文档字符串:说明该类用于保存 RAG 系统运行时配置
    """RAG 系统运行时配置。"""

    # 向量数据库存储路径,默认值为当前目录下的 chroma_db
    db_path: str = "./chroma_db"
    # 向量集合名称,默认值为 rag
    collection_name: str = "rag"
    # 嵌入模型名称,默认值为多语言 MiniLM 模型
    embedding_model: str = (
        "C:/Users/83687/.cache/modelscope/models/Liudef--paraphrase-multilingual-MiniLM-L12-v2/snapshots/master"
    )
    # 分块大小,默认值为 500
    chunk_size: int = 500
    # 分块重叠长度,默认值为 80
    chunk_overlap: int = 80
    # 检索返回的相似文档数量,默认值为 5
    top_k: int = 5
    # 相似度阈值,默认值为 0.2
    # cosine 距离越小越相似;此处用 1 - distance 作为相似度
    score_threshold: float = 0.2
    # OpenAI 兼容接口的基础 URL,默认指向 DeepSeek API
+   openai_base_url: str = "https://api.deepseek.com"
    # OpenAI 兼容接口的 API 密钥,默认值为空字符串
+   openai_api_key: str = ""
    # 大语言模型名称,默认值为 deepseek-v4-flash
+   openai_model: str = "deepseek-v4-flash"
    # 生成温度参数,默认值为 0.2,数值越低输出越稳定
+   temperature: float = 0.2
    # 最大上下文字符数,默认值为 6000
+   max_context_chars: int = 6000

    # 将 from_env 声明为类方法
    @classmethod
    # 定义从环境变量加载配置的类方法
    def from_env(cls):
        # 使用环境变量(若不存在则回退到类默认值)创建并返回配置实例
        return cls(
            # 从环境变量 RAG_DB_PATH 读取数据库路径,缺省使用类默认值
            db_path=os.getenv("RAG_DB_PATH", cls.db_path),
            # 从环境变量 RAG_COLLECTION 读取集合名称,缺省使用类默认值
            collection_name=os.getenv("RAG_COLLECTION", cls.collection_name),
            # 从环境变量 RAG_EMBEDDING_MODEL 读取嵌入模型名称,缺省使用类默认值
            embedding_model=os.getenv("RAG_EMBEDDING_MODEL", cls.embedding_model),
            # 从环境变量 RAG_CHUNK_SIZE 读取分块大小,缺省使用类默认值
            chunk_size=int(os.getenv("RAG_CHUNK_SIZE", str(cls.chunk_size))),
            # 从环境变量 RAG_CHUNK_OVERLAP 读取分块重叠长度,缺省使用类默认值
            chunk_overlap=int(os.getenv("RAG_CHUNK_OVERLAP", str(cls.chunk_overlap))),
            # 从环境变量 RAG_TOP_K 读取检索返回的相似文档数量,缺省使用类默认值
            top_k=int(os.getenv("RAG_TOP_K", str(cls.top_k))),
            # 从环境变量 RAG_SCORE_THRESHOLD 读取相似度阈值,缺省使用类默认值
            score_threshold=float(
                # 读取环境变量 RAG_SCORE_THRESHOLD,若不存在则回退到类默认值并转为字符串
                os.getenv("RAG_SCORE_THRESHOLD", str(cls.score_threshold))
                # 结束 float 转换与 score_threshold 参数赋值
            ),
            # 从环境变量 OPENAI_API_BASE 读取 OpenAI 兼容接口的基础 URL
+           openai_base_url=os.getenv("OPENAI_API_BASE")
            # 若 OPENAI_API_BASE 为空则回退到 OPENAI_BASE_URL,再缺省使用类默认值
+           or os.getenv("OPENAI_BASE_URL", cls.openai_base_url),
            # 从环境变量 DEEPSEEK_API_KEY 读取 API 密钥
+           openai_api_key=os.getenv("DEEPSEEK_API_KEY")
            # 若 DEEPSEEK_API_KEY 为空则回退到 OPENAI_API_KEY,再缺省为空字符串
+           or os.getenv("OPENAI_API_KEY", ""),
            # 从环境变量 OPENAI_MODEL_NAME 读取大语言模型名称,缺省使用类默认值
+           openai_model=os.getenv("OPENAI_MODEL_NAME", cls.openai_model),
            # 从环境变量 RAG_TEMPERATURE 读取生成温度并转为浮点数,缺省使用类默认值
+           temperature=float(os.getenv("RAG_TEMPERATURE", str(cls.temperature))),
            # 从环境变量 RAG_MAX_CONTEXT_CHARS 读取最大上下文字符数并转为整数
+           max_context_chars=int(
                # 读取环境变量 RAG_MAX_CONTEXT_CHARS,若不存在则回退到类默认值并转为字符串
+               os.getenv("RAG_MAX_CONTEXT_CHARS", str(cls.max_context_chars))
                # 结束 int 转换与 max_context_chars 参数赋值
+           ),
        )

7.4 pipeline.py #

pipeline.py

# 从 pathlib 导入 Path,用于处理与解析文件路径
from pathlib import Path
import logging

# 从 loader 模块导入 DocumentLoader,用于加载多格式非结构化文档
from loader import DocumentLoader

# 从 rich 库导入增强版 print,支持彩色与富文本输出
from rich import print

# 从 store 模块导入 VectorStore,用于向量存储与检索
from store import VectorStore

# 从 config 模块导入 RAGConfig,用于配置 RAG 系统
from config import RAGConfig

# 从 embeddings 模块导入 EmbeddingService,用于向量化服务
from embeddings import EmbeddingService

# 从 chunker 模块导入 TextChunker,用于文本分块
from chunker import TextChunker
+from llm import LLMClient

logger = logging.getLogger(__name__)


# 定义 RAG 主流程编排器类
class RAGPipeline:
    # 类文档字符串:说明该类负责编排 RAG 主流程
    """RAG 主流程编排器。"""

    # 初始化方法,可选接收配置参数
    def __init__(self, config=None):
        self.config = config or RAGConfig.from_env()
        # 创建文档加载器实例并保存为实例属性
        self.loader = DocumentLoader()
        # 创建向量化服务实例并保存为实例属性
        self.embedder = EmbeddingService(self.config.embedding_model)
        # 创建向量存储实例并保存为实例属性
        self.store = VectorStore(
            db_path=self.config.db_path,  # 数据库路径
            collection_name=self.config.collection_name,  # 集合名称
            embedding_service=self.embedder,  # 向量化服务
        )
        # 创建文本分块器实例并保存为实例属性
        self.chunker = TextChunker(
            chunk_size=self.config.chunk_size,  # 分块大小
            chunk_overlap=self.config.chunk_overlap,  # 分块重叠长度
        )
        # 创建大语言模型客户端实例并保存为实例属性
+       self.llm = LLMClient(self.config)

    # 定义文件入库方法,接收待入库的文件路径
    def ingest_file(self, file_path):
        # 将输入路径转换为绝对路径对象
        path = Path(file_path).resolve()
        # 校验路径是否指向一个真实存在的文件
        if not path.is_file():
            # 若不是文件则抛出参数错误,提示需指定单个文件路径
            raise ValueError(f"请指定单个文件路径: {path}")
        # 调用文档加载器解析文件内容为文本
        text = self.loader.load(path)
        # 判断清洗后的文本是否为空
        if not text.strip():
            # 文本为空时返回 0,表示未生成有效内容
            return 0
        # 重置向量存储
        self.store.reset()
        # 调用文本分块器将文本切分为文本块列表
        chunks = self.chunker.split(text, source=str(path.name))
        # 调用向量存储实例将文本块列表插入到向量存储中
        n = self.store.upsert_chunks(chunks)
        # 记录入库成功日志,包含文件名与文本块数量
        logger.info("入库成功: %s → %d chunks", path.name, n)
        return n

    # 定义检索方法,接收问题文本与可选的返回数量
    def retrieve(self, question, top_k=None):
        # 调用向量存储的相似度检索,并返回命中结果
        return self.store.similarity_search(
            # 传入查询问题文本
            query=question,
            # 传入返回数量,未指定时回退为配置中的默认 top_k
            top_k=top_k or self.config.top_k,
            # 传入配置中的相似度分数阈值
            score_threshold=self.config.score_threshold,
            # 结束 similarity_search 调用
        )

    # 定义构建提示词方法,接收用户问题与检索命中结果
+   def build_prompt(self, question, hits):
        # 判断检索命中结果是否为空
+       if not hits:
            # 无命中时返回固定提示,说明知识库暂无相关信息
+           return (
                # 标明参考资料为空
+               "参考资料:无\n\n"
                # 拼接用户问题文本
+               f"用户问题:{question}\n\n"
                # 提示模型说明知识库中暂无相关信息
+               "请说明知识库中暂无相关信息。"
                # 结束无命中时的返回元组
+           )

        # 初始化用于拼接参考资料片段的列表
+       parts = []
        # 初始化已使用的上下文字符数计数器
+       used = 0
        # 从 1 开始枚举每条检索命中结果
+       for i, hit in enumerate(hits, start=1):
            # 构造单条参考资料文本块,包含序号、来源、相似度与正文
+           block = (
                # 写入序号、来源与四位小数相似度分数
+               f"[{i}] 来源: {hit.source} | 相似度: {hit.score:.4f}\n"
                # 写入命中文本块正文并换行
+               f"{hit.content}\n"
                # 结束单条参考资料文本块构造
+           )
            # 判断追加当前文本块后是否会超出最大上下文字符数限制
+           if used + len(block) > self.config.max_context_chars:
                # 超出限制则停止继续追加更多命中结果
+               break
            # 将当前文本块追加到参考资料片段列表
+           parts.append(block)
            # 累加当前文本块占用的字符数
+           used += len(block)

        # 将各参考资料片段用空行连接为完整上下文
+       context = "\n".join(parts)
        # 返回包含参考资料与用户问题的完整提示词
+       return (
            # 写入参考资料引导语
+           "以下是从知识库检索到的参考资料:\n"
            # 写入拼接后的参考资料上下文
+           f"{context}\n"
            # 写入分隔线
+           "--------------------------------\n"
            # 写入用户问题文本
+           f"用户问题:{question}\n"
            # 写入作答要求说明
+           "请基于上述资料作答。"
            # 结束完整提示词返回元组
+       )

    # 定义问答方法,接收问题文本与可选的返回数量
    def ask(self, question, top_k=None):
        # 去除问题文本首尾空白字符
        question = question.strip()
        # 判断清洗后的问题是否为空
        if not question:
            # 问题为空时抛出参数错误
            raise ValueError("问题不能为空")

        # 记录开始 RAG 问答的信息日志,包含问题内容
        logger.info("开始 RAG 问答: %s", question)
        # 调用检索方法获取与问题相关的命中结果
        hits = self.retrieve(question, top_k=top_k)
        # 根据问题和检索命中结果构建提示词
+       prompt = self.build_prompt(question, hits)
        # 判断检索结果是否为空
+       if not hits:
            # 无命中时使用固定提示作为答案
+           answer = "根据现有知识库无法确定相关答案,请补充文档后重试。"
        # 存在检索命中结果时走生成分支
+       else:
            # 调用大语言模型根据提示词生成答案
+           answer = self.llm.generate(prompt)
        # 将最终答案输出到标准输出
+       print(answer)

8. 结构化答案与引用展示 #

本节把第 6 节「只打印答案字符串」升级为可追溯的结构化结果:RAGAnswer 聚合问题、答案、检索引用(citations)、所用模型与 prompt_preview(前 500 字预览),并提供 to_dict 便于序列化;ask 生成后封装并返回该对象,不再在 Pipeline 内直接打印。CLI 侧用 print_answer 格式化输出:问题 / 答案 / 逐条引用(来源、分数、正文摘要),无命中时明确显示「引用证据: (无)」。

整体流程:检索与生成照旧 → 组装 RAGAnswer → CLI 美化打印(含引用证据)。

sequenceDiagram autonumber actor User as 用户 participant CLI as cmd_query() participant Pipe as RAGPipeline.ask participant LLM as LLMClient participant Out as print_answer() User->>CLI: query -q "年假有多少天?" CLI->>Pipe: ask(question) Pipe->>Pipe: retrieve → build_prompt alt 无命中 Pipe->>Pipe: answer = 固定「无法确定」提示 else 有命中 Pipe->>LLM: generate(prompt) LLM-->>Pipe: answer end Pipe->>Pipe: 构造 RAGAnswer<br/>question / answer / citations / model / prompt_preview Pipe-->>CLI: result CLI->>Out: print_answer(result) Out->>Out: 打印问题与答案 alt 有 citations Out->>Out: 逐条输出来源、score、正文摘要(≤120 字) else 无引用 Out->>Out: 引用证据: (无) end Out-->>User: 格式化终端输出

8.1 utils.py #

utils.py

def print_answer(result):
    print("\n" + "=" * 60)
    print(f"问题: {result.question}")
    print("-" * 60)
    print(result.answer)
    print("-" * 60)
    if result.citations:
        print("引用证据:")
        for i, hit in enumerate(result.citations, start=1):
            snippet = hit.content.replace("\n", " ")
            if len(snippet) > 120:
                snippet = snippet[:120] + "..."
            print(f"  [{i}] {hit.source}  score={hit.score:.4f}")
            print(f"      {snippet}")
    else:
        print("引用证据: (无)")
    print("=" * 60 + "\n")

8.2 cli.py #

cli.py

# 导入日志模块,用于记录程序运行信息
import logging

# 导入系统模块,用于访问标准错误输出等系统功能
import sys

# 导入命令行参数解析模块
import argparse

# 导入 JSON 模块,用于序列化输出数据
import json

# 从 rich 库导入增强版 print,支持彩色与富文本输出
from rich import print

# 导入 RAG 主流程编排器
from pipeline import RAGPipeline

# 导入 RAG 配置
from config import RAGConfig

# 从 utils 模块导入 print_answer,用于打印 RAG 问答结果
+from utils import print_answer

# 配置全局日志的基本格式与级别
logging.basicConfig(
    # 设置日志级别为 INFO,输出信息及以上级别日志
    level=logging.INFO,
    # 定义日志输出格式:时间、级别、记录器名称与消息内容
    format="%(asctime)s [%(levelname)s] %(name)s - %(message)s",
    # 定义时间戳的显示格式为年-月-日 时:分:秒
    datefmt="%Y-%m-%d %H:%M:%S",
    # 结束 logging.basicConfig 的参数配置
)
# 获取当前模块对应的日志记录器实例
logger = logging.getLogger(__name__)


# 定义 ingest 子命令的处理函数,接收解析后的参数对象
def cmd_ingest(args):
    # 创建 RAG 主流程编排器实例
    pipeline = RAGPipeline()
    # 调用 ingest_file 方法入库文件
    chunks = pipeline.ingest_file(args.path)
    # 将文件路径封装为字典并以格式化 JSON 打印输出
    print(
        print(
            json.dumps(
                {"file": args.path, "chunks": chunks}, ensure_ascii=False, indent=2
            )
        )
    )


# 定义 query 子命令的处理函数,接收解析后的参数对象
def cmd_query(args):
    # 从环境变量加载配置并创建 RAG 主流程编排器实例
    pipeline = RAGPipeline(RAGConfig.from_env())
    # 调用 ask 方法对用户问题进行检索增强问答
+   result = pipeline.ask(args.question)
+   print_answer(result)


# 构建并返回命令行参数解析器
def build_parser():
    # 创建顶层 ArgumentParser,设置程序描述与帮助格式化器
    parser = argparse.ArgumentParser(
        # 设置程序用途说明:RAG 文件入库与问答
        description="RAG:文件入库 + 问答",
        # 使用原始描述帮助格式化器,保留描述中的换行与缩进
        formatter_class=argparse.RawDescriptionHelpFormatter,
        # 结束 ArgumentParser 的参数配置
    )
    # 添加子命令解析器,结果写入 args.command,且必须指定子命令
    sub = parser.add_subparsers(dest="command", required=True)

    # 添加名为 ingest 的子命令,用于文件入库
    p_ingest = sub.add_parser("ingest", help="入库文件")
    # 为 ingest 子命令添加必填的 --path 参数,表示待入库文件路径
    p_ingest.add_argument("--path", required=True, help="文件路径")
    # 将 ingest 子命令的默认处理函数绑定为 cmd_ingest
    p_ingest.set_defaults(func=cmd_ingest)

    # 添加名为 query 的子命令,用于问答
    p_query = sub.add_parser("query", help="问答")
    # 为 query 子命令添加必填的 --question/-q 参数,表示用户问题
    p_query.add_argument("--question", "-q", required=True, help="用户问题")
    # 将 query 子命令的默认处理函数绑定为 cmd_query
    p_query.set_defaults(func=cmd_query)
    # 返回配置完成的参数解析器
    return parser


# 定义程序主入口函数
def main():
    # 调用 build_parser 构建命令行参数解析器
    parser = build_parser()
    # 解析命令行参数,得到命名空间对象 args
    args = parser.parse_args()
    # 打印解析得到的参数对象,便于调试查看
    print(args)
    # 尝试执行子命令对应的处理函数
    try:
        # 调用当前子命令绑定的处理函数并传入参数
        args.func(args)
        # 执行成功时返回退出码 0
        return 0
    # 捕获用户按下 Ctrl+C 产生的键盘中断异常
    except KeyboardInterrupt:
        # 打印已中断提示信息
        print("\n已中断")
        # 返回标准的键盘中断退出码 130
        return 130
    # 捕获其余所有异常(忽略过宽捕获的静态检查警告)
    except Exception as exc:  # noqa: BLE001
        # 记录完整异常堆栈,提示执行失败
        logger.exception("执行失败")
        # 将错误信息输出到标准错误流
        print(f"[错误] {exc}", file=sys.stderr)
        # 返回表示失败的退出码 1
        return 1


# 判断当前模块是否作为脚本直接运行
if __name__ == "__main__":
    # 调用主函数启动程序
    main()

8.3 models.py #

models.py

# 从 dataclasses 模块导入 asdict、dataclass 与 field,用于定义数据类与字段默认值
from dataclasses import asdict, dataclass, field


# 使用 dataclass 装饰器,将类自动转换为数据类
@dataclass
# 定义入库后的文本块数据模型类
class DocumentChunk:
    # 类文档字符串:说明该类表示入库后的文本块
    """入库后的文本块。"""

    # 文本块的唯一标识符
    id: str
    # 文本块的正文内容
    content: str
    # 文本块来源(如文件名)
    source: str
    # 文本块在原文中的序号索引
    chunk_index: int
    # 附加元数据字典,默认创建空字典
    metadata: dict = field(default_factory=dict)


# 使用 dataclass 装饰器,将类自动转换为数据类
@dataclass
# 定义单条检索命中数据模型类
class RetrievalHit:
    # 类文档字符串:说明该类表示单条检索命中
    """单条检索命中。"""

    # 检索命中的正文内容
    content: str
    # 检索命中的来源(如文件名)
    source: str
    # 检索命中对应文本块的唯一标识符
    chunk_id: str
    # 检索相似度得分
    score: float
    # 附加元数据字典,默认创建空字典
    metadata: dict = field(default_factory=dict)


# 使用 dataclass 装饰器,将类自动转换为数据类
+@dataclass
# 定义面向调用方的完整问答结果数据模型类
+class RAGAnswer:
    # 类文档字符串:说明该类表示面向调用方的完整问答结果(含可追溯引用)
+   """面向调用方的完整问答结果(含可追溯引用)。"""

    # 用户提出的问题文本
+   question: str
    # 大模型生成的最终答案文本
+   answer: str
    # 可追溯的检索命中引用列表
+   citations: list
    # 生成答案所使用的大语言模型名称
+   model: str
    # 提示词预览文本,默认空字符串
+   prompt_preview: str = ""

    # 定义将数据类实例转换为字典的方法
+   def to_dict(self):
        # 调用 asdict 将当前实例转为普通字典并返回
+       return asdict(self)

8.4 pipeline.py #

pipeline.py

# 从 pathlib 导入 Path,用于处理与解析文件路径
from pathlib import Path
import logging

# 从 loader 模块导入 DocumentLoader,用于加载多格式非结构化文档
from loader import DocumentLoader

# 从 rich 库导入增强版 print,支持彩色与富文本输出
from rich import print

# 从 store 模块导入 VectorStore,用于向量存储与检索
from store import VectorStore

# 从 config 模块导入 RAGConfig,用于配置 RAG 系统
from config import RAGConfig

# 从 embeddings 模块导入 EmbeddingService,用于向量化服务
from embeddings import EmbeddingService

# 从 chunker 模块导入 TextChunker,用于文本分块
from chunker import TextChunker

# 从 llm 模块导入 LLMClient,用于大语言模型调用
from llm import LLMClient

# 从 models 模块导入 RAGAnswer,用于封装 RAG 问答结果
+from models import RAGAnswer

logger = logging.getLogger(__name__)


# 定义 RAG 主流程编排器类
class RAGPipeline:
    # 类文档字符串:说明该类负责编排 RAG 主流程
    """RAG 主流程编排器。"""

    # 初始化方法,可选接收配置参数
    def __init__(self, config=None):
        self.config = config or RAGConfig.from_env()
        # 创建文档加载器实例并保存为实例属性
        self.loader = DocumentLoader()
        # 创建向量化服务实例并保存为实例属性
        self.embedder = EmbeddingService(self.config.embedding_model)
        # 创建向量存储实例并保存为实例属性
        self.store = VectorStore(
            db_path=self.config.db_path,  # 数据库路径
            collection_name=self.config.collection_name,  # 集合名称
            embedding_service=self.embedder,  # 向量化服务
        )
        # 创建文本分块器实例并保存为实例属性
        self.chunker = TextChunker(
            chunk_size=self.config.chunk_size,  # 分块大小
            chunk_overlap=self.config.chunk_overlap,  # 分块重叠长度
        )
        # 创建大语言模型客户端实例并保存为实例属性
        self.llm = LLMClient(self.config)

    # 定义文件入库方法,接收待入库的文件路径
    def ingest_file(self, file_path):
        # 将输入路径转换为绝对路径对象
        path = Path(file_path).resolve()
        # 校验路径是否指向一个真实存在的文件
        if not path.is_file():
            # 若不是文件则抛出参数错误,提示需指定单个文件路径
            raise ValueError(f"请指定单个文件路径: {path}")
        # 调用文档加载器解析文件内容为文本
        text = self.loader.load(path)
        # 判断清洗后的文本是否为空
        if not text.strip():
            # 文本为空时返回 0,表示未生成有效内容
            return 0
        # 重置向量存储
        self.store.reset()
        # 调用文本分块器将文本切分为文本块列表
        chunks = self.chunker.split(text, source=str(path.name))
        # 调用向量存储实例将文本块列表插入到向量存储中
        n = self.store.upsert_chunks(chunks)
        # 记录入库成功日志,包含文件名与文本块数量
        logger.info("入库成功: %s → %d chunks", path.name, n)
        return n

    # 定义检索方法,接收问题文本与可选的返回数量
    def retrieve(self, question, top_k=None):
        # 调用向量存储的相似度检索,并返回命中结果
        return self.store.similarity_search(
            # 传入查询问题文本
            query=question,
            # 传入返回数量,未指定时回退为配置中的默认 top_k
            top_k=top_k or self.config.top_k,
            # 传入配置中的相似度分数阈值
            score_threshold=self.config.score_threshold,
            # 结束 similarity_search 调用
        )

    # 定义构建提示词方法,接收用户问题与检索命中结果
    def build_prompt(self, question, hits):
        # 判断检索命中结果是否为空
        if not hits:
            # 无命中时返回固定提示,说明知识库暂无相关信息
            return (
                # 标明参考资料为空
                "参考资料:无\n\n"
                # 拼接用户问题文本
                f"用户问题:{question}\n\n"
                # 提示模型说明知识库中暂无相关信息
                "请说明知识库中暂无相关信息。"
                # 结束无命中时的返回元组
            )

        # 初始化用于拼接参考资料片段的列表
        parts = []
        # 初始化已使用的上下文字符数计数器
        used = 0
        # 从 1 开始枚举每条检索命中结果
        for i, hit in enumerate(hits, start=1):
            # 构造单条参考资料文本块,包含序号、来源、相似度与正文
            block = (
                # 写入序号、来源与四位小数相似度分数
                f"[{i}] 来源: {hit.source} | 相似度: {hit.score:.4f}\n"
                # 写入命中文本块正文并换行
                f"{hit.content}\n"
                # 结束单条参考资料文本块构造
            )
            # 判断追加当前文本块后是否会超出最大上下文字符数限制
            if used + len(block) > self.config.max_context_chars:
                # 超出限制则停止继续追加更多命中结果
                break
            # 将当前文本块追加到参考资料片段列表
            parts.append(block)
            # 累加当前文本块占用的字符数
            used += len(block)

        # 将各参考资料片段用空行连接为完整上下文
        context = "\n".join(parts)
        # 返回包含参考资料与用户问题的完整提示词
        return (
            # 写入参考资料引导语
            "以下是从知识库检索到的参考资料:\n"
            # 写入拼接后的参考资料上下文
            f"{context}\n"
            # 写入分隔线
            "--------------------------------\n"
            # 写入用户问题文本
            f"用户问题:{question}\n"
            # 写入作答要求说明
            "请基于上述资料作答。"
            # 结束完整提示词返回元组
        )

    # 定义问答方法,接收问题文本与可选的返回数量
    def ask(self, question, top_k=None):
        # 去除问题文本首尾空白字符
        question = question.strip()
        # 判断清洗后的问题是否为空
        if not question:
            # 问题为空时抛出参数错误
            raise ValueError("问题不能为空")

        # 记录开始 RAG 问答的信息日志,包含问题内容
        logger.info("开始 RAG 问答: %s", question)
        # 调用检索方法获取与问题相关的命中结果
        hits = self.retrieve(question, top_k=top_k)
        # 根据问题和检索命中结果构建提示词
        prompt = self.build_prompt(question, hits)
        # 判断检索结果是否为空
        if not hits:
            # 无命中时使用固定提示作为答案
            answer = "根据现有知识库无法确定相关答案,请补充文档后重试。"
        # 存在检索命中结果时走生成分支
        else:
            # 调用大语言模型根据提示词生成答案
            answer = self.llm.generate(prompt)
        # 将最终答案输出到标准输出
+       result = RAGAnswer(
+           question=question,
+           answer=answer,
+           citations=hits,
+           model=self.config.openai_model,
+           prompt_preview=prompt[:500] + ("..." if len(prompt) > 500 else ""),
+       )
+       logger.info("问答完成: citations=%d", len(hits))
+       return result