RAG实战开发

一、RAG介绍

二、文件上传Web服务

实现app_file_uploader.py

安装streamlit

pip install streamlit

编写第一个简单页面

app_file_uploader.py

"""
基于Streamlit完成Web网页上传服务

pip install streamlit
"""

import streamlit as st

# 添加网页标题
st.title("知识库更新服务")

运行

命令行输入:

streamlit run .\app_file_uploader.py

继续完善代码

"""
基于Streamlit完成Web网页上传服务

pip install streamlit
"""

import streamlit as st
from fsspec.implementations.http import file_size

# 添加网页标题 - 一级大标题(页面主标题)
st.title("知识库更新服务")


# file_uploader
uploader_file = st.file_uploader(
    "请上传txt文件",                   # label文字
    type=["txt"],                          # 支持文件上传类型
    accept_multiple_files=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)

三、MD5工具函数开发

"""
知识库
"""
from flatbuffers.flexbuffers import Object
import os
import config_data as config
import hashlib


def check_md5(md5_str: str):
    """
    检查传入的md5字符串师傅已经被处理过了
    :param md5_str:
    :return: False(md5未处理过) True(md5处理过)
    """
    if not os.path.exists(config.md5_path):
        # if进入表示文件不存在,那肯定没有处理过这个md5值了, 创建文本并返回False
        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 md5_str == line:
                return True       #  已处理过

        return False      #  未处理过

def save_md5(md5_str: str):
    """
    将传入的md5字符串,记录(追加)到文件内保存
    :param md5_str:
    :return:
    """
    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字符串
   :param input_str: 
   :param encoding: 
   :return: 
   """
   # 将字符串转换为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_stream(self, data, filename):
        """
        将传入的字符串存入向量数据库
        :param data:
        :param filename:
        :return:
        """

if __name__ == '__main__':
    """测试方法"""
    check_md5("f6c9b63afc14875864a1f7ac3a7dcfe5")
    print(check_md5("f6c9b63afc14875864a1f7ac3a7dcfe5"))

四、知识库更新服务

"""
知识库
"""
from flatbuffers.flexbuffers import Object
import os

from langchain_community.embeddings import DashScopeEmbeddings
from langchain_text_splitters import RecursiveCharacterTextSplitter
from datetime import datetime

import config_data as config
import hashlib
from langchain_chroma import Chroma


def check_md5(md5_str: str):
    """
    检查传入的md5字符串师傅已经被处理过了
    :param md5_str:
    :return: False(md5未处理过) True(md5处理过)
    """
    if not os.path.exists(config.md5_path):
        # if进入表示文件不存在,那肯定没有处理过这个md5值了, 创建文本并返回False
        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 md5_str == line:
                return True       #  已处理过

        return False      #  未处理过

def save_md5(md5_str: str):
    """
    将传入的md5字符串,记录(追加)到文件内保存
    :param md5_str:
    :return:
    """
    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字符串
   :param input_str: 
   :param encoding: 
   :return: 
   """
   # 将字符串转换为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):
        os.makedirs(config.persist_directory, exist_ok=True)   # 如果文件夹不存在则创建,存在则跳过
        self.chroma = Chroma(
            collection_name=config.collection_name,                                # 向量数据库表名
            embedding_function=DashScopeEmbeddings(model="text-embedding-v4"),     # 默认用的v1,我们显示指定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  # 统计字符的依据函数

        )     # 文本分割器的对象

    def upload_by_str(self, data: str, filename):
        """
        将传入的字符串存入向量数据库
        :param data:
        :param filename:
        :return:
        """
        # 先得到传入字符串的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,
            # 2025-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_chunks,
            metadatas=[metadata for i in knowledge_chunks],           # 元数据要与知识片段保持一对一的关系, 有多少知识片段,就有多少元数据
        )

        # 把处理过的md5保存到md5存储文件中, 这样下次再遇到相等的md5就不再处理了
        save_md5(md5_hex)

        return "[成功]内容已经成功载入向量库"

