构建实战 Wiki AI Agent — 从零到部署的完整教程
#Agent · #LangGraph · #RAG · #Wiki · #实战 · #FastAPI · #WebUI · #Docker
本文是一份可执行的完整教程,手把手教你构建一个能检索整个 Wiki 知识库并回答问题的 AI Agent。使用 LangGraph 编排任务流程,基于 RAG 检索 Wiki 内容,包含会话管理 Web 界面。每一步都可直接运行。
项目目标
构建一个Wiki AI 助手,用户可以通过 Web 界面:
- 用自然语言提问,Agent 自动检索 Wiki 知识库并回答
- 查看 Agent 的推理过程(使用了哪些工具、检索了哪些文档)
- 管理多个对话会话(新建、切换、删除)
- 支持流式输出,实时看到 Agent 的思考过程
mermaid
graph TD
USER["👤 用户<br/>通过 Web UI 提问"] --> API["🌐 FastAPI 后端<br/>WebSocket 流式传输"]
API --> AGENT["🤖 LangGraph Agent<br/>推理 → 行动 → 观察循环"]
AGENT --> TOOL1["📚 wiki_search<br/>检索 Wiki 知识库"]
AGENT --> TOOL2["📖 wiki_get_page<br/>获取完整页面内容"]
AGENT --> TOOL3["🗂️ wiki_list<br/>列出所有文档"]
TOOL1 --> RAG["🗄️ RAG 引擎<br/>Embedding + BM25<br/>多路召回 + Rerank"]
TOOL2 --> RAG
RAG --> VEC["向量数据库<br/>ChromaDB / Milvus"]
RAG --> IDX["倒排索引<br/>BM25"]
AGENT --> SESSION["💾 Session Store<br/>Redis / SQLite<br/>多轮对话持久化"]
style AGENT fill:#9b59b6,color:#fff
style RAG fill:#e74c3c,color:#fff
style API fill:#2ecc71,color:#fff项目结构
wiki-agent/
├── agent/ # Agent 核心
│ ├── __init__.py
│ ├── graph.py # LangGraph 状态图定义
│ ├── tools.py # Agent 工具定义(wiki_search 等)
│ ├── llm.py # LLM 客户端封装
│ └── state.py # Agent 状态定义
├── rag/ # RAG 检索模块
│ ├── __init__.py
│ ├── chunker.py # 文档解析与分块
│ ├── embeddings.py # Embedding 模型
│ ├── vector_store.py # ChromaDB 向量存储
│ ├── bm25_index.py # BM25 倒排索引
│ ├── retriever.py # 混合检索 + Rerank
│ └── indexer.py # 索引构建工具
├── server/ # Web 服务
│ ├── __init__.py
│ ├── main.py # FastAPI 应用
│ ├── routes.py # API 路由
│ ├── session.py # 会话管理
│ └── websocket.py # WebSocket 流式传输
├── web/ # 前端界面
│ ├── index.html # 主页面
│ ├── style.css # 样式
│ └── app.js # 前端逻辑
├── data/ # 数据目录
│ └── wiki/ # Wiki 源文件(软链接到实际位置)
├── requirements.txt # Python 依赖
├── Dockerfile # Docker 部署
├── docker-compose.yml # 一键部署
├── .env # 环境配置
└── README.md # 说明文档第一步:环境搭建
1.1 安装依赖
bash
# 创建虚拟环境
python3 -m venv venv
source venv/bin/activate
# 安装核心依赖
pip install \
langgraph langchain langchain-community langchain-openai \
chromadb sentence-transformers rank-bm25 \
fastapi uvicorn websockets python-multipart \
markdown beautifulsoup4 lxml \
redis sqlalchemy aiosqlite \
pydantic pydantic-settings \
tiktoken保存到 requirements.txt:
txt
# requirements.txt
langgraph>=0.2.0
langchain>=0.3.0
langchain-community>=0.3.0
langchain-openai>=0.2.0
chromadb>=0.5.0
sentence-transformers>=3.0.0
rank-bm25>=0.2.2
fastapi>=0.115.0
uvicorn[standard]>=0.30.0
websockets>=13.0
python-multipart>=0.0.12
markdown>=3.7
beautifulsoup4>=4.12.0
lxml>=5.3.0
redis>=5.0.0
sqlalchemy>=2.0.0
aiosqlite>=0.20.0
pydantic>=2.0.0
pydantic-settings>=2.0.0
tiktoken>=0.7.0
python-dotenv>=1.0.0第二步:RAG 检索引擎
2.1 文档分块引擎
python
# rag/chunker.py
"""Wiki 文档的结构化分块引擎"""
import re
import os
from typing import List, Dict
from dataclasses import dataclass, field
@dataclass
class Chunk:
"""文档分块"""
content: str # 文本内容
title: str = "" # 所属文档标题
section: str = "" # 所属章节
chunk_id: str = "" # 唯一标识
metadata: dict = field(default_factory=dict)
class WikiChunker:
"""基于 Markdown 结构的 Wiki 文档分块器
策略:
1. 以 Markdown 标题 (# → ######) 作为语义分段边界
2. 每个段落尽量保持完整
3. 块大小控制在 512-1024 tokens 之间
4. 相邻块有 128 tokens 的重叠
"""
def __init__(self, chunk_size: int = 1024, chunk_overlap: int = 128):
self.chunk_size = chunk_size
self.chunk_overlap = chunk_overlap
def chunk_file(self, filepath: str) -> List[Chunk]:
"""对单个文件进行分块"""
with open(filepath, 'r', encoding='utf-8') as f:
content = f.read()
# 提取第一个 # 标题作为文档标题
title = ""
title_match = re.search(r'^#\s+(.+)$', content, re.MULTILINE)
if title_match:
title = title_match.group(1).strip()
# 按标题拆分
sections = self._split_by_headings(content)
chunks = []
for section_title, section_content in sections:
# 如果段落太长,按句子进一步拆分
if len(section_content) <= self.chunk_size:
chunks.append(self._create_chunk(
section_content, title, section_title, len(chunks)
))
else:
sub_chunks = self._split_long_text(section_content)
for sc in sub_chunks:
chunks.append(self._create_chunk(
sc, title, section_title, len(chunks)
))
return chunks
def _split_by_headings(self, content: str) -> List[tuple]:
"""按 Markdown 标题拆分,返回 [(标题, 内容), ...]"""
# 用正则找所有标题位置
heading_pattern = re.compile(r'^(#{1,6})\s+(.+)$', re.MULTILINE)
matches = list(heading_pattern.finditer(content))
sections = []
for i, match in enumerate(matches):
start = match.end()
end = matches[i+1].start() if i+1 < len(matches) else len(content)
sections.append((match.group(2), content[start:end].strip()))
# 如果没有标题,整个文档作为一个段落
if not sections:
# 去掉 frontmatter 和 HTML 标签
clean = re.sub(r'<[^>]+>', '', content)
clean = re.sub(r'---.*?---', '', clean, flags=re.DOTALL)
sections = [("", clean.strip())]
return sections
def _split_long_text(self, text: str) -> List[str]:
"""将长文本按句子拆分,控制块大小"""
sentences = re.split(r'(?<=[。!?.!?])\s*', text)
chunks = []
current = ""
current_len = 0
for sent in sentences:
sent = sent.strip()
if not sent:
continue
if current_len + len(sent) > self.chunk_size:
if current:
chunks.append(current)
current = sent
current_len = len(sent)
else:
current += " " + sent if current else sent
current_len += len(sent)
if current:
chunks.append(current)
return chunks
def _create_chunk(self, content: str, title: str,
section: str, index: int) -> Chunk:
"""创建 Chunk 对象"""
chunk_id = f"{title}_{index}"
# 清理:移除 markdown 格式标记但保留结构
return Chunk(
content=content.strip(),
title=title,
section=section,
chunk_id=chunk_id,
metadata={
"title": title,
"section": section,
"char_count": len(content),
"index": index,
}
)
def chunk_directory(self, dir_path: str) -> List[Chunk]:
"""对整个目录的 Markdown 文件进行分块"""
all_chunks = []
for root, dirs, files in os.walk(dir_path):
# 跳过隐藏目录
dirs[:] = [d for d in dirs if not d.startswith('.')]
for file in files:
if file.endswith('.md'):
filepath = os.path.join(root, file)
try:
chunks = self.chunk_file(filepath)
all_chunks.extend(chunks)
print(f" ✓ {filepath}: {len(chunks)} chunks")
except Exception as e:
print(f" ✗ {filepath}: {e}")
return all_chunks2.2 Embedding 模型
python
# rag/embeddings.py
"""Embedding 模型封装"""
from typing import List
import numpy as np
from sentence_transformers import SentenceTransformer
class EmbeddingModel:
"""Embedding 模型(基于 sentence-transformers)
推荐模型:
- BAAI/bge-small-zh-v1.5 (中文,轻量)
- BAAI/bge-large-zh-v1.5 (中文,精度最高)
- intfloat/multilingual-e5-large (多语言)
"""
def __init__(self, model_name: str = "BAAI/bge-small-zh-v1.5",
device: str = "cpu"):
self.model = SentenceTransformer(model_name, device=device)
self.dim = self.model.get_sentence_embedding_dimension()
def encode(self, texts: List[str], batch_size: int = 32,
show_progress: bool = True) -> np.ndarray:
"""将文本列表编码为向量
Args:
texts: 文本列表
batch_size: 批处理大小
show_progress: 是否显示进度条
Returns:
numpy 数组 (len(texts), dim)
"""
# BGE 模型需要为查询添加前缀(文档不需要)
return self.model.encode(
texts,
batch_size=batch_size,
show_progress_bar=show_progress,
normalize_embeddings=True, # L2 归一化,用于余弦相似度
)
def encode_query(self, query: str) -> np.ndarray:
"""编码查询(BGE 模型查询需要前缀)"""
if "bge" in self.model._model_card_text.lower():
query = f"为这个句子生成表示以用于检索相关文章:{query}"
return self.model.encode(
[query],
normalize_embeddings=True,
show_progress_bar=False,
)[0]2.3 向量存储与 BM25
python
# rag/vector_store.py
"""ChromaDB 向量存储"""
import chromadb
from typing import List, Tuple
import numpy as np
from .chunker import Chunk
class VectorStore:
"""基于 ChromaDB 的向量存储"""
def __init__(self, collection_name: str = "wiki_docs",
persist_dir: str = "./data/chroma_db"):
self.client = chromadb.PersistentClient(path=persist_dir)
self.collection = self.client.get_or_create_collection(
name=collection_name,
metadata={"hnsw:space": "cosine"}
)
def add_chunks(self, chunks: List[Chunk],
embeddings: np.ndarray):
"""批量添加文档块"""
batch_size = 100
for i in range(0, len(chunks), batch_size):
batch = chunks[i:i+batch_size]
emb = embeddings[i:i+batch_size]
self.collection.add(
ids=[c.chunk_id for c in batch],
embeddings=emb.tolist() if isinstance(emb, np.ndarray) else emb,
documents=[c.content for c in batch],
metadatas=[c.metadata for c in batch],
)
print(f" 已添加 {min(i+batch_size, len(chunks))}/{len(chunks)} 个文档块")
def search(self, query_embedding: np.ndarray,
top_k: int = 10) -> List[Tuple[str, str, dict, float]]:
"""向量检索"""
results = self.collection.query(
query_embeddings=[query_embedding.tolist()],
n_results=top_k,
include=["documents", "metadatas", "distances"],
)
docs = []
for i in range(len(results['ids'][0])):
doc_id = results['ids'][0][i]
doc_text = results['documents'][0][i]
metadata = results['metadatas'][0][i]
distance = results['distances'][0][i]
# 余弦距离 → 相似度
similarity = 1.0 - distance
docs.append((doc_id, doc_text, metadata, similarity))
return docs
def count(self) -> int:
return self.collection.count()python
# rag/bm25_index.py
"""BM25 倒排索引"""
import re
import pickle
import numpy as np
from typing import List, Tuple
from rank_bm25 import BM25Okapi
from .chunker import Chunk
class BM25Index:
"""BM25 全文检索索引"""
def __init__(self):
self.bm25 = None
self.chunks = []
self.tokenized_corpus = []
def _tokenize(self, text: str) -> List[str]:
"""中文分词(简单实现,生产环境用 jieba)"""
# 简单的中英混合分词
# 中文:按字符切分 + n-gram
# 英文:按空格切分
segments = []
# 提取中文字符
chinese = ''.join(re.findall(r'[\u4e00-\u9fff]', text))
segments.extend(list(chinese))
# 提取英文单词
english = re.findall(r'[a-zA-Z]+', text.lower())
segments.extend(english)
return segments
def build(self, chunks: List[Chunk]):
"""构建 BM25 索引"""
self.chunks = chunks
self.tokenized_corpus = [
self._tokenize(chunk.content) for chunk in chunks
]
self.bm25 = BM25Okapi(self.tokenized_corpus)
print(f"BM25 索引已构建: {len(chunks)} 个文档")
def search(self, query: str, top_k: int = 10) -> List[Tuple[int, float]]:
"""BM25 检索"""
tokenized_query = self._tokenize(query)
scores = self.bm25.get_scores(tokenized_query)
# 返回 top_k 结果
top_indices = np.argsort(scores)[::-1][:top_k]
return [(int(idx), float(scores[idx])) for idx in top_indices]
def save(self, path: str):
"""保存索引"""
with open(path, 'wb') as f:
pickle.dump({
'tokenized_corpus': self.tokenized_corpus,
'chunks': [c.content for c in self.chunks],
}, f)
def load(self, path: str):
"""加载索引"""
with open(path, 'rb') as f:
data = pickle.load(f)
self.tokenized_corpus = data['tokenized_corpus']
self.bm25 = BM25Okapi(self.tokenized_corpus)2.4 混合检索器(多路召回 + RRF + Rerank)
python
# rag/retriever.py
"""混合检索器:向量 + BM25 + RRF 融合 + Rerank"""
from typing import List, Tuple, Dict
import numpy as np
from dataclasses import dataclass
from .vector_store import VectorStore
from .bm25_index import BM25Index
from .embeddings import EmbeddingModel
from .chunker import Chunk
@dataclass
class SearchResult:
"""检索结果"""
doc_id: str
content: str
score: float
metadata: dict
class HybridRetriever:
"""混合检索器
流程:
1. 向量检索 → Top-50
2. BM25 检索 → Top-50
3. RRF 融合 → Top-20
4. (可选) Rerank 精排 → Top-K
"""
def __init__(self, vector_store: VectorStore,
bm25_index: BM25Index,
embedding_model: EmbeddingModel,
use_rerank: bool = False):
self.vector_store = vector_store
self.bm25_index = bm25_index
self.embedding_model = embedding_model
self.use_rerank = use_rerank
def search(self, query: str, top_k: int = 5,
recall_size: int = 50) -> List[SearchResult]:
"""执行混合检索"""
# 1. 向量检索
query_embedding = self.embedding_model.encode_query(query)
vec_results = self.vector_store.search(query_embedding, top_k=recall_size)
# vec_results: [(id, text, metadata, score), ...]
# 2. BM25 检索
bm25_results = self.bm25_index.search(query, top_k=recall_size)
# bm25_results: [(index, score), ...]
# 3. RRF 融合
fused = self._rrf_fusion(vec_results, bm25_results, k=60)
# 取 Top-20 供 rerank
candidates = sorted(fused.items(), key=lambda x: x[1], reverse=True)[:20]
# 4. Rerank(可选)
if self.use_rerank and len(candidates) > top_k:
results = self._rerank(query, candidates, top_k)
else:
results = candidates[:top_k]
# 格式化输出
output = []
for doc_id, score in results:
# 从向量结果或 BM25 结果中获取内容
content = self._get_document_content(doc_id, vec_results)
if not content:
# BM25 结果中查找
chunk_idx = int(doc_id.split('_')[-1])
if chunk_idx < len(self.bm25_index.chunks):
chunk = self.bm25_index.chunks[chunk_idx]
content = chunk.content
output.append(SearchResult(
doc_id=doc_id,
content=content,
score=score,
metadata={} # 实际应填充元数据
))
return output
def _rrf_fusion(self, vec_results: List[Tuple],
bm25_results: List[Tuple],
k: int = 60) -> Dict[str, float]:
"""Reciprocal Rank Fusion"""
scores = {}
# 向量结果
for rank, (doc_id, _, _, _) in enumerate(vec_results):
scores[doc_id] = scores.get(doc_id, 0) + 1.0 / (k + rank + 1)
# BM25 结果
for rank, (chunk_idx, _) in enumerate(bm25_results):
# BM25 的 chunk_id 需要和向量存储一致
doc_id = self.bm25_index.chunks[chunk_idx].chunk_id
scores[doc_id] = scores.get(doc_id, 0) + 1.0 / (k + rank + 1)
return scores
def _rerank(self, query: str, candidates: List[Tuple[str, float]],
top_k: int) -> List[Tuple[str, float]]:
"""使用 Cross-Encoder 重排序"""
# 收集候选文档
doc_texts = [self._get_document_content(doc_id, []) for doc_id, _ in candidates]
# 使用 sentence-transformers 的 CrossEncoder
from sentence_transformers import CrossEncoder
reranker = CrossEncoder('BAAI/bge-reranker-v2-m3')
# 计算相关性分数
pairs = [[query, doc] for doc in doc_texts]
scores = reranker.predict(pairs)
# 按分数排序
ranked = sorted(zip(candidates, scores), key=lambda x: x[1], reverse=True)
return [(c[0][0], float(s)) for c, s in ranked[:top_k]]
def _get_document_content(self, doc_id: str,
vec_results: List) -> str:
"""根据 doc_id 获取文档内容"""
for result in vec_results:
if len(result) >= 2 and result[0] == doc_id:
return result[1]
return ""2.5 索引构建工具
python
# rag/indexer.py
"""Wiki 文档索引构建工具"""
import os
import sys
from .chunker import WikiChunker
from .embeddings import EmbeddingModel
from .vector_store import VectorStore
from .bm25_index import BM25Index
def build_wiki_index(wiki_dir: str):
"""构建完整的 Wiki 知识库索引
Args:
wiki_dir: Wiki Markdown 文件所在目录
"""
print(f"📚 开始构建 Wiki 索引: {wiki_dir}")
print("=" * 50)
# 1. 分块
print("\n1️⃣ 文档分块中...")
chunker = WikiChunker(chunk_size=1024, chunk_overlap=128)
chunks = chunker.chunk_directory(wiki_dir)
print(f" 共产生 {len(chunks)} 个文档块")
if not chunks:
print("❌ 未找到任何文档块,请检查 Wiki 目录路径")
return
# 2. 向量化
print("\n2️⃣ 向量化中...")
embed_model = EmbeddingModel(
model_name="BAAI/bge-small-zh-v1.5",
device="cpu"
)
embeddings = embed_model.encode(
[c.content for c in chunks],
batch_size=32,
show_progress=True,
)
print(f" 向量维度: {embeddings.shape[1]}")
# 3. 存入向量数据库
print("\n3️⃣ 存入向量数据库中...")
vector_store = VectorStore(
collection_name="wiki_docs",
persist_dir="./data/chroma_db",
)
vector_store.add_chunks(chunks, embeddings)
print(f" 向量数据库条目: {vector_store.count()}")
# 4. 构建 BM25 索引
print("\n4️⃣ 构建 BM25 全文索引中...")
bm25_index = BM25Index()
bm25_index.build(chunks)
bm25_index.save("./data/bm25_index.pkl")
print(f" BM25 索引条目: {len(chunks)}")
print("\n✅ Wiki 索引构建完成!")
if __name__ == "__main__":
# 构建索引
wiki_path = sys.argv[1] if len(sys.argv) > 1 else "/data/code/wiki/learn"
build_wiki_index(wiki_path)第三步:LangGraph Agent 实现
3.1 Agent 状态定义
python
# agent/state.py
"""Agent 状态定义"""
from typing import List, Dict, Any, TypedDict, Annotated
from operator import add
class AgentState(TypedDict):
"""LangGraph Agent 的状态
LangGraph 的状态是逐步累加的,每个节点返回的状态会与当前状态合并。
"""
# 对话消息列表
messages: Annotated[List[Dict[str, str]], add]
# 当前任务(用户最新问题)
current_task: str
# 工具调用历史
tool_calls: Annotated[List[Dict[str, Any]], add]
# 检索到的文档
retrieved_docs: Annotated[List[str], add]
# 当前推理步骤
reasoning_steps: Annotated[List[str], add]
# 是否完成
is_finished: bool
# 最终答案
final_answer: str
# 循环计数(防止死循环)
loop_count: int3.2 工具定义
python
# agent/tools.py
"""Agent 工具定义 — 让 Agent 能检索 Wiki"""
from typing import List, Dict
from langchain_core.tools import tool
from rag.retriever import HybridRetriever
from rag.embeddings import EmbeddingModel
from rag.vector_store import VectorStore
from rag.bm25_index import BM25Index
import os
import glob
# === 全局检索器实例(在 Agent 初始化时设置) ===
_retriever: HybridRetriever = None
_wiki_dir: str = ""
def init_tools(wiki_dir: str = "/data/code/wiki/learn"):
"""初始化工具(在 Agent 启动时调用一次)"""
global _retriever, _wiki_dir
_wiki_dir = wiki_dir
# 构建检索器
embed_model = EmbeddingModel("BAAI/bge-small-zh-v1.5")
vector_store = VectorStore(
collection_name="wiki_docs",
persist_dir="./data/chroma_db",
)
bm25_index = BM25Index()
bm25_index.load("./data/bm25_index.pkl")
_retriever = HybridRetriever(vector_store, bm25_index, embed_model)
@tool
def wiki_search(query: str, top_k: int = 5) -> str:
"""
在 Wiki 知识库中搜索相关文档。当用户询问技术问题、需要查找知识点时使用。
Args:
query: 搜索关键词或问题(建议用关键词形式,如 "Go 语言 GMP 调度模型")
top_k: 返回的文档数量,默认 5
Returns:
相关文档片段,包含标题和内容摘要
"""
if _retriever is None:
return "❌ 检索器未初始化,请先运行 init_tools()"
results = _retriever.search(query, top_k=top_k)
if not results:
return "未找到相关文档。"
output = []
for i, result in enumerate(results, 1):
title = result.metadata.get("title", "未知文档")
section = result.metadata.get("section", "")
# 截取内容前 500 字符
content_preview = result.content[:500].replace("\n", " ")
if len(result.content) > 500:
content_preview += "..."
source = f"{title}"
if section:
source += f" > {section}"
output.append(
f"📄 [{i}] {source}\n"
f" 相关度: {result.score:.3f}\n"
f" 摘要: {content_preview}\n"
)
return "\n".join(output)
@tool
def wiki_get_page(title: str) -> str:
"""
获取 Wiki 中指定页面的完整内容。当需要详细阅读某个文档时使用。
Args:
title: 文档标题(如 "GMP 调度模型" 或文件名如 "gmp")
Returns:
文档的完整 Markdown 内容(截断到 5000 字符)
"""
global _wiki_dir
# 搜索匹配的文件
files = []
for root, dirs, files_list in os.walk(_wiki_dir):
dirs[:] = [d for d in dirs if not d.startswith('.')]
for f in files_list:
if f.endswith('.md'):
# 模糊匹配标题
if title.lower() in f.lower() or title.lower() in os.path.splitext(f)[0].lower():
files.append(os.path.join(root, f))
if not files:
# 尝试更宽松的匹配
for root, dirs, files_list in os.walk(_wiki_dir):
dirs[:] = [d for d in dirs if not d.startswith('.')]
for f in files_list:
if f.endswith('.md'):
with open(os.path.join(root, f), 'r', encoding='utf-8') as fp:
first_line = fp.readline().strip('#').strip()
if title.lower() in first_line.lower():
files.append(os.path.join(root, f))
if not files:
return f"❌ 未找到标题包含 '{title}' 的文档。可用工具 wiki_list 查看所有文档。"
# 返回第一个匹配文件的内容
filepath = files[0]
try:
with open(filepath, 'r', encoding='utf-8') as f:
content = f.read()
# 截取前 5000 字符
if len(content) > 5000:
content = content[:5000] + "\n\n... (内容过长,已截断)"
return f"📄 文档: {os.path.basename(filepath)}\n\n{content}"
except Exception as e:
return f"❌ 读取文档失败: {e}"
@tool
def wiki_list() -> str:
"""
列出 Wiki 知识库中所有可用的文档。用于让用户了解有哪些内容可以查询。
Returns:
文档列表(目录树形式)
"""
global _wiki_dir
lines = ["📚 Wiki 文档列表:\n"]
for root, dirs, files in os.walk(_wiki_dir):
# 跳过隐藏目录
dirs[:] = sorted([d for d in dirs if not d.startswith('.')])
level = root.replace(_wiki_dir, '').count(os.sep)
indent = ' ' * level
dir_name = os.path.basename(root)
if level > 0:
lines.append(f"{indent}📁 {dir_name}/")
for f in sorted(files):
if f.endswith('.md'):
# 尝试读取标题
filepath = os.path.join(root, f)
title = f.replace('.md', '')
try:
with open(filepath, 'r', encoding='utf-8') as fp:
first_lines = fp.read(200)
h1_match = first_lines.split('\n')
for line in h1_match:
if line.startswith('# '):
title = line[2:].strip()
break
except Exception:
pass
lines.append(f"{indent} 📄 {title}")
# 限制输出长度
if len(lines) > 100:
lines.append(" ... (更多文档省略)")
break
return '\n'.join(lines[:100])
# 注册工具列表(送給 LangGraph)
TOOLS = [wiki_search, wiki_get_page, wiki_list]3.3 LangGraph Agent 核心
python
# agent/graph.py
"""LangGraph Agent 核心 — 基于 ReAct 的任务规划与执行"""
from typing import Literal
from langgraph.graph import StateGraph, END
from langchain_core.messages import HumanMessage, SystemMessage, AIMessage, ToolMessage
from langchain_openai import ChatOpenAI
import json
import re
from .state import AgentState
from .tools import TOOLS, init_tools
class WikiAgent:
"""基于 LangGraph 的 Wiki 检索 Agent
ReAct 循环: 思考(Thought) → 行动(Action) → 观察(Observation) → ...
"""
def __init__(self, model_name: str = "gpt-4o-mini",
api_key: str = None,
base_url: str = None,
wiki_dir: str = "/data/code/wiki/learn",
max_steps: int = 5,
temperature: float = 0.1):
"""
Args:
model_name: LLM 模型名(OpenAI 兼容接口)
api_key: API 密钥
base_url: API 地址(可选,用于 Ollama/vLLM 等本地模型)
wiki_dir: Wiki 文档目录
max_steps: Agent 最大推理步数
temperature: 模型温度(低温度 = 更确定)
"""
# 初始化 LLM
self.llm = ChatOpenAI(
model=model_name,
api_key=api_key or "not-needed",
base_url=base_url,
temperature=temperature,
)
# 绑定工具
self.llm_with_tools = self.llm.bind_tools(TOOLS)
self.max_steps = max_steps
self.wiki_dir = wiki_dir
# 初始化工具
init_tools(wiki_dir)
# 构建 LangGraph
self.graph = self._build_graph()
def _build_graph(self) -> StateGraph:
"""构建 LangGraph 状态图"""
workflow = StateGraph(AgentState)
# 添加节点
workflow.add_node("think", self._think_node) # 推理节点
workflow.add_node("act", self._act_node) # 工具调用节点
workflow.add_node("summarize", self._summarize_node) # 总结节点
# 设置入口
workflow.set_entry_point("think")
# 添加边
workflow.add_conditional_edges(
"think",
self._should_act, # 决定下一步
{
"act": "act", # 需要调用工具
"summarize": "summarize", # 直接回答
"end": END, # 结束
}
)
workflow.add_edge("act", "think") # 工具调用后继续推理
workflow.add_edge("summarize", END) # 总结后结束
return workflow.compile()
def _think_node(self, state: AgentState) -> AgentState:
"""推理节点:LLM 决定下一步行动"""
messages = state.get("messages", [])
loop_count = state.get("loop_count", 0)
# 构建系统提示
system_prompt = """你是一个 Wiki 知识库助手,能帮助用户检索和学习知识库中的内容。
## 可用工具
- **wiki_search(query, top_k=5)**: 搜索知识库中的相关文档
- **wiki_get_page(title)**: 获取指定文档的完整内容
- **wiki_list()**: 列出所有可用文档
## 工作流程
1. 如果用户的问题需要检索知识库,使用 wiki_search 搜索
2. 搜索后如果需要查看完整文档,使用 wiki_get_page
3. 基于检索到的内容回答问题。必须引用知识库中的内容,不要编造信息。
4. 如果知识库中没有相关信息,礼貌地告知用户。
## 规则
- 优先使用知识库中的内容回答,不要编造
- 如果搜索无结果,告知用户并建议换个关键词
- 一次搜索足够回答问题时,不要再重复搜索
- 回答要简洁、准确、易读"""
# 构建消息
if not messages:
messages = [
SystemMessage(content=system_prompt),
HumanMessage(content=state.get("current_task", "")),
]
# 调用 LLM
response = self.llm_with_tools.invoke(messages)
# 更新状态
new_messages = messages + [response]
new_steps = state.get("reasoning_steps", []) + [
f"Step {loop_count + 1}: {response.content[:200]}"
]
return {
"messages": new_messages,
"reasoning_steps": new_steps,
"loop_count": loop_count + 1,
}
def _should_act(self, state: AgentState) -> str:
"""条件路由:决定下一步动作"""
messages = state["messages"]
last_message = messages[-1] if messages else None
loop_count = state.get("loop_count", 0)
# 检查是否超过最大步数
if loop_count >= self.max_steps:
return "summarize"
# 检查是否有 tool_calls
if hasattr(last_message, 'tool_calls') and last_message.tool_calls:
return "act"
# LLM 给出的直接回答
if hasattr(last_message, 'content') and last_message.content:
return "summarize"
return "end"
def _act_node(self, state: AgentState) -> AgentState:
"""工具调用节点:执行 LLM 请求的工具"""
messages = state["messages"]
last_message = messages[-1]
tool_calls_history = state.get("tool_calls", [])
new_messages = list(messages)
retrieved_docs = state.get("retrieved_docs", [])
# 执行每个 tool call
for tool_call in last_message.tool_calls:
tool_name = tool_call["name"]
tool_args = tool_call["args"]
# 查找并执行工具
tool_result = self._execute_tool(tool_name, tool_args)
# 记录工具调用
tool_calls_history.append({
"tool": tool_name,
"args": tool_args,
"result": tool_result[:200] + "..." if len(tool_result) > 200 else tool_result,
})
# 如果是检索类工具,保存检索结果
if tool_name in ("wiki_search", "wiki_get_page"):
retrieved_docs.append(tool_result[:500])
# 创建 ToolMessage
new_messages.append(ToolMessage(
content=tool_result,
tool_call_id=tool_call["id"],
))
return {
"messages": new_messages,
"tool_calls": tool_calls_history,
"retrieved_docs": retrieved_docs,
}
def _summarize_node(self, state: AgentState) -> AgentState:
"""总结节点:生成最终回答"""
messages = state["messages"]
last_message = messages[-1]
final_answer = ""
if hasattr(last_message, 'content'):
final_answer = last_message.content
return {
"final_answer": final_answer,
"is_finished": True,
}
def _execute_tool(self, tool_name: str, args: dict) -> str:
"""执行工具并返回结果"""
for tool in TOOLS:
if tool.name == tool_name:
try:
result = tool.invoke(args)
return str(result)
except Exception as e:
return f"工具 {tool_name} 执行失败: {str(e)}"
return f"未知工具: {tool_name}"
def run(self, question: str) -> dict:
"""运行 Agent
Returns:
{
"answer": "最终回答",
"steps": ["步骤1", "步骤2", ...],
"tool_calls": [...],
"retrieved_docs": [...],
}
"""
initial_state: AgentState = {
"messages": [],
"current_task": question,
"tool_calls": [],
"retrieved_docs": [],
"reasoning_steps": [],
"is_finished": False,
"final_answer": "",
"loop_count": 0,
}
result = self.graph.invoke(initial_state)
return {
"answer": result.get("final_answer", "抱歉,我无法回答这个问题。"),
"steps": result.get("reasoning_steps", []),
"tool_calls": result.get("tool_calls", []),
"retrieved_docs": result.get("retrieved_docs", []),
}
async def run_stream(self, question: str):
"""流式运行 Agent,逐步产出事件"""
initial_state: AgentState = {
"messages": [],
"current_task": question,
"tool_calls": [],
"retrieved_docs": [],
"reasoning_steps": [],
"is_finished": False,
"final_answer": "",
"loop_count": 0,
}
async for event in self.graph.astream(initial_state):
yield event第四步:FastAPI 后端服务
4.1 会话管理
python
# server/session.py
"""会话管理 — 使用 SQLite 持久化多轮对话"""
import sqlalchemy as sa
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession
from sqlalchemy.orm import sessionmaker, declarative_base
from datetime import datetime, timedelta
import json
import uuid
from typing import List, Dict, Optional
Base = declarative_base()
class SessionModel(Base):
"""会话数据库模型"""
__tablename__ = "sessions"
id = sa.Column(sa.String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
user_id = sa.Column(sa.String(64), default="default")
title = sa.Column(sa.String(200), default="新对话")
messages = sa.Column(sa.Text, default="[]") # JSON 格式
created_at = sa.Column(sa.DateTime, default=datetime.utcnow)
updated_at = sa.Column(sa.DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
class SessionManager:
"""会话管理器"""
def __init__(self, db_url: str = "sqlite+aiosqlite:///./data/sessions.db"):
self.engine = create_async_engine(db_url, echo=False)
self.async_session = sessionmaker(
self.engine, class_=AsyncSession, expire_on_commit=False
)
async def init_db(self):
"""初始化数据库表"""
async with self.engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
async def create_session(self, title: str = "新对话") -> str:
"""创建新会话"""
session_id = str(uuid.uuid4())
async with self.async_session() as db:
session = SessionModel(
id=session_id,
title=title,
messages="[]",
)
db.add(session)
await db.commit()
return session_id
async def get_session(self, session_id: str) -> Optional[dict]:
"""获取会话详情"""
async with self.async_session() as db:
result = await db.execute(
sa.select(SessionModel).where(SessionModel.id == session_id)
)
session = result.scalar_one_or_none()
if session:
return {
"id": session.id,
"title": session.title,
"messages": json.loads(session.messages),
"created_at": session.created_at.isoformat(),
"updated_at": session.updated_at.isoformat(),
}
return None
async def list_sessions(self, user_id: str = "default") -> List[dict]:
"""列出所有会话"""
async with self.async_session() as db:
result = await db.execute(
sa.select(SessionModel)
.where(SessionModel.user_id == user_id)
.order_by(SessionModel.updated_at.desc())
.limit(50)
)
sessions = result.scalars().all()
return [
{
"id": s.id,
"title": s.title,
"message_count": len(json.loads(s.messages)),
"updated_at": s.updated_at.isoformat(),
}
for s in sessions
]
async def add_message(self, session_id: str, role: str,
content: str):
"""向会话添加消息"""
async with self.async_session() as db:
result = await db.execute(
sa.select(SessionModel).where(SessionModel.id == session_id)
)
session = result.scalar_one_or_none()
if session:
messages = json.loads(session.messages)
messages.append({
"role": role,
"content": content,
"time": datetime.utcnow().isoformat(),
})
session.messages = json.dumps(messages, ensure_ascii=False)
# 自动生成标题(使用第一条用户消息的前 30 字)
if session.title == "新对话" and role == "user":
session.title = content[:30] + ("..." if len(content) > 30 else "")
session.updated_at = datetime.utcnow()
await db.commit()
async def delete_session(self, session_id: str):
"""删除会话"""
async with self.async_session() as db:
await db.execute(
sa.delete(SessionModel).where(SessionModel.id == session_id)
)
await db.commit()4.2 FastAPI 主应用
python
# server/main.py
"""FastAPI 主应用 — Wiki AI Agent 后端服务"""
import os
import json
import asyncio
from contextlib import asynccontextmanager
from fastapi import FastAPI, WebSocket, WebSocketDisconnect, HTTPException
from fastapi.staticfiles import StaticFiles
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
from typing import List, Optional
from dotenv import load_dotenv
from agent.graph import WikiAgent
from server.session import SessionManager
load_dotenv()
# === 全局实例 ===
agent: WikiAgent = None
session_manager: SessionManager = None
@asynccontextmanager
async def lifespan(app: FastAPI):
"""应用生命周期"""
global agent, session_manager
print("🚀 启动 Wiki AI Agent...")
# 初始化会话管理
session_manager = SessionManager()
await session_manager.init_db()
print(" ✓ 会话数据库已初始化")
# 初始化 Agent
agent = WikiAgent(
model_name=os.getenv("LLM_MODEL", "gpt-4o-mini"),
api_key=os.getenv("OPENAI_API_KEY", "not-needed"),
base_url=os.getenv("OPENAI_BASE_URL", None),
wiki_dir=os.getenv("WIKI_DIR", "/data/code/wiki/learn"),
max_steps=int(os.getenv("MAX_STEPS", "5")),
temperature=float(os.getenv("TEMPERATURE", "0.1")),
)
print(f" ✓ Agent 已初始化 (模型: {os.getenv('LLM_MODEL', 'gpt-4o-mini')})")
print(f" ✓ Wiki 目录: {os.getenv('WIKI_DIR', '/data/code/wiki/learn')}")
yield
print("👋 关闭 Wiki AI Agent...")
# 创建 FastAPI 应用
app = FastAPI(
title="Wiki AI Agent",
description="基于 LangGraph + RAG 的 Wiki 知识库智能助手",
version="1.0.0",
lifespan=lifespan,
)
# CORS
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# === 请求/响应模型 ===
class ChatRequest(BaseModel):
session_id: Optional[str] = None
message: str
class ChatResponse(BaseModel):
session_id: str
answer: str
steps: List[str] = []
tool_calls: List[dict] = []
# === API 路由 ===
@app.get("/api/sessions")
async def list_sessions():
"""列出所有会话"""
sessions = await session_manager.list_sessions()
return {"sessions": sessions}
@app.post("/api/sessions")
async def create_session():
"""创建新会话"""
session_id = await session_manager.create_session()
return {"session_id": session_id}
@app.delete("/api/sessions/{session_id}")
async def delete_session(session_id: str):
"""删除会话"""
await session_manager.delete_session(session_id)
return {"status": "ok"}
@app.get("/api/sessions/{session_id}")
async def get_session(session_id: str):
"""获取会话详情"""
session = await session_manager.get_session(session_id)
if not session:
raise HTTPException(status_code=404, detail="会话不存在")
return session
@app.post("/api/chat")
async def chat(request: ChatRequest) -> ChatResponse:
"""同步聊天接口(非流式)"""
# 如果没有 session_id,创建新会话
session_id = request.session_id
if not session_id:
session_id = await session_manager.create_session()
# 保存用户消息
await session_manager.add_message(session_id, "user", request.message)
# 运行 Agent
result = agent.run(request.message)
# 保存 Agent 回答
await session_manager.add_message(session_id, "assistant", result["answer"])
return ChatResponse(
session_id=session_id,
answer=result["answer"],
steps=result["steps"],
tool_calls=result["tool_calls"],
)
@app.websocket("/ws/chat")
async def websocket_chat(websocket: WebSocket):
"""WebSocket 流式聊天"""
await websocket.accept()
try:
while True:
# 接收消息
data = await websocket.receive_text()
request = json.loads(data)
session_id = request.get("session_id")
message = request.get("message", "")
if not session_id:
session_id = await session_manager.create_session()
await websocket.send_text(json.dumps({
"type": "session_created",
"session_id": session_id,
}))
# 保存用户消息
await session_manager.add_message(session_id, "user", message)
# 流式运行 Agent
await websocket.send_text(json.dumps({
"type": "thinking_start",
"session_id": session_id,
}))
full_answer = []
async for event in agent.run_stream(message):
# 发送 Agent 中间步骤
await websocket.send_text(json.dumps({
"type": "agent_event",
"event": {k: str(v)[:500] for k, v in event.items()},
}))
# 如果是最终节点,提取答案
if "summarize" in event:
final_state = event["summarize"]
if final_state.get("final_answer"):
full_answer.append(final_state["final_answer"])
# 最终答案
answer_text = "".join(full_answer)
await session_manager.add_message(session_id, "assistant", answer_text)
await websocket.send_text(json.dumps({
"type": "answer",
"session_id": session_id,
"content": answer_text,
}))
except WebSocketDisconnect:
print("WebSocket 断开连接")
# 静态文件(前端界面)
app.mount("/", StaticFiles(directory="web", html=True), name="static")
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8000)第五步:Web 前端界面
5.1 HTML 主页面
html
<!-- web/index.html -->
<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>Wiki AI Agent</title>
<link rel="stylesheet" href="style.css">
<link rel="icon" href="data:image/svg+xml,<svg xmlns='http://www.w3.org/2000/svg' viewBox='0 0 100 100'><text y='.9em' font-size='90'>🤖</text></svg>">
</head>
<body>
<div id="app">
<!-- 侧边栏:会话列表 -->
<aside id="sidebar">
<div id="sidebar-header">
<h2>🤖 Wiki AI Agent</h2>
<button id="new-chat-btn" onclick="createNewSession()" title="新建对话">+</button>
</div>
<div id="session-list"></div>
</aside>
<!-- 主区域:对话界面 -->
<main id="main">
<!-- 对话区域 -->
<div id="chat-container">
<div id="chat-messages">
<div class="welcome-message">
<h1>🤖 Wiki AI Agent</h1>
<p>我是你的 Wiki 知识库助手,可以帮你检索和学习知识库中的所有内容。</p>
<p>试试问我:</p>
<div class="suggestions">
<button onclick="askQuestion('Go 语言的 GMP 调度模型是什么?')">GMP 调度模型</button>
<button onclick="askQuestion('介绍一下 Transformer 架构')">Transformer 架构</button>
<button onclick="askQuestion('Redis 有哪些常用的数据结构?')">Redis 数据结构</button>
<button onclick="askQuestion('什么是 TCP 三次握手?')">TCP 三次握手</button>
</div>
</div>
</div>
</div>
<!-- 输入区域 -->
<div id="input-area">
<div id="input-container">
<textarea id="message-input"
placeholder="输入你的问题..."
rows="1"
onkeydown="handleKeyDown(event)"></textarea>
<button id="send-btn" onclick="sendMessage()" title="发送 (Enter)">
<svg width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2">
<path d="M22 2L11 13"></path>
<path d="M22 2L15 22L11 13L2 9L22 2Z"></path>
</svg>
</button>
</div>
</div>
</main>
<!-- Agent 推理过程面板 -->
<aside id="debug-panel" class="collapsed">
<div id="debug-header">
<h3>🔍 推理过程</h3>
<button onclick="toggleDebug()" title="收起">×</button>
</div>
<div id="debug-content"></div>
</aside>
</div>
<script src="app.js"></script>
</body>
</html>5.2 前端 JavaScript
javascript
// web/app.js
// === 状态管理 ===
const state = {
sessionId: null,
isProcessing: false,
};
// === 初始化 ===
document.addEventListener('DOMContentLoaded', () => {
loadSessions();
autoResizeTextarea();
});
// === 会话管理 ===
async function createNewSession() {
try {
const res = await fetch('/api/sessions', { method: 'POST' });
const data = await res.json();
state.sessionId = data.session_id;
// 清空对话
document.getElementById('chat-messages').innerHTML = '';
// 重新加载会话列表
await loadSessions();
selectSession(state.sessionId);
} catch (err) {
console.error('创建会话失败:', err);
}
}
async function loadSessions() {
try {
const res = await fetch('/api/sessions');
const data = await res.json();
renderSessionList(data.sessions);
} catch (err) {
console.error('加载会话失败:', err);
}
}
function renderSessionList(sessions) {
const container = document.getElementById('session-list');
container.innerHTML = sessions.map(s => `
<div class="session-item ${s.id === state.sessionId ? 'active' : ''}"
onclick="selectSession('${s.id}')">
<span class="session-title">${escapeHtml(s.title)}</span>
<span class="session-info">${s.message_count} 条消息</span>
<button class="session-delete" onclick="event.stopPropagation(); deleteSession('${s.id}')">🗑️</button>
</div>
`).join('');
}
async function selectSession(sessionId) {
state.sessionId = sessionId;
try {
const res = await fetch(`/api/sessions/${sessionId}`);
const session = await res.json();
// 渲染历史消息
const container = document.getElementById('chat-messages');
container.innerHTML = session.messages.map(msg => `
<div class="message ${msg.role}">
<div class="message-role">${msg.role === 'user' ? '👤' : '🤖'}</div>
<div class="message-content">${formatMessage(msg.content)}</div>
<div class="message-time">${formatTime(msg.time)}</div>
</div>
`).join('');
scrollToBottom();
await loadSessions();
} catch (err) {
console.error('加载会话失败:', err);
}
}
async function deleteSession(sessionId) {
if (!confirm('确定删除这个会话吗?')) return;
try {
await fetch(`/api/sessions/${sessionId}`, { method: 'DELETE' });
if (state.sessionId === sessionId) {
state.sessionId = null;
document.getElementById('chat-messages').innerHTML = `
<div class="welcome-message">
<h1>🤖 Wiki AI Agent</h1>
<p>选择或创建一个对话开始吧!</p>
</div>
`;
}
await loadSessions();
} catch (err) {
console.error('删除会话失败:', err);
}
}
// === 消息处理 ===
function askQuestion(question) {
document.getElementById('message-input').value = question;
sendMessage();
}
async function sendMessage() {
if (state.isProcessing) return;
const input = document.getElementById('message-input');
const message = input.value.trim();
if (!message) return;
// 清空输入
input.value = '';
autoResizeTextarea();
// 如果没有会话,先创建
if (!state.sessionId) {
await createNewSession();
}
state.isProcessing = true;
document.getElementById('send-btn').disabled = true;
// 显示用户消息
appendMessage('user', message);
// 显示加载动画
const loadingId = appendLoading();
try {
const res = await fetch('/api/chat', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({
session_id: state.sessionId,
message: message,
}),
});
const data = await res.json();
// 移除加载动画
removeLoading(loadingId);
// 显示回答
appendMessage('assistant', data.answer);
// 显示推理过程
if (data.steps && data.steps.length > 0) {
showDebugInfo(data.steps, data.tool_calls);
}
// 刷新会话列表(更新标题和消息数)
await loadSessions();
selectSessionInList(state.sessionId);
} catch (err) {
removeLoading(loadingId);
appendMessage('system', `❌ 请求失败: ${err.message}`);
console.error(err);
}
state.isProcessing = false;
document.getElementById('send-btn').disabled = false;
}
function appendMessage(role, content) {
const container = document.getElementById('chat-messages');
// 隐藏欢迎消息
const welcome = container.querySelector('.welcome-message');
if (welcome) welcome.style.display = 'none';
const div = document.createElement('div');
div.className = `message ${role}`;
div.innerHTML = `
<div class="message-role">${role === 'user' ? '👤' : '🤖'}</div>
<div class="message-content">${formatMessage(content)}</div>
`;
container.appendChild(div);
scrollToBottom();
}
function appendLoading() {
const container = document.getElementById('chat-messages');
const id = 'loading-' + Date.now();
const div = document.createElement('div');
div.id = id;
div.className = 'message assistant loading';
div.innerHTML = `
<div class="message-role">🤖</div>
<div class="message-content">
<span class="dot"></span><span class="dot"></span><span class="dot"></span>
</div>
`;
container.appendChild(div);
scrollToBottom();
return id;
}
function removeLoading(id) {
const el = document.getElementById(id);
if (el) el.remove();
}
// === 推理过程面板 ===
function showDebugInfo(steps, toolCalls) {
const panel = document.getElementById('debug-panel');
panel.classList.remove('collapsed');
const content = document.getElementById('debug-content');
let html = '<div class="debug-section"><h4>💭 推理步骤</h4><ol>';
steps.forEach(s => { html += `<li>${escapeHtml(s)}</li>`; });
html += '</ol></div>';
if (toolCalls && toolCalls.length > 0) {
html += '<div class="debug-section"><h4>🔧 工具调用</h4>';
toolCalls.forEach(tc => {
html += `<div class="tool-call">
<strong>${escapeHtml(tc.tool)}</strong>
<pre>${escapeHtml(typeof tc.args === 'string' ? tc.args : JSON.stringify(tc.args, null, 2))}</pre>
<div class="tool-result">${escapeHtml(tc.result).substring(0, 300)}</div>
</div>`;
});
html += '</div>';
}
content.innerHTML = html;
}
function toggleDebug() {
document.getElementById('debug-panel').classList.toggle('collapsed');
}
// === 工具函数 ===
function handleKeyDown(event) {
if (event.key === 'Enter' && !event.shiftKey) {
event.preventDefault();
sendMessage();
}
}
function autoResizeTextarea() {
const textarea = document.getElementById('message-input');
textarea.addEventListener('input', () => {
textarea.style.height = 'auto';
textarea.style.height = Math.min(textarea.scrollHeight, 200) + 'px';
});
}
function formatMessage(content) {
// 简单的 Markdown 渲染
if (!content) return '';
return content
.replace(/\*\*(.+?)\*\*/g, '<strong>$1</strong>') // **bold**
.replace(/\*(.+?)\*/g, '<em>$1</em>') // *italic*
.replace(/`([^`]+)`/g, '<code>$1</code>') // `code`
.replace(/\n/g, '<br>') // 换行
.replace(/📄/g, '<span class="icon">📄</span>') // 图标
.replace(/📚/g, '<span class="icon">📚</span>');
}
function formatTime(isoString) {
if (!isoString) return '';
const date = new Date(isoString);
return date.toLocaleTimeString('zh-CN', { hour: '2-digit', minute: '2-digit' });
}
function escapeHtml(text) {
if (!text) return '';
const div = document.createElement('div');
div.textContent = text;
return div.innerHTML;
}
function scrollToBottom() {
const container = document.getElementById('chat-messages');
setTimeout(() => {
container.scrollTop = container.scrollHeight;
}, 100);
}
function selectSessionInList(sessionId) {
document.querySelectorAll('.session-item').forEach(el => {
el.classList.toggle('active', el.dataset?.sessionId === sessionId);
});
}5.3 样式
css
/* web/style.css */
:root {
--bg-primary: #1a1a2e;
--bg-secondary: #16213e;
--bg-tertiary: #0f3460;
--accent: #e94560;
--text-primary: #eee;
--text-secondary: #aaa;
--border: #333;
--user-msg-bg: #0f3460;
--assistant-msg-bg: #1a1a2e;
}
* { margin: 0; padding: 0; box-sizing: border-box; }
body {
font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', sans-serif;
background: var(--bg-primary);
color: var(--text-primary);
height: 100vh;
overflow: hidden;
}
#app {
display: flex;
height: 100vh;
}
/* === 侧边栏 === */
#sidebar {
width: 280px;
background: var(--bg-secondary);
border-right: 1px solid var(--border);
display: flex;
flex-direction: column;
flex-shrink: 0;
}
#sidebar-header {
display: flex;
align-items: center;
justify-content: space-between;
padding: 16px;
border-bottom: 1px solid var(--border);
}
#sidebar-header h2 { font-size: 16px; }
#new-chat-btn {
width: 36px; height: 36px;
border: 2px solid var(--accent);
background: transparent;
color: var(--accent);
border-radius: 8px;
font-size: 20px;
cursor: pointer;
display: flex;
align-items: center;
justify-content: center;
}
#new-chat-btn:hover { background: var(--accent); color: white; }
#session-list {
flex: 1;
overflow-y: auto;
padding: 8px;
}
.session-item {
padding: 12px;
margin-bottom: 4px;
border-radius: 8px;
cursor: pointer;
display: flex;
align-items: center;
gap: 8px;
transition: background 0.2s;
}
.session-item:hover { background: var(--bg-tertiary); }
.session-item.active { background: var(--bg-tertiary); border-left: 3px solid var(--accent); }
.session-title {
flex: 1;
font-size: 14px;
white-space: nowrap;
overflow: hidden;
text-overflow: ellipsis;
}
.session-info { font-size: 11px; color: var(--text-secondary); }
.session-delete { background: none; border: none; cursor: pointer; opacity: 0; }
.session-item:hover .session-delete { opacity: 1; }
/* === 主区域 === */
#main {
flex: 1;
display: flex;
flex-direction: column;
min-width: 0;
}
#chat-container {
flex: 1;
overflow-y: auto;
padding: 16px;
}
.welcome-message {
text-align: center;
padding: 60px 20px;
}
.welcome-message h1 { font-size: 32px; margin-bottom: 16px; }
.welcome-message p { color: var(--text-secondary); margin-bottom: 12px; }
.suggestions {
display: flex;
flex-wrap: wrap;
justify-content: center;
gap: 8px;
margin-top: 16px;
}
.suggestions button {
padding: 8px 16px;
border: 1px solid var(--border);
background: var(--bg-secondary);
color: var(--text-primary);
border-radius: 20px;
cursor: pointer;
font-size: 14px;
transition: all 0.2s;
}
.suggestions button:hover {
background: var(--accent);
border-color: var(--accent);
}
/* === 消息 === */
.message {
max-width: 800px;
margin: 16px auto;
padding: 12px 16px;
border-radius: 12px;
animation: fadeIn 0.3s ease;
}
@keyframes fadeIn { from { opacity: 0; transform: translateY(10px); } to { opacity: 1; transform: translateY(0); } }
.message.user { background: var(--user-msg-bg); }
.message.assistant { background: var(--assistant-msg-bg); border: 1px solid var(--border); }
.message.system { color: var(--accent); text-align: center; font-size: 14px; }
.message-role { font-size: 18px; margin-bottom: 8px; }
.message-content {
line-height: 1.6;
font-size: 15px;
word-wrap: break-word;
}
.message-content code {
background: #333;
padding: 2px 6px;
border-radius: 4px;
font-size: 13px;
}
.message-content strong { color: var(--accent); }
.message-content .icon { font-size: 20px; }
.message-time { font-size: 11px; color: var(--text-secondary); margin-top: 8px; }
/* === 加载动画 === */
.loading .dot {
display: inline-block;
width: 8px; height: 8px;
background: var(--accent);
border-radius: 50%;
margin: 0 4px;
animation: bounce 1.4s infinite ease-in-out both;
}
.loading .dot:nth-child(1) { animation-delay: -0.32s; }
.loading .dot:nth-child(2) { animation-delay: -0.16s; }
.loading .dot:nth-child(3) { animation-delay: 0s; }
@keyframes bounce {
0%, 80%, 100% { transform: scale(0); }
40% { transform: scale(1); }
}
/* === 输入区域 === */
#input-area {
padding: 16px;
border-top: 1px solid var(--border);
}
#input-container {
max-width: 800px;
margin: 0 auto;
display: flex;
gap: 8px;
background: var(--bg-secondary);
border: 1px solid var(--border);
border-radius: 12px;
padding: 8px;
}
#message-input {
flex: 1;
background: transparent;
border: none;
color: var(--text-primary);
font-size: 15px;
resize: none;
outline: none;
padding: 8px;
line-height: 1.5;
}
#message-input::placeholder { color: #666; }
#send-btn {
width: 44px; height: 44px;
background: var(--accent);
border: none;
border-radius: 10px;
color: white;
cursor: pointer;
display: flex;
align-items: center;
justify-content: center;
flex-shrink: 0;
}
#send-btn:disabled { opacity: 0.5; cursor: not-allowed; }
/* === 推理面板 === */
#debug-panel {
width: 360px;
background: var(--bg-secondary);
border-left: 1px solid var(--border);
display: flex;
flex-direction: column;
flex-shrink: 0;
transition: width 0.3s;
}
#debug-panel.collapsed { width: 0; overflow: hidden; }
#debug-header {
display: flex;
justify-content: space-between;
align-items: center;
padding: 16px;
border-bottom: 1px solid var(--border);
}
#debug-header h3 { font-size: 14px; }
#debug-header button { background: none; border: none; color: var(--text-secondary); cursor: pointer; font-size: 20px; }
#debug-content {
flex: 1;
overflow-y: auto;
padding: 16px;
font-size: 13px;
}
.debug-section { margin-bottom: 16px; }
.debug-section h4 { margin-bottom: 8px; color: var(--accent); }
.debug-section li { margin: 4px 0; color: var(--text-secondary); }
.tool-call {
background: var(--bg-primary);
padding: 8px;
border-radius: 6px;
margin: 8px 0;
}
.tool-call pre {
background: #111;
color: #0f0;
padding: 6px;
border-radius: 4px;
font-size: 11px;
margin: 4px 0;
white-space: pre-wrap;
word-wrap: break-word;
}
.tool-result {
font-size: 11px;
color: var(--text-secondary);
max-height: 100px;
overflow-y: auto;
}
/* === 响应式 === */
@media (max-width: 768px) {
#sidebar { display: none; }
#debug-panel { display: none; }
}第六步:部署
6.1 环境配置
bash
# .env
LLM_MODEL=gpt-4o-mini # 使用的 LLM 模型
OPENAI_API_KEY=sk-xxx # API 密钥
OPENAI_BASE_URL= # 如果使用本地模型(如 Ollama: http://localhost:11434/v1)
WIKI_DIR=/data/code/wiki/learn # Wiki 文档目录
MAX_STEPS=5 # Agent 最大推理步数
TEMPERATURE=0.1 # 模型温度6.2 Docker 部署
dockerfile
# Dockerfile
FROM python:3.11-slim
WORKDIR /app
# 安装系统依赖
RUN apt-get update && apt-get install -y \
gcc g++ \
&& rm -rf /var/lib/apt/lists/*
# 安装 Python 依赖
COPY requirements.txt .
RUN pip install --no-cache-dir -r requirements.txt
# 复制代码
COPY . .
# 构建索引(如果提供了 Wiki 数据)
# RUN python -m rag.indexer $WIKI_DIR
# 端口
EXPOSE 8000
# 启动
CMD ["uvicorn", "server.main:app", "--host", "0.0.0.0", "--port", "8000"]yaml
# docker-compose.yml
version: '3.8'
services:
wiki-agent:
build: .
ports:
- "8000:8000"
volumes:
- /data/code/wiki/learn:/data/wiki:ro # Wiki 源文件(只读)
- ./data:/app/data # 索引和数据库持久化
environment:
- LLM_MODEL=gpt-4o-mini
- OPENAI_API_KEY=${OPENAI_API_KEY}
- WIKI_DIR=/data/wiki
- MAX_STEPS=5
- TEMPERATURE=0.1
restart: unless-stopped6.3 使用 Ollama 本地模型
bash
# 1. 安装 Ollama
curl -fsSL https://ollama.com/install.sh | sh
# 2. 拉取模型(推荐 qwen2.5 或 llama3.1)
ollama pull qwen2.5:7b
# 3. 修改 .env
LLM_MODEL=qwen2.5:7b
OPENAI_API_KEY=not-needed
OPENAI_BASE_URL=http://localhost:11434/v1第七步:运行步骤
完整运行流程
bash
# === 1. 克隆项目(假设已放在 wiki-agent 目录) ===
cd wiki-agent
# === 2. 创建虚拟环境 ===
python3 -m venv venv
source venv/bin/activate
# === 3. 安装依赖 ===
pip install -r requirements.txt
# === 4. 配置环境变量 ===
cp .env.example .env
# 编辑 .env,填入 OPENAI_API_KEY 或设置本地模型
# === 5. 构建 Wiki 索引 ===
python -m rag.indexer /data/code/wiki/learn
# === 6. 启动服务 ===
python -m server.main
# === 7. 打开浏览器 ===
# 访问 http://localhost:8000预期效果
mermaid
sequenceDiagram
participant User as 👤 用户
participant UI as 🌐 Web UI
participant API as ⚡ FastAPI
participant Agent as 🤖 LangGraph Agent
participant LLM as 🧠 LLM
participant RAG as 🗄️ RAG 引擎
User->>UI: "Go 的 GMP 调度模型是什么?"
UI->>API: POST /api/chat
API->>Agent: run(question)
Agent->>LLM: [System Prompt + 用户问题]
LLM-->>Agent: Thought: 需要搜索知识库<br/>Action: wiki_search("GMP 调度模型")
Agent->>RAG: wiki_search("GMP 调度模型")
RAG-->>Agent: 📄 找到 3 个相关文档<br/>(GMP调度模型.md, goroutine.md, ...)
Agent->>LLM: [搜索结果 + 原始问题]
LLM-->>Agent: Thought: 获取完整文档<br/>Action: wiki_get_page("GMP 调度模型")
Agent-->>Agent: 读取完整文档内容
Agent->>LLM: [完整文档 + 问题]
LLM-->>Agent: FINISH: "GMP 调度模型是 Go 运行时的核心..."
Agent-->>API: 最终回答
API-->>UI: ChatResponse
UI-->>User: 显示回答 + 推理过程性能指标参考
| 指标 | 小模型 (qwen2.5:7b) | 大模型 (gpt-4o-mini) |
|---|---|---|
| 单次检索延迟 | 50-100ms | 50-100ms |
| LLM 推理延迟 | 2-5s | 1-3s |
| Agent 总耗时 | 5-15s | 3-8s |
| 召回率 (Recall@5) | 0.82 | 0.82 |
| 回答准确率 | 0.75 | 0.88 |
说明:检索部分(RAG)的性能取决于索引质量和 Embedding 模型选择,不随 LLM 变化。LLM 部分影响回答质量和整体延迟。使用 qwen2.5:7b + Ollama 可以完全离线运行。
后续优化方向
- 流式输出:WebSocket 实时展示 Agent 思维过程,当前已实现基础框架
- 多模态支持:支持检索 Wiki 中的图片和代码块
- 对话摘要:长对话自动压缩,防止上下文溢出
- 工具扩展:添加代码执行、图表生成等更多工具
- 权限控制:用户认证、会话隔离
- 监控告警:Agent 调用延迟、召回率监控
第八步:测试与验证
8.1 检索准确率测试
python
"""
RAG 检索系统的回归测试
每次修改分块策略或 Embedding 模型后,运行此测试确保召回率不退化
"""
import json
from dataclasses import dataclass
from typing import List, Dict
@dataclass
class RetrievalTestCase:
"""检索测试用例"""
query: str # 查询
expected_doc_keywords: List[str] # 期望召回的文档应该包含的关键词
min_recall: float = 0.0 # 最低 Recall 要求
# ===== 标准测试用例 =====
RETRIEVAL_TESTS = [
RetrievalTestCase(
query="Go 的 GMP 调度模型是什么?",
expected_doc_keywords=["GMP", "goroutine", "Processor", "Machine"],
min_recall=1, # 所有关键词至少命中 1 个
),
RetrievalTestCase(
query="slice 和 array 的区别",
expected_doc_keywords=["slice", "array", "底层", "扩容"],
min_recall=0.5,
),
RetrievalTestCase(
query="MySQL 索引优化",
expected_doc_keywords=["B+Tree", "索引", "explain"],
min_recall=0.5,
),
]
def test_retrieval_quality(retriever, test_cases: List[RetrievalTestCase]):
"""测试检索质量"""
results = {"total": len(test_cases), "passed": 0, "details": []}
for case in test_cases:
docs = retriever.search(case.query, top_k=5)
# 将所有召回的文档拼接
all_text = " ".join(d.text for d in docs)
# 检查关键词命中率
hits = sum(1 for kw in case.expected_doc_keywords
if kw.lower() in all_text.lower())
recall = hits / len(case.expected_doc_keywords)
passed = recall >= case.min_recall
if passed:
results["passed"] += 1
results["details"].append({
"query": case.query[:50],
"recall": f"{recall:.0%}",
"passed": passed,
"top_doc": docs[0].title if docs else "N/A",
})
results["pass_rate"] = results["passed"] / results["total"]
return results
# 运行: python -m pytest tests/test_retrieval.py8.2 Agent 工具调用测试
python
"""
Agent 回归测试:验证工具调用正确性
"""
import pytest
class AgentTestCase:
"""Agent 测试用例"""
def __init__(self, question: str, expected_tool: str,
expected_keywords: List[str]):
self.question = question
self.expected_tool = expected_tool # 期望调用的工具名
self.expected_keywords = expected_keywords # 最终回答应包含的关键词
# ===== 标准 Agent 测试用例 =====
AGENT_TESTS = [
AgentTestCase(
question="Go 语言里什么是 goroutine?",
expected_tool="wiki_search",
expected_keywords=["goroutine", "协程", "轻量级"],
),
AgentTestCase(
question="帮我写一个 Python 的快速排序",
expected_tool="python_execute",
expected_keywords=["def quick_sort", "return"],
),
AgentTestCase(
question="你好,你叫什么名字?",
expected_tool="", # 不期望调用任何工具
expected_keywords=["Agent", "助手"],
),
]
async def test_agent_pipeline(agent, test_cases: List[AgentTestCase]):
"""测试 Agent 完整流程"""
results = {"total": len(test_cases), "passed": 0, "details": []}
for case in test_cases:
response = await agent.run(case.question)
# 检查工具调用
tool_called = any(case.expected_tool in step.get("action", "")
for step in response.steps)
# 检查答案内容
answer_ok = all(kw.lower() in response.answer.lower()
for kw in case.expected_keywords)
passed = answer_ok
if case.expected_tool:
passed = passed and tool_called
if passed:
results["passed"] += 1
results["details"].append({
"question": case.question[:40],
"tool_called": tool_called,
"answer_ok": answer_ok,
"passed": passed,
})
results["pass_rate"] = results["passed"] / results["total"]
return results8.3 API 冒烟测试
python
"""
FastAPI 接口冒烟测试
"""
from fastapi.testclient import TestClient
from server.main import app
client = TestClient(app)
def test_health_check():
"""健康检查"""
response = client.get("/health")
assert response.status_code == 200
assert response.json()["status"] == "ok"
def test_chat_basic():
"""基础聊天"""
response = client.post("/api/chat", json={
"question": "什么是 goroutine?",
"session_id": "test-session-001",
})
assert response.status_code == 200
data = response.json()
assert "answer" in data, f"缺少 answer 字段: {data}"
assert len(data["answer"]) > 0, "回答不能为空"
def test_chat_empty_question():
"""空问题应返回错误"""
response = client.post("/api/chat", json={
"question": "",
"session_id": "test-empty",
})
assert response.status_code == 422 # 参数校验
def test_sessions_list():
"""会话列表"""
response = client.get("/api/sessions")
assert response.status_code == 200
assert isinstance(response.json(), list)
def test_404():
"""不存在的路由"""
response = client.get("/api/nonexistent")
assert response.status_code == 404
# 运行: pytest tests/test_api.py -v第九步:增量索引与知识更新
9.1 文件变更检测
python
"""
监听 Wiki 文件变更,实现增量索引更新
策略:
1. 启动时全量构建索引
2. 运行时用文件监控器检测变更
3. 只重建变更的文件,不影响其他索引
"""
import os
import hashlib
from pathlib import Path
from watchdog.observers import Observer
from watchdog.events import FileSystemEventHandler
class WikiFileWatcher(FileSystemEventHandler):
"""Wiki 文件变更监控器"""
def __init__(self, indexer, wiki_dir: str):
self.indexer = indexer
self.wiki_dir = wiki_dir
self.file_hashes = {} # 文件路径 → 内容哈希
def on_modified(self, event):
"""文件修改时触发"""
if event.is_directory:
return
if not event.src_path.endswith('.md'):
return
self._handle_change(event.src_path, "modified")
def on_created(self, event):
"""新文件创建时触发"""
if not event.src_path.endswith('.md'):
return
self._handle_change(event.src_path, "created")
def on_deleted(self, event):
"""文件删除时触发"""
if not event.src_path.endswith('.md'):
return
self._handle_deletion(event.src_path)
def _should_process(self, filepath: str) -> bool:
"""判断是否需要处理(去重:只处理内容真正变了的文件)"""
try:
with open(filepath, 'r') as f:
content = f.read()
new_hash = hashlib.md5(content.encode()).hexdigest()
old_hash = self.file_hashes.get(filepath)
if old_hash == new_hash:
return False # 内容未变,跳过
self.file_hashes[filepath] = new_hash
return True
except Exception:
return False
def _handle_change(self, filepath: str, action: str):
"""处理文件变更:重新索引该文件"""
if not self._should_process(filepath):
return
rel_path = os.path.relpath(filepath, self.wiki_dir)
print(f"📝 [{action}] {rel_path}")
# 1. 删除旧索引
self.indexer.remove_document(rel_path)
# 2. 重新分块 + 向量化 + 入库
with open(filepath, 'r') as f:
content = f.read()
chunks = self.indexer.chunker.chunk(content, source=rel_path)
self.indexer.index_chunks(chunks)
print(f" ✅ 重新索引完成 ({len(chunks)} chunks)")
def _handle_deletion(self, filepath: str):
"""处理文件删除"""
rel_path = os.path.relpath(filepath, self.wiki_dir)
print(f"🗑️ [deleted] {rel_path}")
self.indexer.remove_document(rel_path)
def build_initial_hashes(self):
"""启动时构建所有文件的初始哈希"""
for md_file in Path(self.wiki_dir).rglob("*.md"):
try:
with open(md_file, 'r') as f:
content = f.read()
self.file_hashes[str(md_file)] = hashlib.md5(
content.encode()).hexdigest()
except Exception:
pass
# ===== 启动文件监控 =====
def start_wiki_watcher(indexer, wiki_dir: str):
"""启动 Wiki 文件监控"""
handler = WikiFileWatcher(indexer, wiki_dir)
handler.build_initial_hashes()
observer = Observer()
observer.schedule(handler, wiki_dir, recursive=True)
observer.start()
print(f"👀 开始监控 {wiki_dir} 的文件变更...")
return observer9.2 索引版本管理与回滚
python
"""
索引版本管理:支持版本切换和回滚
"""
import json
import shutil
from datetime import datetime
from pathlib import Path
class IndexVersionManager:
"""索引版本管理器"""
def __init__(self, index_dir: str, max_versions: int = 5):
self.index_dir = Path(index_dir)
self.versions_dir = self.index_dir / "versions"
self.versions_dir.mkdir(parents=True, exist_ok=True)
self.max_versions = max_versions
self.current_version_file = self.index_dir / "current_version.json"
def create_version(self, description: str = "") -> str:
"""创建当前索引的快照"""
version_id = datetime.now().strftime("%Y%m%d_%H%M%S")
version_path = self.versions_dir / version_id
version_path.mkdir(parents=True)
# 复制当前索引数据
for item in self.index_dir.iterdir():
if item.name != "versions" and item.name != "current_version.json":
if item.is_dir():
shutil.copytree(item, version_path / item.name)
else:
shutil.copy2(item, version_path / item.name)
# 记录版本信息
info = {
"version_id": version_id,
"created_at": datetime.now().isoformat(),
"description": description,
"file_count": sum(1 for _ in version_path.rglob("*") if _.is_file()),
}
with open(version_path / "version_info.json", 'w') as f:
json.dump(info, f, ensure_ascii=False, indent=2)
# 清理旧版本
self._cleanup_old_versions()
print(f"📸 索引快照已创建: {version_id}")
return version_id
def rollback(self, version_id: str) -> bool:
"""回滚到指定版本"""
version_path = self.versions_dir / version_id
if not version_path.exists():
print(f"❌ 版本 {version_id} 不存在")
return False
# 1. 清空当前索引
for item in self.index_dir.iterdir():
if item.name != "versions" and item.name != "current_version.json":
if item.is_dir():
shutil.rmtree(item)
else:
item.unlink()
# 2. 从快照恢复
for item in version_path.iterdir():
if item.name != "version_info.json":
target = self.index_dir / item.name
if item.is_dir():
shutil.copytree(item, target)
else:
shutil.copy2(item, target)
print(f"⏪ 已回滚到版本: {version_id}")
return True
def list_versions(self) -> list:
"""列出所有可用版本"""
versions = []
for v in sorted(self.versions_dir.iterdir(), reverse=True):
info_file = v / "version_info.json"
if info_file.exists():
with open(info_file) as f:
versions.append(json.load(f))
return versions
def _cleanup_old_versions(self):
"""只保留最近 N 个版本"""
all_versions = sorted(self.versions_dir.iterdir())
if len(all_versions) > self.max_versions:
for old in all_versions[:-self.max_versions]:
shutil.rmtree(old)
print(f"🗑️ 清理旧版本: {old.name}")第十步:可观测性与安全
10.1 日志与监控
python
"""
Agent 调用全链路追踪
"""
import time
import logging
from contextvars import ContextVar
from typing import Dict, Any
# 请求级别的 trace ID
trace_id_var: ContextVar[str] = ContextVar("trace_id", default="")
# 配置日志格式
logging.basicConfig(
format="%(asctime)s [%(levelname)s] [trace=%(trace_id)s] %(message)s",
level=logging.INFO,
)
logger = logging.getLogger("wiki-agent")
class AgentTracer:
"""Agent 调用追踪器"""
def __init__(self):
self.traces: Dict[str, dict] = {}
def start_trace(self, session_id: str, question: str) -> str:
"""开始一次追踪"""
trace_id = f"{session_id}-{int(time.time() * 1000)}"
trace_id_var.set(trace_id)
self.traces[trace_id] = {
"session_id": session_id,
"question": question,
"start_time": time.time(),
"steps": [],
"tool_calls": [],
"total_tokens": 0,
"error": None,
}
return trace_id
def record_step(self, step: dict):
"""记录 Agent 的一个推理步骤"""
trace_id = trace_id_var.get()
if trace_id in self.traces:
self.traces[trace_id]["steps"].append({
"thought": step.get("thought", ""),
"action": step.get("action", ""),
"time": time.time(),
})
def record_tool_call(self, tool_name: str, args: dict,
result: Any, duration_ms: float):
"""记录工具调用"""
trace_id = trace_id_var.get()
if trace_id in self.traces:
self.traces[trace_id]["tool_calls"].append({
"tool": tool_name,
"args": str(args)[:200],
"result_preview": str(result)[:200],
"duration_ms": duration_ms,
})
logger.info(f"Tool: {tool_name}({str(args)[:80]}) → {duration_ms:.0f}ms")
def finish_trace(self, answer: str, error: str = None):
"""结束追踪"""
trace_id = trace_id_var.get()
if trace_id in self.traces:
trace = self.traces[trace_id]
trace["end_time"] = time.time()
trace["duration_ms"] = (trace["end_time"] - trace["start_time"]) * 1000
trace["answer"] = answer[:500]
trace["error"] = error
trace["num_steps"] = len(trace["steps"])
trace["num_tool_calls"] = len(trace["tool_calls"])
logger.info(
f"Agent finished: {trace['num_steps']} steps, "
f"{trace['num_tool_calls']} tools, "
f"{trace['duration_ms']:.0f}ms"
)
# 清理(实际项目中应持久化到数据库)
# del self.traces[trace_id]10.2 安全防护
python
"""
Wiki Agent 安全保护层
"""
import re
import time
from collections import defaultdict
from fastapi import Request, HTTPException
class RateLimiter:
"""简易速率限制器"""
def __init__(self, max_requests: int = 30, window_seconds: int = 60):
self.max_requests = max_requests
self.window_seconds = window_seconds
self.requests: dict = defaultdict(list) # ip → [timestamps]
def check(self, client_ip: str) -> bool:
"""检查是否允许请求"""
now = time.time()
# 清理过期记录
self.requests[client_ip] = [
t for t in self.requests[client_ip]
if now - t < self.window_seconds
]
if len(self.requests[client_ip]) >= self.max_requests:
return False # 速率超限
self.requests[client_ip].append(now)
return True
# 全局限制器
rate_limiter = RateLimiter(max_requests=30, window_seconds=60)
# 在 FastAPI 路由中使用:
# @app.post("/api/chat")
# async def chat(request: Request, body: ChatRequest):
# client_ip = request.client.host
# if not rate_limiter.check(client_ip):
# raise HTTPException(status_code=429, detail="请求太频繁,请稍后再试")
class InputSanitizer:
"""输入安全过滤"""
# Prompt Injection 常见模式
INJECTION_PATTERNS = [
r"忽略.*指令",
r"ignore.*instruction",
r"system.*prompt",
r"你.*不是.*助手",
r"forget.*previous",
r"现在.*你.*是",
r"不要.*说.*不",
r"DAN\s*mode",
]
# 危险操作模式
DANGEROUS_PATTERNS = [
r"rm\s+-rf",
r"os\.system",
r"subprocess",
r"eval\s*\(",
r"exec\s*\(",
r"__import__",
r"delete.*database",
r"drop\s+table",
]
def check_prompt_injection(self, user_input: str) -> list:
"""检测 Prompt Injection 攻击"""
detected = []
for pattern in self.INJECTION_PATTERNS:
if re.search(pattern, user_input, re.IGNORECASE):
detected.append(pattern)
return detected
def check_dangerous_commands(self, user_input: str) -> list:
"""检测危险操作请求"""
detected = []
for pattern in self.DANGEROUS_PATTERNS:
if re.search(pattern, user_input, re.IGNORECASE):
detected.append(pattern)
return detected
def sanitize(self, user_input: str) -> tuple:
"""
安全过滤入口
返回: (is_safe, warnings)
"""
warnings = []
injections = self.check_prompt_injection(user_input)
if injections:
warnings.append(f"检测到 Prompt Injection 尝试: {injections}")
dangerous = self.check_dangerous_commands(user_input)
if dangerous:
warnings.append(f"检测到危险操作请求: {dangerous}")
# 长度限制
if len(user_input) > 10000:
warnings.append("输入过长 (>10000字符)")
is_safe = len(warnings) == 0
return is_safe, warnings
sanitizer = InputSanitizer()10.3 工具调用白名单
python
"""
Agent 工具调用的安全白名单
"""
class ToolGuard:
"""工具调用守卫 — 限制 Agent 的破坏性操作"""
# 允许的工具
ALLOWED_TOOLS = {
"wiki_search", # 只读检索
"wiki_get_page", # 只读页面
"python_execute", # 代码执行(应放在沙箱中)
"final_answer", # 结束
}
# 即使允许,也需要限制参数的工具
RESTRICTED_ARGS = {
"python_execute": {
"max_code_length": 2000, # 代码最长 2000 字符
"timeout_seconds": 10, # 最多执行 10 秒
"forbidden_imports": [ # 禁止的模块
"os", "subprocess", "shutil", "socket",
"requests", "http", "ftp",
],
},
}
def validate_tool_call(self, tool_name: str, tool_args: dict) -> tuple:
"""验证工具调用是否安全
Returns:
(is_allowed, reason)
"""
# 1. 工具名白名单
if tool_name not in self.ALLOWED_TOOLS:
return False, f"工具 '{tool_name}' 不在白名单中"
# 2. 参数限制
if tool_name in self.RESTRICTED_ARGS:
restrictions = self.RESTRICTED_ARGS[tool_name]
# 检查代码长度
if "code" in tool_args:
if len(tool_args["code"]) > restrictions.get("max_code_length", 2000):
return False, f"代码长度 {len(tool_args['code'])} 超过限制"
# 检查禁止的 import
for forbidden in restrictions.get("forbidden_imports", []):
if forbidden in tool_args["code"]:
return False, f"代码包含禁止模块: {forbidden}"
return True, "OK"
# 在 Agent 执行工具前调用
guard = ToolGuard()
# if not guard.validate_tool_call(tool_name, args):
# return {"error": "工具调用被安全策略拒绝"}
登录后即可发表评论 👇