[RAG开发]-智能客服项目
·
项目介绍
RAG回顾
RAG即检索、增强和生成,其主要分为2条线:
- 离线处理: 向私有知识库(向量存储)源源不断添加私有知识文档。
- 向知识库添加来自未来的知识文档(基于模型训练完成时间)
- 向模型添加私有知识文档
- 给出模型参考资料,规避模型幻觉(一本正经的胡说八道)
- 在线处理: 用户提问会先基于私有知识库做检索,获取参考资料,同步组装新提示词询问大模型获取结果。

项目需求和思路
本次项目以“某东商品衣服"为例,以衣服属性构建本地知识。使用者可以自由更新本地知识,用户问题的答案也是基于本地知识生成的。

项目主要会实现如下代码


离线模块
离线流程

- app_file_uploader.py是文件服务
- uploader_file.get_value(): 拿到用户上传文件的内容
- KnowledgetBaseService: 创建知识库服务实例,把文件内容存入向量数据库
- knowledge_base.py是知识库服务
- check_md5: 读取记录文件,对比md5值
- save_md5: 把md5值写入记录文件
- get_string_md5: 把文件内容转md5值
- chroma: 向量数据库实例
- spliter: 字符串分割器实例
- upload_by_str: 数据入库
- config_data.py是全局配置文件
完成app_file_uploader.py文件的编写
"""
基于Streamlit完成Web页面上传服务
pip install streamlit
"""
import streamlit as st
# 添加网页标题
st.title("知识库更新服务")
# 文件上传
uploaded_file = st.file_uploader(
"请上传TXT文件", # 标题
type=["txt"], # 允许上传的文件类型
accept_multiple_files=False, # 是否允许上传多个文件
)
if uploaded_file is not None:
# 提取文件信息
file_name = uploaded_file.name
file_type = uploaded_file.type
file_size = uploaded_file.size / 1024 # KB
# 展示文件信息
st.subheader(f"文件名称:{file_name}")
st.write(f"类型:{file_type} | 大小:{file_size:.2f} KB")
# 提取文件内容
# getvalue()默认返回bytes, 使用decode()转码为str
text = uploaded_file.getvalue().decode("utf-8")
st.write(text)
- 通过命令行启动服务


- 到此,完成了一部分工作

完成knowledge_base.py文件的编写
"""
知识库
"""
import hashlib
import os
from datetime import datetime
from langchain_community.embeddings import DashScopeEmbeddings
from langchain_community.vectorstores import Chroma
from langchain_text_splitters import RecursiveCharacterTextSplitter
import config_data as config
def check_md5(md5_str: str):
"""
检查md5字符串是否存在
"""
if not os.path.exists(config.md5_path):
# 文件不存在
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)
# 创建MD5对象
md5_obj = hashlib.md5()
# 更新转换内容
md5_obj.update(str_bytes)
# 得到md5的十六进制字符串
md5_str = md5_obj.hexdigest()
return md5_str
class KnowLedgeBaseService(object):
"""
知识库服务
"""
def __init__(self):
# 如果文件夹不存在则创建
os.makedirs(config.persist_directory, exist_ok=True)
# Chroma向量库实例
self.chroma = Chroma(
collection_name=config.collection_name, # 数据库表名
embedding_function=DashScopeEmbeddings( # 创建嵌入模型实例
model="text-embedding-v4",
dashscope_api_key="sk-0af497cdd462495aa2f08dddaead1ded",
),
persist_directory=config.persist_directory, # 持久化存储目录
)
# 文本分割器实例
self.spliter = RecursiveCharacterTextSplitter(
chunk_size=config.chunk_size, # 分割后的文本段最大长度
chunk_overlap=config.chunk_overlap, # 相邻段落间之间的字符重叠数
separators=config.separators, # 分割的分隔符
length_function=len, # 使用py自带的len函数做字符长度统计
)
def upload_by_str(self, data, filename):
"""将字符串向量化,存入数据库"""
md5_hex = get_string_md5(data)
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]
metadate = {
"source": filename,
"create_time": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
"operator": "王",
}
# 内容加载到向量库
self.chroma.add_texts(
# 要求传入Iterable[str], Iterable即可迭代对象,list数组/tuple字典都是可迭代对象,然后容器中是str
texts=knowledge_chunks, # 待加载的文本
metadate=[metadate for _ in knowledge_chunks], # 数据的元数据(额外信息)
)
# 记录md5值
save_md5(md5_hex)
return "[完成]内容已经加载到向量库中"
if __name__ == '__main__':
# 单元测试
# print(get_string_md5("hello world"))
# print(check_md5("5eb63bbbe01eeed093cb22bb8f5acdc3"))
# print(save_md5("5eb63bbbe01eeed093cb22bb8f5acdc3"))
service = KnowLedgeBaseService()
r = service.upload_by_str("周杰伦", "testfile")
print(r)
- 配置文件
md5_path = "./md5.text"
# Chroma
collection_name = "rag"
persist_directory = "./chroma_db"
# spliter
chunk_size = 1000
chunk_overlap = 100
separators = ["\n\n", "\n", " ", "", ".", "!", "?", "。","!","?"]
max_split_char_number = 1000 # 文本分隔的阈值 (过短的文本没必要分割)
- 执行单元测试




