如何在Streamlit中正确使用LangChain的ConversationChain实现对话记忆?
问题分析与解决方案
你的代码核心问题是没有将对话历史注入到模型的输入Prompt中——虽然你把对话保存到了记忆组件,但生成回答时只传入了当前用户输入,模型无法获取之前的对话记录;同时你创建的ConversationChain没有被真正利用,仅调用了它的LLM实例,浪费了其记忆整合能力。
以下是两种可行的修复方案:
方案一:手动控制记忆注入(灵活适配RAG场景)
这种方式更适合需要结合检索增强(RAG)的场景,手动将历史对话传入Prompt,完全控制流程:
步骤1:添加历史格式化辅助函数
def get_formatted_history(memory): # 从记忆组件中加载并返回格式化后的历史对话字符串 memory_vars = memory.load_memory_variables({}) return memory_vars.get("history", "")
步骤2:修改完整代码
import streamlit as st from langchain.memory import ConversationBufferMemory from langchain_core.runnables import RemoteRunnable, RunnablePassthrough from langchain_core.prompts import ChatPromptTemplate from langchain_core.output_parsers import StrOutputParser # 假设以下变量已提前定义:LANGSERVE_ENDPOINT, RAG_PROMPT_TEMPLATE, retriever, format_docs, add_history, file # 初始化会话记忆 if 'conversation_memory' not in st.session_state: st.session_state.conversation_memory = ConversationBufferMemory(human_prefix="user", ai_prefix="ai") memory = st.session_state.conversation_memory llm = RemoteRunnable(LANGSERVE_ENDPOINT) if user_input := st.chat_input(): add_history("user", user_input) st.chat_message("user").write(user_input) with st.chat_message("assistant"): chat_container = st.empty() full_answer = "" formatted_history = get_formatted_history(memory) if file is not None: # 带历史的RAG Prompt模板 rag_prompt = ChatPromptTemplate.from_template("""根据历史对话和上下文回答用户问题: 历史对话: {history} 上下文信息: {context} 当前问题:{question} """) rag_chain = ( { "context": retriever | format_docs, "question": RunnablePassthrough(), "history": lambda _: formatted_history } | rag_prompt | llm | StrOutputParser() ) # 流式生成回答 for chunk in rag_chain.stream(user_input): full_answer += chunk chat_container.markdown(full_answer) else: # 带历史的普通对话Prompt模板 general_prompt = ChatPromptTemplate.from_template("""根据历史对话简要回答当前问题: 历史对话: {history} 当前问题:{input} """) general_chain = ( { "input": RunnablePassthrough(), "history": lambda _: formatted_history } | general_prompt | llm | StrOutputParser() ) for chunk in general_chain.stream(user_input): full_answer += chunk chat_container.markdown(full_answer) # 统一保存对话到记忆 memory.save_context( inputs={memory.human_prefix: user_input}, outputs={memory.ai_prefix: full_answer} ) add_history("ai", full_answer)
方案二:利用ConversationChain自动管理记忆(适合纯对话场景)
如果不需要RAG,直接用ConversationChain可以自动处理记忆的注入和保存,减少手动代码:
修改后的代码
import streamlit as st from langchain.memory import ConversationBufferMemory from langchain.chains import ConversationChain from langchain_core.runnables import RemoteRunnable from langchain_core.prompts import ChatPromptTemplate # 初始化会话记忆和ConversationChain if 'conversation_memory' not in st.session_state: st.session_state.conversation_memory = ConversationBufferMemory(human_prefix="user", ai_prefix="ai") memory = st.session_state.conversation_memory llm = RemoteRunnable(LANGSERVE_ENDPOINT) # 定义带历史的对话Prompt general_prompt = ChatPromptTemplate.from_template("""根据历史对话回答当前问题: 历史对话:{history} 当前问题:{input} """) conversation = ConversationChain( memory=memory, llm=llm, prompt=general_prompt, verbose=False ) if user_input := st.chat_input(): add_history("user", user_input) st.chat_message("user").write(user_input) with st.chat_message("assistant"): chat_container = st.empty() full_answer = "" # 调用ConversationChain,自动注入历史并保存结果到记忆 result = conversation({"input": user_input}) full_answer = result["response"] chat_container.markdown(full_answer) add_history("ai", full_answer)
关键修复点说明
- Prompt必须包含历史变量:所有生成回答的Prompt都要加入
{history}占位符,让模型能获取之前的对话记录。 - 历史注入到链输入:通过链式结构将格式化后的历史对话传入Prompt,确保模型生成回答时参考历史。
- 记忆保存的一致性:使用记忆组件的
human_prefix和ai_prefix属性,避免手动写死键名导致的不匹配。
内容的提问来源于stack exchange,提问作者Hojoon Kim
相关产品推荐
相关产品推荐

