diff --git a/reme/agent/memory/default/personal_retriever.py b/reme/agent/memory/default/personal_retriever.py index 7fc8a0c8..3ff50eaf 100644 --- a/reme/agent/memory/default/personal_retriever.py +++ b/reme/agent/memory/default/personal_retriever.py @@ -22,8 +22,10 @@ class PersonalRetriever(BaseMemoryAgent): read_all_profiles_tool: BaseTool | None = self.pop_tool("read_all_profiles") if read_all_profiles_tool is not None: - all_profiles = await read_all_profiles_tool.call(memory_target=self.memory_target, - service_context=self.service_context) + all_profiles = await read_all_profiles_tool.call( + memory_target=self.memory_target, + service_context=self.service_context, + ) else: all_profiles = "" diff --git a/reme/agent/memory/default/personal_retriever.yaml b/reme/agent/memory/default/personal_retriever.yaml index da8eead2..47de7c6d 100644 --- a/reme/agent/memory/default/personal_retriever.yaml +++ b/reme/agent/memory/default/personal_retriever.yaml @@ -6,7 +6,7 @@ system_prompt: | ## User Question {context} - + ## Retrieval Strategy ### Phase 1 `retrieve_memory` - Purpose: Search for relevant memories using semantic similarity diff --git a/reme/agent/memory/default/personal_summarizer.py b/reme/agent/memory/default/personal_summarizer.py index 5885ed23..9f8b85de 100644 --- a/reme/agent/memory/default/personal_summarizer.py +++ b/reme/agent/memory/default/personal_summarizer.py @@ -1,4 +1,5 @@ """Personal memory summarizer agent for two-phase personal memory processing.""" + from loguru import logger from ..base_memory_agent import BaseMemoryAgent diff --git a/reme/agent/memory/default/personal_summarizer.yaml b/reme/agent/memory/default/personal_summarizer.yaml index ada57fd4..014da6e5 100644 --- a/reme/agent/memory/default/personal_summarizer.yaml +++ b/reme/agent/memory/default/personal_summarizer.yaml @@ -1,15 +1,15 @@ system_prompt_s1_zh: | 你是一个记忆Agent,负责管理关于 {memory_target} 的 {memory_type} 类型记忆。 - + ## 最新对话 Format: round [] : {context} - + ## 任务 ### 步骤1 根据`最新对话`的内容,在 `add_draft_and_retrieve_similar_memory` 中创建记忆草稿 `memory_draft`。 工具会根据memory_draft的内容行向量检索,返回历史相似记忆,确保在第二步的时候更好的管理记忆库记忆。 - + ### 步骤2 使用`update_memory`更新向量库记忆。 通过`memory_ids_to_delete`删除历史记忆,`memories_to_add`添加新记忆,包括message_time和memory_content。 @@ -23,16 +23,16 @@ user_message_s1_zh: | system_prompt_s2_zh: | 你是一个Profile Agent,负责管理关于 {memory_target} 的 Profile。 - + ## 最新对话 Format: round [] : {context} - + ## 任务 ### 步骤1 根据`最新对话`的内容,在 `add_draft_and_read_all_profiles` 中创建记忆草稿 `profile_draft`。 工具会直接返回所有的Profile,确保在第二步的时候更好的管理Profile。 - + ### 步骤2 使用`update_profile`更新profile库。 通过`profile_ids_to_delete`删除历史Profile,`profiles_to_add`添加新Profile,包括message_time、profile_key和profile_value。 @@ -46,16 +46,16 @@ user_message_s2_zh: | system_prompt_s1: | You are a Memory Agent responsible for managing {memory_type} type memories about {memory_target}. - + ## Latest Conversation Format: round [] : {context} - + ## Task ### Step 1 Based on the content of `Latest Conversation`, create a memory draft `memory_draft` in `add_draft_and_retrieve_similar_memory`. The tool will perform vector retrieval based on the content of memory_draft and return historically similar memories to better manage the memory store in Step 2. - + ### Step 2 Use `update_memory` to update the vector store memories. Delete historical memories through `memory_ids_to_delete`, add new memories through `memories_to_add`, including message_time and memory_content. @@ -69,16 +69,16 @@ user_message_s1: | system_prompt_s2: | You are a Profile Agent responsible for managing the Profile about {memory_target}. - + ## Latest Conversation Format: round [] : {context} - + ## Task ### Step 1 Based on the content of `Latest Conversation`, create a profile draft `profile_draft` in `add_draft_and_read_all_profiles`. The tool will directly return all Profiles to better manage the Profile store in Step 2. - + ### Step 2 Use `update_profile` to update the profile store. Delete historical Profiles through `profile_ids_to_delete`, add new Profiles through `profiles_to_add`, including message_time, profile_key, and profile_value. diff --git a/reme/agent/memory/default/procedural_retriever.py b/reme/agent/memory/default/procedural_retriever.py index 63d4856b..090d66f0 100644 --- a/reme/agent/memory/default/procedural_retriever.py +++ b/reme/agent/memory/default/procedural_retriever.py @@ -1,6 +1,10 @@ +"""Procedural memory retriever agent implementation.""" + from ..base_memory_agent import BaseMemoryAgent from ....core.enumeration import MemoryType class ProceduralRetriever(BaseMemoryAgent): + """Agent responsible for retrieving procedural memories.""" + memory_type: MemoryType = MemoryType.PROCEDURAL diff --git a/reme/agent/memory/default/procedural_summarizer.py b/reme/agent/memory/default/procedural_summarizer.py index 0be6bd98..72876568 100644 --- a/reme/agent/memory/default/procedural_summarizer.py +++ b/reme/agent/memory/default/procedural_summarizer.py @@ -1,6 +1,10 @@ +"""Procedural memory summarizer agent implementation.""" + from ..base_memory_agent import BaseMemoryAgent from ....core.enumeration import MemoryType class ProceduralSummarizer(BaseMemoryAgent): + """Agent responsible for summarizing procedural memories.""" + memory_type: MemoryType = MemoryType.PROCEDURAL diff --git a/reme/agent/memory/default/tool_retriever.py b/reme/agent/memory/default/tool_retriever.py index 53e603fc..6a2d206b 100644 --- a/reme/agent/memory/default/tool_retriever.py +++ b/reme/agent/memory/default/tool_retriever.py @@ -1,6 +1,10 @@ +"""Tool memory retriever agent implementation.""" + from ..base_memory_agent import BaseMemoryAgent from ....core.enumeration import MemoryType class ToolRetriever(BaseMemoryAgent): + """Agent responsible for retrieving tool-related memories.""" + memory_type: MemoryType = MemoryType.TOOL diff --git a/reme/agent/memory/default/tool_summarizer.py b/reme/agent/memory/default/tool_summarizer.py index 5064b6db..85222d4f 100644 --- a/reme/agent/memory/default/tool_summarizer.py +++ b/reme/agent/memory/default/tool_summarizer.py @@ -1,6 +1,10 @@ +"""Tool memory summarizer agent implementation.""" + from ..base_memory_agent import BaseMemoryAgent from ....core.enumeration import MemoryType class ToolSummarizer(BaseMemoryAgent): + """Agent responsible for summarizing tool-related memories.""" + memory_type: MemoryType = MemoryType.TOOL diff --git a/reme/core/application.py b/reme/core/application.py index a03c3e73..f96b1b03 100644 --- a/reme/core/application.py +++ b/reme/core/application.py @@ -1,3 +1,5 @@ +"""High-level entry point for configuring and running ReMe services and flows.""" + import asyncio from .context import PromptHandler, ServiceContext @@ -11,21 +13,22 @@ from .vector_store import BaseVectorStore class Application: + """Application wrapper that wires together service context, flows, and runtimes.""" def __init__( - self, - *args, - llm_api_key: str | None = None, - llm_api_base: str | None = None, - embedding_api_key: str | None = None, - embedding_api_base: str | None = None, - enable_logo: bool = True, - parser: type[PydanticConfigParser] | None = None, - llm: dict | None = None, - embedding_model: dict | None = None, - vector_store: dict | None = None, - token_counter: dict | None = None, - **kwargs, + self, + *args, + llm_api_key: str | None = None, + llm_api_base: str | None = None, + embedding_api_key: str | None = None, + embedding_api_base: str | None = None, + enable_logo: bool = True, + parser: type[PydanticConfigParser] | None = None, + llm: dict | None = None, + embedding_model: dict | None = None, + vector_store: dict | None = None, + token_counter: dict | None = None, + **kwargs, ): # ServiceContext self.service_context = ServiceContext( @@ -94,10 +97,10 @@ class Application: stream_queue = asyncio.Queue() task = asyncio.create_task(flow.call(stream_queue=stream_queue, **kwargs)) async for chunk in execute_stream_task( - stream_queue=stream_queue, - task=task, - task_name=name, - as_bytes=False, + stream_queue=stream_queue, + task=task, + task_name=name, + as_bytes=False, ): yield chunk diff --git a/reme/core/schema/memory_node.py b/reme/core/schema/memory_node.py index 3cf096ac..e6fde012 100644 --- a/reme/core/schema/memory_node.py +++ b/reme/core/schema/memory_node.py @@ -149,12 +149,12 @@ class MemoryNode(BaseModel): ) def format( - self, - include_memory_id: bool = True, - include_when_to_use: bool = True, - include_content: bool = True, - include_message_time: bool = True, - ref_memory_id_key: str = "", + self, + include_memory_id: bool = True, + include_when_to_use: bool = True, + include_content: bool = True, + include_message_time: bool = True, + ref_memory_id_key: str = "", ) -> str: """Format memory node as string with configurable fields.""" line = "" diff --git a/reme/reme.py b/reme/reme.py index bca18f80..08ca92ca 100644 --- a/reme/reme.py +++ b/reme/reme.py @@ -1,25 +1,36 @@ """ReMe classes for simplified configuration and execution.""" -import asyncio import sys from pathlib import Path -from loguru import logger - -from .core import Application -from .agent.memory.default import ReMeSummarizer, PersonalSummarizer, PersonalRetriever, ReMeRetriever +from reme.agent.memory import BaseMemoryAgent +from .agent.memory.default import ( + ReMeSummarizer, + PersonalSummarizer, + PersonalRetriever, + ReMeRetriever, + ProceduralSummarizer, + ToolSummarizer, + ProceduralRetriever, + ToolRetriever, +) from .config import ReMeConfigParser -from .core.context import PromptHandler, ServiceContext -from .core.embedding import BaseEmbeddingModel +from .core import Application from .core.enumeration import MemoryType -from .core.flow import BaseFlow -from .core.llm import BaseLLM -from .core.schema import Response, Message, MemoryNode, VectorNode -from .core.token_counter import BaseTokenCounter -from .core.utils import execute_stream_task, get_now_time -from .core.vector_store import BaseVectorStore -from .tool.memory import UpdateUserProfile, RetrieveMemory, AddMemory, DelegateTask, ReadHistory, ReadUserProfile, \ - ProfileHandler +from .core.schema import Message +from .tool.memory import ( + RetrieveMemory, + DelegateTask, + ReadHistory, + ProfileHandler, + MemoryHandler, + AddDraftAndRetrieveSimilarMemory, + UpdateMemoryV2, + AddDraftAndReadAllProfiles, + UpdateProfile, + AddHistory, + ReadAllProfiles, +) class ReMe(Application): @@ -37,18 +48,10 @@ class ReMe(Application): embedding_model: dict | None = None, vector_store: dict | None = None, token_counter: dict | None = None, - personal_memory_target: list[str] | None = None, - procedural_memory_target: list[str] | None = None, - tool_memory_target: list[str] | None = None, - profile_path: str = "reme_profile", - main_summary_version: str = "default", - personal_summary_version: str = "default", - procedural_summary_version: str = "default", - tool_summary_version: str = "default", - main_retrieve_version: str = "default", - personal_retrieve_version: str = "default", - procedural_retrieve_version: str = "default", - tool_retrieve_version: str = "default", + personal_memory_target: list[str] | None = None, + procedural_memory_target: list[str] | None = None, + tool_memory_target: list[str] | None = None, + profile_dir: str = "reme_profile", **kwargs, ): super().__init__( @@ -80,300 +83,239 @@ class ReMe(Application): for name in tool_memory_target: assert name not in memory_target_type_mapping, f"Memory target name {name} is already used." memory_target_type_mapping[name] = MemoryType.TOOL + self.service_context.memory_target_type_mapping = memory_target_type_mapping - self.profile_path: str = profile_path - - @property - def memory_target_type_mapping(self) -> dict[str, MemoryType]: - mapping = {} - if self.service_context.personal_memory_target: - for name in self.service_context.personal_memory_target: - assert name not in mapping, f"Memory target name {name} is already used." - mapping[name] = MemoryType.PERSONAL - - if self.service_context.procedural_memory_target: - for name in self.service_context.procedural_memory_target: - assert name not in mapping, f"Memory target name {name} is already used." - mapping[name] = MemoryType.PROCEDURAL - - if self.service_context.tool_memory_target: - for name in self.service_context.tool_memory_target: - assert name not in mapping, f"Memory target name {name} is already used." - mapping[name] = MemoryType.TOOL - return mapping + self.profile_dir: str = profile_dir def add_meta_memory(self, memory_type: str | MemoryType, memory_target: str): - memory_type = MemoryType(memory_type) - if memory_type is MemoryType.PERSONAL: - personal_memory_target = self.service_context.personal_memory_target - if memory_target not in personal_memory_target: - personal_memory_target.append(memory_target) - else: - logger.warning(f"Memory target {memory_target} is already added.") - - elif memory_type is MemoryType.PROCEDURAL: - procedural_memory_target = self.service_context.procedural_memory_target - if memory_target not in procedural_memory_target: - procedural_memory_target.append(memory_target) - else: - logger.warning(f"Memory target {memory_target} is already added.") - - elif memory_type is MemoryType.TOOL: - tool_memory_target = self.service_context.tool_memory_target - if memory_target not in tool_memory_target: - tool_memory_target.append(memory_target) - else: - logger.warning(f"Memory target {memory_target} is already added.") - - + """Register or validate a memory target with the given memory type.""" + if memory_target in self.service_context.memory_target_type_mapping: + assert self.service_context.memory_target_type_mapping[memory_target] is memory_type + else: + self.service_context.memory_target_type_mapping[memory_target] = MemoryType(memory_type) async def summary_memory( self, messages: list[Message | dict], description: str = "", - user_name: str = "", - task_name: str = "", - tool_name: str = "", + user_name: str | list[str] = "", + task_name: str | list[str] = "", + tool_name: str | list[str] = "", enable_thinking_params: bool = False, version: str = "default", + retrieve_top_k: int = 20, return_dict: bool = False, **kwargs, ) -> str | dict: - """Summarize messages and store them in memory for the specified user(s).""" - if user_name: - if isinstance(user_name, str): - for message in messages: - if isinstance(message, dict) and not message.get("name"): - message["name"] = user_name - elif isinstance(message, Message) and not message.name: - message.name = user_name - user_name = [user_name] + """Summarize personal, procedural and tool memories for the given context.""" + format_messages: list[Message] = [] + for message in messages: + if isinstance(message, dict): + assert message.get("time_created"), "message must have time_created field." + message = Message(**message) + format_messages.append(message) - if not meta_memories: - meta_memories = [ - { - "memory_type": "personal", - "memory_target": name, - } - for name in user_name - ] - - if version == "default": - reme_summarizer = ReMeSummarizer( - meta_memories=meta_memories, - tools=[ - DelegateTask( - memory_agents=[ - PersonalSummarizer( - tools=[ - RetrieveMemory(enable_thinking_params=enable_thinking_params), - AddMemory(enable_thinking_params=enable_thinking_params), - UpdateUserProfile(enable_thinking_params=enable_thinking_params), - ], - ), - ], - ), - ], - ) - - else: - raise NotImplementedError - - result = await reme_summarizer.call( - messages=messages, - description=description, - service_context=self.service_context, - **kwargs, + personal_summarizer: BaseMemoryAgent + if version: + personal_summarizer = PersonalSummarizer( + tools=[ + AddDraftAndRetrieveSimilarMemory( + enable_thinking_params=enable_thinking_params, + top_k=retrieve_top_k, + ), + UpdateMemoryV2(enable_thinking_params=enable_thinking_params), + AddDraftAndReadAllProfiles( + enable_thinking_params=enable_thinking_params, + profile_dir=self.profile_dir, + ), + UpdateProfile( + enable_thinking_params=enable_thinking_params, + profile_dir=self.profile_dir, + ), + ], ) - - if return_dict: - return result - else: - return result["answer"] - else: raise NotImplementedError + procedural_summarizer: BaseMemoryAgent + if version == "default": + procedural_summarizer = ProceduralSummarizer(tools=[]) + else: + raise NotImplementedError + + tool_summarizer: BaseMemoryAgent + if version == "default": + tool_summarizer = ToolSummarizer(tools=[]) + else: + raise NotImplementedError + + memory_agents = [] + if user_name: + if isinstance(user_name, str): + for message in format_messages: + message.name = user_name + self.add_meta_memory(MemoryType.PERSONAL, user_name) + elif isinstance(user_name, list): + for name in user_name: + self.add_meta_memory(MemoryType.PERSONAL, name) + else: + raise RuntimeError("user_name must be str or list[str]") + memory_agents.append(personal_summarizer) + + if task_name: + if isinstance(task_name, str): + self.add_meta_memory(MemoryType.PROCEDURAL, task_name) + elif isinstance(task_name, list): + for name in task_name: + self.add_meta_memory(MemoryType.PROCEDURAL, name) + else: + raise RuntimeError("task_name must be str or list[str]") + memory_agents.append(procedural_summarizer) + + if tool_name: + if isinstance(tool_name, str): + self.add_meta_memory(MemoryType.TOOL, tool_name) + elif isinstance(tool_name, list): + for name in tool_name: + self.add_meta_memory(MemoryType.TOOL, name) + else: + raise RuntimeError("tool_name must be str or list[str]") + memory_agents.append(tool_summarizer) + + if not memory_agents: + memory_agents = [personal_summarizer, procedural_summarizer, tool_summarizer] + + reme_summarizer: BaseMemoryAgent + if version == "default": + reme_summarizer = ReMeSummarizer(tools=[AddHistory(), DelegateTask(memory_agents=memory_agents)]) + else: + raise NotImplementedError + + result = await reme_summarizer.call( + messages=messages, + description=description, + service_context=self.service_context, + **kwargs, + ) + + if return_dict: + return result + else: + return result["answer"] + async def retrieve_memory( self, query: str = "", - top_k: int = 20, description: str = "", messages: list[dict] | None = None, user_name: str | list[str] = "", + task_name: str | list[str] = "", + tool_name: str | list[str] = "", enable_thinking_params: bool = False, - meta_memories: list[dict] = None, version: str = "default", + retrieve_top_k: int = 20, + enable_memory_target: bool = True, return_dict: bool = False, **kwargs, ) -> str | dict: - """Retrieve relevant memories for the specified user(s) based on query or messages.""" - if user_name: - if isinstance(user_name, str): - if messages: - for message in messages: - if isinstance(message, dict) and not message.get("name"): - message["name"] = user_name - elif isinstance(message, Message) and not message.name: - message.name = user_name - user_name = [user_name] + """Retrieve relevant personal, procedural and tool memories for a query.""" - if not meta_memories: - meta_memories = [ - { - "memory_type": "personal", - "memory_target": name, - } - for name in user_name - ] - - if version == "default": - reme_retriever = ReMeRetriever( - meta_memories=meta_memories, - tools=[ - DelegateTask( - memory_agents=[ - PersonalRetriever( - tools=[ - RetrieveMemory(enable_thinking_params=enable_thinking_params, top_k=top_k), - ReadHistory(enable_thinking_params=enable_thinking_params), - ], - ), - ], - ), - ], - ) - - else: - raise NotImplementedError - - result = await reme_retriever.call( - query=query, - messages=messages, - description=description, - service_context=self.service_context, - **kwargs, + personal_retriever: BaseMemoryAgent + if version: + personal_retriever = PersonalRetriever( + tools=[ + ReadAllProfiles( + enable_thinking_params=enable_thinking_params, + profile_dir=self.profile_dir, + ), + RetrieveMemory( + enable_thinking_params=enable_thinking_params, + top_k=retrieve_top_k, + enable_memory_target=enable_memory_target, + ), + ReadHistory(enable_thinking_params=enable_thinking_params), + ], ) - - if return_dict: - return result - else: - return result["answer"] - else: raise NotImplementedError - async def add_memory( - self, - memory_content: str, - user_name: str, - memory_type: str | MemoryType | None = None, - memory_target: str = "", - when_to_use: str = "", - ref_memory_id: str = "", - author: str = "", - score: float = 0, - conversation_time: str = "", - **kwargs, - ) -> MemoryNode: - """Add a new memory to the vector store for the specified user.""" + procedural_retriever: BaseMemoryAgent + if version == "default": + procedural_retriever = ProceduralRetriever(tools=[]) + else: + raise NotImplementedError + tool_retriever: BaseMemoryAgent + if version == "default": + tool_retriever = ToolRetriever(tools=[]) + else: + raise NotImplementedError + + memory_agents = [] if user_name: - memory_type = MemoryType.PERSONAL - memory_target = user_name + if isinstance(user_name, str): + self.add_meta_memory(MemoryType.PERSONAL, user_name) + elif isinstance(user_name, list): + for name in user_name: + self.add_meta_memory(MemoryType.PERSONAL, name) + else: + raise RuntimeError("user_name must be str or list[str]") + memory_agents.append(personal_retriever) + + if task_name: + if isinstance(task_name, str): + self.add_meta_memory(MemoryType.PROCEDURAL, task_name) + elif isinstance(task_name, list): + for name in task_name: + self.add_meta_memory(MemoryType.PROCEDURAL, name) + else: + raise RuntimeError("task_name must be str or list[str]") + memory_agents.append(procedural_retriever) + + if tool_name: + if isinstance(tool_name, str): + self.add_meta_memory(MemoryType.TOOL, tool_name) + elif isinstance(tool_name, list): + for name in tool_name: + self.add_meta_memory(MemoryType.TOOL, name) + else: + raise RuntimeError("tool_name must be str or list[str]") + memory_agents.append(tool_retriever) + + if not memory_agents: + memory_agents = [personal_retriever, procedural_retriever, tool_retriever] + + reme_retriever: BaseMemoryAgent + if version == "default": + reme_retriever = ReMeRetriever(tools=[DelegateTask(memory_agents=memory_agents)]) else: - memory_type = MemoryType(memory_type) - assert memory_target, "memory_target is required" + raise NotImplementedError - metadata = kwargs.copy() - if conversation_time: - metadata["conversation_time"] = conversation_time - - memory_node = MemoryNode( - memory_type=memory_type, - memory_target=memory_target, - when_to_use=when_to_use, - content=memory_content, - ref_memory_id=ref_memory_id, - author=author, - score=score, - metadata=metadata, + result = await reme_retriever.call( + query=query, + messages=messages, + description=description, + service_context=self.service_context, + **kwargs, ) - vector_node = memory_node.to_vector_node() - await self.vector_store.delete([vector_node.vector_id]) - await self.vector_store.insert([vector_node]) - return memory_node - - async def update_memory( - self, - memory_id: str, - memory_content: str, - user_name: str, - memory_type: str | MemoryType | None = None, - memory_target: str = "", - when_to_use: str = "", - ref_memory_id: str = "", - author: str = "", - score: float = 0, - conversation_time: str = "", - **kwargs, - ) -> MemoryNode: - """Update an existing memory in the vector store by its ID.""" - - if user_name: - memory_type = MemoryType.PERSONAL - memory_target = user_name + if return_dict: + return result else: - memory_type = MemoryType(memory_type) - assert memory_target, "memory_target is required" + return result["answer"] - metadata = kwargs.copy() - if conversation_time: - metadata["conversation_time"] = conversation_time + @property + def profile_path(self) -> Path: + """Get the path to the profile directory.""" + return Path(self.profile_dir) / self.vector_store.collection_name - memory_node = MemoryNode( - memory_type=memory_type, - memory_target=memory_target, - when_to_use=when_to_use, - content=memory_content, - ref_memory_id=ref_memory_id, - author=author, - score=score, - metadata=metadata, - ) - vector_node = memory_node.to_vector_node() - await self.vector_store.delete(list(set([memory_id, vector_node.vector_id]))) - await self.vector_store.insert([vector_node]) - - return memory_node - - async def delete_memory(self, memory_id: str | list[str]): - """Delete one or more memories from the vector store by their IDs.""" - vector_ids = [memory_id] if isinstance(memory_id, str) else memory_id - await self.vector_store.delete(list(set(vector_ids))) - - async def delete_all_memories(self): - """Delete all memories from the vector store.""" - await self.vector_store.delete_all() - - async def get_memory(self, memory_id: str | list[str]) -> MemoryNode | list[MemoryNode]: - """Retrieve one or more memories from the vector store by their IDs.""" - vector_ids = [memory_id] if isinstance(memory_id, str) else memory_id - vector_nodes = await self.vector_store.get(vector_ids) - if isinstance(vector_nodes, VectorNode): - return vector_nodes.to_memory_node() - else: - return [node.to_memory_node() for node in vector_nodes] - - async def get_all_memories(self) -> list[MemoryNode]: - """Retrieve all memories from the vector store.""" - return [node.to_memory_node() for node in await self.vector_store.list()] + def get_memory_handler(self, memory_target: str) -> MemoryHandler: + """Get the memory handler for the specified memory target.""" + return MemoryHandler(memory_target=memory_target, service_context=self.service_context) def get_profile_handler(self, user_name: str) -> ProfileHandler: """Get the profile handler for the specified user.""" - profile_path = Path(self.profile_path) / self.vector_store.collection_name - return ProfileHandler(memory_target=user_name, profile_path=profile_path) + return ProfileHandler(memory_target=user_name, profile_path=self.profile_path) async def context_offload(self): """working memory summary""" diff --git a/reme/tool/memory/__init__.py b/reme/tool/memory/__init__.py index f0f61870..cab968dd 100644 --- a/reme/tool/memory/__init__.py +++ b/reme/tool/memory/__init__.py @@ -39,4 +39,4 @@ __all__ = [ for name in __all__: tool_class = globals()[name] - R.op.register()(tool_class) \ No newline at end of file + R.op.register()(tool_class) diff --git a/reme/tool/memory/add_draft_and_read_all_profiles.py b/reme/tool/memory/add_draft_and_read_all_profiles.py index c11fc82d..da4e559c 100644 --- a/reme/tool/memory/add_draft_and_read_all_profiles.py +++ b/reme/tool/memory/add_draft_and_read_all_profiles.py @@ -1,5 +1,4 @@ """Add draft profile and read all profiles from local storage""" -from pathlib import Path from loguru import logger @@ -11,9 +10,8 @@ from ...core.schema import ToolCall class AddDraftAndReadAllProfiles(BaseMemoryTool): """Tool to add draft profile and read all profiles""" - def __init__(self, profile_path: str, enable_memory_target: bool = False, **kwargs): + def __init__(self, enable_memory_target: bool = False, **kwargs): super().__init__(**kwargs) - self.profile_path: str = profile_path self.enable_memory_target: bool = enable_memory_target def _build_query_parameters(self) -> dict: @@ -86,10 +84,7 @@ class AddDraftAndReadAllProfiles(BaseMemoryTool): continue targets_processed.add(target) - profile_handler = ProfileHandler( - profile_path=Path(self.profile_path) / self.vector_store.collection_name, - memory_target=target, - ) + profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=target) profiles_str = profile_handler.read_all(add_profile_id=True) if profiles_str: diff --git a/reme/tool/memory/add_draft_and_retrieve_similar_memory.py b/reme/tool/memory/add_draft_and_retrieve_similar_memory.py index 8be42636..cb537fc4 100644 --- a/reme/tool/memory/add_draft_and_retrieve_similar_memory.py +++ b/reme/tool/memory/add_draft_and_retrieve_similar_memory.py @@ -80,11 +80,13 @@ class AddDraftAndRetrieveSimilarMemory(BaseMemoryTool): if target not in queries_by_target: queries_by_target[target] = [] - queries_by_target[target].append({ - "query": item["memory_draft"], - "limit": self.top_k, - "filters": {}, - }) + queries_by_target[target].append( + { + "query": item["memory_draft"], + "limit": self.top_k, + "filters": {}, + }, + ) # Execute batch searches for each target memory_nodes: list[MemoryNode] = [] diff --git a/reme/tool/memory/add_memory.py b/reme/tool/memory/add_memory.py index 1293634f..833c2b54 100644 --- a/reme/tool/memory/add_memory.py +++ b/reme/tool/memory/add_memory.py @@ -11,10 +11,10 @@ class AddMemory(BaseMemoryTool): """Tool to add memories to vector store""" def __init__( - self, - enable_memory_target: bool = False, - enable_when_to_use: bool = False, - **kwargs, + self, + enable_memory_target: bool = False, + enable_when_to_use: bool = False, + **kwargs, ): super().__init__(**kwargs) self.enable_memory_target: bool = enable_memory_target @@ -112,14 +112,16 @@ class AddMemory(BaseMemoryTool): 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, - }) + 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, + }, + ) if memory_dicts: handler = MemoryHandler(target, self.service_context) diff --git a/reme/tool/memory/base_memory_tool.py b/reme/tool/memory/base_memory_tool.py index 3d46e272..d3cb4541 100644 --- a/reme/tool/memory/base_memory_tool.py +++ b/reme/tool/memory/base_memory_tool.py @@ -1,6 +1,7 @@ """Base class for memory tool""" from abc import ABCMeta +from pathlib import Path from ...core.enumeration import MemoryType from ...core.op import BaseTool @@ -14,11 +15,13 @@ class BaseMemoryTool(BaseTool, metaclass=ABCMeta): self, enable_multiple: bool = True, enable_thinking_params: bool = False, + profile_dir: str = "", **kwargs, ): super().__init__(**kwargs) self.enable_multiple: bool = enable_multiple self.enable_thinking_params: bool = enable_thinking_params + self.profile_dir: str = profile_dir def _build_tool_call(self) -> ToolCall: """Build and return the tool call schema""" @@ -98,3 +101,8 @@ class BaseMemoryTool(BaseTool, metaclass=ABCMeta): def memory_target_type_mapping(self) -> dict[str, MemoryType]: """Get the memory target type mapping from context.""" return self.context.memory_target_type_mapping + + @property + def profile_path(self) -> Path: + """Get the path to the profile directory for the current collection.""" + return Path(self.profile_dir) / self.vector_store.collection_name diff --git a/reme/tool/memory/memory_handler.py b/reme/tool/memory/memory_handler.py index 78151325..5b896449 100644 --- a/reme/tool/memory/memory_handler.py +++ b/reme/tool/memory/memory_handler.py @@ -1,3 +1,5 @@ +"""Memory handler""" + from ...core.context import ServiceContext from ...core.enumeration import MemoryType from ...core.schema import MemoryNode @@ -45,14 +47,14 @@ class MemoryHandler: return memory_nodes async def add( - self, - content: str, - when_to_use: str = "", - message_time: str = "", - ref_memory_id: str = "", - author: str = "", - score: float = 0.0, - **kwargs, + self, + content: str, + when_to_use: str = "", + message_time: str = "", + ref_memory_id: str = "", + author: str = "", + score: float = 0.0, + **kwargs, ) -> MemoryNode: """Add a single memory node and return its memory_id.""" memory_dict = { @@ -123,15 +125,15 @@ class MemoryHandler: return updated_nodes async def update( - self, - memory_id: str, - content: str | None = None, - when_to_use: str | None = None, - message_time: str | None = None, - ref_memory_id: str | None = None, - author: str | None = None, - score: float | None = None, - **kwargs, + self, + memory_id: str, + content: str | None = None, + when_to_use: str | None = None, + message_time: str | None = None, + ref_memory_id: str | None = None, + author: str | None = None, + score: float | None = None, + **kwargs, ) -> MemoryNode: """Update a memory node's content, when_to_use, or other fields.""" update_dict: dict = {"memory_id": memory_id} @@ -154,11 +156,11 @@ class MemoryHandler: return memory_nodes[0] async def search( - self, - query: str | list[str], - limit: int = 5, - filters: dict | None = None, - **kwargs, + self, + query: str | list[str], + limit: int = 5, + filters: dict | None = None, + **kwargs, ) -> list[MemoryNode]: """Search for similar memory nodes based on query text.""" filters = filters or {} @@ -195,11 +197,11 @@ class MemoryHandler: return list(seen_ids.values()) async def list( - self, - filters: dict | None = None, - limit: int | None = None, - sort_key: str | None = None, - reverse: bool = True, + self, + filters: dict | None = None, + limit: int | None = None, + sort_key: str | None = None, + reverse: bool = True, ) -> list[MemoryNode]: """List memory nodes with optional filtering and sorting.""" filters = filters or {} diff --git a/reme/tool/memory/profile_handler.py b/reme/tool/memory/profile_handler.py index 4a950171..0f51d014 100644 --- a/reme/tool/memory/profile_handler.py +++ b/reme/tool/memory/profile_handler.py @@ -1,4 +1,5 @@ """Profile Handler for managing user profiles in local memory""" + from pathlib import Path from loguru import logger @@ -36,7 +37,9 @@ class ProfileHandler: removed_count = len(sorted_nodes) - self.max_capacity nodes = sorted_nodes[removed_count:] logger.info( - f"Capacity limit reached: removed {removed_count} oldest profiles (kept {len(nodes)}/{self.max_capacity})") + f"Capacity limit reached: removed {removed_count} oldest profiles " + f"(kept {len(nodes)}/{self.max_capacity})", + ) nodes_data = [node.model_dump(exclude_none=True) for node in nodes] self.cache_handler.save(self.cache_key, nodes_data) @@ -191,9 +194,6 @@ class ProfileHandler: def read_all(self, add_profile_id: bool = False, add_history_id: bool = False) -> str: """Read all profiles and return formatted string""" nodes = self.get_all() - formatted_profiles = [ - self._format_node(node, add_profile_id, add_history_id) - for node in nodes - ] + formatted_profiles = [self._format_node(node, add_profile_id, add_history_id) for node in nodes] logger.info(f"Read {len(formatted_profiles)} profiles from {self.cache_key}") return "\n".join(formatted_profiles).strip() diff --git a/reme/tool/memory/read_all_profiles.py b/reme/tool/memory/read_all_profiles.py index 3d492ef7..d18be220 100644 --- a/reme/tool/memory/read_all_profiles.py +++ b/reme/tool/memory/read_all_profiles.py @@ -1,5 +1,4 @@ """Read user profile tool""" -from pathlib import Path from loguru import logger @@ -11,10 +10,9 @@ from ...core.schema import ToolCall class ReadAllProfiles(BaseMemoryTool): """Tool to read all user profiles""" - def __init__(self, profile_path: str, **kwargs): + def __init__(self, **kwargs): kwargs["enable_multiple"] = False super().__init__(**kwargs) - self.profile_path: str = profile_path def _build_tool_call(self) -> ToolCall: """Build and return the tool call schema""" @@ -30,16 +28,12 @@ class ReadAllProfiles(BaseMemoryTool): ) async def execute(self): - profile_handler = ProfileHandler( - profile_path=Path(self.profile_path) / self.vector_store.collection_name, - memory_target=self.memory_target, - ) - + profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=self.memory_target) profiles_str = profile_handler.read_all(add_profile_id=True) if not profiles_str: output = "No profiles found." logger.info(output) return output - logger.info(f"Successfully read profiles") + logger.info("Successfully read profiles") return profiles_str diff --git a/reme/tool/memory/retrieve_memory.py b/reme/tool/memory/retrieve_memory.py index 48fa47dc..db542643 100644 --- a/reme/tool/memory/retrieve_memory.py +++ b/reme/tool/memory/retrieve_memory.py @@ -31,8 +31,8 @@ class RetrieveMemory(BaseMemoryTool): properties["time_filter"] = { "type": "string", "description": "Optional time filter to narrow down search results by date. " - "Format: single date '20200101' for exact date match, " - "or date range '20200101,20200102' for inclusive range filtering.", + "Format: single date '20200101' for exact date match, " + "or date range '20200101,20200102' for inclusive range filtering.", } if self.enable_memory_target: @@ -99,11 +99,13 @@ class RetrieveMemory(BaseMemoryTool): else: filters = {"time_int": [int(time_filter), int(time_filter)]} - queries_by_target[target].append({ - "query": item["query"], - "limit": self.top_k, - "filters": filters, - }) + queries_by_target[target].append( + { + "query": item["query"], + "limit": self.top_k, + "filters": filters, + }, + ) # Execute batch searches for each target memory_nodes: list[MemoryNode] = [] diff --git a/reme/tool/memory/update_memory.py b/reme/tool/memory/update_memory.py index 1370e180..ba4f8d5b 100644 --- a/reme/tool/memory/update_memory.py +++ b/reme/tool/memory/update_memory.py @@ -11,10 +11,10 @@ class UpdateMemory(BaseMemoryTool): """Tool to update memories in vector store""" def __init__( - self, - enable_memory_target: bool = False, - enable_when_to_use: bool = False, - **kwargs, + self, + enable_memory_target: bool = False, + enable_when_to_use: bool = False, + **kwargs, ): super().__init__(**kwargs) self.enable_memory_target: bool = enable_memory_target @@ -116,14 +116,16 @@ class UpdateMemory(BaseMemoryTool): except Exception: logger.warning(f"Invalid message time format: {message_time}") - update_dicts.append({ - "memory_id": mem.get("memory_id", ""), - "content": memory_content, - "when_to_use": when_to_use, - "message_time": message_time, - "author": self.author, - "metadata": metadata, - }) + update_dicts.append( + { + "memory_id": mem.get("memory_id", ""), + "content": memory_content, + "when_to_use": when_to_use, + "message_time": message_time, + "author": self.author, + "metadata": metadata, + }, + ) if update_dicts: handler = MemoryHandler(target, self.service_context) diff --git a/reme/tool/memory/update_memory_v2.py b/reme/tool/memory/update_memory_v2.py index f9167e69..b855c15b 100644 --- a/reme/tool/memory/update_memory_v2.py +++ b/reme/tool/memory/update_memory_v2.py @@ -11,11 +11,11 @@ class UpdateMemoryV2(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, + 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) @@ -68,7 +68,7 @@ class UpdateMemoryV2(BaseMemoryTool): "type": "array", "description": "List of memory IDs to delete", "items": { - "type": "string" + "type": "string", }, }, "memories_to_add": { @@ -85,7 +85,7 @@ class UpdateMemoryV2(BaseMemoryTool): async def execute(self): # Get parameters memory_ids_to_delete = self.context.get("memory_ids_to_delete", []) - memory_ids_to_delete = sorted(set([mid for mid in memory_ids_to_delete if mid])) + memory_ids_to_delete = sorted({mid for mid in memory_ids_to_delete if mid}) memories_to_add = self.context.get("memories_to_add", []) if not memory_ids_to_delete and not memories_to_add: @@ -126,14 +126,16 @@ class UpdateMemoryV2(BaseMemoryTool): except Exception: logger.warning(f"Invalid message time format: {message_time}") - add_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, - }) + add_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, + }, + ) if add_dicts: handler = MemoryHandler(target, self.service_context) diff --git a/reme/tool/memory/update_profile.py b/reme/tool/memory/update_profile.py index f854f3ab..5ab7c962 100644 --- a/reme/tool/memory/update_profile.py +++ b/reme/tool/memory/update_profile.py @@ -1,5 +1,4 @@ """Update user profile tool""" -from pathlib import Path from loguru import logger @@ -11,12 +10,10 @@ from ...core.schema import ToolCall class UpdateProfile(BaseMemoryTool): """Tool to update user profile by adding or removing profile entries""" - def __init__(self, profile_path: str, **kwargs): + def __init__(self, **kwargs): kwargs["enable_multiple"] = True super().__init__(**kwargs) - self.profile_path: str = profile_path - def _build_multiple_tool_call(self) -> ToolCall: """Build and return the multiple tool call schema""" return ToolCall( @@ -29,7 +26,7 @@ class UpdateProfile(BaseMemoryTool): "type": "array", "description": "List of profile IDs to delete", "items": { - "type": "string" + "type": "string", }, }, "profiles_to_add": { @@ -61,14 +58,11 @@ class UpdateProfile(BaseMemoryTool): ) async def execute(self): - profile_handler = ProfileHandler( - profile_path=Path(self.profile_path) / self.vector_store.collection_name, - memory_target=self.memory_target, - ) + profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=self.memory_target) # Get parameters profile_ids_to_delete = self.context.get("profile_ids_to_delete", []) - profile_ids_to_delete = sorted(set([pid for pid in profile_ids_to_delete if pid])) + profile_ids_to_delete = sorted({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: