mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
refactor(memory): move memory tools to dedicated module and update imports
This commit is contained in:
parent
29a6ee1fba
commit
53aaad2a28
37 changed files with 67 additions and 62 deletions
|
|
@ -1,4 +1,4 @@
|
|||
"""Common constants module.
|
||||
"""Common constants' module.
|
||||
|
||||
This module defines constants used as keys throughout the application to maintain
|
||||
a consistent reference for data structures related to workflow management, chat
|
||||
|
|
|
|||
|
|
@ -57,7 +57,7 @@ class BaseRayOp(BaseOp, metaclass=ABCMeta):
|
|||
for k, v in kwargs.items()
|
||||
}
|
||||
|
||||
# Submit sliced chunks to reduce inter-node data transfer
|
||||
# Submit sliced chunks to reduce internode data transfer
|
||||
remote_task_loop = ray.remote(self._ray_task_loop)
|
||||
for i in range(max_workers):
|
||||
chunk = parallel_list[i::max_workers]
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@
|
|||
|
||||
import datetime
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
|
@ -44,14 +43,14 @@ class ContentBlock(BaseModel):
|
|||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def init_block(cls, data: dict[str, Any]) -> dict[str, Any]:
|
||||
def init_block(cls, data: dict) -> dict:
|
||||
"""Dynamically maps the type-specific key to the content field."""
|
||||
content_type = data.get("type", "")
|
||||
if content_type and content_type in data:
|
||||
data["content"] = data[content_type]
|
||||
return data
|
||||
|
||||
def simple_dump(self) -> dict[str, Any]:
|
||||
def simple_dump(self) -> dict:
|
||||
"""Serializes the block into an API-compatible dictionary format."""
|
||||
return {
|
||||
"type": self.type,
|
||||
|
|
@ -69,28 +68,45 @@ class Message(BaseModel):
|
|||
reasoning_content: str = Field(default="")
|
||||
tool_calls: list[ToolCall] = Field(default_factory=list)
|
||||
tool_call_id: str = Field(default="")
|
||||
time_created: str = Field(
|
||||
default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
|
||||
)
|
||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
time_created: str = Field(default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S"))
|
||||
metadata: dict = Field(default_factory=dict)
|
||||
|
||||
def dump_content(self) -> str | list[dict[str, Any]]:
|
||||
def dump_content(self) -> str | list[dict]:
|
||||
"""Returns content as a raw string or a list of serialized blocks."""
|
||||
if isinstance(self.content, str):
|
||||
return self.content
|
||||
return [block.simple_dump() for block in self.content]
|
||||
|
||||
def simple_dump(self, add_reasoning: bool = True) -> dict[str, Any]:
|
||||
def simple_dump(
|
||||
self,
|
||||
add_name: bool = False,
|
||||
add_reasoning: bool = True,
|
||||
add_time_created: bool = False,
|
||||
add_metadata: bool = False,
|
||||
) -> dict:
|
||||
"""Transforms the message into a simplified dictionary for standard APIs."""
|
||||
result = {"role": self.role.value, "content": self.dump_content()}
|
||||
result = {}
|
||||
if add_name and self.name:
|
||||
result["name"] = self.name
|
||||
|
||||
result["role"] = self.role.value
|
||||
result["content"] = self.dump_content()
|
||||
|
||||
if add_reasoning and self.reasoning_content:
|
||||
result["reasoning_content"] = self.reasoning_content
|
||||
|
||||
if self.tool_calls:
|
||||
result["tool_calls"] = [tc.simple_output_dump() for tc in self.tool_calls]
|
||||
|
||||
if self.tool_call_id:
|
||||
result["tool_call_id"] = self.tool_call_id
|
||||
|
||||
if add_time_created:
|
||||
result["time_created"] = self.time_created
|
||||
|
||||
if add_metadata:
|
||||
result["metadata"] = self.metadata
|
||||
|
||||
return result
|
||||
|
||||
def format_message(
|
||||
|
|
@ -133,4 +149,4 @@ class Trajectory(BaseModel):
|
|||
task_id: str = Field(default="")
|
||||
messages: list[Message] = Field(default_factory=list)
|
||||
score: float = Field(default=0.0)
|
||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
metadata: dict = Field(default_factory=dict)
|
||||
|
|
|
|||
|
|
@ -34,8 +34,8 @@ class MCPService(BaseService):
|
|||
|
||||
self.mcp.add_tool(
|
||||
FunctionTool(
|
||||
name=tool_call.name,
|
||||
description=tool_call.description,
|
||||
name=tool_call.name, # noqa
|
||||
description=tool_call.description, # noqa
|
||||
fn=execute_tool,
|
||||
parameters=tool_call.parameters.simple_input_dump(),
|
||||
),
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from typing import Any
|
|||
from mcp import ClientSession, StdioServerParameters, Tool
|
||||
from mcp.client.sse import sse_client
|
||||
from mcp.client.stdio import stdio_client
|
||||
from mcp.client.streamable_http import streamablehttp_client
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
|
||||
from ..schema import ToolCall
|
||||
|
|
@ -64,7 +64,7 @@ class MCPClient:
|
|||
async with sse_client(**cfg) as transport:
|
||||
yield transport
|
||||
elif t_type == "streamable-http":
|
||||
async with streamablehttp_client(**cfg) as transport:
|
||||
async with streamable_http_client(**cfg) as transport:
|
||||
yield transport
|
||||
else:
|
||||
raise NotImplementedError(f"Unsupported transport: {t_type}")
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ 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
|
||||
from ..mem_tool import BaseMemoryTool, ThinkTool
|
||||
|
||||
|
||||
class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
|
||||
|
|
@ -141,35 +141,26 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
|
|||
|
||||
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
|
||||
return messages
|
||||
|
||||
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
|
||||
messages = await self.react(messages)
|
||||
self.output = [
|
||||
m.simple_dump(add_name=True, add_reasoning=True, add_time_created=True, add_metadata=True) for m in messages
|
||||
]
|
||||
|
||||
@property
|
||||
def memory_target(self) -> str:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
@ -3,10 +3,10 @@
|
|||
from loguru import logger
|
||||
|
||||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ....core.context import C
|
||||
from ....core.enumeration import MemoryType
|
||||
from ....core.schema import ToolCall, Message
|
||||
from ....core.utils import format_messages
|
||||
from ...core.context import C
|
||||
from ...core.enumeration import MemoryType
|
||||
from ...core.schema import ToolCall, Message
|
||||
from ...core.utils import format_messages
|
||||
|
||||
|
||||
@C.register_op()
|
||||
|
|
@ -3,8 +3,8 @@
|
|||
from loguru import logger
|
||||
|
||||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ....core.context import C
|
||||
from ....core.schema import MemoryNode
|
||||
from ...core.context import C
|
||||
from ...core.schema import MemoryNode
|
||||
|
||||
|
||||
@C.register_op()
|
||||
|
|
@ -3,7 +3,7 @@
|
|||
from loguru import logger
|
||||
|
||||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ....core.context import C
|
||||
from ...core.context import C
|
||||
|
||||
|
||||
@C.register_op()
|
||||
|
|
@ -3,7 +3,7 @@
|
|||
from loguru import logger
|
||||
|
||||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ....core.context import C
|
||||
from ...core.context import C
|
||||
|
||||
|
||||
@C.register_op()
|
||||
|
|
@ -5,8 +5,8 @@ import json
|
|||
from loguru import logger
|
||||
|
||||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ....core.context import C
|
||||
from ....core.enumeration import MemoryType
|
||||
from ...core.context import C
|
||||
from ...core.enumeration import MemoryType
|
||||
|
||||
|
||||
@C.register_op()
|
||||
|
|
@ -3,8 +3,8 @@
|
|||
from loguru import logger
|
||||
|
||||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ....core.context import C
|
||||
from ....core.enumeration import MemoryType
|
||||
from ...core.context import C
|
||||
from ...core.enumeration import MemoryType
|
||||
|
||||
|
||||
@C.register_op()
|
||||
|
|
@ -5,8 +5,8 @@ before taking actions, helping agents reason about their next steps.
|
|||
"""
|
||||
|
||||
from .base_memory_tool import BaseMemoryTool
|
||||
from ...core.context import C
|
||||
from ...core.schema import ToolCall
|
||||
from ..core.context import C
|
||||
from ..core.schema import ToolCall
|
||||
|
||||
|
||||
@C.register_op()
|
||||
|
|
@ -3,8 +3,8 @@
|
|||
from loguru import logger
|
||||
|
||||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ....core.context import C
|
||||
from ....core.schema import MemoryNode
|
||||
from ...core.context import C
|
||||
from ...core.schema import MemoryNode
|
||||
|
||||
|
||||
@C.register_op()
|
||||
|
|
@ -3,7 +3,7 @@
|
|||
from loguru import logger
|
||||
|
||||
from .add_memory import AddMemory
|
||||
from ....core.context import C
|
||||
from ...core.context import C
|
||||
|
||||
|
||||
@C.register_op()
|
||||
|
|
@ -3,7 +3,7 @@
|
|||
from loguru import logger
|
||||
|
||||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ....core.context import C
|
||||
from ...core.context import C
|
||||
|
||||
|
||||
@C.register_op()
|
||||
|
|
@ -3,8 +3,8 @@
|
|||
from loguru import logger
|
||||
|
||||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ....core.context import C
|
||||
from ....core.schema import MemoryNode
|
||||
from ...core.context import C
|
||||
from ...core.schema import MemoryNode
|
||||
|
||||
|
||||
@C.register_op()
|
||||
|
|
@ -3,10 +3,10 @@
|
|||
from loguru import logger
|
||||
|
||||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ....core.context import C
|
||||
from ....core.enumeration import MemoryType
|
||||
from ....core.schema import MemoryNode, VectorNode
|
||||
from ....core.utils import deduplicate_memories
|
||||
from ...core.context import C
|
||||
from ...core.enumeration import MemoryType
|
||||
from ...core.schema import MemoryNode, VectorNode
|
||||
from ...core.utils import deduplicate_memories
|
||||
|
||||
|
||||
@C.register_op()
|
||||
|
|
@ -1,11 +1,9 @@
|
|||
"""tool"""
|
||||
|
||||
from . import execute
|
||||
from . import memory
|
||||
from . import search
|
||||
|
||||
__all__ = [
|
||||
"execute",
|
||||
"memory",
|
||||
"search",
|
||||
]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue