大模型学习·第45天:RAG实战——基于LangChain+Chroma+Streamlit搭建知识库问答系统
·
一、项目背景
大模型存在知识滞后、缺乏专业领域知识、容易产生幻觉等问题。RAG(检索增强生成)通过外挂知识库,先检索相关资料再让模型回答,是解决这些问题的核心方案。
二、技术栈
| 技术 | 作用 |
|---|---|
| LangChain | 统一接口串联各组件 |
| Chroma | 向量数据库,存储文档向量 |
| Ollama | 本地部署嵌入模型和聊天模型 |
| Streamlit | 搭建可视化Web界面 |
三、项目架构
app_file_uploader.py # 文件上传页面
app_qa.py # 问答页面
rag.py # RAG核心服务
vector_stores.py # 向量存储封装
knowledge_base.py # 知识库管理(MD5去重)
file_history_store.py # 多轮对话记忆
config_data.py # 全局配置
四、核心流程
上传流程:文件上传 → MD5去重检查 → 文本分割 → 向量化 → 存入Chroma。
问答流程:用户提问 → 向量检索 → 拼接提示词 → 模型生成 → 流式输出。
五、关键代码
RAG链式调用:用 RunnablePassthrough 和 RunnableLambda 把检索和生成串联起来,检索上下文自动注入提示词模板。
MD5去重:上传文件时先计算内容MD5值,已存在则跳过,防止重复上传浪费向量库空间。
会话记忆:继承 BaseChatMessageHistory 实现文件存储,结合 RunnableWithMessageHistory 实现多轮对话。
六、项目截图

