From de28c8e243a2b623eb1c5e2b7e12c87cea9b0d38 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sat, 7 Feb 2026 03:37:12 +0800 Subject: [PATCH] refactor(core): rename memory_storage to memory_store and add base fs tool --- reme/__init__.py | 2 + reme/agent/fs/fs_compactor.py | 198 +++-- reme/agent/fs/fs_compactor.yaml | 16 + reme/agent/fs/fs_summarizer.py | 35 +- reme/agent/fs/fs_summarizer.yaml | 14 +- reme/core/__init__.py | 4 + reme/core/application.py | 2 +- reme/core/context/service_context.py | 116 +-- reme/core/file_watcher/base_file_watcher.py | 2 +- reme/core/llm/lite_llm.py | 3 +- reme/core/llm/lite_llm_sync.py | 4 +- .../__init__.py | 2 +- .../base_memory_store.py | 0 .../sqlite_memory_store.py | 26 +- reme/core/op/base_op.py | 2 +- reme/core/schema/__init__.py | 2 + reme/core/schema/compaction_result.py | 15 + reme/core/schema/memory_search_result.py | 4 +- reme/reme.py | 2 +- reme/reme_fs.py | 106 ++- reme/tool/fs/__init__.py | 6 + reme/tool/fs/base_fs_tool.py | 41 + reme/tool/fs/bash_tool.py | 62 +- reme/tool/fs/edit_tool.py | 18 +- reme/tool/fs/find_tool.py | 41 +- reme/tool/fs/fs_memory_get.py | 42 +- reme/tool/fs/fs_memory_search.py | 12 +- reme/tool/fs/grep_tool.py | 19 +- reme/tool/fs/ls_tool.py | 9 +- reme/tool/fs/read_tool.py | 42 +- reme/tool/fs/write_tool.py | 4 +- tests/demo_memory_search.py | 2 +- tests/test_fs_agent.py | 323 -------- tests/test_fs_compact.py | 301 ++++++++ tests/test_fs_memory_get.py | 369 +++++++++ tests/test_fs_memory_search.py | 704 ++++++++++++++++++ tests/test_fs_summary.py | 234 ++++++ ...st_file_system_tool.py => test_fs_tool.py} | 0 tests/test_memory_store.py | 4 +- 39 files changed, 2096 insertions(+), 692 deletions(-) rename reme/core/{memory_storage => memory_store}/__init__.py (86%) rename reme/core/{memory_storage => memory_store}/base_memory_store.py (100%) rename reme/core/{memory_storage => memory_store}/sqlite_memory_store.py (96%) create mode 100644 reme/core/schema/compaction_result.py create mode 100644 reme/tool/fs/base_fs_tool.py delete mode 100644 tests/test_fs_agent.py create mode 100644 tests/test_fs_compact.py create mode 100644 tests/test_fs_memory_get.py create mode 100644 tests/test_fs_memory_search.py create mode 100644 tests/test_fs_summary.py rename tests/{test_file_system_tool.py => test_fs_tool.py} (100%) diff --git a/reme/__init__.py b/reme/__init__.py index cd5d4e6c..81fe52cf 100644 --- a/reme/__init__.py +++ b/reme/__init__.py @@ -6,6 +6,7 @@ from . import core from . import tool from . import workflow from .reme import ReMe +from .reme_fs import ReMeFs __all__ = [ "agent", @@ -14,6 +15,7 @@ __all__ = [ "tool", "workflow", "ReMe", + "ReMeFs", ] __version__ = "0.3.0.0a1" diff --git a/reme/agent/fs/fs_compactor.py b/reme/agent/fs/fs_compactor.py index 79d70c6f..07c2dd18 100644 --- a/reme/agent/fs/fs_compactor.py +++ b/reme/agent/fs/fs_compactor.py @@ -4,7 +4,7 @@ from loguru import logger from ...core.enumeration import Role, MemoryType from ...core.op import BaseReact -from ...core.schema import Message +from ...core.schema import CutPointResult, Message class FsCompactor(BaseReact): @@ -24,44 +24,39 @@ class FsCompactor(BaseReact): self.reserve_tokens: int = reserve_tokens self.keep_recent_tokens: int = keep_recent_tokens + @staticmethod + def _normalize_messages(messages: list[Message | dict]) -> list[Message]: + """Convert dict messages to Message objects.""" + return [Message(**m) if isinstance(m, dict) else m for m in messages] + @staticmethod def _is_user_message(message: Message) -> bool: - """Check if a message is a user-initiated message (user or tool result).""" + """Check if message is user role.""" return message.role is Role.USER def _find_turn_start_index(self, messages: list[Message], entry_index: int) -> int: - """ - Find the user message that starts the turn containing the given entry index. - Returns -1 if no turn start found before the index. - """ + """Find user message that starts the turn. Returns -1 if not found.""" + if not messages or entry_index < 0 or entry_index >= len(messages): + return -1 + for i in range(entry_index, -1, -1): if self._is_user_message(messages[i]): return i return -1 - def _find_cut_point(self, messages: list[Message]) -> dict: + def _find_cut_point(self, messages: list[Message]) -> CutPointResult: """ Find cut point with split turn detection. - A "split turn" occurs when the cut point falls in the middle of a conversation turn - rather than at a clean user message boundary. For example: - User → Assistant → [CUT HERE] → Assistant continues → User - - In this case, we need to: - 1. Summarize complete history (before turn start) - 2. Separately summarize the turn prefix (turn start to cut point) - 3. Keep the turn suffix (cut point onwards) in full - - Returns dict with: - - messages_to_summarize: Complete turns before the current turn - - turn_prefix_messages: If split turn, messages from turn start to cut point - - is_split_turn: Whether this is a split turn - - cut_index: The actual cut point index + Split turn: User → Assistant → [CUT] → Assistant → User + Clean cut: User → [CUT] → Assistant → User """ + if not messages: + return CutPointResult() + accumulated_tokens = 0 cut_index = 0 - # Walk backwards from the newest messages, accumulating tokens until we hit the keep threshold for i in range(len(messages) - 1, -1, -1): msg = messages[i] msg_tokens = self.token_counter.count_token([msg]) @@ -69,51 +64,51 @@ class FsCompactor(BaseReact): if accumulated_tokens >= self.keep_recent_tokens: cut_index = i + logger.debug(f"Cut point at index {cut_index}, {accumulated_tokens} tokens") break if cut_index == 0: - return { - "messages_to_summarize": [], - "turn_prefix_messages": [], - "is_split_turn": False, - "cut_index": 0, - } + return CutPointResult(left_messages=messages) - # Check if cut point is a user message (clean turn boundary) or assistant/other (mid-turn) cut_message = messages[cut_index] is_user_cut = self._is_user_message(cut_message) if is_user_cut: - # Clean cut: cut point is at a turn boundary, summarize everything before - return { - "messages_to_summarize": messages[:cut_index], - "turn_prefix_messages": [], - "is_split_turn": False, - "cut_index": cut_index, - } + return CutPointResult( + messages_to_summarize=messages[:cut_index], + left_messages=messages[cut_index:], + cut_index=cut_index, + ) - # Split turn detected: find where the current turn started turn_start_index = self._find_turn_start_index(messages, cut_index) if turn_start_index == -1: - # No turn start found (shouldn't happen), treat as clean cut - return { - "messages_to_summarize": messages[:cut_index], - "turn_prefix_messages": [], - "is_split_turn": False, - "cut_index": cut_index, - } + logger.warning("Split turn detected but no turn start found, treating as clean cut") + return CutPointResult( + messages_to_summarize=messages[:cut_index], + left_messages=messages[cut_index:], + cut_index=cut_index, + ) - # Split turn: separate complete history from turn prefix - # History: [0, turn_start_index) - complete turns to summarize - # Turn prefix: [turn_start_index, cut_index) - needs special context summary - # Turn suffix: [cut_index, end) - kept in full (recent work) - return { - "messages_to_summarize": messages[:turn_start_index], - "turn_prefix_messages": messages[turn_start_index:cut_index], - "is_split_turn": True, - "cut_index": cut_index, - } + return CutPointResult( + messages_to_summarize=messages[:turn_start_index], + turn_prefix_messages=messages[turn_start_index:cut_index], + left_messages=messages[cut_index:], + is_split_turn=True, + cut_index=cut_index, + ) + + async def _generate_summary(self, prompt_messages: list[Message]) -> str: + """Generate summary via LLM. Returns empty string if no messages.""" + if not prompt_messages: + return "" + + try: + assistant_message = await self.llm.chat(prompt_messages) + return assistant_message.content if assistant_message.content else "" + except Exception as e: + logger.error(f"Failed to generate summary: {e}") + raise RuntimeError(f"Summarization failed: {e}") from e @staticmethod def _serialize_conversation(messages: list[Message]) -> str: @@ -135,20 +130,15 @@ class FsCompactor(BaseReact): return "\n".join(lines) def build_messages_s1(self) -> list[Message]: - """ - Build messages for compaction summarization. - - This creates the prompt for the main history summary. If split turn is detected, - a separate turn prefix summary will be generated later in execute(). - """ - messages: list[Message] = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages] - + """Build prompt for main history summary.""" + messages = self._normalize_messages(self.context.messages) cut_result = self._find_cut_point(messages) - messages_to_summarize = cut_result["messages_to_summarize"] - self.context.is_split_turn = cut_result["is_split_turn"] - self.context.turn_prefix_messages = cut_result["turn_prefix_messages"] - if not messages_to_summarize: + self.context.is_split_turn = cut_result.is_split_turn + self.context.turn_prefix_messages = cut_result.turn_prefix_messages + self.context.left_messages = cut_result.left_messages + + if not cut_result.messages_to_summarize: logger.info("No messages to summarize") return [] @@ -157,7 +147,7 @@ class FsCompactor(BaseReact): user_prompt = self.prompt_format("update_user_message", previous_summary=self.context.previous_summary) else: user_prompt = self.get_prompt("initial_user_message") - conversation_text = self._serialize_conversation(messages_to_summarize) + conversation_text = self._serialize_conversation(cut_result.messages_to_summarize) return [ Message(role=Role.SYSTEM, content=system_prompt), @@ -165,16 +155,7 @@ class FsCompactor(BaseReact): ] def build_messages_s2(self) -> list[Message]: - """ - Generate summary for turn prefix when splitting a turn. - - This provides context for the retained turn suffix. The summary focuses on: - - What the user originally asked for in this turn - - Key decisions and early progress made in the prefix - - Information needed to understand the kept suffix - - This is shorter and more focused than the full history summary. - """ + """Build prompt for turn prefix summary (split turn only).""" if not self.context.turn_prefix_messages: return [] @@ -191,55 +172,54 @@ class FsCompactor(BaseReact): """ Execute compaction if needed. - Compaction process: - 1. Check if token count exceeds threshold - 2. Find cut point and detect if it's a split turn - 3. Generate history summary (complete turns before cut point) - 4. If split turn: generate turn prefix summary (partial turn before cut point) - 5. Merge summaries and update context - - Final context structure after compaction: - - Summary (history + optional turn prefix context) - - Recent messages kept in full (from cut point onwards) + Returns: [summary_message, ...left_messages] if compacted, else original messages. """ - messages: list[Message] = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages] - token_count: int = self.token_counter.count_token(messages) + original_messages = self._normalize_messages(self.context.messages) + token_count: int = self.token_counter.count_token(original_messages) threshold = self.context_window_tokens - self.reserve_tokens + if token_count < threshold: - logger.info(f"Token count {token_count} below threshold, skipping compaction") + logger.info(f"Token count {token_count} below threshold ({threshold}), skipping compaction") return { - "answer": "", - "success": True, - "messages": [], - "tools": [], - "skipped": True, + "compacted": False, + "tokens_before": token_count, + "is_split_turn": False, + "messages": original_messages, } - logger.info(f"Starting compaction, token count: {token_count}") - messages: list[Message] = self.build_messages_s1() - if messages: - assistant_message = await self.llm.chat(messages) - history_summary = assistant_message.content - else: - history_summary = "" + logger.info(f"Starting compaction, token count: {token_count}, threshold: {threshold}") + + history_prompt_messages = self.build_messages_s1() + + if not history_prompt_messages and not self.context.get("is_split_turn"): + logger.warning("No messages to summarize and not a split turn, returning original messages") + return { + "compacted": False, + "tokens_before": token_count, + "is_split_turn": False, + "messages": original_messages, + } + + history_summary = await self._generate_summary(history_prompt_messages) if history_prompt_messages else "" if self.context.is_split_turn and self.context.turn_prefix_messages: logger.info("Split turn detected, generating turn prefix summary") - messages: list[Message] = self.build_messages_s2() - if messages: - assistant_message = await self.llm.chat(messages) - turn_prefix_summary = assistant_message.content - else: - turn_prefix_summary = "" + turn_prefix_prompt_messages = self.build_messages_s2() + turn_prefix_summary = await self._generate_summary(turn_prefix_prompt_messages) summary = f"{history_summary}\n\n---\n\n**Turn Context (split turn):**\n\n{turn_prefix_summary}" else: summary = history_summary logger.info(f"Compaction complete, summary length: {len(summary)}, split_turn: {self.context.is_split_turn}") + summary_content = self.prompt_format("compaction_summary_format", summary=summary) + summary_message = Message(role=Role.USER, content=summary_content) + left_messages = self.context.get("left_messages", []) + final_messages = [summary_message] + left_messages + return { "compacted": True, "tokens_before": token_count, - "summary": summary, "is_split_turn": self.context.is_split_turn, + "messages": final_messages, } diff --git a/reme/agent/fs/fs_compactor.yaml b/reme/agent/fs/fs_compactor.yaml index 40fb6ec4..dbad84f2 100644 --- a/reme/agent/fs/fs_compactor.yaml +++ b/reme/agent/fs/fs_compactor.yaml @@ -210,3 +210,19 @@ turn_prefix_summarization_zh: | - [理解保留的最近工作所需的信息] 保持简洁。专注于理解保留后缀所需的内容。 + +# Format for wrapping the compaction summary when presenting to LLM +# This matches the TypeScript format: COMPACTION_SUMMARY_PREFIX + summary + COMPACTION_SUMMARY_SUFFIX +compaction_summary_format: | + The conversation history before this point was compacted into the following summary: + + + {summary} + + +compaction_summary_format_zh: | + 此点之前的对话历史已被压缩为以下摘要: + + + {summary} + diff --git a/reme/agent/fs/fs_summarizer.py b/reme/agent/fs/fs_summarizer.py index 25176dc1..0b88ad5c 100644 --- a/reme/agent/fs/fs_summarizer.py +++ b/reme/agent/fs/fs_summarizer.py @@ -1,5 +1,7 @@ """Personal memory retriever agent for retrieving personal memories through vector search.""" +import datetime + from loguru import logger from ...core.enumeration import Role, MemoryType @@ -14,28 +16,25 @@ class FsSummarizer(BaseReact): def __init__( self, - memory_dir: str, + memory_dir: str = "memory", version: str = "default", - context_window_tokens: int = 128000, - reserve_tokens: int = 32000, - soft_threshold_tokens: int = 4000, **kwargs, ): super().__init__(**kwargs) self.memory_dir: str = memory_dir self.version: str = version - self.context_window_tokens: int = context_window_tokens - self.reserve_tokens: int = reserve_tokens - self.soft_threshold_tokens: int = soft_threshold_tokens async def build_messages(self) -> list[Message]: messages: list[Message] = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages] + date_str: str = self.context.get("date", datetime.datetime.now().strftime("%Y-%m-%d")) + if self.version == "default": messages.append( Message( role=Role.USER, content=self.prompt_format( "user_message_v2", + date=date_str, memory_dir=self.memory_dir, ), ), @@ -47,6 +46,7 @@ class FsSummarizer(BaseReact): role=Role.USER, content=self.prompt_format( "user_message", + date=date_str, memory_dir=self.memory_dir, ), ), @@ -54,27 +54,6 @@ class FsSummarizer(BaseReact): return messages async def execute(self): - context_window = max(1, int(self.context_window_tokens)) - reserve_tokens = max(0, int(self.reserve_tokens)) - soft_threshold = max(0, int(self.soft_threshold_tokens)) - threshold = max(0, context_window - reserve_tokens - soft_threshold) - messages: list[Message] = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages] - token_count: int = self.token_counter.count_token(messages) - - if token_count >= threshold: - logger.info(f"[{self.__class__.__name__}] Skipping summary execution based on threshold check") - return { - "answer": "", - "success": True, - "messages": [], - "tools": [], - "skipped": True, - } - - # Mark that we're executing a summary in this cycle - summary_count = self.context.get("summary_count", 0) - self.context["last_summary_at"] = summary_count - result = await super().execute() answer = result["answer"] logger.info(f"[{self.__class__.__name__}] answer={answer}") diff --git a/reme/agent/fs/fs_summarizer.yaml b/reme/agent/fs/fs_summarizer.yaml index f4673750..ef3de8ae 100644 --- a/reme/agent/fs/fs_summarizer.yaml +++ b/reme/agent/fs/fs_summarizer.yaml @@ -5,11 +5,19 @@ system_prompt: | user_message: | Pre-compaction memory flush. + Current Date: {date} Store durable memories now (use {memory_dir}/YYYY-MM-DD.md; create {memory_dir}/ if needed). If nothing to store, reply with [SILENT]. user_message_v2: | - Pre-compaction memory flush. + Current Date: {date} The session is near auto-compaction; capture durable memories to disk. - Store durable memories now (use {memory_dir}/YYYY-MM-DD.md; create {memory_dir}/ if needed). - If nothing to store, reply with [SILENT]. \ No newline at end of file + + Memory storage workflow: + 1. Check if {memory_dir}/ exists; if not, create it via bash + 2. Check if {memory_dir}/YYYY-MM-DD.md exists (use actual date) + 3. If file is NEW: Write memories directly (be concise) + 4. If file EXISTS: Read it first, then UPDATE with new memories (keep concise, merge/deduplicate) + 5. If NO valuable information to store: Reply with reason and [SILENT] + + Store durable memories. Keep entries concise and well-organized. \ No newline at end of file diff --git a/reme/core/__init__.py b/reme/core/__init__.py index 2b56fb08..4128cbd9 100644 --- a/reme/core/__init__.py +++ b/reme/core/__init__.py @@ -3,8 +3,10 @@ from . import context from . import embedding from . import enumeration +from . import file_watcher from . import flow from . import llm +from . import memory_store from . import op from . import schema from . import service @@ -18,8 +20,10 @@ __all__ = [ "context", "embedding", "enumeration", + "file_watcher", "flow", "llm", + "memory_store", "op", "schema", "service", diff --git a/reme/core/application.py b/reme/core/application.py index ff093692..fde84afe 100644 --- a/reme/core/application.py +++ b/reme/core/application.py @@ -7,7 +7,7 @@ from .embedding import BaseEmbeddingModel from .file_watcher import BaseFileWatcher from .flow import BaseFlow from .llm import BaseLLM -from .memory_storage import BaseMemoryStore +from .memory_store import BaseMemoryStore from .schema import Response from .token_counter import BaseTokenCounter from .utils import execute_stream_task, PydanticConfigParser diff --git a/reme/core/context/service_context.py b/reme/core/context/service_context.py index a74eac40..f6f07efa 100644 --- a/reme/core/context/service_context.py +++ b/reme/core/context/service_context.py @@ -15,7 +15,7 @@ if TYPE_CHECKING: from ..llm import BaseLLM from ..embedding import BaseEmbeddingModel from ..vector_store import BaseVectorStore - from ..memory_storage import BaseMemoryStore + from ..memory_store import BaseMemoryStore from ..token_counter import BaseTokenCounter from ..flow import BaseFlow from ..service import BaseService @@ -45,24 +45,39 @@ class ServiceContext(BaseContext): **kwargs, ): super().__init__() - self.service_config: ServiceConfig = self._build_service_config( - *args, - llm_api_key=llm_api_key, - llm_api_base=llm_api_base, - embedding_api_key=embedding_api_key, - embedding_api_base=embedding_api_base, - service_config=service_config, - parser=parser, - config_path=config_path, - enable_logo=enable_logo, - llm=llm, - embedding_model=embedding_model, - vector_store=vector_store, - memory_store=memory_store, - token_counter=token_counter, - file_watcher=file_watcher, - **kwargs, - ) + + load_env() + self._update_env("REME_LLM_API_KEY", llm_api_key) + self._update_env("REME_LLM_BASE_URL", llm_api_base) + self._update_env("REME_EMBEDDING_API_KEY", embedding_api_key) + self._update_env("REME_EMBEDDING_BASE_URL", embedding_api_base) + + if service_config is None: + parser_class = parser if parser is not None else PydanticConfigParser + parser = parser_class(ServiceConfig) + input_args = [] + if config_path: + input_args.append(f"config={config_path}") + if args: + input_args.extend(args) + + if llm: + self._update_section_config(kwargs, "llm", **llm) + if embedding_model: + self._update_section_config(kwargs, "embedding_model", **embedding_model) + if token_counter: + self._update_section_config(kwargs, "token_counter", **token_counter) + if vector_store: + self._update_section_config(kwargs, "vector_store", **vector_store) + if memory_store: + self._update_section_config(kwargs, "memory_store", **memory_store) + if file_watcher: + self._update_section_config(kwargs, "file_watcher", **file_watcher) + kwargs["enable_logo"] = enable_logo + logger.info(f"update with args: {input_args} kwargs: {kwargs}") + service_config = parser.parse_args(*input_args, **kwargs) + + self.service_config: ServiceConfig = service_config if self.service_config.init_logger: init_logger() @@ -92,57 +107,6 @@ class ServiceContext(BaseContext): self._build_flows() - def _build_service_config( - 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, - service_config: ServiceConfig | None = None, - parser: type[PydanticConfigParser] | None = None, - config_path: str | None = None, - enable_logo: bool = True, - llm: dict | None = None, - embedding_model: dict | None = None, - vector_store: dict | None = None, - memory_store: dict | None = None, - token_counter: dict | None = None, - file_watcher: dict | None = None, - **kwargs, - ) -> ServiceConfig: - - load_env() - self._update_env("REME_LLM_API_KEY", llm_api_key) - self._update_env("REME_LLM_BASE_URL", llm_api_base) - self._update_env("REME_EMBEDDING_API_KEY", embedding_api_key) - self._update_env("REME_EMBEDDING_BASE_URL", embedding_api_base) - - if service_config is None: - parser_class = parser if parser is not None else PydanticConfigParser - parser = parser_class(ServiceConfig) - input_args = [] - if config_path: - input_args.append(f"config={config_path}") - if args: - input_args.extend(args) - service_config = parser.parse_args(*input_args, **kwargs) - - service_config.enable_logo = enable_logo - if llm: - self._update_section_config(service_config, "llm", **llm) - if embedding_model: - self._update_section_config(service_config, "embedding_model", **embedding_model) - if token_counter: - self._update_section_config(service_config, "token_counter", **token_counter) - if vector_store: - self._update_section_config(service_config, "vector_store", **vector_store) - if memory_store: - self._update_section_config(service_config, "memory_store", **memory_store) - if file_watcher: - self._update_section_config(service_config, "file_watcher", **file_watcher) - return service_config - @staticmethod def _update_env(key: str, value: str | None): """Update environment variable if value is provided.""" @@ -150,13 +114,13 @@ class ServiceContext(BaseContext): os.environ[key] = value @staticmethod - def _update_section_config(service_config: ServiceConfig, section_name: str, **kwargs): + def _update_section_config(config: dict, section_name: str, **kwargs): """Update a specific section of the service config with new values.""" - section_dict: dict = getattr(service_config, section_name) - if "default" not in section_dict: - raise KeyError(f"Default `{section_name}` config not found") - current_config = section_dict["default"] - section_dict["default"] = current_config.model_copy(update=kwargs, deep=True) + if section_name not in config: + config[section_name] = {} + if "default" not in config[section_name]: + config[section_name]["default"] = {} + config[section_name]["default"].update(kwargs) def _build_flows(self): expression_flow_cls = None diff --git a/reme/core/file_watcher/base_file_watcher.py b/reme/core/file_watcher/base_file_watcher.py index 620d2e17..a4e8af36 100644 --- a/reme/core/file_watcher/base_file_watcher.py +++ b/reme/core/file_watcher/base_file_watcher.py @@ -11,7 +11,7 @@ from typing import Any, Callable from loguru import logger from watchfiles import awatch, Change -from ..memory_storage import BaseMemoryStore +from ..memory_store import BaseMemoryStore class BaseFileWatcher: diff --git a/reme/core/llm/lite_llm.py b/reme/core/llm/lite_llm.py index 5663702b..7950a8c2 100644 --- a/reme/core/llm/lite_llm.py +++ b/reme/core/llm/lite_llm.py @@ -3,7 +3,6 @@ import os from typing import AsyncGenerator -import litellm from loguru import logger from .base_llm import BaseLLM @@ -88,6 +87,8 @@ class LiteLLM(BaseLLM): stream_kwargs: dict | None = None, ) -> AsyncGenerator[StreamChunk, None]: """Execute async streaming chat requests and yield processed response chunks.""" + import litellm + stream_kwargs = stream_kwargs or {} completion = await litellm.acompletion(**stream_kwargs) ret_tool_calls: list[ToolCall] = [] diff --git a/reme/core/llm/lite_llm_sync.py b/reme/core/llm/lite_llm_sync.py index 778eaed0..9951d223 100644 --- a/reme/core/llm/lite_llm_sync.py +++ b/reme/core/llm/lite_llm_sync.py @@ -2,8 +2,6 @@ from typing import Generator -import litellm - from .lite_llm import LiteLLM from ..enumeration import ChunkEnum from ..schema import Message @@ -21,6 +19,8 @@ class LiteLLMSync(LiteLLM): stream_kwargs: dict | None = None, ) -> Generator[StreamChunk, None, None]: """Internal synchronous generator for processing streaming chat completion chunks.""" + import litellm + stream_kwargs = stream_kwargs or {} completion = litellm.completion(**stream_kwargs) ret_tool_calls: list[ToolCall] = [] diff --git a/reme/core/memory_storage/__init__.py b/reme/core/memory_store/__init__.py similarity index 86% rename from reme/core/memory_storage/__init__.py rename to reme/core/memory_store/__init__.py index 371d4e0f..6d43e423 100644 --- a/reme/core/memory_storage/__init__.py +++ b/reme/core/memory_store/__init__.py @@ -1,4 +1,4 @@ -"""Memory storage module for persistent memory management. +"""Memory store module for persistent memory management. This module provides storage backends for memory chunks and file metadata, including SQLite-based implementations with vector and full-text search. diff --git a/reme/core/memory_storage/base_memory_store.py b/reme/core/memory_store/base_memory_store.py similarity index 100% rename from reme/core/memory_storage/base_memory_store.py rename to reme/core/memory_store/base_memory_store.py diff --git a/reme/core/memory_storage/sqlite_memory_store.py b/reme/core/memory_store/sqlite_memory_store.py similarity index 96% rename from reme/core/memory_storage/sqlite_memory_store.py rename to reme/core/memory_store/sqlite_memory_store.py index 756c2b3b..7e1519b8 100644 --- a/reme/core/memory_storage/sqlite_memory_store.py +++ b/reme/core/memory_store/sqlite_memory_store.py @@ -208,14 +208,17 @@ class SqliteMemoryStore(BaseMemoryStore): ), ) - # Insert vector + # Insert vector (vec0 doesn't support OR REPLACE, use DELETE + INSERT) if self.vector_available: assert chunk.embedding, "Embedding is required for vector insert" + # Delete existing vector first cursor.execute( - f""" - INSERT OR REPLACE INTO {self.vector_table_name} (id, embedding) - VALUES (?, ?) - """, + f"DELETE FROM {self.vector_table_name} WHERE id = ?", + (chunk.id,), + ) + # Then insert new vector + cursor.execute( + f"INSERT INTO {self.vector_table_name} (id, embedding) VALUES (?, ?)", (chunk.id, self.vector_to_blob(chunk.embedding)), ) @@ -372,14 +375,17 @@ class SqliteMemoryStore(BaseMemoryStore): ), ) - # Insert/update vector + # Insert/update vector (vec0 doesn't support OR REPLACE, use DELETE + INSERT) if self.vector_available: assert chunk.embedding, "Embedding is required for vector insert" + # Delete existing vector first cursor.execute( - f""" - INSERT OR REPLACE INTO {self.vector_table_name} (id, embedding) - VALUES (?, ?) - """, + f"DELETE FROM {self.vector_table_name} WHERE id = ?", + (chunk.id,), + ) + # Then insert new vector + cursor.execute( + f"INSERT INTO {self.vector_table_name} (id, embedding) VALUES (?, ?)", (chunk.id, self.vector_to_blob(chunk.embedding)), ) diff --git a/reme/core/op/base_op.py b/reme/core/op/base_op.py index 015d0fbb..8911cead 100644 --- a/reme/core/op/base_op.py +++ b/reme/core/op/base_op.py @@ -13,7 +13,7 @@ from tqdm import tqdm from ..context import RuntimeContext, PromptHandler, ServiceContext from ..embedding import BaseEmbeddingModel from ..llm import BaseLLM -from ..memory_storage import BaseMemoryStore +from ..memory_store import BaseMemoryStore from ..schema import Response from ..token_counter import BaseTokenCounter from ..utils import camel_to_snake, CacheHandler, timer diff --git a/reme/core/schema/__init__.py b/reme/core/schema/__init__.py index d4c5c051..ae6b8392 100644 --- a/reme/core/schema/__init__.py +++ b/reme/core/schema/__init__.py @@ -1,5 +1,6 @@ """schema""" +from .compaction_result import CutPointResult from .file_metadata import FileMetadata from .memory_chunk import MemoryChunk from .memory_node import MemoryNode @@ -26,6 +27,7 @@ from .vector_node import VectorNode __all__ = [ "CmdConfig", "ContentBlock", + "CutPointResult", "EmbeddingModelConfig", "FileMetadata", "FlowConfig", diff --git a/reme/core/schema/compaction_result.py b/reme/core/schema/compaction_result.py new file mode 100644 index 00000000..0cc1d86e --- /dev/null +++ b/reme/core/schema/compaction_result.py @@ -0,0 +1,15 @@ +"""Compaction result schemas for context window management.""" + +from pydantic import BaseModel, Field + +from .message import Message + + +class CutPointResult(BaseModel): + """Cut point detection result for conversation compaction.""" + + messages_to_summarize: list[Message] = Field(default_factory=list, description="Complete turns before cut point") + turn_prefix_messages: list[Message] = Field(default_factory=list, description="Turn prefix if split turn") + left_messages: list[Message] = Field(default_factory=list, description="Messages to keep from cut point onwards") + is_split_turn: bool = Field(default=False, description="Whether cut point is mid-turn") + cut_index: int = Field(default=0, description="Index of cut point in original message list") diff --git a/reme/core/schema/memory_search_result.py b/reme/core/schema/memory_search_result.py index d08ca12b..784fe7c7 100644 --- a/reme/core/schema/memory_search_result.py +++ b/reme/core/schema/memory_search_result.py @@ -1,7 +1,5 @@ """Memory search result schema.""" -from typing import Any, Dict - from pydantic import BaseModel, Field from ..enumeration import MemorySource @@ -16,7 +14,7 @@ class MemorySearchResult(BaseModel): score: float = Field(..., description="Relevance score of the search result") snippet: str = Field(..., description="Text snippet from the matched content") source: MemorySource = Field(..., description="Source of the memory data") - metadata: Dict[str, Any] = Field(default_factory=dict, description="Additional metadata") + metadata: dict = Field(default_factory=dict, description="Additional metadata") @property def merge_key(self) -> str: diff --git a/reme/reme.py b/reme/reme.py index f08858ec..4f197dd6 100644 --- a/reme/reme.py +++ b/reme/reme.py @@ -251,7 +251,7 @@ class ReMe(Application): enable_when_to_use=False, enable_multiple=True, ), - UpdateMemoryV1( + AddMemory( enable_thinking_params=enable_thinking_params, enable_memory_target=False, enable_when_to_use=False, diff --git a/reme/reme_fs.py b/reme/reme_fs.py index f5ef42d2..51875ffc 100644 --- a/reme/reme_fs.py +++ b/reme/reme_fs.py @@ -1,7 +1,24 @@ """ReMe File System""" +from pathlib import Path + +from .agent.fs import FsCompactor, FsSummarizer from .config import ReMeConfigParser from .core import Application +from .core.enumeration import MemorySource +from .core.op import BaseTool +from .core.schema import Message +from .tool.fs import ( + BashTool, + EditTool, + FindTool, + FsMemoryGet, + FsMemorySearch, + GrepTool, + LsTool, + ReadTool, + WriteTool, +) class ReMeFs(Application): @@ -17,9 +34,10 @@ class ReMeFs(Application): enable_logo: bool = True, llm: dict | None = None, embedding_model: dict | None = None, - vector_store: dict | None = None, + memory_store: dict | None = None, token_counter: dict | None = None, - working_dir: str = "./agent", + file_watcher: dict | None = None, + working_dir: str = ".reme", **kwargs, ): """Initialize ReMe with config.""" @@ -33,9 +51,91 @@ class ReMeFs(Application): parser=ReMeConfigParser, llm=llm, embedding_model=embedding_model, - vector_store=vector_store, + memory_store=memory_store, token_counter=token_counter, + file_watcher=file_watcher, **kwargs, ) self.working_dir: str = working_dir + self.fs_tools: list[BaseTool] = [ + BashTool(cwd=self.working_dir), + EditTool(cwd=self.working_dir), + FindTool(cwd=self.working_dir), + GrepTool(cwd=self.working_dir), + LsTool(cwd=self.working_dir), + ReadTool(cwd=self.working_dir), + WriteTool(cwd=self.working_dir), + ] + self.working_path: Path = Path(self.working_dir) + self.working_path.mkdir(parents=True, exist_ok=True) + + async def compact( + self, + messages: list[Message | dict], + context_window_tokens: int = 128000, + reserve_tokens: int = 36000, + keep_recent_tokens: int = 20000, + ): + """Compact messages.""" + messages = [Message(**message) if isinstance(message, dict) else message for message in messages] + compactor = FsCompactor( + context_window_tokens=context_window_tokens, + reserve_tokens=reserve_tokens, + keep_recent_tokens=keep_recent_tokens, + ) + + return await compactor.call(messages=messages, service_context=self.service_context) + + async def summary( + self, + messages: list[Message | dict], + date: str, + version: str = "default", + context_window_tokens: int = 128000, + reserve_tokens: int = 32000, + soft_threshold_tokens: int = 4000, + ): + """Summarize messages.""" + messages = [Message(**message) if isinstance(message, dict) else message for message in messages] + summarizer = FsSummarizer( + tools=self.fs_tools, + version=version, + context_window_tokens=context_window_tokens, + reserve_tokens=reserve_tokens, + soft_threshold_tokens=soft_threshold_tokens, + ) + + return await summarizer.call(messages=messages, date=date, service_context=self.service_context) + + async def memory_search( + self, + query: str, + max_results: int = 20, + min_score: float = 0.1, + sources: list[MemorySource] | None = None, + hybrid_enabled: bool = True, + hybrid_vector_weight: float = 0.7, + hybrid_text_weight: float = 0.3, + hybrid_candidate_multiplier: float = 3.0, + ) -> str: + """Semantically search memory files.""" + search_tool = FsMemorySearch( + sources=sources, + hybrid_enabled=hybrid_enabled, + hybrid_vector_weight=hybrid_vector_weight, + hybrid_text_weight=hybrid_text_weight, + hybrid_candidate_multiplier=hybrid_candidate_multiplier, + ) + + return await search_tool.call( + query=query, + max_results=max_results, + min_score=min_score, + service_context=self.service_context, + ) + + async def memory_get(self, path: str, offset: int | None = None, limit: int | None = None) -> str: + """Read specific snippets from memory files.""" + get_tool = FsMemoryGet(workspace_dir=self.working_dir, memory_store=self.memory_store) + return await get_tool.call(path=path, offset=offset, limit=limit, service_context=self.service_context) diff --git a/reme/tool/fs/__init__.py b/reme/tool/fs/__init__.py index 090e1a1f..3a85838d 100644 --- a/reme/tool/fs/__init__.py +++ b/reme/tool/fs/__init__.py @@ -1,8 +1,11 @@ """File system tools.""" +from .base_fs_tool import BaseFsTool from .bash_tool import BashTool from .edit_tool import EditTool from .find_tool import FindTool +from .fs_memory_get import FsMemoryGet +from .fs_memory_search import FsMemorySearch from .grep_tool import GrepTool from .ls_tool import LsTool from .read_tool import ReadTool @@ -10,9 +13,12 @@ from .write_tool import WriteTool from ...core import R __all__ = [ + "BaseFsTool", "BashTool", "EditTool", "FindTool", + "FsMemoryGet", + "FsMemorySearch", "GrepTool", "LsTool", "ReadTool", diff --git a/reme/tool/fs/base_fs_tool.py b/reme/tool/fs/base_fs_tool.py new file mode 100644 index 00000000..92c9092a --- /dev/null +++ b/reme/tool/fs/base_fs_tool.py @@ -0,0 +1,41 @@ +"""Base class for file system tools with unified error handling.""" + +from loguru import logger + +from ...core.context import RuntimeContext +from ...core.op import BaseTool + + +class BaseFsTool(BaseTool): + """Base class for file system tools. + + Features: + - No retry logic (max_retries=1) + - Catches all exceptions and returns error messages to LLM + - Simplifies error handling in subclasses + """ + + def __init__(self, **kwargs): + """Initialize fs tool with no retry.""" + kwargs.setdefault("max_retries", 1) + kwargs.setdefault("raise_exception", False) + super().__init__(**kwargs) + + async def call(self, context: RuntimeContext = None, **kwargs): + """Execute the tool with unified error handling. + + This method catches all exceptions and returns error messages + to the LLM instead of raising them. + """ + self.context = RuntimeContext.from_context(context, **kwargs) + + try: + await self.before_execute() + response = await self.execute() + response = await self.after_execute(response) + return response + except Exception as e: + # Return error message to LLM instead of raising + error_msg = f"{self.__class__.__name__} failed: {str(e)}" + logger.error(error_msg) + return await self.after_execute(error_msg) diff --git a/reme/tool/fs/bash_tool.py b/reme/tool/fs/bash_tool.py index d8aea81c..1b084323 100644 --- a/reme/tool/fs/bash_tool.py +++ b/reme/tool/fs/bash_tool.py @@ -11,8 +11,8 @@ import platform import signal from pathlib import Path +from .base_fs_tool import BaseFsTool from .truncate import DEFAULT_MAX_BYTES, DEFAULT_MAX_LINES, truncate_tail -from ...core.op import BaseTool from ...core.schema import ToolCall, TruncationResult @@ -53,7 +53,7 @@ def kill_process_tree(pid: int) -> None: pass # Best effort -class BashTool(BaseTool): +class BashTool(BaseFsTool): """Production-grade tool for executing bash commands. Features: @@ -118,41 +118,35 @@ class BashTool(BaseTool): shell, shell_args = get_shell_config() # Start process - try: - process = await asyncio.create_subprocess_exec( - shell, - *shell_args, - command, - stdout=asyncio.subprocess.PIPE, - stderr=asyncio.subprocess.PIPE, - cwd=self.cwd, - # Create process group for clean termination - preexec_fn=os.setpgrp if platform.system() != "Windows" else None, - ) - except Exception as e: - raise RuntimeError(f"Failed to start process: {e}") from e + process = await asyncio.create_subprocess_exec( + shell, + *shell_args, + command, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + cwd=self.cwd, + # Create process group for clean termination + preexec_fn=os.setpgrp if platform.system() != "Windows" else None, + ) # Execute command with optional timeout - try: - if timeout and timeout > 0: + if timeout and timeout > 0: + try: + stdout, stderr = await asyncio.wait_for( + process.communicate(), + timeout=timeout, + ) + except asyncio.TimeoutError as e: + # Kill process tree on timeout + if process.pid: + kill_process_tree(process.pid) try: - stdout, stderr = await asyncio.wait_for( - process.communicate(), - timeout=timeout, - ) - except asyncio.TimeoutError as e: - # Kill process tree on timeout - if process.pid: - kill_process_tree(process.pid) - try: - await asyncio.wait_for(process.wait(), timeout=1.0) - except asyncio.TimeoutError: - process.kill() - raise TimeoutError(f"Command timed out after {timeout} seconds") from e - else: - stdout, stderr = await process.communicate() - except TimeoutError as e: - raise RuntimeError(str(e)) from e + await asyncio.wait_for(process.wait(), timeout=1.0) + except asyncio.TimeoutError: + process.kill() + raise TimeoutError(f"Command timed out after {timeout} seconds") from e + else: + stdout, stderr = await process.communicate() # Decode output full_output = stdout.decode("utf-8", errors="ignore") diff --git a/reme/tool/fs/edit_tool.py b/reme/tool/fs/edit_tool.py index fe73bd1f..8733c2af 100644 --- a/reme/tool/fs/edit_tool.py +++ b/reme/tool/fs/edit_tool.py @@ -3,6 +3,7 @@ import os from pathlib import Path +from .base_fs_tool import BaseFsTool from .edit_diff import ( detect_line_ending, fuzzy_find_text, @@ -12,11 +13,10 @@ from .edit_diff import ( restore_line_endings, strip_bom, ) -from ...core.op import BaseTool from ...core.schema import ToolCall -class EditTool(BaseTool): +class EditTool(BaseFsTool): """Edit a file by replacing exact text.""" def __init__(self, cwd: str | None = None): @@ -77,11 +77,8 @@ class EditTool(BaseTool): raise PermissionError(f"File not readable/writable: {path}") # Read file - try: - with open(absolute_path, "r", encoding="utf-8") as f: - raw_content = f.read() - except Exception as e: - raise IOError(f"Failed to read file {path}: {e}") from e + with open(absolute_path, "r", encoding="utf-8") as f: + raw_content = f.read() # Strip BOM (LLM won't include invisible BOM in oldText) bom, content = strip_bom(raw_content) @@ -129,11 +126,8 @@ class EditTool(BaseTool): # Write file final_content = bom + restore_line_endings(new_content, original_ending) - try: - with open(absolute_path, "w", encoding="utf-8") as f: - f.write(final_content) - except Exception as e: - raise IOError(f"Failed to write file {path}: {e}") from e + with open(absolute_path, "w", encoding="utf-8") as f: + f.write(final_content) # Generate diff diff_result = generate_diff_string(base_content, new_content) diff --git a/reme/tool/fs/find_tool.py b/reme/tool/fs/find_tool.py index 7b855586..731722cd 100644 --- a/reme/tool/fs/find_tool.py +++ b/reme/tool/fs/find_tool.py @@ -3,12 +3,12 @@ import os from pathlib import Path +from .base_fs_tool import BaseFsTool from .truncate import FIND_MAX_BYTES, FIND_MAX_LINES, format_size, truncate_head -from ...core.op import BaseTool from ...core.schema import ToolCall -class FindTool(BaseTool): +class FindTool(BaseFsTool): """Search for files by glob pattern, respecting .gitignore.""" def __init__(self, cwd: str | None = None): @@ -131,28 +131,25 @@ class FindTool(BaseTool): # Search for files results = [] - try: - for file_path in search_path.glob(pattern): - if len(results) >= limit: - break + for file_path in search_path.glob(pattern): + if len(results) >= limit: + break - # Skip if matches ignore patterns - if self._should_ignore(file_path, ignore_patterns): - continue + # Skip if matches ignore patterns + if self._should_ignore(file_path, ignore_patterns): + continue - # Get relative path - try: - rel_path = file_path.relative_to(search_path) - # Add trailing slash for directories - if file_path.is_dir(): - results.append(f"{rel_path}/") - else: - results.append(str(rel_path)) - except ValueError: - # If relative_to fails, use the path as-is - results.append(str(file_path)) - except Exception as e: - raise RuntimeError(f"Error searching for files: {e}") from e + # Get relative path + try: + rel_path = file_path.relative_to(search_path) + # Add trailing slash for directories + if file_path.is_dir(): + results.append(f"{rel_path}/") + else: + results.append(str(rel_path)) + except ValueError: + # If relative_to fails, use the path as-is + results.append(str(file_path)) # Handle no results if not results: diff --git a/reme/tool/fs/fs_memory_get.py b/reme/tool/fs/fs_memory_get.py index c872aabf..c4100507 100644 --- a/reme/tool/fs/fs_memory_get.py +++ b/reme/tool/fs/fs_memory_get.py @@ -3,11 +3,11 @@ import os from pathlib import Path -from reme.core.op import BaseTool from reme.core.schema import ToolCall +from .base_fs_tool import BaseFsTool -class FsMemoryGet(BaseTool): +class FsMemoryGet(BaseFsTool): """Read specific snippets from memory files.""" def __init__(self, workspace_dir: str | None = None, **kwargs): @@ -20,7 +20,7 @@ class FsMemoryGet(BaseTool): return ToolCall( **{ "description": ( - "Safe snippet read from MEMORY.md, memory/*.md with optional from/lines; " + "Safe snippet read from MEMORY.md, memory/*.md with optional offset/limit; " "use after memory_search to pull only the needed lines and keep context small." ), "parameters": { @@ -30,11 +30,11 @@ class FsMemoryGet(BaseTool): "type": "string", "description": "Path to the memory file to read (relative or absolute)", }, - "from": { + "offset": { "type": "integer", "description": "Starting line number (1-indexed, optional)", }, - "lines": { + "limit": { "type": "integer", "description": "Number of lines to read from the starting line (optional)", }, @@ -47,8 +47,8 @@ class FsMemoryGet(BaseTool): async def execute(self) -> str: """Execute the memory get operation.""" raw_path: str = self.context.path.strip() - from_param: int | None = self.context.get("from", None) - lines_param: int | None = self.context.get("lines", None) + offset: int | None = self.context.get("offset", None) + limit: int | None = self.context.get("limit", None) if os.path.isabs(raw_path): abs_path = os.path.abspath(raw_path) @@ -65,15 +65,25 @@ class FsMemoryGet(BaseTool): with open(abs_path, "r", encoding="utf-8") as f: content = f.read() - if from_param is None and lines_param is None: + if offset is None and limit is None: return content - else: - lines = content.split("\n") - start = max(1, from_param if from_param is not None else 1) - count = max(1, lines_param if lines_param is not None else len(lines)) + lines = content.split("\n") + total_lines = len(lines) - # Extract slice (1-indexed to 0-indexed conversion) - selected = lines[start - 1 : start - 1 + count] - text = "\n".join(selected) - return text + # Validate and normalize offset (1-indexed) + start = offset if offset is not None else 1 + assert start >= 1, f"offset must be >= 1, got {start}" + assert start <= total_lines, f"offset {start} exceeds total lines {total_lines}" + + # Validate and calculate count + if limit is not None: + assert limit > 0, f"limit must be positive, got {limit}" + count = limit + else: + # Read from start to end of file + count = total_lines - start + 1 + + # Extract slice (1-indexed to 0-indexed conversion) + selected = lines[start - 1 : start - 1 + count] + return "\n".join(selected) diff --git a/reme/tool/fs/fs_memory_search.py b/reme/tool/fs/fs_memory_search.py index af245a99..924be3d2 100644 --- a/reme/tool/fs/fs_memory_search.py +++ b/reme/tool/fs/fs_memory_search.py @@ -3,11 +3,11 @@ import json from reme.core.enumeration import MemorySource -from reme.core.op import BaseTool from reme.core.schema import MemorySearchResult, ToolCall +from .base_fs_tool import BaseFsTool -class FsMemorySearch(BaseTool): +class FsMemorySearch(BaseFsTool): """Semantically search MEMORY.md and memory files.""" def __init__( @@ -48,11 +48,11 @@ class FsMemorySearch(BaseTool): "type": "string", "description": "The semantic search query to find relevant memory snippets", }, - "maxResults": { + "max_results": { "type": "integer", "description": "Maximum number of search results to return (optional)", }, - "minScore": { + "min_score": { "type": "number", "description": "Minimum similarity score threshold for results (optional)", }, @@ -65,8 +65,8 @@ class FsMemorySearch(BaseTool): async def execute(self) -> str: """Execute the memory search operation.""" query: str = self.context.query.strip() - min_score = self.context.get("minScore", self.min_score) - max_results = self.context.get("maxResults", self.max_results) + min_score = self.context.get("min_score", self.min_score) + max_results = self.context.get("max_results", self.max_results) candidates = min(200, max(1, int(max_results * self.hybrid_candidate_multiplier))) # Perform hybrid search (vector + keyword) diff --git a/reme/tool/fs/grep_tool.py b/reme/tool/fs/grep_tool.py index 9df9f8f2..0546a495 100644 --- a/reme/tool/fs/grep_tool.py +++ b/reme/tool/fs/grep_tool.py @@ -13,6 +13,7 @@ import os import shutil from pathlib import Path +from .base_fs_tool import BaseFsTool from .truncate import ( DEFAULT_MAX_BYTES, GREP_MAX_LINE_LENGTH, @@ -20,14 +21,13 @@ from .truncate import ( truncate_head, truncate_line, ) -from ...core.op import BaseTool from ...core.schema import ToolCall # Default limits DEFAULT_LIMIT = 100 # Maximum number of matches -class GrepTool(BaseTool): +class GrepTool(BaseFsTool): """Tool for searching file contents using ripgrep. Features: @@ -145,15 +145,12 @@ class GrepTool(BaseTool): args.extend([pattern, search_path]) # Execute ripgrep - try: - process = await asyncio.create_subprocess_exec( - *args, - stdout=asyncio.subprocess.PIPE, - stderr=asyncio.subprocess.PIPE, - cwd=self.cwd, - ) - except Exception as e: - raise RuntimeError(f"Failed to run ripgrep: {e}") from e + process = await asyncio.create_subprocess_exec( + *args, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + cwd=self.cwd, + ) stdout, stderr = await process.communicate() diff --git a/reme/tool/fs/ls_tool.py b/reme/tool/fs/ls_tool.py index 1c054c62..f231119c 100644 --- a/reme/tool/fs/ls_tool.py +++ b/reme/tool/fs/ls_tool.py @@ -3,14 +3,14 @@ import os from pathlib import Path +from .base_fs_tool import BaseFsTool from .truncate import DEFAULT_MAX_BYTES, truncate_head -from ...core.op import BaseTool from ...core.schema import ToolCall DEFAULT_LIMIT = 500 -class LsTool(BaseTool): +class LsTool(BaseFsTool): """List directory contents with smart truncation. Features: @@ -75,10 +75,7 @@ class LsTool(BaseTool): raise NotADirectoryError(f"Not a directory: {dir_path}") # Read directory entries - try: - entries = list(dir_path.iterdir()) - except Exception as e: - raise PermissionError(f"Cannot read directory: {e}") from e + entries = list(dir_path.iterdir()) # Sort alphabetically (case-insensitive) entries.sort(key=lambda e: e.name.lower()) diff --git a/reme/tool/fs/read_tool.py b/reme/tool/fs/read_tool.py index b236d4a4..ace2339a 100644 --- a/reme/tool/fs/read_tool.py +++ b/reme/tool/fs/read_tool.py @@ -9,8 +9,8 @@ Features: import os from pathlib import Path +from .base_fs_tool import BaseFsTool from .truncate import DEFAULT_MAX_BYTES, DEFAULT_MAX_LINES, format_size, truncate_head -from ...core.op import BaseTool from ...core.schema import ToolCall # Supported image extensions @@ -29,7 +29,7 @@ def is_image_file(path: str) -> bool: return Path(path).suffix.lower() in IMAGE_EXTENSIONS -class ReadTool(BaseTool): +class ReadTool(BaseFsTool): """Read file contents with smart truncation. Features: @@ -163,24 +163,32 @@ class ReadTool(BaseTool): all_lines = content.split("\n") total_file_lines = len(all_lines) - # Apply offset if specified (convert 1-indexed to 0-indexed) - start_line = max(0, (offset - 1)) if offset else 0 + # Validate and apply offset (1-indexed to 0-indexed) + if offset is not None: + if offset < 1: + raise ValueError(f"offset must be >= 1, got {offset}") + start_line = offset - 1 + else: + start_line = 0 + start_line_display = start_line + 1 # Check offset bounds - if start_line >= len(all_lines): + if start_line >= total_file_lines: raise IndexError( - f"Offset {offset} is beyond end of file ({len(all_lines)} lines total)", + f"Offset {offset} is beyond end of file ({total_file_lines} lines total)", ) - # Apply user limit if specified + # Validate and apply limit if limit is not None: - end_line = min(start_line + limit, len(all_lines)) - selected_content = "\n".join(all_lines[start_line:end_line]) - user_limited_lines = end_line - start_line + if limit <= 0: + raise ValueError(f"limit must be positive, got {limit}") + end_line = min(start_line + limit, total_file_lines) else: - selected_content = "\n".join(all_lines[start_line:]) - user_limited_lines = None + end_line = total_file_lines + + # Extract selected lines + selected_content = "\n".join(all_lines[start_line:end_line]) # Apply truncation truncation = truncate_head(selected_content) @@ -205,15 +213,15 @@ class ReadTool(BaseTool): f"of {total_file_lines} ({max_kb}KB limit). " f"Use offset={next_offset} to continue.]" ) - elif user_limited_lines is not None and start_line + user_limited_lines < len(all_lines): - # User limit exceeded, but no truncation - remaining = len(all_lines) - (start_line + user_limited_lines) - next_offset = start_line + user_limited_lines + 1 + elif end_line < total_file_lines: + # User limit reached but no truncation + remaining = total_file_lines - end_line + next_offset = end_line + 1 output_text = truncation.content output_text += f"\n\n[{remaining} more lines in file. " f"Use offset={next_offset} to continue.]" else: - # No truncation or user limit exceeded + # No truncation or limit output_text = truncation.content return output_text diff --git a/reme/tool/fs/write_tool.py b/reme/tool/fs/write_tool.py index 55a4a0c1..a291e0ac 100644 --- a/reme/tool/fs/write_tool.py +++ b/reme/tool/fs/write_tool.py @@ -8,11 +8,11 @@ This module provides a tool for writing content to files with: import os -from ...core.op import BaseTool +from .base_fs_tool import BaseFsTool from ...core.schema import ToolCall -class WriteTool(BaseTool): +class WriteTool(BaseFsTool): """Tool for writing content to files. Features: diff --git a/tests/demo_memory_search.py b/tests/demo_memory_search.py index 104d333a..a1d37091 100644 --- a/tests/demo_memory_search.py +++ b/tests/demo_memory_search.py @@ -24,7 +24,7 @@ from reme.core.embedding import OpenAIEmbeddingModel from reme.core.enumeration import MemorySource from reme.core.file_watcher.delta_file_watcher import DeltaFileWatcher from reme.core.file_watcher.full_file_watcher import FullFileWatcher -from reme.core.memory_storage import SqliteMemoryStore +from reme.core.memory_store import SqliteMemoryStore from reme.core.utils import load_env load_env() diff --git a/tests/test_fs_agent.py b/tests/test_fs_agent.py deleted file mode 100644 index 47995785..00000000 --- a/tests/test_fs_agent.py +++ /dev/null @@ -1,323 +0,0 @@ -"""Tests for fs (full-session) agents including compactor and summarizer. - -This module contains test functions for FsCompactor and FsSummarizer operations. -""" - -import asyncio -import os -import tempfile -from pathlib import Path - -from reme import ReMe -from reme.agent.fs.fs_compactor import FsCompactor -from reme.agent.fs.fs_summarizer import FsSummarizer -from reme.core.enumeration import Role -from reme.core.schema import Message -from reme.tool.fs import ReadTool, WriteTool, EditTool - - -def create_test_messages(num_messages: int = 10) -> list[Message]: - """Create a list of test messages for testing. - - Args: - num_messages: Number of messages to create - - Returns: - List of Message objects alternating between user and assistant - """ - messages = [] - for i in range(num_messages): - if i % 2 == 0: - # User messages - messages.append( - Message( - role=Role.USER, - content=f"User message {i}: Can you help me with task {i}?", - ), - ) - else: - # Assistant messages - messages.append( - Message( - role=Role.ASSISTANT, - content=f"Assistant message {i}: Sure, I'd be happy to help you with task {i - 1}. " - f"Let me explain the solution in detail. " * 10, # Make it longer - ), - ) - return messages - - -def create_long_conversation() -> list[Message]: - """Create a long conversation that exceeds token thresholds.""" - messages = [ - Message( - role=Role.USER, - content="I need help building a complete web application with authentication, database, and API endpoints.", - ), - Message( - role=Role.ASSISTANT, - content="""I'll help you build a complete web application. Here's what we'll do: - -1. Set up the project structure -2. Implement authentication system -3. Design and create database schema -4. Build API endpoints -5. Add frontend components -6. Test and deploy - -Let me start with the project structure...""", - ), - ] - - # Initial user request - - # Assistant response with detailed steps - - # Continue with multiple turns - for i in range(15): - messages.append( - Message( - role=Role.USER, - content=f"What about step {i + 1}? Can you provide more details?", - ), - ) - messages.append( - Message( - role=Role.ASSISTANT, - content=f"""For step {i + 1}, here's a detailed explanation: - -First, we need to consider the architecture. """ - + "This is important context. " * 50 - + """ - -Then we implement the following: -- Component A -- Component B -- Component C - -Let me show you the code for this part...""" - + "\n\ncode_example = 'example'" * 20, - ), - ) - - return messages - - -async def test_compactor_basic(reme: ReMe): - """Test basic FsCompactor functionality without triggering compaction. - - Tests that the compactor correctly skips compaction when token count - is below the threshold. - """ - print("\n" + "=" * 60) - print("Testing FsCompactor - Basic (Below Threshold)") - print("=" * 60) - - # Create a small conversation that won't trigger compaction - messages = create_test_messages(num_messages=6) - - # Create compactor with high threshold so it won't trigger - compactor = FsCompactor( - context_window_tokens=128000, - reserve_tokens=10000, - keep_recent_tokens=5000, - ) - - print(f"Number of messages: {len(messages)}") - output = await compactor.call(messages=messages, service_context=reme.service_context) - print(f"test_compactor_basic output: {output}") - - -async def test_compactor_with_compaction(reme: ReMe): - """Test FsCompactor with a long conversation that triggers compaction. - - Tests that the compactor correctly summarizes old messages when - the conversation exceeds the token threshold. - """ - print("\n" + "=" * 60) - print("Testing FsCompactor - With Compaction") - print("=" * 60) - - # Create a long conversation - messages = create_long_conversation() - - # Create compactor with low threshold to trigger compaction - compactor = FsCompactor( - context_window_tokens=10000, # Low threshold - reserve_tokens=2000, - keep_recent_tokens=2000, - ) - - print(f"Number of messages: {len(messages)}") - output = await compactor.call(messages=messages, service_context=reme.service_context) - print(f"test_compactor_with_compaction output: {output}") - - -async def test_compactor_split_turn(reme: ReMe): - """Test FsCompactor with a split turn scenario. - - Tests the scenario where the cut point falls in the middle of a turn, - requiring special handling to maintain context. - """ - print("\n" + "=" * 60) - print("Testing FsCompactor - Split Turn Detection") - print("=" * 60) - - messages = [] - - # Add some initial conversation - for i in range(5): - messages.append(Message(role=Role.USER, content=f"Question {i}")) - messages.append(Message(role=Role.ASSISTANT, content=f"Answer {i}. " * 30)) - - # Add a very long assistant response that will be split - messages.append(Message(role=Role.USER, content="Please explain this in great detail.")) - messages.append( - Message( - role=Role.ASSISTANT, - content="This is the first part of a very long response. " * 100, - ), - ) - messages.append( - Message( - role=Role.ASSISTANT, - content="This is the continuation of the response. " * 100, - ), - ) - messages.append( - Message( - role=Role.ASSISTANT, - content="And here's the final part with the conclusion. " * 50, - ), - ) - - compactor = FsCompactor( - context_window_tokens=8000, - reserve_tokens=1000, - keep_recent_tokens=2000, - ) - - print(f"Number of messages: {len(messages)}") - output = await compactor.call(messages=messages, service_context=reme.service_context) - print(f"test_compactor_split_turn output: {output}") - - -async def test_summarizer_basic(reme: ReMe): - """Test basic FsSummarizer functionality. - - Tests that the summarizer correctly skips when below threshold - and executes when above threshold. - """ - print("\n" + "=" * 60) - print("Testing FsSummarizer - Basic") - print("=" * 60) - - # Create a temporary directory for memory storage - with tempfile.TemporaryDirectory() as temp_dir: - memory_dir = os.path.join(temp_dir, "memories") - Path(memory_dir).mkdir(parents=True, exist_ok=True) - - # Create a small conversation (below threshold) - messages = create_test_messages(num_messages=4) - - summarizer = FsSummarizer( - tools=[ReadTool(), WriteTool(), EditTool()], - memory_dir=memory_dir, - context_window_tokens=128000, - reserve_tokens=32000, - soft_threshold_tokens=4000, - ) - - print(f"Memory directory: {memory_dir}") - print(f"Number of messages: {len(messages)}") - output = await summarizer.call(messages=messages, service_context=reme.service_context) - print(f"test_summarizer_basic output: {output}") - - -async def test_summarizer_with_execution(reme: ReMe): - """Test FsSummarizer with execution triggered. - - Tests that the summarizer executes when token count is within - the soft threshold range before compaction. - """ - print("\n" + "=" * 60) - print("Testing FsSummarizer - With Execution") - print("=" * 60) - - with tempfile.TemporaryDirectory() as temp_dir: - memory_dir = os.path.join(temp_dir, "memories") - Path(memory_dir).mkdir(parents=True, exist_ok=True) - - # Create messages that will trigger summarizer but not compactor - messages = create_test_messages(num_messages=10) - - # Set low thresholds to trigger execution - summarizer = FsSummarizer( - tools=[ReadTool(), WriteTool(), EditTool()], - memory_dir=memory_dir, - context_window_tokens=5000, - reserve_tokens=1000, - soft_threshold_tokens=500, - ) - - print(f"Memory directory: {memory_dir}") - print(f"Number of messages: {len(messages)}") - output = await summarizer.call(messages=messages, service_context=reme.service_context) - print(f"test_summarizer_with_execution output: {output}") - - -def test_compactor_serialization(): - """Test message serialization in FsCompactor. - - Tests that messages are correctly serialized to text format - for summarization. - """ - print("\n" + "=" * 60) - print("Testing FsCompactor - Message Serialization") - print("=" * 60) - - messages = [ - Message(role=Role.USER, content="Hello, how are you?", name="Alice"), - Message(role=Role.ASSISTANT, content="I'm doing great, thanks!"), - Message(role=Role.USER, content="Can you help me?"), - ] - - # Access static method for testing serialization - serialized = FsCompactor._serialize_conversation(messages) # pylint: disable=protected-access - - print("Serialized conversation:") - print(serialized) - print("\n✓ Serialization completed") - - # Check that it contains expected markers - assert "[Alice]" in serialized - assert "[assistant]" in serialized - assert "Hello, how are you?" in serialized - print("✓ Serialization format is correct") - - -async def main(): - """Run all tests.""" - # Run basic tests first - reme = ReMe() - await reme.start() - test_compactor_serialization() - await test_compactor_basic(reme) - await test_summarizer_basic(reme) - - # Run tests that require LLM calls (commented out by default) - # Uncomment these if you want to test with actual LLM calls - # await test_compactor_with_compaction(reme) - # await test_compactor_split_turn(reme) - # await test_summarizer_with_execution(reme) - - print("\n" + "=" * 60) - print("All basic tests completed!") - print("=" * 60) - print("\nNote: Tests requiring LLM calls are commented out.") - print("Uncomment them in the main() function to run with actual LLM.") - await reme.close() - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/tests/test_fs_compact.py b/tests/test_fs_compact.py new file mode 100644 index 00000000..8824ea4c --- /dev/null +++ b/tests/test_fs_compact.py @@ -0,0 +1,301 @@ +"""Tests for ReMeFs compact interface. + +This module tests the compact() method of ReMeFs class which provides +a high-level interface for conversation compaction. +""" + +import asyncio + +from reme import ReMeFs +from reme.core.enumeration import Role +from reme.core.schema import Message + + +def print_messages(messages: list[Message], title: str = "Messages", max_content_len: int = 150): + """Print messages with their role and content. + + Args: + messages: List of messages to print + title: Title for the message list + max_content_len: Maximum content length to display (truncate if longer) + """ + print(f"\n{title}: (count: {len(messages)})") + print("-" * 80) + for i, msg in enumerate(messages): + content = str(msg.content) + if len(content) > max_content_len: + content = content[:max_content_len] + "..." + print(f" [{i}] {msg.role.value:10s}: {content}") + print("-" * 80) + + +def create_test_messages(num_messages: int = 10) -> list[Message]: + """Create a list of test messages. + + Args: + num_messages: Number of messages to create + + Returns: + List of Message objects alternating between user and assistant + """ + messages = [] + for i in range(num_messages): + if i % 2 == 0: + messages.append( + Message( + role=Role.USER, + content=f"User message {i}: Can you help me with task {i}?", + ), + ) + else: + messages.append( + Message( + role=Role.ASSISTANT, + content=f"Assistant message {i}: Sure, I'd be happy to help you with task {i - 1}. " + f"Let me explain the solution in detail. " * 10, + ), + ) + return messages + + +def create_long_conversation() -> list[Message]: + """Create a long conversation that exceeds token thresholds.""" + messages = [ + Message( + role=Role.USER, + content="I need help building a complete web application with authentication, database, and API endpoints.", + ), + Message( + role=Role.ASSISTANT, + content="""I'll help you build a complete web application. Here's what we'll do: + +1. Set up the project structure +2. Implement authentication system +3. Design and create database schema +4. Build API endpoints +5. Add frontend components +6. Test and deploy + +Let me start with the project structure...""", + ), + ] + + for i in range(15): + messages.append( + Message( + role=Role.USER, + content=f"What about step {i + 1}? Can you provide more details?", + ), + ) + messages.append( + Message( + role=Role.ASSISTANT, + content=f"""For step {i + 1}, here's a detailed explanation: + +First, we need to consider the architecture. """ + + "This is important context. " * 50 + + """ + +Then we implement the following: +- Component A +- Component B +- Component C + +Let me show you the code for this part...""" + + "\n\ncode_example = 'example'" * 20, + ), + ) + + return messages + + +async def test_compact_below_threshold(): + """Test compact() when messages are below threshold. + + Expects: compacted=False, returns original messages + """ + print("\n" + "=" * 80) + print("TEST 1: Compact - Below Threshold (No Compaction)") + print("=" * 80) + + reme_fs = ReMeFs(enable_logo=False, vector_store=None) + await reme_fs.start() + + messages = create_test_messages(num_messages=4) + print_messages(messages, "INPUT MESSAGES", max_content_len=80) + + print("\nParameters:") + print(" context_window_tokens: 5000") + print(" reserve_tokens: 2000 (threshold = 3000)") + print(" keep_recent_tokens: 1000") + + result = await reme_fs.compact( + messages=messages, + context_window_tokens=5000, + reserve_tokens=2000, + keep_recent_tokens=1000, + ) + + print(f"\n{'='*80}") + print("RESULT:") + print(f" compacted: {result.get('compacted')}") + print(f" tokens_before: {result.get('tokens_before')}") + print(f" is_split_turn: {result.get('is_split_turn')}") + + result_messages = result.get("messages", []) + print_messages(result_messages, "OUTPUT MESSAGES", max_content_len=80) + + assert result.get("compacted") is False, "Should not compact below threshold" + assert len(result_messages) == len(messages), "Should return all original messages" + print("\n✓ TEST PASSED: No compaction below threshold\n") + + await reme_fs.close() + + +async def test_compact_above_threshold(): + """Test compact() when messages exceed threshold. + + Expects: compacted=True, returns summary + left_messages + """ + print("\n" + "=" * 80) + print("TEST 2: Compact - Above Threshold (With Compaction & LLM Summary)") + print("=" * 80) + + reme_fs = ReMeFs(enable_logo=False, vector_store=None) + await reme_fs.start() + + messages = create_test_messages(num_messages=12) + print_messages(messages, "INPUT MESSAGES", max_content_len=60) + + print("\nParameters:") + print(" context_window_tokens: 3000") + print(" reserve_tokens: 1500 (threshold = 1500)") + print(" keep_recent_tokens: 500 (keep only recent messages)") + + result = await reme_fs.compact( + messages=messages, + context_window_tokens=3000, + reserve_tokens=1500, + keep_recent_tokens=500, + ) + + print(f"\n{'='*80}") + print("RESULT:") + print(f" compacted: {result.get('compacted')}") + print(f" tokens_before: {result.get('tokens_before')}") + print(f" is_split_turn: {result.get('is_split_turn')}") + + result_messages = result.get("messages", []) + if result.get("compacted") and result_messages: + has_summary = "" in str(result_messages[0].content) + print(f"\n *** First message contains summary: {has_summary}") + + print_messages(result_messages, "OUTPUT MESSAGES (Summary + Recent)", max_content_len=1500) + + assert result.get("compacted") is True, "Should compact above threshold" + assert len(result_messages) < len(messages), "Should reduce message count" + print("\n✓ TEST PASSED: Compaction triggered and summary generated\n") + + await reme_fs.close() + + +async def test_compact_split_turn_scenario(): + """Test compact() with split turn scenario. + + Expects: is_split_turn=True when cut point is mid-turn + """ + print("\n" + "=" * 80) + print("TEST 3: Compact - Split Turn Scenario (Cut in Middle of Assistant Response)") + print("=" * 80) + + reme_fs = ReMeFs(enable_logo=False, vector_store=None) + await reme_fs.start() + + messages = [] + + # Add initial conversation + for i in range(3): + messages.append(Message(role=Role.USER, content=f"Question {i}")) + messages.append(Message(role=Role.ASSISTANT, content=f"Answer {i}. " * 30)) + + # Add a very long multi-part assistant response + messages.append(Message(role=Role.USER, content="Please explain this in great detail.")) + messages.append( + Message( + role=Role.ASSISTANT, + content="This is the first part of a very long response. " * 50, + ), + ) + messages.append( + Message( + role=Role.ASSISTANT, + content="This is the continuation of the response. " * 50, + ), + ) + messages.append( + Message( + role=Role.ASSISTANT, + content="And here's the final part with the conclusion. " * 30, + ), + ) + + print_messages(messages, "INPUT MESSAGES", max_content_len=80) + + print("\nParameters:") + print(" context_window_tokens: 3000") + print(" reserve_tokens: 1000 (threshold = 2000)") + print(" keep_recent_tokens: 800 (should cut in middle of assistant responses)") + + result = await reme_fs.compact( + messages=messages, + context_window_tokens=3000, + reserve_tokens=1000, + keep_recent_tokens=800, + ) + + print(f"\n{'='*80}") + print("RESULT:") + print(f" compacted: {result.get('compacted')}") + print(f" tokens_before: {result.get('tokens_before')}") + print(f" is_split_turn: {result.get('is_split_turn')} *** (should be True)") + + result_messages = result.get("messages", []) + print_messages(result_messages, "OUTPUT MESSAGES (Summary with Turn Context + Recent)", max_content_len=150) + + if result.get("is_split_turn"): + print("\n✓ TEST PASSED: Split turn correctly detected and handled\n") + else: + print("\n⚠ WARNING: Split turn not detected (parameters may need adjustment)\n") + + await reme_fs.close() + + +async def main(): + """Run core compact interface tests.""" + print("\n" + "=" * 80) + print("ReMeFs Compact Interface - Core Test Suite") + print("=" * 80) + print("\nThis test suite demonstrates the three key scenarios of conversation compaction:") + print(" 1. Below threshold - no compaction needed") + print(" 2. Above threshold - full compaction with LLM summary") + print(" 3. Split turn - cut point falls in middle of assistant response") + print("=" * 80) + + # Test 1: No compaction (below threshold) + await test_compact_below_threshold() + + # Test 2: Full compaction (requires LLM) + await test_compact_above_threshold() + + # Test 3: Split turn compaction (requires LLM) + await test_compact_split_turn_scenario() + + print("\n" + "=" * 80) + print("All basic tests completed!") + print("=" * 80) + print("\nNote: Tests requiring LLM calls are commented out.") + print("Uncomment them in the main() function to run with actual LLM.") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/test_fs_memory_get.py b/tests/test_fs_memory_get.py new file mode 100644 index 00000000..246a695e --- /dev/null +++ b/tests/test_fs_memory_get.py @@ -0,0 +1,369 @@ +"""Tests for ReMeFs memory_get interface. + +This module tests the memory_get() method of ReMeFs class which provides +a high-level interface for reading specific snippets from memory files. + +The memory_get function should enable the LLM to: +1. Read entire memory files (MEMORY.md, memory/*.md) +2. Read specific line ranges using offset and limit parameters +3. Extract only the needed content to keep context small +""" + +import asyncio +import os +from pathlib import Path + +from reme import ReMeFs + + +def print_result(content: str, title: str = "RESULT", max_len: int = 300): + """Print the result of memory_get() call. + + Args: + content: Content returned from memory_get() + title: Title for the result section + max_len: Maximum content length to display (truncate if longer) + """ + print(f"\n{'=' * 80}") + print(f"{title}:") + + lines = content.split("\n") + print(f" total_lines: {len(lines)}") + print(f" total_chars: {len(content)}") + + if len(content) > max_len: + preview = content[:max_len] + "..." + else: + preview = content + + print("\n Content Preview:") + print("-" * 80) + print(preview) + print("-" * 80) + print(f"{'=' * 80}") + + +def create_test_memory_file(workspace_dir: str) -> str: + """Create a test memory file with numbered lines. + + Args: + workspace_dir: Directory to create the test file in + + Returns: + Path to the created test file (relative to workspace_dir) + """ + workspace_path = Path(workspace_dir) + workspace_path.mkdir(parents=True, exist_ok=True) + + memory_dir = workspace_path / "memory" + memory_dir.mkdir(parents=True, exist_ok=True) + + test_file = memory_dir / "test_profile.md" + + content = """# User Profile + +## Personal Information +Name: Alice Johnson +Age: 28 +Location: San Francisco, CA + +## Professional Background +Occupation: Software Engineer +Company: Tech Innovations Inc. +Years of Experience: 5 + +## Skills +- Python Programming +- Machine Learning +- Natural Language Processing +- Docker & Kubernetes + +## Interests +- Reading sci-fi novels +- Hiking in national parks +- Photography +- Cooking international cuisines + +## Preferences +- Prefers detailed technical explanations +- Likes to see code examples +- Values efficiency and clean code +- Appreciates constructive feedback +""" + + test_file.write_text(content, encoding="utf-8") + return "memory/test_profile.md" + + +async def test_memory_get_full_file(): + """Test memory_get() reads entire file without offset/limit. + + Expects: Returns complete file content + """ + print("\n" + "=" * 80) + print("TEST 1: Memory Get - Read Entire File") + print("=" * 80) + + workspace_dir = ".reme_test_get" + reme_fs = ReMeFs(enable_logo=False, working_dir=workspace_dir) + await reme_fs.start() + + # Create test file + test_file_path = create_test_memory_file(workspace_dir) + print("\nTest Setup:") + print(f" test_file: {test_file_path}") + print(f" workspace: {workspace_dir}") + + print("\nParameters:") + print(f" path: {test_file_path}") + print(" offset: None") + print(" limit: None") + print(" Expected: Read entire file content") + + # Call memory_get + result = await reme_fs.memory_get(path=test_file_path) + + print_result(result, "MEMORY_GET RESULT", max_len=500) + + # Verify + lines = result.split("\n") + print("\nVerification:") + print(f" ✓ Got {len(lines)} lines") + print(f" ✓ Content starts with: {result[:50].strip()}") + + await reme_fs.close() + + +async def test_memory_get_with_offset(): + """Test memory_get() reads from specific line to end. + + Expects: Returns content from line 10 to end of file + """ + print("\n" + "=" * 80) + print("TEST 2: Memory Get - Read with Offset") + print("=" * 80) + + workspace_dir = ".reme_test_get" + reme_fs = ReMeFs(enable_logo=False, working_dir=workspace_dir) + await reme_fs.start() + + test_file_path = "memory/test_profile.md" + offset = 10 + + print("\nParameters:") + print(f" path: {test_file_path}") + print(f" offset: {offset}") + print(" limit: None") + print(f" Expected: Read from line {offset} to end of file") + + # Call memory_get + result = await reme_fs.memory_get(path=test_file_path, offset=offset) + + print_result(result, "MEMORY_GET RESULT", max_len=400) + + # Verify + lines = result.split("\n") + print("\nVerification:") + print(f" ✓ Got {len(lines)} lines starting from line {offset}") + print(f" ✓ First line of result: {lines[0]}") + + await reme_fs.close() + + +async def test_memory_get_with_offset_and_limit(): + """Test memory_get() reads specific line range. + + Expects: Returns exactly 5 lines starting from line 5 + """ + print("\n" + "=" * 80) + print("TEST 3: Memory Get - Read with Offset and Limit") + print("=" * 80) + + workspace_dir = ".reme_test_get" + reme_fs = ReMeFs(enable_logo=False, working_dir=workspace_dir) + await reme_fs.start() + + test_file_path = "memory/test_profile.md" + offset = 5 + limit = 5 + + print("\nParameters:") + print(f" path: {test_file_path}") + print(f" offset: {offset}") + print(f" limit: {limit}") + print(f" Expected: Read exactly {limit} lines starting from line {offset}") + + # Call memory_get + result = await reme_fs.memory_get(path=test_file_path, offset=offset, limit=limit) + + print_result(result, "MEMORY_GET RESULT", max_len=400) + + # Verify + lines = result.split("\n") + print("\nVerification:") + print(f" ✓ Got {len(lines)} lines (expected {limit})") + print(" ✓ Lines content:") + for i, line in enumerate(lines, start=offset): + print(f" Line {i}: {line}") + + await reme_fs.close() + + +async def test_memory_get_beginning_lines(): + """Test memory_get() reads first few lines. + + Expects: Returns first 3 lines of the file + """ + print("\n" + "=" * 80) + print("TEST 4: Memory Get - Read Beginning Lines") + print("=" * 80) + + workspace_dir = ".reme_test_get" + reme_fs = ReMeFs(enable_logo=False, working_dir=workspace_dir) + await reme_fs.start() + + test_file_path = "memory/test_profile.md" + offset = 1 + limit = 3 + + print("\nParameters:") + print(f" path: {test_file_path}") + print(f" offset: {offset}") + print(f" limit: {limit}") + print(f" Expected: Read first {limit} lines") + + # Call memory_get + result = await reme_fs.memory_get(path=test_file_path, offset=offset, limit=limit) + + print_result(result, "MEMORY_GET RESULT", max_len=400) + + # Verify + lines = result.split("\n") + print("\nVerification:") + print(f" ✓ Got {len(lines)} lines (expected {limit})") + print(" ✓ Should contain '# User Profile' header") + assert "# User Profile" in result, "Expected header not found" + + await reme_fs.close() + + +async def test_memory_get_single_line(): + """Test memory_get() reads a single specific line. + + Expects: Returns exactly 1 line + """ + print("\n" + "=" * 80) + print("TEST 5: Memory Get - Read Single Line") + print("=" * 80) + + workspace_dir = ".reme_test_get" + reme_fs = ReMeFs(enable_logo=False, working_dir=workspace_dir) + await reme_fs.start() + + test_file_path = "memory/test_profile.md" + offset = 3 + limit = 1 + + print("\nParameters:") + print(f" path: {test_file_path}") + print(f" offset: {offset}") + print(f" limit: {limit}") + print(f" Expected: Read exactly 1 line at position {offset}") + + # Call memory_get + result = await reme_fs.memory_get(path=test_file_path, offset=offset, limit=limit) + + print_result(result, "MEMORY_GET RESULT", max_len=400) + + # Verify + lines = result.split("\n") + print("\nVerification:") + print(f" ✓ Got {len(lines)} line (expected {limit})") + print(f" ✓ Line {offset}: {result}") + + await reme_fs.close() + + +async def test_memory_get_with_absolute_path(): + """Test memory_get() with absolute path. + + Expects: Works with both relative and absolute paths + """ + print("\n" + "=" * 80) + print("TEST 6: Memory Get - Read with Absolute Path") + print("=" * 80) + + workspace_dir = ".reme_test_get" + reme_fs = ReMeFs(enable_logo=False, working_dir=workspace_dir) + await reme_fs.start() + + # Get absolute path + abs_path = os.path.abspath(os.path.join(workspace_dir, "memory/test_profile.md")) + limit = 5 + + print("\nParameters:") + print(f" path: {abs_path}") + print(" offset: None") + print(f" limit: {limit}") + print(f" Expected: Read first {limit} lines using absolute path") + + # Call memory_get with absolute path + result = await reme_fs.memory_get(path=abs_path, limit=limit) + + print_result(result, "MEMORY_GET RESULT", max_len=400) + + # Verify + lines = result.split("\n") + print("\nVerification:") + print(f" ✓ Got {len(lines)} lines using absolute path") + print(" ✓ Absolute path handling works correctly") + + await reme_fs.close() + + +async def main(): + """Run core memory_get interface tests.""" + print("\n" + "=" * 80) + print("ReMeFs Memory Get Interface - Tests") + print("=" * 80) + print("\nThis test suite validates that the memory_get() function:") + print(" 1. Reads entire memory files without parameters") + print(" 2. Reads from specific line (offset) to end of file") + print(" 3. Reads specific line ranges (offset + limit)") + print(" 4. Handles both relative and absolute paths") + print("\nTest Scenarios:") + print(" 1. Full file read (no offset/limit)") + print(" 2. Read with offset (from line N to end)") + print(" 3. Read with offset and limit (specific range)") + print(" 4. Read beginning lines (first N lines)") + print(" 5. Read single line") + print(" 6. Read with absolute path") + print("=" * 80) + + # Test 1: Read entire file + await test_memory_get_full_file() + + # Test 2: Read with offset + await test_memory_get_with_offset() + + # Test 3: Read with offset and limit + await test_memory_get_with_offset_and_limit() + + # Test 4: Read beginning lines + await test_memory_get_beginning_lines() + + # Test 5: Read single line + await test_memory_get_single_line() + + # Test 6: Read with absolute path + await test_memory_get_with_absolute_path() + + print("\n" + "=" * 80) + print("All memory_get tests completed!") + print("=" * 80) + print("\nNote: Test files are created in .reme_test_get/ directory") + print("You can manually inspect them or delete the directory after testing.") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/test_fs_memory_search.py b/tests/test_fs_memory_search.py new file mode 100644 index 00000000..88e6da8b --- /dev/null +++ b/tests/test_fs_memory_search.py @@ -0,0 +1,704 @@ +"""Tests for ReMeFs memory_search interface. + +This module tests the memory_search() method of ReMeFs class which provides +a high-level interface for searching personal information stored in memory files. + +The memory_search function should enable: +1. Vector similarity search across memory chunks +2. Keyword/FTS (full-text search) if enabled +3. Hybrid search combining vector and keyword results +4. Source filtering (MEMORY, SESSIONS, etc.) +5. Score-based filtering and result limiting +""" + +import asyncio +import hashlib +import shutil +from pathlib import Path + +from reme import ReMeFs +from reme.core.enumeration import MemorySource +from reme.core.schema import FileMetadata, MemoryChunk + + +# ==================== Test Configuration ==================== + + +class TestConfig: + """Test configuration settings.""" + + WORKING_DIR = ".reme_test_search" + + +# ==================== Sample Data Generator ==================== + + +class SampleDataGenerator: + """Generator for sample test data.""" + + @staticmethod + def create_personal_info_chunks(test_name: str = "") -> list[MemoryChunk]: + """Create sample chunks with personal information.""" + base_path = "memory/personal_info.md" + prefix = f"{test_name}_" if test_name else "" + return [ + MemoryChunk( + id=f"{prefix}personal_1", + path=base_path, + source=MemorySource.MEMORY, + start_line=1, + end_line=3, + text="My name is Alice Chen. I am a software engineer working at TechCorp.", + hash=hashlib.md5(b"personal_1").hexdigest(), + embedding=None, + metadata={"category": "personal", "type": "basic_info"}, + ), + MemoryChunk( + id=f"{prefix}personal_2", + path=base_path, + source=MemorySource.MEMORY, + start_line=4, + end_line=6, + text=( + "I love Python programming and machine learning. " + "My favorite frameworks are PyTorch and scikit-learn." + ), + hash=hashlib.md5(b"personal_2").hexdigest(), + embedding=None, + metadata={"category": "personal", "type": "interests"}, + ), + MemoryChunk( + id=f"{prefix}personal_3", + path=base_path, + source=MemorySource.MEMORY, + start_line=7, + end_line=9, + text="In my free time, I enjoy reading science fiction novels and hiking in the mountains.", + hash=hashlib.md5(b"personal_3").hexdigest(), + embedding=None, + metadata={"category": "personal", "type": "hobbies"}, + ), + ] + + @staticmethod + def create_technical_chunks(test_name: str = "") -> list[MemoryChunk]: + """Create sample chunks with technical information.""" + base_path = "memory/technical_notes.md" + prefix = f"{test_name}_" if test_name else "" + return [ + MemoryChunk( + id=f"{prefix}tech_1", + path=base_path, + source=MemorySource.MEMORY, + start_line=1, + end_line=3, + text=( + "Artificial intelligence is transforming software development " + "with automated code generation and testing." + ), + hash=hashlib.md5(b"tech_1").hexdigest(), + embedding=None, + metadata={"category": "tech", "topic": "AI"}, + ), + MemoryChunk( + id=f"{prefix}tech_2", + path=base_path, + source=MemorySource.MEMORY, + start_line=4, + end_line=6, + text=( + "Machine learning models require careful tuning of hyperparameters " + "to achieve optimal performance." + ), + hash=hashlib.md5(b"tech_2").hexdigest(), + embedding=None, + metadata={"category": "tech", "topic": "ML"}, + ), + MemoryChunk( + id=f"{prefix}tech_3", + path=base_path, + source=MemorySource.MEMORY, + start_line=7, + end_line=9, + text="Deep learning neural networks excel at image recognition and natural language processing tasks.", + hash=hashlib.md5(b"tech_3").hexdigest(), + embedding=None, + metadata={"category": "tech", "topic": "DL"}, + ), + ] + + @staticmethod + def create_session_chunks(test_name: str = "") -> list[MemoryChunk]: + """Create sample session chunks.""" + base_path = "sessions/2024-01-15.jsonl" + prefix = f"{test_name}_" if test_name else "" + return [ + MemoryChunk( + id=f"{prefix}session_1", + path=base_path, + source=MemorySource.SESSIONS, + start_line=1, + end_line=2, + text="User asked about Python best practices for async programming.", + hash=hashlib.md5(b"session_1").hexdigest(), + embedding=None, + metadata={"session_id": "sess_001", "date": "2024-01-15"}, + ), + MemoryChunk( + id=f"{prefix}session_2", + path=base_path, + source=MemorySource.SESSIONS, + start_line=3, + end_line=4, + text="Discussed asyncio event loop and common pitfalls in concurrent Python code.", + hash=hashlib.md5(b"session_2").hexdigest(), + embedding=None, + metadata={"session_id": "sess_001", "date": "2024-01-15"}, + ), + ] + + @staticmethod + def create_file_metadata(path: str, chunk_count: int) -> FileMetadata: + """Create file metadata.""" + return FileMetadata( + path=path, + hash=hashlib.md5(path.encode()).hexdigest(), + mtime_ms=1704067200000, # 2024-01-01 00:00:00 + size=1000, + chunk_count=chunk_count, + ) + + +# ==================== Helper Functions ==================== + + +def print_search_results(results: list[dict], query: str, title: str = "SEARCH RESULTS"): + """Pretty print search results.""" + print(f"\n{'=' * 80}") + print(f"{title}") + print(f"Query: '{query}'") + print(f"Found {len(results)} results") + print(f"{'=' * 80}") + + for i, result in enumerate(results, 1): + print(f"\n[{i}] Source: {result.get('source', 'N/A')}") + print(f" Path: {result.get('path', 'N/A')}") + print(f" Lines: {result.get('start_line', 'N/A')}-{result.get('end_line', 'N/A')}") + print(f" Score: {result.get('score', 0):.6f}") + snippet = result.get("snippet", result.get("text", "")) + if len(snippet) > 100: + snippet = snippet[:100] + "..." + print(f" Snippet: {snippet}") + if result.get("metadata"): + print(f" Metadata: {result.get('metadata')}") + + print(f"{'=' * 80}\n") + + +# ==================== Test Functions ==================== + + +async def test_memory_search_basic(): + """Test basic memory search functionality. + + Insert sample data and perform a simple search query. + """ + print("\n" + "=" * 80) + print("TEST 1: Basic Memory Search") + print("=" * 80) + + # Initialize ReMeFs with unique store name + reme_fs = ReMeFs( + enable_logo=False, + working_dir=TestConfig.WORKING_DIR, + memory_store={ + "backend": "sqlite", + "store_name": "test_basic", + "embedding_model": "default", + "fts_enabled": True, + "snippet_max_chars": 700, + }, + ) + await reme_fs.start() + + # Insert personal info chunks + personal_chunks = SampleDataGenerator.create_personal_info_chunks("test_basic") + personal_chunks = await reme_fs.memory_store.get_chunk_embeddings(personal_chunks) + + file_meta = SampleDataGenerator.create_file_metadata( + "memory/personal_info.md", + len(personal_chunks), + ) + await reme_fs.memory_store.upsert_file( + file_meta, + MemorySource.MEMORY, + personal_chunks, + ) + print(f"✓ Inserted {len(personal_chunks)} personal info chunks") + + # Perform search + query = "What programming languages does the user like?" + print(f"\nSearching for: '{query}'") + + result_json = await reme_fs.memory_search( + query=query, + max_results=5, + min_score=0.0, + ) + + # Parse results + import json + + results = json.loads(result_json) + print_search_results(results, query, "BASIC SEARCH RESULTS") + + # Verify results + assert len(results) > 0, "Should find at least one result" + assert results[0]["score"] > 0, "Top result should have positive score" + print("✓ Basic memory search test passed") + + await reme_fs.close() + + +async def test_memory_search_technical_content(): + """Test memory search with technical content. + + Insert technical chunks and search for ML/AI related queries. + """ + print("\n" + "=" * 80) + print("TEST 2: Technical Content Search") + print("=" * 80) + + # Initialize ReMeFs with unique store name + reme_fs = ReMeFs( + enable_logo=False, + working_dir=TestConfig.WORKING_DIR, + memory_store={ + "backend": "sqlite", + "store_name": "test_technical", + "embedding_model": "default", + "fts_enabled": True, + "snippet_max_chars": 700, + }, + ) + await reme_fs.start() + + # Insert technical chunks + tech_chunks = SampleDataGenerator.create_technical_chunks("test_technical") + tech_chunks = await reme_fs.memory_store.get_chunk_embeddings(tech_chunks) + + file_meta = SampleDataGenerator.create_file_metadata( + "memory/technical_notes.md", + len(tech_chunks), + ) + await reme_fs.memory_store.upsert_file( + file_meta, + MemorySource.MEMORY, + tech_chunks, + ) + print(f"✓ Inserted {len(tech_chunks)} technical chunks") + + # Test multiple queries + queries = [ + "artificial intelligence and machine learning", + "neural networks for image processing", + "hyperparameter tuning in ML models", + ] + + for query in queries: + print(f"\n--- Searching for: '{query}' ---") + result_json = await reme_fs.memory_search( + query=query, + max_results=3, + min_score=0.0, + ) + + import json + + results = json.loads(result_json) + print(f"Found {len(results)} results") + + for i, result in enumerate(results, 1): + print(f" [{i}] Score: {result['score']:.6f} | {result['path']}") + + assert len(results) > 0, f"Should find results for query: {query}" + + print("\n✓ Technical content search test passed") + await reme_fs.close() + + +async def test_memory_search_with_source_filter(): + """Test memory search with source filtering. + + Insert data from different sources and test source-specific searches. + """ + print("\n" + "=" * 80) + print("TEST 3: Memory Search with Source Filter") + print("=" * 80) + + # Initialize ReMeFs with unique store name + reme_fs = ReMeFs( + enable_logo=False, + working_dir=TestConfig.WORKING_DIR, + memory_store={ + "backend": "sqlite", + "store_name": "test_source_filter", + "embedding_model": "default", + "fts_enabled": True, + "snippet_max_chars": 700, + }, + ) + await reme_fs.start() + + # Insert MEMORY source data + personal_chunks = SampleDataGenerator.create_personal_info_chunks("test_source") + personal_chunks = await reme_fs.memory_store.get_chunk_embeddings(personal_chunks) + personal_meta = SampleDataGenerator.create_file_metadata( + "memory/personal_info.md", + len(personal_chunks), + ) + await reme_fs.memory_store.upsert_file( + personal_meta, + MemorySource.MEMORY, + personal_chunks, + ) + print(f"✓ Inserted {len(personal_chunks)} MEMORY chunks") + + # Insert SESSIONS source data + session_chunks = SampleDataGenerator.create_session_chunks("test_source") + session_chunks = await reme_fs.memory_store.get_chunk_embeddings(session_chunks) + session_meta = SampleDataGenerator.create_file_metadata( + "sessions/2024-01-15.jsonl", + len(session_chunks), + ) + await reme_fs.memory_store.upsert_file( + session_meta, + MemorySource.SESSIONS, + session_chunks, + ) + print(f"✓ Inserted {len(session_chunks)} SESSIONS chunks") + + query = "Python programming and async" + + # Search only MEMORY source + print(f"\n--- Searching MEMORY source for: '{query}' ---") + result_json = await reme_fs.memory_search( + query=query, + max_results=5, + sources=[MemorySource.MEMORY], + ) + import json + + memory_results = json.loads(result_json) + print(f"Found {len(memory_results)} results in MEMORY source") + for result in memory_results: + assert result["source"] == MemorySource.MEMORY.value, "Should only return MEMORY source results" + + # Search only SESSIONS source + print(f"\n--- Searching SESSIONS source for: '{query}' ---") + result_json = await reme_fs.memory_search( + query=query, + max_results=5, + sources=[MemorySource.SESSIONS], + ) + session_results = json.loads(result_json) + print(f"Found {len(session_results)} results in SESSIONS source") + for result in session_results: + assert result["source"] == MemorySource.SESSIONS.value, "Should only return SESSIONS source results" + + # Search all sources + print(f"\n--- Searching ALL sources for: '{query}' ---") + result_json = await reme_fs.memory_search( + query=query, + max_results=10, + ) + all_results = json.loads(result_json) + print(f"Found {len(all_results)} results across all sources") + + sources_found = {result["source"] for result in all_results} + print(f"Sources found: {sources_found}") + + print("\n✓ Source filter search test passed") + await reme_fs.close() + + +async def test_memory_search_score_filtering(): + """Test memory search with score threshold. + + Test min_score parameter to filter low-relevance results. + """ + print("\n" + "=" * 80) + print("TEST 4: Memory Search with Score Filtering") + print("=" * 80) + + # Initialize ReMeFs with unique store name + reme_fs = ReMeFs( + enable_logo=False, + working_dir=TestConfig.WORKING_DIR, + memory_store={ + "backend": "sqlite", + "store_name": "test_score_filter", + "embedding_model": "default", + "fts_enabled": True, + "snippet_max_chars": 700, + }, + ) + await reme_fs.start() + + # Insert test data + chunks = SampleDataGenerator.create_technical_chunks("test_score") + chunks = await reme_fs.memory_store.get_chunk_embeddings(chunks) + file_meta = SampleDataGenerator.create_file_metadata( + "memory/technical_notes.md", + len(chunks), + ) + await reme_fs.memory_store.upsert_file( + file_meta, + MemorySource.MEMORY, + chunks, + ) + print(f"✓ Inserted {len(chunks)} test chunks") + + query = "machine learning algorithms" + + # Search with different min_score thresholds + thresholds = [0.0, 0.1, 0.3, 0.5] + + for min_score in thresholds: + print(f"\n--- Searching with min_score={min_score} ---") + result_json = await reme_fs.memory_search( + query=query, + max_results=10, + min_score=min_score, + ) + + import json + + results = json.loads(result_json) + print(f"Found {len(results)} results with min_score >= {min_score}") + + # Verify all results meet threshold + for result in results: + assert result["score"] >= min_score, f"Result score {result['score']:.6f} should be >= {min_score}" + + if results: + print(f" Top score: {results[0]['score']:.6f}") + print(f" Lowest score: {results[-1]['score']:.6f}") + + print("\n✓ Score filtering search test passed") + await reme_fs.close() + + +async def test_memory_search_max_results(): + """Test memory search with result limiting. + + Test max_results parameter to limit number of returned results. + """ + print("\n" + "=" * 80) + print("TEST 5: Memory Search with Result Limiting") + print("=" * 80) + + # Initialize ReMeFs with unique store name + reme_fs = ReMeFs( + enable_logo=False, + working_dir=TestConfig.WORKING_DIR, + memory_store={ + "backend": "sqlite", + "store_name": "test_max_results", + "embedding_model": "default", + "fts_enabled": True, + "snippet_max_chars": 700, + }, + ) + await reme_fs.start() + + # Insert multiple chunks + personal_chunks = SampleDataGenerator.create_personal_info_chunks("test_max") + tech_chunks = SampleDataGenerator.create_technical_chunks("test_max") + all_chunks = personal_chunks + tech_chunks + + all_chunks = await reme_fs.memory_store.get_chunk_embeddings(all_chunks) + + # Insert as one file for simplicity + combined_meta = SampleDataGenerator.create_file_metadata( + "memory/combined.md", + len(all_chunks), + ) + await reme_fs.memory_store.upsert_file( + combined_meta, + MemorySource.MEMORY, + all_chunks, + ) + + print(f"✓ Inserted {len(all_chunks)} total chunks") + + query = "programming and technology" + + # Test different max_results values + result_limits = [1, 2, 3, 5, 20] + + for max_results in result_limits: + print(f"\n--- Searching with max_results={max_results} ---") + result_json = await reme_fs.memory_search( + query=query, + max_results=max_results, + min_score=0.0, + ) + + import json + + results = json.loads(result_json) + print(f"Requested {max_results}, got {len(results)} results") + + assert len(results) <= max_results, f"Should return at most {max_results} results, got {len(results)}" + + print("\n✓ Result limiting search test passed") + await reme_fs.close() + + +async def test_memory_search_hybrid_mode(): + """Test memory search with hybrid mode (vector + keyword). + + Test different hybrid configurations and weights. + """ + print("\n" + "=" * 80) + print("TEST 6: Memory Search with Hybrid Mode") + print("=" * 80) + + # Initialize ReMeFs with unique store name + reme_fs = ReMeFs( + enable_logo=False, + working_dir=TestConfig.WORKING_DIR, + memory_store={ + "backend": "sqlite", + "store_name": "test_hybrid", + "embedding_model": "default", + "fts_enabled": True, + "snippet_max_chars": 700, + }, + ) + await reme_fs.start() + + # Insert test data + chunks = SampleDataGenerator.create_technical_chunks("test_hybrid") + chunks = await reme_fs.memory_store.get_chunk_embeddings(chunks) + file_meta = SampleDataGenerator.create_file_metadata( + "memory/technical_notes.md", + len(chunks), + ) + await reme_fs.memory_store.upsert_file( + file_meta, + MemorySource.MEMORY, + chunks, + ) + print(f"✓ Inserted {len(chunks)} test chunks") + + query = "neural networks" + + # Test with hybrid enabled + print(f"\n--- Hybrid search (enabled) for: '{query}' ---") + result_json_hybrid = await reme_fs.memory_search( + query=query, + max_results=5, + hybrid_enabled=True, + hybrid_vector_weight=0.7, + hybrid_text_weight=0.3, + ) + + import json + + hybrid_results = json.loads(result_json_hybrid) + print(f"Hybrid search found {len(hybrid_results)} results") + print_search_results(hybrid_results, query, "HYBRID SEARCH RESULTS") + + # Test with hybrid disabled (vector only) + print(f"\n--- Vector-only search for: '{query}' ---") + result_json_vector = await reme_fs.memory_search( + query=query, + max_results=5, + hybrid_enabled=False, + ) + + vector_results = json.loads(result_json_vector) + print(f"Vector search found {len(vector_results)} results") + print_search_results(vector_results, query, "VECTOR-ONLY SEARCH RESULTS") + + # Test different weight configurations + print("\n--- Testing different hybrid weights ---") + weight_configs = [ + (0.9, 0.1), # Mostly vector + (0.5, 0.5), # Balanced + (0.3, 0.7), # Mostly text + ] + + for vec_weight, text_weight in weight_configs: + result_json = await reme_fs.memory_search( + query=query, + max_results=5, + hybrid_enabled=True, + hybrid_vector_weight=vec_weight, + hybrid_text_weight=text_weight, + ) + results = json.loads(result_json) + print(f" Vector:{vec_weight}/Text:{text_weight} -> {len(results)} results") + + print("\n✓ Hybrid mode search test passed") + await reme_fs.close() + + +async def cleanup_test_data(): + """Clean up test data directory.""" + print("\n" + "=" * 80) + print("CLEANUP: Removing test data") + print("=" * 80) + + test_dir = Path(TestConfig.WORKING_DIR) + if test_dir.exists(): + shutil.rmtree(test_dir) + print(f"✓ Removed test directory: {test_dir}") + else: + print(f"⊘ Test directory does not exist: {test_dir}") + + +# ==================== Main Entry Point ==================== + + +async def main(): + """Run all memory search tests.""" + print("\n" + "=" * 80) + print("ReMeFs Memory Search Interface Tests") + print("=" * 80) + print("\nThis test suite validates the memory_search() function:") + print(" 1. Basic semantic search functionality") + print(" 2. Technical content search") + print(" 3. Source filtering (MEMORY, SESSIONS)") + print(" 4. Score threshold filtering") + print(" 5. Result limiting (max_results)") + print(" 6. Hybrid mode (vector + keyword search)") + print("=" * 80) + + try: + # Run tests + await test_memory_search_basic() + await test_memory_search_technical_content() + await test_memory_search_with_source_filter() + await test_memory_search_score_filtering() + await test_memory_search_max_results() + await test_memory_search_hybrid_mode() + + print("\n" + "=" * 80) + print("✓ All memory search tests passed!") + print("=" * 80) + + finally: + # Cleanup + await cleanup_test_data() + + print("\nNote: These tests require:") + print(" - Valid API keys for embedding model") + print(" - sqlite-vec extension for vector search") + print(" - FTS5 enabled for keyword search") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/test_fs_summary.py b/tests/test_fs_summary.py new file mode 100644 index 00000000..ff9ec210 --- /dev/null +++ b/tests/test_fs_summary.py @@ -0,0 +1,234 @@ +"""Tests for ReMeFs summary interface. + +This module tests the summary() method of ReMeFs class which provides +a high-level interface for storing user's personal information into memory files. + +The summary function should enable the LLM to: +1. Extract personal information from user messages (name, preferences, requirements) +2. Call file system tools (WriteTool, EditTool) to store this information +3. Maintain personalized memory for future conversations +""" + +import asyncio + +from reme import ReMeFs +from reme.core.enumeration import Role +from reme.core.schema import Message + + +def print_messages(messages: list[Message], title: str = "Messages", max_content_len: int = 150): + """Print messages with their role and content. + + Args: + messages: List of messages to print + title: Title for the message list + max_content_len: Maximum content length to display (truncate if longer) + """ + print(f"\n{title}: (count: {len(messages)})") + print("-" * 80) + for i, msg in enumerate(messages): + content = str(msg.content) + if len(content) > max_content_len: + content = content[:max_content_len] + "..." + print(f" [{i}] {msg.role.value:10s}: {content}") + print("-" * 80) + + +def print_result(result: dict, title: str = "RESULT"): + """Print the result of summary() call. + + Args: + result: Result dictionary from summary() + title: Title for the result section + """ + print(f"\n{'=' * 80}") + print(f"{title}:") + print(f" success: {result.get('success')}") + print(f" skipped: {result.get('skipped', False)}") + + tools_used = result.get("tools", []) + print(f" tools_called: {len(tools_used)}") + + if tools_used: + print("\n Tool Usage Details:") + for i, tool in enumerate(tools_used): + print(f" [{i}] Tool: {tool.name} Arguments: {tool.tool_call.arguments}") + + answer = result.get("answer", "") + if answer: + answer_preview = answer[:300] + "..." if len(answer) > 300 else answer + print(f"\n answer: {answer_preview}") + + print(f"{'=' * 80}") + + +def create_personal_info_introduction() -> list[Message]: + """Create a conversation where user introduces personal information.""" + return [ + Message( + role=Role.USER, + content="Hi! My name is Alice, and I'm a software engineer.", + ), + Message( + role=Role.ASSISTANT, + content="Nice to meet you, Alice! How can I help you today?", + ), + Message( + role=Role.USER, + content=( + "I love Python programming and working on AI projects. " + "I also enjoy reading sci-fi novels in my free time." + ), + ), + Message( + role=Role.ASSISTANT, + content="That's great! Python and AI are exciting fields. What kind of AI projects are you interested in?", + ), + ] + + +def create_detailed_profile_conversation() -> list[Message]: + """Create a comprehensive conversation with multiple personal details.""" + return [ + Message( + role=Role.USER, + content=( + "Let me tell you about myself. My name is Charlie Chen, " "I'm a data scientist based in San Francisco." + ), + ), + Message( + role=Role.ASSISTANT, + content="Hello Charlie! It's nice to meet you. Tell me more about your work.", + ), + Message( + role=Role.USER, + content=( + "I specialize in machine learning and natural language processing. " + "My favorite tools are PyTorch and Hugging Face transformers." + ), + ), + Message( + role=Role.ASSISTANT, + content="Those are excellent tools for NLP work. What kind of projects do you work on?", + ), + Message( + role=Role.USER, + content="I work on chatbots and sentiment analysis. Outside of work, I love hiking and photography.", + ), + Message( + role=Role.ASSISTANT, + content="That's a great combination of technical and creative interests!", + ), + Message( + role=Role.USER, + content=( + "Oh, and one more thing - please be more assertive when you think " + "I'm making a mistake. I want you to challenge my ideas." + ), + ), + Message( + role=Role.ASSISTANT, + content="Absolutely, I'll make sure to provide critical feedback when needed.", + ), + ] + + +async def test_summary_personal_info_storage(): + """Test summary() stores user's personal information to memory. + + User introduces name and basic info. + Expects: LLM should call WriteTool to store this information + """ + print("\n" + "=" * 80) + print("TEST 1: Summary - Personal Information Storage") + print("=" * 80) + + reme_fs = ReMeFs(enable_logo=False) + await reme_fs.start() + + messages = create_personal_info_introduction() + print_messages(messages, "INPUT MESSAGES", max_content_len=200) + + print("\nParameters:") + print(" version: default") + print(" Expected: LLM should call WriteTool to save user's name and interests") + + result = await reme_fs.summary( + messages=messages, + version="default", + date="2023-09-01", + ) + + print_result(result, "SUMMARY RESULT") + await reme_fs.close() + + +async def test_summary_detailed_profile(): + """Test summary() with comprehensive user profile information. + + User provides detailed personal, professional, and preference information. + Expects: LLM should organize and store all relevant information + """ + print("\n" + "=" * 80) + print("TEST 2: Summary - Detailed User Profile Storage") + print("=" * 80) + + reme_fs = ReMeFs(enable_logo=False, vector_store=None) + await reme_fs.start() + + messages = create_detailed_profile_conversation() + print_messages(messages, "INPUT MESSAGES", max_content_len=150) + + print("\nParameters:") + print(" version: default") + print(" Expected: LLM should extract and store:") + print(" - Name: Charlie Chen") + print(" - Profession: Data Scientist") + print(" - Location: San Francisco") + print(" - Skills: ML, NLP, PyTorch, Hugging Face") + print(" - Hobbies: Hiking, Photography") + print(" - Assistant behavior: Be assertive and critical") + + result = await reme_fs.summary( + messages=messages, + version="default", + date="2023-10-01", + ) + + print_result(result, "SUMMARY RESULT") + await reme_fs.close() + + +async def main(): + """Run core summary interface tests for personal memory storage.""" + print("\n" + "=" * 80) + print("ReMeFs Summary Interface - Personal Memory Storage Tests") + print("=" * 80) + print("\nThis test suite validates that the summary() function:") + print(" 1. Extracts personal information from user messages") + print(" 2. Calls file system tools (WriteTool/EditTool) to store the info") + print(" 3. Organizes information for future retrieval") + print("\nTest Scenarios:") + print(" 1. Personal info - name, profession, interests") + print(" 2. Detailed profile - comprehensive user information") + print("=" * 80) + + # Test 1: Basic personal information storage + await test_summary_personal_info_storage() + + # Test 2: Comprehensive user profile + await test_summary_detailed_profile() + + print("\n" + "=" * 80) + print("All summary tests completed!") + print("=" * 80) + print("\nNote: These tests require LLM calls to:") + print(" - Analyze user messages for personal information") + print(" - Decide what information to store") + print(" - Call appropriate tools (WriteTool/EditTool) to save to memory") + print("\nMake sure your API keys are properly configured before running.") + print("The LLM will autonomously decide what to store based on the conversation.") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/test_file_system_tool.py b/tests/test_fs_tool.py similarity index 100% rename from tests/test_file_system_tool.py rename to tests/test_fs_tool.py diff --git a/tests/test_memory_store.py b/tests/test_memory_store.py index a057c897..0052049c 100644 --- a/tests/test_memory_store.py +++ b/tests/test_memory_store.py @@ -21,8 +21,8 @@ from loguru import logger from reme.core.embedding import OpenAIEmbeddingModel from reme.core.enumeration.memory_source import MemorySource -from reme.core.memory_storage.base_memory_store import BaseMemoryStore -from reme.core.memory_storage.sqlite_memory_store import SqliteMemoryStore +from reme.core.memory_store.base_memory_store import BaseMemoryStore +from reme.core.memory_store.sqlite_memory_store import SqliteMemoryStore from reme.core.schema.file_metadata import FileMetadata from reme.core.schema.memory_chunk import MemoryChunk from reme.core.utils import load_env