diff --git a/.gitignore b/.gitignore index 51164732..efc9bea1 100644 --- a/.gitignore +++ b/.gitignore @@ -34,6 +34,7 @@ test_compact_storage/* test_working_memory/* *.code-workspace local_vector_store/* +reme_local_memory/* chroma_vector_store/* bench_results/* meta_memory/* diff --git a/reme/agent/memory/base_memory_agent.py b/reme/agent/memory/base_memory_agent.py index 72c90634..361a6676 100644 --- a/reme/agent/memory/base_memory_agent.py +++ b/reme/agent/memory/base_memory_agent.py @@ -29,16 +29,19 @@ class BaseMemoryAgent(BaseReact, metaclass=ABCMeta): from ...tool.memory import ReadUserProfile read_tool = ReadUserProfile(show_id=show_id) - await read_tool.call(memory_target=self.memory_target) + await read_tool.call(memory_target=self.memory_target, service_context=self.service_context) return str(read_tool.response.answer) - @staticmethod - async def read_history_node() -> MemoryNode: - """Read and return the current history node from the context.""" + async def add_history_node(self) -> MemoryNode: + """Add history node""" from ...tool.memory import AddHistory add_history_tool = AddHistory() - await add_history_tool.call() + await add_history_tool.call( + messages=self.messages, + description=self.description, + service_context=self.service_context, + ) return add_history_tool.context.history_node @property diff --git a/reme/agent/memory/default/personal_retriever.py b/reme/agent/memory/default/personal_retriever.py index 1bafa317..13197304 100644 --- a/reme/agent/memory/default/personal_retriever.py +++ b/reme/agent/memory/default/personal_retriever.py @@ -3,7 +3,7 @@ from ..base_memory_agent import BaseMemoryAgent from ....core.enumeration import Role, MemoryType from ....core.op import BaseTool -from ....core.schema import Message +from ....core.schema import Message, MemoryNode from ....core.utils import format_messages @@ -12,6 +12,10 @@ class PersonalRetriever(BaseMemoryAgent): memory_type: MemoryType = MemoryType.PERSONAL + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.retrieved_nodes: list[MemoryNode] = [] + async def build_messages(self) -> list[Message]: if self.context.get("query"): context = self.context.query @@ -52,5 +56,6 @@ class PersonalRetriever(BaseMemoryAgent): step, memory_type=self.memory_type.value, memory_target=self.memory_target, + retrieved_nodes=self.retrieved_nodes, **kwargs, ) diff --git a/reme/agent/memory/default/personal_summarizer.py b/reme/agent/memory/default/personal_summarizer.py index a317b713..ef14a1d0 100644 --- a/reme/agent/memory/default/personal_summarizer.py +++ b/reme/agent/memory/default/personal_summarizer.py @@ -5,7 +5,7 @@ from loguru import logger from ..base_memory_agent import BaseMemoryAgent from ....core.enumeration import Role, MemoryType from ....core.op import BaseTool -from ....core.schema import Message +from ....core.schema import Message, MemoryNode class PersonalSummarizer(BaseMemoryAgent): @@ -13,6 +13,10 @@ class PersonalSummarizer(BaseMemoryAgent): memory_type: MemoryType = MemoryType.PERSONAL + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.retrieved_nodes: list[MemoryNode] = [] + async def _build_phase1_messages(self) -> list[Message]: """Build messages for phase 1: retrieve and add memory.""" return [ @@ -68,6 +72,7 @@ class PersonalSummarizer(BaseMemoryAgent): memory_target=self.memory_target, history_node=self.history_node, author=self.author, + retrieved_nodes=self.retrieved_nodes, **kwargs, ) @@ -75,7 +80,7 @@ class PersonalSummarizer(BaseMemoryAgent): """Execute two-phase memory processing: retrieve/add -> update profile.""" tools = self.tools for i, tool in enumerate(tools): - logger.info(f"[{self.__class__.__name__}] tool_call[{i}]={tool.tool_call.name}") + logger.info(f"[{self.__class__.__name__}] tool_call[{i}]={tool.tool_call.simple_input_dump(as_dict=False)}") messages_phase1 = await self._build_phase1_messages() for i, message in enumerate(messages_phase1): diff --git a/reme/agent/memory/default/reme_retriever.py b/reme/agent/memory/default/reme_retriever.py index 4b4c3e22..46657344 100644 --- a/reme/agent/memory/default/reme_retriever.py +++ b/reme/agent/memory/default/reme_retriever.py @@ -52,15 +52,13 @@ class ReMeRetriever(BaseMemoryAgent): description=self.description, messages=self.messages, query=self.query, - history_node=self.history_node, author=self.author, **kwargs, ) async def execute(self): - await super().execute() - - tools: list[BaseTool] = self.response.metadata["tools"] + result = await super().execute() + tools: list[BaseTool] = result["tools"] hands_off_tool = tools[0] agents: list[BaseMemoryAgent] = hands_off_tool.response.metadata["agents"] @@ -70,7 +68,7 @@ class ReMeRetriever(BaseMemoryAgent): tools = [] for agent in agents: answer += "\n" + agent.response.answer - success = success and agent.response.metadata["success"] + success = success and agent.response.success messages += agent.response.metadata["messages"] tools += agent.response.metadata["tools"] diff --git a/reme/agent/memory/default/reme_summarizer.py b/reme/agent/memory/default/reme_summarizer.py index e802f908..9cae5cbd 100644 --- a/reme/agent/memory/default/reme_summarizer.py +++ b/reme/agent/memory/default/reme_summarizer.py @@ -14,7 +14,7 @@ class ReMeSummarizer(BaseMemoryAgent): self.meta_memories: list[dict] = meta_memories or [] async def build_messages(self) -> list[Message]: - self.context.history_node = await self.read_history_node() + self.context.history_node = await self.add_history_node() messages = [ Message( @@ -53,9 +53,8 @@ class ReMeSummarizer(BaseMemoryAgent): ) async def execute(self): - await super().execute() - - tools: list[BaseTool] = self.response.metadata["tools"] + result = await super().execute() + tools: list[BaseTool] = result["tools"] hands_off_tool = tools[0] agents: list[BaseMemoryAgent] = hands_off_tool.response.metadata["agents"] @@ -65,7 +64,7 @@ class ReMeSummarizer(BaseMemoryAgent): tools = [] for agent in agents: answer += "\n" + agent.response.answer - success = success and agent.response.metadata["success"] + success = success and agent.response.success messages += agent.response.metadata["messages"] tools += agent.response.metadata["tools"] diff --git a/reme/core/context/prompt_handler.py b/reme/core/context/prompt_handler.py index e6b6d737..e93ee4c2 100644 --- a/reme/core/context/prompt_handler.py +++ b/reme/core/context/prompt_handler.py @@ -101,7 +101,6 @@ class PromptHandler(BaseContext): prompt_file_path = Path(prompt_file_path) if not prompt_file_path.exists(): - logger.warning(f"Prompt file not found: {prompt_file_path}") return self suffix = prompt_file_path.suffix.lower() @@ -117,7 +116,6 @@ class PromptHandler(BaseContext): f"Unsupported file format: {suffix}. " f"Supported formats: .yaml, .yml, .json", ) - logger.info(f"Loaded {len(prompt_dict or {})} prompts from {prompt_file_path}") self.load_prompt_dict(prompt_dict, overwrite=overwrite) except (yaml.YAMLError, json.JSONDecodeError) as e: diff --git a/reme/core/op/base_op.py b/reme/core/op/base_op.py index ffa9f14b..9c2dff71 100644 --- a/reme/core/op/base_op.py +++ b/reme/core/op/base_op.py @@ -94,12 +94,12 @@ class BaseOp(metaclass=ABCMeta): def _handle_failure(self, e: Exception, attempt: int) -> str | None: """Log failures and handle final retry logic.""" - message = f"[{self.__class__.__name__}] {self.name} failed (attempt {attempt + 1}): {e}" + message = f"[{self.__class__.__name__}] failed (attempt {attempt + 1}): {e}" if attempt == self.max_retries - 1: logger.exception(message) if self.raise_exception: raise e - return f"{self.name} failed: {e}" + return f"[{self.__class__.__name__}] failed: {e}" else: logger.warning(message) return None @@ -178,7 +178,7 @@ class BaseOp(metaclass=ABCMeta): if k == "answer": self.response.answer = v elif k == "success": - self.response.success = v.lower() == "true" + self.response.success = v if isinstance(v, bool) else v.lower() == "true" else: self.response.metadata[k] = v else: @@ -262,7 +262,7 @@ class BaseOp(metaclass=ABCMeta): if isinstance(result, list): results.extend(result) else: - result.append(result) + results.append(result) self._pending_tasks.clear() return results diff --git a/reme/core/op/base_react.py b/reme/core/op/base_react.py index 5b24e30d..95302138 100644 --- a/reme/core/op/base_react.py +++ b/reme/core/op/base_react.py @@ -102,7 +102,7 @@ class BaseReact(BaseOp): # Create isolated kwargs for each tool call to avoid parameter conflicts tool_kwargs = {**kwargs, **tool_call.argument_dict} - self.submit_async_task(tool_copy.call, **tool_kwargs) + self.submit_async_task(tool_copy.call, service_context=self.service_context, **tool_kwargs) if self.tool_call_interval > 0: await asyncio.sleep(self.tool_call_interval) diff --git a/reme/core/vector_store/base_vector_store.py b/reme/core/vector_store/base_vector_store.py index 40a84e99..cc15101b 100644 --- a/reme/core/vector_store/base_vector_store.py +++ b/reme/core/vector_store/base_vector_store.py @@ -91,6 +91,16 @@ class BaseVectorStore(ABC): async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode]: """Fetch specific vector nodes from the collection by their IDs.""" + async def dump(self) -> list[VectorNode]: + """Dump the vector store to a list of vector nodes.""" + return await self.list() + + async def load(self, nodes: list[VectorNode]): + """Load the vector store from a list of vector nodes.""" + await self.delete_all() + if nodes: + await self.insert(nodes) + @abstractmethod async def list( self, diff --git a/reme/reme.py b/reme/reme.py index d304b351..255c6a1d 100644 --- a/reme/reme.py +++ b/reme/reme.py @@ -108,13 +108,13 @@ class ReMe: ): """Summarize messages and store them in memory for the specified user(s).""" if user_name: - for message in messages and isinstance(user_name, str): - 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 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] if not meta_memories: @@ -147,7 +147,12 @@ class ReMe: else: raise NotImplementedError - return await reme_summarizer.call(messages=messages, description=description, **kwargs) + return await reme_summarizer.call( + messages=messages, + description=description, + service_context=self.service_context, + **kwargs, + ) else: raise NotImplementedError @@ -166,14 +171,13 @@ class ReMe: ): """Retrieve relevant memories for the specified user(s) based on query or messages.""" if user_name: - if messages: - for message in messages and isinstance(user_name, str): - 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 - 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] if not meta_memories: @@ -206,7 +210,13 @@ class ReMe: else: raise NotImplementedError - return await reme_retriever.call(query=query, messages=messages, description=description, **kwargs) + return await reme_retriever.call( + query=query, + messages=messages, + description=description, + service_context=self.service_context, + **kwargs, + ) else: raise NotImplementedError diff --git a/reme/tool/memory/hands_off/hands_off.py b/reme/tool/memory/hands_off/hands_off.py index c2a53872..45822381 100644 --- a/reme/tool/memory/hands_off/hands_off.py +++ b/reme/tool/memory/hands_off/hands_off.py @@ -90,7 +90,7 @@ class HandsOff(BaseMemoryTool): for k in ["query", "messages", "description", "history_node"]: if k in self.context: task_kwargs[k] = self.context[k] - self.submit_async_task(agent.call, **task_kwargs) + self.submit_async_task(agent.call, service_context=self.service_context, **task_kwargs) await self.join_async_tasks() diff --git a/reme/tool/memory/history/add_history.py b/reme/tool/memory/history/add_history.py index 588d42dc..d9527ff8 100644 --- a/reme/tool/memory/history/add_history.py +++ b/reme/tool/memory/history/add_history.py @@ -39,10 +39,10 @@ class AddHistory(BaseMemoryTool): author=self.author, ) self.context.history_node = history_node - logger.info(f"Adding history node: {history_node.model_dump_json(indent=2, exclude={'content'})}") + logger.info(f"Adding history node: {history_node.model_dump_json(indent=2)}") vector_node = history_node.to_vector_node() - await self.vector_store.delete(vector_node.memory_id) + await self.vector_store.delete(vector_node.vector_id) await self.vector_store.insert([vector_node]) return f"Successfully added history: {history_node.memory_id}" diff --git a/reme/tool/memory/user_profile/update_user_profile.py b/reme/tool/memory/user_profile/update_user_profile.py index 93f3a5dc..da561d8e 100644 --- a/reme/tool/memory/user_profile/update_user_profile.py +++ b/reme/tool/memory/user_profile/update_user_profile.py @@ -106,4 +106,4 @@ class UpdateUserProfile(BaseMemoryTool): operations.append(f"added {len(new_nodes)} new profiles.") operations.append("Operation completed.") logger.info("\n".join(operations)) - return operations + return "\n".join(operations) diff --git a/reme/tool/memory/vector/retrieve_memory.py b/reme/tool/memory/vector/retrieve_memory.py index de20a1df..a83a1dae 100644 --- a/reme/tool/memory/vector/retrieve_memory.py +++ b/reme/tool/memory/vector/retrieve_memory.py @@ -122,7 +122,16 @@ class RetrieveMemory(BaseMemoryTool): if not new_memory_nodes: output = "No new memory_nodes found matching the query (duplicates removed)." else: - output = "\n".join([m.format_memory() for m in new_memory_nodes]) + outputs = [] + for node in new_memory_nodes: + line = "" + if "conversation_time" in node.metadata and node.metadata["conversation_time"]: + line += f"conversation_time={node.metadata['conversation_time']} " + line += node.content.strip() + " " + if node.ref_memory_id: + line += f"history_id={node.ref_memory_id} " + outputs.append(line.strip()) + output = "\n".join(outputs) logger.info(f"Retrieved {len(memory_nodes)} memory_nodes, {len(new_memory_nodes)} new after deduplication") return output diff --git a/tests/test_reme.py b/tests/test_reme.py index f5929122..a4f383cf 100644 --- a/tests/test_reme.py +++ b/tests/test_reme.py @@ -2,12 +2,10 @@ import asyncio +from reme import ReMe from reme.core.schema import VectorNode, MemoryNode -from reme.reme import ReMe -reme = ReMe( - vector_store={"collection_name": "reme"}, -) +reme = ReMe(vector_store={"collection_name": "reme"}) async def test_reme(): @@ -71,7 +69,7 @@ async def test_reme(): nodes: list[VectorNode] = await reme.default_vector_store.list() for i, node in enumerate(nodes, 1): memory_node = MemoryNode.from_vector_node(node) - print(f"{i} {memory_node.memory_type} {memory_node.memory_target} {memory_node.format_memory()}") + print(f"{i} {memory_node.model_dump_json()}") print("\n" + "=" * 60) print("步骤3: 测试记忆检索 - 验证个人信息")