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

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

项目主要会实现如下代码:


离线流程

-
项目基础数据 - /data/xxx.txt
-
app_file_uploader.py
可以先按照以下步骤运行网页:
1.编写测试代码
1
2
3
4
|
import streamlit as st
# 添加网页标题
st.title("知识库更新服务")
|
2.复制文件地址到命令提示符(cmd)中,进入到当前项目目录(c盘切到d盘:d:)
D:\RAG\AI大模型RAG与智能体开发\P4_RAG项目案例>
3.在cmd中输入streamlit run app_file_uploader.py则会自动跳转到网页中

4.使用 st.file_uploader 创建文件上传组件
1
2
3
4
5
6
|
# file_uploader
uploader_file = st.file_uploader(
"请上传TXT文件",
type=['txt'],
accept_multiple_files=False, # False表示仅接受一个文件的上传[不接受多文件]
)
|
5.提取文件信息
1
2
3
4
5
6
7
8
9
10
11
12
|
if uploader_file is not None:
# 提取文件的信息
file_name = uploader_file.name
file_type = uploader_file.type
file_size = uploader_file.size / 1024 # KB
st.subheader(f"文件名:{file_name}")
st.write(f"格式:{file_type} | 大小:{file_size:.2f} KB")
# get_value -> bytes -> decode('utf-8')
text = uploader_file.getvalue().decode("utf-8")
st.write(text)
|

完整代码如下:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
|
"""
基于Streamlit完成WEB网页上传服务
pip install streamlit
Streamlit:当WEB页面元素发生变化,则代码重新执行一遍
"""
import streamlit as st
# 添加网页标题
st.title("知识库更新服务")
# file_uploader
uploader_file = st.file_uploader(
"请上传TXT文件",
type=['txt'],
accept_multiple_files=False, # False表示仅接受一个文件的上传[不接受多文件]
)
if uploader_file is not None:
# 提取文件的信息
file_name = uploader_file.name
file_type = uploader_file.type
file_size = uploader_file.size / 1024 # KB
st.subheader(f"文件名:{file_name}")
st.write(f"格式:{file_type} | 大小:{file_size:.2f} KB")
# get_value -> bytes -> decode('utf-8')
text = uploader_file.getvalue().decode("utf-8")
st.write(text)
|
知识库服务

