refactor(memory): restructure memory tools and add base memory agent

This commit is contained in:
jinli.yl 2026-01-07 00:56:22 +08:00
parent d29aff4e6c
commit 29a6ee1fba
19 changed files with 238 additions and 44 deletions

View file

@ -2,12 +2,14 @@
from .base_op import BaseOp
from .base_ray_op import BaseRayOp
from .mcp_tool import MCPTool
from .parallel_op import ParallelOp
from .sequential_op import SequentialOp
__all__ = [
"BaseOp",
"BaseRayOp",
"MCPTool",
"ParallelOp",
"SequentialOp",
]

View file

@ -126,21 +126,6 @@ class BaseOp:
}
return self._tool_call
def set_tool_call(self, tool_call: ToolCall | dict):
"""Set the tool call."""
if isinstance(tool_call, dict):
self._tool_call = ToolCall(**tool_call)
elif isinstance(tool_call, ToolCall):
self._tool_call = tool_call
else:
raise ValueError(f"Invalid tool call: {tool_call}")
self._tool_call.name = self._tool_call.name or self.name
if not self._tool_call.output.properties:
self._tool_call.output.properties = {
f"{self.name}_result": ToolAttr(type="string", description=f"Execution result of {self.name}"),
}
@property
def input_dict(self) -> dict:
"""Extract required and optional inputs from context based on schema."""
@ -219,6 +204,26 @@ class BaseOp:
"""Get the response object."""
return self.context.response
def set_tool_call(self, tool_call: ToolCall | dict):
"""Set the tool call."""
if isinstance(tool_call, dict):
self._tool_call = ToolCall(**tool_call)
elif isinstance(tool_call, ToolCall):
self._tool_call = tool_call
else:
raise ValueError(f"Invalid tool call: {tool_call}")
self._tool_call.name = self._tool_call.name or self.name
if not self._tool_call.output.properties:
self._tool_call.output.properties = {
f"{self.name}_result": ToolAttr(type="string", description=f"Execution result of {self.name}"),
}
def set_language(self, language: str):
"""Set the language."""
self.language = language
return self
def before_execute_sync(self):
"""Prepare context and validate before sync execution."""
self.context.apply_mapping(self.input_mapping)
@ -318,10 +323,12 @@ class BaseOp:
assert self.async_mode == op.async_mode, "Async mode mismatch!"
op.name = name
self.sub_ops.append(op)
elif isinstance(sub_ops, list):
for op in sub_ops:
assert self.async_mode == op.async_mode, "Async mode mismatch!"
self.sub_ops.append(op)
else:
assert self.async_mode == sub_ops.async_mode, "Async mode mismatch!"
self.sub_ops.append(sub_ops)

View file

@ -2,10 +2,10 @@
from typing import List
from ..core.context import C
from ..core.op import BaseOp
from ..core.schema import ToolCall
from ..core.utils import MCPClient
from .base_op import BaseOp
from ..context import C
from ..schema import ToolCall
from ..utils import MCPClient
@C.register_op()

View file

