From 3600c7f9b7f06675bd56482b808fa089be0dd91f Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Sun, 12 Oct 2025 16:04:46 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E4=BF=AE=E5=A4=8D=E9=83=A8=E5=88=86?= =?UTF-8?q?=E5=9C=BA=E6=99=AF=E4=B8=8B=E5=8E=86=E5=8F=B2=E8=AE=B0=E5=BD=95?= =?UTF-8?q?=E4=BF=9D=E5=AD=98=E5=BC=82=E5=B8=B8=E7=9A=84=E9=97=AE=E9=A2=98?= =?UTF-8?q?=E3=80=82=E5=9B=9E=E9=80=80=E5=88=B0=E4=BD=BF=E7=94=A8=20AsyncS?= =?UTF-8?q?aver=E7=9A=84=E5=8F=8C=E7=BA=BF=E5=A4=87=E4=BB=BD=E6=B6=88?= =?UTF-8?q?=E6=81=AF=E6=A8=A1=E5=BC=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- server/routers/chat_router.py | 2 +- src/agents/chatbot/graph.py | 15 +++------------ src/agents/common/base.py | 31 ++++++++++++++++++++++++++++++- src/agents/react/graph.py | 2 ++ src/storage/db/models.py | 6 ++++++ 5 files changed, 42 insertions(+), 14 deletions(-) diff --git a/server/routers/chat_router.py b/server/routers/chat_router.py index 32a742d8..bbd5157f 100644 --- a/server/routers/chat_router.py +++ b/server/routers/chat_router.py @@ -166,7 +166,7 @@ async def chat_agent( existing_count = len(existing_messages) # 只保存新增的消息 - new_messages = messages + new_messages = messages[existing_count :] for msg in new_messages: msg_dict = msg.model_dump() if hasattr(msg, "model_dump") else {} diff --git a/src/agents/chatbot/graph.py b/src/agents/chatbot/graph.py index 766c6a02..d94552a1 100644 --- a/src/agents/chatbot/graph.py +++ b/src/agents/chatbot/graph.py @@ -26,10 +26,8 @@ class ChatbotAgent(BaseAgent): def __init__(self, **kwargs): super().__init__(**kwargs) self.graph = None - self.checkpointer = InMemorySaver() + self.checkpointer = None self.context_schema = Context - self.workdir = Path(sys_config.save_dir) / "agents" / self.module_name - self.workdir.mkdir(parents=True, exist_ok=True) self.agent_tools = None def get_tools(self): @@ -103,21 +101,14 @@ class ChatbotAgent(BaseAgent): builder.add_edge("tools", "chatbot") builder.add_edge("chatbot", END) + self.checkpointer = await self._get_checkpointer() graph = builder.compile(checkpointer=self.checkpointer, name=self.name) self.graph = graph return graph def main(): - agent = ChatbotAgent(Context) - - thread_id = str(uuid.uuid4()) - config = {"configurable": {"thread_id": thread_id}} - - from src.agents.utils import agent_cli - - agent_cli(agent, config) - + pass if __name__ == "__main__": main() diff --git a/src/agents/common/base.py b/src/agents/common/base.py index 99a37f07..48e87c83 100644 --- a/src/agents/common/base.py +++ b/src/agents/common/base.py @@ -1,10 +1,14 @@ from __future__ import annotations +import os +from pathlib import Path from abc import abstractmethod from langgraph.graph.state import CompiledStateGraph from langgraph.checkpoint.memory import InMemorySaver +from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver, aiosqlite +from src import config as sys_config from src.agents.common.context import BaseContext from src.utils import logger @@ -19,8 +23,10 @@ class BaseAgent: def __init__(self, **kwargs): self.graph = None # will be covered by get_graph - self.checkpointer = InMemorySaver() + self.checkpointer = None self.context_schema = BaseContext + self.workdir = Path(sys_config.save_dir) / "agents" / self.module_name + self.workdir.mkdir(parents=True, exist_ok=True) @property def module_name(self) -> str: @@ -103,3 +109,26 @@ class BaseAgent: 例如: graph = workflow.compile(checkpointer=sqlite_checkpointer) """ pass + + async def _get_checkpointer(self): + # 创建数据库连接并确保设置 checkpointer + checkpointer = None + + try: + checkpointer = AsyncSqliteSaver(await self.get_async_conn()) + + except Exception as e: + logger.error(f"构建 Graph 设置 checkpointer 时出错: {e}, 尝试使用内存存储") + checkpointer = InMemorySaver() + + return checkpointer + + + 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()) + diff --git a/src/agents/react/graph.py b/src/agents/react/graph.py index 3dc0e494..33e55b21 100644 --- a/src/agents/react/graph.py +++ b/src/agents/react/graph.py @@ -35,7 +35,9 @@ class ReActAgent(BaseAgent): return self.graph available_tools = get_buildin_tools() + self.checkpointer = await self._get_checkpointer() + # 创建 ReActAgent graph = create_react_agent(model, tools=available_tools, prompt=prompt, checkpointer=self.checkpointer) self.graph = graph logger.info("ReActAgent 使用内存 checkpointer 构建成功") diff --git a/src/storage/db/models.py b/src/storage/db/models.py index e6823680..b0490490 100644 --- a/src/storage/db/models.py +++ b/src/storage/db/models.py @@ -83,6 +83,12 @@ class Message(Base): "tool_calls": [tc.to_dict() for tc in self.tool_calls] if self.tool_calls else [], } + def to_simple_dict(self): + return { + "role": self.role, + "content": self.content, + } + class ToolCall(Base): """ToolCall table - stores tool invocations"""