RAG 实战教程:基于 FastAPI + LangChain + Milvus 构建知识库问答系统

一句话总结:本文基于一个完整的政务政策 RAG 项目,从文档上传、父子块切分、混合检索(向量 + BM25 + RRF 融合 + Reranker 重排)、查询改写、引用校验到 LLM 生成回答,再到评估体系,完整拆解生产级 RAG 系统的每个核心模块。

一、项目概述

这是一个政务政策 RAG(检索增强生成)问答系统,用户上传政策文档后,系统自动完成文档解析、切分、向量化入库,用户提问时通过混合检索找到最相关的政策依据,再由大模型生成带引用标注的回答。

核心能力

能力说明
文档管理支持 PDF、DOCX、TXT 上传,自动解析、切分、入库
父子块切分大块保留上下文,小块精准检索,命中后回溯父块
混合检索向量检索 + BM25 关键词检索,RRF 融合排序
Reranker 重排Cross-Encoder 对候选证据精排,提升准确率
查询改写LLM 将口语化问题改写为检索友好的查询
引用校验自动验证回答中的引用编号是否合法
评估体系召回率、MRR、引用准确率、事实准确率、幻觉率
流式回答SSE 推送检索进度和逐 token 生成

二、技术栈与项目结构

技术栈

组件技术说明
Web 框架FastAPI异步 API,支持 SSE 流式推送
ORMSQLAlchemy 2.0声明式映射,支持异步 Session
关系数据库MySQL存储文档元数据、分块记录、会话历史
向量数据库Milvus (milvus-lite)本地文件模式,存储子块向量
EmbeddingDashScope (阿里云)qwen3.7-text-embedding 模型
LLMDeepSeek / Ollama支持云端 API 和本地部署切换
Rerankersentence-transformersBAAI/bge-reranker-base Cross-Encoder
中文分词jiebaBM25 检索的中文分词
配置管理pydantic-settings类型安全的 .env 配置

项目结构

backend/
├── .env                          # 环境变量配置
├── main.py                       # FastAPI 入口
├── requirements.txt              # Python 依赖
├── app/
│   ├── config.py                 # Pydantic Settings 配置类
│   ├── database.py               # SQLAlchemy 引擎与 Session
│   ├── factory.py                # FastAPI 应用工厂
│   ├── models.py                 # SQLAlchemy ORM 模型
│   ├── schemas.py                # Pydantic 请求/响应模型
│   ├── api/                      # API 路由层
│   │   ├── chat.py               # 聊天接口(SSE 流式)
│   │   ├── documents.py          # 文档上传/删除/重建索引
│   │   ├── configuration.py       # 模型配置管理
│   │   └── evaluation.py         # 评估任务管理
│   └── services/                 # 核心业务逻辑
│       ├── chunking.py           # 父子块切分
│       ├── vector_store.py       # Milvus 向量存储
│       ├── retrieval.py          # 混合检索引擎
│       ├── rag_service.py        # RAG 生成服务
│       ├── model_factory.py      # LLM 模型工厂
│       ├── document_service.py   # 文档入库服务
│       ├── pipeline_utils.py     # 工具函数(归一化、引用校验)
│       ├── evaluation_service.py # 评估执行服务
│       └── metrics.py            # 评估指标计算

配置文件 (.env)

APP_NAME=政务政策 RAG 系统
DATABASE_URL=mysql+pymysql://root:123456@127.0.0.1:3306/gov_rag?charset=utf8mb4

# LLM 模型配置
MODEL_PROVIDER=deepseek
DEEPSEEK_API_KEY=sk-your-key
DEEPSEEK_BASE_URL=https://dashscope.aliyuncs.com/compatible-mode/v1
DEEPSEEK_MODEL=deepseek-v4-flash-0731

# 本地 Ollama 备选
OLLAMA_BASE_URL=http://127.0.0.1:11434
OLLAMA_MODEL=qwen2.5:7b

# Embedding 配置
EMBEDDING_MODEL=qwen3.7-text-embedding
DASHSCOPE_API_KEY=sk-your-key

# Reranker 配置
RERANKER_MODEL=BAAI/bge-reranker-base
ENABLE_RERANKER=true

# 检索参数
RETRIEVAL_TOP_K=8
RERANK_TOP_K=4
MIN_EVIDENCE_SCORE=0.25

配置类 (config.py)

使用 pydantic-settings 实现类型安全的配置管理:

from pydantic_settings import BaseSettings, SettingsConfigDict

class Settings(BaseSettings):
    app_name: str = '政务政策 RAG 系统'
    debug: bool = True
    api_prefix: str = '/api'

    database_url: str = "mysql+pymysql://root:password@127.0.0.1:3306/gov_rag"
    milvus_path: Path = Path('data/milvus')
    embedding_model: str = 'qwen3.7-text-embedding'
    dashscope_api_key: str = ''

    model_provider: Literal["ollama", "deepseek"] = "ollama"
    deepseek_api_key: str = ''
    deepseek_base_url: str = "https://api.siliconflow.cn/v1"
    deepseek_model: str = "deepseek-ai/DeepSeek-V4-Flash"

    # Reranker
    reranker_model: str = "BAAI/bge-reranker-base"
    enable_reranker: bool = True
    rerank_top_k: int = 4
    min_evidence_score: float = 0.25

    # 检索参数
    retrieval_top_k: int = 8

    model_config = SettingsConfigDict(
        env_file='.env',
        env_file_encoding='utf-8',
        extra='ignore'
    )