if __name__ == '__main__':
    """测试方法"""
    service = KnowledgeBaseService()
    r = service.upload_by_str("周杰2222111", "testfile")
    print(r)
C:\Program_Files\Python\Python31210\python.exe D:\PythonProject\AI_RAG_AGENT\RAG项目案例\knowledge_base.py 
D:\PythonProject\AI_RAG_AGENT\RAG项目案例\knowledge_base.py:7: DeprecationWarning: `langchain-community` is being sunset and is no longer actively maintained. See https://github.com/langchain-ai/langchain-community/issues/674 for details and migration guidance toward standalone integration packages.
  from langchain_community.embeddings import DashScopeEmbeddings
[跳过]内容已经存在知识库中

Process finished with exit code 0

五、完成离线流程开发

完成文件web上传和知识库服务整合

其实,现在存在一个问题,页面每次刷新,streamlit就会重新执行一遍。

后果,一些状态得不到保存。

修改app_file_uploader.py

"""
基于Streamlit完成Web网页上传服务

pip install streamlit
"""

import streamlit as st
from fsspec.implementations.http import file_size

# 添加网页标题 - 一级大标题(页面主标题)
st.title("知识库更新服务")


# file_uploader
uploader_file = st.file_uploader(
    "请上传txt文件",                   # label文字
    type=["txt"],                          # 支持文件上传类型
    accept_multiple_files=False,           # 是否接受多文件上传
)

# 解决页面刷新后,会话数据丢失问题
# session_state就是一个字典
if "counter" not in st.session_state:
    st.session_state["counter"] = 0

st.session_state.uploader_file = uploader_file
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)

    st.session_state["counter"] += 1
    print(f"上传了{st.session_state["counter"]}个文件")

继续完成离线流程整合

修改app_file_uploader.py

六、在线流程向量存储服务代码

开发vector_stores.py

from langchain_chroma import Chroma
from langchain_community.embeddings import DashScopeEmbeddings

import config_data as config

class VectorStoreService(object):
    def __init__(self, embeddings):
        """
        嵌入模型的传入
        :param embeddings:
        """
        self.embedding = embeddings
        self.vector_store = Chroma(
            collection_name=config.collection_name,
            embedding_function=self.embedding,
            persist_directory=config.persist_directory,
        )

    def get_retriver(self):
        """
        返回向量检索器,方便加入chain
        :return:
        """
        return self.vector_store.as_retriever(search_kwargs={"k": config.simplify_threshold})



if __name__ == '__main__':
    retriver = VectorStoreService(DashScopeEmbeddings(model="text-embedding-v4")).get_retriver()
    res = retriver.invoke("我的体重180斤,尺码推荐")
    print(res)
C:\Program_Files\Python\Python31210\python.exe D:\PythonProject\AI_RAG_AGENT\RAG项目案例\vector_stores.py 
D:\PythonProject\AI_RAG_AGENT\RAG项目案例\vector_stores.py:2: DeprecationWarning: `langchain-community` is being sunset and is no longer actively maintained. See https://github.com/langchain-ai/langchain-community/issues/674 for details and migration guidance toward standalone integration packages.
  from langchain_community.embeddings import DashScopeEmbeddings
[Document(id='afe08c36-6b1c-410e-bf56-41ca0e2abdcf', metadata={'operator': '小草', 'source': '尺码推荐.txt', 'create_time': '2026-07-26 17:20:32'}, page_content='身高:155-165cm, 体重:75-95 斤,建议尺码S。\n身高:160-170cm, 体重:90-115斤,建议尺码M。\n身高:165-175cm, 体重:115-135斤,建议尺码L。\n身高:170-178cm, 体重:130-150斤,建议尺码XL。\n身高:175-182cm, 体重:145-165斤,建议尺码2XL。\n身高:178-185cm, 体重:160-180斤,建议尺码3XL。\n身高:180-190cm, 体重:180-210斤,建议尺码4XL。\n身高:190cm+,体重:210斤+,建议尺码5XL。')]

