mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-08 03:10:24 +00:00
refactor(memory): restructure memory tools and add base memory agent
This commit is contained in:
parent
d29aff4e6c
commit
29a6ee1fba
19 changed files with 238 additions and 44 deletions
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
187
reme_ai/mem_agent/base_memory_agent.py
Normal file
187
reme_ai/mem_agent/base_memory_agent.py
Normal 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
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue