跳转至

rag_kb · RAG + ChromaDB 检索管线

本页收录可运行示例,配套本文档「RAG 检索管线」(RAG 检索管线)。 代码为随框架交付的示例,已一并收录到本站供直接对照使用。

说明与运行

示例: 知识库(RAG)检索 —— CSV 文档 → 切分/向量化 → ChromaDB 存储 → 检索. 演示 agentframework.retrieval 能力包对接真实第三方向量库 ChromaDB 的全流程: 1. 切分: 读取《知识库文档.csv》(公积金办事指南问答, 每行一问答, 每行视为一个文档), 长文本经 RecursiveCharacterSplitter 递归切分; 2. 向量化: embedder 将每个块编码为向量(默认 HashEmbedder 离线可跑, 语义检索请替换 _build_embedder 为真实中文 embedding 实现); 3. 向量存储: ChromaVectorStore 实现框架 VectorStore 协议, 持久化到 chromadb.PersistentClient(默认 examples/rag/.chroma_db); 4. 检索: create_retriever + search_with_citations 输出带引用 (source/chunk_index/score) 的命中, 支持交互提问. 运行: cd agentframework

首次: 建 venv(需 python>=3.11)并安装框架依赖 + 示例依赖 chromadb

python3 -m venv .venv && source .venv/bin/activate pip install -e ".[dev]" chromadb

之后(在已安装依赖的 venv 内执行; 缺 pydantic/chromadb 说明当前解释器未装依赖)

PYTHONPATH=src python examples/rag/rag_kb.py # 幂等重建 + 演示查询(交互输入 Ctrl-C 退出) PYTHONPATH=src python examples/rag/rag_kb.py --query "租房提取住房公积金办理材料" PYTHONPATH=src python examples/rag/rag_kb.py --no-ingest --query "..." # 只查不重建 PYTHONPATH=src python examples/rag/rag_kb.py --fresh # 更换 embedder / 维度变更时清空重建

运行前需在已安装 agentframeworkPython ≥ 3.11 环境;示例多为自带本地模型,无需真实 API Key;需要外部依赖(如 chromadb / 数据库驱动)会在说明中注明。

完整源码

"""示例: 知识库(RAG)检索 —— CSV 文档 → 切分/向量化 → ChromaDB 存储 → 检索.

演示 agentframework.retrieval 能力包对接真实第三方向量库 ChromaDB 的全流程:
1. 切分:  读取《知识库文档.csv》(公积金办事指南问答, 每行一问答, 每行视为一个文档),
          长文本经 RecursiveCharacterSplitter 递归切分;
2. 向量化: embedder 将每个块编码为向量(默认 HashEmbedder 离线可跑,
          语义检索请替换 _build_embedder 为真实中文 embedding 实现);
3. 向量存储: ChromaVectorStore 实现框架 VectorStore 协议, 持久化到
           chromadb.PersistentClient(默认 examples/rag/.chroma_db);
4. 检索:   create_retriever + search_with_citations 输出带引用
           (source/chunk_index/score) 的命中, 支持交互提问.

运行:
    cd agentframework
    # 首次: 建 venv(需 python>=3.11)并安装框架依赖 + 示例依赖 chromadb
    python3 -m venv .venv && source .venv/bin/activate
    pip install -e ".[dev]" chromadb
    # 之后(在已安装依赖的 venv 内执行; 缺 pydantic/chromadb 说明当前解释器未装依赖)
    PYTHONPATH=src python examples/rag/rag_kb.py         # 幂等重建 + 演示查询(交互输入 Ctrl-C 退出)
    PYTHONPATH=src python examples/rag/rag_kb.py --query "租房提取住房公积金办理材料"
    PYTHONPATH=src python examples/rag/rag_kb.py --no-ingest --query "..."   # 只查不重建
    PYTHONPATH=src python examples/rag/rag_kb.py --fresh # 更换 embedder / 维度变更时清空重建
"""