@ -0,0 +1,187 @@
"""Base memory agent for handling memory operations with tool-based reasoning."""
import asyncio
from abc import ABCMeta
from loguru import logger
from ..core.enumeration import Role, MemoryType
from ..core.op import BaseOp
from ..core.schema import Message, ToolCall
from ..tool.memory import BaseMemoryTool, ThinkTool
class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
"""Base class for memory agents that perform reasoning and acting with memory tools."""
memory_type: MemoryType | None = None
def __init__(
self,
max_steps: int = 20,
tool_call_interval: float = 0,
add_think_tool: bool = False, # only for instruct model
tools: list[BaseMemoryTool] | None = None,
**kwargs,
):
super().__init__(**kwargs)
self.max_steps: int = max_steps
self.tool_call_interval: float = tool_call_interval
self.add_think_tool: bool = add_think_tool
assert not self.sub_ops, "sub_ops must be empty, use `tools`~"
if tools:
self.sub_ops.extend([t.set_language(self.language) for t in tools])
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": self.get_prompt("tool"),
"parameters": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "query",
},
"messages": {
"type": "array",
"items": {
"type": "object",
"properties": {
"role": {
"type": "string",
"description": "role",
},
"content": {
"type": "string",
"description": "content",
},
},
"required": ["role", "content"],
},
},
},
"required": [],
},
},
)
@property
def tools(self) -> list[BaseMemoryTool]:
"""Returns the list of memory tools available to the agent."""
tools: list[BaseMemoryTool] = [o for o in self.sub_ops if isinstance(o, BaseMemoryTool)]
if self.add_think_tool:
tools.append(ThinkTool(language=self.language))
return tools
@tools.setter
def tools(self, tools: list[BaseMemoryTool] | BaseMemoryTool):
"""Sets the memory tools for the agent."""
self.sub_ops = tools
def get_messages(self) -> list[Message]:
"""Extracts and returns messages from the context query or messages."""
if self.context.get("query"):
messages = [Message(role=Role.USER, content=self.context.query)]
elif self.context.get("messages"):
messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages]
else:
raise ValueError("input must have either `query` or `messages`")
return messages
async def build_messages(self) -> list[Message]:
"""Builds and returns the initial messages for the agent."""
return self.get_messages()
async def _reasoning_step(self, messages: list[Message], step: int, **kwargs) -> tuple[Message, bool]:
assistant_message: Message = await self.llm.chat(
messages=messages,
tools=[t.tool_call for t in self.tools],
**kwargs,
)
messages.append(assistant_message)
logger.info(f"step{step + 1}.assistant={assistant_message.model_dump_json()}")
should_act = bool(assistant_message.tool_calls)
return assistant_message, should_act
async def _acting_step(self, assistant_message: Message, step: int, **kwargs) -> list[Message]:
if not assistant_message.tool_calls:
return []
tool_list: list[BaseMemoryTool] = []
tool_result_messages: list[Message] = []
tool_dict = {t.tool_call.name: t for t in self.tools}
for j, tool_call in enumerate(assistant_message.tool_calls):
if tool_call.name not in tool_dict:
logger.warning(f"unknown tool_call.name={tool_call.name}")
continue
logger.info(f"step{step + 1}.{j} submit tool_calls={tool_call.name} argument={tool_call.argument_dict}")
tool_copy: BaseMemoryTool = tool_dict[tool_call.name].copy()
tool_copy.tool_call.id = tool_call.id
tool_list.append(tool_copy)
kwargs.update(tool_call.argument_dict)
self.submit_async_task(tool_copy.call, **kwargs)
if self.tool_call_interval > 0:
await asyncio.sleep(self.tool_call_interval)
await self.join_async_tasks()
for j, op in enumerate(tool_list):
tool_result = str(op.output)
tool_message = Message(
role=Role.TOOL,
content=tool_result,
tool_call_id=op.tool_call.id,
)
tool_result_messages.append(tool_message)
logger.info(f"step{step + 1}.{j} join tool_result={tool_result[:200]}...\n\n")
return tool_result_messages
async def react(self, messages: list[Message]):
"""Performs reasoning and acting steps until completion or max steps reached."""
success: bool = False
for step in range(self.max_steps):
assistant_message, should_act = await self._reasoning_step(messages, step)
if not should_act:
success = True
break
tool_result_messages = await self._acting_step(assistant_message, step)
messages.extend(tool_result_messages)
return messages, success
async def execute(self):
messages = await self.build_messages()
for i, message in enumerate(messages):
logger.info(f"step0.{i} {message.role} {message.name or ''} {message.simple_dump()}")
messages, success = await self.react(messages)
if messages:
if success:
self.output = messages[-1].content
else:
self.output = f"react is not complete with content:\n{messages[-1].content}"
else:
self.output = "empty messages"
self.context.response.metadata["messages"] = messages
self.context.response.metadata["success"] = success
@property
def memory_target(self) -> str:
"""Returns the target memory identifier from context."""
return self.context.get("memory_target", "")
@property
def ref_memory_id(self) -> str:
"""Returns the reference memory ID from context."""
return self.context.get("ref_memory_id", "")
@property
def author(self) -> str:
"""Returns the LLM model name as the author identifier."""
return self.llm.model_name

View file

@ -3,15 +3,9 @@
from . import execute
from . import memory
from . import search
from .base_memory_tool import BaseMemoryTool
from .mcp_tool import MCPTool
from .think_tool import ThinkTool
__all__ = [
"execute",
"memory",
"search",
"BaseMemoryTool",
"MCPTool",
"ThinkTool",
]

View file