@lru_cache
def get_settings() -> Settings:
    return Settings()

关键设计get_settings() 使用 @lru_cache 装饰器,全局只创建一个 Settings 实例,避免重复读取 .env 文件。


三、数据模型设计

系统使用 SQLAlchemy ORM 定义了 5 张核心表:

核心表结构

表名说明关键字段
policy_documents政策文档id, filename, topic, effective_date, status
document_chunks文档分块(父子)id, document_id, parent_id, is_parent, content
conversations会话id, title
chat_messages聊天消息id, conversation_id, role, content, trace
citations引用记录id, message_id, chunk_id, label, quote, score

DocumentChunk:父子块自引用

class DocumentChunk(TimestampMixin, Base):
    __tablename__ = "document_chunks"

    id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id)
    document_id: Mapped[str] = mapped_column(
        ForeignKey("policy_documents.id", ondelete="CASCADE"), index=True
    )
    parent_id: Mapped[str | None] = mapped_column(
        ForeignKey("document_chunks.id", ondelete="CASCADE"), index=True  # 自引用外键
    )
    ordinal: Mapped[int] = mapped_column(Integer)       # 全局序号
    section: Mapped[str | None] = mapped_column(String(255))  # 章节标题
    page: Mapped[int | None] = mapped_column(Integer)         # 页码
    content: Mapped[str] = mapped_column(Text)                 # 块文本内容
    is_parent: Mapped[bool] = mapped_column(Boolean, default=False, index=True)
    token_count: Mapped[int] = mapped_column(Integer, default=0)
    metadata_json: Mapped[dict[str, Any]] = mapped_column(JSON, default=dict)

关联关系全景

PolicyDocument (1) ──→ (N) DocumentChunk
                                ↑
                    parent_id (自引用)
                                ↓
                         DocumentChunk (子块)

DocumentChunk (子块, is_parent=False) ──→ (1:1) Milvus 向量记录

核心设计parent_id 是自引用外键,指向同表中的另一条记录。设置了 ON DELETE CASCADE,父块被删除时子块自动级联删除。只有子块(is_parent=False)写入 Milvus,父块仅存 MySQL。


四、文档处理与父子块切分

4.1 文档上传与解析

文档上传接口支持 PDF、DOCX、TXT 三种格式:

ALLOWED_EXTENSIONS = {'.pdf', '.docx', '.txt'}

async def ingest_upload(session: Session, upload: UploadFile,
                        topic: str | None = None,
                        effective_date: date | None = None):
    # 1. 校验文件类型和大小
    suffix = Path(upload.filename).suffix.lower()
    if suffix not in ALLOWED_EXTENSIONS:
        raise ValueError("仅支持 PDF、DOCX、TXT 文件")

    content = await upload.read()
    settings = get_settings()
    if len(content) > settings.max_upload_mb * 1024 * 1024:
        raise ValueError("文件大小超出限制")

    # 2. 保存文件记录到 MySQL
    document = PolicyDocument(
        filename=upload.filename,
        stored_path="",
        content_type=upload.content_type,
        topic=topic,
        effective_date=effective_date,
        status=DocumentStatus.indexing,
    )
    session.add(document)
    session.flush()  # 获取数据库生成的 UUID

    # 3. 写入磁盘
    upload_path = settings.upload_path / f"{document.id}{suffix}"
    upload_path.write_bytes(content)
    document.stored_path = str(upload_path)

    # 4. 逐页解析 + 父子块切分 + 入库(见下文)
    # 5. 子块写入 Milvus 向量库
    # ...

多格式解析extract_pages 函数):

def extract_pages(path: Path) -> list[tuple[int | None, str]]:
    suffix = path.suffix.lower()
    if suffix == ".pdf":
        pages = PdfReader(path).pages
        return [(i, page.extract_text()) for i, page in enumerate(pages, 1)]
    elif suffix == ".docx":
        doc = DocxDocument(str(path))
        content = "\n".join([p.text for p in doc.paragraphs])
        return [(None, content)]
    else:  # .txt
        content = path.read_text(encoding="utf-8")
        return [(None, content)]

4.2 父子块切分策略

核心思想:子块负责"精准命中",父块负责"提供上下文"。

原始文档
    ↓
按章节标题切分(正则识别"第X章"、"一、"、"1."等标题)
    ↓
每个章节 → _windows(正文, size=1200, overlap=100) → 父块
    ↓
每个父块 → _windows(父块, size=420, overlap=80) → 子块
切分阶段数据类:ChunkDraft
@dataclass
class ChunkDraft:
    content: str              # 块文本内容
    section: str | None       # 所属章节标题
    page: int | None          # 来源页码
    is_parent: bool            # True=父块, False=子块
    parent_index: int | None = None  # 子块指向父块在列表中的下标
    metadata: dict = field(default_factory=dict)

