前 言

后面我想将目前做的这个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

Logo

这里是“一人公司”的成长家园。我们提供从产品曝光、技术变现到法律财税的全栈内容,并连接云服务、办公空间等稀缺资源,助你专注创造,无忧运营。

更多推荐