- 至此,完成一部分工作

完成离线流程的开发
"""
基于Streamlit完成Web页面上传服务
pip install streamlit
Streamlit的重要特点
1. 当web网页元素发生变化,则代码重新执行一遍,
2. 会造成页面状态的丢失
3. 如果有数据需要保存,要借助session_state实现
4. session_state就是一个持久化存储的字典
"""
import time
import streamlit as st
from knowledge_base import KnowLedgeBaseService
# 添加网页标题
st.title("知识库更新服务")
# 文件上传
uploaded_file = st.file_uploader(
"请上传TXT文件", # 标题
type=["txt"], # 允许上传的文件类型
accept_multiple_files=False, # 是否允许上传多个文件
)
if "service" not in st.session_state:
st.session_state["service"] = KnowLedgeBaseService()
if uploaded_file is not None:
# 提取文件信息
file_name = uploaded_file.name
file_type = uploaded_file.type
file_size = uploaded_file.size / 1024 # KB
# 展示文件信息
st.subheader(f"文件名称:{file_name}")
st.write(f"类型:{file_type} | 大小:{file_size:.2f} KB")
# 提取文件内容
# getvalue()默认返回bytes, 使用decode()转码为str
text = uploaded_file.getvalue().decode("utf-8")
# st.write(text)
# 存储文件 (加一个loding效果)
with st.spinner("载入知识库中..."):
time.sleep(1)
# 调用知识库服务完成文件内容的向量化以及存储
result = st.session_state["service"].upload_by_str(text, file_name)
st.write(result)
- 调试

- 至此,完成的工作,整个离线流程全部打通

在线模块
在线流程

- vector_stores: 向量存储服务
- 通过get_retriever方法获取检索器,用于向量检索
- file_history_store.py: 历史信息存储
- rag.py:检索增强服务
- vector_service: 向量检索器的实例
- prompt_template: 提示词模板
- chain: 链实例
- chat_model: 模型对象
- __get_chain(): 获取链
- invoke: 向模型提问
- app_qa.py 用户交互服务
vector_stores 向量存储服务开发
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):
"""
获取向量检索器
"""
return self.vector_store.as_retriever(search_kwargs={"k": config.simplify_threshold})
if __name__ == '__main__':
from langchain_community.embeddings import DashScopeEmbeddings
# retriever就是检索器对象
retriever = service = VectorStoreService(DashScopeEmbeddings(model="text-embedding-v4", dashscope_api_key="sk-0af497cdd462495aa2f08dddaead1ded")).get_retriever()
res = retriever.invoke("我的体重180斤,尺码推荐")
print(res)
- 测试:已经把匹配到的 尺码推荐.text 的知识拿到了

