第 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-dotenv19.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 result19.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 全链路可观测性
关键知识点
createOpenAICompatible可接入任何 OpenAI 兼容的嵌入模型StateGraph的条件边实现动态路由MemorySaver用于开发,PostgresSaver用于生产astream_events提供流式追踪能力- Langfuse 的
@observe装饰器自动捕获执行细节
🛠️ 动手实践
- 扩展检索策略:在
retrieve_node中加入重排序(rerank),对检索结果按相关性二次排序 - 添加文件上传:在 FastAPI 中添加
/documents/upload端点,支持 PDF/Markdown 文件上传与解析 - 实现中断恢复:在
generate_node前加入interrupt(),实现人机协作的审批流程