from __future__ import annotations

import argparse
import csv
import hashlib
import math
import os
import sys
from pathlib import Path

from agentframework.retrieval import (
    BaseEmbedder,
    DocumentChunk,
    HashEmbedder,
    RecursiveCharacterSplitter,
    SearchResult,
    VectorStore,
    create_retriever,
    ingest_document,
)

# 演示查询: 前三条取自知识库原文标题(HashEmbedder 精确匹配命中),
# 最后一条为自然口语改写(近似匹配分数较低, 语义 embedding 下效果更佳).
DEMO_QUERIES = [
    "租房提取住房公积金办理材料",
    "提前偿还住房公积金贷款办理材料",
    "商业性个人住房贷款转住房公积金贷款的申请条件",
    "离职很久了想一次性取公积金需要什么?",
]

# 默认 chroma 持久化目录(与 csv 同目录, 已 gitignore), collection 名。
DEFAULT_CHROMA_DIR = Path(__file__).resolve().parent / ".chroma_db"
DEFAULT_COLLECTION = "zhijin_kb"


class ChromaVectorStore:
    """ChromaDB 向量存储(实现框架 retrieval.VectorStore 协议).

    框架 VectorStore 只约定 add/search/delete_by_source/count 四个方法
    (见 src/agentframework/retrieval/store.py), 存哪个向量库由使用方决定;
    本例对接 chromadb 并落盘持久化。要点:
    - add/search 前对向量做 L2 归一化, 使 chroma 的 L2 距离与余弦相似度同序,
      返回分数换算: 余弦相似度 = 1 - d^2 / 2 (0~1, 越大越相关, 与 MemoryVectorStore 对齐);
    - 块元数据只写 str/int/bool(chroma 约束), id 用 source_index_hash 保证稳定唯一。
    """

    def __init__(self, path: str, collection_name: str = DEFAULT_COLLECTION):
        """创建持久化客户端并获取/创建 collection.

        Args:
            path: chroma 落盘目录(不存在会自动创建).
            collection_name: collection 名(更换 embedder/维度时用 drop_collection 重建).

        Raises:
            RuntimeError: chromadb 未安装时给出安装指引.
        """
        try:
            import chromadb
            from chromadb.config import Settings
        except ImportError as exc:  # pragma: no cover - 依赖缺失提示
            raise RuntimeError("示例依赖 chromadb 未安装, 请先执行: pip install chromadb") from exc
        os.makedirs(path, exist_ok=True)
        self._name = collection_name
        self._client = chromadb.PersistentClient(
            path=path, settings=Settings(anonymized_telemetry=False)
        )
        self._coll = self._get_or_create_collection()

    def _get_or_create_collection(self):
        """取已有 collection; 不存在则创建(显式 l2 空间, 配合向量归一化做余弦)."""
        try:
            return self._client.get_collection(name=self._name)
        except Exception:  # chroma 缺 collection 时抛异常, 捕获后创建
            return self._client.create_collection(name=self._name, metadata={"hnsw:space": "l2"})

    # ---------- VectorStore 协议实现 ----------

    def add(self, chunks: list[DocumentChunk], vectors: list[list[float]]) -> None:
        """批量写入块 + 归一化向量到 chroma."""
        if len(chunks) != len(vectors):
            raise ValueError("chunks 与 vectors 数量不一致")
        if not chunks:
            return
        ids: list[str] = []
        documents: list[str] = []
        metadatas: list[dict] = []
        embeddings: list[list[float]] = []
        for index, (chunk, vector) in enumerate(zip(chunks, vectors, strict=True)):
            source = chunk.metadata.get("source", f"chunk-{index}")
            chunk_index = chunk.metadata.get("chunk_index", index)
            ids.append(_stable_id(source, chunk_index, chunk.text))
            documents.append(chunk.text)
            metadatas.append(dict(chunk.metadata))
            embeddings.append(_normalize_vector(vector))
        self._coll.add(ids=ids, embeddings=embeddings, documents=documents, metadatas=metadatas)

    def search(self, vector: list[float], *, top_k: int = 5) -> list[SearchResult]:
        """按归一化向量在 chroma 检索最相似的块(余弦相似度排序)."""
        if top_k <= 0:
            return []
        total = self._coll.count()
        if total == 0:
            return []
        query_vector = _normalize_vector(vector)
        n_results = min(top_k, total)
        result = self._coll.query(query_embeddings=[query_vector], n_results=n_results)
        results: list[SearchResult] = []
        documents = (result.get("documents") or [[]])[0] or []
        metadatas = (result.get("metadatas") or [[]])[0] or []
        distances = (result.get("distances") or [[]])[0] or []
        for text, meta, distance in zip(documents, metadatas, distances, strict=True):
            if text is None or meta is None:
                continue
            # 单位向量下: ||u-v||^2 = 2 - 2cos => cos = 1 - d^2/2
            score = 1.0 - (float(distance) ** 2) / 2.0
            score = max(0.0, min(1.0, score))
            results.append(
                SearchResult(chunk=DocumentChunk(text=text, metadata=dict(meta)), score=score)
            )
        return results

    def delete_by_source(self, source: str) -> int:
        """按来源删除该文档全部块(重导入时先删旧), 返回删除条数."""
        found = self._coll.get(where={"source": source}, include=[])
        ids = found.get("ids") or []
        if ids:
            self._coll.delete(ids=ids)
        return len(ids)

    def count(self) -> int:
        """当前块总数."""
        return self._coll.count()

    def drop_collection(self) -> None:
        """删除并重建 collection(--fresh / embedder 维度变更用)."""
        try:
            self._client.delete_collection(name=self._name)
        except Exception:  # 不存在时忽略
            pass
        self._coll = self._get_or_create_collection()


