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- 作用:把指定文档解析、分块并写入向量库,供后续检索使用。
- 参数:
--path指向待入库文件(如handbook.md);也可换成目录,批量处理其下支持的文档。 - 预期:终端输出入库结果(如文件名、chunk 数量);成功后知识库中已有该文档的向量索引。
1.2.2 问答(query) #
uv run cli.py query --question 公司年假怎么申请?- 作用:对问题做相似度检索,拼装提示词并调用大模型生成答案,同时展示引用证据。
- 参数:
--question(或简写-q)为用户问题文本。 - 预期:打印问题、模型答案,以及命中的来源、相似度与正文摘要;若无相关文档,会提示知识库暂无答案并显示「引用证据: (无)」。
建议先执行 ingest 完成入库,再执行 query 验证检索与生成是否正常。
1.3 参考链接 #
- paraphrase-multilingual-MiniLM-L12-v2
- openai
- ChromaDB
- langchain-text-splitters
- pymupdf
- python-docx
- python-pptx
- openpyxl
- beautifulsoup4
- rich
2. 命令行入口(cli.py) #
本节介绍 RAG 项目的命令行入口:如何用 argparse 定义子命令、解析参数,并在 main 中统一调度与异常处理。当前实现已支持 ingest(入库)子命令,后续可按相同模式扩展 ask 等问答能力。
整体流程:用户在终端输入命令 → main 构建解析器并解析参数 → 按子命令调用对应处理函数 → 返回退出码。
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.py 的 cmd_ingest 改为委托 RAGPipeline.ingest_file,从而把命令行与具体解析逻辑解耦。
整体流程:ingest 子命令 → RAGPipeline.ingest_file → DocumentLoader.load(校验 / 选解析器 / 规范化)→ 打印文本;空文本时返回 0。
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 加载文本 → 非空则重置集合。
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 → 返回写入数量。
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_k 与 score_threshold;RetrievalHit 描述单条命中;embed_query 将问题编码为查询向量;VectorStore.similarity_search 在 Chroma 中查询,把 cosine 距离转为相似度(1 - distance)并按阈值过滤、按分数排序。CLI 新增 query -q,经 ask → retrieve 打印命中列表。
整体流程:校验问题非空 → 问题向量化 → 向量库 top-k 检索 → 阈值过滤与排序 → 输出 RetrievalHit[]。
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.26.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 加载 .env。build_prompt 把命中块拼成带来源与相似度的上下文(受字符上限截断);ask 在检索后:无命中直接返回固定提示,有命中则 llm.generate 产出最终答案。
整体流程:检索 hits → 组装 user prompt →(有命中)调用聊天补全 → 打印答案。
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=60007.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 美化打印(含引用证据)。
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