Skip to content

第 19 章 · 实战一:RAG 知识库 Agent

本章目标:

  • 从零构建一个完整的 RAG Agent 服务
  • 使用 Vercel AI SDK 的嵌入向量检索文本
  • 用 LangGraph StateGraph 编排检索 + 生成流程
  • 通过 Checkpointer 实现多轮对话持久化
  • 用 FastAPI 暴露 HTTP 接口
  • 接入 Langfuse 实现全链路追踪

本章将逐步构建一个生产级 RAG Agent。我们从最基础的语义检索开始,逐步加入 LangGraph 状态图、对话持久化、API 接口和可观测性。每阶段都是完整可运行的代码。

📌 前置知识:本章需要理解 Vercel AI SDK 嵌入向量LangGraph 基础概念

19.1 阶段一:基础 RAG — 嵌入检索

RAG(Retrieval-Augmented Generation)的核心是「先检索,再生成」。我们先用一个简单的向量数据库存储文档,然后用语义相似度检索相关片段。

19.1.1 项目初始化

首先创建项目目录并安装依赖:

bash
mkdir rag-agent && cd rag-agent
pip install "ai[all]" langgraph langchain-core langchain-text-splitters langfuse fastapi uvicorn python-dotenv

19.1.2 Provider 配置

参考 Vercel AI SDK ch03,我们使用统一的 Provider 管理:

python
# provider.py
import os
from ai import createGateway, createOpenAICompatible

def get_embedding_model():
    """根据环境变量选择嵌入模型"""
    if os.getenv("AI_GATEWAY_API_KEY"):
        gateway = createGateway(api_key=os.getenv("AI_GATEWAY_API_KEY"))
        return gateway.embed("openai/text-embedding-3-small")
    
    # 自定义 OpenAI 兼容 Provider
    if os.getenv("EMBEDDING_BASE_URL"):
        provider = createOpenAICompatible(
            name="embedding-provider",
            baseURL=os.getenv("EMBEDDING_BASE_URL"),
            apiKey=os.getenv("EMBEDDING_API_KEY"),
        )
        return provider.embedding(os.getenv("EMBEDDING_MODEL", "text-embedding-3-small"))
    
    raise ValueError("请设置 AI_GATEWAY_API_KEY 或 EMBEDDING_BASE_URL")

19.1.3 文档分块与向量化

我们将文档切成小块,并为每个块计算嵌入向量:

python
# rag_engine.py
import json
import sqlite3
from pathlib import Path
from typing import List, Tuple
import numpy as np
from langchain_text_splitters import RecursiveCharacterTextSplitter
from provider import get_embedding_model

# 使用 SQLite 作为轻量级向量存储
DB_PATH = Path("rag.db")

def init_db():
    """初始化数据库表"""
    conn = sqlite3.connect(DB_PATH)
    conn.execute("""
        CREATE TABLE IF NOT EXISTS documents (
            id INTEGER PRIMARY KEY AUTOINCREMENT,
            text TEXT NOT NULL,
            embedding BLOB,
            metadata TEXT
        )
    """)
    conn.execute("""
        CREATE INDEX IF NOT EXISTS idx_embedding 
        ON documents (embedding)
    """)
    conn.commit()
    return conn

def chunk_document(text: str, chunk_size: int = 500, chunk_overlap: int = 50) -> List[str]:
    """将文档分割成小块"""
    splitter = RecursiveCharacterTextSplitter(
        chunk_size=chunk_size,
        chunk_overlap=chunk_overlap,
        separators=["\n\n", "\n", ". ", " ", ""]
    )
    return splitter.split_text(text)

def compute_embeddings(chunks: List[str]) -> List[List[float]]:
    """批量计算嵌入向量"""
    model = get_embedding_model()
    embeddings = []
    for chunk in chunks:
        emb = model.embed(chunk)
        embeddings.append(emb.embedding)
    return embeddings

def cosine_similarity(a: List[float], b: List[float]) -> float:
    """计算余弦相似度"""
    import numpy as np
    a, b = np.array(a), np.array(b)
    return float(np.dot(a, b) / (np.linalg.norm(a) * np.linalg.norm(b)))