def _normalize_vector(vector: list[float]) -> list[float]:
    """L2 归一化(使 L2 距离序等价于余弦相似度序)."""
    norm = math.sqrt(sum(v * v for v in vector))
    if norm == 0:
        return list(vector)
    return [v / norm for v in vector]


def _stable_id(source: str, chunk_index: int, text: str) -> str:
    """生成稳定唯一 id: {source}_{chunk_index}_{text sha256 前 12 位}."""
    digest = hashlib.sha256(text.encode("utf-8")).hexdigest()[:12]
    return f"{source}_{chunk_index}_{digest}"


def _build_embedder() -> BaseEmbedder:
    """构建 embedding 实现(默认离线确定性哈希, 维度 64).

    HashEmbedder 仅适合链路演示(相同文本→相同向量, 精确/词重叠匹配);
    生产/语义检索请替换为真实中文 embedding, 例如 OpenAI 兼容 API:

        class OpenAIEmbedder(BaseEmbedder):
            def __init__(self, api_key, base_url, model, dimension):
                self._client = httpx.Client(base_url=base_url,
                                            headers={"Authorization": f"Bearer {api_key}"})
                self._model, self._dimension = model, dimension

            @property
            def dimension(self) -> int:
                return self._dimension

            def embed(self, texts):
                resp = self._client.post("/embeddings",
                                         json={"model": self._model, "input": list(texts)})
                return [d["embedding"] for d in resp.json()["data"]]

    注意: 更换 embedder 后向量维度可能变化, 需用 --fresh 清空重建 collection.
    """
    return HashEmbedder(dimension=64)


