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) logger.debug(f"stream_messages: {context}") # 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