一、项目背景

大模型存在知识滞后、缺乏专业领域知识、容易产生幻觉等问题。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链式调用:用 RunnablePassthroughRunnableLambda 把检索和生成串联起来,检索上下文自动注入提示词模板。

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)

更多推荐