@ -1,11 +1,13 @@
"""Memory tool operations."""
from .base_memory_tool import BaseMemoryTool
from .history.add_history_memory import AddHistoryMemory
from .history.read_history_memory import ReadHistoryMemory
from .identity.read_identity_memory import ReadIdentityMemory
from .identity.update_identity_memory import UpdateIdentityMemory
from .meta.add_meta_memory import AddMetaMemory
from .meta.read_meta_memory import ReadMetaMemory
from .think_tool import ThinkTool
from .vector.add_memory import AddMemory
from .vector.add_summary_memory import AddSummaryMemory
from .vector.delete_memory import DeleteMemory
@ -13,12 +15,14 @@ from .vector.update_memory import UpdateMemory
from .vector.vector_retrieve_memory import VectorRetrieveMemory
__all__ = [
"BaseMemoryTool",
"AddHistoryMemory",
"ReadHistoryMemory",
"ReadIdentityMemory",
"UpdateIdentityMemory",
"AddMetaMemory",
"ReadMetaMemory",
"ThinkTool",
"AddMemory",
"AddSummaryMemory",
"DeleteMemory",

View file

@ -3,10 +3,10 @@
from abc import ABCMeta
from pathlib import Path
from ..core.enumeration import MemoryType
from ..core.op import BaseOp
from ..core.schema import ToolCall, MemoryNode
from ..core.utils import CacheHandler
from ...core.enumeration import MemoryType
from ...core.op import BaseOp
from ...core.schema import ToolCall, MemoryNode
from ...core.utils import CacheHandler
class BaseMemoryTool(BaseOp, metaclass=ABCMeta):

View file

@ -2,7 +2,7 @@
from loguru import logger
from ...base_memory_tool import BaseMemoryTool
from ..base_memory_tool import BaseMemoryTool
from ....core.context import C
from ....core.enumeration import MemoryType
from ....core.schema import ToolCall, Message

View file

@ -2,7 +2,7 @@
from loguru import logger
from ...base_memory_tool import BaseMemoryTool
from ..base_memory_tool import BaseMemoryTool
from ....core.context import C
from ....core.schema import MemoryNode

View file

@ -2,7 +2,7 @@
from loguru import logger
from ...base_memory_tool import BaseMemoryTool
from ..base_memory_tool import BaseMemoryTool
from ....core.context import C

View file

@ -2,7 +2,7 @@
from loguru import logger
from ...base_memory_tool import BaseMemoryTool
from ..base_memory_tool import BaseMemoryTool
from ....core.context import C

View file

@ -4,7 +4,7 @@ import json
from loguru import logger
from ...base_memory_tool import BaseMemoryTool
from ..base_memory_tool import BaseMemoryTool
from ....core.context import C
from ....core.enumeration import MemoryType

View file

@ -2,7 +2,7 @@
from loguru import logger
from ...base_memory_tool import BaseMemoryTool
from ..base_memory_tool import BaseMemoryTool
from ....core.context import C
from ....core.enumeration import MemoryType

View file

@ -4,13 +4,13 @@ This module provides a tool that prompts the model for explicit reflection
before taking actions, helping agents reason about their next steps.
"""
from ..core.context import C
from ..core.op import BaseOp
from ..core.schema import ToolCall
from .base_memory_tool import BaseMemoryTool
from ...core.context import C
from ...core.schema import ToolCall
@C.register_op()
class ThinkTool(BaseOp):
class ThinkTool(BaseMemoryTool):
"""Utility that prompts the model for explicit reflection text.
This tool provides a thinking mechanism for agents to reflect on:
@ -20,7 +20,7 @@ class ThinkTool(BaseOp):
"""
def __init__(self, add_output_reflection: bool = False, **kwargs):
"""Initialize the think tool tool.
"""Initialize the think tool.
Args:
add_output_reflection: If True, outputs the reflection content;

View file

@ -2,7 +2,7 @@
from loguru import logger
from ...base_memory_tool import BaseMemoryTool
from ..base_memory_tool import BaseMemoryTool
from ....core.context import C
from ....core.schema import MemoryNode

View file

@ -2,7 +2,7 @@
from loguru import logger
from ...base_memory_tool import BaseMemoryTool
from ..base_memory_tool import BaseMemoryTool
from ....core.context import C

View file

@ -2,7 +2,7 @@
from loguru import logger
from ...base_memory_tool import BaseMemoryTool
from ..base_memory_tool import BaseMemoryTool
from ....core.context import C
from ....core.schema import MemoryNode

View file

@ -2,7 +2,7 @@
from loguru import logger
from ...base_memory_tool import BaseMemoryTool
from ..base_memory_tool import BaseMemoryTool
from ....core.context import C
from ....core.enumeration import MemoryType
from ....core.schema import MemoryNode, VectorNode