Process finished with exit code 0

七、rag服务核心代码开发

开发rag.py

from langchain_core.documents import Document
from langchain_core.runnables import RunnablePassthrough

from vector_stores import VectorStoreService
import config_data as config
from langchain_community.embeddings import DashScopeEmbeddings
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.output_parsers import StrOutputParser
from langchain_community.chat_models 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(
            embeddings=DashScopeEmbeddings(model=config.embedding_model_name),
        )

        self.prompt_template = ChatPromptTemplate.from_messages(
            [
                ("system", "以我提供的一致参考资料为主,简介和专业的回答用户问题。参考资料:{context}。"),
                ("user", "请回答用户提问:{input}")
            ]
        )

        self.chain_model = ChatTongyi(model=config.chat_model)
        self.chain = self.__get_chain()

    def __get_chain(self):
        """
        获取最终的执行链
        :return:
        """
        retriver = self.vector_service.get_retriver()

        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}"

            return formatted_str

        chain = (
            {
                "input": RunnablePassthrough(), "context": retriver | format_document,
            } | self.prompt_template | print_prompt | self.chain_model | StrOutputParser()
        )

        return chain


if __name__ == '__main__':
    res = RagService().chain.invoke("我体重180斤,尺码推荐")
    print(res)
C:\Program_Files\Python\Python31210\python.exe D:\PythonProject\AI_RAG_AGENT\RAG项目案例\rag.py 
====================
System: 以我提供的一致参考资料为主,简介和专业的回答用户问题。参考资料:文档片段:身高:155-165cm, 体重:75-95 斤,建议尺码S。
身高:160-170cm, 体重:90-115斤,建议尺码M。
身高:165-175cm, 体重:115-135斤,建议尺码L。
身高:170-178cm, 体重:130-150斤,建议尺码XL。
身高:175-182cm, 体重:145-165斤,建议尺码2XL。
身高:178-185cm, 体重:160-180斤,建议尺码3XL。
身高:180-190cm, 体重:180-210斤,建议尺码4XL。
身高:190cm+,体重:210斤+,建议尺码5XL。
文档元数据:{'operator': '小草', 'create_time': '2026-07-26 17:20:32', 'source': '尺码推荐.txt'}。
Human: 请回答用户提问:我体重180斤,尺码推荐
====================
根据您提供的体重180斤,结合参考资料中的尺码推荐标准:

- 体重180斤对应的是 **身高180–190cm** 区间,建议尺码为 **4XL**;
- 若您的身高超过190cm且体重仍在180斤或以上,也可能适用 **5XL**。

因此,在一般情况下(身高在180–190cm之间),**推荐尺码为4XL**。如身高超出此范围,请结合具体身高进一步判断。

Process finished with exit code 0

八、历史会话记录功能的实现

创建file_history_store.py

import os, json
from typing import Sequence

from langchain_community.chat_models import ChatTongyi
from langchain_core.messages import message_to_dict, messages_from_dict, BaseMessage
from langchain_core.chat_history import BaseChatMessageHistory
from langchain_core.output_parsers import StrOutputParser
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
from langchain_core.runnables import RunnableWithMessageHistory


def get_history(session_id):
    return FileChatMessageHistory(session_id, "./chat_history")

# message_to_dict:  单个消息对象(BaseMessage类实例) ->  字典
# messages_from_dict: [字典、字典...] -> [消息、消息...]
# 像之前的AIMessage、HumanMessage、SystemMessage都是 BaseMessage的子类
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:
        #     new_messages.append(message_to_dict(message))

        # 有更优的写法,推导式写法
        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装饰器将message方法变成成员属性用
    def messages(self) -> list[BaseMessage]:
        #  当前文件内: list[字典]
        try:
            with open(self.file_path, "r", encoding="utf-8") as f:
                messsage_data = json.load(f)   # 返回值就是: list[字典]
                return messages_from_dict(messsage_data)
        except FileNotFoundError:
            return []


    def clear(self) -> None:
        with open(self.file_path, "w", encoding="utf-8") as f:
            json.dump([], f)