关键设计parent_index 不是数据库外键,而是切分阶段的临时内存下标,指向 drafts 列表中父块的位置。这避免了切分阶段就依赖数据库 ID。

章节识别正则
HEADING_PATTERN = re.compile(
    r"(?m)^([第][一二三四五六七八九十百]+[章节条]\s*.*"
    r"|[一二三四五六七八九十]+[、.]\s*.*"
    r"|\d+[.、]\s*.*)$"
)

匹配三类标题:

  • 第X章/节/条:如"第一章 总则"
  • 中文数字编号:如"一、概述"
  • 阿拉伯数字编号:如"1. 总则"
父子块生成:split_policy_text
def split_policy_text(text, page=None, parent_size=1200,
                      child_size=420, child_overlap=80) -> list[ChunkDraft]:
    # 1. 正则识别所有章节标题
    matches = list(HEADING_PATTERN.finditer(text))
    sections = []
    if not matches:
        sections = [(None, text)]  # 无标题则整体作为一个章节
    else:
        if matches[0].start() > 0:
            sections.append((None, text[:matches[0].start()]))  # 前置文本
        for i, match in enumerate(matches):
            end = matches[i + 1].start() if i + 1 < len(matches) else len(text)
            sections.append((match.group(0).strip(), text[match.start():end]))

    # 2. 双层切分:先切父块,再切子块
    drafts: list[ChunkDraft] = []
    for sec_title, sec_body in sections:
        for draft in _windows(sec_body, parent_size, 100):
            parent_index = len(drafts)  # 记录父块即将插入的位置
            # 生成父块
            drafts.append(ChunkDraft(
                content=draft, section=sec_title, page=page, is_parent=True
            ))
            # 对父块再切子块
            for child in _windows(draft, child_size, child_overlap):
                drafts.append(ChunkDraft(
                    content=child, section=sec_title, page=page,
                    is_parent=False, parent_index=parent_index
                ))
    return drafts

核心行parent_index = len(drafts) —— 在父块被追加之前记录当前列表长度,这个值恰好是父块即将被插入的位置。后续所有子块都通过这个下标找到自己的父块。

4.3 窗口切分函数

_windows 函数支持标点感知截断,避免把句子从中间切断:

def _windows(text: str, size: int, overlap: int) -> list[str]:
    text = text.strip()
    if not text:
        return []
    chunks = []
    start = 0
    while start < len(text):
        end = min(len(text), start + size)
        if end < len(text):
            # 在窗口范围内从右往左找句末标点
            punctuation = max(text.rfind(mark, start, end) for mark in '。;!?\n')
            if punctuation > start + size // 2:  # 标点超过窗口中点才截断
                end = punctuation + 1
        chunks.append(text[start:end].strip())
        if end >= len(text):
            break
        start = max(start + 1, end - overlap)  # 带重叠的滑动窗口
    return [chunk for chunk in chunks if chunk]

4.4 切分参数选择

参数说明
parent_size1200父块窗口,保证完整段落语义
parent_overlap100父块重叠,避免跨块丢信息
child_size420子块窗口,聚焦单个要点,适合精准检索
child_overlap80子块重叠,避免截断关键句

4.5 入库时的关联建立

切分完成后,ChunkDraft 被转换为 DocumentChunk ORM 对象写入 MySQL。关联从内存下标变为数据库 UUID 的转换在 ingest_upload 中完成:

chunk_rows: list[DocumentChunk] = []
ordinal = 0
for page, text in extract_pages(upload_path):
    drafts = split_policy_text(text, page=page)
    local_rows: dict[int, DocumentChunk] = {}  # 桥接字典

    for index, draft in enumerate(drafts):
        chunk = DocumentChunk(
            document_id=document.id,
            parent_id=None,
            ordinal=ordinal,
            section=draft.section,
            page=draft.page,
            content=draft.content,
            is_parent=draft.is_parent,
            token_count=len(draft.content),
            metadata_json=draft.metadata,
        )
        session.add(chunk)
        session.flush()  # 触发 UUID 生成
        local_rows[index] = chunk

        # 关键:将 parent_index(内存下标)转换为 parent_id(数据库 UUID)
        if draft.parent_index is not None:
            chunk.parent_id = local_rows[draft.parent_index].id
        ordinal += 1

local_rows 字典 是整个关联设计的桥梁:key 是 ChunkDraftdrafts 列表中的下标,value 是已 flush 获得 UUID 的 DocumentChunk 对象。通过 local_rows[parent_index].id 拿到父块的 UUID。


五、向量存储服务

5.1 Embedding 模型初始化

使用阿里云 DashScope 的 Embedding 模型,单例模式避免重复初始化:

from langchain_community.embeddings import DashScopeEmbeddings

@lru_cache
def get_embeddings():
    settings = get_settings()
    return DashScopeEmbeddings(
        model=settings.embedding_model,          # qwen3.7-text-embedding
        dashscope_api_key=settings.dashscope_api_key
    )

5.2 Milvus 向量库连接

项目使用 milvus-lite(本地文件模式),无需部署 Docker:

from langchain_milvus import Milvus

@lru_cache
def get_vector_store():
    settings = get_settings()
    settings.milvus_path.mkdir(parents=True, exist_ok=True)

    db_url = r"D:\PythonProject\langchain_project\backend\data\milvus\milvus_demo.db"
    vector_store = Milvus(
        embedding_function=get_embeddings(),
        connection_args={"uri": db_url},      # milvus-lite 文件路径
        collection_name="policy_chunks",
        auto_id=False,                         # 使用自定义 ID(chunk_id)
        consistency_level="Strong",
        index_params={
            "metric_type": "COSINE",           # 余弦相似度
            "index_type": "HNSW",              # 图索引,查询快
            "params": {"M": 8, "efConstruction": 64}
        }
    )
    return vector_store
配置项说明
metric_typeCOSINE余弦相似度,适合语义搜索
index_typeHNSW层次可导航小世界图,查询效率高
auto_idFalse使用 chunk_id 作为主键,方便与 MySQL 关联
consistency_levelStrong强一致性,保证写入立即可查

5.3 DashScope 分批处理

DashScope Embedding API 单次最多 20 条文本,需要分批处理:

def batch_split(lst: list, batch_size: int = 20) -> List[list]:
    """切分列表为多个小批次,dashscope embedding 最大 20"""
    for i in range(0, len(lst), batch_size):
        yield lst[i:i + batch_size]

5.4 写入与检索

写入子块到 Milvus

class VectorStoreService:
    async def add_chunks(self, document: PolicyDocument,
                         chunks: list[DocumentChunk]) -> None:
        if not chunks:
            return
        embeddings = get_embeddings()
        content_list = [chunk.content for chunk in chunks]

        # DashScope 单批最大 20 条,分批 embedding
        all_vectors = []
        for batch in batch_split(content_list, 20):
            batch_vectors = await asyncio.to_thread(
                embeddings.embed_documents, batch
            )
            all_vectors.extend(batch_vectors)

        # 组装 LangChain Document 对象
        doc_list = []
        for chunk in chunks:
            metadata = {
                "document_id": document.id,
                "filename": document.filename,
                "topic": document.topic or "",
                "effective_date": document.effective_date.isoformat() if document.effective_date else "",
                "section": chunk.section or "",
                "page": chunk.page or 0,
                "parent_id": chunk.parent_id or "",
                "chunk_id": chunk.id,
            }
            doc = Document(
                page_content=chunk.content,
                metadata=metadata,
                id=chunk.id
            )
            doc_list.append(doc)

        # 分批写入 Milvus
        vector_store = get_vector_store()
        for batch in batch_split(doc_list, 20):
            vector_store.add_documents(batch)

向量检索(支持按 document_id 过滤):

    async def search(self, query: str, top_k: int,
                     document_ids: list[str] | None = None) -> list[dict]:
        vector_store = get_vector_store()

        if document_ids:
            # Milvus expr 表达式过滤
            ids_list = [f'"{x}"' for x in document_ids]
            ids = ",".join(ids_list)
            expr = f'document_id in [{ids}]'
        else:
            expr = None

        results = vector_store.similarity_search_with_score(query, top_k, expr=expr)

        return [
            {
                "chunk_id": result[0].metadata["chunk_id"],
                "content": result[0].page_content,
                "metadata": result[0].metadata,
                "score": result[1],
                "source": result[0].metadata["filename"],
            }
            for result in results
        ]

六、混合检索引擎

6.1 检索流程总览

这是整个项目最核心的模块,完整流程如下:

用户提问
    ↓
Query Rewrite(LLM 改写查询)
    ↓
    ┌─────────────────┬──────────────────┐
    ↓                 ↓                  │
向量检索          BM25 关键词检索         │
(Milvus top-16)   (MySQL + jieba)        │
    ↓                 ↓                  │
    └────────┬────────┘                  │
             ↓                           │
      RRF 倒数排名融合                    │
             ↓                           │
      Metadata 过滤(topic/date)        │
             ↓                           │
      取 Top-K 候选证据                   │
             ↓                           │
      父块上下文扩展(parent_id 回溯)     │
             ↓                           │
      Reranker 重排(Cross-Encoder)      │
             ↓                           │
      证据打分 + 低分过滤                  │
             ↓                           │
      返回最终证据列表                     │

6.2 查询改写

将口语化问题改写为更适合检索的查询:

async def rewrite_query(question: str) -> str:
    """把口语化问题改写成适合检索的查询;失败时退回原问题。"""
    prompt = f"""
    将下面的政务咨询改写为一个适合知识库检索的简洁中文查询。
    不得添加问题中没有的地区、政策或条件,只输出改写结果。
    问题:{question}
    """
    try:
        llm = create_chat_model(settings=get_settings(), temperature=0)
        content, _ = await invoke_text(llm, prompt)
        return content
    except Exception:
        return question  # 失败时退回原问题

设计要点:改写失败时回退到原问题,不影响主流程。temperature=0 保证改写结果稳定。

6.3 向量检索

vector_results = await VectorStoreService().search(
    new_question,
    top_k * 2,  # 多召回一倍,给 RRF 融合留余量
    document_ids=filters.document_ids or None
)