def load_kb_csv(path: Path) -> list[dict]:
    """解析《知识库文档.csv》为问答文档列表(每行一个文档).

    Args:
        path: csv 文件路径(utf-8-sig, 每行至多 2 列: 标题/问题, 答复/材料).

    Returns:
        [{index, title, source, text}] 列表, 已跳过空行; text = "标题\\n答复".

    Raises:
        FileNotFoundError: 文件不存在.
        RuntimeError: 无有效问答行.
    """
    if not path.exists():
        raise FileNotFoundError(f"知识库文件不存在: {path}")
    docs: list[dict] = []
    with open(path, encoding="utf-8-sig", newline="") as fh:
        for row_number, row in enumerate(csv.reader(fh), start=1):
            title = (row[0] if row else "").strip()
            body = (row[1] if len(row) > 1 else "").strip()
            if not title and not body:
                continue
            text = f"{title}\n{body}" if body else title
            docs.append(
                {
                    "index": row_number,
                    "title": title,
                    "source": f"row-{row_number}",
                    "text": text,
                }
            )
    if not docs:
        raise RuntimeError(f"知识库文件无有效内容: {path}")
    return docs


def build_kb(
    store: VectorStore,
    embedder: BaseEmbedder,
    docs: list[dict],
    *,
    chunk_size: int,
    chunk_overlap: int,
) -> tuple[int, int]:
    """逐条把问答文档导入 chroma(幂等: 同 source 重导入先删旧块).

    Args:
        store: 向量存储(ChromaVectorStore 实例).
        embedder: embedding 实现.
        docs: load_kb_csv 返回的问答文档列表.
        chunk_size / chunk_overlap: RecursiveCharacterSplitter 切分参数.

    Returns:
        (入库总块数, 被切成多块的长文档数).
    """
    splitter = RecursiveCharacterSplitter(chunk_size=chunk_size, chunk_overlap=chunk_overlap)
    total_chunks = 0
    multi_chunk_docs = 0
    for doc in docs:
        chunks = ingest_document(
            doc["text"],
            store=store,
            embedder=embedder,
            source=doc["source"],
            splitter=splitter,
            metadata={"title": doc["title"]},
        )
        total_chunks += len(chunks)
        if len(chunks) > 1:
            multi_chunk_docs += 1
    return total_chunks, multi_chunk_docs


def print_citations(citations: list[dict]) -> None:
    """打印检索命中的引用结构 [{source, chunk_index, text, score}]."""
    if not citations:
        print("  未检索到相关内容(若为空库请先不带 --no-ingest 运行入库).")
        return
    print(f"  命中 {len(citations)} 条:")
    for i, citation in enumerate(citations, start=1):
        source = citation["source"]
        chunk_index = citation["chunk_index"]
        score = citation["score"]
        text = citation["text"].strip().replace("\n", " / ")
        if len(text) > 120:
            text = text[:120] + "…"
        print(f"    [{i}] {source}#{chunk_index}  score={score:.4f}")
        print(f"        {text}")


def run_queries(retriever, queries: list[str], *, top_k: int) -> None:
    """对一组查询执行检索并打印引用结果."""
    for query in queries:
        print(f"\n查询: {query}")
        citations = retriever.search_with_citations(query, top_k=top_k)
        print_citations(citations)


def _read_line() -> str | None:
    """从标准输入读一行, 兼容非 UTF-8 终端粘贴(如 GBK), 避免 input() 抛 UnicodeDecodeError.

    优先按 UTF-8 解码; 失败则退回 GB18030(覆盖 GBK/GB2312); 仍失败则以替换符收尾
    (不中断交互)。注意: 建议终端保持 UTF-8 编码以获得最佳检索效果。

    Returns:
        去掉换行的行文本; 读到 EOF(Ctrl-D)返回 None.
    """
    buffer = getattr(sys.stdin, "buffer", None)
    if buffer is None:  # 非缓冲 stdin 时退回 input()
        try:
            line = input("")
        except EOFError:
            return None
        return line.rstrip("\r\n")
    raw = buffer.readline()
    if raw == b"":  # EOF
        return None
    for encoding in ("utf-8", "gb18030"):
        try:
            return raw.decode(encoding).rstrip("\r\n")
        except UnicodeDecodeError:
            continue
    return raw.decode("utf-8", errors="replace").rstrip("\r\n")