model = ChatTongyi(model="qwen3-max")
# prompt = PromptTemplate.from_template(
#     "你需要根据历史会话回应用户问题。对话历史:{chat_history}, 用户提问:{input},请回答"
# )
prompt = ChatPromptTemplate.from_messages(
    [
        ("system", "你需要根据历史会话回应用户问题。对话历史:"),
        MessagesPlaceholder("chat_history"),
        ("human", "请回答如下问题:{input}")
    ]
)

str_parser = StrOutputParser()

def print_prompt(full_prompt):
    print("="*20, full_prompt.to_string(), "="*20)
    return full_prompt

base_chain = prompt | print_prompt | model | str_parser

store = {}   # key就是session, value就是InMemoryChatMessageHistory类对象

# 实现通过会话id获取InMemoryChatMessageHistory类对象
def get_history(session_id):
    return FileChatMessageHistory(session_id, "./chat_history")


# 创建一个新的链,对原有链增强功能: 自动附加历史消息
conversation_chain = RunnableWithMessageHistory(
    base_chain,   #  被增强的原有chain
    get_history,   #  通过会话id获取InMemoryChatMessageHistory
    input_messages_key="input",          #  表示用户输入在模板中的占位符
    history_messages_key="chat_history"  #  表示用户输入在模板中的占位符
)

if __name__ == '__main__':
    #  固定格式,添加LangChain的配置,为当前程序配置所属的session_id
    session_config = {
        "configurable": {
            "session_id": "user_001"
        }
    }

    # res = conversation_chain.invoke({"input":"小明有2个猫"}, session_config)
    # print("第1次执行", res)
    # res = conversation_chain.invoke({"input": "小张有2只猪"}, session_config)
    # print("第2次执行", res)
    res = conversation_chain.invoke({"input": "总共有几个宠物"}, session_config)
    print("第3次执行", res)

修改rag.py

from langchain_core.documents import Document
from langchain_core.runnables import RunnablePassthrough, RunnableWithMessageHistory, RunnableLambda

from vector_stores import VectorStoreService
import config_data as config
from langchain_community.embeddings import DashScopeEmbeddings
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
from langchain_core.output_parsers import StrOutputParser
from langchain_community.chat_models import ChatTongyi
from file_history_store import get_history

def print_prompt(prompt):
    print("="*20)
    print(prompt.to_string())
    print("=" * 20)
    return prompt

class RagService(object):
    def __init__(self):
        self.vector_service = VectorStoreService(
            embeddings=DashScopeEmbeddings(model=config.embedding_model_name),
        )

        self.prompt_template = ChatPromptTemplate.from_messages(
            [
                ("system", "以我提供的一致参考资料为主,简介和专业的回答用户问题。参考资料:{context}。"),
                ("system", "并且我提供用户的对话历史记录,如下:"),
                MessagesPlaceholder("history"),
                ("user", "请回答用户提问:{input}")
            ]
        )

        self.chain_model = ChatTongyi(model=config.chat_model)
        self.chain = self.__get_chain()

    def __get_chain(self):
        """
        获取最终的执行链
        :return:
        """
        retriver = self.vector_service.get_retriver()

        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}"

            return formatted_str

        def format_for_retriver(value: dict) -> str:
            print("------", value)
            return value["input"]

        def format_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_retriver) | retriver | format_document
            }  | RunnableLambda(format_prompt_template) | self.prompt_template | print_prompt | self.chain_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)
C:\Program_Files\Python\Python31210\python.exe D:\PythonProject\AI_RAG_AGENT\RAG项目案例\rag.py 
D:\PythonProject\AI_RAG_AGENT\RAG项目案例\rag.py:10: LangChainDeprecationWarning: RunnableWithMessageHistory is deprecated. Use LangGraph's built-in persistence instead.
  from file_history_store import get_history