class RagEngine:
    def __init__(self):
        self.conn = init_db()
    
    def add_document(self, text: str, metadata: dict = None):
        """添加文档到知识库"""
        chunks = chunk_document(text)
        embeddings = compute_embeddings(chunks)
        
        for i, (chunk, emb) in enumerate(zip(chunks, embeddings)):
            self.conn.execute(
                "INSERT INTO documents (text, embedding, metadata) VALUES (?, ?, ?)",
                (chunk, json.dumps(emb).encode(), json.dumps(metadata or {}))
            )
        self.conn.commit()
    
    def retrieve(self, query: str, top_k: int = 3) -> List[Tuple[str, float]]:
        """检索最相关的文档片段"""
        model = get_embedding_model()
        query_emb = model.embed(query).embedding
        
        cursor = self.conn.execute("SELECT id, text, embedding FROM documents")
        docs = []
        for doc_id, text, emb_bytes in cursor:
            emb = json.loads(emb_bytes)
            sim = cosine_similarity(query_emb, emb)
            docs.append((text, sim))
        
        # 按相似度排序,返回 top_k
        docs.sort(key=lambda x: x[1], reverse=True)
        return docs[:top_k]
    
    def close(self):
        self.conn.close()

19.1.4 基础检索示例

python
# demo_basic_rag.py
from rag_engine import RagEngine

# 初始化引擎
engine = RagEngine()

# 添加示例文档
documents = [
    ("Python 是一种高级编程语言,由 Guido van Rossum 于 1991 年创建。它以简洁易读的语法著称,广泛应用于数据科学、机器学习和 Web 开发。", {"source": "百科"}),
    ("LangGraph 是 LangChain 团队开发的图结构 Agent 框架,支持状态管理、持久化和人机协作。它是构建复杂 Agent 工作流的首选工具。", {"source": "技术文档"}),
    ("RAG(检索增强生成)结合检索系统和 LLM,能够减少幻觉并提供准确的知识回答。核心步骤包括:文档分块、向量化、检索和生成。", {"source": "论文"}),
]

for text, meta in documents:
    engine.add_document(text, meta)

# 测试检索
query = "什么是 LangGraph?"
results = engine.retrieve(query)

print(f"查询: {query}")
print("\n检索结果:")
for i, (text, score) in enumerate(results, 1):
    print(f"{i}. [{score:.3f}] {text[:80]}...")

engine.close()

运行后输出类似:

查询: 什么是 LangGraph?
检索结果:
1. [0.872] LangGraph 是 LangChain 团队开发的图结构 Agent 框架...
2. [0.651] RAG(检索增强生成)结合检索系统和 LLM...
3. [0.523] Python 是一种高级编程语言...

19.2 阶段二:LangGraph StateGraph 编排

基础 RAG 只是线性流程。实际应用中,我们需要根据检索结果动态决定生成策略。LangGraph 的 StateGraph 让我们可以构建有向图,精确控制 Agent 的执行路径。

19.2.1 定义状态与节点

python
# rag_agent.py
import os
from typing import TypedDict, Annotated, List
import operator
from langgraph.graph import StateGraph, START, END
from langgraph.checkpoint.memory import MemorySaver
from langchain_core.messages import HumanMessage, AIMessage, SystemMessage
from provider import get_language_model, get_embedding_model
from rag_engine import RagEngine

# 定义 Agent 状态
class AgentState(TypedDict):
    messages: Annotated[List, operator.add]  # 消息历史(自动追加)
    context: str  # 检索到的上下文
    needs_retry: bool  # 是否需要重试检索

# 创建状态图
workflow = StateGraph(AgentState)

# 检索节点
def retrieve_node(state: AgentState) -> dict:
    """检索相关文档"""
    last_message = state["messages"][-1]
    query = last_message.content
    
    engine = RagEngine()
    results = engine.retrieve(query, top_k=3)
    engine.close()
    
    context = "\n\n".join([f"[{i+1}] {text}" for i, (text, _) in enumerate(results)])
    
    return {"context": context, "messages": state["messages"]}

# 生成节点
def generate_node(state: AgentState) -> dict:
    """生成回答"""
    model = get_language_model()
    messages = state["messages"] + [
        SystemMessage(content=f"""你是一个智能助手。请根据以下检索到的上下文回答问题。
如果上下文中没有相关信息,请诚实地说"我不知道"。

## 检索到的上下文:
{state['context']}
""")
    ]
    
    response = model.generate(messages)
    return {"messages": [AIMessage(content=response.choices[0].message.content)]}

# 条件边:判断是否需要重试
def should_retry(state: AgentState) -> str:
    """如果上下文为空或相关性低,返回 retry;否则返回 generate"""
    if not state["context"]:
        return "retry"
    return "generate"

