From 5a5e06a43dbb2a07b39547bd6b022b83e8b08c94 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sat, 31 Jan 2026 02:36:10 +0800 Subject: [PATCH] feat(memory): add draft memory functionality and update memory/profile management tools --- .../memory/personal/personal_v1_retriever.py | 8 +- .../personal/personal_v1_retriever.yaml | 34 ++- .../memory/personal/personal_v1_summarizer.py | 20 +- .../personal/personal_v1_summarizer.yaml | 82 ++++--- reme/reme.py | 19 +- reme/tool/memory/__init__.py | 4 + reme/tool/memory/history/read_history.py | 56 +++-- .../memory/profiles/update_profiles_v1.py | 176 +++++++++++++++ .../add_draft_and_retrieve_similar_memory.py | 127 +++++++++++ reme/tool/memory/vector/update_memory_v1.py | 212 ++++++++++++++++++ 10 files changed, 638 insertions(+), 100 deletions(-) create mode 100644 reme/tool/memory/profiles/update_profiles_v1.py create mode 100644 reme/tool/memory/vector/add_draft_and_retrieve_similar_memory.py create mode 100644 reme/tool/memory/vector/update_memory_v1.py diff --git a/reme/agent/memory/personal/personal_v1_retriever.py b/reme/agent/memory/personal/personal_v1_retriever.py index 37b2fa41..749aa564 100644 --- a/reme/agent/memory/personal/personal_v1_retriever.py +++ b/reme/agent/memory/personal/personal_v1_retriever.py @@ -35,19 +35,15 @@ class PersonalV1Retriever(BaseMemoryAgent): return [ Message( - role=Role.SYSTEM, + role=Role.USER, content=self.prompt_format( - prompt_name="system_prompt", + prompt_name="user_message", memory_type=self.memory_type.value, memory_target=self.memory_target, user_profile=all_profiles, context=context.strip(), ), ), - Message( - role=Role.USER, - content=self.get_prompt("user_message"), - ), ] async def _acting_step( diff --git a/reme/agent/memory/personal/personal_v1_retriever.yaml b/reme/agent/memory/personal/personal_v1_retriever.yaml index 32b85e45..54ece8d0 100644 --- a/reme/agent/memory/personal/personal_v1_retriever.yaml +++ b/reme/agent/memory/personal/personal_v1_retriever.yaml @@ -1,4 +1,4 @@ -system_prompt: | +user_message: | You are a Memory Retrieval Agent specialized in retrieving {memory_type} memories about {memory_target}. ## User Profile @@ -22,9 +22,9 @@ system_prompt: | * Related context queries (broader themes) - Review all results before proceeding to next phase - ### Phase 2: Temporal Search (Optional) + ### Phase 2(Optional): Temporal Search **Tool**: `retrieve_memory` (with time filter) - **When to use**: Only if the user question contains temporal references + **When to use**: Only if the user's question includes a time-related reference; otherwise, go straight to Phase 3. **Time Filter Format**: - Single date: `20200101` - Date range: `20200101,20200102` (inclusive: 20200101 ≤ time ≤ 20200102) @@ -38,29 +38,25 @@ system_prompt: | ### Phase 3: Deep Dive into History **Tool**: `read_history` **When to use**: After exhausting retrieval attempts OR when specific conversation context is needed + **Important Constraints**: + - Each history is very long and resource-intensive to read + - **Maximum limit: Read no more than 3 histories total** + - Only use this phase when absolutely necessary for answering the question **Approach**: - Extract `history_id` from retrieved memory references - - Prioritize histories that are most relevant or recent - - Read multiple histories if necessary for complete context + - Prioritize the most relevant or recent histories + - Can read multiple histories at once by passing multiple history_ids + - Be selective: choose only the top 1-3 most promising histories - Use this to understand the full conversation surrounding a memory ## Response Guidelines - **Critical Rules**: - - Base your answer EXCLUSIVELY on retrieved memories, user profile, and history data + - Base your answer EXCLUSIVELY on user profile, retrieved memories, and history data - Never infer, assume, or hallucinate information - Always cite sources with timestamps: `[timestamp] Memory content` - Present conflicting information transparently with respective timestamps + - If you find sufficient information to answer the user's question, you may output directly without exhausting all search phases - Exhaust all search strategies before concluding information doesn't exist - **Output Format**: - - When information is found: - - [timestamp] All relevant memory/profile/history content - - - When no information is found after thorough search (5+ queries across phases): - - No relevant information found after exhaustive search using multiple query strategies and retrieval phases. - - -user_message: | - Retrieve relevant memories following the multi-phase strategy outlined above. \ No newline at end of file + ### Output any tangentially related findings, Format: + [timestamp] [memory/profile/history] [relevant content1] + [timestamp] [memory/profile/history] [relevant content2] diff --git a/reme/agent/memory/personal/personal_v1_summarizer.py b/reme/agent/memory/personal/personal_v1_summarizer.py index 864cbd3b..a87b183a 100644 --- a/reme/agent/memory/personal/personal_v1_summarizer.py +++ b/reme/agent/memory/personal/personal_v1_summarizer.py @@ -16,35 +16,27 @@ class PersonalV1Summarizer(BaseMemoryAgent): async def _build_s1_messages(self) -> list[Message]: return [ Message( - role=Role.SYSTEM, + role=Role.USER, content=self.prompt_format( - prompt_name="system_prompt_s1", + prompt_name="user_message_s1", context=self.context.history_node.content, memory_type=self.memory_type.value, memory_target=self.memory_target, ), ), - Message( - role=Role.USER, - content=self.get_prompt("user_message_s1"), - ), ] async def _build_s2_messages(self) -> list[Message]: return [ Message( - role=Role.SYSTEM, + role=Role.USER, content=self.prompt_format( - prompt_name="system_prompt_s2", + prompt_name="user_message_s2", context=self.context.history_node.content, memory_type=self.memory_type.value, memory_target=self.memory_target, ), ), - Message( - role=Role.USER, - content=self.get_prompt("user_message_s2"), - ), ] async def _acting_step( @@ -99,7 +91,9 @@ class PersonalV1Summarizer(BaseMemoryAgent): else: tools_s2, messages_s2, success_s2 = [], [], True - answer = (messages_s1[-1].content if success_s1 else "") + (messages_s2[-1].content if success_s2 else "") + answer = (messages_s1[-1].content if success_s1 and messages_s1 else "") + ( + messages_s2[-1].content if success_s2 and messages_s2 else "" + ) success = success_s1 and success_s2 messages = messages_s1 + messages_s2 tools = tools_s1 + tools_s2 diff --git a/reme/agent/memory/personal/personal_v1_summarizer.yaml b/reme/agent/memory/personal/personal_v1_summarizer.yaml index f30dabf6..c988b15f 100644 --- a/reme/agent/memory/personal/personal_v1_summarizer.yaml +++ b/reme/agent/memory/personal/personal_v1_summarizer.yaml @@ -1,4 +1,4 @@ -system_prompt_s1: | +user_message_s1: | You are a Memory Agent responsible for managing {memory_type} memories about {memory_target}. ## Latest Conversation @@ -7,59 +7,57 @@ system_prompt_s1: | ## Task ### Step 1: Create Memory Draft - Create a memory draft in `add_draft_and_retrieve_similar_memory` based on the latest conversation. + Use `add_draft_and_retrieve_similar_memory` to create a memory draft list based on the latest conversation. + - For each memory draft, fill in the required parameters: + * `message_time`: timestamp from the conversation (e.g., '2020-01-01 00:00:00') + * `memory_content`: concise memory content extracted from the conversation - Use actual names from the conversation (e.g., "Bob likes apples") instead of generic references (e.g., "user likes apples") - - Always record memories with real names - - The tool will retrieve similar historical memories via vector search to help you consolidate in Step 2 - - ### Step 2: Update Memory Store - Update the vector store using `update_memory` to keep it well-organized and consolidated: - - **What to Delete** (via `memory_ids_to_delete`): - - Duplicate memories with identical or highly similar content - - Memories that should be merged into a single consolidated entry - - **What to Add** (via `memories_to_add` with message_time and memory_content): - - For each topic with changes: add ONE consolidated memory that merges related information - - New distinct memories that don't overlap with existing ones - - Updated memories that capture the latest state while preserving temporal evolution - - ## Requirements - Extract only what's explicitly stated—no inferences, assumptions, or fabrications - - Preserve temporal evolution: capture how things change over time within the same topic - - Maintain organization: group related memories by topic and eliminate all redundancy + - The tool will retrieve similar historical memories via vector search to help you in Step 2 -user_message_s1: | - Complete the task by following Step 1 and Step 2 in order + ### Step 2: Add New Memories + Review each memory draft from Step 1 and compare it with the retrieved historical memories: + - **Skip** the draft if its content is already fully covered by historical memories (avoid redundancy) + - **Add** the draft using `add_memory` if it contains new or additional information not present in historical memories + - Required parameters for each memory to add: + * `message_time`: timestamp from the conversation + * `memory_content`: consolidated memory content -system_prompt_s2: | +user_message_s2: | You are a Profile Agent responsible for managing profiles about {memory_target}. ## Latest Conversation Format: round [] : {context} + ## Current Profiles + {profiles} + ## Task - ### Step 1: Create Profile Draft - Create a profile draft in `add_draft_and_read_all_profiles` based on the latest conversation. - - The tool will return all existing profiles to help you maintain the profile store in Step 2 + Analyze the Latest Conversation and use `update_profiles` to manage profiles (both updates and additions in one call): - ### Step 2: Update Profile Store - Update the profile store using `update_profile` to keep it well-organized and consolidated: + **For profiles_to_update** (updating existing profiles): + - For each profile to update, fill in the required parameters: + * `profile_id`: ID of the profile to update (from Current Profiles) + * `message_time`: timestamp from the conversation (e.g., '2020-01-01 00:00:00') + * `profile_key`: profile key or category (e.g., 'name', 'age', 'occupation') + * `profile_value`: updated profile value (e.g., 'John Smith') + - Update profiles when: + * Information in the conversation conflicts with or supersedes existing profiles + * Profiles need to be consolidated or merged with new information + * Existing profile values need to be corrected or refined - **What to Delete** (via `profile_ids_to_delete`): - - Duplicate profiles with identical keys or values - - Conflicting profiles that contradict the new information - - Profiles that should be merged into a single consolidated entry + **For profiles_to_add** (adding new profiles): + - For each new profile, fill in the required parameters: + * `message_time`: timestamp from the conversation (e.g., '2020-01-01 00:00:00') + * `profile_key`: profile key or category (e.g., 'name', 'age', 'occupation') + * `profile_value`: profile value (e.g., 'John Smith') + - Add profiles when: + * The information represents a new distinct profile not present in Current Profiles + * The profile key doesn't exist in Current Profiles + * The information cannot be merged into existing profiles - **What to Add** (via `profiles_to_add` with message_time, profile_key, and profile_value): - - For each profile key with changes: add ONE consolidated profile that merges related information - - New distinct profiles that don't overlap with existing ones - - Updated profiles that capture the latest state - - ## Requirements + **General Guidelines:** + - Use actual names from the conversation (e.g., "Bob") instead of generic references (e.g., "user") - Extract only what's explicitly stated—no inferences, assumptions, or fabrications - - Maintain organization: group related profiles by key and eliminate all redundancy - -user_message_s2: | - Complete the task by following Step 1 and Step 2 in order + - You can update and add profiles in a single tool call diff --git a/reme/reme.py b/reme/reme.py index 0f41ecb9..7511b0bf 100644 --- a/reme/reme.py +++ b/reme/reme.py @@ -29,12 +29,14 @@ from .tool.memory import ( ProfileHandler, MemoryHandler, AddAndRetrieveSimilarMemory, + AddDraftAndRetrieveSimilarMemory, UpdateMemoryV2, AddDraftAndReadAllProfiles, UpdateProfile, AddHistory, ReadAllProfiles, - AddMemory, + UpdateProfilesV1, + UpdateMemoryV1, ) @@ -143,25 +145,24 @@ class ReMe(Application): elif version == "v1": personal_summarizer = PersonalV1Summarizer( tools=[ - AddAndRetrieveSimilarMemory( + AddDraftAndRetrieveSimilarMemory( enable_thinking_params=enable_thinking_params, enable_memory_target=False, enable_when_to_use=False, enable_multiple=True, ), - AddMemory( + UpdateMemoryV1( enable_thinking_params=enable_thinking_params, enable_memory_target=False, enable_when_to_use=False, enable_multiple=True, ), - AddDraftAndReadAllProfiles( + ReadAllProfiles( enable_thinking_params=enable_thinking_params, enable_memory_target=False, - enable_multiple=True, profile_dir=self.profile_dir, ), - UpdateProfile( + UpdateProfilesV1( enable_thinking_params=enable_thinking_params, enable_memory_target=False, enable_multiple=True, @@ -308,6 +309,7 @@ class ReMe(Application): elif version == "v1": personal_retriever = PersonalV1Retriever( + return_memory_nodes=False, tools=[ ReadAllProfiles( enable_thinking_params=enable_thinking_params, @@ -320,7 +322,10 @@ class ReMe(Application): enable_time_filter=enable_time_filter, enable_multiple=True, ), - ReadHistory(enable_thinking_params=enable_thinking_params), + ReadHistory( + enable_thinking_params=enable_thinking_params, + enable_multiple=True, + ), ], ) elif version == "halumem": diff --git a/reme/tool/memory/__init__.py b/reme/tool/memory/__init__.py index 705487c1..3c7bac99 100644 --- a/reme/tool/memory/__init__.py +++ b/reme/tool/memory/__init__.py @@ -10,6 +10,7 @@ from .profiles.delete_profile import DeleteProfile from .profiles.profile_handler import ProfileHandler from .profiles.read_all_profiles import ReadAllProfiles from .profiles.update_profile import UpdateProfile +from .profiles.update_profiles_v1 import UpdateProfilesV1 from .vector.add_and_retrieve_similar_memory import AddAndRetrieveSimilarMemory from .vector.add_draft_and_retrieve_similar_memory import AddDraftAndRetrieveSimilarMemory from .vector.add_memory import AddMemory @@ -18,6 +19,7 @@ from .vector.memory_handler import MemoryHandler from .vector.retrieve_memory import RetrieveMemory from .vector.retrieve_recent_memory import RetrieveRecentMemory from .vector.update_memory import UpdateMemory +from .vector.update_memory_v1 import UpdateMemoryV1 from .vector.update_memory_v2 import UpdateMemoryV2 from ...core import R @@ -35,6 +37,7 @@ __all__ = [ "ReadAllProfiles", "UpdateProfile", "DeleteProfile", + "UpdateProfilesV1", # Vector "AddAndRetrieveSimilarMemory", "AddDraftAndRetrieveSimilarMemory", @@ -44,6 +47,7 @@ __all__ = [ "RetrieveMemory", "RetrieveRecentMemory", "UpdateMemory", + "UpdateMemoryV1", "UpdateMemoryV2", ] diff --git a/reme/tool/memory/history/read_history.py b/reme/tool/memory/history/read_history.py index 1f328eb2..105db721 100644 --- a/reme/tool/memory/history/read_history.py +++ b/reme/tool/memory/history/read_history.py @@ -10,20 +10,20 @@ class ReadHistory(BaseMemoryTool): """Read history memory tool""" def __init__(self, **kwargs): - kwargs["enable_multiple"] = False + kwargs.setdefault("enable_multiple", False) super().__init__(**kwargs) def _build_tool_call(self) -> ToolCall: - """Build and return the tool call schema""" + """Build and return the tool call schema for single history""" return ToolCall( **{ - "description": "Read original history dialogue.", + "description": "Read a single original history dialogue by its ID.", "parameters": { "type": "object", "properties": { "history_id": { "type": "string", - "description": "history_id", + "description": "The history ID to read", }, }, "required": ["history_id"], @@ -31,17 +31,47 @@ class ReadHistory(BaseMemoryTool): }, ) - async def execute(self): - history_id = self.context.history_id - nodes = await self.vector_store.get(vector_ids=[history_id]) + def _build_multiple_tool_call(self) -> ToolCall: + """Build and return the tool call schema for multiple histories""" + return ToolCall( + **{ + "description": "Read multiple original history dialogues by their IDs.", + "parameters": { + "type": "object", + "properties": { + "history_ids": { + "type": "array", + "items": {"type": "string"}, + "description": "List of history IDs to read", + }, + }, + "required": ["history_ids"], + }, + }, + ) - if not nodes: - output = f"No history_id={history_id} data." + async def execute(self): + """Execute the tool call""" + history_ids = self.context.history_ids if self.enable_multiple else [self.context.history_id] + + if not history_ids or (len(history_ids) == 1 and not history_ids[0]): + output = "No history_ids provided." logger.warning(output) return output - memory_node: MemoryNode = MemoryNode.from_vector_node(nodes[0]) - self.retrieved_nodes.append(memory_node) - output = f"Historical Dialogue[{history_id}]\n{memory_node.content}" - logger.info(f"Successfully read history memory_node: {history_id}") + nodes = await self.vector_store.get(vector_ids=history_ids) + + if not nodes: + output = f"No data found for history_ids={history_ids}." + logger.warning(output) + return output + + results = [] + for node in nodes: + memory_node: MemoryNode = MemoryNode.from_vector_node(node) + self.retrieved_nodes.append(memory_node) + results.append(f"Historical Dialogue[{memory_node.memory_id}]\n{memory_node.content}") + + output = "\n\n".join(results) if self.enable_multiple else results[0] + logger.info(f"Successfully read {len(nodes)} history memory_node(s): {history_ids}") return output diff --git a/reme/tool/memory/profiles/update_profiles_v1.py b/reme/tool/memory/profiles/update_profiles_v1.py new file mode 100644 index 00000000..d38b3747 --- /dev/null +++ b/reme/tool/memory/profiles/update_profiles_v1.py @@ -0,0 +1,176 @@ +"""Update user profile tool""" + +from loguru import logger + +from .profile_handler import ProfileHandler +from ..base_memory_tool import BaseMemoryTool +from ....core.schema import ToolCall + + +class UpdateProfilesV1(BaseMemoryTool): + """Tool to update user profile by adding or removing profile entries""" + + def __init__(self, name="update_profiles", enable_memory_target: bool = False, **kwargs): + kwargs["enable_multiple"] = True + super().__init__(name=name, **kwargs) + self.enable_memory_target: bool = enable_memory_target + + def _build_profile_parameters(self, include_profile_id: bool = False) -> dict: + """Build the profile parameters schema based on enabled features.""" + properties = {} + required = [] + + if include_profile_id: + properties["profile_id"] = { + "type": "string", + "description": "ID of the profile to update", + } + required.append("profile_id") + + properties.update( + { + "message_time": { + "type": "string", + "description": "Message 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.extend(["message_time", "profile_key", "profile_value"]) + + if self.enable_memory_target: + properties["memory_target"] = { + "type": "string", + "description": "memory_target", + } + required.append("memory_target") + + return { + "type": "object", + "properties": properties, + "required": required, + } + + def _build_multiple_tool_call(self) -> ToolCall: + """Build and return the multiple tool call schema""" + return ToolCall( + **{ + "description": "Update existing profiles and add new profiles.", + "parameters": { + "type": "object", + "properties": { + "profiles_to_update": { + "type": "array", + "description": "List of profiles to update", + "items": self._build_profile_parameters(include_profile_id=True), + }, + "profiles_to_add": { + "type": "array", + "description": "List of profiles to add", + "items": self._build_profile_parameters(include_profile_id=False), + }, + }, + "required": ["profiles_to_update", "profiles_to_add"], + }, + }, + ) + + async def execute(self): + # Get parameters + profiles_to_update = self.context.get("profiles_to_update", []) + profiles_to_add = self.context.get("profiles_to_add", []) + + if not profiles_to_update and not profiles_to_add: + return "No profiles to update or add, operation completed." + + # Step 1: Collect and delete all old profiles that need to be updated + if profiles_to_update: + # Group deletion IDs by memory_target if enabled + if self.enable_memory_target: + delete_by_target = {} + for profile in profiles_to_update: + target = profile.get("memory_target", self.memory_target) + profile_id = profile.get("profile_id") + if profile_id: + if target not in delete_by_target: + delete_by_target[target] = [] + delete_by_target[target].append(profile_id) + else: + delete_by_target = { + self.memory_target: [ + profile.get("profile_id") for profile in profiles_to_update if profile.get("profile_id") + ], + } + + # Delete old profiles for each target + for target, profile_ids in delete_by_target.items(): + if profile_ids: + profile_ids = sorted(set(profile_ids)) # Remove duplicates and sort + profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=target) + profile_handler.delete(profile_ids) + + # Step 2: Prepare all profiles to add (both updated and new) + all_profiles_to_add = [] + + # Add profiles from updates + if profiles_to_update: + for profile in profiles_to_update: + target = ( + profile.get( + "memory_target", + self.memory_target, + ) + if self.enable_memory_target + else self.memory_target + ) + all_profiles_to_add.append((target, profile)) + + # Add new profiles + if profiles_to_add: + for profile in profiles_to_add: + target = ( + profile.get( + "memory_target", + self.memory_target, + ) + if self.enable_memory_target + else self.memory_target + ) + all_profiles_to_add.append((target, profile)) + + # Step 3: Group all profiles by target and add them in batch + from collections import defaultdict + + profiles_by_target = defaultdict(list) + for target, profile in all_profiles_to_add: + profiles_by_target[target].append(profile) + + # Process each target and add profiles + all_memory_nodes = [] + updated_count = len(profiles_to_update) + added_count = len(profiles_to_add) + + for target, target_profiles in profiles_by_target.items(): + profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=target) + new_nodes = profile_handler.add_batch(profiles=target_profiles, ref_memory_id=self.history_id) + all_memory_nodes.extend(new_nodes) + + # Extend memory_nodes for tracking + self.memory_nodes.extend(all_memory_nodes) + + # Build output message + operations = [] + if updated_count > 0: + operations.append(f"updated {updated_count} profiles.") + if added_count > 0: + operations.append(f"added {added_count} new profiles.") + operations.append("Operation completed.") + logger.info("\n".join(operations)) + return "\n".join(operations) diff --git a/reme/tool/memory/vector/add_draft_and_retrieve_similar_memory.py b/reme/tool/memory/vector/add_draft_and_retrieve_similar_memory.py new file mode 100644 index 00000000..040c475a --- /dev/null +++ b/reme/tool/memory/vector/add_draft_and_retrieve_similar_memory.py @@ -0,0 +1,127 @@ +"""Add draft memory and retrieve similar memories from vector store""" + +from loguru import logger + +from .memory_handler import MemoryHandler +from ..base_memory_tool import BaseMemoryTool +from ....core.schema import ToolCall, MemoryNode +from ....core.utils import deduplicate_memories + + +class AddDraftAndRetrieveSimilarMemory(BaseMemoryTool): + """Tool to add draft memory and retrieve similar memories""" + + def __init__( + self, + top_k: int = 20, + enable_memory_target: bool = False, + enable_when_to_use: bool = False, + **kwargs, + ): + super().__init__(**kwargs) + self.top_k: int = top_k + self.enable_memory_target: bool = enable_memory_target + self.enable_when_to_use: bool = enable_when_to_use + + def _build_query_parameters(self) -> dict: + """Build the query parameters schema""" + properties = { + "message_time": { + "type": "string", + "description": "message time, e.g. '2020-01-01 00:00:00'", + }, + "memory_content": { + "type": "string", + "description": "content of the memory.", + }, + } + required = ["message_time", "memory_content"] + + if self.enable_when_to_use: + properties["when_to_use"] = { + "type": "string", + "description": "description of when to use this memory.", + } + required.append("when_to_use") + + if self.enable_memory_target: + properties["memory_target"] = { + "type": "string", + "description": "target memory type for this memory.", + } + required.append("memory_target") + + return { + "type": "object", + "properties": properties, + "required": required, + } + + def _build_tool_call(self) -> ToolCall: + return ToolCall( + **{ + "description": "Add draft memory and retrieve similar memories from the vector store.", + "parameters": self._build_query_parameters(), + }, + ) + + def _build_multiple_tool_call(self) -> ToolCall: + return ToolCall( + **{ + "description": "Add draft memory and retrieve similar memories from the vector store.", + "parameters": { + "type": "object", + "properties": { + "draft_items": { + "type": "array", + "description": "draft_items", + "items": self._build_query_parameters(), + }, + }, + "required": ["draft_items"], + }, + }, + ) + + async def execute(self): + if self.enable_multiple: + draft_items = self.context.get("draft_items", []) + else: + draft_items = [self.context] + + queries_by_target: dict[str, list[dict]] = {} + for item in draft_items: + if self.enable_memory_target: + target = item["memory_target"] + else: + target = self.memory_target + if target not in queries_by_target: + queries_by_target[target] = [] + + queries_by_target[target].append( + { + "query": item["memory_content"], + "limit": self.top_k, + "filters": {}, + }, + ) + + # Execute batch searches for each target + memory_nodes: list[MemoryNode] = [] + for target, searches in queries_by_target.items(): + handler = MemoryHandler(target, self.service_context) + nodes = await handler.batch_search(searches) + memory_nodes.extend(nodes) + + memory_nodes = deduplicate_memories(memory_nodes) + retrieved_ids = {n.memory_id for n in self.retrieved_nodes if n.memory_id} + new_nodes = [n for n in memory_nodes if n.memory_id not in retrieved_ids] + self.retrieved_nodes.extend(new_nodes) + + if not new_nodes: + output = "No similar memories found." + else: + output = "\n".join([n.format(ref_memory_id_key="history_id") for n in new_nodes]) + + logger.info(f"Retrieved {len(memory_nodes)} similar memories, {len(new_nodes)} new after deduplication") + return output diff --git a/reme/tool/memory/vector/update_memory_v1.py b/reme/tool/memory/vector/update_memory_v1.py new file mode 100644 index 00000000..5d1551e7 --- /dev/null +++ b/reme/tool/memory/vector/update_memory_v1.py @@ -0,0 +1,212 @@ +"""Update memory in vector store""" + +from loguru import logger + +from .memory_handler import MemoryHandler +from ..base_memory_tool import BaseMemoryTool +from ....core.schema import ToolCall + + +class UpdateMemoryV1(BaseMemoryTool): + """Tool to update memories in vector store by deleting and adding memory entries""" + + def __init__( + self, + name="update_memory", + enable_memory_target: bool = False, + enable_when_to_use: bool = False, + **kwargs, + ): + kwargs["enable_multiple"] = True + super().__init__(name=name, **kwargs) + self.enable_memory_target: bool = enable_memory_target + self.enable_when_to_use: bool = enable_when_to_use + + def _build_memory_parameters(self, include_memory_id: bool = False) -> dict: + """Build the memory parameters schema based on enabled features. + + Args: + include_memory_id: If True, include memory_id field (for updates) + """ + properties = {} + required = [] + + if include_memory_id: + properties["memory_id"] = { + "type": "string", + "description": "ID of the memory to update", + } + required.append("memory_id") + + properties.update( + { + "message_time": { + "type": "string", + "description": "message time, e.g. '2020-01-01 00:00:00'", + }, + "memory_content": { + "type": "string", + "description": "content of the memory.", + }, + }, + ) + required.extend(["message_time", "memory_content"]) + + if self.enable_when_to_use: + properties["when_to_use"] = { + "type": "string", + "description": "description of when to use this memory.", + } + required.append("when_to_use") + + if self.enable_memory_target: + properties["memory_target"] = { + "type": "string", + "description": "target memory type for this memory.", + } + required.append("memory_target") + + return { + "type": "object", + "properties": properties, + "required": required, + } + + def _build_multiple_tool_call(self) -> ToolCall: + """Build and return the multiple tool call schema""" + return ToolCall( + **{ + "description": "update memories by updating existing memories and adding new memory entries.", + "parameters": { + "type": "object", + "properties": { + "memories_to_update": { + "type": "array", + "description": "List of memories to update", + "items": self._build_memory_parameters(include_memory_id=True), + }, + "memories_to_add": { + "type": "array", + "description": "List of memories to add", + "items": self._build_memory_parameters(include_memory_id=False), + }, + }, + "required": ["memories_to_update", "memories_to_add"], + }, + }, + ) + + async def execute(self): + # Get parameters + memories_to_update = self.context.get("memories_to_update", []) + memories_to_add = self.context.get("memories_to_add", []) + + if not memories_to_update and not memories_to_add: + return "No memories to update or add, operation completed." + + # Step 1: Collect and delete all old memories that need to be updated + if memories_to_update: + # Group deletion IDs by memory_target if enabled + if self.enable_memory_target: + delete_by_target = {} + for mem in memories_to_update: + target = mem.get("memory_target", self.memory_target) + memory_id = mem.get("memory_id") + if memory_id: + if target not in delete_by_target: + delete_by_target[target] = [] + delete_by_target[target].append(memory_id) + else: + delete_by_target = { + self.memory_target: [mem.get("memory_id") for mem in memories_to_update if mem.get("memory_id")], + } + + # Delete old memories for each target + for target, memory_ids in delete_by_target.items(): + if memory_ids: + handler = MemoryHandler(target, self.service_context) + await handler.delete(memory_ids) + + # Step 2: Prepare all memories to add (both updated and new) + all_memories_to_add = [] + + # Add memories from updates + if memories_to_update: + for mem in memories_to_update: + target = ( + mem.get( + "memory_target", + self.memory_target, + ) + if self.enable_memory_target + else self.memory_target + ) + all_memories_to_add.append((target, mem)) + + # Add new memories + if memories_to_add: + for mem in memories_to_add: + target = ( + mem.get( + "memory_target", + self.memory_target, + ) + if self.enable_memory_target + else self.memory_target + ) + all_memories_to_add.append((target, mem)) + + # Step 3: Group all memories by target and add them in batch + memories_by_target = {} + for target, mem in all_memories_to_add: + if target not in memories_by_target: + memories_by_target[target] = [] + memories_by_target[target].append(mem) + + # Process each target and add memories + all_memory_nodes = [] + updated_count = len(memories_to_update) + added_count = len(memories_to_add) + + for target, target_memories in memories_by_target.items(): + # Prepare memory data for batch add + memory_dicts = [] + for mem in target_memories: + memory_content = mem.get("memory_content", "") + message_time = mem.get("message_time", "") + when_to_use = mem.get("when_to_use", "") if self.enable_when_to_use else "" + metadata = {} + try: + metadata["time_int"] = int(message_time.split(" ")[0].replace("-", "")) + except Exception: + logger.warning(f"Invalid message time format: {message_time}") + + memory_dicts.append( + { + "content": memory_content, + "when_to_use": when_to_use, + "message_time": message_time, + "ref_memory_id": self.history_id, + "author": self.author, + "metadata": metadata, + }, + ) + + # Batch add all memories for this target + if memory_dicts: + handler = MemoryHandler(target, self.service_context) + memory_nodes = await handler.add_batch(memory_dicts) + all_memory_nodes.extend(memory_nodes) + + # Extend memory_nodes for tracking + self.memory_nodes.extend(all_memory_nodes) + + # Build output message + operations = [] + if updated_count > 0: + operations.append(f"updated {updated_count} memories.") + if added_count > 0: + operations.append(f"added {added_count} new memories.") + operations.append("Operation completed.") + logger.info("\n".join(operations)) + return "\n".join(operations)