6.4 BM25 关键词检索

BM25 是经典的关键词检索算法,项目用 rank_bm25 库实现,配合 jieba 中文分词:

中文分词
import jieba
import re

STOP_WORDS = {"的", "了", "是", "在", "和", "有", "就", "都", "而", "及"}
valid_pat = re.compile(r"^([\u4e00-\u9fff]+|[A-Za-z0-9]+)$")

def tokenize_chinese(text: str) -> list[str]:
    """中文文本分词:jieba 搜索模式 + 停用词过滤 + 正则校验"""
    text_low = text.lower()
    words = jieba.lcut_for_search(text_low)  # 搜索模式,长词再细切
    word_list = []
    for w in words:
        if valid_pat.fullmatch(w) and w not in STOP_WORDS and len(w) > 1:
            word_list.append(w)
    return word_list
BM25 检索
from rank_bm25 import BM25Okapi

async def keyword_search(session: Session, query: str, top_k: int,
                         filters: RetrievalFilters) -> list[dict]:
    # 1. 从 MySQL 查询所有子块(支持过滤条件)
    statement = (
        select(DocumentChunk, PolicyDocument)
        .join(PolicyDocument, DocumentChunk.document_id == PolicyDocument.id)
        .where(DocumentChunk.is_parent.is_(False))
    )
    if filters.document_ids:
        statement = statement.where(DocumentChunk.document_id.in_(filters.document_ids))
    if filters.topic:
        statement = statement.where(PolicyDocument.topic == filters.topic)

    rows = list(session.execute(statement).all())
    if not rows:
        return []

    # 2. 构建语料库(每个子块分词后的 token 列表)
    corpus = [tokenize_chinese(chunk[0].content) for chunk in rows]

    # 3. BM25 打分
    bm25_model = BM25Okapi(corpus)
    scores = bm25_model.get_scores(tokenize_chinese(query))

    # 4. 取 Top-K
    ranks = sorted(enumerate(scores), key=lambda x: x[1], reverse=True)[:top_k]

    result_list = []
    for index, score in ranks:
        if score > 0:
            document_chunk = rows[index][0]
            policy_document = rows[index][1]
            result_list.append({
                'chunk_id': document_chunk.id,
                'content': document_chunk.content,
                'score': float(score),
                'source': 'bm25',
                'metadata': {
                    "document_id": policy_document.id,
                    "filename": policy_document.filename,
                    "section": document_chunk.section or "",
                    "page": document_chunk.page or 0,
                    "parent_id": document_chunk.parent_id or "",
                }
            })
    return result_list

6.5 RRF 倒数排名融合

向量检索和 BM25 检索各有优劣,RRF(Reciprocal Rank Fusion)将两个排序结果融合:

def reciprocal_rank_fusion(result_lists: list[list[dict]], k: int = 60) -> list[dict]:
    """
    RRF 公式: rrf_score = sum(1 / (k + rank)) for each list where item appears
    用 content 做唯一标识匹配同一 chunk。
    """
    fused: dict[str, dict[str, Any]] = {}

    for result_list in result_lists:
        for rank, item in enumerate(result_list):
            key = item.get("content", "")
            if not key:
                continue
            rrf_score = 1.0 / (k + rank + 1)
            if key not in fused:
                fused[key] = {**item}
                fused[key]["rrf_score"] = 0.0
                fused[key]["sources"] = []
            fused[key]["rrf_score"] += rrf_score
            if item.get("source") not in fused[key]["sources"]:
                fused[key]["sources"].append(item["source"])

    # 按 RRF 分数降序排列
    merged = sorted(fused.values(), key=lambda x: x["rrf_score"], reverse=True)

    # 归一化到 0-1
    max_score = max((item["rrf_score"] for item in merged), default=1.0)
    for item in merged:
        item["score"] = item["rrf_score"] / max_score if max_score > 0 else 0.0

    return merged

RRF 的优势:不需要校准两个检索系统的分数尺度,只看排名位置。k=60 是经验值,平滑排名差异。

6.6 父块上下文扩展

检索命中的是子块(420字),但 LLM 需要更完整的上下文。通过 parent_id 回溯到父块(1200字):

async def expand_parent_context(session: Session, evidence: list[dict]) -> list:
    # 收集所有 parent_id
    parent_ids = [item['metadata']['parent_id'] for item in evidence]
    if not parent_ids:
        return evidence

    # 从 MySQL 查询父块信息
    chunks_list = session.query(DocumentChunk).filter(
        DocumentChunk.id.in_(parent_ids)
    ).all()
    parents = {chunk.id: chunk for chunk in chunks_list}

    # 用父块内容替换子块内容(保留子块元数据)
    seen_parents = set()
    expanded = []
    for item in evidence:
        parent_id = item["metadata"].get("parent_id")
        parent = parents.get(parent_id)
        if parent and parent_id not in seen_parents:
            item = {
                **item,
                "content": parent.content,  # 用父块的完整内容
                "metadata": {
                    **item["metadata"],
                    "section": parent.section,
                    "page": parent.page,
                    "context_chunk_id": parent.id,
                },
            }
            seen_parents.add(parent_id)
            expanded.append(item)
    return expanded