# 重试节点(可选的二次检索)
def retry_node(state: AgentState) -> dict:
    """重新检索,扩大范围"""
    last_message = state["messages"][-1]
    engine = RagEngine()
    results = engine.retrieve(last_message.content, top_k=5)  # 扩大检索范围
    engine.close()
    
    context = "\n\n".join([f"[{i+1}] {text}" for i, (text, _) in enumerate(results)])
    return {"context": context, "messages": state["messages"]}

# 构建图
workflow.add_node("retrieve", retrieve_node)
workflow.add_node("retry", retry_node)
workflow.add_node("generate", generate_node)

workflow.add_edge(START, "retrieve")
workflow.add_conditional_edges(
    "retrieve",
    should_retry,
    {"retry": "retry", "generate": "generate"}
)
workflow.add_edge("retry", "generate")
workflow.add_edge("generate", END)

# 编译图
graph = workflow.compile()

19.2.2 调用 Agent

python
# test_agent.py
from rag_agent import graph

# 发起对话
response = graph.invoke({
    "messages": [
        {"role": "user", "content": "什么是 LangGraph?"}
    ]
})

print("Agent 回答:")
print(response["messages"][-1]["content"])

19.3 阶段三:持久化对话历史

状态图需要 Checkpointer 才能在多轮对话中保持上下文。LangGraph 支持多种 Checkpointer,我们先使用内存版,然后展示 PostgreSQL 版。

19.3.1 内存 Checkpointer

python
from langgraph.checkpoint.memory import MemorySaver

# 使用内存持久化
memory = MemorySaver()
graph = workflow.compile(checkpointer=memory)

# 第一次对话
thread1 = {"configurable": {"thread_id": "user_001"}}
response1 = graph.invoke(
    {"messages": [{"role": "user", "content": "解释一下 RAG"}]},
    config=thread1
)
print("第一轮:", response1["messages"][-1]["content"])

# 第二次对话(自动携带第一轮上下文)
response2 = graph.invoke(
    {"messages": [{"role": "user", "content": "它有什么优点?"}]},
    config=thread1  # 同一个 thread_id
)
print("第二轮:", response2["messages"][-1]["content"])

19.3.2 PostgreSQL Checkpointer(生产环境)

python
from langgraph.checkpoint.postgres import PostgresSaver
import psycopg2

# 连接 PostgreSQL
conn_str = os.getenv("DATABASE_URL", "postgresql://user:pass@localhost/db")
postgres_conn = psycopg2.connect(conn_str)

# 创建 Checkpointer
checkpointer = PostgresSaver(postgres_conn)
checkpointer.setup()

# 使用数据库持久化
graph = workflow.compile(checkpointer=checkpointer)

19.4 阶段四:FastAPI HTTP 接口

将 Agent 封装为 REST API,支持流式输出。

19.4.1 完整 API 服务

python
# main.py
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
from typing import List, Optional
import asyncio
from rag_agent import graph
from rag_engine import RagEngine

app = FastAPI(title="RAG Agent API")

app.add_middleware(
    CORSMiddleware,
    allow_origins=["*"],
    allow_methods=["*"],
    allow_headers=["*"],
)

# 请求模型
class ChatRequest(BaseModel):
    message: str
    thread_id: Optional[str] = "default"
    top_k: int = 3

class DocumentRequest(BaseModel):
    text: str
    metadata: Optional[dict] = None

class AddDocResponse(BaseModel):
    chunks: int
    status: str

# 聊天接口
@app.post("/chat")
async def chat(req: ChatRequest):
    """发送消息,获取 Agent 回答"""
    try:
        config = {
            "configurable": {"thread_id": req.thread_id}
        }
        
        response = await asyncio.to_thread(
            graph.invoke,
            {"messages": [{"role": "user", "content": req.message}]},
            config=config
        )
        
        return {
            "answer": response["messages"][-1]["content"],
            "thread_id": req.thread_id
        }
    except Exception as e:
        raise HTTPException(status_code=500, detail=str(e))

# 流式聊天接口
@app.post("/chat/stream")
async def chat_stream(req: ChatRequest):
    """流式输出"""
    config = {"configurable": {"thread_id": req.thread_id}}
    
    async def stream_events():
        async for event in graph.astream_events(
            {"messages": [{"role": "user", "content": req.message}]},
            config=config,
            version="v2"
        ):
            yield f"data: {event}\n\n"
    
    from fastapi.responses import StreamingResponse
    return StreamingResponse(
        stream_events(),
        media_type="text/event-stream"
    )

