12.RAG/Agent系统升级:基于session的会话状态管理与多轮对话隔离
·
前 言
后面我想将目前做的这个RAG智能体系统通过Django框架做成一个网站部署,但是在将这个系统上Django之前,还需要做一步会话级状态管理(Session-based chat state),这样做的原因在于:希望未来登录到系统的用户,他们之间的会话不会互相污染。
第一步:建立会话管理
首先第一步,需要建立一个会话管理类,在这个类中需要做绕不开的基础业务:增、删、改、查。
class SessionManager:
def __init__(self, max_turns=3):
self.sessions = {}
self.max_turns = max_turns
def get_history(self, session_id: str): # 通过传入的会话id查询会话
if session_id not in self.sessions:
self.sessions[session_id] = []
return self.sessions[session_id]
def append_turn(self, session_id: str, user_message: str, assistant_message: str): #查询时候又会话,将用户问题还有历史会话拼接
history = self.get_history(session_id)
history.append({"role": "user", "content": user_message})
history.append({"role": "assistant", "content": assistant_message})
self.trim_history(session_id)
def trim_history(self, session_id: str): # 做会话限制,避免历史信息太多
history = self.get_history(session_id)
max_messages = self.max_turns * 2
if len(history) > max_messages:
self.sessions[session_id] = history[-max_messages:]
def clear_session(self, session_id: str): # 清楚会话
self.sessions[session_id] = []
第二步:改RAGSystem类
前面咱们写RAGSystem类,历史对话都是保存在类内的,这样等于是写死的,导致同一个类不能处理不同的历史会话,所以咱们需要将self.chat_history相关的内容都去掉,改成这样:
def ask(self, question, chat_history=None):
if chat_history is None:
chat_history = []
retrieved = self.retrieve(question, k=self.top_k)
texts = [c["text"] for c in retrieved]
sorted_indices = self.rerank(question, texts)
best_chunks = [retrieved[i] for i in sorted_indices[:self.rerank_k]]
context = ""
for c in best_chunks:
context += f"[Source: {c['source']}]\n{c['text']}\n\n"
messages = [
{
"role": "system",
"content": "You are a helpful assistant. Answer based on context and conversation history."
}
]
messages.extend(chat_history)
messages.append({
"role": "user",
"content": f"{context}\n\nQuestion: {question}"
})
response = client.chat.completions.create(
model=CHAT_MODEL,
messages=messages
)
answer = response.choices[0].message.content
return answer
第三步:改工具
类方法已经修改了,咱们在tools.py文件中对应的工具调用方法也需要进行修改,尤其是函数调用入口,传参的位置要传入来自不同会话的历史消息:
from datetime import datetime
from llm_utils import client
from config import CHAT_MODEL
def rag_tool(query, rag, chat_history=None):
return rag.ask(query, chat_history=chat_history)
def calculator_tool(expression):
try:
return str(eval(expression))
except:
return "Invalid expression"
def time_tool(_):
return datetime.now().strftime("%Y-%m-%d %H:%M:%S")
def llm_tool(query, chat_history=None):
messages = [
{"role": "system", "content": "You are a helpful assistant."}
]
if chat_history:
messages.extend(chat_history)
messages.append({"role": "user", "content": query})
response = client.chat.completions.create(
model=CHAT_MODEL,
messages=messages
)
return response.choices[0].message.content
TOOLS = [
{
"name": "rag",
"description": "Use for paper/document questions",
"func": rag_tool
},
{
"name": "calculator",
"description": "Use for math calculations",
"func": calculator_tool
},
{
"name": "time",
"description": "Use to get current time",
"func": time_tool
},
{
"name": "llm",
"description": "Use for general questions",
"func": llm_tool
}
]
第四步:改智能体
相应的智能体相关的调用部分,尤其是函数入口的变量部分也需要进行修改,传入来自不同会话的历史消息:
def execute_tool(decision, tools, rag=None, chat_history=None):
tool_name = decision["tool"]
tool_input = decision["input"]
for t in tools:
if t["name"] == tool_name:
if tool_name == "rag":
result = t["func"](tool_input, rag, chat_history=chat_history)
elif tool_name == "llm":
result = t["func"](tool_input, chat_history=chat_history)
else:
result = t["func"](tool_input)
return {
"tool_name": tool_name,
"tool_input": tool_input,
"tool_output": result
}
return {
"tool_name": "none",
"tool_input": tool_input,
"tool_output": "No valid tool found."
}
def run_agent(query, tools, rag=None, chat_history=None):
logger.info("Agent Start")
logger.info(f"User query: {query}")
decision = choose_tool(query, tools)
logger.info(f"Decision: {decision}")
tool_result = execute_tool(decision, tools, rag=rag, chat_history=chat_history)
logger.info(f"Tool result: {tool_result}")
final_answer = generate_final_answer(query, tool_result)
logger.info("Agent End")
return final_answer
第五步:改FastAPI
最后需要修改咱们的app.py,首先在开头导入咱们写的会话管理类:
from session_manager import SessionManager
然后初始化该类:
session_manager = SessionManager(max_turns=3)
这里有个坑要注意,咱们这回要通过post请求不止传输问题了,还要传会话id,所以咱们的请求模型也要修改:
class QueryRequest(BaseModel):
session_id: str
question: str
然后就是修改/ask请求逻辑了:
@app.post("/ask")
def ask_question(req: QueryRequest):
try:
history = session_manager.get_history(req.session_id) # 先获取历史对话列表
answer = run_agent(
req.question,
TOOLS,
rag=rag,
chat_history=history
)
session_manager.append_turn(
req.session_id,
req.question,
answer
)
return {
"session_id": req.session_id,
"question": req.question,
"answer": answer
}
except Exception as e:
logger.exception("Error occurred in /ask")
raise HTTPException(status_code=500, detail=str(e))
最后,再加一个清空会话的接口:
@app.post("/clear/{session_id}")
def clear_session(session_id: str):
session_manager.clear_session(session_id)
return {
"session_id": session_id,
"message": "session cleared"
}
注意到了吗? 咱们原来写到RAGSystem类中的那些会话操作,咱们现在都给拆出来放到了会话管理类中了,实现了解耦,这样咱们就可以用同一个RAGSystem处理来自不同会话的请求了,而且会话之间不会存在信息污染。
如果这篇文章对你有帮助,可以点个赞~
完整代码地址:https://github.com/1186141415/A-Paper-Rag-Agent
更多推荐



所有评论(0)