From b072a14ae23260b5e72789dc4d41af778dd5c27d Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Thu, 22 Jan 2026 02:38:37 +0800 Subject: [PATCH] =?UTF-8?q?refactor(agents):=20=E7=A7=BB=E9=99=A4graph?= =?UTF-8?q?=E7=BC=93=E5=AD=98=E5=B9=B6=E4=BC=98=E5=8C=96=E8=BF=9E=E6=8E=A5?= =?UTF-8?q?=E5=92=8C=E6=A3=80=E6=9F=A5=E7=82=B9=E7=BC=93=E5=AD=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 移除所有agent子类中graph的缓存逻辑,改为每次都重新构建 在BaseAgent中添加_async_conn缓存并优化checkpointer缓存逻辑 --- src/agents/chatbot/graph.py | 4 ---- src/agents/common/base.py | 13 +++++++++++-- src/agents/deep_agent/graph.py | 4 ---- src/agents/mini_agent/graph.py | 4 ---- src/agents/reporter/graph.py | 4 ---- 5 files changed, 11 insertions(+), 18 deletions(-) diff --git a/src/agents/chatbot/graph.py b/src/agents/chatbot/graph.py index 810bdb1a..09b50aa9 100644 --- a/src/agents/chatbot/graph.py +++ b/src/agents/chatbot/graph.py @@ -18,9 +18,6 @@ class ChatbotAgent(BaseAgent): async def get_graph(self, **kwargs): """构建图""" - if self.graph: - return self.graph - # 获取上下文配置 context = self.context_schema.from_file(module_name=self.module_name) @@ -36,7 +33,6 @@ class ChatbotAgent(BaseAgent): checkpointer=await self._get_checkpointer(), ) - self.graph = graph return graph diff --git a/src/agents/common/base.py b/src/agents/common/base.py index 25a61b9a..d3b8c04a 100644 --- a/src/agents/common/base.py +++ b/src/agents/common/base.py @@ -28,6 +28,7 @@ class BaseAgent: def __init__(self, **kwargs): self.graph = None # will be covered by get_graph self.checkpointer = None + self._async_conn = None self.workdir = Path(sys_config.save_dir) / "agents" / self.module_name self.workdir.mkdir(parents=True, exist_ok=True) self._metadata_cache = None # Cache for metadata to avoid repeated file reads @@ -147,6 +148,9 @@ class BaseAgent: pass async def _get_checkpointer(self): + if self.checkpointer is not None: + return self.checkpointer + # 创建数据库连接并确保设置 checkpointer checkpointer = None @@ -157,15 +161,20 @@ class BaseAgent: logger.error(f"构建 Graph 设置 checkpointer 时出错: {e}, 尝试使用内存存储") checkpointer = InMemorySaver() - return checkpointer + self.checkpointer = checkpointer + return self.checkpointer async def get_async_conn(self) -> aiosqlite.Connection: """获取异步数据库连接""" + if self._async_conn is not None: + return self._async_conn + conn = await aiosqlite.connect(os.path.join(self.workdir, "aio_history.db")) # Patch: langgraph's AsyncSqliteSaver expects is_alive() method which aiosqlite may not have if not hasattr(conn, "is_alive"): conn.is_alive = lambda: True - return conn + self._async_conn = conn + return self._async_conn async def get_aio_memory(self) -> AsyncSqliteSaver: """获取异步存储实例""" diff --git a/src/agents/deep_agent/graph.py b/src/agents/deep_agent/graph.py index 8a63cb01..d865ee0d 100644 --- a/src/agents/deep_agent/graph.py +++ b/src/agents/deep_agent/graph.py @@ -88,9 +88,6 @@ class DeepAgent(BaseAgent): async def get_graph(self, **kwargs): """构建 Deep Agent 的图""" - if self.graph: - return self.graph - # 获取上下文配置 context = self.context_schema.from_file(module_name=self.module_name) @@ -139,5 +136,4 @@ class DeepAgent(BaseAgent): checkpointer=await self._get_checkpointer(), ) - self.graph = graph return graph diff --git a/src/agents/mini_agent/graph.py b/src/agents/mini_agent/graph.py index 6271508b..ff38019e 100644 --- a/src/agents/mini_agent/graph.py +++ b/src/agents/mini_agent/graph.py @@ -12,9 +12,6 @@ class MiniAgent(BaseAgent): super().__init__(**kwargs) async def get_graph(self, **kwargs): - if self.graph: - return self.graph - context = self.context_schema.from_file(module_name=self.module_name) # 创建 MiniAgent @@ -25,5 +22,4 @@ class MiniAgent(BaseAgent): checkpointer=await self._get_checkpointer(), ) - self.graph = graph return graph diff --git a/src/agents/reporter/graph.py b/src/agents/reporter/graph.py index 15cbcdd2..d72a9a00 100644 --- a/src/agents/reporter/graph.py +++ b/src/agents/reporter/graph.py @@ -35,9 +35,6 @@ class SqlReporterAgent(BaseAgent): super().__init__(**kwargs) async def get_graph(self, **kwargs): - if self.graph: - return self.graph - context = self.context_schema.from_file(module_name=self.module_name) # 创建 SqlReporterAgent @@ -48,6 +45,5 @@ class SqlReporterAgent(BaseAgent): checkpointer=await self._get_checkpointer(), ) - self.graph = graph logger.info("SqlReporterAgent 构建成功") return graph