knowledge_base.py
工具类的实现:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
|
"""
知识库
"""
import os
import config_data as config
import hashlib
def check_md5(md5_str: str):
"""
检查传入的md5字符串是否已经被处理过了
return False(md5未处理过) 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_hex = md5_obj.hexdigest() # 得到md5的十六进制字符串
return md5_hex
class KnowledgeBaseService(object):
"""知识库服务"""
def __init__(self):
self.chroma = None # 向量存储的实例 Chroma向量库对象
self.spliter = None # 文本分割器的对象
def upload_by_str(self,data,filename):
"""将传入的字符串,进行向量化,存入向量数据"""
|
同样的内容计算的md5值是一样的(16进制字符串,长度32位),不管体积多大计算出来的结果都是32位,假设文本非常非常长,只要内容相同,计算的结果都是一样的(32位),方便做记录【节省空间,效率高】
以下为测试代码(粘贴至上述代码下方)
1
2
3
4
5
6
7
8
9
|
if __name__ == '__main__':
# 无论r1和r2里面内容多大,只要内容相同,输出结果就是一样的32位的16进制字符串
r1 = get_string_md5("周杰伦")
r2 = get_string_md5("周杰伦")
r3 = get_string_md5("周杰伦2")
print(r1)
print(r2)
print(r3)
|

测试check_md5:
1
2
|
if __name__ == '__main__':
print(check_md5("7a8941058aaf4df5147042ce104568da"))
|
无论运行几次都是False,怎样才能得到True呢
只需要调用save方法将内容保存至md5.text中即可
1
2
3
|
if __name__ == '__main__':
save_md5("7a8941058aaf4df5147042ce104568da")
print(check_md5("7a8941058aaf4df5147042ce104568da"))
|

配置文件:config_data.py
接下来继续完成知识库服务部分**(KnowledgeBaseService)**的代码编写
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
|
from langchain_chroma import Chroma
from langchain_community.embeddings import DashScopeEmbeddings
from langchain_text_splitters import RecursiveCharacterTextSplitter
from datetime import datetime
```
此处省略mt5部分的代码
```
class KnowledgeBaseService(object):
"""知识库服务"""
def __init__(self):
# 如果文件夹不存在则创建,如果存在则跳过
os.makedirs(config.persist_directory,exist_ok=True)
self.chroma = Chroma(
collection_name=config.collection_name, # 数据库的表名
embedding_function=DashScopeEmbeddings(model="text-embedding-v4"),
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,filename):
"""将传入的字符串,进行向量化,存入向量数据"""
# 先得到传入字符串的md5值
md5_hex = get_string_md5(data)
if check_md5(md5_hex):
return "[跳过]内容已经存在知识库中"
if len(data) > config.max_split_char_number:
# 如果传入的字符串长度大于阈值,则进行文本分割
knowledge_chunk: list[str] = self.spliter.split_text(data)
else:
knowledge_chunk = [data]
metadata = {
"source": filename,
# 2026-01-01 10:00:00
"create_time": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
"operator": "小曹",
}
self.chroma.add_texts( # 内容就加载到向量库中了
# iterable -> list \ tuple
knowledge_chunk,
# 列表推导式 :通过这个列表推导式,为每个文本块生成一份相同的元数据字典
metadatas=[metadata for _ in knowledge_chunk],
)
# 保存md5
save_md5(md5_hex)
return "[成功]内容已经成功载入到向量库中"
|
补充config_data.py
1
2
3
4
5
6
7
8
9
10
|
# Chroma
collection_name = "rag"
persist_directory = "./chroma_db"
# spliter
chunk_size = 1000
chunk_overlap = 100
separators = ["\n\n", "\n", ".", "!", "?", "。", "!", "?", " ", ""]
max_split_char_number = 1000 # 文本分割的阈值
|
测试代码(删除先前创建的text文件)
1
2
3
4
|
if __name__ == '__main__':
service = KnowledgeBaseService()
r = service.upload_by_str("周杰伦","testfile")
print(r)
|


再次运行

接下来将前面两个文件合并(app_file+knowledge)
我们可以在app_file文件中最下方写一行print语句
然后打开网页进行测试,可以发现不管是刷新还是添加文件,控制台都会输出一遍print里的内容,说明Streamlit:当WEB页面元素发生变化,则代码重新执行一遍
那么这个有什么问题呢?
其实问题还是挺大的,因为有了这个问题,就相当于我们整个代码都重跑了一遍,这会造成状态的丢失
我们可以来进行测试一下:
将前面的app文件代码修改为:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
|
import streamlit as st
# 添加网页标题
st.title("知识库更新服务")
# file_uploader
uploader_file = st.file_uploader(
"请上传TXT文件",
type=['txt'],
accept_multiple_files=False, # False表示仅接受一个文件的上传[不接受多文件]
)
count = 0
if uploader_file is not None:
# 提取文件的信息
file_name = uploader_file.name
file_type = uploader_file.type
file_size = uploader_file.size / 1024 # KB
st.subheader(f"文件名:{file_name}")
st.write(f"格式:{file_type} | 大小:{file_size:.2f} KB")
# get_value -> bytes -> decode('utf-8')
text = uploader_file.getvalue().decode("utf-8")
st.write(text)
count += 1
print(f"上传了{count}个文件")
|

第一次上传文件是正常的,但是后续再上传都是相当于每次将代码从头跑了一遍,count都是从0开始,可以发现,只要每次页面发生变化,代码都是从头开始执行一遍
将content相关的代码替换为:
1
2
3
|
# session_state就是一个字典
if "count" not in st.session_state:
st.session_state["count"] = 0
|
1
|
st.session_state["count"] += 1
|
1
|
print(f'上传了{st.session_state["count"]}个文件')
|
打开网页进行测试

这样就解决了状态维护的问题
但是这样还是有缺陷,因为只要页面一刷新,KnowledgeBaseService对象还是留不下来(接下来完整代码中出现),需要重新创建,性能消耗大
接下来就需要达成一个这样的效果:程序不停,只需要创建一个对象即可
完整代码:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
|
"""
基于Streamlit完成WEB网页上传服务
pip install streamlit
Streamlit:当WEB页面元素发生变化,则代码重新执行一遍
"""
import time
import streamlit as st
from knowledge_base import KnowledgeBaseService
# 添加网页标题
st.title("知识库更新服务")
# file_uploader
uploader_file = 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 uploader_file is not None:
# 提取文件的信息
file_name = uploader_file.name
file_type = uploader_file.type
file_size = uploader_file.size / 1024 # KB
st.subheader(f"文件名:{file_name}")
st.write(f"格式:{file_type} | 大小:{file_size:.2f} KB")
# get_value -> bytes -> decode('utf-8')
text = uploader_file.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) # 可以直接在页面中显示结果
|

在线流程

首先先创创建一个向量存储文件
vector_stores.py
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
|
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})
if __name__ == '__main__':
from langchain_community.embeddings import DashScopeEmbeddings
retriever = VectorStoreService(DashScopeEmbeddings(model="text-embedding-v4")).get_retriever()
res = retriever.invoke("我的体重120斤,尺码推荐")
print(res)
|
config_data中添加下面这一行代码:
1
|
similarity_threshold = 1 # 检索返回匹配的文档数量
|

接下来实现rag.py
基础实现:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
|
from langchain_core.documents import Document
from langchain_core.output_parsers import StrOutputParser
from langchain_core.runnables import RunnablePassthrough, RunnableWithMessageHistory, RunnableLambda
from vector_stores import VectorStoreService
from langchain_community.embeddings import DashScopeEmbeddings
import config_data as config
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
from langchain_community.chat_models.tongyi import ChatTongyi
def print_prompt(prompt):
print("="*20)
print(prompt.to_string())
print("="*20)
return prompt
class RagService(object):
def __init__(self):
# 向量数据库服务 - 用于检索
self.vector_service = VectorStoreService(
embedding=DashScopeEmbeddings(model=config.embedding_model_name)
)
self.prompt_template = ChatPromptTemplate.from_messages(
[
("system", "以我提供的已知参考资料为主,"
"简洁和专业的回答用户问题。参考资料:{context}。"),
("user", "请回答用户提问:{input}")
]
)
self.chat_model = ChatTongyi(model=config.chat_model_name)
self.chain = self.__get_chain()
def __get_chain(self):
"""获取最终的执行链"""
# 获取检索器
retriever = self.vector_service.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({"我身高170厘米,尺码推荐"}, )
print(res)
|

历史会话记录功能
完整代码
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
|
from langchain_core.documents import Document
from langchain_core.output_parsers import StrOutputParser
from langchain_core.runnables import RunnablePassthrough, RunnableWithMessageHistory, RunnableLambda
from file_history_store import get_history
from vector_stores import VectorStoreService
from langchain_community.embeddings import DashScopeEmbeddings
import config_data as config
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
from langchain_community.chat_models.tongyi import ChatTongyi
def print_prompt(prompt):
print("="*20)
print(prompt.to_string())
print("="*20)
return prompt
class RagService(object):
def __init__(self):
# 向量数据库服务 - 用于检索
self.vector_service = VectorStoreService(
embedding=DashScopeEmbeddings(model=config.embedding_model_name)
)
self.prompt_template = ChatPromptTemplate.from_messages(
[
("system", "以我提供的已知参考资料为主,"
"简洁和专业的回答用户问题。参考资料:{context}。"),
("system", "并且我提供用户的对话历史记录,如下:"),
MessagesPlaceholder("history"),
("user", "请回答用户提问:{input}")
]
)
self.chat_model = ChatTongyi(model=config.chat_model_name)
self.chain = self.__get_chain()
def __get_chain(self):
"""获取最终的执行链"""
# 获取检索器
retriever = self.vector_service.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: dict) -> str:
return value["input"]
# 格式化Prompt模板所需的输入变量
def format_for_prompt_template(value):
# 将嵌套字典展平为 {input, context, history} 格式
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_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": "春天穿什么颜色的衣服?"}, session_config)
print(res)
|

聊天页面开发
进入到“智能客服”网页:
1.编写测试代码:
1
2
3
4
5
|
import streamlit as st
# 标题
st.title("智能客服")
st.divider()
|
2.复制文件地址到命令提示符(cmd)中,进入到当前项目目录:
D:\RAG\AI大模型RAG与智能体开发\P4_RAG项目案例\app_qa.py>
3.启动项目
1
|
streamlit run app_qa.py
|
完整代码:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
|
import time
from rag import RagService # 导入RAG服务类
import streamlit as st # 导入Streamlit用于构建Web界面
import config_data as config # 导入配置文件
# ========== 页面标题 ==========
st.title("智能客服") # 设置页面标题
st.divider() # 添加分隔符
# ========== 初始化会话状态 ==========
# 存储聊天记录,初始时添加一条欢迎消息
if "message" not in st.session_state:
st.session_state["message"] = [{"role": "assistant", "content": "你好,有什么可以帮助你?"}]
# 初始化RAG服务实例(只创建一次,避免重复初始化)
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响应 ==========
ai_res_list = [] # 用于缓存AI的流式响应
with st.spinner("AI思考中..."):
# 调用RAG链进行流式响应(res_stream:迭代器)
res_stream = st.session_state["rag"].chain.stream({"input": prompt}, config.session_config)
#yield
# 定义捕获生成器:将流式数据同时返回并缓存
def capture(generator, cache_list):
for chunk in generator:
cache_list.append(chunk) # 缓存每个chunk
yield chunk # 返回给Streamlit显示
# 流式输出AI响应
st.chat_message("assistant").write_stream(capture(res_stream, ai_res_list))
# 将AI完整响应添加到聊天记录
st.session_state["message"].append({"role": "assistant", "content": "".join(ai_res_list)})
# ["a", "b", "c"] "".join(list) -> abc
# ["a", "b", "c"] ",".join(list) -> a,b,c
|
配置文件config_data.py:
在原本的文件上添加以下代码:
1
2
3
4
5
|
session_config = {
"configurable": {
"session_id": "user_001",
}
}
|
