2025-05-15 22:34:32 +08:00
|
|
|
|
import os
|
2025-03-24 23:00:14 +08:00
|
|
|
|
import uuid
|
|
|
|
|
|
from typing import Any
|
2025-05-16 23:46:25 +08:00
|
|
|
|
from pathlib import Path
|
2025-04-02 13:00:25 +08:00
|
|
|
|
from datetime import datetime, timezone
|
2025-03-24 19:07:51 +08:00
|
|
|
|
|
2025-05-16 23:46:25 +08:00
|
|
|
|
import sqlite3
|
2025-03-24 19:07:51 +08:00
|
|
|
|
from langchain_core.runnables import RunnableConfig
|
2025-03-24 23:00:14 +08:00
|
|
|
|
from langgraph.graph import StateGraph, START, END
|
|
|
|
|
|
from langgraph.prebuilt import ToolNode, tools_condition
|
2025-05-16 23:46:25 +08:00
|
|
|
|
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver, aiosqlite
|
2025-03-24 19:07:51 +08:00
|
|
|
|
|
2025-05-16 23:46:25 +08:00
|
|
|
|
from src import config as sys_config
|
2025-04-02 13:00:25 +08:00
|
|
|
|
from src.utils import logger
|
2025-03-24 19:07:51 +08:00
|
|
|
|
from src.agents.registry import State, BaseAgent
|
2025-04-02 13:00:25 +08:00
|
|
|
|
from src.agents.utils import load_chat_model, get_cur_time_with_utc
|
2025-03-31 22:20:29 +08:00
|
|
|
|
from src.agents.chatbot.configuration import ChatbotConfiguration
|
2025-04-05 17:27:52 +08:00
|
|
|
|
from src.agents.tools_factory import get_all_tools
|
2025-03-24 19:07:51 +08:00
|
|
|
|
|
|
|
|
|
|
class ChatbotAgent(BaseAgent):
|
2025-07-22 17:29:38 +08:00
|
|
|
|
name = "对话机器人(Chatbot)"
|
2025-04-05 17:27:52 +08:00
|
|
|
|
description = "基础的对话机器人,可以回答问题,默认不使用任何工具,可在配置中启用需要的工具。"
|
2025-03-31 22:20:29 +08:00
|
|
|
|
requirements = ["TAVILY_API_KEY", "ZHIPUAI_API_KEY"]
|
|
|
|
|
|
config_schema = ChatbotConfiguration
|
2025-03-24 23:00:14 +08:00
|
|
|
|
|
2025-03-28 11:40:46 +08:00
|
|
|
|
def __init__(self, **kwargs):
|
|
|
|
|
|
super().__init__(**kwargs)
|
2025-05-15 22:34:32 +08:00
|
|
|
|
self.graph = None
|
2025-07-22 17:30:00 +08:00
|
|
|
|
self.workdir = Path(sys_config.save_dir) / "agents" / self.id
|
2025-05-16 23:46:25 +08:00
|
|
|
|
self.workdir.mkdir(parents=True, exist_ok=True)
|
2025-03-25 05:40:07 +08:00
|
|
|
|
|
2025-04-05 17:27:52 +08:00
|
|
|
|
def _get_tools(self, tools: list[str]):
|
|
|
|
|
|
"""根据配置获取工具。
|
|
|
|
|
|
默认不使用任何工具。
|
|
|
|
|
|
如果配置为列表,则使用列表中的工具。
|
|
|
|
|
|
"""
|
|
|
|
|
|
platform_tools = get_all_tools()
|
|
|
|
|
|
if tools is None or not isinstance(tools, list) or len(tools) == 0:
|
|
|
|
|
|
# 默认不使用任何工具
|
|
|
|
|
|
logger.info("未配置工具或配置为空,不使用任何工具")
|
|
|
|
|
|
return []
|
2025-04-02 13:00:25 +08:00
|
|
|
|
else:
|
2025-04-05 17:27:52 +08:00
|
|
|
|
# 使用配置中指定的工具
|
|
|
|
|
|
tool_names = [tool for tool in platform_tools.keys() if tool in tools]
|
|
|
|
|
|
logger.info(f"使用工具: {tool_names}")
|
|
|
|
|
|
return [platform_tools[tool] for tool in tool_names]
|
2025-03-25 05:40:07 +08:00
|
|
|
|
|
2025-06-27 01:49:31 +08:00
|
|
|
|
async def llm_call(self, state: State, config: RunnableConfig = None) -> dict[str, Any]:
|
|
|
|
|
|
"""调用 llm 模型 - 异步版本以支持异步工具"""
|
2025-07-22 17:30:00 +08:00
|
|
|
|
conf = self.config_schema.from_runnable_config(config, module_name=self.module_name)
|
2025-04-02 13:00:25 +08:00
|
|
|
|
|
|
|
|
|
|
system_prompt = f"{conf.system_prompt} Now is {get_cur_time_with_utc()}"
|
|
|
|
|
|
model = load_chat_model(conf.model)
|
2025-03-25 05:40:07 +08:00
|
|
|
|
|
2025-06-24 00:13:11 +08:00
|
|
|
|
if tools := self._get_tools(conf.tools):
|
|
|
|
|
|
model = model.bind_tools(tools)
|
|
|
|
|
|
|
2025-06-27 01:49:31 +08:00
|
|
|
|
# 使用异步调用
|
|
|
|
|
|
res = await model.ainvoke(
|
2025-04-02 13:00:25 +08:00
|
|
|
|
[{"role": "system", "content": system_prompt}, *state["messages"]]
|
2025-03-29 17:33:09 +08:00
|
|
|
|
)
|
2025-03-24 23:00:14 +08:00
|
|
|
|
return {"messages": [res]}
|
|
|
|
|
|
|
2025-05-16 23:46:25 +08:00
|
|
|
|
async def get_graph(self, config_schema: RunnableConfig = None, **kwargs):
|
2025-03-24 23:00:14 +08:00
|
|
|
|
"""构建图"""
|
2025-05-15 22:34:32 +08:00
|
|
|
|
if self.graph:
|
|
|
|
|
|
return self.graph
|
|
|
|
|
|
|
2025-03-31 22:20:29 +08:00
|
|
|
|
workflow = StateGraph(State, config_schema=self.config_schema)
|
2025-03-25 05:40:07 +08:00
|
|
|
|
workflow.add_node("chatbot", self.llm_call)
|
2025-05-17 10:06:56 +08:00
|
|
|
|
workflow.add_node("tools", ToolNode(tools=list(get_all_tools().values())))
|
2025-03-25 05:40:07 +08:00
|
|
|
|
workflow.add_edge(START, "chatbot")
|
|
|
|
|
|
workflow.add_conditional_edges(
|
|
|
|
|
|
"chatbot",
|
|
|
|
|
|
tools_condition,
|
|
|
|
|
|
)
|
|
|
|
|
|
workflow.add_edge("tools", "chatbot")
|
|
|
|
|
|
workflow.add_edge("chatbot", END)
|
|
|
|
|
|
|
2025-05-20 22:21:51 +08:00
|
|
|
|
# 创建数据库连接并确保设置 checkpointer
|
|
|
|
|
|
try:
|
|
|
|
|
|
sqlite_checkpointer = AsyncSqliteSaver(await self.get_async_conn())
|
|
|
|
|
|
graph = workflow.compile(checkpointer=sqlite_checkpointer)
|
|
|
|
|
|
self.graph = graph
|
|
|
|
|
|
return graph
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.error(f"构建 Graph 设置 checkpointer 时出错: {e}")
|
|
|
|
|
|
# 即使出错也返回一个可用的图实例,只是无法保存历史
|
|
|
|
|
|
graph = workflow.compile()
|
|
|
|
|
|
self.graph = graph
|
|
|
|
|
|
return graph
|
2025-03-24 23:00:14 +08:00
|
|
|
|
|
2025-05-16 23:46:25 +08:00
|
|
|
|
async def get_async_conn(self) -> aiosqlite.Connection:
|
|
|
|
|
|
"""获取异步数据库连接"""
|
|
|
|
|
|
return await aiosqlite.connect(os.path.join(self.workdir, "aio_history.db"))
|
2025-05-15 22:34:32 +08:00
|
|
|
|
|
2025-05-16 23:46:25 +08:00
|
|
|
|
async def get_aio_memory(self) -> AsyncSqliteSaver:
|
|
|
|
|
|
"""获取异步存储实例"""
|
|
|
|
|
|
return AsyncSqliteSaver(await self.get_async_conn())
|
2025-03-24 23:00:14 +08:00
|
|
|
|
|
|
|
|
|
|
def main():
|
|
|
|
|
|
agent = ChatbotAgent(ChatbotConfiguration())
|
2025-03-24 19:07:51 +08:00
|
|
|
|
|
2025-03-24 23:00:14 +08:00
|
|
|
|
thread_id = str(uuid.uuid4())
|
|
|
|
|
|
config = {"configurable": {"thread_id": thread_id}}
|
2025-03-24 19:07:51 +08:00
|
|
|
|
|
2025-03-25 05:40:07 +08:00
|
|
|
|
from src.agents.utils import agent_cli
|
2025-03-24 23:00:14 +08:00
|
|
|
|
agent_cli(agent, config)
|
2025-03-24 19:07:51 +08:00
|
|
|
|
|
|
|
|
|
|
|
2025-03-24 23:00:14 +08:00
|
|
|
|
if __name__ == "__main__":
|
|
|
|
|
|
main()
|
|
|
|
|
|
# asyncio.run(main())
|