diff --git a/reme_ai/core/op/__init__.py b/reme_ai/core/op/__init__.py index c009e89f..a48bf276 100644 --- a/reme_ai/core/op/__init__.py +++ b/reme_ai/core/op/__init__.py @@ -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", ] diff --git a/reme_ai/core/op/base_op.py b/reme_ai/core/op/base_op.py index 4c82b64c..5cb5a0a2 100644 --- a/reme_ai/core/op/base_op.py +++ b/reme_ai/core/op/base_op.py @@ -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) diff --git a/reme_ai/tool/mcp_tool.py b/reme_ai/core/op/mcp_tool.py similarity index 95% rename from reme_ai/tool/mcp_tool.py rename to reme_ai/core/op/mcp_tool.py index f06275b0..57cbcbd4 100644 --- a/reme_ai/tool/mcp_tool.py +++ b/reme_ai/core/op/mcp_tool.py @@ -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() diff --git a/reme_ai/mem_agent/base_memory_agent.py b/reme_ai/mem_agent/base_memory_agent.py new file mode 100644 index 00000000..b7913ed1 --- /dev/null +++ b/reme_ai/mem_agent/base_memory_agent.py @@ -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 diff --git a/reme_ai/tool/__init__.py b/reme_ai/tool/__init__.py index d7d25f85..f1dcedd3 100644 --- a/reme_ai/tool/__init__.py +++ b/reme_ai/tool/__init__.py @@ -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", ] diff --git a/reme_ai/tool/memory/__init__.py b/reme_ai/tool/memory/__init__.py index b6dfb163..15392685 100644 --- a/reme_ai/tool/memory/__init__.py +++ b/reme_ai/tool/memory/__init__.py @@ -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", diff --git a/reme_ai/tool/base_memory_tool.py b/reme_ai/tool/memory/base_memory_tool.py similarity index 95% rename from reme_ai/tool/base_memory_tool.py rename to reme_ai/tool/memory/base_memory_tool.py index 7ec1a133..f87286a0 100644 --- a/reme_ai/tool/base_memory_tool.py +++ b/reme_ai/tool/memory/base_memory_tool.py @@ -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): diff --git a/reme_ai/tool/memory/history/add_history_memory.py b/reme_ai/tool/memory/history/add_history_memory.py index 85d0928d..125ebe86 100644 --- a/reme_ai/tool/memory/history/add_history_memory.py +++ b/reme_ai/tool/memory/history/add_history_memory.py @@ -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 diff --git a/reme_ai/tool/memory/history/read_history_memory.py b/reme_ai/tool/memory/history/read_history_memory.py index a4c7a91b..25c71f8d 100644 --- a/reme_ai/tool/memory/history/read_history_memory.py +++ b/reme_ai/tool/memory/history/read_history_memory.py @@ -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 diff --git a/reme_ai/tool/memory/identity/read_identity_memory.py b/reme_ai/tool/memory/identity/read_identity_memory.py index 9247ab69..0e1e5df4 100644 --- a/reme_ai/tool/memory/identity/read_identity_memory.py +++ b/reme_ai/tool/memory/identity/read_identity_memory.py @@ -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 diff --git a/reme_ai/tool/memory/identity/update_identity_memory.py b/reme_ai/tool/memory/identity/update_identity_memory.py index ba0e2de7..dec7444b 100644 --- a/reme_ai/tool/memory/identity/update_identity_memory.py +++ b/reme_ai/tool/memory/identity/update_identity_memory.py @@ -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 diff --git a/reme_ai/tool/memory/meta/add_meta_memory.py b/reme_ai/tool/memory/meta/add_meta_memory.py index cc71f640..2de6cc74 100644 --- a/reme_ai/tool/memory/meta/add_meta_memory.py +++ b/reme_ai/tool/memory/meta/add_meta_memory.py @@ -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 diff --git a/reme_ai/tool/memory/meta/read_meta_memory.py b/reme_ai/tool/memory/meta/read_meta_memory.py index 2a69d799..0aee30c8 100644 --- a/reme_ai/tool/memory/meta/read_meta_memory.py +++ b/reme_ai/tool/memory/meta/read_meta_memory.py @@ -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 diff --git a/reme_ai/tool/think_tool.py b/reme_ai/tool/memory/think_tool.py similarity index 90% rename from reme_ai/tool/think_tool.py rename to reme_ai/tool/memory/think_tool.py index 2e83b3ad..3172010e 100644 --- a/reme_ai/tool/think_tool.py +++ b/reme_ai/tool/memory/think_tool.py @@ -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; diff --git a/reme_ai/tool/think_tool.yaml b/reme_ai/tool/memory/think_tool.yaml similarity index 100% rename from reme_ai/tool/think_tool.yaml rename to reme_ai/tool/memory/think_tool.yaml diff --git a/reme_ai/tool/memory/vector/add_memory.py b/reme_ai/tool/memory/vector/add_memory.py index 9a27b370..15fc86c4 100644 --- a/reme_ai/tool/memory/vector/add_memory.py +++ b/reme_ai/tool/memory/vector/add_memory.py @@ -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 diff --git a/reme_ai/tool/memory/vector/delete_memory.py b/reme_ai/tool/memory/vector/delete_memory.py index 8c747b57..aa529283 100644 --- a/reme_ai/tool/memory/vector/delete_memory.py +++ b/reme_ai/tool/memory/vector/delete_memory.py @@ -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 diff --git a/reme_ai/tool/memory/vector/update_memory.py b/reme_ai/tool/memory/vector/update_memory.py index d427891d..ad7203bd 100644 --- a/reme_ai/tool/memory/vector/update_memory.py +++ b/reme_ai/tool/memory/vector/update_memory.py @@ -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 diff --git a/reme_ai/tool/memory/vector/vector_retrieve_memory.py b/reme_ai/tool/memory/vector/vector_retrieve_memory.py index 9e0cc6c7..f6c7c271 100644 --- a/reme_ai/tool/memory/vector/vector_retrieve_memory.py +++ b/reme_ai/tool/memory/vector/vector_retrieve_memory.py @@ -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