From ff955bdda53bdb59c45f9ad4fe2810a2a96f14f5 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Sat, 11 Oct 2025 21:10:39 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E7=A7=BB=E9=99=A4=20Async=E7=9A=84Chec?= =?UTF-8?q?kpointer=EF=BC=8C=E5=B9=B6=E9=BB=98=E8=AE=A4=E9=85=8D=E7=BD=AE?= =?UTF-8?q?=E4=BA=86=20InMemorySaver?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/agents/chatbot/graph.py | 31 ++++++------------------------- src/agents/common/base.py | 4 +++- src/agents/react/graph.py | 6 ++---- 3 files changed, 11 insertions(+), 30 deletions(-) diff --git a/src/agents/chatbot/graph.py b/src/agents/chatbot/graph.py index 5b138695..766c6a02 100644 --- a/src/agents/chatbot/graph.py +++ b/src/agents/chatbot/graph.py @@ -1,11 +1,9 @@ -import os import uuid from pathlib import Path from typing import Any, cast from langchain_core.messages import AIMessage, ToolMessage from langgraph.checkpoint.memory import InMemorySaver -from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver, aiosqlite from langgraph.graph import END, START, StateGraph from langgraph.prebuilt import ToolNode, tools_condition from langgraph.runtime import Runtime @@ -28,6 +26,7 @@ class ChatbotAgent(BaseAgent): def __init__(self, **kwargs): super().__init__(**kwargs) self.graph = None + self.checkpointer = InMemorySaver() self.context_schema = Context self.workdir = Path(sys_config.save_dir) / "agents" / self.module_name self.workdir.mkdir(parents=True, exist_ok=True) @@ -90,8 +89,8 @@ class ChatbotAgent(BaseAgent): async def get_graph(self, **kwargs): """构建图""" - # if self.graph: - # return self.graph + if self.graph: + return self.graph builder = StateGraph(State, context_schema=self.context_schema) builder.add_node("chatbot", self.llm_call) @@ -104,27 +103,9 @@ class ChatbotAgent(BaseAgent): builder.add_edge("tools", "chatbot") builder.add_edge("chatbot", END) - # 创建数据库连接并确保设置 checkpointer - try: - sqlite_checkpointer = AsyncSqliteSaver(await self.get_async_conn()) - graph = builder.compile(checkpointer=sqlite_checkpointer, name=self.name) - self.graph = graph - return graph - except Exception as e: - logger.error(f"构建 Graph 设置 checkpointer 时出错: {e}, 尝试使用内存存储") - # 即使出错也返回一个可用的图实例,只是无法保存历史 - checkpointer = InMemorySaver() - graph = builder.compile(checkpointer=checkpointer, name=self.name) - self.graph = graph - return graph - - async def get_async_conn(self) -> aiosqlite.Connection: - """获取异步数据库连接""" - return await aiosqlite.connect(os.path.join(self.workdir, "aio_history.db")) - - async def get_aio_memory(self) -> AsyncSqliteSaver: - """获取异步存储实例""" - return AsyncSqliteSaver(await self.get_async_conn()) + graph = builder.compile(checkpointer=self.checkpointer, name=self.name) + self.graph = graph + return graph def main(): diff --git a/src/agents/common/base.py b/src/agents/common/base.py index c784071f..99a37f07 100644 --- a/src/agents/common/base.py +++ b/src/agents/common/base.py @@ -3,6 +3,7 @@ from __future__ import annotations from abc import abstractmethod from langgraph.graph.state import CompiledStateGraph +from langgraph.checkpoint.memory import InMemorySaver from src.agents.common.context import BaseContext from src.utils import logger @@ -18,6 +19,7 @@ class BaseAgent: def __init__(self, **kwargs): self.graph = None # will be covered by get_graph + self.checkpointer = InMemorySaver() self.context_schema = BaseContext @property @@ -52,7 +54,7 @@ class BaseAgent: graph = await self.get_graph() context = self.context_schema.from_file(module_name=self.module_name, input_context=input_context) logger.debug(f"stream_messages: {context}") - # TODO 的 Checkpointer 似乎还没有适配最新的 Context API + # TODO Checkpointer 似乎还没有适配最新的 1.0 Context API input_config = {"configurable": input_context, "recursion_limit": 100} async for msg, metadata in graph.astream( {"messages": messages}, stream_mode="messages", context=context, config=input_config diff --git a/src/agents/react/graph.py b/src/agents/react/graph.py index 774b6a64..3dc0e494 100644 --- a/src/agents/react/graph.py +++ b/src/agents/react/graph.py @@ -1,7 +1,6 @@ from pathlib import Path from langchain_core.messages import AnyMessage, SystemMessage -from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver, aiosqlite from langgraph.prebuilt import create_react_agent from langgraph.runtime import get_runtime @@ -37,8 +36,7 @@ class ReActAgent(BaseAgent): available_tools = get_buildin_tools() - sqlite_checkpointer = AsyncSqliteSaver(await aiosqlite.connect(self.workdir / "react_history.db")) - graph = create_react_agent(model, tools=available_tools, checkpointer=sqlite_checkpointer, prompt=prompt) + graph = create_react_agent(model, tools=available_tools, prompt=prompt, checkpointer=self.checkpointer) self.graph = graph - logger.info("ReActAgent使用SQLite checkpointer构建成功") + logger.info("ReActAgent 使用内存 checkpointer 构建成功") return graph