You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为LangChain中的SQL Agent添加对话记忆功能?

解决LangChain SQL Agent记忆功能失效问题

核心问题分析

你的代码里虽然初始化了ConversationBufferMemory并传入Agent,但存在两个关键缺失:

  • 自定义full_prompt中未添加对话历史占位符,导致Agent无法读取过往交互内容
  • 调用agent_executor.invoke时仅传入当前用户输入,未传递历史上下文给Agent

具体修复步骤

1. 在Prompt中添加对话历史占位符

修改full_prompt,在用户输入前加入MessagesPlaceholder来加载对话历史:

full_prompt = ChatPromptTemplate.from_messages(
    [
        SystemMessagePromptTemplate(prompt=few_shot_prompt),
        MessagesPlaceholder(variable_name="chat_history"),  # 新增:加载对话历史
        ("human", "{input}"),
        MessagesPlaceholder("agent_scratchpad"),
    ]
)

2. 绑定记忆到Streamlit会话状态

将memory对象存入st.session_state,避免页面刷新时重置记忆:

# 替换原memory初始化代码
if "memory" not in st.session_state:
    st.session_state.memory = ConversationBufferMemory(
        memory_key="chat_history", 
        input_key="input",
        return_messages=True,
        k=5
    )
memory = st.session_state.memory

3. 调用Agent时传入完整上下文

修改safe_run_query函数,调用invoke时传入包含历史的字典,并手动保存对话到记忆:

def safe_run_query(query):
    try:
        # 传入对话历史和当前输入
        result = st.session_state.agent_executor.invoke({
            "input": query,
            "chat_history": st.session_state.memory.chat_memory.messages
        })
        # 更新记忆
        st.session_state.memory.save_context({"input": query}, {"output": result['output']})
        return result
    except Exception as e:
        st.error("Failed to execute query: " + str(e))
        return None

4. 移除冗余的prefix参数

创建Agent时的prefix=sql_agent_query会覆盖系统提示,需删除:

st.session_state.agent_executor = create_sql_agent(
    llm=llm,
    toolkit=toolkit,
    prompt=full_prompt,
    agent_type="openai-tools",
    agent_executor_kwargs={"memory": memory,  "return_intermediate_steps": True},
    verbose=True,
)

完整修复后的代码

from langchain.agents import create_sql_agent
from langchain_community.agent_toolkits import SQLDatabaseToolkit
from langchain_community.utilities.sql_database import SQLDatabase
from langchain_openai import ChatOpenAI
import streamlit as st
from langchain_core.messages import AIMessage, HumanMessage
from langchain.memory import ConversationBufferMemory
from langchain_core.prompts import (
    ChatPromptTemplate,
    FewShotPromptTemplate,
    MessagesPlaceholder,
    PromptTemplate,
    SystemMessagePromptTemplate,
)
from examples import get_example_selector
from dotenv import load_dotenv
 
load_dotenv()
 
# 初始化数据库连接
def init_database(username, password, server, database, port) -> SQLDatabase:
    db_uri = f"mysql+mysqlconnector://{username}:{password}@{server}:{port}/{database}"
    return SQLDatabase.from_uri(db_uri)

# Streamlit界面初始化
st.title("LangChain SQL Agent - With Memory")
with st.sidebar:
    on_show_query = st.toggle('Show Query Mode')
    st.subheader("Database Configuration")
    username = st.text_input("User", key="User")
    password = st.text_input("Password",type="password", key="Password")
    server = st.text_input("Server", key="Server")
    port = st.text_input("Port", key="Port")
    database = st.text_input("DB Name", key="Database")
    connect_button = st.button("Connect")

# 初始化会话状态变量
if "agent_executor" not in st.session_state:
    st.session_state.agent_executor = None

if "chat_history" not in st.session_state:
        st.session_state.chat_history = [
            AIMessage(content="Hello! I'm the SQL Agent. How can I help you today?")
        ]

