refactor(memory): move memory tools to dedicated module and update imports

This commit is contained in:
jinli.yl 2026-01-07 10:10:03 +08:00
parent 29a6ee1fba
commit 53aaad2a28
37 changed files with 67 additions and 62 deletions

View file

@ -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

View file

@ -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]

View file

@ -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)

View file

@ -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(),
),

View file

@ -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}")

View file

@ -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:

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

@ -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()

View file

@ -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()

View file

@ -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()

View file

@ -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()

View file

@ -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()

View file

@ -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()

View file

@ -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()

View file

@ -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()

View file

@ -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()

View file

@ -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()

View file

@ -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()

View file

@ -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()

View file

@ -1,11 +1,9 @@
"""tool"""
from . import execute
from . import memory
from . import search
__all__ = [
"execute",
"memory",
"search",
]