D:\PythonProject\AI_RAG_AGENT\RAG项目案例\rag.py:34: LangChainDeprecationWarning: RunnableWithMessageHistory is deprecated. Use LangGraph's built-in persistence instead.
  self.chain = self.__get_chain()
------ {'input': '我体重180斤,尺码推荐', 'history': [HumanMessage(content='我体重180斤,尺码推荐', additional_kwargs={}, response_metadata={}), AIMessage(content='根据您提供的体重180斤,结合参考资料中的尺码推荐:\n\n- 若您的身高在 **180–190cm** 之间,建议选择 **4XL**;\n- 若您的身高 **超过190cm**,则建议选择 **5XL**。\n\n请根据您的实际身高进一步确认合适尺码。', additional_kwargs={}, response_metadata={}, tool_calls=[], invalid_tool_calls=[])]}
====================
System: 以我提供的一致参考资料为主,简介和专业的回答用户问题。参考资料:文档片段:身高:155-165cm, 体重:75-95 斤,建议尺码S。
身高:160-170cm, 体重:90-115斤,建议尺码M。
身高:165-175cm, 体重:115-135斤,建议尺码L。
身高:170-178cm, 体重:130-150斤,建议尺码XL。
身高:175-182cm, 体重:145-165斤,建议尺码2XL。
身高:178-185cm, 体重:160-180斤,建议尺码3XL。
身高:180-190cm, 体重:180-210斤,建议尺码4XL。
身高:190cm+,体重:210斤+,建议尺码5XL。
文档元数据:{'source': '尺码推荐.txt', 'create_time': '2026-07-26 17:20:32', 'operator': '小草'}。
System: 并且我提供用户的对话历史记录,如下:
Human: 我体重180斤,尺码推荐
AI: 根据您提供的体重180斤,结合参考资料中的尺码推荐:

- 若您的身高在 **180–190cm** 之间,建议选择 **4XL**;
- 若您的身高 **超过190cm**,则建议选择 **5XL**。

请根据您的实际身高进一步确认合适尺码。
Human: 请回答用户提问:我体重180斤,尺码推荐
====================
根据您提供的体重180斤,参考尺码推荐标准:

- 如果您的身高在 **180–190cm** 之间,建议选择 **4XL**;
- 如果您的身高 **超过190cm**,建议选择 **5XL**。

请结合您的实际身高选择最合适的尺码。

Process finished with exit code 0

九、聊天页面开发

import time

import  streamlit as st
from streamlit import session_state
from rag import  RagService
import config_data as config

# 标题
st.title("智能客服")
st.divider()    # 分隔符

# 在页面最下方提供用户输入栏
prompt = st.chat_input()

if "rag" not in st.session_state:
    st.session_state["rag"] = RagService()

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 prompt:

    # 在页面输出用户的提问
    st.chat_message("user").write(prompt)
    st.session_state["message"].append({"role": "user", "content": prompt})

    # 助理给出回答
    ai_res_list = []   # 存ai返回的str数据
    with st.spinner("AI思考中..."):
        # 直接一次性输出
        # res = st.session_state["rag"].chain.invoke({"input": prompt}, config.session_config)
        # st.chat_message("assistant").write(res)
        # st.session_state["message"].append({"role": "assistant", "content": res})
        # 流式输出
        res_stream = st.session_state["rag"].chain.stream({"input": prompt}, config.session_config)

        # res_stream 是Interator, 而content存放的是str,所以,这里就需要特殊处理了
        # 这里只能通过抓包的操作,完成流式数据转str, 使用yield表达式
        def capture(generator, cache_list):
            for chunk in generator:
                cache_list.append(chunk)
                yield chunk

        st.chat_message("assistant").write_stream(capture(res_stream, ai_res_list))
        st.session_state["message"].append({"role": "assistant", "content": "".join(ai_res_list)})  # ["a", "b", "c"] -> "".join(list) -> abc

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值