From e0d0e3e568e6d2163c068ad05af2cf4536c42ad2 Mon Sep 17 00:00:00 2001 From: Zhouwk <57825291+nitwtog@users.noreply.github.com> Date: Tue, 28 Apr 2026 15:11:45 +0800 Subject: [PATCH] =?UTF-8?q?=E6=8F=90=E4=BE=9B=E6=94=AF=E6=8C=81=E5=90=91?= =?UTF-8?q?=E9=87=8F=E6=95=B0=E6=8D=AE=E5=BA=93=E7=9A=84profile=E5=8A=9F?= =?UTF-8?q?=E8=83=BD=20(#221)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(reme): 添加配置选项以启用或禁用个人资料功能 - 在 ReMe 初始化方法中添加 enable_profile 参数,默认值为 True - 根据 enable_profile 设置决定是否创建 profile 目录和设置 profile_dir - 在 PersonalSummarizer 中根据 enable_profile 条件性地添加个人资料相关工具 - 在 PersonalRetriever 中根据 enable_profile 条件性地添加 ReadAllProfiles 工具 - 修改 profile_path 属性以在禁用个人资料时返回 None - 修改 get_profile_handler 方法以在禁用个人资料时返回 None - 为 enable_profile 参数添加文档说明其用于云向量存储场景 * refactor(benchmark): 重构LongMemEval基准测试中的ReMe实例管理 - 移除未使用的shutil导入 - 将固定的ReMe实例改为每个问题创建独立实例以实现隔离 - 更新LLM配置名称从qwen3-max-think到qwen-max-t - 修改模型调用逻辑使用正确的model_name参数 - 添加qwen-flash和GPT-4o-mini等新模型配置 - 统一使用"User"作为用户名,通过集合名实现隔离 - 调整并发处理数从4降至1,批处理大小从10增至30 - 每个问题类型采样数从2增至4 - 添加异步上下文管理确保资源正确释放 * reformat 2 files * refactor(benchmark): 重构长记忆评估中的模型配置 - 将原有的 eval_model_name 替换为专门的 retrieve_model_name 用于检索操作 - 添加对 qwen-max 模型配置的支持 - 更新参数解析器以支持新的检索模型参数 - 修改最大并发数默认值从 1 提升到 4 - 调整样本数量默认值从 4 减少到 1 - 统一模型参数命名规范,区分摘要、检索和评估模型 - 优化内存处理器初始化逻辑,支持独立的检索模型配置 * fix(benchmark): 移除数据路径默认值并设为必填参数 - 将LongMemEval评估脚本中的data_path参数改为必需参数 - 将HaluMem评估脚本中的data_path参数改为必需参数 - 删除了硬编码的默认文件路径配置 - 强制用户显式指定数据集文件路径以避免路径错误 * Update __init__.py * Update __init__.py * fix(benchmark): 修复ReMe评估中的模型配置和空值处理问题 - 移除了retrieve_memory调用中不需要的llm_config_name参数 - 修复了长字符串打印的换行格式问题 - 添加了eval_result为空时的初始化处理 - 在accuracy评估中加入了eval_model_name参数传递 * style(benchmark): 格式化模型名称打印输出 - 移除了多行字符串中的换行符和多余空格 - 将模型名称信息合并为单行连续显示 - 保持了原有的打印格式和信息完整性 * docs(readme): 更新文档添加实验结果表格 - 在英文版 README 中添加 🧪 Experiments 章节 - 添加 LoCoMo 和 HaluMem 两个基准测试的结果表格 - 在中文版 README_ZH 中添加 🧪 实验 章节 - 添加 LoCoMo 和 HaluMem 测试集的实验配置说明 - 添加完整的实验数据对比表格和评估协议说明 * docs(readme): 更新文档中的内存系统链接 - 为基于文件的记忆系统添加锚点链接 - 为基于向量库的记忆系统添加锚点链接 - 修复英文文档中的链接格式 - 修复中文文档中的链接格式和空行问题 * docs(readme): update experimental results section in documentation - Remove outdated experimental data placeholder "Coming soon..." - Add complete evaluation results for LoCoMo and HaluMem benchmarks - Include detailed performance metrics tables for all memory methods - Update experimental settings description with ReMe backbone details - Align evaluation protocol information with LLM-as-a-Judge approach - Maintain consistent formatting between English and Chinese documentation * docs(benchmark): add quick start guides for halumem and longmemeval experiments - Created HaluMem experiment quick start guide with ReMe integration setup - Added detailed steps for installing ReMe environment using conda - Included repository cloning instructions for HaluMem benchmark - Provided complete command examples for running HaluMem experiments - Created LongMeMEval quick start guide with data download procedures - Added wget commands for downloading cleaned dataset files - Included evaluation script instructions for computing experiment statistics - Documented parameter configurations for different model types and batch sizes * docs(longmemeval): update quickstart guide documentation - Changed project name from Halumem to Longmemeval in title - Updated description to reference Longmemeval experiments instead of Halumem - Maintained existing ReMe integration instructions unchanged * chore(logger): add test comment to logger configuration - Added test comment in logger utility function - Removed duplicate log handling by keeping the remove() call * chore(logger): add test comment to logger configuration - Added test comment in logger utility function - Removed duplicate log handling by keeping the remove() call * feat(core): add file logging capability to application - Added log_to_file parameter to Application class constructor - Integrated log_to_file option in logger initialization - Updated ServiceContext to support file logging configuration - Modified init_logger function to conditionally enable file logging - Added log_to_file field to ServiceConfig schema - Updated ReMe class to include file logging option - Wrapped file logging setup in conditional check to prevent unnecessary operations * docs(benchmark): update HaluMem quickstart guide with dataset download instructions - Replace repository cloning with direct dataset download using curl - Add commands to download HaluMem-Medium.jsonl and HaluMem-Long.jsonl files - Include both official Hugging Face and mirror download sources - Update data path reference from nested directory to local data folder - Add dataset page link and mirror usage instructions for mainland China access * feat(memory): add profile retrieval tool and refactor profile management - Introduce RetrieveProfile tool for fetching specific user profiles - Refactor ProfileHandler to support both filesystem and vector backends - Add async methods to ProfileHandler with synchronous fallbacks - Update PersonalRetriever to support two-stage profile and memory retrieval - Enhance PersonalSummarizer with improved tool partitioning logic - Add profile_backend, profile_store_name, and profile_max_capacity configuration options - Replace direct ProfileHandler imports with get_profile_handler method - Implement profile search functionality with dedicated prompts and workflows - Add FileProfileBackend and VectorProfileBackend implementations - Update base memory tool with new profile configuration parameters * feat(profile): add custom profile collection name support - Add profile_collection_name parameter to Application constructor - Allow custom database collection name for vector profiles instead of default suffix - Update profile vector store configuration logic to use custom collection name - Modify _ensure_profile_vector_store_config to handle custom collection names - Update docstring with detailed parameter descriptions for profile configuration options * test(history): add single history id acceptance test for multiple mode - Add test case to verify multiple-mode history lookup accepts a single history_id string - Create FakeVectorStore stub with minimal implementation for ReadHistory tests - Return requested history node from vector store mock - Initialize ReadHistory tool with multiple mode enabled - Add pylint disable comment for protected access to vector store property * refactor(memory): update profile handler and vector tools with improved formatting and error handling - Add module docstring to profiles/__init__.py - Add pylint disable comments for no-name-in-module and missing-function-docstring - Format long error message in ProfileHandler.sync_run method for better readability - Reformat parameters in ProfileHandler.aadd method to separate lines - Update model_copy call in reme.py to span multiple lines for better readability - Format aadd_batch call in update_profile.py to span multiple lines --- .../personal/personal_retriever.py | 130 ++++--- .../personal/personal_retriever.yaml | 33 +- .../personal/personal_summarizer.py | 85 +++-- reme/memory/vector_tools/__init__.py | 4 + reme/memory/vector_tools/base_memory_tool.py | 22 +- reme/memory/vector_tools/profiles/__init__.py | 1 + .../add_draft_and_read_all_profiles.py | 8 +- .../vector_tools/profiles/add_profile.py | 7 +- .../vector_tools/profiles/delete_profile.py | 7 +- .../vector_tools/profiles/profile_handler.py | 318 +++++++++--------- .../profiles/read_all_profiles.py | 7 +- .../vector_tools/profiles/update_profile.py | 18 +- .../profiles/update_profiles_v1.py | 11 +- reme/reme.py | 164 +++++++-- tests/test_reme_memory_error_handling.py | 38 ++- 15 files changed, 544 insertions(+), 309 deletions(-) diff --git a/reme/memory/vector_based/personal/personal_retriever.py b/reme/memory/vector_based/personal/personal_retriever.py index b8803717..8b4965f1 100644 --- a/reme/memory/vector_based/personal/personal_retriever.py +++ b/reme/memory/vector_based/personal/personal_retriever.py @@ -1,39 +1,19 @@ """Personal memory retriever agent for retrieving personal memories through vector search.""" +from loguru import logger + from ..base_memory_agent import BaseMemoryAgent -from ....core.enumeration import Role, MemoryType +from ....core.enumeration import MemoryType, Role from ....core.op import BaseTool from ....core.schema import Message from ....core.utils import format_messages +_PROFILE_TOOL_NAMES: tuple[str, ...] = ("retrieve_profile", "read_all_profiles") +_EMPTY_PROFILE_RESULTS: tuple[str, ...] = ("", "No profiles found.", "No new profiles found.") + class PersonalRetriever(BaseMemoryAgent): - """Retrieve personal memories through vector search and history reading. - - clear && python benchmark/halumem/eval_reme.py \ - --data_path /Users/yuli/workspace/HaluMem/data/HaluMem-Medium.jsonl \ - --reme_model_name qwen3.5-plus \ - --batch_size 10000 \ - --algo_version default - - 📊 Question Answering (with LLM answer): - Correct (all): 0.8537 - Hallucination (all): 0.1159 - Omission (all): 0.0305 - Correct (valid): 0.8537 - Hallucination (valid): 0.1159 - Omission (valid): 0.0305 - Valid/Total: 164/164 - - 📊 Question Answering (with original memories): - Correct (all): 0.9085 - Hallucination (all): 0.0671 - Omission (all): 0.0244 - Correct (valid): 0.9085 - Hallucination (valid): 0.0671 - Omission (valid): 0.0244 - Valid/Total: 164/164 - """ + """Retrieve personal memories through vector search and history reading.""" memory_type: MemoryType = MemoryType.PERSONAL @@ -41,36 +21,73 @@ class PersonalRetriever(BaseMemoryAgent): super().__init__(**kwargs) self.return_memory_nodes: bool = return_memory_nodes - async def build_messages(self) -> list[Message]: + def _get_context(self) -> str: if self.context.get("query"): - context = self.context.query - elif self.context.get("messages"): - context = self.description + "\n" + format_messages(self.context.messages) - else: - raise ValueError("input must have either `query` or `messages`") - - 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, - ) - else: - all_profiles = "" + return self.context.query.strip() + if self.context.get("messages"): + return (self.description + "\n" + format_messages(self.context.messages)).strip() + raise ValueError("input must have either `query` or `messages`") + async def _build_s1_messages(self, context: str) -> list[Message]: return [ Message( role=Role.USER, content=self.prompt_format( - prompt_name="user_message", + prompt_name="user_message_s1", memory_type=self.memory_type.value, memory_target=self.memory_target, - user_profile=all_profiles, - context=context.strip(), + context=context, ), ), ] + async def _build_s2_messages(self, context: str, profiles: str) -> list[Message]: + return [ + Message( + role=Role.USER, + content=self.prompt_format( + prompt_name="user_message_s2", + memory_type=self.memory_type.value, + memory_target=self.memory_target, + profiles=profiles, + context=context, + ), + ), + ] + + def _partition_tools(self) -> tuple[list[BaseTool], list[BaseTool]]: + profile_tools: list[BaseTool] = [] + memory_tools: list[BaseTool] = [] + for i, tool in enumerate(self.tools): + name = tool.tool_call.name + if name in _PROFILE_TOOL_NAMES: + profile_tools.append(tool) + else: + memory_tools.append(tool) + logger.info(f"[{self.__class__.__name__}] tool_call[{i}]={tool.tool_call.simple_input_dump(as_dict=False)}") + return profile_tools, memory_tools + + @staticmethod + def _extract_profile_context(tools: list[BaseTool]) -> str: + outputs = [] + for tool in tools: + response = getattr(tool, "response", None) + answer = getattr(response, "answer", "") + if answer and answer not in _EMPTY_PROFILE_RESULTS: + outputs.append(answer) + return "\n".join(outputs) + + async def _run_stage( + self, + stage: str, + messages: list[Message], + tools: list[BaseTool], + ) -> tuple[list[BaseTool], list[Message], bool]: + for message in messages: + role = message.name or message.role + logger.info(f"[{self.__class__.__name__} {stage}] role={role} {message.simple_dump(as_dict=False)}") + return await self.react(messages, tools, stage=stage) + async def _acting_step( self, assistant_message: Message, @@ -91,7 +108,28 @@ class PersonalRetriever(BaseMemoryAgent): ) async def execute(self): - result = await super().execute() + context = self._get_context() + profile_tools, memory_tools = self._partition_tools() + + tools_s1: list[BaseTool] = [] + messages_s1: list[Message] = [] + success_s1 = True + profiles = "" + if profile_tools: + messages_s1 = await self._build_s1_messages(context) + tools_s1, messages_s1, success_s1 = await self._run_stage("s1-profile", messages_s1, profile_tools) + profiles = self._extract_profile_context(tools_s1) + + messages_s2 = await self._build_s2_messages(context, profiles) + tools_s2, messages_s2, success_s2 = await self._run_stage("s2-memory", messages_s2, memory_tools) + + answer = messages_s2[-1].content if success_s2 and messages_s2 else "" + result = { + "answer": answer, + "success": success_s1 and success_s2, + "messages": messages_s1 + messages_s2, + "tools": tools_s1 + tools_s2, + } if self.return_memory_nodes: result["answer"] = "\n".join( [ diff --git a/reme/memory/vector_based/personal/personal_retriever.yaml b/reme/memory/vector_based/personal/personal_retriever.yaml index 99df1374..c52ae16f 100644 --- a/reme/memory/vector_based/personal/personal_retriever.yaml +++ b/reme/memory/vector_based/personal/personal_retriever.yaml @@ -1,8 +1,26 @@ -user_message: | +user_message_s1: | + You are a Profile Retrieval Agent specialized in finding profile information about {memory_target}. + + ## User Question + {context} + + ## Task + Use the available profile tool to search for profile content that is relevant to the user question. + + ## Instructions + - If `retrieve_profile` is available, use it to search with focused profile queries derived from the question + - If `read_all_profiles` is available, use it to inspect the full profile list and identify relevant rows + - Focus on profile attributes such as identity, location, work, education, preferences, relationships, and other long-term facts + - Only retrieve information that is directly relevant to the user question + - If no relevant profile information exists, say so clearly + + Output a concise summary of the relevant profile information you found. + +user_message_s2: | You are a Memory Retrieval Agent specialized in retrieving {memory_type} memories about {memory_target}. - ## User Profile - {user_profile} + ## Profile Search Results + {profiles} ## User Question {context} @@ -14,11 +32,12 @@ user_message: | **Tool**: `retrieve_memory` (without time constraints) **Objective**: Cast a wide net to find potentially relevant memories **Approach**: + - Use the profile search results above as supporting context when forming retrieval queries - Execute 3-5 diverse search queries using different formulations: * Original question verbatim * Rephrased variations (different wording, synonyms) * Entity-focused queries (extract and search specific names, places, events) - * Keyword-based searches (core concepts, topics) + * Keyword-based searches (core concepts and profile facts) * Related context queries (broader themes) ### Phase 2(Optional): Temporal Search @@ -31,7 +50,7 @@ user_message: | - After date: `20200101,99999999` (from 20200101 onwards) **Approach**: - Identify temporal constraints from the user question - - Refine Phase 1 queries with 3-5 diverse appropriate different time filters + - Refine Phase 1 queries with 3-5 diverse appropriate time filters ### Phase 3: Deep Dive into History **Tool**: `read_history` @@ -48,11 +67,11 @@ user_message: | - Use this to understand the full conversation surrounding a memory ## Response Guidelines - - Base your answer EXCLUSIVELY on user profile, retrieved memories, and history data + - Base your answer EXCLUSIVELY on the profile search results, 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 a summary of all retrieved memories, user profile, and history data. + Output a summary of all retrieved memories and relevant history data. diff --git a/reme/memory/vector_based/personal/personal_summarizer.py b/reme/memory/vector_based/personal/personal_summarizer.py index 8e570718..89b1316f 100644 --- a/reme/memory/vector_based/personal/personal_summarizer.py +++ b/reme/memory/vector_based/personal/personal_summarizer.py @@ -3,13 +3,17 @@ from loguru import logger from ..base_memory_agent import BaseMemoryAgent -from ....core.enumeration import Role, MemoryType +from ....core.enumeration import MemoryType, Role from ....core.op import BaseTool from ....core.schema import Message +# Optional profile tools used to pre-load profile context; consumed by the +# summarizer itself and never exposed to the stage-two ReAct loop. +_PROFILE_CONTEXT_TOOLS: tuple[str, ...] = ("retrieve_profile", "read_all_profiles") + class PersonalSummarizer(BaseMemoryAgent): - """Two-phase personal memory processor: retrieve/add memories then update profile.""" + """Two-phase personal memory processor: add memories, then update profiles.""" memory_type: MemoryType = MemoryType.PERSONAL @@ -62,62 +66,71 @@ class PersonalSummarizer(BaseMemoryAgent): **kwargs, ) - async def execute(self): - memory_tools = [] - profile_tools = [] - read_all_profiles_tool: BaseTool | None = None + def _partition_tools(self) -> tuple[list[BaseTool], list[BaseTool], BaseTool | None]: + """Split attached tools into memory tools, profile tools, and a profile context tool.""" + memory_tools: list[BaseTool] = [] + profile_tools: list[BaseTool] = [] + profile_context_tool: BaseTool | None = None for i, tool in enumerate(self.tools): - tool_name = tool.tool_call.name - if tool_name == "read_all_profiles": - read_all_profiles_tool = tool - elif "_memory" in tool_name: + name = tool.tool_call.name + if name in _PROFILE_CONTEXT_TOOLS: + profile_context_tool = tool + elif "_memory" in name: memory_tools.append(tool) - elif "_profile" in tool_name: + elif "_profile" in name: profile_tools.append(tool) else: - raise ValueError(f"[{self.__class__.__name__}] unknown tool_name={tool_name}") + raise ValueError(f"[{self.__class__.__name__}] unknown tool_name={name}") logger.info(f"[{self.__class__.__name__}] tool_call[{i}]={tool.tool_call.simple_input_dump(as_dict=False)}") + return memory_tools, profile_tools, profile_context_tool - stage = "s1-memory" - messages_s1 = await self._build_s1_messages() - for i, message in enumerate(messages_s1): + async def _preload_user_profile(self, tool: BaseTool | None) -> str: + """Invoke the profile context tool to obtain inline profile text.""" + if tool is None: + return "" + call_kwargs: dict = { + "memory_target": self.memory_target, + "service_context": self.service_context, + "retrieved_nodes": self.retrieved_nodes, + } + if tool.tool_call.name == "retrieve_profile": + call_kwargs["query"] = self.context.history_node.content + return await tool.call(**call_kwargs) + + async def _run_stage( + self, + stage: str, + messages: list[Message], + tools: list[BaseTool], + ) -> tuple[list[BaseTool], list[Message], bool]: + for message in messages: role = message.name or message.role logger.info(f"[{self.__class__.__name__} {stage}] role={role} {message.simple_dump(as_dict=False)}") - tools_s1, messages_s1, success_s1 = await self.react(messages_s1, memory_tools, stage=stage) + return await self.react(messages, tools, stage=stage) - if read_all_profiles_tool is not None: - profiles = await read_all_profiles_tool.call( - memory_target=self.memory_target, - service_context=self.service_context, - ) - else: - profiles = "" + async def execute(self): + memory_tools, profile_tools, profile_context_tool = self._partition_tools() + + messages_s1 = await self._build_s1_messages() + tools_s1, messages_s1, success_s1 = await self._run_stage("s1-memory", messages_s1, memory_tools) if profile_tools: - stage = "s2-profile" + profiles = await self._preload_user_profile(profile_context_tool) messages_s2 = await self._build_s2_messages(profiles) - for i, message in enumerate(messages_s2): - role = message.name or message.role - logger.info(f"[{self.__class__.__name__} {stage}] role={role} {message.simple_dump(as_dict=False)}") - tools_s2, messages_s2, success_s2 = await self.react(messages_s2, profile_tools, stage=stage) + tools_s2, messages_s2, success_s2 = await self._run_stage("s2-profile", messages_s2, profile_tools) else: tools_s2, messages_s2, success_s2 = [], [], True 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 - memory_nodes = [] - for tool in tools: - if tool.memory_nodes: - memory_nodes.extend(tool.memory_nodes) + memory_nodes = [node for tool in tools for node in (tool.memory_nodes or [])] return { "answer": answer, - "success": success, - "messages": messages, + "success": success_s1 and success_s2, + "messages": messages_s1 + messages_s2, "tools": tools, "memory_nodes": memory_nodes, } diff --git a/reme/memory/vector_tools/__init__.py b/reme/memory/vector_tools/__init__.py index fd5e721f..9fb04923 100644 --- a/reme/memory/vector_tools/__init__.py +++ b/reme/memory/vector_tools/__init__.py @@ -1,5 +1,7 @@ """memory tools""" +# pylint: disable=no-name-in-module + from .base_memory_tool import BaseMemoryTool # chunk tools @@ -15,6 +17,7 @@ from .profiles.add_draft_and_read_all_profiles import AddDraftAndReadAllProfiles from .profiles.add_profile import AddProfile from .profiles.delete_profile import DeleteProfile from .profiles.read_all_profiles import ReadAllProfiles +from .profiles.retrieve_profile import RetrieveProfile from .profiles.update_profile import UpdateProfile from .profiles.update_profiles_v1 import UpdateProfilesV1 @@ -43,6 +46,7 @@ __all__ = [ "AddProfile", "DeleteProfile", "ReadAllProfiles", + "RetrieveProfile", "UpdateProfile", "UpdateProfilesV1", # record tools diff --git a/reme/memory/vector_tools/base_memory_tool.py b/reme/memory/vector_tools/base_memory_tool.py index 3431b164..f0781bd3 100644 --- a/reme/memory/vector_tools/base_memory_tool.py +++ b/reme/memory/vector_tools/base_memory_tool.py @@ -3,6 +3,7 @@ from abc import ABCMeta from pathlib import Path +from .profiles.profile_handler import ProfileHandler from ...core.enumeration import MemoryType from ...core.op import BaseTool from ...core.schema import ToolCall, MemoryNode, ToolAttr @@ -16,12 +17,18 @@ class BaseMemoryTool(BaseTool, metaclass=ABCMeta): enable_multiple: bool = True, enable_thinking_params: bool = False, profile_dir: str = "", + profile_backend: str = "filesystem", + profile_store_name: str = "profile", + profile_max_capacity: int = 50, **kwargs, ): super().__init__(**kwargs) self.enable_multiple: bool = enable_multiple self.enable_thinking_params: bool = enable_thinking_params self.profile_dir: str = profile_dir + self.profile_backend: str = profile_backend + self.profile_store_name: str = profile_store_name + self.profile_max_capacity: int = profile_max_capacity def _build_tool_call(self) -> ToolCall: """Build and return the tool call schema""" @@ -103,6 +110,19 @@ class BaseMemoryTool(BaseTool, metaclass=ABCMeta): return self.context.service_context.memory_target_type_mapping @property - def profile_path(self) -> Path: + def profile_path(self) -> Path | None: """Get the path to the profile directory for the current collection.""" + if not self.profile_dir: + return None return Path(self.profile_dir) / self.vector_store.collection_name + + def get_profile_handler(self, memory_target: str) -> ProfileHandler: + """Build a profile handler for the current backend configuration.""" + return ProfileHandler( + memory_target=memory_target, + profile_path=self.profile_path, + service_context=self.service_context, + profile_backend=self.profile_backend, + profile_store_name=self.profile_store_name, + max_capacity=self.profile_max_capacity, + ) diff --git a/reme/memory/vector_tools/profiles/__init__.py b/reme/memory/vector_tools/profiles/__init__.py index e69de29b..150455db 100644 --- a/reme/memory/vector_tools/profiles/__init__.py +++ b/reme/memory/vector_tools/profiles/__init__.py @@ -0,0 +1 @@ +"""Profile memory tools.""" diff --git a/reme/memory/vector_tools/profiles/add_draft_and_read_all_profiles.py b/reme/memory/vector_tools/profiles/add_draft_and_read_all_profiles.py index 716c4e5e..c8e4889a 100644 --- a/reme/memory/vector_tools/profiles/add_draft_and_read_all_profiles.py +++ b/reme/memory/vector_tools/profiles/add_draft_and_read_all_profiles.py @@ -1,8 +1,7 @@ -"""Add draft profile and read all profiles from local storage""" +"""Add draft profile and read all profiles from the configured backend.""" from loguru import logger -from .profile_handler import ProfileHandler from ..base_memory_tool import BaseMemoryTool from ....core.schema import ToolCall @@ -92,9 +91,8 @@ class AddDraftAndReadAllProfiles(BaseMemoryTool): continue targets_processed.add(target) - profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=target) - - profiles_str = profile_handler.read_all(add_profile_id=True) + profile_handler = self.get_profile_handler(target) + profiles_str = await profile_handler.aread_all(add_profile_id=True) if profiles_str: all_profiles.append(f"## Profiles for {target}:\n{profiles_str}") diff --git a/reme/memory/vector_tools/profiles/add_profile.py b/reme/memory/vector_tools/profiles/add_profile.py index 880bcf45..5654ebb7 100644 --- a/reme/memory/vector_tools/profiles/add_profile.py +++ b/reme/memory/vector_tools/profiles/add_profile.py @@ -1,8 +1,7 @@ -"""Add user profile tool""" +"""Add user profile tool.""" from loguru import logger -from .profile_handler import ProfileHandler from ..base_memory_tool import BaseMemoryTool from ....core.schema import ToolCall @@ -40,7 +39,7 @@ class AddProfile(BaseMemoryTool): ) async def execute(self): - profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=self.memory_target) + profile_handler = self.get_profile_handler(self.memory_target) # Get parameters message_time = self.context.get("message_time", "") @@ -58,7 +57,7 @@ class AddProfile(BaseMemoryTool): } # Add profile using ProfileHandler - new_nodes = profile_handler.add_batch(profiles=[profile], ref_memory_id=self.history_id) + new_nodes = await profile_handler.aadd_batch(profiles=[profile], ref_memory_id=self.history_id) self.memory_nodes.extend(new_nodes) if new_nodes: diff --git a/reme/memory/vector_tools/profiles/delete_profile.py b/reme/memory/vector_tools/profiles/delete_profile.py index ae5bbc05..1c5f845b 100644 --- a/reme/memory/vector_tools/profiles/delete_profile.py +++ b/reme/memory/vector_tools/profiles/delete_profile.py @@ -1,8 +1,7 @@ -"""Delete user profile tool""" +"""Delete user profile tool.""" from loguru import logger -from .profile_handler import ProfileHandler from ..base_memory_tool import BaseMemoryTool from ....core.schema import ToolCall @@ -32,7 +31,7 @@ class DeleteProfile(BaseMemoryTool): ) async def execute(self): - profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=self.memory_target) + profile_handler = self.get_profile_handler(self.memory_target) # Get profile_id parameter profile_id = self.context.get("profile_id", "") @@ -41,7 +40,7 @@ class DeleteProfile(BaseMemoryTool): return "No profile_id provided, operation cancelled." # Delete profile using ProfileHandler - success = profile_handler.delete(profile_id) + success = await profile_handler.adelete(profile_id) if success: output = f"Successfully deleted profile with ID: {profile_id}" diff --git a/reme/memory/vector_tools/profiles/profile_handler.py b/reme/memory/vector_tools/profiles/profile_handler.py index 3d0df400..0df3707a 100644 --- a/reme/memory/vector_tools/profiles/profile_handler.py +++ b/reme/memory/vector_tools/profiles/profile_handler.py @@ -1,195 +1,123 @@ -"""Profile Handler for managing user profiles in local memory""" +"""Profile handler facade for filesystem and vector backends.""" +# pylint: disable=missing-function-docstring + +import asyncio from pathlib import Path from loguru import logger -from ....core.enumeration import MemoryType +from .file_profile_backend import FileProfileBackend +from .profile_backend import BaseProfileBackend +from .vector_profile_backend import VectorProfileBackend +from ....core import ServiceContext from ....core.schema import MemoryNode -from ....core.utils import CacheHandler, deduplicate_memories class ProfileHandler: - """User profile CRUD handler""" + """User profile facade with pluggable storage backends.""" - def __init__(self, profile_path: str | Path, memory_target: str, max_capacity: int = 50): - """init""" - self.memory_target: str = memory_target - self.cache_key: str = self.memory_target.replace(" ", "_").lower() - self.cache_handler: CacheHandler = CacheHandler(profile_path) - self.max_capacity: int = max_capacity - - def _load_nodes(self) -> list[MemoryNode]: - """Load profile nodes""" - cached_data = self.cache_handler.load(self.cache_key, auto_clean=False) - if not cached_data: - return [] - return [MemoryNode(**data) for data in cached_data] - - def _save_nodes(self, nodes: list[MemoryNode], apply_limits: bool = True): - """Save nodes with optional deduplication and capacity enforcement""" - if apply_limits: - nodes = deduplicate_memories(nodes) - - # Enforce capacity limit by removing the oldest profiles - if len(nodes) > self.max_capacity: - sorted_nodes = sorted(nodes, key=lambda n: n.message_time) - removed_count = len(sorted_nodes) - self.max_capacity - nodes = sorted_nodes[removed_count:] - logger.info( - 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) - logger.info(f"Saved {len(nodes)} profiles to {self.cache_key}") - - def delete(self, profile_id: str | list[str]) -> bool | int: - """Delete profile by ID(s), returns True/False for single ID or count for batch delete""" - nodes = self._load_nodes() - original_count = len(nodes) - - # Batch delete mode - if isinstance(profile_id, list): - profile_ids_set = set(profile_id) - nodes = [n for n in nodes if n.memory_id not in profile_ids_set] - deleted_count = original_count - len(nodes) - - if deleted_count == 0: - logger.warning(f"No profiles found to delete from {len(profile_id)} IDs") - return 0 - - self._save_nodes(nodes, apply_limits=False) - logger.info(f"Batch deleted {deleted_count} profiles") - return deleted_count - - # Single delete mode - nodes = [n for n in nodes if n.memory_id != profile_id] - - if len(nodes) == original_count: - logger.warning(f"Profile {profile_id} not found") - return False - - self._save_nodes(nodes, apply_limits=False) - logger.info(f"Deleted profile {profile_id}") - return True - - def delete_all(self) -> int: - """Delete all profiles, returns count deleted""" - nodes = self._load_nodes() - count = len(nodes) - self._save_nodes([], apply_limits=False) - logger.info(f"Deleted all {count} profiles") - return count - - def add(self, message_time: str, profile_key: str, profile_value: str, ref_memory_id: str = "") -> MemoryNode: - """Add new profile, returns created MemoryNode""" - nodes = self._load_nodes() - - new_node = MemoryNode( - memory_type=MemoryType.PERSONAL, - memory_target=self.memory_target, - when_to_use=profile_key, - content=profile_value, - message_time=message_time, - ref_memory_id=ref_memory_id, + def __init__( + self, + memory_target: str, + profile_path: str | Path | None = None, + service_context: ServiceContext | None = None, + profile_backend: str = "filesystem", + profile_store_name: str = "profile", + max_capacity: int = 50, + ): + self.memory_target = memory_target + self.profile_backend = profile_backend + self.profile_store_name = profile_store_name + self.max_capacity = max_capacity + self.cache_key = self.memory_target.replace(" ", "_").lower() + self.backend = self._build_backend( + profile_path=profile_path, + service_context=service_context, ) - # Remove existing nodes with the same when_to_use (profile_key) - original_count = len(nodes) - nodes = [n for n in nodes if n.when_to_use != profile_key] - if len(nodes) < original_count: - logger.info(f"Removed {original_count - len(nodes)} duplicate profile(s) with key: {profile_key}") - - nodes.append(new_node) - self._save_nodes(nodes) - logger.info(f"Added profile: {profile_key}={profile_value}") - return new_node - - def add_batch(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]: - """Add multiple profiles in batch, returns list of created MemoryNodes""" - if not profiles: - return [] - - nodes = self._load_nodes() - - new_nodes = [ - MemoryNode( - memory_type=MemoryType.PERSONAL, + def _build_backend( + self, + profile_path: str | Path | None, + service_context: ServiceContext | None, + ) -> BaseProfileBackend: + if self.profile_backend == "filesystem": + if profile_path is None: + raise ValueError("profile_path is required for filesystem profile backend") + return FileProfileBackend( + profile_path=profile_path, memory_target=self.memory_target, - when_to_use=p.get("profile_key", ""), - content=p.get("profile_value", ""), - message_time=p.get("message_time", ""), - ref_memory_id=ref_memory_id, + max_capacity=self.max_capacity, ) - for p in profiles - ] - # Remove existing nodes with the same when_to_use (profile_key) - new_keys = {n.when_to_use for n in new_nodes} - original_count = len(nodes) - nodes = [n for n in nodes if n.when_to_use not in new_keys] - if len(nodes) < original_count: - logger.info(f"Removed {original_count - len(nodes)} duplicate profile(s) with matching keys") + if self.profile_backend == "vector": + if service_context is None: + raise ValueError("service_context is required for vector profile backend") + return VectorProfileBackend( + memory_target=self.memory_target, + service_context=service_context, + vector_store_name=self.profile_store_name, + max_capacity=self.max_capacity, + ) - nodes.extend(new_nodes) - self._save_nodes(nodes) - logger.info(f"Batch added {len(new_nodes)} profiles") - return new_nodes + raise ValueError(f"Unsupported profile backend: {self.profile_backend}") - def update(self, profile_id: str, message_time: str, profile_key: str, profile_value: str) -> MemoryNode | None: - """Update profile by ID, returns updated node or None if not found""" - nodes = self._load_nodes() + @staticmethod + def _run_sync(coro): + try: + asyncio.get_running_loop() + except RuntimeError: + return asyncio.run(coro) + raise RuntimeError( + "Synchronous profile access is not available in an active event loop. Use async methods instead.", + ) - target_node = None - for node in nodes: - if node.memory_id == profile_id: - node.when_to_use = profile_key - node.content = profile_value - node.message_time = message_time - target_node = node - break + async def adelete(self, profile_id: str | list[str]) -> bool | int: + return await self.backend.delete(profile_id) - if target_node is None: - logger.warning(f"Profile {profile_id} not found") - return None + async def adelete_all(self) -> int: + return await self.backend.delete_all() - self._save_nodes(nodes, apply_limits=False) - logger.info(f"Updated profile {profile_id}: {profile_key}={profile_value}") - return target_node + async def aadd( + self, + message_time: str, + profile_key: str, + profile_value: str, + ref_memory_id: str = "", + ) -> MemoryNode: + return await self.backend.add(message_time, profile_key, profile_value, ref_memory_id) - def get_by(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None: - """Get profile by ID or key""" - if not profile_id and not profile_key: - raise ValueError("Must provide either profile_id or profile_key") + async def aadd_batch(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]: + return await self.backend.add_batch(profiles, ref_memory_id) - nodes = self._load_nodes() - for node in nodes: - if profile_id and node.memory_id == profile_id: - return node - if profile_key and node.when_to_use == profile_key: - return node - return None + async def aupdate( + self, + profile_id: str, + message_time: str, + profile_key: str, + profile_value: str, + ) -> MemoryNode | None: + return await self.backend.update(profile_id, message_time, profile_key, profile_value) - def get_by_id(self, profile_id: str) -> MemoryNode | None: - """Get profile by ID (convenience method)""" - return self.get_by(profile_id=profile_id) + async def aget_by(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None: + return await self.backend.get_by(profile_id=profile_id, profile_key=profile_key) - def get_by_key(self, profile_key: str) -> MemoryNode | None: - """Get profile by key (convenience method)""" - return self.get_by(profile_key=profile_key) + async def aget_by_id(self, profile_id: str) -> MemoryNode | None: + return await self.aget_by(profile_id=profile_id) - def get_all(self) -> list[MemoryNode]: - """Get all profiles, sorted by message_time""" - nodes = self._load_nodes() - nodes.sort(key=lambda n: n.message_time) - return nodes + async def aget_by_key(self, profile_key: str) -> MemoryNode | None: + return await self.aget_by(profile_key=profile_key) + + async def aget_all(self) -> list[MemoryNode]: + return await self.backend.get_all() + + async def asearch(self, query: str | list[str], limit: int = 5) -> list[MemoryNode]: + return await self.backend.search(query=query, limit=limit) @staticmethod def _format_node(node: MemoryNode, add_profile_id: bool = False, add_history_id: bool = False) -> str: - """Format a single node to string""" parts = [] + profile_key = str(node.metadata.get("profile_key", node.when_to_use)) if add_profile_id: parts.append(f"profile_id={node.memory_id}") @@ -197,16 +125,70 @@ class ProfileHandler: if node.message_time: parts.append(f"[{node.message_time}]") - parts.append(f"{node.when_to_use}: {node.content}") + parts.append(f"{profile_key}: {node.content}") - if add_history_id: + if add_history_id and node.ref_memory_id: parts.append(f"history_id={node.ref_memory_id}") return " ".join(parts) - 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() + async def aread_all(self, add_profile_id: bool = False, add_history_id: bool = False) -> str: + nodes = await self.aget_all() 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() + + async def aretrieve( + self, + query: str | list[str], + limit: int = 5, + add_profile_id: bool = True, + add_history_id: bool = False, + ) -> tuple[list[MemoryNode], str]: + nodes = await self.asearch(query=query, limit=limit) + formatted_profiles = [self._format_node(node, add_profile_id, add_history_id) for node in nodes] + return nodes, "\n".join(formatted_profiles).strip() + + def delete(self, profile_id: str | list[str]) -> bool | int: + if isinstance(self.backend, FileProfileBackend): + return self.backend.delete_sync(profile_id) + return self._run_sync(self.adelete(profile_id)) + + def delete_all(self) -> int: + if isinstance(self.backend, FileProfileBackend): + return self.backend.delete_all_sync() + return self._run_sync(self.adelete_all()) + + def add(self, message_time: str, profile_key: str, profile_value: str, ref_memory_id: str = "") -> MemoryNode: + if isinstance(self.backend, FileProfileBackend): + return self.backend.add_sync(message_time, profile_key, profile_value, ref_memory_id) + return self._run_sync(self.aadd(message_time, profile_key, profile_value, ref_memory_id)) + + def add_batch(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]: + if isinstance(self.backend, FileProfileBackend): + return self.backend.add_batch_sync(profiles, ref_memory_id) + return self._run_sync(self.aadd_batch(profiles, ref_memory_id)) + + def update(self, profile_id: str, message_time: str, profile_key: str, profile_value: str) -> MemoryNode | None: + if isinstance(self.backend, FileProfileBackend): + return self.backend.update_sync(profile_id, message_time, profile_key, profile_value) + return self._run_sync(self.aupdate(profile_id, message_time, profile_key, profile_value)) + + def get_by(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None: + if isinstance(self.backend, FileProfileBackend): + return self.backend.get_by_sync(profile_id=profile_id, profile_key=profile_key) + return self._run_sync(self.aget_by(profile_id=profile_id, profile_key=profile_key)) + + def get_by_id(self, profile_id: str) -> MemoryNode | None: + return self._run_sync(self.aget_by_id(profile_id)) + + def get_by_key(self, profile_key: str) -> MemoryNode | None: + return self._run_sync(self.aget_by_key(profile_key)) + + def get_all(self) -> list[MemoryNode]: + if isinstance(self.backend, FileProfileBackend): + return self.backend.get_all_sync() + return self._run_sync(self.aget_all()) + + def read_all(self, add_profile_id: bool = False, add_history_id: bool = False) -> str: + return self._run_sync(self.aread_all(add_profile_id, add_history_id)) diff --git a/reme/memory/vector_tools/profiles/read_all_profiles.py b/reme/memory/vector_tools/profiles/read_all_profiles.py index 1d7ddb1a..50eab684 100644 --- a/reme/memory/vector_tools/profiles/read_all_profiles.py +++ b/reme/memory/vector_tools/profiles/read_all_profiles.py @@ -1,8 +1,7 @@ -"""Read user profile tool""" +"""Read user profile tool.""" from loguru import logger -from .profile_handler import ProfileHandler from ..base_memory_tool import BaseMemoryTool from ....core.schema import ToolCall @@ -44,8 +43,8 @@ class ReadAllProfiles(BaseMemoryTool): else: target = self.memory_target - profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=target) - profiles_str = profile_handler.read_all(add_profile_id=True) + profile_handler = self.get_profile_handler(target) + profiles_str = await profile_handler.aread_all(add_profile_id=True) if not profiles_str: output = "No profiles found." logger.info(output) diff --git a/reme/memory/vector_tools/profiles/update_profile.py b/reme/memory/vector_tools/profiles/update_profile.py index 37512714..04c932bf 100644 --- a/reme/memory/vector_tools/profiles/update_profile.py +++ b/reme/memory/vector_tools/profiles/update_profile.py @@ -1,8 +1,7 @@ -"""Update user profile tool""" +"""Update user profile tool.""" from loguru import logger -from .profile_handler import ProfileHandler from ..base_memory_tool import BaseMemoryTool from ....core.schema import ToolCall @@ -82,8 +81,8 @@ class UpdateProfile(BaseMemoryTool): # Delete profiles (using self.memory_target) if profile_ids_to_delete: - profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=self.memory_target) - removed_count = profile_handler.delete(profile_ids_to_delete) + profile_handler = self.get_profile_handler(self.memory_target) + removed_count = await profile_handler.adelete(profile_ids_to_delete) # Add new profiles if profiles_to_add: @@ -98,14 +97,17 @@ class UpdateProfile(BaseMemoryTool): # Add profiles for each target 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) + profile_handler = self.get_profile_handler(target) + new_nodes = await profile_handler.aadd_batch( + profiles=target_profiles, + ref_memory_id=self.history_id, + ) self.memory_nodes.extend(new_nodes) added_count += len(new_nodes) else: # Use self.memory_target for all profiles - profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=self.memory_target) - new_nodes = profile_handler.add_batch(profiles=profiles_to_add, ref_memory_id=self.history_id) + profile_handler = self.get_profile_handler(self.memory_target) + new_nodes = await profile_handler.aadd_batch(profiles=profiles_to_add, ref_memory_id=self.history_id) self.memory_nodes.extend(new_nodes) added_count = len(new_nodes) diff --git a/reme/memory/vector_tools/profiles/update_profiles_v1.py b/reme/memory/vector_tools/profiles/update_profiles_v1.py index d38b3747..128188fb 100644 --- a/reme/memory/vector_tools/profiles/update_profiles_v1.py +++ b/reme/memory/vector_tools/profiles/update_profiles_v1.py @@ -1,8 +1,7 @@ -"""Update user profile tool""" +"""Update user profile tool.""" from loguru import logger -from .profile_handler import ProfileHandler from ..base_memory_tool import BaseMemoryTool from ....core.schema import ToolCall @@ -113,8 +112,8 @@ class UpdateProfilesV1(BaseMemoryTool): 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) + profile_handler = self.get_profile_handler(target) + await profile_handler.adelete(profile_ids) # Step 2: Prepare all profiles to add (both updated and new) all_profiles_to_add = [] @@ -158,8 +157,8 @@ class UpdateProfilesV1(BaseMemoryTool): 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) + profile_handler = self.get_profile_handler(target) + new_nodes = await profile_handler.aadd_batch(profiles=target_profiles, ref_memory_id=self.history_id) all_memory_nodes.extend(new_nodes) # Extend memory_nodes for tracking diff --git a/reme/reme.py b/reme/reme.py index f559ad71..c7c04d2e 100644 --- a/reme/reme.py +++ b/reme/reme.py @@ -14,6 +14,7 @@ from .memory.vector_tools import ( DelegateTask, ReadAllProfiles, ReadHistory, + RetrieveProfile, RetrieveMemory, UpdateProfilesV1, ) @@ -55,6 +56,10 @@ class ReMe(Application): target_task_names: list[str] | None = None, target_tool_names: list[str] | None = None, enable_profile: bool = True, + profile_backend: str = "filesystem", + profile_store_name: str = "profile", + profile_collection_name: str | None = None, + profile_max_capacity: int = 50, **kwargs, ): """Initialize ReMe with config. @@ -69,8 +74,37 @@ class ReMe(Application): ``` Args: + *args: Positional arguments forwarded to the base `Application`. + llm_api_key: API key used by the default LLM backend when provided. + llm_base_url: Base URL used by the default LLM backend when provided. + embedding_api_key: API key used by the default embedding backend when provided. + embedding_base_url: Base URL used by the default embedding backend when provided. + working_dir: Directory for generated config, logs, caches, and local stores. + config_path: Built-in config name or config file path used to initialize services. + enable_logo: Whether to print the ReMe logo during startup. + log_to_console: Whether to emit logs to the console. + log_to_file: Whether to write logs under `working_dir`. + default_llm_config: Overrides for the default LLM configuration. + default_embedding_model_config: Overrides for the default embedding model configuration. + default_vector_store_config: Configuration for the default memory vector store. + Its `collection_name` is used for normal memory storage. + default_token_counter_config: Overrides for the default token counter configuration. + target_user_names: Personal memory targets to register at initialization. + target_task_names: Procedural memory targets to register at initialization. + target_tool_names: Tool memory targets to register at initialization. enable_profile: Whether to enable profile functionality. Set to False when using - cloud-based vector stores to avoid local file operations. Default is True. + profile-free memory flows. + profile_backend: Profile storage backend. Use "filesystem" for local JSONL profile + files or "vector" for a dedicated profile vector collection. + profile_store_name: Internal vector store key used to register and look up the + profile vector store in `service_context.vector_stores`. This is not the + database collection name. + profile_collection_name: Dedicated database collection/table name for vector + profiles. When unset, vector profiles use the default memory collection name + with a "_profile" suffix. + profile_max_capacity: Maximum number of profile rows to keep per memory target. + When the limit is exceeded, the oldest profile rows are removed. + **kwargs: Additional keyword arguments forwarded to the base `Application`. """ super().__init__( *args, @@ -92,6 +126,10 @@ class ReMe(Application): ) self.enable_profile = enable_profile + self.profile_backend = profile_backend + self.profile_store_name = profile_store_name + self.profile_collection_name = profile_collection_name + self.profile_max_capacity = profile_max_capacity memory_target_type_mapping: dict[str, MemoryType] = {} if target_user_names: @@ -111,13 +149,16 @@ class ReMe(Application): self.service_context.memory_target_type_mapping = memory_target_type_mapping - if self.enable_profile: + if self.enable_profile and self.profile_backend == "filesystem": profile_path = Path(self.service_context.service_config.working_dir) / "profile" profile_path.mkdir(parents=True, exist_ok=True) self.profile_dir: str = str(profile_path) else: self.profile_dir: str = "" + if self.enable_profile and self.profile_backend == "vector": + self._ensure_profile_vector_store_config() + def _add_meta_memory(self, memory_type: str | MemoryType, memory_target: str): """Register or validate a memory target with the given memory type.""" if memory_target in self.service_context.memory_target_type_mapping: @@ -186,6 +227,38 @@ class ReMe(Application): return result return result["answer"] + def _ensure_profile_vector_store_config(self) -> None: + """Ensure the dedicated profile vector store exists in service config.""" + vector_store_configs = self.service_context.service_config.vector_stores + if "default" not in vector_store_configs: + raise RuntimeError("Vector profile backend requires a default vector store configuration") + + default_config = vector_store_configs["default"] + profile_collection_name = self.profile_collection_name or f"{default_config.collection_name}_profile" + + if self.profile_store_name in vector_store_configs: + if self.profile_collection_name: + vector_store_configs[self.profile_store_name] = vector_store_configs[ + self.profile_store_name + ].model_copy( + update={"collection_name": profile_collection_name}, + ) + return + + vector_store_configs[self.profile_store_name] = default_config.model_copy( + update={"collection_name": profile_collection_name}, + ) + + def _get_profile_tool_kwargs(self, raise_exception: bool) -> dict: + """Shared profile tool configuration.""" + return { + "profile_dir": self.profile_dir, + "profile_backend": self.profile_backend, + "profile_store_name": self.profile_store_name, + "profile_max_capacity": self.profile_max_capacity, + "raise_exception": raise_exception, + } + async def summarize_memory( self, messages: list[Message | dict], @@ -211,6 +284,7 @@ class ReMe(Application): format_messages.append(message) if version == "default": + profile_tool_kwargs = self._get_profile_tool_kwargs(raise_exception) personal_summarizer_tools: list = [ AddDraftAndRetrieveSimilarMemory( enable_thinking_params=enable_thinking_params, @@ -229,20 +303,28 @@ class ReMe(Application): ), ] if self.enable_profile: + if self.profile_backend == "vector": + profile_context_tool = RetrieveProfile( + top_k=min(5, retrieve_top_k), + enable_thinking_params=False, + enable_memory_target=False, + enable_multiple=False, + **profile_tool_kwargs, + ) + else: + profile_context_tool = ReadAllProfiles( + enable_thinking_params=False, + enable_memory_target=False, + **profile_tool_kwargs, + ) personal_summarizer_tools.extend( [ - ReadAllProfiles( - enable_thinking_params=False, - enable_memory_target=False, - profile_dir=self.profile_dir, - raise_exception=raise_exception, - ), + profile_context_tool, UpdateProfilesV1( enable_thinking_params=enable_thinking_params, enable_memory_target=False, enable_multiple=True, - profile_dir=self.profile_dir, - raise_exception=raise_exception, + **profile_tool_kwargs, ), ], ) @@ -381,16 +463,24 @@ class ReMe(Application): self._ensure_started() if version == "default": + profile_tool_kwargs = self._get_profile_tool_kwargs(raise_exception) personal_retriever_tools = [] if self.enable_profile: - personal_retriever_tools.append( - ReadAllProfiles( + if self.profile_backend == "vector": + profile_context_tool = RetrieveProfile( + top_k=min(5, retrieve_top_k), enable_thinking_params=False, enable_memory_target=False, - profile_dir=self.profile_dir, - raise_exception=raise_exception, - ), - ) + enable_multiple=False, + **profile_tool_kwargs, + ) + else: + profile_context_tool = ReadAllProfiles( + enable_thinking_params=False, + enable_memory_target=False, + **profile_tool_kwargs, + ) + personal_retriever_tools.append(profile_context_tool) personal_retriever_tools.extend( [ RetrieveMemory( @@ -509,6 +599,34 @@ class ReMe(Application): return self._unwrap_memory_result(result, "retrieve_memory", return_dict) + async def retrieve_profile( + self, + query: str | list[str], + user_name: str, + top_k: int = 5, + return_dict: bool = False, + ) -> str | dict: + """Retrieve relevant profile rows for a user.""" + self._ensure_started() + if not self.enable_profile: + raise RuntimeError("Profile functionality is disabled.") + + profile_handler = self.get_profile_handler(user_name) + if profile_handler is None: + raise RuntimeError("Profile functionality is disabled.") + + retrieved_nodes, output = await profile_handler.aretrieve( + query=query, + limit=top_k, + add_profile_id=True, + add_history_id=True, + ) + result = { + "answer": output or "No matching profiles found.", + "retrieved_nodes": retrieved_nodes, + } + return self._unwrap_memory_result(result, "retrieve_profile", return_dict) + async def add_memory( self, memory_content: str, @@ -675,15 +793,23 @@ class ReMe(Application): @property def profile_path(self) -> Path | None: """Get the path to the profile directory. Returns None if profile is disabled.""" - if not self.enable_profile: + if not self.enable_profile or self.profile_backend != "filesystem": return None - return Path(self.profile_dir) / self.default_vector_store.collection_name + collection_name = self.service_context.service_config.vector_stores["default"].collection_name + return Path(self.profile_dir) / collection_name def get_profile_handler(self, user_name: str) -> ProfileHandler | None: """Get the profile handler for the specified user. Returns None if profile is disabled.""" if not self.enable_profile: return None - return ProfileHandler(memory_target=user_name, profile_path=self.profile_path) + return ProfileHandler( + memory_target=user_name, + profile_path=self.profile_path, + service_context=self.service_context, + profile_backend=self.profile_backend, + profile_store_name=self.profile_store_name, + max_capacity=self.profile_max_capacity, + ) def main(): diff --git a/tests/test_reme_memory_error_handling.py b/tests/test_reme_memory_error_handling.py index 99953efd..563e8221 100644 --- a/tests/test_reme_memory_error_handling.py +++ b/tests/test_reme_memory_error_handling.py @@ -1,9 +1,14 @@ """Tests for ReMe memory error handling and raise_exception propagation.""" +from types import SimpleNamespace + import pytest import reme.reme as reme_module -from reme import ReMe +from reme.core.runtime_context import RuntimeContext +from reme.core.schema import MemoryNode +from reme.memory.vector_tools.history.read_history import ReadHistory +from reme.reme import ReMe class Recorder: @@ -130,3 +135,34 @@ async def test_summarize_memory_raises_runtime_error_for_unstructured_result(mon messages=[{"role": "user", "content": "hi", "time_created": "2026-03-20 10:00:00"}], task_name="demo-task", ) + + +@pytest.mark.asyncio +async def test_read_history_accepts_single_history_id_in_multiple_mode(): + """Verify multiple-mode history lookup accepts a single history_id string.""" + + class FakeVectorStore: + """Minimal vector store stub for ReadHistory tests.""" + + async def get(self, vector_ids): + """Return the requested history node.""" + assert vector_ids == ["history_123"] + node = MemoryNode( + memory_id="history_123", + memory_type="history", + memory_target="alice", + content="Alice said hello.", + ) + return [node.to_vector_node()] + + tool = ReadHistory(enable_multiple=True) + tool._vector_store = FakeVectorStore() # pylint: disable=protected-access + tool.context = RuntimeContext( + history_id="history_123", + retrieved_nodes=[], + service_context=SimpleNamespace(memory_target_type_mapping={"alice": "personal"}), + ) + + result = await tool.execute() + + assert "Historical Dialogue[history_123]" in result