设计要点seen_parents 集合避免同一个父块被多次展开。多个子块命中同一父块时,只保留一条。

6.7 Reranker 重排

向量检索和 BM25 都是双塔模型(分别编码 query 和 doc),精度有限。Reranker 使用 Cross-Encoder 对 query 和 doc 做联合编码,精度更高:

from sentence_transformers import CrossEncoder

@lru_cache
def get_reranker():
    model_name = get_settings().reranker_model  # BAAI/bge-reranker-base
    return CrossEncoder(model_name, local_files_only=True)

async def rerank(query: str, candidates: list[dict], top_k: int) -> list[dict]:
    """用 Cross-Encoder 对融合后的候选证据重新排序。"""
    if not candidates:
        return []
    settings = get_settings()
    if not settings.enable_reranker:
        return candidates[:top_k]

    try:
        # Cross-Encoder 对 (query, doc) 对打分
        scores = await asyncio.to_thread(
            get_reranker().predict,
            [(query, item["content"]) for item in candidates],
        )
        for item, score in zip(candidates, scores):
            raw_score = float(score)
            item["rerank_raw_score"] = raw_score
            item["rerank_score"] = normalize_logit(raw_score)  # sigmoid 归一化

        return sorted(candidates, key=lambda x: x["rerank_score"], reverse=True)[:top_k]
    except Exception as e:
        # 降级:reranker 失败时用 RRF 分数
        for item in candidates:
            item["rerank_score"] = min(1.0, item.get("rrf_score", 0) * 30)
        return candidates[:top_k]

Sigmoid 归一化normalize_logit):

import math

def normalize_logit(value: float) -> float:
    return 1 / (1 + math.exp(-max(-30, min(30, value))))

6.8 证据过滤与打分

Reranker 完成后,统一打分并过滤低分证据:

# 统一打分:rerank_score > score > rrf_score*30
for item in evidence:
    normalized = item.get("rerank_score")
    if normalized is None:
        normalized = item.get("score", item.get("rrf_score", 0) * 30)
    item["evidence_score"] = max(0.0, min(1.0, float(normalized)))

# full 模式下过滤低分证据
if mode == "full":
    evidence = filter_evidence(evidence, min_score)

七、RAG 生成服务

7.1 回答生成流程

async def answer_question(session, question, filters=None, mode="full",
                          top_k=None, rerank_top_k=None, min_score=None):
    # 兜底默认值
    top_k = top_k or 8
    rerank_top_k = rerank_top_k or 4
    min_score = min_score or 0.25

    # Step 1: 检索证据
    retrieve_result = await retrieve(session, question, filters, top_k, 4, min_score, mode)

    # 无证据 → 拒绝回答
    if retrieve_result.refused:
        return RAGAnswer(REFUSAL_TEXT, True, [], retrieve_result.trace, ...)

    # Step 2: 组装上下文(带编号)
    evidences = retrieve_result.evidence
    contexts = []
    for index, evidence in enumerate(evidences, start=1):
        content = evidence['content']
        page = evidence['metadata']['page']
        section = evidence['metadata']['section']
        filename = evidence['metadata']['filename']
        contexts.append(f"[{index}] 文档:{filename},章节:{section},页码:{page}\n{content}")
    context = '\n\n'.join(contexts)

    # Step 3: 构建 Prompt
    prompt = f"""你是政务政策知识库助手。必须遵守:
    1. 只能依据"政策证据"回答,不得使用外部知识补充政策内容。
    2. 每个包含政策事实、数字、日期、对象或条件的句子末尾标注引用编号,如[1]。
    3. 证据没有说明的内容,明确回答"资料未说明"。
    4. 不得编造政策名称、办理条件、部门、时间或引用。

    政策证据:
    {context}

    用户问题:{question}
    """

    # Step 4: LLM 生成回答
    llm = create_chat_model(temperature=0, settings=get_settings())
    answer, usage = await invoke_text(llm, prompt)

    # Step 5: 引用校验
    answer, cited_indexes = validate_citations(answer, evidences)

    # Step 6: 收集引用信息
    citations = []
    for index in cited_indexes:
        item = evidences[index - 1]
        citations.append({
            "chunk_id": item["chunk_id"],
            "filename": item["metadata"].get("filename", ""),
            "section": item["metadata"].get("section"),
            "page": item["metadata"].get("page"),
            "quote": item["content"][:500],
            "score": item["evidence_score"],
            "label": str(index),
        })

    return RAGAnswer(answer, refused, citations, trace, latency_ms, usage)

7.2 引用校验机制

LLM 可能在回答中生成不存在的引用编号(如证据只有 3 条但模型写了 [5])。validate_citations 负责清理:

CITATION_PATTERN = re.compile(r"\[(\d+)\]")