rag.py 知识检索开发
from langchain_community.chat_models import ChatTongyi
from langchain_community.embeddings import DashScopeEmbeddings
from langchain_core.documents import Document
from langchain_core.output_parsers import StrOutputParser
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.runnables import RunnablePassthrough
import config_data as config
from vector_stores import VectorStoreService
def print_prompt(prompt):
print("=" * 20)
print(prompt.to_string())
print("=" * 20)
return prompt
class RagService(object):
def __init__(self):
# 构建向量存储对象
self.vector_store = VectorStoreService(
embedding=DashScopeEmbeddings(model=config.embedding_model_name,
dashscope_api_key="sk-0af497cdd462495aa2f08dddaead1ded")
)
# 构建提示词模板
self.prompt_template = ChatPromptTemplate.from_messages(
[
("system", "以我提供的已知参考资料为主,简洁专业的回答用户问题。参考资料:{context}"),
("user", "请回答用户提问:{input}"),
]
)
# 构建模型对象
self.chat_model = ChatTongyi(model=config.chat_model_name, api_key="sk-0af497cdd462495aa2f08dddaead1ded")
# 构建执行链
self.chain = self.__get_chain()
def __get_chain(self):
"""获取最终的执行链"""
retriever = self.vector_store.get_retriever()
def format_document(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
chain = (
{
"input": RunnablePassthrough(),
"context": retriever | format_document,
} | self.prompt_template | print_prompt | self.chat_model | StrOutputParser()
)
return chain
if __name__ == '__main__':
res = RagService().chain.invoke("我身高180厘米,尺码推荐")
print(res)

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, messages_from_dict, message_to_dict
# 自定义会话存储的方法
class FileChatMessageHistory(BaseChatMessageHistory):
def __init__(self, session_id, storage_path):
self.session_id = session_id # 会话id
self.storage_path = storage_path # 存储路径
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]): # Sequence序列 类似list/tuple
all_messages = list(self.messages) # 已有的消息列表, messages来自父类
all_messages.extend(messages) # 把新的消息和旧的消息合并为一个list
# 将数据同步写入到本地文件
# 类对象不能直接写入文件, 如果直接写入是一堆二进制数据
# 要将BaseMessage对象转成字典, 再借助json模块以json格式写入文件
# 完整代码
# new_message = []
# for message in all_messages:
# new_message.append(message.to_dict())
# 列表推导式
# 流程: 遍历all_messages, 赋值给message, 调用message_to_dict()方法, 把返回值放入list
new_message = [message_to_dict(message) for message in all_messages]
# 将数据写入文件
with open(self.file_path, "w", encoding="utf-8") as f:
json.dump(new_message, f)
@property # 该注解不能省略, 作用是把该方法变成成员属性, LangChain调用该方法是以属性的方式调用的
def messages(self) -> list[BaseMessage]:
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)
# 历史会话容器 (由内存存储变成文件存储)
def get_history(session_id):
return FileChatMessageHistory(session_id, "./chat_history")
from langchain_community.chat_models import ChatTongyi
from langchain_community.embeddings import DashScopeEmbeddings
from langchain_core.documents import Document
from langchain_core.output_parsers import StrOutputParser
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
from langchain_core.runnables import RunnablePassthrough, RunnableWithMessageHistory, RunnableLambda
import config_data as config
from file_history_store import get_history
from vector_stores import VectorStoreService
def print_prompt(prompt):
print("=" * 20)
print(prompt.to_string())
print("=" * 20)
return prompt
class RagService(object):
def __init__(self):
# 构建向量存储对象
self.vector_store = VectorStoreService(
embedding=DashScopeEmbeddings(model=config.embedding_model_name,
dashscope_api_key="sk-0af497cdd462495aa2f08dddaead1ded")
)
# 构建提示词模板
self.prompt_template = ChatPromptTemplate.from_messages(
[
("system", "以我提供的已知参考资料为主,简洁专业的回答用户问题。参考资料:{context}"),
("system", "并且我提供用户的对话历史记录,如下:"),
MessagesPlaceholder("history"),
("user", "请回答用户提问:{input}"),
]
)
# 构建模型对象
self.chat_model = ChatTongyi(model=config.chat_model_name, api_key="sk-0af497cdd462495aa2f08dddaead1ded")
# 构建执行链
self.chain = self.__get_chain()
def __get_chain(self):
"""获取最终的执行链"""
retriever = self.vector_store.get_retriever()
def format_document(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):
return value["input"]
# 重组参数
def format_for_prompt_template(value):
new_value = {}
new_value["input"] = value["input"]["input"]
new_value["history"] = value["input"]["history"]
new_value["context"] = value["context"]
return new_value
# 基础链
chain = (
{
"input": RunnablePassthrough(),
"context": RunnableLambda(format_for_retriever) | retriever | format_document,
} | 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"
}
}
res = RagService().chain.invoke({"input": "我身高180厘米,尺码推荐"}, session_config)
print(res)

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": "你好,有什么可以帮你?"}]
for message in st.session_state["message"]:
st.chat_message(message["role"]).write(message["content"])
if "rag" not in st.session_state:
st.session_state["rag"] = RagService()
# 用户输入栏
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_stream = 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(capture(res_stream, ai_res_list))
st.session_state["message"].append({"role": "assistant", "content": "".join(ai_res_list)})


更多推荐



所有评论(0)