七、项目亮点
- MD5去重机制,防止重复上传
- 流式输出,用户体验更好
- 多轮对话记忆,上下文连贯
- 模块化设计,配置统一管理
八、项目地址
https://github.com/OldCat263/rag-knowledge-base.git
app_file_uploader.py
"""
基于Streamlit
Streamlit:WEB页面元素发生变化,则代码重新执行一遍
"""
import time
import streamlit as st
from knowledge_base import KnowledgeBaseService
#添加网页标题
st.title("知识库更新服务")
#file uploader
uploaded_files = st.file_uploader("请上传TXT文件",
type=["txt"],
accept_multiple_files=False,#False,表示仅接受一个文件的上传
)
#session_state就是一个字典
if "service" not in st.session_state:
st.session_state["service"] = KnowledgeBaseService()
if uploaded_files is not None:
#提取文件的信息
file_name = uploaded_files.name
file_size = uploaded_files.size / 1024 # 转换为KB
file_type = uploaded_files.type
st.subheader("文件名: {}".format(file_name))
st.write("文件大小: {file_size:.2f} KB | 文件类型: {file_type}".format(file_size=file_size, file_type=file_type))
#getvalue
text = uploaded_files.getvalue().decode("utf-8") # 将字节流解码为字符串
with st.spinner("正在处理文件,请稍等..."): #在spinner内的代码执行过程中,会有一个转圈动画
time.sleep(1) # 模拟处理时间
result = st.session_state["service"].upload_by_str(text,file_name)
st.write(result)
app_qa.py
import time
from rag import RagService
import streamlit as st
import config_data as config
# 标题
st.title("智能客服")
st.divider()
# 在页面最下方提供用户输入栏
if "message" not in st.session_state:
st.session_state["message"] = [{"role":"assistant","content":"你好请问有什么可以帮你?"}]
if "rag" not in st.session_state:
st.session_state["rag"] = RagService()
for message in st.session_state["message"]:
st.chat_message(message["role"]).write(message["content"])
prompt = st.chat_input()
if prompt:
# 在页面输出用户的提问
st.chat_message("user").write(prompt)
st.session_state["message"].append({"role":"user","content":prompt})
ai_res_list = []
with st.spinner("Ai思考中..."):
res = st.session_state["rag"].chain.stream({"input":prompt},config.session_config)
def capture(generator,cache_list):
for chunk in generator:
cache_list.append(chunk)
yield chunk
st.chat_message("assistant").write_stream(capture(res,ai_res_list))
st.session_state["message"].append({"role":"assistant","content":"".join(ai_res_list)})
config_data.py
md5_path = "./md5.text"
# Chroma
collection_name = "rag" # 数据库的表名
persist_directory = "./chroma_db" # 数据库本地存储文件夹
# spliter
chunk_size = 1000 # 分割后的段大小
chunk_overlap = 200 # 分割后的文本段之间的重叠长度
separators = ["\n\n", "\n", " ", ""] # 自然段落划分的符号
max_split_char_number = 1000 # 文本分割的阈值
#
similarity_threshold = 2 # 检索返回匹配的文档数量
embedding_model_name = "nomic-embed-text" # 嵌入模型
chat_model_name = "qwen:1.8b" # 聊天模型
session_config = {
"configurable":{
"session_id":"user_001"
}
}
file_history_store.py
import json
import os
from typing import Sequence
from langchain_core.chat_history import BaseChatMessageHistory
from langchain_core.messages import BaseMessage, message_to_dict, messages_from_dict
def get_history(session_id):
return FileChatMessageHistory(session_id, "./chat_history")
class FileChatMessageHistory(BaseChatMessageHistory):
def __init__(self, session_id, storage_path):
self.session_id = session_id # 会话id
self.storage_path = storage_path # 不同会话id的存储文件,所在的文件夹路径
# 完整的文件路径
self.file_path = os.path.join(self.storage_path, self.session_id)
# 确保文件夹是存在的
os.makedirs(os.path.dirname(self.file_path), exist_ok=True)
def add_messages(self, messages: Sequence[BaseMessage]) -> None:
# Sequence序列 类似list、tuple
all_messages = list(self.messages) # 已有的消息列表
all_messages.extend(messages) # 新的和已有的融合成一个list
# 将数据同步写入到本地文件中
# 类对象写入文件 -> 一堆二进制
# 为了方便,可以将BaseMessage消息转为字典(借助json模块以json字符串写入文件)
# 官方message_to_dict:单个消息对象(BaseMessage类实例) -> 字典
# new_messages = []
# for message in all_messages:
# d = message_to_dict(message)
# new_messages.append(d)
new_messages = [message_to_dict(message) for message in all_messages]
# 将数据写入文件
with open(self.file_path, "w", encoding="utf-8") as f:
json.dump(new_messages, f)
@property # @property装饰器将messages方法变成成员属性用
def messages(self) -> list[BaseMessage]:
# 当前文件内: list[字典]
try:
with open(self.file_path, "r", encoding="utf-8") as f:
messages_data = json.load(f) # 返回值就是:list[字典]
return messages_from_dict(messages_data)
except FileNotFoundError:
return []
def clear(self) -> None:
with open(self.file_path, "w", encoding="utf-8") as f:
json.dump([], f)
knowledge_base.py
"""
知识库
"""
import os
import config_data as config
import hashlib
from langchain_chroma import Chroma
from langchain_ollama import OllamaEmbeddings
from langchain_text_splitters import RecursiveCharacterTextSplitter
from datetime import datetime
def check_md5(md5_str:str):
"""检查完传去的md5字符串是否被处理过
return false表示没有处理过,true表示处理过"""
if not os.path.exists(config.md5_path):
#if进入表示文件不存在,那肯定没有处理过这个md5
open(config.md5_path, "w",encoding = "utf-8").close() # 创建空文件
return False
else:
for line in open(config.md5_path, "r",encoding = "utf-8").readlines(): # 打开文件,检查是否存在
line = line.strip() # 去掉换行符
if line == md5_str:
return True # 已经处理过
return False
def save_md5(md5_str:str):
"""保存md5字符串"""
with open(config.md5_path, "a", encoding = "utf-8") as f:
f.write(md5_str + "\n")
def get_string_md5(input_str:str,encoding='utf-8'):
"""将传入的字符串转为md5字符串"""
# 将字符串转换为bytes字节数组
str_bytes = input_str.encode(encoding = encoding)
#创建md5对象
md5_obj = hashlib.md5() #得到MD5对象
md5_obj.update(str_bytes) #更新MD5对象(传入即将要转换的字节数组)
md5_hex = md5_obj.hexdigest() #得到MD5的十六进制字符串
return md5_hex
class KnowledgeBaseService:
"""知识库服务类"""
def __init__(self):
# 如果文件夹不存在则创建,如果存在则跳过
os.makedirs(config.persist_directory, exist_ok=True)
self.chroma = Chroma(
collection_name=config.collection_name, # 数据库的表名
embedding_function=OllamaEmbeddings(model="nomic-embed-text"), # 嵌入模型
persist_directory=config.persist_directory # 数据库本地存储文件夹
) # 向量存储的实例Chroma向量库对象
self.spliter = RecursiveCharacterTextSplitter(
chunk_size=config.chunk_size, # 分割后的文本段最大长度
chunk_overlap=config.chunk_overlap, # 连续文本段之间的字符重叠数量
separators=config.separators, # 自然段落划分的符号
length_function=len # 使用python自带的len函数计算文本长度
) # 文本分割器的对象
def upload_by_str(self, data:str,filename):
"""将传入的字符串,进行向量化,存入向量数据库中"""
#先得到传入字符串的md5值
md5_hex = get_string_md5(data)
#检查md5值是否已经处理过
if check_md5(md5_hex):
return "文件已经处理过了,请勿重复上传"
if len(data) > config.max_split_char_number:
knowledge_chunks: list[str]= self.spliter.split_text(data) # 将文本分割为多个段落
else:
knowledge_chunks = [data] # 如果文本长度小于阈值,则不进行分割,直接将整个文本作为一个段落
metadata = {
"source": filename,
"create_time":datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
"operator": "admin"
}
self.chroma.add_texts( #内容就加载到了向量库中了
knowledge_chunks,
metadatas=[metadata for _ in knowledge_chunks]
) # 将文本段落添加到向量数据库中,并为每个段落添加元数据
#
save_md5(md5_hex) # 将md5值保存到文件中,表示已经处理过了
return "[成功] 文件上传成功,知识库已更新"
if __name__ == "__main__":
service = KnowledgeBaseService()
r = service.upload_by_str("这是一个测试字符串", "test")
print(r)
rag.py
from vector_stores import VectorStoreService
import config_data as config
from langchain_core.prompts import ChatPromptTemplate,MessagesPlaceholder
from langchain_ollama import OllamaEmbeddings, ChatOllama
from langchain_core.runnables import RunnableLambda, RunnablePassthrough, RunnableWithMessageHistory
from langchain_core.documents import Document
from langchain_core.output_parsers import StrOutputParser
from file_history_store import get_history
def print_prompt(prompt):
print("Prompt内容如下:")
print("="*20)
print(prompt.to_string())
print("="*20)
return prompt
class RagService:
def __init__(self):
self.vector_service = VectorStoreService(
embedding=OllamaEmbeddings(model=config.embedding_model_name)
)
self.prompt_template = ChatPromptTemplate.from_messages(
[
("system", "以我提供的已知参考资料为主,"
"简介和专业的回答用户提问,参考资料:{context}"),
MessagesPlaceholder("history"),
("user", "请回答用户提问:{input}")
]
)
self.chat_model = ChatOllama(model=config.chat_model_name)
self.chain = self.__get_chain()
def __get_chain(self):
# 获取最终的执行链
retriever = self.vector_service.get_retriever()
def format_documents(docs:list[Document]):
# 格式化文档为字符串
if not docs:
return "没有找到相关的参考资料。"
formatted_str = ""
for doc in docs:
formatted_str += f"文档片段:{doc.page_content}\n文档元数据:{doc.metadata}\n\n"
return formatted_str
def format_for_retriever(value: dict) -> str:
return value["input"]
def format_for_prompt_template(value):
new_value={}
new_value["input"] = value["input"]["input"]
new_value["context"] = value["context"]
new_value["history"] = value["input"]["history"]
return new_value
chain = (
{
"input":RunnablePassthrough(),
"context": RunnableLambda(format_for_retriever) | retriever | format_documents
} | RunnableLambda(format_for_prompt_template) | self.prompt_template | print_prompt | self.chat_model | StrOutputParser()
)
conversation_chain = RunnableWithMessageHistory(
chain,
get_history,
input_messages_key="input",
history_messages_key="history",
)
return conversation_chain
if __name__ == "__main__":
#session id 配置
session_config = {
"configurable":{
"session_id":"user_001"
}
}
rag_service = RagService()
question = {"input":"春天穿什么颜色的衣服"}
result = rag_service.chain.invoke(question, session_config)
print(result)
vector_stores.py
from langchain_chroma import Chroma
import config_data as config
class VectorStoreService(object):
def __init__(self,embedding):
# param embedding: 嵌入模型对象
self.embedding = embedding
self.vector_store = Chroma(
collection_name=config.collection_name, # 数据库的表名
embedding_function=self.embedding, # 嵌入模型
persist_directory=config.persist_directory # 数据库本地存储文件夹
)
def get_retriever(self):
#返回向量检索器,方便加入chain
return self.vector_store.as_retriever(search_kwargs={"k": config.similarity_threshold}) # 返回向量检索器,方便加入chain
if __name__ == "__main__":
from langchain_ollama import OllamaEmbeddings
retriever = VectorStoreService(OllamaEmbeddings(model="nomic-embed-text")).get_retriever()
res = retriever.invoke("我的体重180斤,尺码推荐") # 测试检索器是否正常工作
print(res)
更多推荐
所有评论(0)