From 30990308d866fc7628cfa6d3a590d30d05112d7b Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 22 Jan 2026 23:46:11 +0800 Subject: [PATCH] refactor(memory): migrate and optimize user profile memory tools --- reme/reme_app.py | 2 +- reme/tool/memory/base_memory_tool.py | 5 + reme/tool/memory/read_user_profile.py | 60 ++++++++++++ reme/tool/memory/update_user_profile.py | 109 +++++++++++++++++++++ reme_ai/mem_tool/v4/read_user_profile.py | 81 --------------- reme_ai/mem_tool/v4/update_user_profile.py | 105 -------------------- 6 files changed, 175 insertions(+), 187 deletions(-) create mode 100644 reme/tool/memory/read_user_profile.py create mode 100644 reme/tool/memory/update_user_profile.py delete mode 100644 reme_ai/mem_tool/v4/read_user_profile.py delete mode 100644 reme_ai/mem_tool/v4/update_user_profile.py diff --git a/reme/reme_app.py b/reme/reme_app.py index 483cd5bb..e41aa8a6 100644 --- a/reme/reme_app.py +++ b/reme/reme_app.py @@ -3,11 +3,11 @@ import asyncio import sys -from reme.core.utils import execute_stream_task from .config import ReMeConfigParser from .core.context import ServiceContext from .core.flow import BaseFlow from .core.schema import Response +from .core.utils import execute_stream_task class ReMeApp: diff --git a/reme/tool/memory/base_memory_tool.py b/reme/tool/memory/base_memory_tool.py index aaef5542..adc2f6b2 100644 --- a/reme/tool/memory/base_memory_tool.py +++ b/reme/tool/memory/base_memory_tool.py @@ -75,6 +75,11 @@ class BaseMemoryTool(BaseTool, metaclass=ABCMeta): """Get the memory target from context.""" return self.context.get("memory_target", "") + @property + def memory_cache_key(self) -> str: + """Get the memory cache key from context.""" + return f"{self.memory_type.value}_{self.memory_target}".replace(" ", "_").lower() + @property def history_node(self) -> MemoryNode: """Get the history node from context.""" diff --git a/reme/tool/memory/read_user_profile.py b/reme/tool/memory/read_user_profile.py new file mode 100644 index 00000000..6b3118cf --- /dev/null +++ b/reme/tool/memory/read_user_profile.py @@ -0,0 +1,60 @@ +"""Read user profile tool""" + +from typing import Literal + +from loguru import logger + +from .base_memory_tool import BaseMemoryTool +from ...core.schema import ToolCall +from ...core.schema.memory_node import MemoryNode + + +class ReadUserProfile(BaseMemoryTool): + """Tool to read user profile from local memory""" + + def __init__(self, show_id: Literal["profile", "history"] = "profile", **kwargs): + kwargs["enable_multiple"] = False + super().__init__(**kwargs) + self.show_id = show_id + + def _build_tool_call(self) -> ToolCall: + return ToolCall( + **{ + "description": "read user profile.", + "parameters": { + "type": "object", + "properties": {}, + "required": [], + }, + }, + ) + + async def execute(self): + cached_data = self.local_memory.load(self.memory_cache_key, auto_clean=False) + + if not cached_data: + logger.info(f"No cached data found for {self.memory_cache_key}") + return "" + + nodes = [MemoryNode(**data) for data in cached_data] + nodes.sort(key=lambda n: n.metadata.get("conversation_time", "")) + + formatted_profiles = [] + for node in nodes: + parts = [] + if self.show_id == "profile": + parts.append(f"profile_id={node.memory_id}") + + if conv_time := node.metadata.get("conversation_time"): + parts.append(f"conversation_time={conv_time}") + + parts.append(f"{node.when_to_use}: {node.content}") + + if self.show_id == "history": + parts.append(f"history_id={node.ref_memory_id}") + + formatted_profiles.append(" ".join(parts)) + + logger.info(f"Read {len(formatted_profiles)} profiles from cache key: {self.memory_cache_key}") + + return "### User Profile\n" + "\n".join(formatted_profiles).strip() diff --git a/reme/tool/memory/update_user_profile.py b/reme/tool/memory/update_user_profile.py new file mode 100644 index 00000000..e2032d45 --- /dev/null +++ b/reme/tool/memory/update_user_profile.py @@ -0,0 +1,109 @@ +"""Update user profile tool""" + +from loguru import logger + +from .base_memory_tool import BaseMemoryTool +from ...core.schema import ToolCall +from ...core.schema.memory_node import MemoryNode +from ...core.utils import deduplicate_memories + + +class UpdateUserProfile(BaseMemoryTool): + """Tool to update user profile by adding or removing profile entries""" + + def __init__(self, **kwargs): + kwargs["enable_multiple"] = True + super().__init__(**kwargs) + + def _build_multiple_tool_call(self) -> ToolCall: + """Build and return the multiple tool call schema""" + return ToolCall( + **{ + "description": "update user profile by adding or removing profile entries.", + "parameters": { + "type": "object", + "properties": { + "profile_ids_to_delete": { + "type": "array", + "description": "List of profile IDs to delete", + "items": {"type": "string"}, + }, + "profiles_to_add": { + "type": "array", + "description": "List of profiles to add", + "items": { + "type": "object", + "properties": { + "conversation_time": { + "type": "string", + "description": "Conversation time, e.g. '2020-01-01 00:00:00'", + }, + "profile_key": { + "type": "string", + "description": "Profile key or category, e.g. 'name'", + }, + "profile_value": { + "type": "string", + "description": "Profile value or content, e.g. 'John Smith'", + }, + }, + "required": ["conversation_time", "profile_key", "profile_value"], + }, + }, + }, + "required": ["profile_ids_to_delete", "profiles_to_add"], + }, + }, + ) + + async def execute(self): + # Get and deduplicate profile IDs to delete + profile_ids_to_delete = self.context.get("profile_ids_to_delete", []) + profile_ids_to_delete = list(dict.fromkeys([pid for pid in profile_ids_to_delete if pid])) + profiles_to_add = self.context.get("profiles_to_add", []) + + if not profile_ids_to_delete and not profiles_to_add: + return "No profiles to remove or add. Operation completed." + + # Load existing profiles from local memory + cached_data = self.local_memory.load(self.memory_cache_key, auto_clean=False) + existing_nodes = [MemoryNode(**data) for data in cached_data] if cached_data else [] + + # Remove profiles + removed_count = 0 + if profile_ids_to_delete: + original_count = len(existing_nodes) + existing_nodes = [n for n in existing_nodes if n.memory_id not in profile_ids_to_delete] + removed_count = original_count - len(existing_nodes) + logger.info(f"Removed {removed_count} profiles.") + + # Add new profiles + new_nodes = [] + if profiles_to_add: + for profile in profiles_to_add: + node = MemoryNode( + memory_type=self.memory_type, + memory_target=self.memory_target, + when_to_use=profile.get("profile_key", ""), + content=profile.get("profile_value", ""), + ref_memory_id=self.history_node.memory_id, + author=self.author, + metadata={"conversation_time": profile.get("conversation_time", "")}, + ) + new_nodes.append(node) + logger.info(f"Added {len(new_nodes)} new profiles.") + + # Deduplicate and save updated profiles + updated_nodes = deduplicate_memories(existing_nodes + new_nodes) + nodes_data = [node.model_dump(exclude_none=True) for node in updated_nodes] + self.local_memory.save(self.memory_cache_key, nodes_data) + + # Build output message + operations = [] + if removed_count > 0: + operations.append(f"removed {removed_count} old profiles.") + if len(new_nodes) > 0: + operations.append(f"added {len(new_nodes)} new profiles.") + operations.append("Operation completed.") + logger.info("\n".join(operations)) + return operations diff --git a/reme_ai/mem_tool/v4/read_user_profile.py b/reme_ai/mem_tool/v4/read_user_profile.py deleted file mode 100644 index ff963ea2..00000000 --- a/reme_ai/mem_tool/v4/read_user_profile.py +++ /dev/null @@ -1,81 +0,0 @@ -from typing import Literal -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ...core.schema.memory_node import MemoryNode - - -class ReadUserProfile(BaseMemoryTool): - - def __init__(self, add_memory_type_target: bool = False, show_ids: Literal["both", "profile", "history", "none"] = "both", **kwargs): - kwargs["enable_multiple"] = False - super().__init__(**kwargs) - self.add_memory_type_target = add_memory_type_target - self.show_ids = show_ids - - def _build_tool_description(self) -> str: - return "Read user profile." - - def _build_parameters(self) -> dict: - if self.add_memory_type_target: - return { - "type": "object", - "properties": { - "memory_type": { - "type": "string", - "description": "memory_type", - }, - "memory_target": { - "type": "string", - "description": "memory_target", - }, - }, - "required": ["memory_type", "memory_target"], - } - else: - return { - "type": "object", - "properties": {}, - "required": [], - } - - async def execute(self): - # Determine which IDs to show - show_profile_id = self.show_ids in ("both", "profile") - show_history_id = self.show_ids in ("both", "history") - - cache_key = f"{self.memory_type}_{self.memory_target}".replace(" ", "_").lower() - cached_data = self.meta_memory.load(cache_key, auto_clean=False) - - if not cached_data: - self.output = "### User Profile\nNo user profile found." - logger.info(f"empty cached_data={cache_key}") - return - - memory_nodes = [MemoryNode(**node_data) for node_data in cached_data] - memory_nodes.sort(key=lambda n: n.metadata.get("conversation_time", "")) - - memory_formated = [] - for node in memory_nodes: - node_formated_parts = [] - - # Add profile_id if enabled - if show_profile_id: - node_formated_parts.append(f"profile_id={node.memory_id}") - - # Always add profile_content - node_formated_parts.append(f"profile_content={node.content}") - - # Add conversation_time if available - if "conversation_time" in node.metadata and node.metadata["conversation_time"]: - node_formated_parts.append(f"conversation_time={node.metadata['conversation_time']}") - - # Add history_id if enabled and available - if show_history_id and node.ref_memory_id: - node_formated_parts.append(f"history_id={node.ref_memory_id}") - - node_formated = " ".join(node_formated_parts) - memory_formated.append(node_formated.strip()) - - self.output = "### User Profile\n" + "\n".join(memory_formated) - logger.info(f"Read {len(memory_formated)} nodes from cache key: {cache_key}") diff --git a/reme_ai/mem_tool/v4/update_user_profile.py b/reme_ai/mem_tool/v4/update_user_profile.py deleted file mode 100644 index a8fa04f5..00000000 --- a/reme_ai/mem_tool/v4/update_user_profile.py +++ /dev/null @@ -1,105 +0,0 @@ -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ...core.schema.memory_node import MemoryNode -from ...core.utils import deduplicate_memories - - -class UpdateUserProfile(BaseMemoryTool): - - def __init__(self, **kwargs): - kwargs["enable_multiple"] = True - super().__init__(**kwargs) - - def _build_tool_description(self) -> str: - return "Update user profile." - - def _build_multiple_parameters(self) -> dict: - return { - "type": "object", - "properties": { - "profile_ids_to_delete": { - "type": "array", - "description": "profile_ids_to_delete", - "items": { - "type": "string" - }, - }, - "profiles_to_add": { - "type": "array", - "description": "profiles_to_add", - "items": { - "type": "object", - "properties": { - "conversation_time": { - "type": "string", - "description": "conversation_time, e.g. '2020-01-01 00:00:00'", - }, - "profile_content": { - "type": "string", - "description": "profile_content", - }, - }, - "required": ["conversation_time", "profile_content"], - }, - }, - }, - "required": ["profile_ids_to_delete", "profiles_to_add"], - } - - async def execute(self): - profile_ids_to_delete = self.context.get("profile_ids_to_delete", []) - profile_ids_to_delete = [m for m in profile_ids_to_delete if m] - profile_ids_to_delete = list(dict.fromkeys(profile_ids_to_delete)) - profiles_to_add = self.context.get("profiles_to_add", []) - - if not profile_ids_to_delete and not profiles_to_add: - self.output = "No profiles to remove or add. Operation has been done." - return - - cache_key = f"{self.memory_type}_{self.memory_target}".replace(" ", "_").lower() - cached_data = self.meta_memory.load(cache_key, auto_clean=False) - if cached_data: - existing_memory_nodes = [MemoryNode(**node_data) for node_data in cached_data] - else: - existing_memory_nodes = [] - - removed_count = 0 - if profile_ids_to_delete: - original_count = len(existing_memory_nodes) - existing_memory_nodes = [n for n in existing_memory_nodes if n.memory_id not in profile_ids_to_delete] - removed_count = original_count - len(existing_memory_nodes) - logger.info(f"Removed {removed_count} profiles.") - - added_count = 0 - new_memory_nodes = [] - if profiles_to_add: - for mem in profiles_to_add: - memory_node = MemoryNode( - memory_type=self.memory_type, - memory_target=self.memory_target, - when_to_use="", - content=mem.get("profile_content", ""), - ref_memory_id=self.history_node.memory_id, - author=self.author, - metadata={"conversation_time": mem.get("conversation_time", "")}, - ) - new_memory_nodes.append(memory_node) - added_count = len(new_memory_nodes) - logger.info(f"Added {added_count} new profiles.") - - updated_memory_nodes = deduplicate_memories(existing_memory_nodes + new_memory_nodes) - nodes_data = [node.model_dump(exclude_none=True) for node in updated_memory_nodes] - self.meta_memory.save(cache_key, nodes_data) - - operations = [] - if removed_count > 0: - operations.append(f"removed {removed_count} old profiles") - if added_count > 0: - operations.append(f"added {added_count} new profiles") - - if operations: - self.output = f"Successfully {' and '.join(operations)} in user profile." - else: - self.output = "Operation has been done." - logger.info(self.output)