2025-08-31 00:34:26 +08:00
|
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
|
|
|
|
|
from abc import abstractmethod
|
|
|
|
|
|
from langgraph.graph.state import CompiledStateGraph
|
|
|
|
|
|
|
|
|
|
|
|
from src.utils import logger
|
|
|
|
|
|
from src.agents.common.context import BaseContext
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class BaseAgent:
|
|
|
|
|
|
|
|
|
|
|
|
"""
|
|
|
|
|
|
定义一个基础 Agent 供 各类 graph 继承
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
name = "base_agent"
|
|
|
|
|
|
description = "base_agent"
|
|
|
|
|
|
|
|
|
|
|
|
def __init__(self, **kwargs):
|
|
|
|
|
|
self.graph = None # will be covered by get_graph
|
|
|
|
|
|
self.context_schema = BaseContext
|
|
|
|
|
|
|
|
|
|
|
|
@property
|
|
|
|
|
|
def module_name(self) -> str:
|
|
|
|
|
|
"""Get the module name of the agent class."""
|
|
|
|
|
|
return self.__class__.__module__.split('.')[-2]
|
|
|
|
|
|
|
|
|
|
|
|
@property
|
|
|
|
|
|
def id(self) -> str:
|
|
|
|
|
|
"""Get the agent's class name."""
|
|
|
|
|
|
return self.__class__.__name__
|
|
|
|
|
|
|
|
|
|
|
|
async def get_info(self):
|
|
|
|
|
|
return {
|
|
|
|
|
|
"id": self.id,
|
|
|
|
|
|
"name": self.name if hasattr(self, "name") else "Unknown",
|
|
|
|
|
|
"description": self.description if hasattr(self, "description") else "Unknown",
|
|
|
|
|
|
"configurable_items": self.context_schema.get_configurable_items(),
|
|
|
|
|
|
"has_checkpointer": await self.check_checkpointer(),
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
async def get_config(self):
|
|
|
|
|
|
return self.context_schema.from_file(module_name=self.module_name)
|
|
|
|
|
|
|
|
|
|
|
|
async def stream_values(self, messages: list[str], input_context = None, **kwargs):
|
|
|
|
|
|
graph = await self.get_graph()
|
|
|
|
|
|
context = self.context_schema.from_file(module_name=self.module_name, input_context=input_context)
|
|
|
|
|
|
for event in graph.astream({"messages": messages}, stream_mode="values", context=context):
|
|
|
|
|
|
yield event["messages"]
|
|
|
|
|
|
|
|
|
|
|
|
async def stream_messages(self, messages: list[str], input_context = None, **kwargs):
|
|
|
|
|
|
graph = await self.get_graph()
|
|
|
|
|
|
context = self.context_schema.from_file(module_name=self.module_name, input_context=input_context)
|
2025-09-01 03:38:46 +08:00
|
|
|
|
logger.debug(f"stream_messages: {context}")
|
2025-08-31 00:34:26 +08:00
|
|
|
|
# TODO 的 Checkpointer 似乎还没有适配最新的 Context API
|
|
|
|
|
|
async for msg, metadata in graph.astream({"messages": messages}, stream_mode="messages", context=context, config={"configurable": input_context}):
|
|
|
|
|
|
yield msg, metadata
|
|
|
|
|
|
|
|
|
|
|
|
async def check_checkpointer(self):
|
|
|
|
|
|
app = await self.get_graph()
|
|
|
|
|
|
if not hasattr(app, "checkpointer") or app.checkpointer is None:
|
|
|
|
|
|
logger.warning(f"智能体 {self.name} 的 Graph 未配置 checkpointer,无法获取历史记录")
|
|
|
|
|
|
return False
|
|
|
|
|
|
return True
|
|
|
|
|
|
|
|
|
|
|
|
async def get_history(self, user_id, thread_id) -> list[dict]:
|
|
|
|
|
|
"""获取历史消息"""
|
|
|
|
|
|
try:
|
|
|
|
|
|
app = await self.get_graph()
|
|
|
|
|
|
|
|
|
|
|
|
if not await self.check_checkpointer():
|
|
|
|
|
|
return []
|
|
|
|
|
|
|
|
|
|
|
|
config = {"configurable": {"thread_id": thread_id, "user_id": user_id}}
|
|
|
|
|
|
state = await app.aget_state(config)
|
|
|
|
|
|
|
|
|
|
|
|
result = []
|
|
|
|
|
|
if state:
|
|
|
|
|
|
messages = state.values.get('messages', [])
|
|
|
|
|
|
for msg in messages:
|
|
|
|
|
|
if hasattr(msg, 'model_dump'):
|
|
|
|
|
|
msg_dict = msg.model_dump() # 转换成字典
|
|
|
|
|
|
else:
|
|
|
|
|
|
msg_dict = dict(msg) if hasattr(msg, '__dict__') else {"content": str(msg)}
|
|
|
|
|
|
result.append(msg_dict)
|
|
|
|
|
|
|
|
|
|
|
|
return result
|
|
|
|
|
|
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.error(f"获取智能体 {self.name} 历史消息出错: {e}")
|
|
|
|
|
|
return []
|
|
|
|
|
|
|
|
|
|
|
|
@abstractmethod
|
|
|
|
|
|
async def get_graph(self, **kwargs) -> CompiledStateGraph:
|
|
|
|
|
|
"""
|
|
|
|
|
|
获取并编译对话图实例。
|
|
|
|
|
|
必须确保在编译时设置 checkpointer,否则将无法获取历史记录。
|
|
|
|
|
|
例如: graph = workflow.compile(checkpointer=sqlite_checkpointer)
|
|
|
|
|
|
"""
|
|
|
|
|
|
pass
|