feat: Agent 消息保存在SQLite 中

This commit is contained in:
Wenjie Zhang 2025-05-16 23:46:25 +08:00
parent 6212f5c456
commit 77ce5a25da
7 changed files with 61 additions and 30 deletions

View File

@ -13,6 +13,7 @@ dependencies = [
"langchain-huggingface>=0.2.0", "langchain-huggingface>=0.2.0",
"langchain-openai>=0.3.14", "langchain-openai>=0.3.14",
"langgraph>=0.3.34", "langgraph>=0.3.34",
"langgraph-checkpoint-sqlite>=2.0.7",
"langsmith>=0.3.37", "langsmith>=0.3.37",
"llama-index>=0.12.33", "llama-index>=0.12.33",
"llama-index-readers-file>=0.4.7", "llama-index-readers-file>=0.4.7",

View File

@ -165,7 +165,7 @@ async def get_agent(current_user: User = Depends(get_required_user)):
return {"agents": agents} return {"agents": agents}
@chat.post("/agent/{agent_name}") @chat.post("/agent/{agent_name}")
def chat_agent(agent_name: str, async def chat_agent(agent_name: str,
query: str = Body(...), query: str = Body(...),
config: dict = Body({}), config: dict = Body({}),
meta: dict = Body({}), meta: dict = Body({}),
@ -189,7 +189,7 @@ def chat_agent(agent_name: str,
**kwargs **kwargs
}, ensure_ascii=False).encode('utf-8') + b"\n" }, ensure_ascii=False).encode('utf-8') + b"\n"
def stream_messages(): async def stream_messages():
# 代表服务端已经收到了请求 # 代表服务端已经收到了请求
yield make_chunk(status="init", meta=meta, msg=HumanMessage(content=query).model_dump()) yield make_chunk(status="init", meta=meta, msg=HumanMessage(content=query).model_dump())
@ -212,7 +212,7 @@ def chat_agent(agent_name: str,
runnable_config = {"configurable": {**config}} runnable_config = {"configurable": {**config}}
try: try:
for msg, metadata in agent.stream_messages(messages, config_schema=runnable_config): async for msg, metadata in agent.stream_messages(messages, config_schema=runnable_config):
logger.debug(f"msg: {msg.model_dump()}, metadata: {metadata}") logger.debug(f"msg: {msg.model_dump()}, metadata: {metadata}")
if isinstance(msg, AIMessageChunk): if isinstance(msg, AIMessageChunk):
yield make_chunk(content=msg.content, yield make_chunk(content=msg.content,
@ -289,7 +289,7 @@ async def get_agent_history(
raise HTTPException(status_code=404, detail=f"智能体 {agent_name} 不存在") raise HTTPException(status_code=404, detail=f"智能体 {agent_name} 不存在")
# 获取历史消息 # 获取历史消息
history = agent.get_history(user_id=current_user.id, thread_id=thread_id) history = await agent.get_history(user_id=current_user.id, thread_id=thread_id)
return {"history": history} return {"history": history}
except Exception as e: except Exception as e:

View File

@ -1,16 +1,16 @@
import os import os
import uuid import uuid
from typing import Any from typing import Any
from pathlib import Path
from datetime import datetime, timezone from datetime import datetime, timezone
import sqlite3
from langchain_core.runnables import RunnableConfig from langchain_core.runnables import RunnableConfig
from langgraph.graph import StateGraph, START, END from langgraph.graph import StateGraph, START, END
from langgraph.prebuilt import ToolNode, tools_condition from langgraph.prebuilt import ToolNode, tools_condition
from langgraph.checkpoint.memory import InMemorySaver from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver, aiosqlite
# from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver, aiosqlite
from src import config as sys_config
from src.utils import logger from src.utils import logger
from src.agents.registry import State, BaseAgent from src.agents.registry import State, BaseAgent
from src.agents.utils import load_chat_model, get_cur_time_with_utc from src.agents.utils import load_chat_model, get_cur_time_with_utc
@ -26,6 +26,8 @@ class ChatbotAgent(BaseAgent):
def __init__(self, **kwargs): def __init__(self, **kwargs):
super().__init__(**kwargs) super().__init__(**kwargs)
self.graph = None self.graph = None
self.workdir = Path(sys_config.save_dir) / "agents" / self.name
self.workdir.mkdir(parents=True, exist_ok=True)
def _get_tools(self, tools: list[str]): def _get_tools(self, tools: list[str]):
"""根据配置获取工具。 """根据配置获取工具。
@ -56,7 +58,7 @@ class ChatbotAgent(BaseAgent):
) )
return {"messages": [res]} return {"messages": [res]}
def get_graph(self, config_schema: RunnableConfig = None, **kwargs): async def get_graph(self, config_schema: RunnableConfig = None, **kwargs):
"""构建图""" """构建图"""
if self.graph: if self.graph:
return self.graph return self.graph
@ -73,18 +75,19 @@ class ChatbotAgent(BaseAgent):
workflow.add_edge("tools", "chatbot") workflow.add_edge("tools", "chatbot")
workflow.add_edge("chatbot", END) workflow.add_edge("chatbot", END)
mem_checkpointer = InMemorySaver() # 创建数据库连接
graph = workflow.compile(checkpointer=mem_checkpointer) sqlite_checkpointer = AsyncSqliteSaver(await self.get_async_conn())
graph = workflow.compile(checkpointer=sqlite_checkpointer)
self.graph = graph self.graph = graph
return graph return graph
# async def get_async_conn(self) -> aiosqlite.Connection: async def get_async_conn(self) -> aiosqlite.Connection:
# """获取异步数据库连接""" """获取异步数据库连接"""
# return await aiosqlite.connect(os.path.join(self.db_dir, "aio_history.db")) return await aiosqlite.connect(os.path.join(self.workdir, "aio_history.db"))
# async def get_aio_memory(self) -> AsyncSqliteSaver: async def get_aio_memory(self) -> AsyncSqliteSaver:
# """获取异步存储实例""" """获取异步存储实例"""
# return AsyncSqliteSaver(await self.get_async_conn()) return AsyncSqliteSaver(await self.get_async_conn())
def main(): def main():
agent = ChatbotAgent(ChatbotConfiguration()) agent = ChatbotAgent(ChatbotConfiguration())

View File

@ -30,7 +30,7 @@ def main():
config = {"configurable": {"thread_id": thread_id}} config = {"configurable": {"thread_id": thread_id}}
from src.agents.utils import agent_cli from src.agents.utils import agent_cli
agent_cli(agent, config) asyncio.run(agent_cli(agent, config))
if __name__ == "__main__": if __name__ == "__main__":

View File

@ -171,28 +171,27 @@ class BaseAgent():
if requirement not in os.environ: if requirement not in os.environ:
raise ValueError(f"没有配置{requirement} 环境变量,请在 src/.env 文件中配置,并重新启动服务") raise ValueError(f"没有配置{requirement} 环境变量,请在 src/.env 文件中配置,并重新启动服务")
def stream_values(self, messages: list[str], config_schema: RunnableConfig = None, **kwargs): async def stream_values(self, messages: list[str], config_schema: RunnableConfig = None, **kwargs):
graph = self.get_graph(config_schema=config_schema, **kwargs) graph = await self.get_graph(config_schema=config_schema, **kwargs)
logger.debug(f"stream_values: {config_schema}") logger.debug(f"stream_values: {config_schema}")
for event in graph.stream({"messages": messages}, stream_mode="values", config=config_schema): for event in graph.astream({"messages": messages}, stream_mode="values", config=config_schema):
yield event["messages"] yield event["messages"]
def stream_messages(self, messages: list[str], config_schema: RunnableConfig = None, **kwargs): async def stream_messages(self, messages: list[str], config_schema: RunnableConfig = None, **kwargs):
graph = self.get_graph(config_schema=config_schema, **kwargs) graph = await self.get_graph(config_schema=config_schema, **kwargs)
logger.debug(f"stream_messages: {config_schema}") logger.debug(f"stream_messages: {config_schema}")
for msg, metadata in graph.stream({"messages": messages}, stream_mode="messages", config=config_schema): async for msg, metadata in graph.astream({"messages": messages}, stream_mode="messages", config=config_schema):
yield msg, metadata yield msg, metadata
def get_history(self, user_id, thread_id) -> list[dict]: async def get_history(self, user_id, thread_id) -> list[dict]:
"""获取历史消息""" """获取历史消息"""
# 获取LangGraph应用实例 # 获取LangGraph应用实例
app = self.get_graph() app = await self.get_graph()
# 构建配置信息 # 构建配置信息
config = {"configurable": {"thread_id": thread_id, "user_id": user_id}} config = {"configurable": {"thread_id": thread_id, "user_id": user_id}}
# 获取状态 # 获取状态
state = app.get_state(config) state = await app.aget_state(config)
# logger.debug(f"获取历史消息: {state}")
result = [] result = []
if state: if state:

View File

@ -1,4 +1,5 @@
from datetime import datetime, timezone from datetime import datetime, timezone
import asyncio
from src.models import select_model from src.models import select_model
from src.agents.registry import BaseAgent from src.agents.registry import BaseAgent
@ -22,7 +23,7 @@ def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel:
return model_instance.chat_open_ai return model_instance.chat_open_ai
def agent_cli(agent: BaseAgent, config: RunnableConfig = None): async def agent_cli(agent: BaseAgent, config: RunnableConfig = None):
config = config or {} config = config or {}
if "configurable" not in config: if "configurable" not in config:
config["configurable"] = {} config["configurable"] = {}
@ -34,7 +35,7 @@ def agent_cli(agent: BaseAgent, config: RunnableConfig = None):
break break
stream_flag = False stream_flag = False
for msg, metadata in agent.stream_messages([{"role": "user", "content": user_input}], config): async for msg, metadata in agent.stream_messages([{"role": "user", "content": user_input}], config):
if isinstance(msg, AIMessageChunk): if isinstance(msg, AIMessageChunk):
content = msg.content or msg.tool_calls content = msg.content or msg.tool_calls

27
uv.lock
View File

@ -138,6 +138,18 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/ec/6a/bc7e17a3e87a2985d3e8f4da4cd0f481060eb78fb08596c42be62c90a4d9/aiosignal-1.3.2-py2.py3-none-any.whl", hash = "sha256:45cde58e409a301715980c2b01d0c28bdde3770d8290b5eb2173759d9acb31a5", size = 7597, upload-time = "2024-12-13T17:10:38.469Z" }, { url = "https://files.pythonhosted.org/packages/ec/6a/bc7e17a3e87a2985d3e8f4da4cd0f481060eb78fb08596c42be62c90a4d9/aiosignal-1.3.2-py2.py3-none-any.whl", hash = "sha256:45cde58e409a301715980c2b01d0c28bdde3770d8290b5eb2173759d9acb31a5", size = 7597, upload-time = "2024-12-13T17:10:38.469Z" },
] ]
[[package]]
name = "aiosqlite"
version = "0.21.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/13/7d/8bca2bf9a247c2c5dfeec1d7a5f40db6518f88d314b8bca9da29670d2671/aiosqlite-0.21.0.tar.gz", hash = "sha256:131bb8056daa3bc875608c631c678cda73922a2d4ba8aec373b19f18c17e7aa3", size = 13454, upload-time = "2025-02-03T07:30:16.235Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/f5/10/6c25ed6de94c49f88a91fa5018cb4c0f3625f31d5be9f771ebe5cc7cd506/aiosqlite-0.21.0-py3-none-any.whl", hash = "sha256:2549cf4057f95f53dcba16f2b64e8e2791d7e1adedb13197dd8ed77bb226d7d0", size = 15792, upload-time = "2025-02-03T07:30:13.6Z" },
]
[[package]] [[package]]
name = "albucore" name = "albucore"
version = "0.0.23" version = "0.0.23"
@ -1350,6 +1362,19 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/12/52/bceb5b5348c7a60ef0625ab0a0a0a9ff5d78f0e12aed8cc55c49d5e8a8c9/langgraph_checkpoint-2.0.25-py3-none-any.whl", hash = "sha256:23416a0f5bc9dd712ac10918fc13e8c9c4530c419d2985a441df71a38fc81602", size = 42312, upload-time = "2025-04-26T21:00:42.242Z" }, { url = "https://files.pythonhosted.org/packages/12/52/bceb5b5348c7a60ef0625ab0a0a0a9ff5d78f0e12aed8cc55c49d5e8a8c9/langgraph_checkpoint-2.0.25-py3-none-any.whl", hash = "sha256:23416a0f5bc9dd712ac10918fc13e8c9c4530c419d2985a441df71a38fc81602", size = 42312, upload-time = "2025-04-26T21:00:42.242Z" },
] ]
[[package]]
name = "langgraph-checkpoint-sqlite"
version = "2.0.7"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "aiosqlite" },
{ name = "langgraph-checkpoint" },
]
sdist = { url = "https://files.pythonhosted.org/packages/9c/89/125b80e41ddeb8476654a8cda4b76ee999e4fb3d5913ff9c2b45fb9bfff7/langgraph_checkpoint_sqlite-2.0.7.tar.gz", hash = "sha256:344f307c0840a1cbd85a18dcd6daac8e989947979c1a43c2bdc6c6f4ed12084a", size = 9584, upload-time = "2025-05-02T05:44:24.853Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/25/72/86354caec3bd546ea596422dcaf8495a052287400a9961c2b45264536d1e/langgraph_checkpoint_sqlite-2.0.7-py3-none-any.whl", hash = "sha256:b04decd8c3f7c2966ca63b4fa11eb789a03b27001e4d855ccd132c50da59812b", size = 12958, upload-time = "2025-05-02T05:44:23.665Z" },
]
[[package]] [[package]]
name = "langgraph-prebuilt" name = "langgraph-prebuilt"
version = "0.1.8" version = "0.1.8"
@ -4498,6 +4523,7 @@ dependencies = [
{ name = "langchain-huggingface" }, { name = "langchain-huggingface" },
{ name = "langchain-openai" }, { name = "langchain-openai" },
{ name = "langgraph" }, { name = "langgraph" },
{ name = "langgraph-checkpoint-sqlite" },
{ name = "langsmith" }, { name = "langsmith" },
{ name = "llama-index" }, { name = "llama-index" },
{ name = "llama-index-readers-file" }, { name = "llama-index-readers-file" },
@ -4531,6 +4557,7 @@ requires-dist = [
{ name = "langchain-huggingface", specifier = ">=0.2.0" }, { name = "langchain-huggingface", specifier = ">=0.2.0" },
{ name = "langchain-openai", specifier = ">=0.3.14" }, { name = "langchain-openai", specifier = ">=0.3.14" },
{ name = "langgraph", specifier = ">=0.3.34" }, { name = "langgraph", specifier = ">=0.3.34" },
{ name = "langgraph-checkpoint-sqlite", specifier = ">=2.0.7" },
{ name = "langsmith", specifier = ">=0.3.37" }, { name = "langsmith", specifier = ">=0.3.37" },
{ name = "llama-index", specifier = ">=0.12.33" }, { name = "llama-index", specifier = ">=0.12.33" },
{ name = "llama-index-readers-file", specifier = ">=0.4.7" }, { name = "llama-index-readers-file", specifier = ">=0.4.7" },