def validate_citations(answer: str, evidence: list[dict]) -> tuple[str, list[int]]:
    """
    校验回答中的 [数字] 引用标记,过滤无效引用。
    返回: (清理后的回答, 合法引用编号列表)
    """
    # 提取回答中所有引用编号
    cited = set()
    for value in CITATION_PATTERN.findall(answer):
        cited.add(int(value))

    # 合法编号范围: 1 ~ len(evidence)
    valid = set(range(1, len(evidence) + 1))

    # 移除不合法的编号
    invalid = cited - valid
    if invalid:
        for index in invalid:
            answer = answer.replace(f'[{index}]', '')

    # 返回合法的引用编号(有序)
    return answer, sorted(cited & valid)

7.3 会话保存

回答生成后,将问题、回答、引用保存到 MySQL:

async def save_chat(session, question, result, conversation_id):
    # 获取或创建会话
    conversation = session.get(Conversation, conversation_id) if conversation_id else None
    if not conversation:
        conversation = Conversation(title=question[:50])
        session.add(conversation)
        session.flush()

    # 保存用户消息
    session.add(ChatMessage(
        conversation_id=conversation.id,
        role='user',
        content=question
    ))

    # 保存 AI 回答
    assistant_message = ChatMessage(
        conversation_id=conversation.id,
        role='assistant',
        content=result.answer,
        refused=result.refused,
        latency_ms=result.latency_ms,
        token_usage=result.token_usage,
        trace=result.trace
    )
    session.add(assistant_message)
    session.flush()

    # 保存引用记录
    for item in result.citations:
        session.add(Citation(
            message_id=assistant_message.id,
            chunk_id=item['chunk_id'],
            label=item['label'],
            quote=item["quote"],
            score=item["score"]
        ))

    session.commit()
    return conversation, assistant_message

八、评估体系

8.1 评估指标

系统实现了 6 个核心评估指标:

指标函数说明
Recall@Krecall_at_k()检索召回率:期望命中的 chunk 有多少被检索到
MRRreciprocal_rank()平均倒数排名:第一个命中的 chunk 排在第几位
引用准确率citation_precision()回答中的引用是否都来自检索结果
拒绝准确率refusal_accuracy应该拒绝的问题是否正确拒绝
事实准确率factual_accuracy()关键数字/日期是否与参考答案一致
幻觉率citation_coverage()含事实的句子中,有多少带引用标记
# Recall@K
def recall_at_k(expected: list[str], retrieved: list[str], k: int) -> float:
    if not expected:
        return 1.0
    return len(set(expected) & set(retrieved[:k])) / len(set(expected))

# MRR
def reciprocal_rank(expected: list[str], retrieved: list[str]) -> float:
    expected_set = set(expected)
    for rank, chunk_id in enumerate(retrieved, start=1):
        if chunk_id in expected_set:
            return 1.0 / rank
    return 0.0

# 关键事实提取(数字、日期、金额)
def extract_facts(text: str) -> set[str]:
    return set(re.findall(
        r"\d+(?:\.\d+)?(?:%|万元|元|天|日|个月|年)?|"
        r"\d{4}年\d{1,2}月\d{1,2}日", text
    ))

# 事实准确率
def factual_accuracy(reference: str, answer: str) -> float:
    s1 = extract_facts(reference)
    if not s1:
        return 1.0
    s2 = extract_facts(answer)
    return len(s1 & s2) / len(s1)

# 引用覆盖率(幻觉率的反向指标)
def citation_coverage(answer: str) -> float:
    sentences = re.split(r"[。!?\n]+", answer)
    factual = [s for s in sentences if re.search(r"\d|条件|材料|期限|标准|应当|不得|可以|补贴|申请", s)]
    if not factual:
        return 1.0
    supported = [s for s in factual if re.search(r"\[\d+\]", s)]
    return len(supported) / len(factual)

8.2 LLM Judge 评估器

除了规则指标,系统还支持用 LLM 做主观评估:

async def judge_answer(question, reference, answer, context) -> dict:
    prompt = f"""你是独立的 RAG 评估器。只输出 JSON:
    {{"faithfulness":0到1,"relevance":0到1,"factual_accuracy":0到1,
    "hallucination_rate":0到1,"reason":"简短理由"}}
    忠实度仅判断答案是否由上下文支持。
    问题:{question}
    参考答案:{reference}
    上下文:{context}
    待评答案:{answer}
    """
    llm = create_chat_model(get_settings(), temperature=0)
    output, _ = await invoke_text(llm, prompt)
    match = re.search(r"\{.*\}", output, re.S)
    return json.loads(match.group(0)) if match else {}

8.3 评估执行流程