# 添加文档接口
@app.post("/documents", response_model=AddDocResponse)
async def add_document(req: DocumentRequest):
    """添加文档到知识库"""
    engine = RagEngine()
    engine.add_document(req.text, req.metadata)
    engine.close()
    return AddDocResponse(chunks=5, status="added")

# 健康检查
@app.get("/health")
async def health():
    return {"status": "ok"}

if __name__ == "__main__":
    import uvicorn
    uvicorn.run(app, host="0.0.0.0", port=8000)

19.4.2 客户端调用示例

python
# test_api.py
import requests

# 添加文档
requests.post("http://localhost:8000/documents", json={
    "text": "Deep Agents 是 LangChain 的最新 Agent 框架...",
    "metadata": {"source": "news"}
})

# 发送聊天请求
response = requests.post("http://localhost:8000/chat", json={
    "message": "什么是 Deep Agents?",
    "thread_id": "user_123"
})

print(response.json()["answer"])

19.5 阶段五:Langfuse 可观测性

接入 Langfuse 实现全链路追踪,监控 Token 使用、延迟和成本。

19.5.1 初始化 Langfuse

python
# observability.py
import os
from langfuse import Langfuse
from langfuse.decorators import observe
from contextlib import contextmanager

langfuse = Langfuse(
    public_key=os.getenv("LANGFUSE_PUBLIC_KEY"),
    secret_key=os.getenv("LANGFUSE_SECRET_KEY"),
    host=os.getenv("LANGFUSE_HOST", "https://cloud.langfuse.com")
)

@contextmanager
def trace_rag_call(user_id: str, session_id: str = None):
    """追踪一次 RAG 调用"""
    trace = langfuse.trace(
        name="rag-agent-call",
        user_id=user_id,
        session_id=session_id
    )
    try:
        yield trace
    finally:
        trace.flush()

19.5.2 集成到 Agent

python
# rag_agent.py (更新版)
from observability import trace_rag_call

@observe
def retrieve_node_traced(state: AgentState, user_id: str) -> dict:
    with trace_rag_call(user_id) as trace:
        trace.update(input={"query": state["messages"][-1]["content"]})
        
        result = retrieve_node(state)
        
        trace.update(
            output={"context_length": len(result["context"])},
            metadata={"retrieved_docs": len(result["context"].split("\n\n"))}
        )
        return result

19.5.3 Langfuse Dashboard

部署后,你可以在 Langfuse Dashboard 看到:

  • 每次调用的完整链路追踪
  • Token 使用量和成本统计
  • 检索相关性分布
  • 生成延迟分析
Trace: rag-agent-call-abc123
├── retrieve_node (12ms, 45 tokens)
│   ├── query: "什么是 LangGraph?"
│   └── results: 3 documents found
├── generate_node (2.3s, 128 tokens)
│   ├── context length: 850 chars
│   └── response: "LangGraph 是..."
└── Total cost: $0.0023

本章小结

本章完整实现了生产级 RAG Agent 服务:

  • 阶段一:基础 RAG 引擎,使用嵌入向量进行语义检索
  • 阶段二:LangGraph StateGraph 编排检索与生成逻辑
  • 阶段三:Checkpointer 实现多轮对话持久化
  • 阶段四:FastAPI 暴露 HTTP 接口,支持流式输出
  • 阶段五:Langfuse 全链路可观测性

关键知识点

  1. createOpenAICompatible 可接入任何 OpenAI 兼容的嵌入模型
  2. StateGraph 的条件边实现动态路由
  3. MemorySaver 用于开发,PostgresSaver 用于生产
  4. astream_events 提供流式追踪能力
  5. Langfuse 的 @observe 装饰器自动捕获执行细节

🛠️ 动手实践

  1. 扩展检索策略:在 retrieve_node 中加入重排序(rerank),对检索结果按相关性二次排序
  2. 添加文件上传:在 FastAPI 中添加 /documents/upload 端点,支持 PDF/Markdown 文件上传与解析
  3. 实现中断恢复:在 generate_node 前加入 interrupt(),实现人机协作的审批流程

🧪 随堂测验

点击你认为正确的选项。答错时会展示正确答案与原因解析。

1. 在 LangGraph 中,如何实现在生成节点前先检索知识库?

2. MemorySaver 和 PostgresSaver 的主要区别是什么?

3. Langfuse 的 @observe 装饰器主要功能是什么?

4. 在 RAG 系统中,文档分块(chunking)的主要目的是什么?