RAG项目案例

RAG项目案例

RAG回顾

RAG即检索、增强和生成,其主要分为2条线:

  • 离线处理:向私有知识库(向量存储)源源不断添加私有知识文档。
    • 向知识库添加来自未来的知识文档(基于模型训练完成时间)
    • 向模型添加私有知识文档
    • 给出模型参考资料,规避模型幻觉(一本正经的胡说八道)
  • 在线处理:用户提问会先基于私有知识库做检索,获取参考资料,同步组装新提示词询问大模型获取结果。

项目需求和思路

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

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

离线流程

  1. 项目基础数据 - /data/xxx.txt

  2. 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

1
md5_path = "./md5.txt"

接下来继续完成知识库服务部分**(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",
        }
    }

本站于2026年3月31日建立
使用 Hugo 构建
主题 StackJimmy 设计