# 连接数据库逻辑
if connect_button:
    sql_database = init_database(username, password, server, database, port)
    llm = ChatOpenAI(temperature=0, model_name='gpt-4')
    toolkit = SQLDatabaseToolkit(db=sql_database, llm=llm)
    
    system_prefix="""You are an agent designed to interact with a SQL database.
    Given an input question, create a syntactically correct MariaDB query to run, then look at the results of the query and return the answer.
    You can order the results by a relevant column to return the most interesting examples in the database.
    Never query for all the columns from a specific table, only ask for the relevant columns given the question.
    You have access to tools for interacting with the database.
    Only use the given tools. Only use the information returned by the tools to construct your final answer.
    You MUST double check your query before executing it. If you get an error while executing a query, rewrite the query and try again.

    DO NOT make any DML statements (INSERT, UPDATE, DELETE, DROP etc.) to the database.

    If the question does not seem related to the database, just return "I don't know" as the answer.

    Here are some examples of user inputs and their corresponding SQL queries:"""
 
 
    few_shot_prompt = FewShotPromptTemplate(
        example_selector=get_example_selector(),
        example_prompt=PromptTemplate.from_template(
            "User input: {input}\nSQL query: {query}"
        ),
        input_variables=["input"],
        prefix=system_prefix,
        suffix="If the user's question is not related to any of the examples, you do not have to use them",
    )

    full_prompt = ChatPromptTemplate.from_messages(
        [
            SystemMessagePromptTemplate(prompt=few_shot_prompt),
            MessagesPlaceholder(variable_name="chat_history"),  # 新增对话历史占位符
            ("human", "{input}"),
            MessagesPlaceholder("agent_scratchpad"),
        ]
    )
 
    # 初始化记忆并存入session_state
    if "memory" not in st.session_state:
        st.session_state.memory = ConversationBufferMemory(
            memory_key="chat_history", 
            input_key="input",
            return_messages=True,
            k=5
        )
    memory = st.session_state.memory
    
    st.session_state.agent_executor = create_sql_agent(
        llm=llm,
        toolkit=toolkit,
        prompt=full_prompt,
        agent_type="openai-tools",
        agent_executor_kwargs={"memory": memory,  "return_intermediate_steps": True},
        verbose=True,
    )
    st.session_state["db_connected"] = True
    st.success("Database connected successfully!")
 
# 安全执行查询函数
def safe_run_query(query):
    try:
        # 传入完整上下文
        result = st.session_state.agent_executor.invoke({
            "input": query,
            "chat_history": st.session_state.memory.chat_memory.messages
        })
        # 手动保存对话到记忆
        st.session_state.memory.save_context({"input": query}, {"output": result['output']})
        return result
    except Exception as e:
        st.error("Failed to execute query: " + str(e))
        return None
 
# 渲染聊天历史
for message in st.session_state.chat_history:
  if isinstance(message, AIMessage):
    with st.chat_message("AI"):
      st.markdown(message.content)
  elif isinstance(message, HumanMessage):
    with st.chat_message("Human"):
      st.markdown(message.content)
 
# 处理用户输入
if "db_connected" in st.session_state and st.session_state["db_connected"]:
    
    user_input = st.chat_input("Enter your question:")
    if user_input is not None and user_input.strip() != "":
        st.session_state.chat_history.append(HumanMessage(content=user_input))
    
        with st.chat_message("Human"):
            st.markdown(user_input)
            
        with st.spinner("Generating response..."):
            with st.chat_message("AI"):
                    response = safe_run_query(user_input)
                    if not response:
                        st.markdown("Failed to get response.")
                        continue
                    # 提取输出内容
                    output_message = response.get('output', 'No output found')
                    
                    st.markdown(output_message)
 
                    intermediate_steps = response['intermediate_steps']
 
                    # 提取SQL查询(如果开启显示模式)
                    output_query = None
                    if on_show_query:
                        for step in intermediate_steps:
                            action, result = step
                            st.text(action.log)
                            action, tool_input = step
                            if action.tool == 'sql_db_query':
                                output_query = action.tool_input
                                break
 
                        if output_query:
                            st.code(output_query)
                          
            st.session_state.chat_history.append(AIMessage(content=output_message))
 
else:
    st.warning("Please connect to the database first.")

验证方法

  1. 先问基础问题,比如"用户表有多少条数据?"
  2. 接着问关联问题,比如"其中活跃用户有多少?"
  3. 观察Agent是否能基于上一个问题的上下文生成正确的SQL查询

内容的提问来源于stack exchange,提问作者Reon

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.25 11:08:09