async def execute_evaluation_run(run_id):
    with SessionLocal() as session:
        run = session.get(EvaluationRun, run_id)
        run.status = RunStatus.running
        session.commit()

        cases = session.query(EvaluationCase).filter(
            EvaluationCase.dataset_id == run.dataset_id
        ).all()
        config = run.rag_config

        for case in cases:
            # 用测试用例的问题调用 RAG
            rag_answer = await answer_question(
                session, case.question, RetrievalFilters(),
                mode=config["mode"], top_k=config["top_k"],
                rerank_top_k=config["rerank_top_k"],
                min_score=config["min_evidence_score"],
            )

            # 计算指标
            metrics = {
                'recall_at_k': recall_at_k(expected, retrieved, top_k),
                'mrr': reciprocal_rank(expected, retrieved),
                'citation_precision': citation_precision(cited, retrieved),
                'refusal_accuracy': float(rag_answer.refused == case.should_refuse),
                'factual_accuracy': factual_accuracy(reference, answer),
                'hallucination_rate': 1 - citation_coverage(answer) if not refused else 0.0,
            }

            # 可选:LLM Judge
            if config.get('enable_judge'):
                judge_score = await judge_answer(question, reference, answer, context)
                metrics.update(judge_score)

            # 保存评估结果
            session.add(EvaluationResult(..., metrics=metrics, ...))

        # 汇总指标
        run.metrics = aggregate_metrics(stored_results)
        run.status = RunStatus.completed
        session.commit()

九、API 层设计

FastAPI 应用工厂

def create_app():
    app = FastAPI(
        title=settings.app_name,
        description="FastAPI + LangChain 政策知识库问答系统",
        version="1.0.0",
    )
    for router in (documents.router, chat.router,
                   configuration.router, evaluation.router):
        app.include_router(router, prefix=settings.api_prefix)
    return app

API 接口列表

方法路径说明
POST/api/documents上传文档(自动切分入库)
GET/api/documents查询文档列表
DELETE/api/documents/{id}删除文档(级联删除分块+向量)
POST/api/documents/rebuild重建所有向量索引
POST/api/chat/stream流式问答(SSE 推送)
GET/api/chat/conversations查询会话列表
GET/api/chat/conversations/{id}查询会话详情(含消息和引用)
PUT/api/config/model更新模型配置
GET/api/config/model查询模型配置
POST/api/evaluation/datasets创建测试集
POST/api/evaluation/runs启动评估任务

SSE 流式回答

聊天接口使用 Server-Sent Events 推送检索进度和逐 token 生成:

@router.post("/stream")
async def chat_stream(payload: ChatRequest):
    return StreamingResponse(generate(payload), media_type="text/event-stream")

async def generate(chatRequest: ChatRequest):
    yield event("status", {"stage": "retrieving", "message": "正在检索政策依据"})

    result = await answer_question(session, chatRequest.question, chatRequest.filters)

    yield event("retrieval", {
        "rewritten_query": result.trace.get("rewritten_query"),
        "evidence_count": len(result.trace.get("final_evidence", [])),
    })

    # 逐段推送回答文本(每 24 字一段)
    for index in range(0, len(result.answer), 24):
        yield event("token", {"content": result.answer[index: index + 24]})

    # 保存到数据库
    conversation, message = await save_chat(session, chatRequest.question, result, chatRequest.conversation_id)

    yield event("done", {
        "conversation_id": conversation.id,
        "citations": result.citations,
        "latency_ms": result.latency_ms,
    })

十、总结

本文完整拆解了一个生产级政务政策 RAG 系统的核心实现,关键设计要点如下:

架构全景

用户上传文档
    ↓
PDF/DOCX/TXT 解析 → 父子块切分(1200/420字)
    ↓
子块写入 MySQL + Milvus(DashScope Embedding 分批 20 条)
    ↓
用户提问
    ↓
LLM 查询改写 → 向量检索(top-16) + BM25 检索(top-16)
    ↓
RRF 融合 → Top-8 候选 → 父块扩展 → Reranker 重排 → Top-4 证据
    ↓
组装 Prompt → LLM 生成回答(带 [1][2] 引用标注)
    ↓
引用校验 → 过滤无效编号 → 返回答案 + 引用来源
    ↓
SSE 流式推送给前端 + 保存到 MySQL

核心设计亮点

设计点方案价值
父子块切分1200字父块 + 420字子块小块精准检索,大块完整上下文
混合检索向量 + BM25 + RRF 融合语义匹配 + 关键词匹配双保险
RerankerCross-Encoder 精排双塔召回后精排,大幅提升准确率
查询改写LLM 改写口语化问题提升检索召回率
引用校验正则提取 + 合法性验证防止 LLM 编造引用
父块扩展parent_id 回溯子块命中后获取完整上下文
评估体系6 个规则指标 + LLM Judge量化 RAG 质量可持续优化
拒绝机制无证据时拒绝回答降低幻觉风险

关键参数速查

参数默认值说明
parent_size1200父块窗口大小
child_size420子块窗口大小
RETRIEVAL_TOP_K8候选证据数量
RERANK_TOP_K4重排后保留的证据数量
MIN_EVIDENCE_SCORE0.25最低证据得分阈值
batch_size20DashScope Embedding 单批最大量

安装依赖

pip install fastapi uvicorn sqlalchemy pymysql pydantic-settings
pip install langchain langchain-core langchain-openai langchain-milvus
pip install dashscope pymilvus sentence-transformers rank-bm25
pip install pypdf python-docx jieba httpx

启动服务

# 初始化数据库
python scripts/init_db.py

# 启动 FastAPI
uvicorn main:app --reload --port 8000

# 访问 API 文档
# http://localhost:8000/docs

如果这篇文章对你有帮助,欢迎点赞、收藏、关注!有问题可以在评论区留言讨论。

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值