mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
refactor(memory): migrate and optimize user profile memory tools
This commit is contained in:
parent
3f0d45c51e
commit
30990308d8
6 changed files with 175 additions and 187 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
60
reme/tool/memory/read_user_profile.py
Normal file
60
reme/tool/memory/read_user_profile.py
Normal file
|
|
@ -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()
|
||||
109
reme/tool/memory/update_user_profile.py
Normal file
109
reme/tool/memory/update_user_profile.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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}")
|
||||
|
|
@ -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)
|
||||
Loading…
Add table
Reference in a new issue