def run_interactive(retriever, *, top_k: int) -> None:
    """交互式提问(仅 tty; Ctrl-C / EOF / exit 退出)."""
    print("\n交互问答(输入 exit/quit 或 Ctrl-C 退出):")
    try:
        while True:
            line = _read_line()
            if line is None:  # EOF(Ctrl-D)
                print()
                break
            query = line.strip()
            if not query:
                continue
            if query.lower() in {"exit", "quit"}:
                break
            citations = retriever.search_with_citations(query, top_k=top_k)
            print_citations(citations)
    except KeyboardInterrupt:
        print("\n再见。")


def _parse_args(argv: list[str] | None) -> argparse.Namespace:
    """解析命令行参数."""
    default_csv = Path(__file__).resolve().parent / "知识库文档.csv"
    parser = argparse.ArgumentParser(
        description="知识库(RAG)检索示例: CSV → 切分/向量化 → ChromaDB → 检索"
    )
    parser.add_argument("--csv", type=Path, default=default_csv, help="知识库问答 CSV")
    parser.add_argument(
        "--chroma-dir", type=str, default=str(DEFAULT_CHROMA_DIR), help="chroma 落盘目录"
    )
    parser.add_argument("--collection", default=DEFAULT_COLLECTION, help="chroma collection 名")
    parser.add_argument("--no-ingest", action="store_true", help="跳过重建, 只查现有库")
    parser.add_argument(
        "--fresh", action="store_true", help="清空重建 collection(更换 embedder/维度变更用)"
    )
    parser.add_argument("--chunk-size", type=int, default=200, help="切分块大小(默认 200)")
    parser.add_argument("--chunk-overlap", type=int, default=50, help="切分块重叠(默认 50)")
    parser.add_argument("--top-k", type=int, default=3, help="检索返回条数(默认 3)")
    parser.add_argument("--query", help="单次检索问题(缺省跑演示查询, tty 下进入交互)")
    return parser.parse_args(argv)


def main(argv: list[str] | None = None) -> int:
    """示例入口: 建库(可选) → 检索 → 交互."""
    args = _parse_args(argv)
    store = ChromaVectorStore(args.chroma_dir, collection_name=args.collection)
    if args.fresh:
        store.drop_collection()
    embedder = _build_embedder()

    if not args.no_ingest:
        docs = load_kb_csv(args.csv)
        print(f"解析知识库: {len(docs)} 行问答  <-  {args.csv}")
        try:
            total_chunks, multi_chunk_docs = build_kb(
                store,
                embedder,
                docs,
                chunk_size=args.chunk_size,
                chunk_overlap=args.chunk_overlap,
            )
        except Exception as exc:  # 向 CLI 用户转友好提示
            message = str(exc)
            if "dimension" in message.lower():
                print("入库失败: 向量维度与现有 collection 不一致(可能更换了 embedder)。")
                print("提示: 请加 --fresh 清空重建后重试。")
            else:
                print(f"入库失败: {message}")
            return 1
        print(
            f"入库完成: {total_chunks} 个块"
            f"({multi_chunk_docs} 行长文本被切为多块), 向量维度={embedder.dimension}"
        )
    else:
        total_chunks = store.count()
        print(f"跳过重建, 读取现有库: {total_chunks} 个块, collection={args.collection!r}")
        if total_chunks == 0:
            print("提示: 知识库为空, 请先不带 --no-ingest 运行一次完成入库。")
            return 0

    retriever = create_retriever(store, embedder, top_k=args.top_k)

    if args.query:
        print(f"\n===== 检索: {args.query} =====")
        run_queries(retriever, [args.query], top_k=args.top_k)
        return 0

    print("\n===== 演示查询 =====")
    run_queries(retriever, DEMO_QUERIES, top_k=args.top_k)

    if sys.stdin.isatty():
        run_interactive(retriever, top_k=args.top_k)
    return 0


if __name__ == "__main__":
    raise SystemExit(main())

回到示例索引