RAG工具升级:WebAPI与Excel导出
·
本篇对上一篇进行 基于 CLI 的 RAG 工具改进,原版能从 Excel 构建 FAISS 向量库,并结合 DeepSeek 做中文问答。但在实际应用中,这种方式存在一些局限:
-
仅支持单用户交互,无法并发访问
-
没有会话管理,无法保存历史对话
-
检索结果无法直接导出为 Excel
-
数据库管理缺失,不易持久化
1、Web API 接口化
-
使用 FastAPI 替代 CLI,提供
/chat/askPOST 接口 -
支持多用户通过
session_id区分会话 -
可直接集成 Web、移动端或其他系统
@app.post("/chat/ask")
async def ask(request: Request, data: ChatRequest):
session_id = data.session_id or str(uuid.uuid4())
result = retrieve_and_answer(
data.message,
request.app.state.vector_store,
request.app.state.embedding,
messages,
top_k=5
)
return {"session_id": session_id, "answer": result["answer"]}

2、会话管理与历史消息存储
-
引入 PostgreSQL + SQLAlchemy 保存会话 (
sessions) 与聊天记录 (chat) -
支持三类消息
role:user、ai、retrieved -
会话历史可用于构建 RAG Prompt,支持多轮对话
storage.store_message(session_id=session_id, role="retrieved", content=json.dumps(retrieved, ensure_ascii=False))
history = storage.get_conversation_history(session_id=session_id, limit=1000)
3、 检索结果 Excel 导出功能
用户输入包含“表格”时,自动:
-
收集当前会话所有检索文档序号
-
从原始 Excel 抽取整行数据
-
生成新 Excel 文件并返回下载
out_df = raw_df[raw_df["序号"].isin(seqs)].copy()
output_path = os.path.join(EXPORT_DIR, f"export_{session_id}.xlsx")
out_df.to_excel(output_path, index=False)
return FileResponse(
output_path,
filename=f"查询结果_{session_id}.xlsx",
media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet"
)
4、启动优化与资源预加载
在 startup_event 中统一加载:
-
Embedding 模型
-
向量库
-
原始 Excel
-
消息存储
@app.on_event("startup")
def startup_event():
app.state.embedding = load_embedding()
app.state.vector_store = load_vector_store(index_dir, app.state.embedding)
app.state.raw_df = pd.read_excel(raw_excel_path, dtype=str).fillna("")
app.state.storage = MessageStorage()
5、RAG 工具函数增强
-
retrieve_and_answer支持返回检索文档 metadata -
异常处理完善,增加调试信息
-
Prompt 构建支持历史会话信息
retrieved_docs = [{"序号": d.metadata.get("id"), "内容": d.page_content} for d in docs]
payload = {"model": DEESEEK_MODEL_NAME, "messages": messages + [{"role": "user", "content": rag_prompt}]}
resp = requests.post(DEESEEK_API_URL, json=payload, headers=headers, timeout=30)
answer = resp.json().get("choices", [{}])[0].get("message", {}).get("content", "")
6、数据库初始化脚本
-
create_db_and_tables.py自动创建 PostgreSQL 数据库和表 -
支持零配置部署
def create_database_if_not_exists():
cur.execute(f"SELECT 1 FROM pg_database WHERE datname = '{TARGET_DB}';")
if not cur.fetchone():
cur.execute(f'CREATE DATABASE "{TARGET_DB}";')
def create_tables():
engine = create_engine(TARGET_DB_URL)
Base.metadata.create_all(bind=engine)
7、完整代码
# 文件: create_db_and_tables.py
# 功能: 初始化 PostgreSQL 数据库和创建 sessions, chat 表
# 说明: 如果数据库或表不存在,则创建
import psycopg2
from psycopg2.extensions import ISOLATION_LEVEL_AUTOCOMMIT
from sqlalchemy import create_engine
from models import Base
# PostgreSQL 配置
PG_USER = "postgres"
PG_PASSWORD = "12345"
PG_HOST = "localhost"
PG_PORT = "5432"
TARGET_DB = "heritage_rag"
DEFAULT_DB_URL = f"postgresql://{PG_USER}:{PG_PASSWORD}@{PG_HOST}:{PG_PORT}/postgres"
TARGET_DB_URL = f"postgresql://{PG_USER}:{PG_PASSWORD}@{PG_HOST}:{PG_PORT}/{TARGET_DB}"
def create_database_if_not_exists():
"""连接 postgres 数据库,检查并创建目标数据库"""
try:
conn = psycopg2.connect(dbname="postgres", user=PG_USER, password=PG_PASSWORD, host=PG_HOST, port=PG_PORT)
conn.set_isolation_level(ISOLATION_LEVEL_AUTOCOMMIT)
cur = conn.cursor()
# 查询数据库是否存在
cur.execute(f"SELECT 1 FROM pg_database WHERE datname = '{TARGET_DB}';")
exists = cur.fetchone()
if not exists:
print(f"📦 未找到数据库 '{TARGET_DB}',正在创建...")
cur.execute(f'CREATE DATABASE "{TARGET_DB}";')
print(f"✅ 数据库 '{TARGET_DB}' 创建成功!")
else:
print(f"✔ 数据库 '{TARGET_DB}' 已存在。")
cur.close()
conn.close()
except Exception as e:
print(f"❌ 创建数据库时出错: {e}")
raise
def create_tables():
"""使用 SQLAlchemy 创建表"""
try:
print("📌 正在连接数据库并创建表...")
engine = create_engine(TARGET_DB_URL)
Base.metadata.create_all(bind=engine)
print("✅ 数据表创建完成(sessions, chat)")
except Exception as e:
print(f"❌ 创建数据表时出错: {e}")
raise
if __name__ == "__main__":
print("🚀 开始初始化数据库与数据表")
create_database_if_not_exists()
create_tables()
print("🎉 初始化成功!")
# 文件: build_index.py
# 功能: 从 Excel 构建 FAISS 向量索引(文档 metadata 中包含原始“序号”)
import os
import pandas as pd
from langchain_community.vectorstores import FAISS
from langchain_core.documents import Document
from rag_utils import load_embedding
EXCEL_FILE = "四普数据.xlsx"
INDEX_DIR = "data/fourth_survey_index"
# Excel 表头列名(用于构造 page_content)
COLUMNS = ["名称", "编号", "调查结果", "省", "市", "县", "调查人", "调查日期", "审定人", "审定日期",
"抽查人", "抽查日期", "地址及位置", "是否整体迁移并在新迁址并在新迁址地域范围内",
"变更消失情况","高程海拔","是否已公布保护范围","是否公布建设控制地带","文物级别",
def build_index(excel_file=EXCEL_FILE, index_dir=INDEX_DIR):
"""构建 FAISS 向量库"""
if not os.path.exists(excel_file):
raise FileNotFoundError(f"找不到 Excel 文件: {excel_file}")
# 读取 Excel 数据,缺失值填空
df = pd.read_excel(excel_file, dtype=str).fillna("")
docs = []
for _, row in df.iterrows():
seq = str(row.get("序号", "")).strip()
if seq == "":
continue
try:
doc_id = int(float(seq))
except Exception:
doc_id = seq
# 将每列都加入内容,便于检索时快速对应
parts = []
for col in COLUMNS:
val = row.get(col, "")
if pd.isna(val) or val == "":
continue
parts.append(f"{col}:{val}")
page_content = "\n".join(parts) if parts else row.get("简介", "")
docs.append(Document(page_content=page_content, metadata={"id": doc_id}))
# 加载向量化模型
embedding = load_embedding()
os.makedirs(index_dir, exist_ok=True)
vector_store = FAISS.from_documents(docs, embedding)
vector_store.save_local(index_dir)
print("🎉 向量库已生成:", index_dir)
if __name__ == "__main__":
build_index()
# 文件: models.py
# 功能: 定义 PostgreSQL 数据库模型(sessions 和 chat 表),并提供消息存储和读取功能
from sqlalchemy import create_engine, Column, Integer, String, Text, DateTime, ForeignKey, func
from sqlalchemy.orm import declarative_base, sessionmaker, relationship
# PostgreSQL 数据库连接(如需修改请在此处替换)
DATABASE_URL = "postgresql://postgres:12345@localhost:5432/heritage_rag"
engine = create_engine(DATABASE_URL, echo=False)
SessionLocal = sessionmaker(bind=engine, autoflush=False, autocommit=False)
Base = declarative_base()
class SessionModel(Base):
"""会话表模型"""
__tablename__ = "sessions"
session_id = Column(String(255), primary_key=True, index=True)
created_at = Column(DateTime, server_default=func.now())
# 一个会话可以有多条聊天记录
chats = relationship("Chat", back_populates="session", cascade="all, delete-orphan")
class Chat(Base):
"""聊天记录表模型"""
__tablename__ = "chat"
id = Column(Integer, primary_key=True, index=True)
session_id = Column(String(255), ForeignKey("sessions.session_id"))
role = Column(String(20), nullable=False) # 'user' / 'ai' / 'retrieved'
content = Column(Text, nullable=False)
replyed_time = Column(DateTime, nullable=True)
created_at = Column(DateTime, server_default=func.now())
session = relationship("SessionModel", back_populates="chats")
class MessageStorage:
"""封装消息存储与读取操作"""
def __init__(self, session_factory=SessionLocal):
self.session_factory = session_factory
def store_message(self, session_id: str, role: str, content: str):
"""存储一条消息"""
db = self.session_factory()
try:
# 查询当前会话是否存在
sess = db.query(SessionModel).filter_by(session_id=session_id).first()
if not sess:
sess = SessionModel(session_id=session_id)
db.add(sess)
db.flush()
chat = Chat(session_id=session_id, role=role, content=content)
db.add(chat)
db.commit()
finally:
db.close()
def get_conversation_history(self, session_id: str, limit: int = 50):
"""获取会话历史消息(按插入顺序升序)"""
db = self.session_factory()
try:
rows = db.query(Chat).filter_by(session_id=session_id).order_by(Chat.id.asc()).limit(limit).all()
return [{"role": r.role, "content": r.content} for r in rows]
finally:
db.close()
# 文件: rag_utils.py
# 功能: RAG 工具函数(返回检索到的文档 metadata,以便导出对应序号)
import os
import requests
from langchain_community.vectorstores import FAISS
from langchain_community.embeddings import HuggingFaceEmbeddings
# 配置
SHIBING_MODEL = "shibing624/text2vec-base-chinese"
DEESEEK_API_KEY = "9b6b8bd0- -4e53-ae63-91f89dbbddb8"
DEESEEK_API_URL = "https://ark.cn-beijing.volces.com/api/v3/chat/completions"
DEESEEK_MODEL_NAME = "deepseek-v3-250324"
def load_embedding():
"""加载句向量模型"""
return HuggingFaceEmbeddings(model_name=SHIBING_MODEL)
def load_vector_store(index_dir, embedding):
"""加载本地向量库"""
if not os.path.exists(index_dir):
raise FileNotFoundError(f"向量库目录不存在: {index_dir}")
return FAISS.load_local(index_dir, embedding, allow_dangerous_deserialization=True)
def prepare_messages(storage, session_id, user_message, limit=50):
"""构建会话历史,映射 role 为 DeepSeek 支持值"""
history = storage.get_conversation_history(session_id, limit=limit)
role_map = {"user": "user", "ai": "assistant"}
return [{"role": role_map.get(h["role"], "user"), "content": h["content"]} for h in history]
def retrieve_and_answer(query, vector_store, embedding, messages, top_k=5):
"""检索 + DeepSeek 回答,增加调试和错误处理
返回结构:
{
"answer": "...",
"context": "...",
"retrieved_docs": [{"序号": ... , "内容": "..."} , ...]
}
"""
# === 1) 文档检索 ===
try:
docs = vector_store.similarity_search(query, k=top_k)
except Exception as e:
print("❌ 向量检索出错:", e)
docs = []
context = "\n".join(d.page_content for d in docs) if docs else "暂无相关文档。"
# 收集检索到的序号信息(来自 Document.metadata)
retrieved_docs = []
for d in docs:
meta_id = None
if isinstance(d, dict):
# 某些向量库可能返回 dict
meta = d.get("metadata", {}) or {}
meta_id = meta.get("id") or meta.get("序号")
page_content = d.get("page_content", "")
else:
meta = getattr(d, "metadata", {}) or {}
meta_id = meta.get("id") or meta.get("序号")
page_content = getattr(d, "page_content", "")
retrieved_docs.append({"序号": meta_id, "内容": page_content})
# === 2) 构建 RAG Prompt ===
rag_prompt = (
"以下是知识库检索到的内容,请结合这些内容回答用户问题。\n\n"
"【检索内容】\n"
f"{context}\n\n"
"【用户问题】\n"
f"{query}"
)
# === 3) 构建 payload ===
payload = {
"model": DEESEEK_MODEL_NAME,
"messages": messages + [{"role": "user", "content": rag_prompt}],
"temperature": 0.3,
"max_tokens": 512,
}
headers = {
"Authorization": f"Bearer {DEESEEK_API_KEY}",
"Content-Type": "application/json"
}
# === 4) 请求 DeepSeek,增加异常处理 ===
try:
print("=== DeepSeek 请求 payload ===")
print(payload)
resp = requests.post(DEESEEK_API_URL, json=payload, headers=headers, timeout=30)
resp.raise_for_status()
resp_json = resp.json()
print("=== DeepSeek 返回 ===")
print(resp_json)
answer = resp_json.get("choices", [{}])[0].get("message", {}).get("content", "") if isinstance(resp_json, dict) else ""
except requests.exceptions.HTTPError as e:
print("❌ DeepSeek API 请求失败:", e)
print("响应内容:", resp.text if 'resp' in locals() else "无响应")
answer = "服务端返回错误,无法回答问题"
except Exception as e:
print("❌ 请求 DeepSeek 出现异常:", e)
answer = "请求 DeepSeek 时出现异常"
return {"answer": answer, "context": context, "retrieved_docs": retrieved_docs}
# 文件: app.py
# 功能: FastAPI 接口服务
# 说明:
# - 提供 /chat/ask 接口进行 RAG 问答
# - 加载向量库和 Embedding 模型
# - 使用 MessageStorage 保存会话和历史记录,实现对话记忆
# 文件: app.py
# 功能: FastAPI 接口服务(含按序号从原 Excel 抽取整行并按需生成 Excel)
from fastapi import FastAPI, Request
from fastapi.responses import FileResponse
from pydantic import BaseModel
import uuid
import os
import json
import uvicorn
import pandas as pd
from rag_utils import load_embedding, load_vector_store, retrieve_and_answer, prepare_messages
from models import MessageStorage
app = FastAPI(
title="HeritageRAG",
description="文化遗产检索增强问答系统(含按序号导出表格)",
version="1.0.0",
)
EXPORT_DIR = "export"
os.makedirs(EXPORT_DIR, exist_ok=True)
@app.on_event("startup")
def startup_event():
"""应用启动事件: 加载模型、向量库、消息存储与原始 Excel"""
print("🔧 [startup] 跳过表创建,直接加载资源...")
print("🔍 [startup] 加载 Embedding 模型 ...")
app.state.embedding = load_embedding()
index_dir = os.environ.get("INDEX_DIR", "data/fourth_survey_index")
print(f"📦 [startup] 加载向量库({index_dir}) ...")
app.state.vector_store = load_vector_store(index_dir, app.state.embedding)
# 加载原始 Excel(用于按序号精确抽取)
raw_excel_path = os.environ.get("RAW_EXCEL", "四普数据.xlsx")
if not os.path.exists(raw_excel_path):
print(f"⚠️ 警告: 原始 Excel 未找到: {raw_excel_path},导出功能将不可用。")
app.state.raw_df = None
else:
print(f"📄 [startup] 加载原始 Excel ({raw_excel_path}) ...")
app.state.raw_df = pd.read_excel(raw_excel_path, dtype=str).fillna("")
# 标准化序号列为字符串形式,便于匹配
if "序号" in app.state.raw_df.columns:
app.state.raw_df["序号"] = app.state.raw_df["序号"].astype(str).str.strip()
print(f"✅ 原始 Excel 加载完成,共 {len(app.state.raw_df)} 条记录")
app.state.storage = MessageStorage()
print("✅ 启动完成,服务准备就绪。")
# 请求/响应模型
class ChatRequest(BaseModel):
message: str
session_id: str | None = None
history_limit: int | None = 50
class ChatResponse(BaseModel):
session_id: str
answer: str
context: str | None = None
@app.post("/chat/ask", response_model=ChatResponse)
async def ask(request: Request, data: ChatRequest):
"""处理问答请求
- 正常返回 RAG 回答
- 若检索有命中文档,会把命中 metadata(序号)以 role='retrieved' 存储
- 若用户 message 中包含 '表格',则汇总当前会话所有 retrieved 序号并生成 Excel 并直接返回文件响应
"""
storage: MessageStorage = request.app.state.storage
vector_store = request.app.state.vector_store
embedding = request.app.state.embedding
raw_df = request.app.state.raw_df
# 如果没有 session_id,则生成新的
session_id = data.session_id or str(uuid.uuid4())
# 存储用户问题
storage.store_message(session_id=session_id, role="user", content=data.message)
# 准备上下文历史
messages = prepare_messages(storage, session_id, data.message, limit=data.history_limit or 50)
# 检索并生成回答(同时返回检索到的文档 metadata)
result = retrieve_and_answer(data.message, vector_store, embedding, messages, top_k=5)
# 存储 AI 回复
storage.store_message(session_id=session_id, role="ai", content=result["answer"])
# 如果有检索到的文档,把序号 list 存储(便于后续汇总导出)
retrieved = result.get("retrieved_docs", [])
if retrieved:
# 存储为 JSON 字符串,role 用 'retrieved'
storage.store_message(session_id=session_id, role="retrieved", content=json.dumps(retrieved, ensure_ascii=False))
# 如果用户要求“表格”,则生成 Excel(基于会话内所有 retrieved 序号)
if "表格" in data.message and raw_df is not None:
# 从会话历史中收集所有 retrieved
history = storage.get_conversation_history(session_id=session_id, limit=1000)
seqs = []
for h in history:
if h["role"] == "retrieved":
try:
items = json.loads(h["content"])
for it in items:
sid = it.get("序号") or it.get("id") or it.get("ID")
if sid is not None:
seqs.append(str(sid).strip())
except Exception:
# 忽略解析错误
continue
seqs = list(dict.fromkeys(seqs)) # 去重并保留顺序
if not seqs:
# 没有序号可导出,直接返回普通回答
return {"session_id": session_id, "answer": result["answer"], "context": result.get("context", "")}
# 直接从 raw_df 中筛选对应序号
df = raw_df.copy()
if "序号" not in df.columns:
# 如果原表没有 '序号' 列,则不能进行精确抽取,返回普通回答
return {"session_id": session_id, "answer": result["answer"], "context": result.get("context", "")}
out_df = df[df["序号"].isin(seqs)].copy()
# 如果没有筛到任何行,也返回普通回答
if out_df.empty:
return {"session_id": session_id, "answer": result["answer"], "context": result.get("context", "")}
# 保存 Excel
output_path = os.path.join(EXPORT_DIR, f"export_{session_id}.xlsx")
out_df.to_excel(output_path, index=False)
# 返回文件响应(下载)
return FileResponse(output_path, filename=f"查询结果_{session_id}.xlsx", media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet")
# 默认返回 JSON 回答
return {"session_id": session_id, "answer": result["answer"], "context": result.get("context", "")}
@app.get("/")
def root():
"""健康检查"""
return {"status": "ok", "msg": "HeritageRAG is running"}
if __name__ == "__main__":
uvicorn.run("app:app", host="0.0.0.0", port=8000, reload=True)
更多推荐


所有评论(0)