diff --git a/.gitignore b/.gitignore index 7ee1a323..07fbdaca 100644 --- a/.gitignore +++ b/.gitignore @@ -39,4 +39,6 @@ chroma_vector_store/* bench_results/* meta_memory/* *.sqlite3 -**/data/*.json \ No newline at end of file +**/data/*.json +*.db +memories/* \ No newline at end of file diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 917a5cd8..e29ebf69 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -78,7 +78,7 @@ repos: --disable=C3001, --disable=R1702, --disable=R0912, - --max-statements=75, + --max-statements=120, --max-line-length=120, ] - repo: https://github.com/regebro/pyroma diff --git a/pyproject.toml b/pyproject.toml index f7b62c4a..9598ab81 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -34,6 +34,7 @@ keywords = ["llm", "memory", "experience", "memoryscope", "ai", "mcp", "http"] dependencies = [ "flowllm[reme]>=0.2.0.10", + "sqlite-vec>=0.1.6", ] [project.optional-dependencies] @@ -85,4 +86,7 @@ Repository = "https://github.com/agentscope-ai/ReMe" reme = "reme_ai.main:main" reme2 = "reme.reme:main" +[tool.pytest.ini_options] +asyncio_default_fixture_loop_scope = "function" + # python -m build && twine upload dist/* diff --git a/reme/agent/fs/__init__.py b/reme/agent/fs/__init__.py new file mode 100644 index 00000000..97dc861a --- /dev/null +++ b/reme/agent/fs/__init__.py @@ -0,0 +1,9 @@ +"""File system agents for memory management.""" + +from .fs_compactor import FsCompactor +from .fs_summarizer import FsSummarizer + +__all__ = [ + "FsSummarizer", + "FsCompactor", +] diff --git a/reme/agent/fs/fs_compactor.py b/reme/agent/fs/fs_compactor.py new file mode 100644 index 00000000..79d70c6f --- /dev/null +++ b/reme/agent/fs/fs_compactor.py @@ -0,0 +1,245 @@ +"""Context compaction agent for long sessions.""" + +from loguru import logger + +from ...core.enumeration import Role, MemoryType +from ...core.op import BaseReact +from ...core.schema import Message + + +class FsCompactor(BaseReact): + """Compact long conversation history into structured summaries.""" + + memory_type: MemoryType = MemoryType.PERSONAL + + def __init__( + self, + context_window_tokens: int = 128000, + reserve_tokens: int = 36000, + keep_recent_tokens: int = 20000, + **kwargs, + ): + super().__init__(tools=[], **kwargs) + self.context_window_tokens: int = context_window_tokens + self.reserve_tokens: int = reserve_tokens + self.keep_recent_tokens: int = keep_recent_tokens + + @staticmethod + def _is_user_message(message: Message) -> bool: + """Check if a message is a user-initiated message (user or tool result).""" + 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. + """ + 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: + """ + 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 + """ + 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]) + accumulated_tokens += msg_tokens + + if accumulated_tokens >= self.keep_recent_tokens: + cut_index = i + break + + if cut_index == 0: + return { + "messages_to_summarize": [], + "turn_prefix_messages": [], + "is_split_turn": False, + "cut_index": 0, + } + + # 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, + } + + # 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, + } + + # 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, + } + + @staticmethod + def _serialize_conversation(messages: list[Message]) -> str: + """Serialize conversation messages to text format.""" + lines = [] + for msg in messages: + role = msg.name if msg.name else msg.role.value + content = msg.content + if isinstance(content, str): + lines.append(f"[{role}]") + lines.append(content) + lines.append("") + elif isinstance(content, list): + lines.append(f"[{role}]") + for block in content: + lines.append(block.model_dump_json()) + lines.append("") + + 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] + + 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: + logger.info("No messages to summarize") + return [] + + system_prompt = self.get_prompt("system_prompt") + if self.context.get("previous_summary", ""): + 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) + + return [ + Message(role=Role.SYSTEM, content=system_prompt), + Message(role=Role.USER, content=f"\n{conversation_text}\n\n\n{user_prompt}"), + ] + + 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. + """ + if not self.context.turn_prefix_messages: + return [] + + system_prompt = self.get_prompt("system_prompt") + conversation_text = self._serialize_conversation(self.context.turn_prefix_messages) + turn_prefix_prompt = self.prompt_format("turn_prefix_summarization", conversation_text=conversation_text) + + return [ + Message(role=Role.SYSTEM, content=system_prompt), + Message(role=Role.USER, content=turn_prefix_prompt), + ] + + async def execute(self): + """ + 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) + """ + 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) + threshold = self.context_window_tokens - self.reserve_tokens + if token_count < threshold: + logger.info(f"Token count {token_count} below threshold, skipping compaction") + return { + "answer": "", + "success": True, + "messages": [], + "tools": [], + "skipped": True, + } + + 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 = "" + + 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 = "" + 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}") + + return { + "compacted": True, + "tokens_before": token_count, + "summary": summary, + "is_split_turn": self.context.is_split_turn, + } diff --git a/reme/agent/fs/fs_compactor.yaml b/reme/agent/fs/fs_compactor.yaml new file mode 100644 index 00000000..40fb6ec4 --- /dev/null +++ b/reme/agent/fs/fs_compactor.yaml @@ -0,0 +1,212 @@ +system_prompt: | + You are a context compaction assistant. Your role is to create structured summaries of conversations + that can be used to restore context in future sessions. Focus on preserving critical information while reducing token count. + +system_prompt_zh: | + 你是一个上下文压缩助手。你的角色是创建对话的结构化摘要, + 这些摘要可以在未来会话中用于恢复上下文。专注于保留关键信息,同时减少token数量。 + +initial_user_message: | + The messages above are a conversation to summarize. Create a structured context checkpoint summary + that another LLM will use to continue the work. + + Use this EXACT format: + + ## Goal + [What is the user trying to accomplish? Can be multiple items if the session covers different tasks.] + + ## Constraints & Preferences + - [Any constraints, preferences, or requirements mentioned by user] + - [Or "(none)" if none were mentioned] + + ## Progress + ### Done + - [x] [Completed tasks/changes] + + ### In Progress + - [ ] [Current work] + + ### Blocked + - [Issues preventing progress, if any] + + ## Key Decisions + - **[Decision]**: [Brief rationale] + + ## Next Steps + 1. [Ordered list of what should happen next] + + ## Critical Context + - [Any data, examples, or references needed to continue] + - [Or "(none)" if not applicable] + + Keep each section concise. Preserve exact file paths, function names, and error messages. + +initial_user_message_zh: | + 上述消息是一场需要总结的对话。创建一个结构化的上下文检查点摘要, + 以便另一个LLM可以用来继续工作。 + + 使用此确切格式: + + ## 目标 + [用户试图完成什么?如果会话涵盖不同任务,可以有多个项目。] + + ## 约束和偏好 + - [任何用户提到的约束、偏好或要求] + - [或者如果没有提到则为"(none)"] + + ## 进展 + ### 已完成 + - [x] [已完成的任务/更改] + + ### 进行中 + - [ ] [当前工作] + + ### 阻塞 + - [如果有任何阻碍进展的问题] + + ## 关键决策 + - **[决策]**: [简短理由] + + ## 下一步 + 1. [接下来应该发生的事情的有序列表] + + ## 关键上下文 + - [任何继续工作所需的数据、示例或参考资料] + - [或者如果不适用则为"(none)"] + + 保持每个部分简洁。保留确切的文件路径、函数名称和错误消息。 + +update_user_message: | + The messages above are NEW conversation messages to incorporate into the existing summary provided in + tags. + + + {previous_summary} + + + Update the existing structured summary with new information. RULES: + - PRESERVE all existing information from the previous summary + - ADD new progress, decisions, and context from the new messages + - UPDATE the Progress section: move items from "In Progress" to "Done" when completed + - UPDATE "Next Steps" based on what was accomplished + - PRESERVE exact file paths, function names, and error messages + - If something is no longer relevant, you may remove it + + Use this EXACT format: + + ## Goal + [Preserve existing goals, add new ones if the task expanded] + + ## Constraints & Preferences + - [Preserve existing, add new ones discovered] + + ## Progress + ### Done + - [x] [Include previously done items AND newly completed items] + + ### In Progress + - [ ] [Current work - update based on progress] + + ### Blocked + - [Current blockers - remove if resolved] + + ## Key Decisions + - **[Decision]**: [Brief rationale] (preserve all previous, add new) + + ## Next Steps + 1. [Update based on current state] + + ## Critical Context + - [Preserve important context, add new if needed] + + Keep each section concise. Preserve exact file paths, function names, and error messages. + +update_user_message_zh: | + 上述消息是要整合到现有摘要中的新对话消息,这些消息在标签中提供。 + + + {previous_summary} + + + 用新信息更新现有的结构化摘要。规则: + - 保留来自先前摘要的所有现有信息 + - 从新消息中添加新的进展、决策和上下文 + - 更新进度部分:当完成时将项目从"进行中"移到"已完成" + - 根据已完成的内容更新"下一步" + - 保留确切的文件路径、函数名称和错误消息 + - 如果某些内容不再相关,您可以删除它 + + 使用此确切格式: + + ## 目标 + [保留现有目标,如果任务扩展则添加新目标] + + ## 约束和偏好 + - [保留现有内容,添加发现的新内容] + + ## 进展 + ### 已完成 + - [x] [包含以前完成的项目和新完成的项目] + + ### 进行中 + - [ ] [当前工作 - 根据进展更新] + + ### 阻塞 + - [当前阻塞问题 - 如果解决则删除] + + ## 关键决策 + - **[决策]**: [简短理由](保留所有之前的内容,添加新的) + + ## 下一步 + 1. [根据当前状态更新] + + ## 关键上下文 + - [保留重要上下文,如需要则添加新的] + + 保持每个部分简洁。保留确切的文件路径、函数名称和错误消息。 + +# Used when cut point falls mid-turn (split turn scenario) +# Example: User → Assistant(part1) → [CUT HERE] → Assistant(part2 kept) → User +# This summarizes part1 to provide context for part2 +turn_prefix_summarization: | + + {conversation_text} + + + This is the PREFIX of a turn that was too large to keep. The SUFFIX (recent work) is retained. + + Summarize the prefix to provide context for the retained suffix: + + ## Original Request + [What did the user ask for in this turn?] + + ## Early Progress + - [Key decisions and work done in the prefix] + + ## Context for Suffix + - [Information needed to understand the retained recent work] + + Be concise. Focus on what's needed to understand the kept suffix. + +# 当切割点落在回合中间时使用(split turn 场景) +# 示例:用户 → 助手(部分1) → [切割点] → 助手(部分2保留) → 用户 +# 这会总结部分1以为部分2提供上下文 +turn_prefix_summarization_zh: | + + {conversation_text} + + + 这是一个过长而无法保留的回合的前缀。后缀(最近的工作)已保留。 + + 总结前缀以为保留的后缀提供上下文: + + ## 原始请求 + [用户在此回合中要求了什么?] + + ## 早期进展 + - [在前缀中做出的关键决策和完成的工作] + + ## 后缀上下文 + - [理解保留的最近工作所需的信息] + + 保持简洁。专注于理解保留后缀所需的内容。 diff --git a/reme/agent/fs/fs_summarizer.py b/reme/agent/fs/fs_summarizer.py new file mode 100644 index 00000000..25176dc1 --- /dev/null +++ b/reme/agent/fs/fs_summarizer.py @@ -0,0 +1,81 @@ +"""Personal memory retriever agent for retrieving personal memories through vector search.""" + +from loguru import logger + +from ...core.enumeration import Role, MemoryType +from ...core.op import BaseReact +from ...core.schema import Message + + +class FsSummarizer(BaseReact): + """Retrieve personal memories through vector search and history reading.""" + + memory_type: MemoryType = MemoryType.PERSONAL + + def __init__( + self, + memory_dir: str, + 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] + if self.version == "default": + messages.append( + Message( + role=Role.USER, + content=self.prompt_format( + "user_message_v2", + memory_dir=self.memory_dir, + ), + ), + ) + else: + messages.append(Message(role=Role.SYSTEM, content=self.get_prompt("system_prompt"))) + messages.append( + Message( + role=Role.USER, + content=self.prompt_format( + "user_message", + memory_dir=self.memory_dir, + ), + ), + ) + 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}") + return result diff --git a/reme/agent/fs/fs_summarizer.yaml b/reme/agent/fs/fs_summarizer.yaml new file mode 100644 index 00000000..f4673750 --- /dev/null +++ b/reme/agent/fs/fs_summarizer.yaml @@ -0,0 +1,15 @@ +system_prompt: | + Pre-compaction memory flush turn. + The session is near auto-compaction; capture durable memories to disk. + You may reply, but usually [SILENT] is correct. + +user_message: | + Pre-compaction memory flush. + 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. + 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 diff --git a/reme/agent/memory/base_memory_agent.py b/reme/agent/memory/base_memory_agent.py index fa6cd33e..dc2a0ed1 100644 --- a/reme/agent/memory/base_memory_agent.py +++ b/reme/agent/memory/base_memory_agent.py @@ -66,7 +66,7 @@ class BaseMemoryAgent(BaseReact, metaclass=ABCMeta): lines = [] for memory_target, memory_type in self.memory_target_type_mapping.items(): line = { - "agent": f"Agent managing {memory_type} memories for {memory_target}", + "agent": f"Agent managing {memory_type.value} memories for {memory_target}", "memory_target": memory_target, } lines.append(json.dumps(line, ensure_ascii=False)) diff --git a/reme/agent/memory/reme_retriever.py b/reme/agent/memory/reme_retriever.py index 9de51aa2..dcdade94 100644 --- a/reme/agent/memory/reme_retriever.py +++ b/reme/agent/memory/reme_retriever.py @@ -68,25 +68,27 @@ class ReMeRetriever(BaseMemoryAgent): async def execute(self): result = await super().execute() tools: list[BaseTool] = result["tools"] - delegate_task_tool = tools[0] - agents: list[BaseMemoryAgent] = delegate_task_tool.response.metadata["agents"] answer = [] success = True messages = [] - tools = [] + tools_result = [] retrieved_nodes = [] - for agent in agents: - answer.append(agent.response.answer) - success = success and agent.response.success - messages.extend(agent.response.metadata["messages"]) - tools.extend(agent.response.metadata["tools"]) - retrieved_nodes.extend(agent.response.metadata["retrieved_nodes"]) + + if tools: + delegate_task_tool = tools[0] + agents: list[BaseMemoryAgent] = delegate_task_tool.response.metadata["agents"] + for agent in agents: + answer.append(agent.response.answer) + success = success and agent.response.success + messages.extend(agent.response.metadata["messages"]) + tools_result.extend(agent.response.metadata["tools"]) + retrieved_nodes.extend(agent.response.metadata["retrieved_nodes"]) return { "answer": "\n".join(answer), "success": True, "messages": messages, - "tools": tools, + "tools": tools_result, "retrieved_nodes": retrieved_nodes, } diff --git a/reme/agent/memory/reme_summarizer.py b/reme/agent/memory/reme_summarizer.py index 8b263f64..3687fe2c 100644 --- a/reme/agent/memory/reme_summarizer.py +++ b/reme/agent/memory/reme_summarizer.py @@ -74,22 +74,25 @@ class ReMeSummarizer(BaseMemoryAgent): async def execute(self): result = await super().execute() tools: list[BaseTool] = result["tools"] - delegate_task_tool = tools[0] - agents: list[BaseMemoryAgent] = delegate_task_tool.response.metadata["agents"] success = True messages = [] - tools = [] + tools_result = [] memory_nodes = [] - for agent in agents: - success = success and agent.response.success - messages.extend(agent.response.metadata["messages"]) - tools.extend(agent.response.metadata["tools"]) - memory_nodes.extend(agent.response.metadata["memory_nodes"]) + + if tools: + delegate_task_tool = tools[0] + agents: list[BaseMemoryAgent] = delegate_task_tool.response.metadata["agents"] + + for agent in agents: + success = success and agent.response.success + messages.extend(agent.response.metadata["messages"]) + tools_result.extend(agent.response.metadata["tools"]) + memory_nodes.extend(agent.response.metadata["memory_nodes"]) return { "answer": memory_nodes, "success": True, "messages": messages, - "tools": tools, + "tools": tools_result, } diff --git a/reme/core/context/registry_factory.py b/reme/core/context/registry_factory.py index 28cb1c82..8cb51e61 100644 --- a/reme/core/context/registry_factory.py +++ b/reme/core/context/registry_factory.py @@ -36,10 +36,12 @@ class RegistryFactory: self.llm = Registry() self.embedding_model = Registry() self.vector_store = Registry() + self.memory_store = Registry() self.op = Registry() self.flow = Registry() self.service = Registry() self.token_counter = Registry() + self.file_watcher = Registry() R = RegistryFactory() diff --git a/reme/core/context/service_context.py b/reme/core/context/service_context.py index 23a9091c..05a45355 100644 --- a/reme/core/context/service_context.py +++ b/reme/core/context/service_context.py @@ -18,6 +18,7 @@ if TYPE_CHECKING: from ..token_counter import BaseTokenCounter from ..flow import BaseFlow from ..service import BaseService + from ..memory_storage import BaseMemoryStore class ServiceContext(BaseContext): @@ -77,6 +78,8 @@ class ServiceContext(BaseContext): self.embedding_models: dict[str, "BaseEmbeddingModel"] = {} self.token_counters: dict[str, "BaseTokenCounter"] = {} self.vector_stores: dict[str, "BaseVectorStore"] = {} + self.memory_stores: dict[str, "BaseMemoryStore"] = {} + self.flows: dict[str, "BaseFlow"] = {} self.mcp_server_mapping: dict[str, dict] = {} self.service: "BaseService" = R.service[self.service_config.backend](service_context=self) @@ -193,6 +196,14 @@ class ServiceContext(BaseContext): ) await self.vector_stores[name].create_collection(config.collection_name) + for name, config in self.service_config.memory_store.items(): + self.memory_stores[name] = R.memory_store[config.backend]( + store_name=config.store_name, + embedding_model=self.embedding_models[config.embedding_model], + **config.model_extra, + ) + await self.memory_stores[name].start() + if self.service_config.mcp_servers: await self.prepare_mcp_servers() @@ -226,6 +237,9 @@ class ServiceContext(BaseContext): for _, vector_store in self.vector_stores.items(): await vector_store.close() + for _, memory_store in self.memory_stores.items(): + await memory_store.close() + for _, llm in self.llms.items(): await llm.close() diff --git a/reme/core/embedding/base_embedding_model.py b/reme/core/embedding/base_embedding_model.py index 8b0916f3..4bd4cbe3 100644 --- a/reme/core/embedding/base_embedding_model.py +++ b/reme/core/embedding/base_embedding_model.py @@ -4,12 +4,15 @@ Defines the abstract base class and standard API for all embedding model impleme """ import asyncio +import hashlib import time from abc import ABC +from collections import OrderedDict from loguru import logger from ..schema import VectorNode +from ..schema.memory_chunk import MemoryChunk class BaseEmbeddingModel(ABC): @@ -27,17 +30,35 @@ class BaseEmbeddingModel(ABC): max_retries: int = 3, raise_exception: bool = True, max_input_length: int = 8192, + max_cache_size: int = 10000, **kwargs, ): - """Initialize model configuration and parameters.""" + """Initialize model configuration and parameters. + + Args: + model_name: Name of the embedding model + dimensions: Vector dimensions of the embeddings + max_batch_size: Maximum batch size for embedding requests + max_retries: Maximum number of retry attempts on failure + raise_exception: Whether to raise exceptions on failure + max_input_length: Maximum input text length + max_cache_size: Maximum number of embeddings to cache in memory (LRU) + **kwargs: Additional model-specific parameters + """ self.model_name = model_name self.dimensions = dimensions self.max_batch_size = max_batch_size self.max_retries = max_retries self.raise_exception = raise_exception self.max_input_length = max_input_length + self.max_cache_size = max_cache_size self.kwargs = kwargs + # Initialize LRU cache for embeddings + self._embedding_cache: OrderedDict[str, list[float]] = OrderedDict() + self._cache_hits = 0 + self._cache_misses = 0 + def _truncate_text(self, text: str) -> str: """Truncate text to max_input_length if it exceeds the limit.""" if len(text) > self.max_input_length: @@ -51,6 +72,76 @@ class BaseEmbeddingModel(ABC): """Truncate a list of texts to max_input_length.""" return [self._truncate_text(text) for text in texts] + def _get_cache_key(self, text: str) -> str: + """Generate a cache key by hashing the input text. + + Args: + text: Input text to hash + + Returns: + SHA256 hash of the text as hexadecimal string + """ + return hashlib.sha256(text.encode("utf-8")).hexdigest() + + def _get_from_cache(self, text: str) -> list[float] | None: + """Retrieve embedding from cache if it exists. + + Args: + text: Input text to look up + + Returns: + Cached embedding vector or None if not found + """ + cache_key = self._get_cache_key(text) + if cache_key in self._embedding_cache: + # Move to end (most recently used) + self._embedding_cache.move_to_end(cache_key) + self._cache_hits += 1 + return self._embedding_cache[cache_key] + self._cache_misses += 1 + return None + + def _put_to_cache(self, text: str, embedding: list[float]) -> None: + """Store embedding in cache with LRU eviction. + + Args: + text: Input text used as cache key + embedding: Embedding vector to cache + """ + if self.max_cache_size <= 0: + return + + cache_key = self._get_cache_key(text) + + # Remove oldest entry if cache is full + if len(self._embedding_cache) >= self.max_cache_size and cache_key not in self._embedding_cache: + self._embedding_cache.popitem(last=False) + + self._embedding_cache[cache_key] = embedding + self._embedding_cache.move_to_end(cache_key) + + def get_cache_stats(self) -> dict[str, int]: + """Get cache statistics. + + Returns: + Dictionary with cache size, hits, misses, and hit rate + """ + total_requests = self._cache_hits + self._cache_misses + hit_rate = self._cache_hits / total_requests if total_requests > 0 else 0.0 + return { + "cache_size": len(self._embedding_cache), + "max_cache_size": self.max_cache_size, + "cache_hits": self._cache_hits, + "cache_misses": self._cache_misses, + "hit_rate": hit_rate, + } + + def clear_cache(self) -> None: + """Clear the embedding cache and reset statistics.""" + self._embedding_cache.clear() + self._cache_hits = 0 + self._cache_misses = 0 + async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]: """Internal async implementation for calling the embedding API with batch input.""" @@ -60,10 +151,20 @@ class BaseEmbeddingModel(ABC): async def get_embedding(self, input_text: str, **kwargs) -> list[float]: """Async get embedding for a single text with exponential backoff retries.""" truncated_text = self._truncate_text(input_text) + + # Check cache first + cached_embedding = self._get_from_cache(truncated_text) + if cached_embedding is not None: + return cached_embedding + + # Cache miss - compute embedding for i in range(self.max_retries): try: result = await self._get_embeddings([truncated_text], **kwargs) - return result[0] + embedding = result[0] + # Store in cache + self._put_to_cache(truncated_text, embedding) + return embedding except Exception as e: logger.error(f"Model {self.model_name} failed: {e}") if i == self.max_retries - 1: @@ -78,16 +179,36 @@ class BaseEmbeddingModel(ABC): # Truncate all input texts first truncated_texts = self._truncate_texts(input_text) - # Split into batches and process sequentially to respect rate limits - results = [] - for i in range(0, len(truncated_texts), self.max_batch_size): - batch = truncated_texts[i : i + self.max_batch_size] + # Check cache for each text and separate cached vs uncached + results: list[list[float] | None] = [None] * len(truncated_texts) + texts_to_compute: list[tuple[int, str]] = [] # (original_index, text) + + for idx, text in enumerate(truncated_texts): + cached = self._get_from_cache(text) + if cached is not None: + results[idx] = cached + else: + texts_to_compute.append((idx, text)) + + # If all texts were cached, return early + if not texts_to_compute: + return [r for r in results if r is not None] + + # Compute embeddings for uncached texts in batches + uncached_texts = [text for _, text in texts_to_compute] + for i in range(0, len(uncached_texts), self.max_batch_size): + batch_texts = uncached_texts[i : i + self.max_batch_size] + batch_indices = [idx for idx, _ in texts_to_compute[i : i + self.max_batch_size]] + # Process each batch with retry logic for retry in range(self.max_retries): try: - batch_res = await self._get_embeddings(batch, **kwargs) - if batch_res: - results.extend(batch_res) + batch_embeddings = await self._get_embeddings(batch_texts, **kwargs) + if batch_embeddings: + # Store results and cache them + for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings): + results[orig_idx] = embedding + self._put_to_cache(text, embedding) break except Exception as e: logger.error(f"Model {self.model_name} batch failed: {e}") @@ -96,15 +217,26 @@ class BaseEmbeddingModel(ABC): raise else: await asyncio.sleep(retry + 1) - return results + + return [r for r in results if r is not None] def get_embedding_sync(self, input_text: str, **kwargs) -> list[float]: """Synchronous get embedding for a single text with retry logic.""" truncated_text = self._truncate_text(input_text) + + # Check cache first + cached_embedding = self._get_from_cache(truncated_text) + if cached_embedding is not None: + return cached_embedding + + # Cache miss - compute embedding for i in range(self.max_retries): try: result = self._get_embeddings_sync([truncated_text], **kwargs) - return result[0] + embedding = result[0] + # Store in cache + self._put_to_cache(truncated_text, embedding) + return embedding except Exception as exc: logger.error(f"Model {self.model_name} failed: {exc}") if i == self.max_retries - 1: @@ -119,15 +251,36 @@ class BaseEmbeddingModel(ABC): # Truncate all input texts first truncated_texts = self._truncate_texts(input_text) - results = [] - for i in range(0, len(truncated_texts), self.max_batch_size): - batch = truncated_texts[i : i + self.max_batch_size] + # Check cache for each text and separate cached vs uncached + results: list[list[float] | None] = [None] * len(truncated_texts) + texts_to_compute: list[tuple[int, str]] = [] # (original_index, text) + + for idx, text in enumerate(truncated_texts): + cached = self._get_from_cache(text) + if cached is not None: + results[idx] = cached + else: + texts_to_compute.append((idx, text)) + + # If all texts were cached, return early + if not texts_to_compute: + return [r for r in results if r is not None] + + # Compute embeddings for uncached texts in batches + uncached_texts = [text for _, text in texts_to_compute] + for i in range(0, len(uncached_texts), self.max_batch_size): + batch_texts = uncached_texts[i : i + self.max_batch_size] + batch_indices = [idx for idx, _ in texts_to_compute[i : i + self.max_batch_size]] + # Process each batch with retry logic for retry in range(self.max_retries): try: - batch_res = self._get_embeddings_sync(batch, **kwargs) - if batch_res: - results.extend(batch_res) + batch_embeddings = self._get_embeddings_sync(batch_texts, **kwargs) + if batch_embeddings: + # Store results and cache them + for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings): + results[orig_idx] = embedding + self._put_to_cache(text, embedding) break except Exception as exc: logger.error(f"Model {self.model_name} batch failed: {exc}") @@ -136,7 +289,8 @@ class BaseEmbeddingModel(ABC): raise else: time.sleep(retry + 1) - return results + + return [r for r in results if r is not None] async def get_node_embedding(self, node: VectorNode, **kwargs) -> VectorNode: """Async generate and populate vector field for a single VectorNode object.""" @@ -172,6 +326,72 @@ class BaseEmbeddingModel(ABC): logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(nodes)} nodes") return nodes + async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk: + """Async generate and populate embedding field for a single MemoryChunk object. + + Args: + chunk: MemoryChunk object containing text to embed + **kwargs: Additional arguments passed to the embedding model + + Returns: + The same MemoryChunk object with populated embedding field + """ + chunk.embedding = await self.get_embedding(chunk.text, **kwargs) + return chunk + + async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]: + """Async generate and populate embedding fields for a batch of MemoryChunk objects. + + Args: + chunks: List of MemoryChunk objects containing text to embed + **kwargs: Additional arguments passed to the embedding model + + Returns: + The same list of MemoryChunk objects with populated embedding fields + """ + texts = [chunk.text for chunk in chunks] + embeddings: list[list[float]] = await self.get_embeddings(texts, **kwargs) + + if len(embeddings) == len(chunks): + for chunk, vec in zip(chunks, embeddings): + chunk.embedding = vec + else: + logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(chunks)} chunks") + return chunks + + def get_chunk_embedding_sync(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk: + """Synchronously generate and populate embedding field for a single MemoryChunk object. + + Args: + chunk: MemoryChunk object containing text to embed + **kwargs: Additional arguments passed to the embedding model + + Returns: + The same MemoryChunk object with populated embedding field + """ + chunk.embedding = self.get_embedding_sync(chunk.text, **kwargs) + return chunk + + def get_chunk_embeddings_sync(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]: + """Synchronously generate embeddings for a batch of MemoryChunk objects. + + Args: + chunks: List of MemoryChunk objects containing text to embed + **kwargs: Additional arguments passed to the embedding model + + Returns: + The same list of MemoryChunk objects with populated embedding fields + """ + texts = [chunk.text for chunk in chunks] + embeddings: list[list[float]] = self.get_embeddings_sync(texts, **kwargs) + + if len(embeddings) == len(chunks): + for chunk, vec in zip(chunks, embeddings): + chunk.embedding = vec + else: + logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(chunks)} chunks") + return chunks + def close_sync(self): """Synchronously release resources and close connections.""" diff --git a/reme/core/enumeration/__init__.py b/reme/core/enumeration/__init__.py index 7202a949..c709c20d 100644 --- a/reme/core/enumeration/__init__.py +++ b/reme/core/enumeration/__init__.py @@ -3,6 +3,7 @@ from .chunk_enum import ChunkEnum from .http_enum import HttpEnum from .json_schema_enum import JsonSchemaEnum +from .memory_source import MemorySource from .memory_type import MemoryType from .registry_enum import RegistryEnum from .role import Role @@ -11,6 +12,7 @@ __all__ = [ "ChunkEnum", "HttpEnum", "JsonSchemaEnum", + "MemorySource", "MemoryType", "RegistryEnum", "Role", diff --git a/reme/core/enumeration/memory_source.py b/reme/core/enumeration/memory_source.py new file mode 100644 index 00000000..05174549 --- /dev/null +++ b/reme/core/enumeration/memory_source.py @@ -0,0 +1,11 @@ +"""Memory source types.""" + +from enum import Enum + + +class MemorySource(str, Enum): + """Source of memory data.""" + + MEMORY = "memory" + + SESSIONS = "sessions" diff --git a/reme/core/enumeration/registry_enum.py b/reme/core/enumeration/registry_enum.py index 876c06b8..aec8ca93 100644 --- a/reme/core/enumeration/registry_enum.py +++ b/reme/core/enumeration/registry_enum.py @@ -15,6 +15,9 @@ class RegistryEnum(str, Enum): # Databases or storage systems for vector search VECTOR_STORE = "vector_store" + # Databases or storage systems for long-term memory storage + MEMORY_STORE = "memory_store" + # Atomic operations or functional units OP = "op" diff --git a/reme/core/file_watcher/__init__.py b/reme/core/file_watcher/__init__.py new file mode 100644 index 00000000..74eb0fd7 --- /dev/null +++ b/reme/core/file_watcher/__init__.py @@ -0,0 +1,19 @@ +"""File watcher module for monitoring file system changes. + +This module provides file watcher implementations for monitoring file changes +and updating memory stores accordingly. +""" + +from .base_file_watcher import BaseFileWatcher +from .delta_file_watcher import DeltaFileWatcher +from .full_file_watcher import FullFileWatcher +from ..context import R + +__all__ = [ + "BaseFileWatcher", + "DeltaFileWatcher", + "FullFileWatcher", +] + +R.file_watcher.register("full")(FullFileWatcher) +R.file_watcher.register("delta")(DeltaFileWatcher) diff --git a/reme/core/file_watcher/base_file_watcher.py b/reme/core/file_watcher/base_file_watcher.py new file mode 100644 index 00000000..8df6b9b4 --- /dev/null +++ b/reme/core/file_watcher/base_file_watcher.py @@ -0,0 +1,126 @@ +"""Base file watcher implementation. + +This module provides the base class for file watcher implementations +that monitor file system changes and trigger callbacks. +""" + +import asyncio +from collections.abc import Coroutine +from typing import Any, Callable + +from loguru import logger +from watchfiles import awatch, Change + +from ..memory_storage import BaseMemoryStore + + +class BaseFileWatcher: + """ + Minimal file watcher base class + + This base class provides basic file monitoring functionality that can be extended + to implement specific file monitoring requirements. + """ + + def __init__( + self, + watch_paths: list[str] | str, + recursive: bool = False, + debounce: int = 500, # Millisecond debounce + suffix_filters: list[str] | None = None, + callback: Callable[[set[tuple[Change, str]]], None | Coroutine[Any, Any, None]] | None = None, + memory_store: BaseMemoryStore | None = None, + **kwargs, + ): + """ + Initialize the file watcher""" + self.watch_paths: list[str] = [watch_paths] if isinstance(watch_paths, str) else watch_paths + self.recursive: bool = recursive + self.debounce: int = debounce + self.suffix_filters: list[str] = suffix_filters or [] + self.callback = callback + self.memory_store: BaseMemoryStore = memory_store + self.kwargs: dict = kwargs + + self._stop_event = asyncio.Event() + self._watch_task: asyncio.Task | None = None + self._running = False + + async def start(self): + """Start the file watcher""" + if self._running: + return + + self._running = True + self._watch_task = asyncio.create_task(self._watch_loop()) + logger.info(f"Started watching: {self.watch_paths}") + + async def close(self): + """Stop the file watcher""" + if not self._running: + return + + self._stop_event.set() + if self._watch_task: + await self._watch_task + self._running = False + logger.info("Stopped watching") + + def watch_filter(self, _change: Change, path: str) -> bool: + """Filter function for file watching.""" + # If no suffix filters are specified, watch all files + if not self.suffix_filters: + return True + + # Check if the file has one of the allowed suffixes + for suffix in self.suffix_filters: + if path.endswith("." + suffix.strip(".")): + return True + + return False + + async def _watch_loop(self): + """Core monitoring loop""" + async for changes in awatch( + *self.watch_paths, + watch_filter=self.watch_filter, + recursive=self.recursive, + debounce=self.debounce, + stop_event=self._stop_event, + ): + if self._stop_event.is_set(): + break + + await self.on_changes(changes) + + async def _on_changes(self, changes: set[tuple[Change, str]]): + """Callback method to handle file changes""" + + async def on_changes(self, changes: set[tuple[Change, str]]): + """Hook method to handle file changes""" + if self.callback: + result = self.callback(changes) + if asyncio.iscoroutine(result): + await result + else: + await self._on_changes(changes) + + def is_running(self) -> bool: + """Check if the watcher is running""" + return self._running + + async def add_path(self, path: str): + """Dynamically add a path to monitor""" + if path not in self.watch_paths: + self.watch_paths.append(path) + if self._running: + await self.close() + await self.start() + + async def remove_path(self, path: str): + """Remove a monitored path""" + if path in self.watch_paths: + self.watch_paths.remove(path) + if self._running: + await self.close() + await self.start() diff --git a/reme/core/file_watcher/delta_file_watcher.py b/reme/core/file_watcher/delta_file_watcher.py new file mode 100644 index 00000000..60807133 --- /dev/null +++ b/reme/core/file_watcher/delta_file_watcher.py @@ -0,0 +1,277 @@ +"""Delta file watcher for incremental file synchronization. + +This module provides a file watcher that detects append-only changes +and only processes newly added content, avoiding redundant operations. +""" + +import asyncio +import os + +from loguru import logger +from watchfiles import Change + +from .base_file_watcher import BaseFileWatcher +from ..enumeration import MemorySource +from ..schema import FileMetadata, MemoryChunk +from ..utils import chunk_markdown, hash_text + + +class DeltaFileWatcher(BaseFileWatcher): + """Delta file watcher implementation for incremental synchronization. + + This watcher detects append-only changes (e.g., log files) and only processes + the newly added content, avoiding redundant embedding requests for unchanged content. + + Strategy: + - Detect if file is append-only (new lines added at end) + - Find the safe cutoff point (considering chunk overlap) + - Only re-chunk and embed content from cutoff to end + - Delete affected old chunks and insert new chunks + """ + + def __init__(self, chunk_tokens: int = 400, chunk_overlap: int = 80, overlap_lines: int = 2, **kwargs): + """ + Initialize delta file watcher. + + Args: + chunk_tokens: Maximum tokens per chunk + chunk_overlap: Overlap tokens between chunks + """ + super().__init__(**kwargs) + self.chunk_tokens = chunk_tokens + self.chunk_overlap = chunk_overlap + self.overlap_lines = overlap_lines + + self.dirty = False + + @staticmethod + async def _build_file_metadata(path: str) -> FileMetadata: + """Build file metadata from filesystem.""" + + def _read_file_sync(): + stat_t = os.stat(path) + with open(path, "r", encoding="utf-8") as f: + content_t = f.read() + return stat_t, content_t + + stat, content = await asyncio.to_thread(_read_file_sync) + return FileMetadata( + hash=hash_text(content), + mtime_ms=stat.st_mtime * 1000, + size=stat.st_size, + path=path, + content=content, + ) + + def _find_cutoff_line( + self, + old_chunks: list[MemoryChunk], + old_file_meta: FileMetadata, + new_file_meta: FileMetadata, + ) -> int | None: + """Find the safe cutoff line for incremental update. + + Uses a heuristic approach: if file size increased and hash changed, + we verify by comparing content. For true append-only files (like logs), + the old content should be a prefix of new content. + + Args: + old_chunks: Existing chunks sorted by start_line + old_file_meta: Previous file metadata + new_file_meta: Current file metadata (with content) + + Returns: + Cutoff line number (1-indexed), or None if not append-only + """ + if not old_chunks: + return None + + # File shrunk - definitely not append-only + if new_file_meta.size < old_file_meta.size: + logger.debug("File shrunk, not append-only") + return None + + # File didn't grow much - might be a modification + size_growth = new_file_meta.size - old_file_meta.size + if size_growth < 10: # Less than 10 bytes growth + logger.debug("Minimal size growth, treating as modification") + return None + + # Verify append-only by checking if old content is prefix + # We need to read old file content from chunks + old_chunks_sorted = sorted(old_chunks, key=lambda c: c.start_line) + + # Simple heuristic: check if first few chunks' content matches + # This avoids reconstructing full old content + new_lines = new_file_meta.content.split("\n") + + # Sample check: verify first chunk still matches + first_chunk = old_chunks_sorted[0] + first_chunk_lines = first_chunk.text.split("\n") + new_first_lines = new_lines[first_chunk.start_line - 1 : first_chunk.end_line] + + # Compare (allowing for minor whitespace differences at boundaries) + if len(first_chunk_lines) > 0 and len(new_first_lines) > 0: + # Check if most of the lines match + matches = sum(1 for old, new in zip(first_chunk_lines, new_first_lines) if old == new) + if matches < len(first_chunk_lines) * 0.8: # Less than 80% match + logger.debug("First chunk content changed, not append-only") + return None + + # File appears to be append-only + # Find the last chunk and set cutoff considering overlap + last_chunk = max(old_chunks_sorted, key=lambda c: c.end_line) + cutoff_line = max(1, last_chunk.end_line - self.overlap_lines) + + logger.debug( + f"Append-only detected: size {old_file_meta.size} -> {new_file_meta.size}, " + f"cutoff at line {cutoff_line}", + ) + + return cutoff_line + + @staticmethod + def _extract_content_from_line(content: str, start_line: int) -> str: + """Extract content starting from a specific line number.""" + lines = content.split("\n") + if start_line <= 1: + return content + if start_line > len(lines): + return "" + # start_line is 1-indexed, array is 0-indexed + return "\n".join(lines[start_line - 1 :]) + + async def _on_changes(self, changes: set[tuple[Change, str]]): + """Handle file changes with incremental synchronization.""" + self.dirty = True + + for change_type, path in changes: + if change_type == Change.added: + # New file: process everything + file_meta = await self._build_file_metadata(path) + chunks = ( + chunk_markdown( + file_meta.content, + file_meta.path, + MemorySource.MEMORY, + self.chunk_tokens, + self.chunk_overlap, + ) + or [] + ) + + if chunks: + chunks = await self.memory_store.get_chunk_embeddings(chunks) + file_meta.chunk_count = len(chunks) + await self.memory_store.upsert_file(file_meta, MemorySource.MEMORY, chunks) + logger.info(f"File added: {path} ({len(chunks)} chunks)") + else: + logger.warning(f"No chunks generated for new file {path}") + + elif change_type == Change.modified: + # Get existing data + old_chunks = await self.memory_store.get_file_chunks(path, MemorySource.MEMORY) + old_file_meta = await self.memory_store.get_file_metadata(path, MemorySource.MEMORY) + + # Read new file + file_meta = await self._build_file_metadata(path) + + # If no old chunks, fallback to full update + if not old_chunks or not old_file_meta: + logger.debug(f"No existing chunks for {path}, doing full update") + chunks = ( + chunk_markdown( + file_meta.content, + file_meta.path, + MemorySource.MEMORY, + self.chunk_tokens, + self.chunk_overlap, + ) + or [] + ) + if chunks: + chunks = await self.memory_store.get_chunk_embeddings(chunks) + file_meta.chunk_count = len(chunks) + await self.memory_store.delete_file(path, MemorySource.MEMORY) + await self.memory_store.upsert_file( + file_meta, + MemorySource.MEMORY, + chunks, + ) + logger.info(f"File modified (full): {path} ({len(chunks)} chunks)") + continue + + # Check if append-only and find cutoff line + old_chunks_sorted = sorted(old_chunks, key=lambda c: c.start_line) + cutoff_line = self._find_cutoff_line(old_chunks_sorted, old_file_meta, file_meta) + + if cutoff_line is None: + # Not append-only, do full update + logger.debug(f"File {path} has modifications, doing full update") + chunks = ( + chunk_markdown( + file_meta.content, + file_meta.path, + MemorySource.MEMORY, + self.chunk_tokens, + self.chunk_overlap, + ) + or [] + ) + if chunks: + chunks = await self.memory_store.get_chunk_embeddings(chunks) + file_meta.chunk_count = len(chunks) + await self.memory_store.delete_file(path, MemorySource.MEMORY) + await self.memory_store.upsert_file(file_meta, MemorySource.MEMORY, chunks) + logger.info(f"File modified (full): {path} ({len(chunks)} chunks)") + else: + # Append-only: incremental update + new_content_part = self._extract_content_from_line(file_meta.content, cutoff_line) + + new_chunks = ( + chunk_markdown( + new_content_part, + file_meta.path, + MemorySource.MEMORY, + self.chunk_tokens, + self.chunk_overlap, + ) + or [] + ) + + if not new_chunks: + logger.debug(f"No new chunks for {path}, skipping") + continue + + for idx, chunk in enumerate(new_chunks): + chunk.start_line += cutoff_line - 1 + chunk.end_line += cutoff_line - 1 + chunk.id = hash_text( + f"{chunk.source}:{chunk.path}:{chunk.start_line}:" f"{chunk.end_line}:{chunk.hash}:{idx}", + ) + + new_chunks = await self.memory_store.get_chunk_embeddings(new_chunks) + + chunks_to_delete = [c.id for c in old_chunks_sorted if c.start_line >= cutoff_line] + + # Apply incremental updates + if chunks_to_delete: + await self.memory_store.delete_file_chunks(path, chunks_to_delete) + + if new_chunks: + await self.memory_store.upsert_chunks(new_chunks, MemorySource.MEMORY) + + logger.info( + f"File modified (incremental): {path} " + f"(cutoff: line {cutoff_line}, " + f"+{len(new_chunks)} chunks, -{len(chunks_to_delete)} chunks)", + ) + + elif change_type == Change.deleted: + await self.memory_store.delete_file(path, MemorySource.MEMORY) + logger.info(f"File deleted: {path}") + + else: + logger.warning(f"Unknown change type: {change_type}") + + self.dirty = False diff --git a/reme/core/file_watcher/full_file_watcher.py b/reme/core/file_watcher/full_file_watcher.py new file mode 100644 index 00000000..69870c2c --- /dev/null +++ b/reme/core/file_watcher/full_file_watcher.py @@ -0,0 +1,76 @@ +"""Full file watcher for complete file synchronization. + +This module provides a file watcher that processes entire files +on any change, ensuring complete synchronization. +""" + +import asyncio +import os + +from loguru import logger +from watchfiles import Change + +from .base_file_watcher import BaseFileWatcher +from ..enumeration import MemorySource +from ..schema import FileMetadata +from ..utils import chunk_markdown, hash_text + + +class FullFileWatcher(BaseFileWatcher): + """Full file watcher implementation for full synchronization""" + + def __init__(self, chunk_tokens: int = 400, chunk_overlap: int = 80, **kwargs): + """ + Initialize full file watcher""" + super().__init__(**kwargs) + self.chunk_tokens = chunk_tokens + self.chunk_overlap = chunk_overlap + self.dirty = False + + @staticmethod + async def _build_file_metadata(path: str) -> FileMetadata: + def _read_file_sync(): + stat_t = os.stat(path) + with open(path, "r", encoding="utf-8") as f: + content_t = f.read() + return stat_t, content_t + + stat, content = await asyncio.to_thread(_read_file_sync) + return FileMetadata( + hash=hash_text(content), + mtime_ms=stat.st_mtime * 1000, + size=stat.st_size, + path=path, + content=content, + ) + + async def _on_changes(self, changes: set[tuple[Change, str]]): + """Handle file changes with full synchronization""" + self.dirty = True + for change_type, path in changes: + if change_type in [Change.added, Change.modified]: + file_meta = await self._build_file_metadata(path) + chunks = ( + chunk_markdown( + file_meta.content, + file_meta.path, + MemorySource.MEMORY, + self.chunk_tokens, + self.chunk_overlap, + ) + or [] + ) + if chunks: + chunks = await self.memory_store.get_chunk_embeddings(chunks) + file_meta.chunk_count = len(chunks) + await self.memory_store.delete_file(file_meta.path, MemorySource.MEMORY) + await self.memory_store.upsert_file(file_meta, MemorySource.MEMORY, chunks) + + elif change_type == Change.deleted: + await self.memory_store.delete_file(path, MemorySource.MEMORY) + + else: + logger.warning(f"Unknown change type: {change_type}") + + logger.info(f"File {change_type} changed: {path}") + self.dirty = False diff --git a/reme/core/memory_storage/__init__.py b/reme/core/memory_storage/__init__.py new file mode 100644 index 00000000..371d4e0f --- /dev/null +++ b/reme/core/memory_storage/__init__.py @@ -0,0 +1,16 @@ +"""Memory storage 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. +""" + +from .base_memory_store import BaseMemoryStore +from .sqlite_memory_store import SqliteMemoryStore +from ..context import R + +__all__ = [ + "BaseMemoryStore", + "SqliteMemoryStore", +] + +R.memory_store.register("sqlite")(SqliteMemoryStore) diff --git a/reme/core/memory_storage/base_memory_store.py b/reme/core/memory_storage/base_memory_store.py new file mode 100644 index 00000000..ed8f59e7 --- /dev/null +++ b/reme/core/memory_storage/base_memory_store.py @@ -0,0 +1,126 @@ +"""Base storage interface for memory manager.""" + +from abc import ABC, abstractmethod + +from ..embedding import BaseEmbeddingModel +from ..enumeration import MemorySource +from ..schema import FileMetadata, MemoryChunk, MemorySearchResult + + +class BaseMemoryStore(ABC): + """Abstract base class for memory storage backends.""" + + def __init__( + self, + store_name: str, + embedding_model: BaseEmbeddingModel, + fts_enabled: bool = True, + snippet_max_chars: int = 700, + **kwargs, + ): + """Initialize""" + self.store_name: str = store_name + self.embedding_model: BaseEmbeddingModel = embedding_model + self.fts_enabled: bool = fts_enabled + self.snippet_max_chars: int = snippet_max_chars + self.kwargs: dict = kwargs + + self.vector_available = False + self.fts_available = False + + @property + def embedding_dim(self) -> int: + """Get the embedding model's dimensionality.""" + return self.embedding_model.dimensions + + async def get_embedding(self, query: str, **kwargs) -> list[float]: + """Get embedding for a single query string.""" + return await self.embedding_model.get_embedding(query, **kwargs) + + async def get_embeddings(self, queries: list[str], **kwargs) -> list[list[float]]: + """Get embeddings for a batch of query strings.""" + return await self.embedding_model.get_embeddings(queries, **kwargs) + + async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk: + """Generate and populate embedding field for a single MemoryChunk object.""" + return await self.embedding_model.get_chunk_embedding(chunk, **kwargs) + + async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]: + """Generate and populate embedding fields for a batch of MemoryChunk objects.""" + return await self.embedding_model.get_chunk_embeddings(chunks, **kwargs) + + @abstractmethod + async def start(self): + """Initialize the storage backend.""" + + @abstractmethod + async def upsert_file(self, file_meta: FileMetadata, source: MemorySource, chunks: list[MemoryChunk]): + """Insert or update a file and its chunks.""" + + @abstractmethod + async def delete_file(self, path: str, source: MemorySource): + """Delete a file and all its chunks.""" + + @abstractmethod + async def delete_file_chunks(self, path: str, chunk_ids: list[str]): + """Delete chunks for a file.""" + + @abstractmethod + async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource): + """Insert or update specific chunks without affecting other chunks.""" + + @abstractmethod + async def list_files(self, source: MemorySource) -> list[str]: + """List all indexed file paths for a source.""" + + @abstractmethod + async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None: + """Get full file metadata with statistics.""" + + @abstractmethod + async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]: + """Get all chunks for a file.""" + + @abstractmethod + async def vector_search( + self, + query: str, + limit: int, + sources: list[MemorySource] | None = None, + ) -> list[MemorySearchResult]: + """Perform vector similarity search. + + Args: + query: Query embedding vector + limit: Maximum number of results + sources: Optional list of sources to filter + + Returns: + List of search results sorted by similarity + """ + + @abstractmethod + async def keyword_search( + self, + query: str, + limit: int, + sources: list[MemorySource] | None = None, + ) -> list[MemorySearchResult]: + """Perform keyword/full-text search. + + Args: + query: Search query text + limit: Maximum number of results + sources: Optional list of sources to filter + + Returns: + List of search results sorted by relevance + """ + + @abstractmethod + async def clear_all(self): + """Clear all indexed data.""" + + @abstractmethod + async def close(self): + """Close storage and release resources.""" diff --git a/reme/core/memory_storage/sqlite_memory_store.py b/reme/core/memory_storage/sqlite_memory_store.py new file mode 100644 index 00000000..756c2b3b --- /dev/null +++ b/reme/core/memory_storage/sqlite_memory_store.py @@ -0,0 +1,652 @@ +"""SQLite storage backend for memory index.""" + +import json +import sqlite3 +import struct +import time +from pathlib import Path + +from loguru import logger + +from .base_memory_store import BaseMemoryStore +from ..enumeration import MemorySource +from ..schema import FileMetadata, MemoryChunk, MemorySearchResult + + +class SqliteMemoryStore(BaseMemoryStore): + """SQLite memory storage with vector and full-text search. + + Inherits embedding methods from BaseMemoryStore: + - get_chunk_embedding / get_chunk_embeddings (async) + - get_chunk_embedding_sync / get_chunk_embeddings_sync (sync) + - get_embedding / get_embeddings (async) + + Provides SQLite-backed persistent storage with: + - Vector similarity search (via sqlite-vec extension) + - Full-text search (via FTS5) + - Efficient chunk and file metadata management + """ + + def __init__(self, db_path: str = ".reme/memory.db", vec_ext_path: str = "", **kwargs): + super().__init__(**kwargs) + self.db_path = db_path + self.vec_ext_path = vec_ext_path + + self.conn: sqlite3.Connection | None = None + + @property + def vector_table_name(self) -> str: + """Get the name of the vector table for this store.""" + return f"chunks_vec_{self.store_name}" + + @property + def fts_table_name(self) -> str: + """Get the name of the FTS table for this store.""" + return f"chunks_fts_{self.store_name}" + + @property + def chunks_table_name(self) -> str: + """Get the name of the chunks table for this store.""" + return f"chunks_{self.store_name}" + + @property + def files_table_name(self) -> str: + """Get the name of the files table for this store.""" + return f"files_{self.store_name}" + + @staticmethod + def vector_to_blob(embedding: list[float]) -> bytes: + """Convert vector to binary blob for sqlite-vec.""" + return struct.pack(f"{len(embedding)}f", *embedding) + + async def start(self) -> None: + """Initialize database and load extensions.""" + if self.conn is not None: + return + + Path(self.db_path).parent.mkdir(parents=True, exist_ok=True) + + self.conn = sqlite3.connect(self.db_path, check_same_thread=False) + self.conn.enable_load_extension(True) + + # Load sqlite-vec extension + if self.vec_ext_path: + try: + self.conn.load_extension(self.vec_ext_path) + self.vector_available = True + logger.info(f"Loaded sqlite-vec: {self.vec_ext_path}") + except Exception as e: + logger.warning(f"Failed to load sqlite-vec: {e}") + + else: + try: + import sqlite_vec + + ext_path = sqlite_vec.loadable_path() + self.conn.load_extension(ext_path) + self.vector_available = True + logger.info(f"Loaded sqlite-vec from package: {ext_path}") + + except Exception as e: + logger.warning(f"Failed to load sqlite-vec from package: {e}") + # Fallback: try common extension names + for name in ["vec0", "sqlite_vec", "vector0"]: + try: + self.conn.load_extension(name) + self.vector_available = True + logger.info(f"Loaded sqlite-vec: {name}") + break + except Exception: + pass + + self.conn.enable_load_extension(False) + await self._create_tables() + + async def _create_tables(self) -> None: + """Create database schema.""" + cursor = self.conn.cursor() + + # Files + cursor.execute( + f""" + CREATE TABLE IF NOT EXISTS {self.files_table_name} ( + path TEXT, + source TEXT, + hash TEXT, + mtime REAL, + size INTEGER, + PRIMARY KEY (path, source) + ) + """, + ) + + # Chunks + cursor.execute( + f""" + CREATE TABLE IF NOT EXISTS {self.chunks_table_name} ( + id TEXT PRIMARY KEY, + path TEXT, + source TEXT, + start_line INTEGER, + end_line INTEGER, + hash TEXT, + text TEXT, + embedding TEXT, + updated_at INTEGER + ) + """, + ) + + # Vector table (sqlite-vec) + if self.vector_available: + cursor.execute( + f""" + CREATE VIRTUAL TABLE IF NOT EXISTS {self.vector_table_name} USING vec0( + id TEXT PRIMARY KEY, + embedding FLOAT[{self.embedding_dim}] + ) + """, + ) + logger.info(f"Created vector table (dims={self.embedding_dim})") + + # FTS table + if self.fts_enabled: + cursor.execute( + f""" + CREATE VIRTUAL TABLE IF NOT EXISTS {self.fts_table_name} USING fts5( + text, + id UNINDEXED, + path UNINDEXED, + source UNINDEXED, + start_line UNINDEXED, + end_line UNINDEXED + ) + """, + ) + self.fts_available = True + logger.info("Created FTS5 table") + + self.conn.commit() + cursor.close() + + async def upsert_file(self, file_meta: FileMetadata, source: MemorySource, chunks: list[MemoryChunk]): + """Insert or update file and its chunks.""" + cursor = self.conn.cursor() + + try: + cursor.execute("BEGIN") + + # Insert file + cursor.execute( + f""" + INSERT OR REPLACE INTO {self.files_table_name} (path, source, hash, mtime, size) + VALUES (?, ?, ?, ?, ?) + """, + (file_meta.path, source.value, file_meta.hash, file_meta.mtime_ms, file_meta.size), + ) + + # Insert chunks + now = int(time.time() * 1000) + for chunk in chunks: + cursor.execute( + f""" + INSERT OR REPLACE INTO {self.chunks_table_name} ( + id, path, source, start_line, end_line, + hash, text, embedding, updated_at + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) + """, + ( + chunk.id, + file_meta.path, + source.value, + chunk.start_line, + chunk.end_line, + chunk.hash, + chunk.text, + json.dumps(chunk.embedding) if chunk.embedding else None, + now, + ), + ) + + # Insert vector + if self.vector_available: + assert chunk.embedding, "Embedding is required for vector insert" + cursor.execute( + f""" + INSERT OR REPLACE INTO {self.vector_table_name} (id, embedding) + VALUES (?, ?) + """, + (chunk.id, self.vector_to_blob(chunk.embedding)), + ) + + # Insert FTS + if self.fts_available: + cursor.execute( + f""" + INSERT OR REPLACE INTO {self.fts_table_name} ( + text, id, path, source, start_line, end_line + ) VALUES (?, ?, ?, ?, ?, ?) + """, + ( + chunk.text, + chunk.id, + file_meta.path, + source.value, + chunk.start_line, + chunk.end_line, + ), + ) + + cursor.execute("COMMIT") + except Exception: + cursor.execute("ROLLBACK") + raise + finally: + cursor.close() + + async def delete_file(self, path: str, source: MemorySource): + """Delete file and all its chunks.""" + cursor = self.conn.cursor() + try: + cursor.execute("BEGIN") + + # Get chunk IDs for vector deletion + cursor.execute( + f"SELECT id FROM {self.chunks_table_name} WHERE path = ? AND source = ?", + (path, source.value), + ) + chunk_ids = [row[0] for row in cursor.fetchall()] + + # Delete vectors + if self.vector_available and chunk_ids: + for chunk_id in chunk_ids: + try: + cursor.execute( + f"DELETE FROM {self.vector_table_name} WHERE id = ?", + (chunk_id,), + ) + except Exception as e: + logger.debug(f"Vector delete failed: {e}") + + # Delete FTS entries + if self.fts_available: + try: + cursor.execute( + f"DELETE FROM {self.fts_table_name} WHERE path = ? AND source = ?", + (path, source.value), + ) + except Exception as e: + logger.debug(f"FTS delete failed: {e}") + + # Delete chunks and file + cursor.execute( + f"DELETE FROM {self.chunks_table_name} WHERE path = ? AND source = ?", + (path, source.value), + ) + cursor.execute( + f"DELETE FROM {self.files_table_name} WHERE path = ? AND source = ?", + (path, source.value), + ) + + cursor.execute("COMMIT") + except Exception: + cursor.execute("ROLLBACK") + raise + finally: + cursor.close() + + async def delete_file_chunks(self, path: str, chunk_ids: list[str]): + """Delete specific chunks for a file.""" + if not chunk_ids: + return + + cursor = self.conn.cursor() + try: + cursor.execute("BEGIN") + + # Delete vectors + if self.vector_available: + for chunk_id in chunk_ids: + try: + cursor.execute( + f"DELETE FROM {self.vector_table_name} WHERE id = ?", + (chunk_id,), + ) + except Exception as e: + logger.debug(f"Vector delete failed for {chunk_id}: {e}") + + # Delete FTS entries + if self.fts_available: + placeholders = ",".join("?" * len(chunk_ids)) + try: + cursor.execute( + f"DELETE FROM {self.fts_table_name} WHERE id IN ({placeholders})", + chunk_ids, + ) + except Exception as e: + logger.debug(f"FTS delete failed: {e}") + + # Delete chunks + placeholders = ",".join("?" * len(chunk_ids)) + cursor.execute( + f"DELETE FROM {self.chunks_table_name} WHERE id IN ({placeholders})", + chunk_ids, + ) + + cursor.execute("COMMIT") + except Exception: + cursor.execute("ROLLBACK") + raise + finally: + cursor.close() + + async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource): + """Insert or update specific chunks without affecting other chunks.""" + if not chunks: + return + + cursor = self.conn.cursor() + try: + cursor.execute("BEGIN") + + now = int(time.time() * 1000) + for chunk in chunks: + # Insert/update chunk + cursor.execute( + f""" + INSERT OR REPLACE INTO {self.chunks_table_name} ( + id, path, source, start_line, end_line, + hash, text, embedding, updated_at + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) + """, + ( + chunk.id, + chunk.path, + source.value, + chunk.start_line, + chunk.end_line, + chunk.hash, + chunk.text, + json.dumps(chunk.embedding) if chunk.embedding else None, + now, + ), + ) + + # Insert/update vector + if self.vector_available: + assert chunk.embedding, "Embedding is required for vector insert" + cursor.execute( + f""" + INSERT OR REPLACE INTO {self.vector_table_name} (id, embedding) + VALUES (?, ?) + """, + (chunk.id, self.vector_to_blob(chunk.embedding)), + ) + + # Insert/update FTS + if self.fts_available: + cursor.execute( + f""" + INSERT OR REPLACE INTO {self.fts_table_name} ( + text, id, path, source, start_line, end_line + ) VALUES (?, ?, ?, ?, ?, ?) + """, + ( + chunk.text, + chunk.id, + chunk.path, + source.value, + chunk.start_line, + chunk.end_line, + ), + ) + + cursor.execute("COMMIT") + except Exception: + cursor.execute("ROLLBACK") + raise + finally: + cursor.close() + + async def list_files(self, source: MemorySource) -> list[str]: + """List all indexed files.""" + cursor = self.conn.cursor() + cursor.execute(f"SELECT path FROM {self.files_table_name} WHERE source = ?", (source.value,)) + paths = [row[0] for row in cursor.fetchall()] + cursor.close() + return paths + + async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None: + """Get file metadata with chunk count.""" + cursor = self.conn.cursor() + cursor.execute( + f"SELECT hash, mtime, size FROM {self.files_table_name} WHERE path = ? AND source = ?", + (path, source.value), + ) + row = cursor.fetchone() + if not row: + cursor.close() + return None + + hash_val, mtime, size = row + cursor.execute( + f"SELECT COUNT(*) FROM {self.chunks_table_name} WHERE path = ? AND source = ?", + (path, source.value), + ) + chunk_count = cursor.fetchone()[0] + cursor.close() + + return FileMetadata( + hash=hash_val, + mtime_ms=mtime, + size=size, + path=path, + chunk_count=chunk_count, + ) + + async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]: + """Get all chunks for a file.""" + cursor = self.conn.cursor() + cursor.execute( + f""" + SELECT id, path, source, start_line, end_line, text, hash, embedding + FROM {self.chunks_table_name} WHERE path = ? AND source = ? + ORDER BY start_line + """, + (path, source.value), + ) + + chunks = [] + for row in cursor.fetchall(): + chunk_id, path_val, source_val, start, end, text, hash_val, emb_str = row + # Parse embedding from JSON string + embedding = None + if emb_str: + try: + embedding = json.loads(emb_str) + except (json.JSONDecodeError, TypeError): + embedding = None + + chunks.append( + MemoryChunk( + id=chunk_id, + path=path_val, + source=MemorySource(source_val), + start_line=start, + end_line=end, + text=text, + hash=hash_val, + embedding=embedding, + ), + ) + + cursor.close() + return chunks + + async def vector_search( + self, + query: str, + limit: int, + sources: list[MemorySource] | None = None, + ) -> list[MemorySearchResult]: + """Perform vector similarity search.""" + if not self.vector_available or not query: + return [] + + # Get query embedding + query_embedding = await self.get_embedding(query) + if not query_embedding: + return [] + + cursor = self.conn.cursor() + source_filter = "" + params: list = [] + if sources: + placeholders = ",".join("?" * len(sources)) + source_filter = f" AND c.source IN ({placeholders})" + params = [s.value for s in sources] + + try: + query_blob = self.vector_to_blob(query_embedding) + + # Correct SQLite-vec syntax for vector search with limit + # vec0 requires 'k = ?' constraint for knn queries + query_sql = f""" + SELECT c.id, c.path, c.start_line, c.end_line, c.source, c.text, v.distance + FROM {self.vector_table_name} v + JOIN {self.chunks_table_name} c ON v.id = c.id + WHERE v.embedding MATCH ? + AND k = ? + """ + query_params: list = [query_blob, limit] + + # Add source filter if specified + if source_filter: + query_sql += source_filter + query_params.extend(params) + + # Order by distance (k constraint already limits results) + query_sql += " ORDER BY v.distance" + + cursor.execute(query_sql, query_params) + + results = [] + for _, path, start, end, src, text, dist in cursor.fetchall(): + score = max(0.0, 1.0 - dist) + snippet = text[: self.snippet_max_chars] if len(text) > self.snippet_max_chars else text + results.append( + MemorySearchResult( + path=path, + start_line=start, + end_line=end, + score=score, + snippet=snippet, + source=MemorySource(src), + ), + ) + + return results + except Exception as e: + logger.error(f"Vector search failed: {e}") + return [] + finally: + cursor.close() + + async def keyword_search( + self, + query: str, + limit: int, + sources: list[MemorySource] | None = None, + ) -> list[MemorySearchResult]: + """Perform full-text search.""" + if not self.fts_available: + return [] + + # Build FTS5 query + # Split query into tokens and join with OR for better recall + # Individual words are automatically stemmed and matched by FTS5 + cleaned = query.strip() + if not cleaned: + return [] + + # Split into words and escape each + words = cleaned.split() + if not words: + return [] + + # Use OR operator for better recall - match any of the query words + escaped_words = [word.replace('"', '""') for word in words] + fts_query = " OR ".join(escaped_words) + + cursor = self.conn.cursor() + source_filter = "" + params: list = [fts_query] + if sources: + placeholders = ",".join("?" * len(sources)) + source_filter = f" AND fts.source IN ({placeholders})" + params.extend([s.value for s in sources]) + params.append(limit) + + try: + cursor.execute( + f""" + SELECT fts.id, fts.path, fts.start_line, fts.end_line, + fts.source, fts.text, rank + FROM {self.fts_table_name} fts + WHERE fts.text MATCH ?{source_filter} + ORDER BY rank + LIMIT ? + """, + params, + ) + + results = [] + for _, path, start, end, src, text, rank in cursor.fetchall(): + # Convert BM25 rank (negative) to 0-1 score (higher=better) + score = max(0.0, 1.0 / (1.0 + abs(rank))) + snippet = text[: self.snippet_max_chars] if len(text) > self.snippet_max_chars else text + results.append( + MemorySearchResult( + path=path, + start_line=start, + end_line=end, + score=score, + snippet=snippet, + source=MemorySource(src), + ), + ) + + return results + except Exception as e: + logger.error(f"Keyword search failed: {e}") + return [] + finally: + cursor.close() + + async def clear_all(self): + """Clear all indexed data.""" + cursor = self.conn.cursor() + cursor.execute("BEGIN") + + try: + cursor.execute(f"DELETE FROM {self.files_table_name}") + cursor.execute(f"DELETE FROM {self.chunks_table_name}") + + if self.vector_available: + cursor.execute(f"DELETE FROM {self.vector_table_name}") + + if self.fts_available: + cursor.execute(f"DELETE FROM {self.fts_table_name}") + + cursor.execute("COMMIT") + except Exception: + cursor.execute("ROLLBACK") + raise + finally: + cursor.close() + + async def close(self): + """Close database connection.""" + if self.conn: + self.conn.close() + self.conn = None diff --git a/reme/core/op/base_op.py b/reme/core/op/base_op.py index 6d32ddc0..015d0fbb 100644 --- a/reme/core/op/base_op.py +++ b/reme/core/op/base_op.py @@ -13,6 +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 ..schema import Response from ..token_counter import BaseTokenCounter from ..utils import camel_to_snake, CacheHandler, timer @@ -41,6 +42,7 @@ class BaseOp(metaclass=ABCMeta): llm: str | BaseLLM = "default", embedding_model: str | BaseEmbeddingModel = "default", vector_store: str | BaseVectorStore = "default", + memory_store: str | BaseMemoryStore = "default", token_counter: str | BaseTokenCounter = "default", enable_cache: bool = False, cache_path: str = "cache/op", @@ -62,6 +64,7 @@ class BaseOp(metaclass=ABCMeta): self._llm = llm self._embedding_model = embedding_model self._vector_store = vector_store + self._memory_store = memory_store self._token_counter = token_counter self.enable_cache = enable_cache @@ -139,6 +142,13 @@ class BaseOp(metaclass=ABCMeta): self._vector_store = self.service_context.vector_stores[self._vector_store] return self._vector_store + @property + def memory_store(self) -> BaseMemoryStore: + """Lazily initialize and return the memory store instance.""" + if isinstance(self._memory_store, str): + self._memory_store = self.service_context.memory_stores[self._memory_store] + return self._memory_store + @property def token_counter(self) -> BaseTokenCounter: """Get the token counter instance from ServiceContext.""" diff --git a/reme/core/schema/__init__.py b/reme/core/schema/__init__.py index b7b73719..a4a4214e 100644 --- a/reme/core/schema/__init__.py +++ b/reme/core/schema/__init__.py @@ -1,6 +1,10 @@ """schema""" +from .file_metadata import FileMetadata +from .memory_chunk import MemoryChunk +from .memory_index_meta import MemoryIndexMeta from .memory_node import MemoryNode +from .memory_search_result import MemorySearchResult from .message import ContentBlock, Message, Trajectory from .request import Request from .response import Response @@ -17,16 +21,22 @@ from .service_config import ( ) from .stream_chunk import StreamChunk from .tool_call import ToolAttr, ToolCall +from .truncation_result import TruncationResult from .vector_node import VectorNode __all__ = [ - "MemoryNode", + "CmdConfig", "ContentBlock", "EmbeddingModelConfig", + "FileMetadata", "FlowConfig", "HttpConfig", "LLMConfig", "MCPConfig", + "MemoryChunk", + "MemoryIndexMeta", + "MemoryNode", + "MemorySearchResult", "Message", "Request", "Response", @@ -36,7 +46,7 @@ __all__ = [ "Trajectory", "ToolAttr", "ToolCall", + "TruncationResult", "VectorNode", "VectorStoreConfig", - "CmdConfig", ] diff --git a/reme/core/schema/file_metadata.py b/reme/core/schema/file_metadata.py new file mode 100644 index 00000000..672e9c27 --- /dev/null +++ b/reme/core/schema/file_metadata.py @@ -0,0 +1,15 @@ +"""File metadata schema.""" + +from pydantic import BaseModel, Field + + +class FileMetadata(BaseModel): + """File metadata with optional extended fields for various use cases.""" + + hash: str = Field(default=..., description="Hash of the file content") + mtime_ms: float = Field(default=..., description="Last modification time in milliseconds") + size: int = Field(default=..., description="File size in bytes") + path: str | None = Field(default=None, description="Relative path to the session file") + content: str | None = Field(default=None, description="Parsed content from the session file") + chunk_count: int | None = Field(default=None, description="Number of chunks in the file") + metadata: dict = Field(default_factory=dict, description="Additional metadata") diff --git a/reme/core/schema/memory_chunk.py b/reme/core/schema/memory_chunk.py new file mode 100644 index 00000000..36959da1 --- /dev/null +++ b/reme/core/schema/memory_chunk.py @@ -0,0 +1,19 @@ +"""Memory chunk schema.""" + +from pydantic import BaseModel, Field + +from ..enumeration import MemorySource + + +class MemoryChunk(BaseModel): + """A chunk of memory content with metadata.""" + + id: str = Field(..., description="Unique identifier for the chunk") + path: str = Field(..., description="File path relative to workspace") + source: MemorySource = Field(..., description="Source of the memory data") + start_line: int = Field(..., description="Starting line number in the source file") + end_line: int = Field(..., description="Ending line number in the source file") + text: str = Field(..., description="Text content of the chunk") + hash: str = Field(..., description="Hash of the chunk content") + embedding: list[float] | None = Field(default=None, description="Vector embedding of the chunk") + metadata: dict = Field(default_factory=dict, description="Additional metadata") diff --git a/reme/core/schema/memory_index_meta.py b/reme/core/schema/memory_index_meta.py new file mode 100644 index 00000000..a4538b77 --- /dev/null +++ b/reme/core/schema/memory_index_meta.py @@ -0,0 +1,14 @@ +"""Memory index metadata schema.""" + +from typing import Optional + +from pydantic import BaseModel, Field + + +class MemoryIndexMeta(BaseModel): + """Metadata for memory index configuration.""" + + model: str = Field(..., description="Name of the embedding model") + chunk_tokens: int = Field(..., description="Maximum tokens per chunk") + chunk_overlap: int = Field(..., description="Number of overlapping tokens between chunks") + vector_dims: Optional[int] = Field(default=None, description="Vector embedding dimensions") diff --git a/reme/core/schema/memory_search_result.py b/reme/core/schema/memory_search_result.py new file mode 100644 index 00000000..d08ca12b --- /dev/null +++ b/reme/core/schema/memory_search_result.py @@ -0,0 +1,24 @@ +"""Memory search result schema.""" + +from typing import Any, Dict + +from pydantic import BaseModel, Field + +from ..enumeration import MemorySource + + +class MemorySearchResult(BaseModel): + """Search result from memory index.""" + + path: str = Field(..., description="File path relative to workspace") + start_line: int = Field(..., description="Starting line number of the match") + end_line: int = Field(..., description="Ending line number of the match") + 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") + + @property + def merge_key(self) -> str: + """Merge key for the search result.""" + return self.path + f":{self.start_line}:{self.end_line}" diff --git a/reme/core/schema/service_config.py b/reme/core/schema/service_config.py index 95461196..16a70b3e 100644 --- a/reme/core/schema/service_config.py +++ b/reme/core/schema/service_config.py @@ -77,6 +77,16 @@ class VectorStoreConfig(BaseModel): embedding_model: str = Field(default="default") +class MemoryStoreConfig(BaseModel): + """Configuration for memory database storage and associated embeddings.""" + + model_config = ConfigDict(extra="allow") + + backend: str = Field(default="sqlite") + store_name: str = Field(default="reme") + embedding_model: str = Field(default="default") + + class TokenCounterConfig(BaseModel): """Configuration for token counting services and model mapping.""" @@ -109,4 +119,5 @@ class ServiceConfig(BaseModel): llm: dict[str, LLMConfig] = Field(default_factory=dict) embedding_model: dict[str, EmbeddingModelConfig] = Field(default_factory=dict) vector_store: dict[str, VectorStoreConfig] = Field(default_factory=dict) + memory_store: dict[str, MemoryStoreConfig] = Field(default_factory=dict) token_counter: dict[str, TokenCounterConfig] = Field(default_factory=dict) diff --git a/reme/core/schema/truncation_result.py b/reme/core/schema/truncation_result.py new file mode 100644 index 00000000..18e8cc1a --- /dev/null +++ b/reme/core/schema/truncation_result.py @@ -0,0 +1,35 @@ +"""Truncation result schema for command output truncation.""" + +from typing import Literal + +from pydantic import BaseModel, Field + + +class TruncationResult(BaseModel): + """Result of output truncation operation. + + Attributes: + content: The truncated content + truncated: Whether truncation occurred + total_lines: Total number of lines in original output + output_lines: Number of lines in truncated output + total_bytes: Total bytes in original output + output_bytes: Bytes in truncated output + truncated_by: What caused truncation ('lines' or 'bytes') + last_line_partial: Whether last line was partially truncated + """ + + content: str = Field(description="The truncated content") + truncated: bool = Field(description="Whether truncation occurred") + total_lines: int = Field(description="Total number of lines in original output") + output_lines: int = Field(description="Number of lines in truncated output") + total_bytes: int = Field(description="Total bytes in original output") + output_bytes: int = Field(description="Bytes in truncated output") + truncated_by: Literal["lines", "bytes"] | None = Field( + default=None, + description="What caused truncation ('lines' or 'bytes')", + ) + last_line_partial: bool = Field( + default=False, + description="Whether last line was partially truncated", + ) diff --git a/reme/core/utils/__init__.py b/reme/core/utils/__init__.py index 3d54e4e4..642238c0 100644 --- a/reme/core/utils/__init__.py +++ b/reme/core/utils/__init__.py @@ -2,7 +2,8 @@ from .cache_handler import CacheHandler from .case_converter import snake_to_camel, camel_to_snake -from .common_utils import run_coro_safely, execute_stream_task +from .chunking_utils import chunk_markdown +from .common_utils import run_coro_safely, execute_stream_task, hash_text from .env_utils import load_env from .execute_utils import exec_code, run_shell_command from .http_client import HttpClient @@ -19,8 +20,10 @@ __all__ = [ "CacheHandler", "snake_to_camel", "camel_to_snake", + "chunk_markdown", "run_coro_safely", "execute_stream_task", + "hash_text", "load_env", "exec_code", "run_shell_command", diff --git a/reme/core/utils/chunking_utils.py b/reme/core/utils/chunking_utils.py new file mode 100644 index 00000000..e1511200 --- /dev/null +++ b/reme/core/utils/chunking_utils.py @@ -0,0 +1,124 @@ +"""Chunking logic for Markdown files.""" + +from .common_utils import hash_text +from ..enumeration import MemorySource +from ..schema import MemoryChunk + + +def chunk_markdown( + text: str, + path: str, + source: MemorySource, + chunk_tokens: int, + overlap: int, +) -> list[MemoryChunk]: + """ + Markdown chunking logic implemented based on the TypeScript version. + + Args: + text: Input text + path: File path + source: Memory source + chunk_tokens: Maximum tokens per chunk + overlap: Overlap tokens between chunks + + Returns: + List of MemoryChunk objects + """ + lines = text.split("\n") + if not lines: + return [] + + # Convert tokens to characters (~1 token = 4 chars) + max_chars = max(32, chunk_tokens * 4) + overlap_chars = max(0, overlap * 4) + + chunks: list[MemoryChunk] = [] + + # Currently building chunk + current: list[dict] = [] # [{'line': str, 'line_no': int}] + current_chars = 0 + + def flush(): + """Add current chunk to results list""" + if not current: + return + + first_entry = current[0] + last_entry = current[-1] + + if not first_entry or not last_entry: + return + + chunk_text = "\n".join([entry["line"] for entry in current]) + start_line = first_entry["line_no"] + end_line = last_entry["line_no"] + + chunk_hash = hash_text(chunk_text) + + chunks.append( + MemoryChunk( + id=hash_text(f"{source}:{path}:{start_line}:{end_line}:{chunk_hash}:{len(chunks)}"), + path=path, + source=source, + start_line=start_line, + end_line=end_line, + text=chunk_text, + hash=chunk_hash, + ), + ) + + def carry_overlap(): + """Keep overlapping part and clear the rest""" + nonlocal current, current_chars + + if overlap_chars <= 0 or not current: + current = [] + current_chars = 0 + return + + acc = 0 + kept = [] + + # Collect lines from the end until reaching overlap size + for j in range(len(current) - 1, -1, -1): + entry = current[j] + if not entry: + continue + + acc += len(entry["line"]) + 1 # +1 for newline + kept.insert(0, entry) # Insert at the beginning to maintain order + + if acc >= overlap_chars: + break + + current = kept + current_chars = sum(len(entry["line"]) + 1 for entry in kept) + + for i, line in enumerate(lines): + line_no = i + 1 + + # Split long lines into multiple segments + segments = [] + if not line: # Empty line + segments.append("") + else: + # If line is too long, split by maximum character count + for start in range(0, len(line), max_chars): + segments.append(line[start : start + max_chars]) + + for segment in segments: + line_size = len(segment) + 1 # +1 for newline + + # If adding current segment would exceed the limit, flush current chunk + if current_chars + line_size > max_chars and current: + flush() + carry_overlap() + + current.append({"line": segment, "line_no": line_no}) + current_chars += line_size + + # Process the final chunk + flush() + + return [c for c in chunks if c.text.strip()] diff --git a/reme/core/utils/common_utils.py b/reme/core/utils/common_utils.py index 3f171f64..bf35d3db 100644 --- a/reme/core/utils/common_utils.py +++ b/reme/core/utils/common_utils.py @@ -1,6 +1,7 @@ """Common utility functions""" import asyncio +import hashlib from collections.abc import AsyncGenerator, Coroutine from typing import Any @@ -81,3 +82,15 @@ async def execute_stream_task( # Ensure task is cancelled if still running to avoid resource leaks if not task.done(): task.cancel() + + +def hash_text(text: str) -> str: + """Generate SHA-256 hash of text content. + + Args: + text: Input text to hash + + Returns: + Hexadecimal representation of the SHA-256 hash + """ + return hashlib.sha256(text.encode("utf-8")).hexdigest() diff --git a/reme/core/utils/mcp_client.py b/reme/core/utils/mcp_client.py index 7f5a1514..4c13be9f 100644 --- a/reme/core/utils/mcp_client.py +++ b/reme/core/utils/mcp_client.py @@ -8,7 +8,7 @@ from typing import Any from mcp import ClientSession, StdioServerParameters, Tool from mcp.client.sse import sse_client from mcp.client.stdio import stdio_client -from mcp.client.streamable_http import streamable_http_client +from mcp.client.streamable_http import streamablehttp_client from mcp.types import CallToolResult, TextContent from ..schema import ToolCall @@ -64,7 +64,7 @@ class MCPClient: async with sse_client(**cfg) as transport: yield transport elif t_type == "streamable-http": - async with streamable_http_client(**cfg) as transport: + async with streamablehttp_client(**cfg) as transport: yield transport else: raise NotImplementedError(f"Unsupported transport: {t_type}") diff --git a/reme/reme.py b/reme/reme.py index 8dfc4b51..30eccf90 100644 --- a/reme/reme.py +++ b/reme/reme.py @@ -20,7 +20,7 @@ from .agent.memory import ( ) from .config import ReMeConfigParser from .core import Application -from .core.enumeration import MemoryType +from .core.enumeration import MemoryType, Role from .core.schema import Message, MemoryNode from .tool.memory import ( RetrieveMemory, @@ -59,7 +59,7 @@ class ReMe(Application): target_user_names: list[str] | None = None, target_task_names: list[str] | None = None, target_tool_names: list[str] | None = None, - profile_dir: str = "reme_profile", + profile_dir: str = ".reme/profile", **kwargs, ): """Initialize ReMe with config. @@ -83,11 +83,14 @@ class ReMe(Application): Example: ```python - reme = await ReMe(...).start() + reme = ReMe(...) + await reme.start() # reme = await ReMe.create(...) # both ok await reme.summarize_memory(...) await reme.retrieve_memory(...) + + await reme.close() ``` """ @@ -308,7 +311,8 @@ class ReMe(Application): if user_name: if isinstance(user_name, str): for message in format_messages: - message.name = user_name + if message.role is Role.USER: + message.name = user_name self._add_meta_memory(MemoryType.PERSONAL, user_name) memory_targets.append(user_name) elif isinstance(user_name, list): diff --git a/reme/tool/fs/__init__.py b/reme/tool/fs/__init__.py new file mode 100644 index 00000000..090e1a1f --- /dev/null +++ b/reme/tool/fs/__init__.py @@ -0,0 +1,24 @@ +"""File system tools.""" + +from .bash_tool import BashTool +from .edit_tool import EditTool +from .find_tool import FindTool +from .grep_tool import GrepTool +from .ls_tool import LsTool +from .read_tool import ReadTool +from .write_tool import WriteTool +from ...core import R + +__all__ = [ + "BashTool", + "EditTool", + "FindTool", + "GrepTool", + "LsTool", + "ReadTool", + "WriteTool", +] + +for name in __all__: + tool_class = globals()[name] + R.op.register(tool_class) diff --git a/reme/tool/fs/bash_tool.py b/reme/tool/fs/bash_tool.py new file mode 100644 index 00000000..d8aea81c --- /dev/null +++ b/reme/tool/fs/bash_tool.py @@ -0,0 +1,191 @@ +"""Bash command execution tool with production-grade features. + +This module provides a production-grade tool for executing bash commands with: +- Smart output truncation (keeps last N lines/bytes to prevent memory issues) +- Process tree termination (prevents orphan processes) +""" + +import asyncio +import os +import platform +import signal +from pathlib import Path + +from .truncate import DEFAULT_MAX_BYTES, DEFAULT_MAX_LINES, truncate_tail +from ...core.op import BaseTool +from ...core.schema import ToolCall, TruncationResult + + +def get_shell_config() -> tuple[str, list[str]]: + """Get the appropriate shell and arguments for the current platform. + + Returns: + Tuple of (shell_path, args) for subprocess execution + """ + system = platform.system() + + if system == "Windows": + # Use PowerShell on Windows + return "powershell.exe", ["-Command"] + else: + # Use bash on Unix-like systems + shell = os.environ.get("SHELL", "/bin/bash") + return shell, ["-c"] + + +def kill_process_tree(pid: int) -> None: + """Kill a process and all its children. + + Args: + pid: Process ID to kill + """ + try: + if platform.system() == "Windows": + # Windows: use taskkill + os.system(f"taskkill /F /T /PID {pid}") + else: + # Unix: kill process group + try: + os.killpg(os.getpgid(pid), signal.SIGTERM) + except ProcessLookupError: + pass # Process already dead + except Exception: + pass # Best effort + + +class BashTool(BaseTool): + """Production-grade tool for executing bash commands. + + Features: + - Smart output truncation (preserves last N lines or M bytes) + - Kills entire process tree on timeout (prevents orphan processes) + """ + + def __init__(self, cwd: str | None = None, command_prefix: str | None = None): + """Initialize bash tool. + + Args: + cwd: Working directory (defaults to current directory) + command_prefix: Optional prefix prepended to every command + """ + super().__init__() + self.cwd = cwd or os.getcwd() + self.command_prefix = command_prefix + + def _build_tool_call(self) -> ToolCall: + max_kb = DEFAULT_MAX_BYTES // 1024 + return ToolCall( + **{ + "description": ( + f"Execute a bash command in the current working directory. " + f"Returns stdout and stderr. Output is truncated to last " + f"{DEFAULT_MAX_LINES} lines or {max_kb}KB (whichever is hit first). " + f"Optionally provide a timeout in seconds." + ), + "parameters": { + "type": "object", + "properties": { + "command": { + "type": "string", + "description": "Bash command to execute", + }, + "timeout": { + "type": "number", + "description": "Timeout in seconds (optional, no default timeout)", + }, + }, + "required": ["command"], + }, + }, + ) + + async def execute(self) -> str: + """Execute the bash command with production-grade features.""" + command: str = self.context.command + timeout: float | None = self.context.get("timeout", None) + + # Apply command prefix if configured + if self.command_prefix: + command = f"{self.command_prefix}\n{command}" + + # Verify working directory exists + if not Path(self.cwd).exists(): + raise FileNotFoundError( + f"Working directory does not exist: {self.cwd}\n" f"Cannot execute bash commands.", + ) + + # Get shell configuration + 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 + + # Execute command with optional timeout + try: + 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: + 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 + + # Decode output + full_output = stdout.decode("utf-8", errors="ignore") + if stderr: + stderr_text = stderr.decode("utf-8", errors="ignore") + if full_output: + full_output += "\n" + full_output += stderr_text + + # Apply tail truncation_result to prevent memory issues + truncation_result: TruncationResult = truncate_tail(full_output) + output_text = truncation_result.content or "(no output)" + + # Build truncation_result notice if needed + if truncation_result.truncated: + start_line = truncation_result.total_lines - truncation_result.output_lines + 1 + end_line = truncation_result.total_lines + + if truncation_result.truncated_by == "lines": + output_text += ( + f"\n\n[Output truncated: showing lines {start_line}-{end_line} " + f"of {truncation_result.total_lines} total lines]" + ) + else: + max_kb = DEFAULT_MAX_BYTES // 1024 + output_text += ( + f"\n\n[Output truncated: showing lines {start_line}-{end_line} " + f"of {truncation_result.total_lines} ({max_kb}KB limit reached)]" + ) + + # Handle non-zero exit code + if process.returncode != 0: + output_text += f"\n\nCommand exited with code {process.returncode}" + raise RuntimeError(output_text) + + return output_text diff --git a/reme/tool/fs/edit_diff.py b/reme/tool/fs/edit_diff.py new file mode 100644 index 00000000..76e2fcdb --- /dev/null +++ b/reme/tool/fs/edit_diff.py @@ -0,0 +1,164 @@ +"""Diff utilities for edit tool.""" + +import re +from dataclasses import dataclass +from difflib import unified_diff + + +def detect_line_ending(content: str) -> str: + """Detect line ending style (CRLF or LF).""" + crlf_idx = content.find("\r\n") + lf_idx = content.find("\n") + if lf_idx == -1: + return "\n" + if crlf_idx == -1: + return "\n" + return "\r\n" if crlf_idx < lf_idx else "\n" + + +def normalize_to_lf(text: str) -> str: + """Normalize line endings to LF.""" + return text.replace("\r\n", "\n").replace("\r", "\n") + + +def restore_line_endings(text: str, ending: str) -> str: + """Restore original line endings.""" + return text.replace("\n", ending) if ending == "\r\n" else text + + +def normalize_for_fuzzy_match(text: str) -> str: + """Normalize text for fuzzy matching: strip trailing whitespace, normalize quotes/dashes.""" + lines = text.split("\n") + normalized = "\n".join(line.rstrip() for line in lines) + + # Smart quotes → ASCII + normalized = re.sub(r"[\u2018\u2019\u201A\u201B]", "'", normalized) + normalized = re.sub(r"[\u201C\u201D\u201E\u201F]", '"', normalized) + + # Dashes → hyphen + normalized = re.sub(r"[\u2010\u2011\u2012\u2013\u2014\u2015\u2212]", "-", normalized) + + # Special spaces → regular space + normalized = re.sub(r"[\u00A0\u2002-\u200A\u202F\u205F\u3000]", " ", normalized) + + return normalized + + +@dataclass +class FuzzyMatchResult: + """Result of fuzzy text matching.""" + + found: bool + index: int + match_length: int + used_fuzzy_match: bool + content_for_replacement: str + + +def fuzzy_find_text(content: str, old_text: str) -> FuzzyMatchResult: + """Find old_text in content, trying exact match first, then fuzzy match.""" + # Try exact match + exact_index = content.find(old_text) + if exact_index != -1: + return FuzzyMatchResult( + found=True, + index=exact_index, + match_length=len(old_text), + used_fuzzy_match=False, + content_for_replacement=content, + ) + + # Try fuzzy match + fuzzy_content = normalize_for_fuzzy_match(content) + fuzzy_old_text = normalize_for_fuzzy_match(old_text) + fuzzy_index = fuzzy_content.find(fuzzy_old_text) + + if fuzzy_index == -1: + return FuzzyMatchResult( + found=False, + index=-1, + match_length=0, + used_fuzzy_match=False, + content_for_replacement=content, + ) + + return FuzzyMatchResult( + found=True, + index=fuzzy_index, + match_length=len(fuzzy_old_text), + used_fuzzy_match=True, + content_for_replacement=fuzzy_content, + ) + + +def strip_bom(content: str) -> tuple[str, str]: + """Strip UTF-8 BOM, return (bom, text_without_bom).""" + if content.startswith("\ufeff"): + return "\ufeff", content[1:] + return "", content + + +@dataclass +class DiffResult: + """Result of diff generation.""" + + diff: str + first_changed_line: int | None + + +def generate_diff_string(old_content: str, new_content: str, context_lines: int = 4) -> DiffResult: + """Generate unified diff with line numbers.""" + old_lines = old_content.split("\n") + new_lines = new_content.split("\n") + + # Use difflib to get the changes + diff_lines = list( + unified_diff( + old_lines, + new_lines, + lineterm="", + n=context_lines, + ), + ) + + if not diff_lines: + return DiffResult(diff="", first_changed_line=None) + + # Parse and format the diff + output = [] + first_changed_line = None + max_line_num = max(len(old_lines), len(new_lines)) + line_num_width = len(str(max_line_num)) + + old_line_num = 1 + new_line_num = 1 + + for line in diff_lines[2:]: # Skip header lines + if line.startswith("@@"): + # Parse hunk header + match = re.match(r"@@ -(\d+),?\d* \+(\d+),?\d* @@", line) + if match: + old_line_num = int(match.group(1)) + new_line_num = int(match.group(2)) + continue + + if line.startswith("+"): + if first_changed_line is None: + first_changed_line = new_line_num + line_num = str(new_line_num).rjust(line_num_width) + output.append(f"+{line_num} {line[1:]}") + new_line_num += 1 + elif line.startswith("-"): + if first_changed_line is None: + first_changed_line = new_line_num + line_num = str(old_line_num).rjust(line_num_width) + output.append(f"-{line_num} {line[1:]}") + old_line_num += 1 + else: + # Context line + line_num = str(old_line_num).rjust(line_num_width) + output.append(f" {line_num} {line[1:] if line.startswith(' ') else line}") + old_line_num += 1 + new_line_num += 1 + + return DiffResult(diff="\n".join(output), first_changed_line=first_changed_line) diff --git a/reme/tool/fs/edit_tool.py b/reme/tool/fs/edit_tool.py new file mode 100644 index 00000000..fe73bd1f --- /dev/null +++ b/reme/tool/fs/edit_tool.py @@ -0,0 +1,141 @@ +"""File editing tool with exact text replacement.""" + +import os +from pathlib import Path + +from .edit_diff import ( + detect_line_ending, + fuzzy_find_text, + generate_diff_string, + normalize_for_fuzzy_match, + normalize_to_lf, + restore_line_endings, + strip_bom, +) +from ...core.op import BaseTool +from ...core.schema import ToolCall + + +class EditTool(BaseTool): + """Edit a file by replacing exact text.""" + + def __init__(self, cwd: str | None = None): + """Initialize edit tool. + + Args: + cwd: Working directory (defaults to current directory) + """ + super().__init__() + self.cwd = cwd or os.getcwd() + + def _build_tool_call(self) -> ToolCall: + return ToolCall( + **{ + "description": ( + "Edit a file by replacing exact text. The oldText must match exactly " + "(including whitespace). Use this for precise, surgical edits." + ), + "parameters": { + "type": "object", + "properties": { + "path": { + "type": "string", + "description": "Path to the file to edit (relative or absolute)", + }, + "oldText": { + "type": "string", + "description": "Exact text to find and replace (must match exactly)", + }, + "newText": { + "type": "string", + "description": "New text to replace the old text with", + }, + }, + "required": ["path", "oldText", "newText"], + }, + }, + ) + + async def execute(self) -> str: + """Execute the edit operation.""" + path: str = self.context.path + old_text: str = self.context.oldText + new_text: str = self.context.newText + + # Resolve path + if not os.path.isabs(path): + absolute_path = os.path.join(self.cwd, path) + else: + absolute_path = path + + # Check file exists and is writable + path_obj = Path(absolute_path) + if not path_obj.exists(): + raise FileNotFoundError(f"File not found: {path}") + + if not os.access(absolute_path, os.R_OK | os.W_OK): + 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 + + # Strip BOM (LLM won't include invisible BOM in oldText) + bom, content = strip_bom(raw_content) + + original_ending = detect_line_ending(content) + normalized_content = normalize_to_lf(content) + normalized_old_text = normalize_to_lf(old_text) + normalized_new_text = normalize_to_lf(new_text) + + # Find old text using fuzzy matching + match_result = fuzzy_find_text(normalized_content, normalized_old_text) + + if not match_result.found: + raise ValueError( + f"Could not find the exact text in {path}. The old text must match " + f"exactly including all whitespace and newlines.", + ) + + # Count occurrences for uniqueness check + fuzzy_content = normalize_for_fuzzy_match(normalized_content) + fuzzy_old_text = normalize_for_fuzzy_match(normalized_old_text) + occurrences = fuzzy_content.count(fuzzy_old_text) + + if occurrences > 1: + raise ValueError( + f"Found {occurrences} occurrences of the text in {path}. " + f"The text must be unique. Please provide more context to make it unique.", + ) + + # Perform replacement + base_content = match_result.content_for_replacement + new_content = ( + base_content[: match_result.index] + + normalized_new_text + + base_content[match_result.index + match_result.match_length :] + ) + + # Verify replacement changed something + if base_content == new_content: + raise ValueError( + f"No changes made to {path}. The replacement produced identical content. " + f"This might indicate an issue with special characters or the text not " + f"exist as expected.", + ) + + # 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 + + # Generate diff + diff_result = generate_diff_string(base_content, new_content) + + return f"Successfully replaced text in {path}.\n\n{diff_result.diff}" diff --git a/reme/tool/fs/find_tool.py b/reme/tool/fs/find_tool.py new file mode 100644 index 00000000..7b855586 --- /dev/null +++ b/reme/tool/fs/find_tool.py @@ -0,0 +1,183 @@ +"""File search tool using glob patterns with gitignore support.""" + +import os +from pathlib import Path + +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): + """Search for files by glob pattern, respecting .gitignore.""" + + def __init__(self, cwd: str | None = None): + """Initialize find tool. + + Args: + cwd: Working directory (defaults to current directory) + """ + super().__init__() + self.cwd = cwd or os.getcwd() + + def _build_tool_call(self) -> ToolCall: + max_kb = FIND_MAX_BYTES // 1024 + return ToolCall( + **{ + "description": ( + f"Search for files by glob pattern. Returns matching file paths relative " + f"to the search directory. Respects .gitignore. Output is truncated to " + f"1000 results or {max_kb}KB (whichever is hit first)." + ), + "parameters": { + "type": "object", + "properties": { + "pattern": { + "type": "string", + "description": "Glob pattern to match files, " + "e.g. '*.ts', '**/*.json', or 'src/**/*.spec.ts'", + }, + "path": { + "type": "string", + "description": "Directory to search in (default: current directory)", + }, + "limit": { + "type": "number", + "description": "Maximum number of results (default: 1000)", + }, + }, + "required": ["pattern"], + }, + }, + ) + + def _load_gitignore_patterns(self, search_path: Path) -> list[str]: + """Load gitignore patterns from directory and subdirectories.""" + patterns = ["**/node_modules/**", "**/.git/**"] + + # Load root .gitignore + gitignore_path = search_path / ".gitignore" + if gitignore_path.exists(): + try: + with open(gitignore_path, "r", encoding="utf-8") as f: + for line in f: + line = line.strip() + if line and not line.startswith("#"): + patterns.append(line) + except Exception: + pass # Ignore errors + + # Load nested .gitignore files + try: + for gitignore in search_path.rglob(".gitignore"): + if gitignore == gitignore_path: + continue + try: + with open(gitignore, "r", encoding="utf-8") as f: + for line in f: + line = line.strip() + if line and not line.startswith("#"): + patterns.append(line) + except Exception: + pass # Ignore errors + except Exception: + pass # Ignore glob errors + + return patterns + + def _should_ignore(self, path: Path, ignore_patterns: list[str]) -> bool: + """Check if path matches any ignore pattern.""" + path_str = str(path) + + for pattern in ignore_patterns: + # Simple pattern matching (not full gitignore spec) + if "**" in pattern: + # Recursive match + clean_pattern = pattern.replace("**/", "").replace("/**", "") + if clean_pattern in path_str: + return True + elif "*" in pattern: + # Wildcard match + from fnmatch import fnmatch + + if fnmatch(path.name, pattern): + return True + elif pattern in path_str: + return True + + return False + + async def execute(self) -> str: + """Execute file search.""" + pattern: str = self.context.pattern + search_dir: str = self.context.get("path", ".") + limit: int = self.context.get("limit", 1000) + + # Resolve search path + if not os.path.isabs(search_dir): + search_path = Path(self.cwd) / search_dir + else: + search_path = Path(search_dir) + + # Check if directory exists + if not search_path.exists(): + raise FileNotFoundError(f"Path not found: {search_dir}") + + if not search_path.is_dir(): + raise NotADirectoryError(f"Path is not a directory: {search_dir}") + + # Load gitignore patterns + ignore_patterns = self._load_gitignore_patterns(search_path) + + # Search for files + results = [] + try: + 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 + + # 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 + + # Handle no results + if not results: + return "No files found matching pattern" + + # Sort results for consistency + results.sort() + + # Apply limit and truncation + result_limit_reached = len(results) >= limit + raw_output = "\n".join(results) + truncation = truncate_head(raw_output, max_lines=FIND_MAX_LINES, max_bytes=FIND_MAX_BYTES) + + output = truncation.content + notices = [] + + if result_limit_reached: + notices.append( + f"{limit} results limit reached. Use limit={limit * 2} for more, or refine pattern", + ) + + if truncation.truncated: + notices.append(f"{format_size(FIND_MAX_BYTES)} limit reached") + + if notices: + output += f"\n\n[{'. '.join(notices)}]" + + return output diff --git a/reme/tool/fs/fs_memory_get.py b/reme/tool/fs/fs_memory_get.py new file mode 100644 index 00000000..c872aabf --- /dev/null +++ b/reme/tool/fs/fs_memory_get.py @@ -0,0 +1,79 @@ +"""Memory get tool for reading specific snippets from memory files.""" + +import os +from pathlib import Path + +from reme.core.op import BaseTool +from reme.core.schema import ToolCall + + +class FsMemoryGet(BaseTool): + """Read specific snippets from memory files.""" + + def __init__(self, workspace_dir: str | None = None, **kwargs): + """Initialize memory get tool.""" + kwargs.setdefault("name", "memory_get") + super().__init__(**kwargs) + self.workspace_dir = workspace_dir or os.getcwd() + + def _build_tool_call(self) -> ToolCall: + return ToolCall( + **{ + "description": ( + "Safe snippet read from MEMORY.md, memory/*.md with optional from/lines; " + "use after memory_search to pull only the needed lines and keep context small." + ), + "parameters": { + "type": "object", + "properties": { + "path": { + "type": "string", + "description": "Path to the memory file to read (relative or absolute)", + }, + "from": { + "type": "integer", + "description": "Starting line number (1-indexed, optional)", + }, + "lines": { + "type": "integer", + "description": "Number of lines to read from the starting line (optional)", + }, + }, + "required": ["path"], + }, + }, + ) + + 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) + + if os.path.isabs(raw_path): + abs_path = os.path.abspath(raw_path) + else: + abs_path = os.path.abspath(os.path.join(self.workspace_dir, raw_path)) + assert abs_path.lower().endswith(".md") + + # Check file exists, is not a symlink, and is a regular file + file_path = Path(abs_path) + assert ( + file_path.exists() and not file_path.is_symlink() and file_path.is_file() + ), f"File not found or not a regular file: {abs_path}" + + with open(abs_path, "r", encoding="utf-8") as f: + content = f.read() + + if from_param is None and lines_param 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)) + + # Extract slice (1-indexed to 0-indexed conversion) + selected = lines[start - 1 : start - 1 + count] + text = "\n".join(selected) + return text diff --git a/reme/tool/fs/fs_memory_search.py b/reme/tool/fs/fs_memory_search.py new file mode 100644 index 00000000..af245a99 --- /dev/null +++ b/reme/tool/fs/fs_memory_search.py @@ -0,0 +1,134 @@ +"""Memory search tool for semantic search in memory files.""" + +import json + +from reme.core.enumeration import MemorySource +from reme.core.op import BaseTool +from reme.core.schema import MemorySearchResult, ToolCall + + +class FsMemorySearch(BaseTool): + """Semantically search MEMORY.md and memory files.""" + + def __init__( + self, + sources: list[MemorySource] | None = None, + min_score: float = 0.1, + max_results: int = 20, + hybrid_enabled: bool = True, + hybrid_vector_weight: float = 0.7, + hybrid_text_weight: float = 0.3, + hybrid_candidate_multiplier: float = 3.0, + **kwargs, + ): + """Initialize memory search tool.""" + kwargs.setdefault("name", "memory_search") + super().__init__(**kwargs) + self.sources = sources or [MemorySource.MEMORY] + self.min_score = min_score + self.max_results = max_results + self.hybrid_enabled = hybrid_enabled + self.hybrid_vector_weight = hybrid_vector_weight + self.hybrid_text_weight = hybrid_text_weight + self.hybrid_candidate_multiplier = hybrid_candidate_multiplier + + def _build_tool_call(self) -> ToolCall: + return ToolCall( + **{ + "description": ( + "Mandatory recall step: semantically search MEMORY.md + memory/*.md " + "(and optional session transcripts) before answering questions about " + "prior work, decisions, dates, people, preferences, or todos; returns " + "top snippets with path + lines." + ), + "parameters": { + "type": "object", + "properties": { + "query": { + "type": "string", + "description": "The semantic search query to find relevant memory snippets", + }, + "maxResults": { + "type": "integer", + "description": "Maximum number of search results to return (optional)", + }, + "minScore": { + "type": "number", + "description": "Minimum similarity score threshold for results (optional)", + }, + }, + "required": ["query"], + }, + }, + ) + + 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) + candidates = min(200, max(1, int(max_results * self.hybrid_candidate_multiplier))) + + # Perform hybrid search (vector + keyword) + if self.hybrid_enabled: + keyword_results = [] + if self.memory_store.fts_enabled: + keyword_results = await self._search_keyword(query, candidates) + vector_results = await self._search_vector(query, candidates) + + if not keyword_results: + results = [r for r in vector_results if r.score >= min_score][:max_results] + elif not vector_results: + results = [r for r in keyword_results if r.score >= min_score][:max_results] + else: + merged = self._merge_hybrid_results( + vector=vector_results, + keyword=keyword_results, + vector_weight=self.hybrid_vector_weight, + text_weight=self.hybrid_text_weight, + ) + results = [r for r in merged if r.score >= min_score][:max_results] + else: + vector_results = await self._search_vector(query, candidates) + results = [r for r in vector_results if r.score >= min_score][:max_results] + + return json.dumps([result.model_dump(exclude_none=True) for result in results], indent=2, ensure_ascii=False) + + async def _search_vector(self, query: str, limit: int) -> list[MemorySearchResult]: + """Perform vector similarity search.""" + return await self.memory_store.vector_search(query, limit, sources=self.sources) + + async def _search_keyword(self, query: str, limit: int) -> list[MemorySearchResult]: + """Perform keyword/FTS search.""" + if not self.memory_store.fts_enabled: + return [] + return await self.memory_store.keyword_search(query, limit, sources=self.sources) + + @staticmethod + def _merge_hybrid_results( + vector: list[MemorySearchResult], + keyword: list[MemorySearchResult], + vector_weight: float, + text_weight: float, + ) -> list[MemorySearchResult]: + """Merge vector and keyword search results with weighted scoring.""" + merged: dict[str, MemorySearchResult] = {} + + # Process vector results + for result in vector: + result.score = result.score * vector_weight + merged[result.merge_key] = result + + # Process keyword results + for result in keyword: + key = result.merge_key + if key in merged: + merged[key].score += result.score * text_weight + else: + result.score = result.score * text_weight + merged[key] = result + + # Sort by score and return + results = list(merged.values()) + results.sort(key=lambda r: r.score, reverse=True) + return results diff --git a/reme/tool/fs/grep_tool.py b/reme/tool/fs/grep_tool.py new file mode 100644 index 00000000..9df9f8f2 --- /dev/null +++ b/reme/tool/fs/grep_tool.py @@ -0,0 +1,276 @@ +"""Grep tool for searching file contents using ripgrep. + +This module provides a tool for searching file contents with: +- Pattern matching (regex or literal string) +- Smart output truncation (prevents memory issues) +- Context lines support +- Respects .gitignore +""" + +import asyncio +import json +import os +import shutil +from pathlib import Path + +from .truncate import ( + DEFAULT_MAX_BYTES, + GREP_MAX_LINE_LENGTH, + format_size, + 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): + """Tool for searching file contents using ripgrep. + + Features: + - Pattern matching with regex or literal string + - Context lines support + - Smart output truncation + - Respects .gitignore + """ + + def __init__(self, cwd: str | None = None): + """Initialize grep tool. + + Args: + cwd: Working directory (defaults to current directory) + """ + super().__init__() + self.cwd = cwd or os.getcwd() + + def _build_tool_call(self) -> ToolCall: + max_kb = DEFAULT_MAX_BYTES // 1024 + return ToolCall( + **{ + "description": ( + f"Search file contents for a pattern. Returns matching lines with " + f"file paths and line numbers. Respects .gitignore. Output is " + f"truncated to {DEFAULT_LIMIT} matches or {max_kb}KB (whichever is " + f"hit first). Long lines are truncated to {GREP_MAX_LINE_LENGTH} chars." + ), + "parameters": { + "type": "object", + "properties": { + "pattern": { + "type": "string", + "description": "Search pattern (regex or literal string)", + }, + "path": { + "type": "string", + "description": "Directory or file to search (default: current directory)", + }, + "glob": { + "type": "string", + "description": "Filter files by glob pattern, e.g. '*.ts' or '**/*.spec.ts'", + }, + "ignoreCase": { + "type": "boolean", + "description": "Case-insensitive search (default: false)", + }, + "literal": { + "type": "boolean", + "description": "Treat pattern as literal string instead of regex (default: false)", + }, + "contextLines": { + "type": "number", + "description": "Number of lines to show before and after each match (default: 0)", + }, + "limit": { + "type": "number", + "description": f"Maximum number of matches to return (default: {DEFAULT_LIMIT})", + }, + }, + "required": ["pattern"], + }, + }, + ) + + async def execute(self) -> str: + """Execute the grep search.""" + pattern: str = self.context.pattern + search_path: str = self.context.get("path", ".") + glob: str | None = self.context.get("glob", None) + ignore_case: bool = self.context.get("ignoreCase", False) + literal: bool = self.context.get("literal", False) + context_lines: int = self.context.get("contextLines", 0) + limit: int = self.context.get("limit", DEFAULT_LIMIT) + + # Check if ripgrep is available + rg_path = shutil.which("rg") + if not rg_path: + raise RuntimeError( + "ripgrep (rg) is not available. Please install it:\n" + " macOS: brew install ripgrep\n" + " Ubuntu: apt-get install ripgrep\n" + " Other: https://github.com/BurntSushi/ripgrep", + ) + + # Resolve search path + if not os.path.isabs(search_path): + search_path = os.path.join(self.cwd, search_path) + + # Check if path exists + if not Path(search_path).exists(): + raise FileNotFoundError(f"Path not found: {search_path}") + + is_directory = Path(search_path).is_dir() + effective_limit = max(1, limit) + + # Build ripgrep arguments + args = [ + rg_path, + "--json", + "--line-number", + "--color=never", + "--hidden", + ] + + if ignore_case: + args.append("--ignore-case") + + if literal: + args.append("--fixed-strings") + + if glob: + args.extend(["--glob", glob]) + + 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 + + stdout, stderr = await process.communicate() + + # Parse JSON output + matches = [] + match_count = 0 + lines_truncated = False + + for line in stdout.decode("utf-8", errors="ignore").splitlines(): + if not line.strip() or match_count >= effective_limit: + break + + try: + event = json.loads(line) + except json.JSONDecodeError: + continue + + if event.get("type") == "match": + match_count += 1 + data = event.get("data", {}) + file_path = data.get("path", {}).get("text", "") + line_number = data.get("line_number", 0) + + if file_path and line_number: + matches.append({"file_path": file_path, "line_number": line_number}) + + if match_count >= effective_limit: + break + + # Check for errors + if process.returncode not in (0, 1) and match_count == 0: + error_msg = stderr.decode("utf-8", errors="ignore").strip() + if error_msg: + raise RuntimeError(error_msg) + raise RuntimeError(f"ripgrep exited with code {process.returncode}") + + # No matches found + if match_count == 0: + return "No matches found" + + # Format matches with context + output_lines = [] + file_cache = {} + + for match in matches: + file_path = match["file_path"] + line_number = match["line_number"] + + # Read file if not cached + if file_path not in file_cache: + try: + with open(file_path, "r", encoding="utf-8", errors="ignore") as f: + file_cache[file_path] = f.read().replace("\r\n", "\n").replace("\r", "\n").split("\n") + except Exception: + file_cache[file_path] = [] + + lines = file_cache[file_path] + + # Format relative path + if is_directory: + relative_path = os.path.relpath(file_path, search_path) + if not relative_path.startswith(".."): + display_path = relative_path.replace("\\", "/") + else: + display_path = os.path.basename(file_path) + else: + display_path = os.path.basename(file_path) + + # Generate context block + if not lines: + output_lines.append(f"{display_path}:{line_number}: (unable to read file)") + continue + + context_value = max(0, context_lines) + start = max(1, line_number - context_value) if context_value > 0 else line_number + end = min(len(lines), line_number + context_value) if context_value > 0 else line_number + + for current in range(start, end + 1): + if current < 1 or current > len(lines): + continue + + line_text = lines[current - 1] + is_match_line = current == line_number + + # Truncate long lines + truncated_text, was_truncated = truncate_line(line_text) + if was_truncated: + lines_truncated = True + + if is_match_line: + output_lines.append(f"{display_path}:{current}: {truncated_text}") + else: + output_lines.append(f"{display_path}-{current}- {truncated_text}") + + # Apply byte truncation + raw_output = "\n".join(output_lines) + truncation = truncate_head(raw_output, max_lines=999999999) + + output = truncation.content + notices = [] + + # Add notices + if match_count >= effective_limit: + notices.append( + f"{effective_limit} matches limit reached. " + f"Use limit={effective_limit * 2} for more, or refine pattern", + ) + + if truncation.truncated: + notices.append(f"{format_size(DEFAULT_MAX_BYTES)} limit reached") + + if lines_truncated: + notices.append( + f"Some lines truncated to {GREP_MAX_LINE_LENGTH} chars. " f"Use read tool to see full lines", + ) + + if notices: + output += f"\n\n[{'. '.join(notices)}]" + + return output diff --git a/reme/tool/fs/ls_tool.py b/reme/tool/fs/ls_tool.py new file mode 100644 index 00000000..1c054c62 --- /dev/null +++ b/reme/tool/fs/ls_tool.py @@ -0,0 +1,128 @@ +"""Directory listing tool with truncation support.""" + +import os +from pathlib import Path + +from .truncate import DEFAULT_MAX_BYTES, truncate_head +from ...core.op import BaseTool +from ...core.schema import ToolCall + +DEFAULT_LIMIT = 500 + + +class LsTool(BaseTool): + """List directory contents with smart truncation. + + Features: + - Returns entries sorted alphabetically (case-insensitive) + - Directory indicators ('/' suffix) + - Includes dotfiles + - Entry count limiting + - Byte truncation + """ + + def __init__(self, cwd: str | None = None): + """Initialize ls tool. + + Args: + cwd: Working directory (defaults to current directory) + """ + super().__init__() + self.cwd = cwd or os.getcwd() + + def _build_tool_call(self) -> ToolCall: + max_kb = DEFAULT_MAX_BYTES // 1024 + return ToolCall( + **{ + "description": ( + f"List directory contents. Returns entries sorted alphabetically, " + f"with '/' suffix for directories. Includes dotfiles. Output is truncated " + f"to {DEFAULT_LIMIT} entries or {max_kb}KB (whichever is hit first)." + ), + "parameters": { + "type": "object", + "properties": { + "path": { + "type": "string", + "description": "Directory to list (default: current directory)", + }, + "limit": { + "type": "number", + "description": f"Maximum number of entries to return (default: {DEFAULT_LIMIT})", + }, + }, + "required": [], + }, + }, + ) + + async def execute(self) -> str: + """List directory contents with production-grade features.""" + path: str | None = self.context.get("path", None) + limit: int | None = self.context.get("limit", None) + + # Resolve directory path + dir_path = Path(self.cwd) / (path or ".") + dir_path = dir_path.resolve() + effective_limit = limit if limit is not None else DEFAULT_LIMIT + + # Check if path exists + if not dir_path.exists(): + raise FileNotFoundError(f"Path not found: {dir_path}") + + # Check if path is a directory + if not dir_path.is_dir(): + 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 + + # Sort alphabetically (case-insensitive) + entries.sort(key=lambda e: e.name.lower()) + + # Format entries with directory indicators + results: list[str] = [] + entry_limit_reached = False + + for entry in entries: + if len(results) >= effective_limit: + entry_limit_reached = True + break + + try: + # Add '/' suffix for directories + suffix = "/" if entry.is_dir() else "" + results.append(entry.name + suffix) + except Exception: + # Skip entries we can't stat + continue + + # Handle empty directory + if len(results) == 0: + return "(empty directory)" + + # Apply byte truncation + raw_output = "\n".join(results) + truncation_result = truncate_head(raw_output, max_lines=float("inf")) + + output_text = truncation_result.content + + # Build notices + notices: list[str] = [] + + if entry_limit_reached: + notices.append( + f"{effective_limit} entries limit reached. Use limit={effective_limit * 2} for more", + ) + + if truncation_result.truncated: + max_kb = DEFAULT_MAX_BYTES // 1024 + notices.append(f"{max_kb}KB limit reached") + + if notices: + output_text += f"\n\n[{'. '.join(notices)}]" + + return output_text diff --git a/reme/tool/fs/read_tool.py b/reme/tool/fs/read_tool.py new file mode 100644 index 00000000..b236d4a4 --- /dev/null +++ b/reme/tool/fs/read_tool.py @@ -0,0 +1,219 @@ +"""Read file tool with smart truncation and image support. + +Features: +- Reads text files with offset/limit support +- Detects and handles image files (jpg, png, gif, webp) +- Smart truncation to prevent memory issues +""" + +import os +from pathlib import Path + +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 +IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".gif", ".webp"} + + +def is_image_file(path: str) -> bool: + """Check if file is a supported image type. + + Args: + path: File path to check + + Returns: + True if file is a supported image + """ + return Path(path).suffix.lower() in IMAGE_EXTENSIONS + + +class ReadTool(BaseTool): + """Read file contents with smart truncation. + + Features: + - Supports text files and images (jpg, png, gif, webp) + - Smart truncation for large files + - Offset/limit for reading specific portions + """ + + def __init__(self, cwd: str | None = None): + """Initialize read tool. + + Args: + cwd: Working directory (defaults to current directory) + """ + super().__init__() + self.cwd = cwd or os.getcwd() + + def _build_tool_call(self) -> ToolCall: + max_kb = DEFAULT_MAX_BYTES // 1024 + return ToolCall( + **{ + "description": ( + f"Read the contents of a file. Supports text files and images " + f"(jpg, png, gif, webp). Images are sent as attachments. For text files, " + f"output is truncated to {DEFAULT_MAX_LINES} lines or {max_kb}KB " + f"(whichever is hit first). Use offset/limit for large files. " + f"When you need the full file, continue with offset until complete." + ), + "parameters": { + "type": "object", + "properties": { + "path": { + "type": "string", + "description": "Path to the file to read (relative or absolute)", + }, + "offset": { + "type": "number", + "description": "Line number to start reading from (1-indexed)", + }, + "limit": { + "type": "number", + "description": "Maximum number of lines to read", + }, + }, + "required": ["path"], + }, + }, + ) + + async def execute(self) -> str: + """Execute the read operation.""" + path: str = self.context.path + offset: int | None = self.context.get("offset", None) + limit: int | None = self.context.get("limit", None) + + # Resolve path + if not os.path.isabs(path): + absolute_path = os.path.join(self.cwd, path) + else: + absolute_path = path + absolute_path = os.path.normpath(absolute_path) + + # Check file exists and is readable + if not os.path.exists(absolute_path): + raise ValueError(f"File not found: {path}") + + if not os.path.isfile(absolute_path): + raise ValueError(f"Not a file: {path}") + + if not os.access(absolute_path, os.R_OK): + raise ValueError(f"File not readable: {path}") + + # Check if image + if is_image_file(absolute_path): + return await self._read_image(absolute_path, path) + else: + return await self._read_text(absolute_path, path, offset, limit) + + @staticmethod + async def _read_image(absolute_path: str, display_path: str) -> str: + """Read and return image file information. + + Args: + absolute_path: Absolute path to image + display_path: Path to display to user + + Returns: + Image information text + """ + # Get file size + file_size = os.path.getsize(absolute_path) + file_ext = Path(absolute_path).suffix.lower() + + # For Python tools, we typically can't return image data directly to LLM + # So we return a descriptive message + return ( + f"Read image file [{file_ext}]\n" + f"Path: {display_path}\n" + f"Size: {format_size(file_size)}\n" + f"Note: Image content cannot be displayed in text format. " + f"Use bash tool or other methods to process the image." + ) + + @staticmethod + async def _read_text( + absolute_path: str, + _display_path: str, + offset: int | None, + limit: int | None, + ) -> str: + """Read text file with smart truncation. + + Args: + absolute_path: Absolute path to file + _display_path: Path to display to user + offset: Starting line (1-indexed) + limit: Maximum lines to read + + Returns: + File contents with truncation notices + """ + # Read file + try: + with open(absolute_path, "r", encoding="utf-8") as f: + content = f.read() + except UnicodeDecodeError: + # Try with error handling for binary files + with open(absolute_path, "r", encoding="utf-8", errors="ignore") as f: + content = f.read() + + 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 + start_line_display = start_line + 1 + + # Check offset bounds + if start_line >= len(all_lines): + raise IndexError( + f"Offset {offset} is beyond end of file ({len(all_lines)} lines total)", + ) + + # Apply user limit if specified + 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 + else: + selected_content = "\n".join(all_lines[start_line:]) + user_limited_lines = None + + # Apply truncation + truncation = truncate_head(selected_content) + + # Build output with truncation notices + if truncation.truncated: + # Truncation occurred + end_line_display = start_line_display + truncation.output_lines - 1 + next_offset = end_line_display + 1 + + output_text = truncation.content + + if truncation.truncated_by == "lines": + output_text += ( + f"\n\n[Showing lines {start_line_display}-{end_line_display} " + f"of {total_file_lines}. Use offset={next_offset} to continue.]" + ) + else: + max_kb = DEFAULT_MAX_BYTES // 1024 + output_text += ( + f"\n\n[Showing lines {start_line_display}-{end_line_display} " + 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 + + 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 + output_text = truncation.content + + return output_text diff --git a/reme/tool/fs/truncate.py b/reme/tool/fs/truncate.py new file mode 100644 index 00000000..b4de0197 --- /dev/null +++ b/reme/tool/fs/truncate.py @@ -0,0 +1,209 @@ +"""fs utils""" + +from typing import Literal + +from ...core.schema import TruncationResult + +# Default limits for output truncation +DEFAULT_MAX_LINES = 1000 # Maximum lines to keep for tail truncation +DEFAULT_MAX_BYTES = 30 * 1024 # Maximum bytes to keep (30KB) + +# Find tool limits +FIND_MAX_LINES = 2000 # Maximum lines for find output +FIND_MAX_BYTES = 50 * 1024 # 50KB for find output + +# Grep tool limits +GREP_MAX_LINE_LENGTH = 500 # Maximum line length for grep output + + +def format_size(num_bytes: int) -> str: + """Format byte size in human-readable format. + + Args: + num_bytes: Number of bytes + + Returns: + Formatted string (e.g., "1.5KB", "2.3MB") + """ + if num_bytes < 1024: + return f"{num_bytes}B" + elif num_bytes < 1024 * 1024: + return f"{num_bytes / 1024:.1f}KB" + else: + return f"{num_bytes / (1024 * 1024):.1f}MB" + + +def truncate_line(text: str, max_length: int = GREP_MAX_LINE_LENGTH) -> tuple[str, bool]: + """Truncate a single line if it exceeds max length. + + Args: + text: Line text + max_length: Maximum line length + + Returns: + Tuple of (truncated_text, was_truncated) + """ + if len(text) <= max_length: + return text, False + return text[:max_length] + "...", True + + +def truncate_tail( + text: str, + max_lines: int = DEFAULT_MAX_LINES, + max_bytes: int = DEFAULT_MAX_BYTES, +) -> TruncationResult: + """Truncate text to keep only the tail (last portion). + + Keeps the last N lines or M bytes, whichever is hit first. + This is useful for command outputs where the end is most relevant. + + Args: + text: The text to truncate + max_lines: Maximum number of lines to keep + max_bytes: Maximum bytes to keep + + Returns: + TruncationResult with truncated content and metadata + """ + if not text: + return TruncationResult( + content="", + truncated=False, + total_lines=0, + output_lines=0, + total_bytes=0, + output_bytes=0, + ) + + total_bytes = len(text.encode("utf-8")) + lines = text.split("\n") + total_lines = len(lines) + + # Check if we need to truncate + if total_lines <= max_lines and total_bytes <= max_bytes: + return TruncationResult( + content=text, + truncated=False, + total_lines=total_lines, + output_lines=total_lines, + total_bytes=total_bytes, + output_bytes=total_bytes, + ) + + # Keep last N lines + kept_lines = lines[-max_lines:] if total_lines > max_lines else lines + truncated_by: Literal["lines", "bytes"] = "lines" if total_lines > max_lines else "bytes" + + # Check byte limit on kept lines + kept_text = "\n".join(kept_lines) + kept_bytes = len(kept_text.encode("utf-8")) + + # If still over byte limit, truncate further + last_line_partial = False + if kept_bytes > max_bytes: + truncated_by = "bytes" + # Keep truncating from the start until under byte limit + while kept_lines and len("\n".join(kept_lines).encode("utf-8")) > max_bytes: + kept_lines.pop(0) + + # If still over (single line > max_bytes), truncate the line itself + if kept_lines and len("\n".join(kept_lines).encode("utf-8")) > max_bytes: + last_line = kept_lines[-1] + # Binary search to find how much of last line fits + encoded = last_line.encode("utf-8") + if len(encoded) > max_bytes: + last_line_partial = True + # Take last max_bytes of the line + kept_lines[-1] = encoded[-max_bytes:].decode("utf-8", errors="ignore") + + kept_text = "\n".join(kept_lines) + kept_bytes = len(kept_text.encode("utf-8")) + + return TruncationResult( + content=kept_text, + truncated=True, + total_lines=total_lines, + output_lines=len(kept_lines), + total_bytes=total_bytes, + output_bytes=kept_bytes, + truncated_by=truncated_by, + last_line_partial=last_line_partial, + ) + + +def truncate_head( + text: str, + max_lines: int = FIND_MAX_LINES, + max_bytes: int = FIND_MAX_BYTES, +) -> TruncationResult: + """Truncate text to keep only the head (first portion). + + Keeps the first N lines or M bytes, whichever is hit first. + Suitable for file reads where you want to see the beginning. + + Args: + text: The text to truncate + max_lines: Maximum number of lines to keep + max_bytes: Maximum bytes to keep + + Returns: + TruncationResult with truncated content and metadata + """ + if not text: + return TruncationResult( + content="", + truncated=False, + total_lines=0, + output_lines=0, + total_bytes=0, + output_bytes=0, + ) + + total_bytes = len(text.encode("utf-8")) + lines = text.split("\n") + total_lines = len(lines) + + # Check if no truncation needed + if total_lines <= max_lines and total_bytes <= max_bytes: + return TruncationResult( + content=text, + truncated=False, + total_lines=total_lines, + output_lines=total_lines, + total_bytes=total_bytes, + output_bytes=total_bytes, + ) + + # Collect complete lines that fit + kept_lines = [] + kept_bytes = 0 + truncated_by: Literal["lines", "bytes"] = "lines" + + for i, line in enumerate(lines): + if i >= max_lines: + truncated_by = "lines" + break + + # Calculate bytes for this line (+1 for newline except first line) + line_bytes = len(line.encode("utf-8")) + (1 if i > 0 else 0) + + if kept_bytes + line_bytes > max_bytes: + truncated_by = "bytes" + break + + kept_lines.append(line) + kept_bytes += line_bytes + + kept_text = "\n".join(kept_lines) + final_bytes = len(kept_text.encode("utf-8")) + + return TruncationResult( + content=kept_text, + truncated=True, + total_lines=total_lines, + output_lines=len(kept_lines), + total_bytes=total_bytes, + output_bytes=final_bytes, + truncated_by=truncated_by, + ) diff --git a/reme/tool/fs/write_tool.py b/reme/tool/fs/write_tool.py new file mode 100644 index 00000000..55a4a0c1 --- /dev/null +++ b/reme/tool/fs/write_tool.py @@ -0,0 +1,81 @@ +"""Write tool for creating and overwriting files. + +This module provides a tool for writing content to files with: +- Automatic parent directory creation +- File overwriting (creates if doesn't exist, overwrites if exists) +- Path resolution (relative to working directory) +""" + +import os + +from ...core.op import BaseTool +from ...core.schema import ToolCall + + +class WriteTool(BaseTool): + """Tool for writing content to files. + + Features: + - Creates file if it doesn't exist, overwrites if it does + - Automatically creates parent directories + - Supports both relative and absolute paths + """ + + def __init__(self, cwd: str | None = None): + """Initialize write tool. + + Args: + cwd: Working directory (defaults to current directory) + """ + super().__init__() + self.cwd = cwd or os.getcwd() + + def _build_tool_call(self) -> ToolCall: + return ToolCall( + **{ + "description": ( + "Write content to a file. Creates the file if it doesn't exist, " + "overwrites if it does. Automatically creates parent directories." + ), + "parameters": { + "type": "object", + "properties": { + "path": { + "type": "string", + "description": "Path to the file to write (relative or absolute)", + }, + "content": { + "type": "string", + "description": "Content to write to the file", + }, + }, + "required": ["path", "content"], + }, + }, + ) + + async def execute(self) -> str: + """Execute the write operation.""" + path: str = self.context.path + content: str = self.context.content + + # Resolve path to absolute + if not os.path.isabs(path): + absolute_path = os.path.join(self.cwd, path) + else: + absolute_path = path + + absolute_path = os.path.normpath(absolute_path) + + # Create parent directories if needed + parent_dir = os.path.dirname(absolute_path) + if parent_dir: + os.makedirs(parent_dir, exist_ok=True) + + # Write the file + with open(absolute_path, "w", encoding="utf-8") as f: + f.write(content) + + # Return success message + content_bytes = len(content.encode("utf-8")) + return f"Successfully wrote {content_bytes} bytes to {path}" diff --git a/tests/demo_memory_search.py b/tests/demo_memory_search.py new file mode 100644 index 00000000..104d333a --- /dev/null +++ b/tests/demo_memory_search.py @@ -0,0 +1,400 @@ +""" +测试场景:对比 DeltaFileWatcher 和 FullFileWatcher 的行为差异 + +场景流程: +1. 测试 FullFileWatcher - 完整更新模式 + - 创建初始文件 + - 修改文件 + - 验证数据库:所有 chunks 被重新创建 + +2. 测试 DeltaFileWatcher - 增量更新模式 + - 创建初始文件 + - 追加内容到文件 + - 验证数据库:只有新增的 chunks,旧 chunks 保留 + +3. 直接查询数据库验证更新结果 +""" + +import asyncio +import os +import tempfile +from datetime import datetime + +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.utils import load_env + +load_env() + +# 配置 +DB_PATH_FULL = "./demo_memory_search/full_watcher.db" +DB_PATH_DELTA = "./demo_memory_search/delta_watcher.db" +EMBEDDING_MODEL = "text-embedding-v4" +EMBEDDING_DIMENSIONS = 64 + + +def print_separator(title): + """打印分隔符""" + print("\n" + "=" * 80) + print(f" {title}") + print("=" * 80 + "\n") + + +def create_file(workspace_dir: str, filename: str, content: str) -> str: + """在工作空间创建文件""" + file_path = os.path.join(workspace_dir, filename) + os.makedirs(os.path.dirname(file_path), exist_ok=True) + with open(file_path, "w", encoding="utf-8") as f: + f.write(content) + print(f"✓ 创建文件: {filename}") + return file_path + + +def append_to_file(file_path: str, content: str): + """追加内容到文件""" + with open(file_path, "a", encoding="utf-8") as f: + f.write(content) + print(f"✓ 追加内容到: {file_path}") + + +async def verify_database(store: SqliteMemoryStore, title: str): + """验证数据库内容""" + print_separator(title) + + # 查询文件列表 + files = await store.list_files(MemorySource.MEMORY) + print(f"📁 数据库中的文件数量: {len(files)}\n") + + for file_path in files: + print(f"文件: {file_path}") + + # 获取文件元数据 + file_meta = await store.get_file_metadata(file_path, MemorySource.MEMORY) + if file_meta: + print(f" - Hash: {file_meta.hash[:16]}...") + print(f" - Size: {file_meta.size} bytes") + print(f" - Chunk count: {file_meta.chunk_count}") + + # 获取文件的所有 chunks + chunks = await store.get_file_chunks(file_path, MemorySource.MEMORY) + print(f" - Chunks in database: {len(chunks)}") + + for i, chunk in enumerate(chunks, 1): + print(f" Chunk #{i}:") + print(f" ID: {chunk.id[:16]}...") + print(f" Lines: {chunk.start_line}-{chunk.end_line}") + print(f" Hash: {chunk.hash[:16]}...") + print(f" Text preview: {chunk.text[:100]}...") + print(f" Has embedding: {chunk.embedding is not None}") + + print() + + +async def test_full_file_watcher(): + """测试 FullFileWatcher - 完整更新模式""" + print_separator("测试 1: FullFileWatcher (完整更新模式)") + + # 创建临时工作目录 + temp_workspace = tempfile.mkdtemp(prefix="full_watcher_") + print(f"工作目录: {temp_workspace}\n") + + # 准备数据库 + os.makedirs(os.path.dirname(DB_PATH_FULL), exist_ok=True) + + # 初始化组件 + embedding_model = OpenAIEmbeddingModel( + model_name=EMBEDDING_MODEL, + dimensions=EMBEDDING_DIMENSIONS, + ) + + store = SqliteMemoryStore( + db_path=DB_PATH_FULL, + vec_ext_path="", + embedding_model=embedding_model, + fts_enabled=True, + ) + await store.start() + await store.clear_all() # 清空数据库 + + # 创建 FullFileWatcher + watcher = FullFileWatcher( + watch_paths=temp_workspace, + memory_store=store, + chunk_tokens=200, + chunk_overlap=20, + recursive=True, + suffix_filters=["md"], + ) + + try: + # 阶段 1: 创建初始文件 + print("📝 阶段 1: 创建初始文件\n") + _ = create_file( + temp_workspace, + "MEMORY.md", + """# Python 基础 + +## 变量和数据类型 +Python 是动态类型语言。 + +## 控制流 +if、for、while 语句。 +""", + ) + + # 启动 watcher + await watcher.start() + print("✓ FullFileWatcher 启动\n") + + # 等待文件被检测和处理 + await asyncio.sleep(2) + + # 验证数据库 - 初始状态 + await verify_database(store, "数据库验证 1.1: 初始文件索引后") + + # 阶段 2: 修改文件(非追加,而是完全修改) + print("📝 阶段 2: 修改文件内容\n") + create_file( + temp_workspace, + "MEMORY.md", + """# Python 进阶 + +## 变量和数据类型 +Python 是动态类型语言,支持多种数据类型。 + +## 控制流 +if、for、while 语句用于控制程序流程。 + +## 函数 +def 关键字定义函数。 + +## 类和对象 +面向对象编程的核心概念。 +""", + ) + + # 等待文件变化被检测和处理 + await asyncio.sleep(2) + + # 验证数据库 - 修改后 + await verify_database(store, "数据库验证 1.2: 文件修改后(完整更新)") + + print("📊 观察要点:") + print(" - FullFileWatcher 每次修改都会删除所有旧 chunks") + print(" - 然后重新创建所有新 chunks") + print(" - Chunk IDs 会完全不同") + + finally: + await watcher.close() + await store.close() + print(f"\n💡 工作目录: {temp_workspace}") + print(f"💡 数据库: {DB_PATH_FULL}") + + +async def test_delta_file_watcher(): + """测试 DeltaFileWatcher - 增量更新模式""" + print_separator("测试 2: DeltaFileWatcher (增量更新模式)") + + # 创建临时工作目录 + temp_workspace = tempfile.mkdtemp(prefix="delta_watcher_") + print(f"工作目录: {temp_workspace}\n") + + # 准备数据库 + os.makedirs(os.path.dirname(DB_PATH_DELTA), exist_ok=True) + + # 初始化组件 + embedding_model = OpenAIEmbeddingModel( + model_name=EMBEDDING_MODEL, + dimensions=EMBEDDING_DIMENSIONS, + ) + + store = SqliteMemoryStore( + db_path=DB_PATH_DELTA, + vec_ext_path="", + embedding_model=embedding_model, + fts_enabled=True, + ) + await store.start() + await store.clear_all() # 清空数据库 + + # 创建 DeltaFileWatcher + watcher = DeltaFileWatcher( + watch_paths=temp_workspace, + memory_store=store, + chunk_tokens=200, + chunk_overlap=20, + overlap_lines=2, + recursive=True, + suffix_filters=["md"], + ) + + try: + # 阶段 1: 创建初始文件 + print("📝 阶段 1: 创建初始文件\n") + test_file = create_file( + temp_workspace, + "MEMORY.md", + """# Python 基础 + +## 变量和数据类型 +Python 是动态类型语言。 + +## 控制流 +if、for、while 语句。 +""", + ) + + # 启动 watcher + await watcher.start() + print("✓ DeltaFileWatcher 启动\n") + + # 等待文件被检测和处理 + await asyncio.sleep(2) + + # 验证数据库 - 初始状态 + await verify_database(store, "数据库验证 2.1: 初始文件索引后") + + # 保存初始 chunk IDs + initial_chunks = await store.get_file_chunks(test_file, MemorySource.MEMORY) + initial_chunk_ids = {chunk.id for chunk in initial_chunks} + print(f"📌 初始 chunk IDs: {len(initial_chunk_ids)} 个\n") + + # 阶段 2: 追加内容到文件(append-only) + print("📝 阶段 2: 追加新内容到文件\n") + append_to_file( + test_file, + """ + +## 函数 +def 关键字定义函数。 + +## 类和对象 +面向对象编程的核心概念。 + +## 模块和包 +代码组织和复用。 +""", + ) + + # 等待文件变化被检测和处理 + await asyncio.sleep(2) + + # 验证数据库 - 追加后 + await verify_database(store, "数据库验证 2.2: 追加内容后(增量更新)") + + # 检查哪些 chunks 被保留 + updated_chunks = await store.get_file_chunks(test_file, MemorySource.MEMORY) + updated_chunk_ids = {chunk.id for chunk in updated_chunks} + + preserved_ids = initial_chunk_ids & updated_chunk_ids + new_ids = updated_chunk_ids - initial_chunk_ids + deleted_ids = initial_chunk_ids - updated_chunk_ids + + print("\n📊 Chunk 变化统计:") + print(f" - 保留的 chunks: {len(preserved_ids)} 个") + print(f" - 新增的 chunks: {len(new_ids)} 个") + print(f" - 删除的 chunks: {len(deleted_ids)} 个") + + print("\n📊 观察要点:") + print(" - DeltaFileWatcher 检测到 append-only 模式") + print(" - 保留了大部分旧 chunks(除了重叠部分)") + print(" - 只处理和嵌入新增的内容") + print(" - 节省了 embedding API 调用") + + # 阶段 3: 非追加式修改(触发完整更新) + print("\n📝 阶段 3: 修改文件中间部分(非追加)\n") + create_file( + temp_workspace, + "MEMORY.md", + """# Python 完全改版 + +## 这是全新的内容 +完全不同的文档结构。 + +## 新的章节 +之前的内容都不见了。 +""", + ) + + # 等待文件变化被检测和处理 + await asyncio.sleep(2) + + # 验证数据库 - 非追加修改后 + await verify_database(store, "数据库验证 2.3: 非追加修改后(回退到完整更新)") + + final_chunks = await store.get_file_chunks(test_file, MemorySource.MEMORY) + final_chunk_ids = {chunk.id for chunk in final_chunks} + + print("\n📊 第二次修改后的 Chunk 统计:") + print(f" - 当前 chunks: {len(final_chunk_ids)} 个") + print(f" - 所有 chunk IDs 都是新的: {len(final_chunk_ids & updated_chunk_ids) == 0}") + + print("\n📊 观察要点:") + print(" - DeltaFileWatcher 检测到非追加式修改") + print(" - 自动回退到完整更新模式") + print(" - 所有 chunks 被重新创建") + + finally: + await watcher.close() + await store.close() + print(f"\n💡 工作目录: {temp_workspace}") + print(f"💡 数据库: {DB_PATH_DELTA}") + + +async def compare_watchers(): + """对比两种 watcher 的性能和行为""" + print_separator("对比分析") + + print("🔍 FullFileWatcher vs DeltaFileWatcher\n") + + print("FullFileWatcher 特点:") + print(" ✓ 实现简单") + print(" ✓ 适合频繁修改的文档") + print(" ✓ 每次修改都保证完整性") + print(" ✗ 每次都重新处理整个文件") + print(" ✗ 更多的 embedding API 调用") + + print("\nDeltaFileWatcher 特点:") + print(" ✓ 支持增量更新") + print(" ✓ 节省 embedding API 调用") + print(" ✓ 适合日志类追加文件") + print(" ✓ 检测到非追加时自动回退") + print(" ✗ 实现复杂") + print(" ✗ 需要维护 cutoff line 逻辑") + + print("\n推荐使用场景:") + print(" - 日志文件、聊天记录 → DeltaFileWatcher") + print(" - 文档、代码文件 → FullFileWatcher") + print(" - 不确定的场景 → DeltaFileWatcher (会自动回退)") + + +async def main(): + """主函数""" + print_separator("FileWatcher 对比测试") + print(f"开始时间: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n") + + # 测试 1: FullFileWatcher + await test_full_file_watcher() + + # 测试 2: DeltaFileWatcher + await test_delta_file_watcher() + + # 对比分析 + await compare_watchers() + + print_separator("测试完成") + print(f"结束时间: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n") + + print("📂 生成的文件:") + print(f" - Full watcher DB: {DB_PATH_FULL}") + print(f" - Delta watcher DB: {DB_PATH_DELTA}") + + print("\n🔧 清理命令:") + print(f" rm -rf {os.path.dirname(DB_PATH_FULL)}") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/test_cache_memory_usage.py b/tests/test_cache_memory_usage.py new file mode 100644 index 00000000..ed0020f9 --- /dev/null +++ b/tests/test_cache_memory_usage.py @@ -0,0 +1,209 @@ +""" +Calculate and measure the memory usage of embedding cache. + +This script estimates and measures the actual memory footprint +of the embedding cache at different sizes. +""" + +import sys +from collections import OrderedDict + + +def calculate_theoretical_memory(): + """Calculate theoretical memory usage for embedding cache.""" + print("=" * 70) + print("Theoretical Memory Usage Calculation") + print("=" * 70) + + # Constants + dimensions = 1024 # Default embedding dimensions + bytes_per_float = 8 # Python float is 64-bit + cache_sizes = [1000, 5000, 10000, 50000, 100000] + + print("\nAssumptions:") + print(f" • Embedding dimensions: {dimensions}") + print(f" • Bytes per float: {bytes_per_float}") + print(" • Hash key (SHA256 hex): 64 characters ≈ 128 bytes (UTF-8)") + print(" • Python list overhead: ~56 bytes") + print(" • Python string overhead: ~50 bytes") + print(" • OrderedDict per-entry overhead: ~230 bytes") + + # Memory per single cache entry + embedding_size = dimensions * bytes_per_float # Vector data + list_overhead = 56 # Python list object overhead + hash_key_size = 64 + 50 # String chars + string object overhead + ordereddict_entry_overhead = 230 # Dict entry + ordering overhead + + entry_size = embedding_size + list_overhead + hash_key_size + ordereddict_entry_overhead + + print(f"\n{'─' * 70}") + print("Memory per cache entry:") + print(f" • Embedding vector: {embedding_size:,} bytes ({embedding_size/1024:.1f} KB)") + print(f" • List overhead: {list_overhead} bytes") + print(f" • Hash key: {hash_key_size} bytes") + print(f" • OrderedDict overhead: {ordereddict_entry_overhead} bytes") + print(f" • Total per entry: {entry_size:,} bytes ({entry_size/1024:.2f} KB)") + + print(f"\n{'─' * 70}") + print("Memory usage at different cache sizes:") + print(f"{'─' * 70}") + print(f"{'Cache Size':>12} | {'Memory (MB)':>12} | {'Memory (GB)':>12}") + print(f"{'─' * 70}") + + for size in cache_sizes: + total_bytes = size * entry_size + mb = total_bytes / (1024 * 1024) + gb = total_bytes / (1024 * 1024 * 1024) + print(f"{size:>12,} | {mb:>12.2f} | {gb:>12.4f}") + + print(f"{'─' * 70}") + + # Highlight default size + default_size = 10000 + default_memory_mb = (default_size * entry_size) / (1024 * 1024) + + print(f"\n✨ Default cache size (max_cache_size={default_size:,}):") + print(f" Estimated memory: ~{default_memory_mb:.1f} MB") + print("\n💡 Recommendation:") + print(" • For memory-constrained environments: 1,000-5,000 (~8-42 MB)") + print(" • For balanced performance: 10,000 (~84 MB) [DEFAULT]") + print(" • For high-throughput applications: 50,000-100,000 (~420-840 MB)") + print(" • To disable cache: max_cache_size=0") + + +def measure_actual_memory(): + """Measure actual memory usage with real data.""" + print(f"\n\n{'=' * 70}") + print("Actual Memory Measurement") + print(f"{'=' * 70}") + + import random + + # Create a sample cache + cache = OrderedDict() + dimensions = 1024 + + test_sizes = [100, 1000, 5000, 10000] + + print("\nCreating embedding cache with random data...") + print(f"{'─' * 70}") + print(f"{'Cache Size':>12} | {'Memory (MB)':>12} | {'Per Entry (KB)':>15}") + print(f"{'─' * 70}") + + for size in test_sizes: + # Clear cache + cache.clear() + + # Fill with sample data + for i in range(size): + hash_key = f"hash_{i:064x}" # 64-char hex string + embedding = [random.random() for _ in range(dimensions)] + cache[hash_key] = embedding + + # Measure memory (approximate) + # Calculate size of all embeddings + total_bytes = 0 + for key, value in cache.items(): + total_bytes += sys.getsizeof(key) # Key size + total_bytes += sys.getsizeof(value) # List object + total_bytes += len(value) * sys.getsizeof(float()) # Float elements + + # Add OrderedDict overhead + total_bytes += sys.getsizeof(cache) + + mb = total_bytes / (1024 * 1024) + per_entry_kb = total_bytes / size / 1024 + + print(f"{size:>12,} | {mb:>12.2f} | {per_entry_kb:>15.2f}") + + print(f"{'─' * 70}") + + # Measure default size + cache.clear() + default_size = 10000 + print(f"\n📊 Measuring default size ({default_size:,} entries)...") + + for i in range(default_size): + hash_key = f"hash_{i:064x}" + embedding = [random.random() for _ in range(dimensions)] + cache[hash_key] = embedding + + total_bytes = sys.getsizeof(cache) + for key, value in cache.items(): + total_bytes += sys.getsizeof(key) + total_bytes += sys.getsizeof(value) + total_bytes += len(value) * sys.getsizeof(float()) + + mb = total_bytes / (1024 * 1024) + + print(f" Actual memory usage: {mb:.2f} MB") + print(f" Per entry: {total_bytes / default_size / 1024:.2f} KB") + + +def print_usage_guidelines(): + """Print guidelines for choosing cache size.""" + print(f"\n\n{'=' * 70}") + print("Cache Size Selection Guidelines") + print(f"{'=' * 70}") + + print("\n📋 How to choose the right cache size:\n") + + print("1️⃣ Estimate your query patterns:") + print(" • How many unique texts will you embed?") + print(" • What's the repetition rate?") + print(" • Example: 1,000 unique texts with 70% repetition → cache_size=1,000\n") + + print("2️⃣ Consider available memory:") + print(" • ~84 MB per 10,000 entries (1024-dim embeddings)") + print(" • ~420 MB per 50,000 entries") + print(" • Scale linearly: ~8.4 MB per 1,000 entries\n") + + print("3️⃣ Monitor cache statistics:") + print(" • Use model.get_cache_stats() to check hit rate") + print(" • If hit rate < 50%, consider reducing cache size") + print(" • If hit rate > 90%, you might benefit from larger cache\n") + + print("4️⃣ Configuration examples:\n") + + examples = [ + ("Embedded device / IoT", "max_cache_size=100", "~840 KB"), + ("Development / Testing", "max_cache_size=1000", "~8.4 MB"), + ("Production (default)", "max_cache_size=10000", "~84 MB"), + ("High-volume service", "max_cache_size=50000", "~420 MB"), + ("Disable cache", "max_cache_size=0", "0 MB"), + ] + + print(f"{'Use Case':<25} | {'Configuration':<22} | {'Memory':<10}") + print(f"{'─' * 25}-+-{'─' * 22}-+-{'─' * 10}") + for use_case, config, memory in examples: + print(f"{use_case:<25} | {config:<22} | {memory:<10}") + + print("\n💡 Code example:\n") + print(" # Default configuration (recommended)") + print(" model = OpenAIEmbeddingModel(") + print(" model_name='text-embedding-v4',") + print(" dimensions=1024,") + print(" max_cache_size=10000, # ~84 MB") + print(" )") + print() + print(" # Memory-constrained configuration") + print(" model = OpenAIEmbeddingModel(") + print(" model_name='text-embedding-v4',") + print(" dimensions=1024,") + print(" max_cache_size=1000, # ~8.4 MB") + print(" )") + print() + print(" # Monitor cache performance") + print(" stats = model.get_cache_stats()") + print(" print(f'Hit rate: {stats[\"hit_rate\"]:.1%}')") + print(" print(f'Memory used: ~{stats[\"cache_size\"] * 8.4:.1f} KB')") + + +if __name__ == "__main__": + calculate_theoretical_memory() + measure_actual_memory() + print_usage_guidelines() + + print(f"\n{'=' * 70}") + print("✅ Memory analysis complete") + print(f"{'=' * 70}\n") diff --git a/tests/test_chunking_utils.py b/tests/test_chunking_utils.py new file mode 100644 index 00000000..c0937021 --- /dev/null +++ b/tests/test_chunking_utils.py @@ -0,0 +1,268 @@ +"""Tests for chunking utilities.""" + +import pytest + +from reme.core.enumeration import MemorySource +from reme.core.utils.chunking_utils import chunk_markdown + + +def test_chunk_markdown_basic(): + """Test basic markdown chunking functionality.""" + text = """# Heading 1 + +This is a paragraph with some content. + +## Heading 2 + +Another paragraph here. +More content in this paragraph.""" + + chunks = chunk_markdown( + text=text, + path="test.md", + source=MemorySource.MEMORY, + chunk_tokens=100, + overlap=10, + ) + + assert len(chunks) > 0 + assert all(chunk.path == "test.md" for chunk in chunks) + assert all(chunk.source == MemorySource.MEMORY for chunk in chunks) + assert all(chunk.hash for chunk in chunks) + assert all(chunk.id for chunk in chunks) + + +def test_chunk_markdown_empty(): + """Test chunking with empty text.""" + chunks = chunk_markdown( + text="", + path="empty.md", + source=MemorySource.MEMORY, + chunk_tokens=100, + overlap=10, + ) + + # Empty string splits to [""] which creates one chunk with empty text + assert len(chunks) == 1 + assert chunks[0].text == "" + assert chunks[0].start_line == 1 + assert chunks[0].end_line == 1 + + +def test_chunk_markdown_single_line(): + """Test chunking with single line.""" + text = "Single line of text" + + chunks = chunk_markdown( + text=text, + path="single.md", + source=MemorySource.MEMORY, + chunk_tokens=100, + overlap=10, + ) + + assert len(chunks) == 1 + assert chunks[0].text == text + assert chunks[0].start_line == 1 + assert chunks[0].end_line == 1 + + +def test_chunk_markdown_long_line(): + """Test chunking with a very long line that exceeds max_chars.""" + # Create a line longer than max_chars (300 tokens * 4 = 1200 chars) + long_text = "x" * 1500 + + chunks = chunk_markdown( + text=long_text, + path="long.md", + source=MemorySource.MEMORY, + chunk_tokens=300, + overlap=30, + ) + + # Should split into multiple chunks + assert len(chunks) > 1 + # All chunks should have the same line number since it's one line + assert all(chunk.start_line == 1 for chunk in chunks) + assert all(chunk.end_line == 1 for chunk in chunks) + + +def test_chunk_markdown_overlap(): + """Test that overlap is working correctly.""" + text = "\n".join([f"Line {i}" for i in range(1, 51)]) + + chunks = chunk_markdown( + text=text, + path="overlap.md", + source=MemorySource.MEMORY, + chunk_tokens=50, + overlap=10, + ) + + # With overlap, consecutive chunks should have some overlapping content + if len(chunks) > 1: + for i in range(len(chunks) - 1): + # Check that there's potential overlap + assert chunks[i].end_line >= chunks[i].start_line + assert chunks[i + 1].start_line <= chunks[i].end_line + 1 + + +def test_chunk_markdown_no_overlap(): + """Test chunking without overlap.""" + text = "\n".join([f"Line {i}" for i in range(1, 51)]) + + chunks = chunk_markdown( + text=text, + path="no_overlap.md", + source=MemorySource.MEMORY, + chunk_tokens=50, + overlap=0, + ) + + assert len(chunks) > 0 + # Verify all chunks are non-empty + assert all(chunk.text for chunk in chunks) + + +def test_chunk_markdown_line_numbers(): + """Test that line numbers are correctly assigned.""" + text = """Line 1 +Line 2 +Line 3 +Line 4 +Line 5""" + + chunks = chunk_markdown( + text=text, + path="lines.md", + source=MemorySource.MEMORY, + chunk_tokens=20, + overlap=5, + ) + + # First chunk should start at line 1 + assert chunks[0].start_line == 1 + # Last chunk should end at the last line + assert chunks[-1].end_line == 5 + # All chunks should have valid line ranges + for chunk in chunks: + assert chunk.start_line <= chunk.end_line + assert chunk.start_line >= 1 + + +def test_chunk_markdown_hash_uniqueness(): + """Test that different chunks have different hashes.""" + text = """# Section 1 + +Content for section 1. + +# Section 2 + +Content for section 2.""" + + chunks = chunk_markdown( + text=text, + path="sections.md", + source=MemorySource.MEMORY, + chunk_tokens=50, + overlap=5, + ) + + # Collect all hashes + hashes = [chunk.hash for chunk in chunks] + + # If we have multiple chunks, they should have different hashes + if len(chunks) > 1: + assert len(set(hashes)) == len(hashes), "Chunks should have unique hashes" + + +def test_chunk_markdown_id_uniqueness(): + """Test that chunk IDs are unique.""" + text = "\n".join([f"Line {i}" for i in range(1, 101)]) + + chunks = chunk_markdown( + text=text, + path="unique.md", + source=MemorySource.MEMORY, + chunk_tokens=50, + overlap=10, + ) + + ids = [chunk.id for chunk in chunks] + assert len(set(ids)) == len(ids), "All chunk IDs should be unique" + + +def test_chunk_markdown_small_tokens(): + """Test chunking with very small token limit.""" + text = "This is a short text with multiple words in it." + + chunks = chunk_markdown( + text=text, + path="small.md", + source=MemorySource.MEMORY, + chunk_tokens=5, + overlap=1, + ) + + # Even with small chunk size, should create at least one chunk + assert len(chunks) >= 1 + + +def test_chunk_markdown_sessions_source(): + """Test chunking with SESSIONS source.""" + text = "Session log content" + + chunks = chunk_markdown( + text=text, + path="session.log", + source=MemorySource.SESSIONS, + chunk_tokens=100, + overlap=10, + ) + + assert len(chunks) == 1 + assert chunks[0].source == MemorySource.SESSIONS + + +def test_chunk_markdown_multiline_paragraph(): + """Test chunking with realistic markdown content.""" + text = """# Introduction + +This is a longer paragraph that spans multiple lines. +It contains various sentences and information. +The chunking algorithm should handle this properly. + +## Details + +Here are some details: +- Point 1 +- Point 2 +- Point 3 + +## Conclusion + +Final thoughts and conclusions go here.""" + + chunks = chunk_markdown( + text=text, + path="document.md", + source=MemorySource.MEMORY, + chunk_tokens=50, + overlap=10, + ) + + # Should create multiple chunks + assert len(chunks) > 0 + + # Verify text reconstruction + all_text_parts = [] + for chunk in chunks: + all_text_parts.append(chunk.text) + + # At least some of the original content should be in the chunks + combined = "\n".join(all_text_parts) + assert "Introduction" in combined or "Introduction" in text + + +if __name__ == "__main__": + pytest.main([__file__, "-v"]) diff --git a/tests/test_embedding_cache.py b/tests/test_embedding_cache.py new file mode 100644 index 00000000..94b6673e --- /dev/null +++ b/tests/test_embedding_cache.py @@ -0,0 +1,394 @@ +""" +Async unit tests for embedding cache functionality. + +Tests cover: +- Cache hit/miss tracking +- LRU eviction policy +- Cache statistics +- Performance improvements with repeated queries +- Cache clearing + +Usage: + python test_embedding_cache.py +""" + +# flake8: noqa: E402 +# pylint: disable=C0413 + +import asyncio +from typing import List + +from reme.core.utils import load_env + +load_env() + +from reme.core.embedding import OpenAIEmbeddingModel + + +def get_test_texts() -> List[str]: + """Create test texts for embedding cache testing.""" + return [ + "What is machine learning?", + "How does neural network work?", + "Explain artificial intelligence", + "Define deep learning", + "What is data science?", + ] + + +async def test_cache_basic_functionality(): + """Test basic cache hit/miss functionality.""" + print(f"\n{'='*60}") + print("Test 1: Basic Cache Functionality") + print(f"{'='*60}") + + model = OpenAIEmbeddingModel( + model_name="text-embedding-v4", + dimensions=1024, + max_cache_size=100, + max_retries=2, + raise_exception=True, + ) + + test_text = "Hello, this is a test sentence for embedding cache." + + print(f"Input text: {test_text}") + + # First call - should be a cache miss + print("\n1️⃣ First embedding call (cold cache):") + embedding1 = await model.get_embedding(test_text) + stats1 = model.get_cache_stats() + + print(f" Embedding dimension: {len(embedding1)}") + print(f" Cache size: {stats1['cache_size']}") + print(f" Cache hits: {stats1['cache_hits']}") + print(f" Cache misses: {stats1['cache_misses']}") + print(f" Hit rate: {stats1['hit_rate']:.2%}") + + assert len(embedding1) == 1024, "Embedding dimension mismatch" + assert stats1["cache_misses"] == 1, "Should have 1 cache miss" + assert stats1["cache_hits"] == 0, "Should have 0 cache hits" + assert stats1["cache_size"] == 1, "Cache should have 1 entry" + + # Second call with same text - should be a cache hit + print("\n2️⃣ Second embedding call (same text):") + embedding2 = await model.get_embedding(test_text) + stats2 = model.get_cache_stats() + + print(f" Cache hits: {stats2['cache_hits']}") + print(f" Cache misses: {stats2['cache_misses']}") + print(f" Hit rate: {stats2['hit_rate']:.2%}") + + assert embedding1 == embedding2, "Cached embedding should be identical" + assert stats2["cache_hits"] == 1, "Should have 1 cache hit" + assert stats2["cache_misses"] == 1, "Should still have 1 cache miss" + assert stats2["hit_rate"] == 0.5, "Hit rate should be 50%" + + await model.close() + print("\n✓ PASSED: Basic cache functionality works correctly") + + +async def test_batch_cache_efficiency(): + """Test cache efficiency with batch embeddings including duplicates.""" + print(f"\n{'='*60}") + print("Test 2: Batch Cache Efficiency") + print(f"{'='*60}") + + model = OpenAIEmbeddingModel( + model_name="text-embedding-v4", + dimensions=1024, + max_cache_size=1000, + max_retries=2, + raise_exception=True, + ) + + texts = get_test_texts() + + # Create a list with duplicates + texts_with_duplicates = texts + texts[:3] # 5 unique + 3 duplicates = 8 total + + print(f"Processing {len(texts_with_duplicates)} texts (5 unique + 3 duplicates)") + + # First batch + print("\n1️⃣ First batch (cold cache):") + embeddings1 = await model.get_embeddings(texts) + stats1 = model.get_cache_stats() + + print(f" Embeddings generated: {len(embeddings1)}") + print(f" Cache size: {stats1['cache_size']}") + print(f" Cache misses: {stats1['cache_misses']}") + print(f" Cache hits: {stats1['cache_hits']}") + + assert len(embeddings1) == len(texts), "Embeddings count mismatch" + assert stats1["cache_size"] == len(texts), f"Cache should have {len(texts)} entries" + assert stats1["cache_misses"] == len(texts), "All should be cache misses" + + # Second batch with duplicates + print("\n2️⃣ Second batch (with duplicates):") + embeddings2 = await model.get_embeddings(texts_with_duplicates) + stats2 = model.get_cache_stats() + + print(f" Embeddings generated: {len(embeddings2)}") + print(f" Cache hits: {stats2['cache_hits']}") + print(f" Cache misses: {stats2['cache_misses']}") + print(f" Hit rate: {stats2['hit_rate']:.2%}") + + assert len(embeddings2) == len(texts_with_duplicates), "Embeddings count mismatch" + assert stats2["cache_hits"] >= 3, "Should have at least 3 cache hits from duplicates" + + # Verify embeddings are identical for duplicated texts + for i in range(3): + assert embeddings2[i] == embeddings2[len(texts) + i], f"Duplicate {i} should have identical embedding" + + await model.close() + print("\n✓ PASSED: Batch cache efficiently handles duplicates") + + +async def test_cache_lru_eviction(): + """Test LRU cache eviction policy.""" + print(f"\n{'='*60}") + print("Test 3: LRU Cache Eviction") + print(f"{'='*60}") + + # Create model with small cache size + model = OpenAIEmbeddingModel( + model_name="text-embedding-v4", + dimensions=1024, + max_cache_size=3, # Small cache for testing eviction + max_retries=2, + raise_exception=True, + ) + + texts = get_test_texts()[:5] # Use 5 texts, cache size is 3 + + print(f"Cache size limit: {model.max_cache_size}") + print(f"Number of unique texts: {len(texts)}") + + # Fill cache beyond capacity + print("\n1️⃣ Filling cache with 5 texts (capacity = 3):") + for i, text in enumerate(texts): + await model.get_embedding(text) + stats = model.get_cache_stats() + print( + f" After text {i+1}: cache_size={stats['cache_size']}, " + f"hits={stats['cache_hits']}, misses={stats['cache_misses']}", + ) + + final_stats = model.get_cache_stats() + assert final_stats["cache_size"] <= 3, "Cache size should not exceed max_cache_size" + assert final_stats["cache_misses"] == 5, "Should have 5 cache misses for 5 unique texts" + + # Access the most recent entries - should be cache hits + print("\n2️⃣ Accessing recent entries (should be cached):") + recent_texts = texts[-3:] # Last 3 texts should still be in cache + + for i, text in enumerate(recent_texts): + await model.get_embedding(text) + stats = model.get_cache_stats() + print(f" Text {len(texts) - 3 + i + 1}: hits={stats['cache_hits']}") + + final_stats = model.get_cache_stats() + assert final_stats["cache_hits"] == 3, "Should have 3 cache hits for recent entries" + + # Access oldest entries - should be cache misses (evicted) + print("\n3️⃣ Accessing oldest entries (should be evicted):") + old_texts = texts[:2] # First 2 texts should have been evicted + + before_misses = final_stats["cache_misses"] + for i, text in enumerate(old_texts): + await model.get_embedding(text) + stats = model.get_cache_stats() + print(f" Text {i + 1}: misses={stats['cache_misses']}") + + final_stats = model.get_cache_stats() + assert final_stats["cache_misses"] == before_misses + 2, "Should have 2 more cache misses for evicted entries" + + await model.close() + print("\n✓ PASSED: LRU eviction works correctly") + + +async def test_cache_stats_and_clear(): + """Test cache statistics tracking and clearing.""" + print(f"\n{'='*60}") + print("Test 4: Cache Statistics and Clearing") + print(f"{'='*60}") + + model = OpenAIEmbeddingModel( + model_name="text-embedding-v4", + dimensions=1024, + max_cache_size=100, + max_retries=2, + raise_exception=True, + ) + + texts = get_test_texts() + + # Generate some cache activity + print("\n1️⃣ Generating cache activity:") + await model.get_embeddings(texts) + await model.get_embeddings(texts[:3]) # Repeat first 3 + + stats = model.get_cache_stats() + print(f" Cache size: {stats['cache_size']}") + print(f" Max cache size: {stats['max_cache_size']}") + print(f" Cache hits: {stats['cache_hits']}") + print(f" Cache misses: {stats['cache_misses']}") + print(f" Hit rate: {stats['hit_rate']:.2%}") + + assert stats["cache_size"] > 0, "Cache should not be empty" + assert stats["cache_hits"] >= 3, "Should have at least 3 cache hits" + assert "hit_rate" in stats, "Stats should include hit_rate" + + # Clear cache + print("\n2️⃣ Clearing cache:") + model.clear_cache() + stats_after_clear = model.get_cache_stats() + + print(f" Cache size after clear: {stats_after_clear['cache_size']}") + print(f" Hits after clear: {stats_after_clear['cache_hits']}") + print(f" Misses after clear: {stats_after_clear['cache_misses']}") + print(f" Hit rate after clear: {stats_after_clear['hit_rate']:.2%}") + + assert stats_after_clear["cache_size"] == 0, "Cache should be empty after clear" + assert stats_after_clear["cache_hits"] == 0, "Hits should be reset" + assert stats_after_clear["cache_misses"] == 0, "Misses should be reset" + assert stats_after_clear["hit_rate"] == 0.0, "Hit rate should be 0" + + await model.close() + print("\n✓ PASSED: Cache statistics and clearing work correctly") + + +async def test_cache_disabled(): + """Test behavior when cache is disabled (max_cache_size=0).""" + print(f"\n{'='*60}") + print("Test 5: Cache Disabled") + print(f"{'='*60}") + + model = OpenAIEmbeddingModel( + model_name="text-embedding-v4", + dimensions=1024, + max_cache_size=0, # Disable cache + max_retries=2, + raise_exception=True, + ) + + test_text = "Test text with cache disabled" + + print(f"Cache size limit: {model.max_cache_size} (disabled)") + print(f"Input text: {test_text}") + + # Call twice with same text + print("\n1️⃣ First call:") + embedding1 = await model.get_embedding(test_text) + stats1 = model.get_cache_stats() + print(f" Cache size: {stats1['cache_size']}") + print(f" Cache misses: {stats1['cache_misses']}") + + print("\n2️⃣ Second call (same text):") + embedding2 = await model.get_embedding(test_text) + stats2 = model.get_cache_stats() + print(f" Cache size: {stats2['cache_size']}") + print(f" Cache misses: {stats2['cache_misses']}") + print(f" Cache hits: {stats2['cache_hits']}") + + assert stats2["cache_size"] == 0, "Cache should remain empty when disabled" + assert stats2["cache_misses"] == 2, "Both calls should be cache misses" + assert stats2["cache_hits"] == 0, "Should have no cache hits when disabled" + assert embedding1 == embedding2, "Embeddings should still be consistent" + + await model.close() + print("\n✓ PASSED: Cache correctly disabled when max_cache_size=0") + + +async def test_cache_performance_demo(): + """Demonstrate cache performance improvements.""" + print(f"\n{'='*60}") + print("Test 6: Cache Performance Demo") + print(f"{'='*60}") + + model = OpenAIEmbeddingModel( + model_name="text-embedding-v4", + dimensions=1024, + max_cache_size=1000, + max_retries=2, + raise_exception=True, + ) + + texts = get_test_texts() + + # Create a realistic workload with many repeated queries + workload = texts * 3 # 15 queries total, 5 unique + + print(f"\nProcessing {len(workload)} queries ({len(texts)} unique texts)") + print("This simulates a realistic scenario with repeated queries\n") + + # Process all queries + for i, text in enumerate(workload, 1): + await model.get_embedding(text) + if i % 5 == 0: # Report every 5 queries + stats = model.get_cache_stats() + print( + f"After {i:2d} queries: hits={stats['cache_hits']:2d}, " + f"misses={stats['cache_misses']:2d}, " + f"hit_rate={stats['hit_rate']:5.1%}", + ) + + final_stats = model.get_cache_stats() + total_requests = final_stats["cache_hits"] + final_stats["cache_misses"] + + print(f"\n{'─'*60}") + print("📊 Final Statistics:") + print(f"{'─'*60}") + print(f" Total queries: {total_requests}") + print(f" Unique texts: {len(texts)}") + print(f" Cache hits: {final_stats['cache_hits']}") + print(f" Cache misses: {final_stats['cache_misses']}") + print(f" Hit rate: {final_stats['hit_rate']:.1%}") + print(f" Cache size: {final_stats['cache_size']}/{final_stats['max_cache_size']}") + print(f"{'─'*60}") + print( + f"💰 API calls saved: {final_stats['cache_hits']} out of {total_requests} " + f"({final_stats['cache_hits']/total_requests*100:.1f}%)", + ) + print(f"{'─'*60}") + + assert final_stats["cache_hits"] == 10, "Should have 10 cache hits (2 repeats × 5 texts)" + assert final_stats["cache_misses"] == 5, "Should have 5 cache misses (5 unique texts)" + assert final_stats["hit_rate"] > 0.6, "Hit rate should be > 60%" + + await model.close() + print("\n✓ PASSED: Cache provides significant performance improvement") + + +async def main(): + """Run all cache tests.""" + print("\n" + "#" * 60) + print("# EMBEDDING CACHE TESTS") + print("#" * 60) + + try: + await test_cache_basic_functionality() + await test_batch_cache_efficiency() + await test_cache_lru_eviction() + await test_cache_stats_and_clear() + await test_cache_disabled() + await test_cache_performance_demo() + + print("\n" + "=" * 60) + print("✅ ALL CACHE TESTS PASSED") + print("=" * 60) + print("\nKey takeaways:") + print(" • Cache correctly tracks hits/misses") + print(" • LRU eviction works as expected") + print(" • Duplicate queries are efficiently cached") + print(" • Cache can be disabled or cleared") + print(" • Significant performance improvement with realistic workloads") + print("=" * 60 + "\n") + + except Exception as e: + print(f"\n✗ TEST FAILED: {type(e).__name__}: {e}") + raise + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/test_file_system_tool.py b/tests/test_file_system_tool.py new file mode 100644 index 00000000..15f2d858 --- /dev/null +++ b/tests/test_file_system_tool.py @@ -0,0 +1,445 @@ +"""Tests for file system tools including bash, edit, find, grep, ls, read, and write tools.""" + +import asyncio +import os +import tempfile +from pathlib import Path + + +async def test_bash_tool(): + """Test BashTool.""" + from reme.tool.fs import BashTool + + print("=== Testing BashTool ===") + bash_tool = BashTool() + result = await bash_tool.call(command="echo 'Hello World'") + print(f"Result: {result}") + assert "Hello World" in result + print("✓ BashTool test passed\n") + + +async def test_edit_tool(): + """Test EditTool.""" + from reme.tool.fs import EditTool + + print("=== Testing EditTool ===") + + # Create temp file + with tempfile.NamedTemporaryFile(mode="w", suffix=".txt", delete=False) as f: + temp_path = f.name + f.write("Hello World\nThis is a test\nGoodbye World\n") + + try: + # Test edit + edit_tool = EditTool() + result = await edit_tool.call( + path=temp_path, + oldText="This is a test", + newText="This is an updated test", + ) + print(f"Result: {result}") + + # Verify content + with open(temp_path, "r", encoding="utf-8") as f: + content = f.read() + + assert "This is an updated test" in content + assert "This is a test" not in content + print("✓ EditTool test passed\n") + + # Test error: file not found + print("=== Testing file not found error ===") + result = await edit_tool.call( + path="/nonexistent/file.txt", + oldText="test", + newText="new", + ) + print(f"Expected error result: {result}") + assert "failed" in result and "File not found" in result + print("✓ File not found error test passed\n") + + # Test error: text not found + print("=== Testing text not found error ===") + result = await edit_tool.call( + path=temp_path, + oldText="nonexistent text", + newText="new", + ) + print(f"Expected error result: {result}") + assert "failed" in result and "Could not find" in result + print("✓ Text not found error test passed\n") + + finally: + # Cleanup + if os.path.exists(temp_path): + os.unlink(temp_path) + + +async def test_find_tool(): + """Test FindTool.""" + from reme.tool.fs import FindTool + + print("=== Testing FindTool ===") + + # Create temp directory with test files + with tempfile.TemporaryDirectory() as temp_dir: + temp_path = Path(temp_dir) + + # Create some test files + (temp_path / "test1.txt").write_text("test file 1") + (temp_path / "test2.txt").write_text("test file 2") + (temp_path / "readme.md").write_text("readme") + + # Create subdirectory with files + sub_dir = temp_path / "subdir" + sub_dir.mkdir() + (sub_dir / "test3.txt").write_text("test file 3") + (sub_dir / "config.json").write_text("{}") + + # Create .gitignore to ignore certain files + (temp_path / ".gitignore").write_text("*.md\n") + + # Test: find all txt files + find_tool = FindTool(cwd=str(temp_path)) + result = await find_tool.call(pattern="*.txt") + print(f"Find *.txt result:\n{result}") + assert "test1.txt" in result + assert "test2.txt" in result + assert "readme.md" not in result # Should be ignored by .gitignore + print("✓ Find *.txt test passed\n") + + # Test: find with recursive pattern + result = await find_tool.call(pattern="**/*.txt") + print(f"Find **/*.txt result:\n{result}") + assert "test1.txt" in result + assert "subdir/test3.txt" in result or "test3.txt" in result + print("✓ Find **/*.txt test passed\n") + + # Test: find with no matches + result = await find_tool.call(pattern="*.nonexistent") + print(f"Find *.nonexistent result:\n{result}") + assert "No files found" in result + print("✓ No matches test passed\n") + + # Test: error - directory not found + print("=== Testing directory not found error ===") + result = await find_tool.call(pattern="*.txt", path="/nonexistent/dir") + print(f"Expected error result: {result}") + assert "failed" in result and "Path not found" in result + print("✓ Directory not found error test passed\n") + + +async def test_grep_tool(): + """Test GrepTool.""" + from reme.tool.fs import GrepTool + + print("=== Testing GrepTool ===") + + # Create temp directory with test files + with tempfile.TemporaryDirectory() as temp_dir: + temp_path = Path(temp_dir) + + # Create test files with searchable content + (temp_path / "file1.txt").write_text("Hello World\nThis is a test\nGoodbye World\n") + (temp_path / "file2.txt").write_text("Another test file\nWith multiple lines\nHello again\n") + (temp_path / "script.py").write_text("def hello():\n print('Hello')\n return True\n") + + # Create subdirectory with files + sub_dir = temp_path / "subdir" + sub_dir.mkdir() + (sub_dir / "nested.txt").write_text("Nested file content\nWith hello keyword\n") + + # Test: search for pattern + grep_tool = GrepTool(cwd=str(temp_path)) + result = await grep_tool.call(pattern="Hello", path=str(temp_path)) + print(f"Search 'Hello' result:\n{result}") + assert "file1.txt" in result + assert "Hello World" in result or "Hello" in result + print("✓ Basic search test passed\n") + + # Test: case-insensitive search + result = await grep_tool.call(pattern="hello", path=str(temp_path), ignoreCase=True) + print(f"Case-insensitive search result:\n{result}") + assert "file1.txt" in result or "Hello" in result.lower() + print("✓ Case-insensitive search test passed\n") + + # Test: literal string search + result = await grep_tool.call(pattern="Hello()", path=str(temp_path), literal=True) + print(f"Literal search result:\n{result}") + # Should not find regex interpretation + print("✓ Literal search test passed\n") + + # Test: glob filter + result = await grep_tool.call(pattern="Hello", path=str(temp_path), glob="*.txt") + print(f"Glob filter *.txt result:\n{result}") + assert "file1.txt" in result or "file2.txt" in result + assert ".py" not in result # Python files should be excluded + print("✓ Glob filter test passed\n") + + # Test: context lines + result = await grep_tool.call(pattern="test", path=str(temp_path), contextLines=1) + print(f"Context lines result:\n{result}") + # Should include lines before and after matches + print("✓ Context lines test passed\n") + + # Test: limit matches + result = await grep_tool.call(pattern="Hello", path=str(temp_path), limit=1) + print(f"Limit to 1 match result:\n{result}") + assert "limit reached" in result or result.count(":") >= 1 + print("✓ Limit test passed\n") + + # Test: no matches + result = await grep_tool.call(pattern="nonexistent_pattern_xyz", path=str(temp_path)) + print(f"No matches result:\n{result}") + assert "No matches found" in result + print("✓ No matches test passed\n") + + # Test: error - path not found + print("=== Testing path not found error ===") + try: + result = await grep_tool.call(pattern="test", path="/nonexistent/path") + assert "failed" in result and "not found" in result + except Exception as e: + assert "not found" in str(e).lower() + print("✓ Path not found error test passed\n") + + +async def test_ls_tool(): + """Test LsTool.""" + from reme.tool.fs import LsTool + + print("=== Testing LsTool ===") + + # Create temp directory with test files + with tempfile.TemporaryDirectory() as temp_dir: + temp_path = Path(temp_dir) + + # Create test files and directories + (temp_path / "file1.txt").write_text("test file 1") + (temp_path / "file2.py").write_text("test file 2") + (temp_path / ".hidden").write_text("hidden file") + (temp_path / "README.md").write_text("readme") + + # Create subdirectories + (temp_path / "subdir1").mkdir() + (temp_path / "subdir2").mkdir() + + # Test: list current directory + ls_tool = LsTool(cwd=str(temp_path)) + result = await ls_tool.call() + print(f"List directory result:\n{result}") + assert ".hidden" in result # Includes dotfiles + assert "file1.txt" in result + assert "file2.py" in result + assert "subdir1/" in result # Directories have '/' suffix + assert "subdir2/" in result + print("✓ Basic ls test passed\n") + + # Test: list specific path + result = await ls_tool.call(path=".") + print(f"List current directory result:\n{result}") + assert "file1.txt" in result + print("✓ Specific path test passed\n") + + # Test: empty directory + empty_dir = temp_path / "empty" + empty_dir.mkdir() + result = await ls_tool.call(path="empty") + print(f"Empty directory result:\n{result}") + assert "(empty directory)" in result + print("✓ Empty directory test passed\n") + + # Test: entry limit + # Create many files + for i in range(10): + (temp_path / f"file{i:03d}.txt").write_text(f"file {i}") + + result = await ls_tool.call(limit=5) + print(f"Limited entries result:\n{result}") + assert "entries limit reached" in result + assert "limit=10" in result # Should suggest doubling the limit + print("✓ Entry limit test passed\n") + + # Test: error - path not found + print("=== Testing path not found error ===") + result = await ls_tool.call(path="/nonexistent/path") + print(f"Expected error result: {result}") + assert "failed" in result and "Path not found" in result + print("✓ Path not found error test passed\n") + + # Test: error - not a directory + print("=== Testing not a directory error ===") + result = await ls_tool.call(path="file1.txt") + print(f"Expected error result: {result}") + assert "failed" in result and "Not a directory" in result + print("✓ Not a directory error test passed\n") + + +async def test_read_tool(): + """Test ReadTool.""" + from reme.tool.fs import ReadTool + + print("=== Testing ReadTool ===") + + # Create temp directory with test files + with tempfile.TemporaryDirectory() as temp_dir: + temp_path = Path(temp_dir) + + # Create test text file + test_file = temp_path / "test.txt" + test_content = "\n".join([f"Line {i}" for i in range(1, 101)]) # 100 lines + test_file.write_text(test_content) + + # Create test image file + image_file = temp_path / "test.jpg" + image_file.write_bytes(b"\xff\xd8\xff\xe0") # Minimal JPEG header + + # Test: read full file + read_tool = ReadTool(cwd=str(temp_path)) + result = await read_tool.call(path="test.txt") + print(f"Read full file result:\n{result[:200]}...") + assert "Line 1" in result + assert "Line 100" in result + print("✓ Read full file test passed\n") + + # Test: read with offset + result = await read_tool.call(path="test.txt", offset=50) + print(f"Read with offset=50 result:\n{result[:200]}...") + # Check that we start from Line 50 (should be first line of content) + assert result.startswith("Line 50"), f"Should start with 'Line 50', got: {result[:50]}" + assert "Line 100" in result + print("✓ Read with offset test passed\n") + + # Test: read with limit + result = await read_tool.call(path="test.txt", limit=10) + print(f"Read with limit=10 result:\n{result}") + assert "Line 1" in result + assert "Line 10" in result or "more lines in file" in result + assert "Line 50" not in result + print("✓ Read with limit test passed\n") + + # Test: read with offset and limit + result = await read_tool.call(path="test.txt", offset=20, limit=5) + print(f"Read with offset=20, limit=5 result:\n{result}") + assert "Line 20" in result + assert "Line 24" in result or "more lines" in result + print("✓ Read with offset and limit test passed\n") + + # Test: read image file + result = await read_tool.call(path="test.jpg") + print(f"Read image result:\n{result}") + assert "image file" in result.lower() or ".jpg" in result.lower() + print("✓ Read image test passed\n") + + # Test: offset beyond file + print("=== Testing offset beyond file error ===") + result = await read_tool.call(path="test.txt", offset=200) + print(f"Expected error result: {result}") + assert "failed" in result and ("beyond end of file" in result or "offset" in result.lower()) + print("✓ Offset beyond file error test passed\n") + + # Test: file not found + print("=== Testing file not found error ===") + result = await read_tool.call(path="nonexistent.txt") + print(f"Expected error result: {result}") + assert "failed" in result and "not found" in result.lower() + print("✓ File not found error test passed\n") + + # Test: read directory (should fail) + print("=== Testing read directory error ===") + sub_dir = temp_path / "subdir" + sub_dir.mkdir() + result = await read_tool.call(path="subdir") + print(f"Expected error result: {result}") + assert "failed" in result and ("Not a file" in result or "directory" in result.lower()) + print("✓ Read directory error test passed\n") + + +async def test_write_tool(): + """Test WriteTool.""" + from reme.tool.fs import WriteTool + + print("=== Testing WriteTool ===") + + # Create temp directory + with tempfile.TemporaryDirectory() as temp_dir: + temp_path = Path(temp_dir) + + # Test: write new file + write_tool = WriteTool(cwd=str(temp_path)) + test_content = "Hello World\nThis is a test file\n" + result = await write_tool.call(path="test.txt", content=test_content) + print(f"Write result: {result}") + assert "Successfully wrote" in result + assert "test.txt" in result + + # Verify file was created + test_file = temp_path / "test.txt" + assert test_file.exists() + assert test_file.read_text() == test_content + print("✓ Write new file test passed\n") + + # Test: overwrite existing file + new_content = "Updated content\n" + result = await write_tool.call(path="test.txt", content=new_content) + print(f"Overwrite result: {result}") + assert "Successfully wrote" in result + + # Verify file was overwritten + assert test_file.read_text() == new_content + assert test_content not in test_file.read_text() + print("✓ Overwrite existing file test passed\n") + + # Test: create file with parent directories + nested_path = "subdir1/subdir2/nested.txt" + nested_content = "Nested file content" + result = await write_tool.call(path=nested_path, content=nested_content) + print(f"Create with parents result: {result}") + assert "Successfully wrote" in result + + # Verify nested file was created + nested_file = temp_path / "subdir1" / "subdir2" / "nested.txt" + assert nested_file.exists() + assert nested_file.read_text() == nested_content + print("✓ Create file with parent directories test passed\n") + + # Test: write empty file + result = await write_tool.call(path="empty.txt", content="") + print(f"Write empty file result: {result}") + assert "Successfully wrote" in result + assert "0 bytes" in result + + # Verify empty file + empty_file = temp_path / "empty.txt" + assert empty_file.exists() + assert empty_file.read_text() == "" + print("✓ Write empty file test passed\n") + + # Test: write file with absolute path + abs_path = str(temp_path / "absolute.txt") + abs_content = "Absolute path content" + result = await write_tool.call(path=abs_path, content=abs_content) + print(f"Write absolute path result: {result}") + assert "Successfully wrote" in result + + # Verify absolute path file + abs_file = Path(abs_path) + assert abs_file.exists() + assert abs_file.read_text(encoding="utf-8") == abs_content + print("✓ Write absolute path test passed\n") + + +async def main(): + """Run all file system tool tests.""" + await test_bash_tool() + await test_edit_tool() + await test_find_tool() + await test_grep_tool() + await test_ls_tool() + await test_read_tool() + await test_write_tool() + print("=== All tests passed! ===") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/test_fs_agent.py b/tests/test_fs_agent.py new file mode 100644 index 00000000..47995785 --- /dev/null +++ b/tests/test_fs_agent.py @@ -0,0 +1,323 @@ +"""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_memory_store.py b/tests/test_memory_store.py new file mode 100644 index 00000000..a057c897 --- /dev/null +++ b/tests/test_memory_store.py @@ -0,0 +1,979 @@ +# pylint: disable=too-many-lines +"""Unified test suite for memory store implementations. + +This module provides comprehensive test coverage for SqliteMemoryStore and future +memory store implementations. Tests can be run for specific stores or all implementations. + +Usage: + python test_memory_store.py --sqlite # Test SqliteMemoryStore only + python test_memory_store.py --all # Test all memory stores +""" + +import argparse +import asyncio +import hashlib +import shutil +import time +from pathlib import Path +from typing import List + +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.schema.file_metadata import FileMetadata +from reme.core.schema.memory_chunk import MemoryChunk +from reme.core.utils import load_env + +# Direct imports to avoid circular dependencies + +load_env() + + +# ==================== Configuration ==================== + + +class TestConfig: + """Configuration for test execution.""" + + # SqliteMemoryStore settings + SQLITE_DB_PATH = "./test_memory_store_sqlite/memory.db" + SQLITE_VEC_EXT_PATH = "" # Empty string to use default vec0/sqlite_vec/vector0 + SQLITE_FTS_ENABLED = True + SQLITE_SNIPPET_MAX_CHARS = 700 + + # Embedding model settings + EMBEDDING_MODEL_NAME = "text-embedding-v4" + EMBEDDING_DIMENSIONS = 64 + + # Test prefix for cleanup + TEST_PATH_PREFIX = "test_memory_" + + +# ==================== Sample Data Generator ==================== + + +class SampleDataGenerator: + """Generator for sample test data.""" + + @staticmethod + def create_sample_chunks(file_path: str, prefix: str = "") -> List[MemoryChunk]: + """Create sample MemoryChunk instances for testing. + + Args: + file_path: Path to the file + prefix: Optional prefix for chunk_id to avoid conflicts + + Returns: + List[MemoryChunk]: List of sample chunks with diverse content + """ + id_prefix = f"{prefix}_" if prefix else "" + base_hash = hashlib.md5(file_path.encode()).hexdigest()[:8] + + return [ + MemoryChunk( + id=f"{id_prefix}chunk1_{base_hash}", + path=file_path, + source=MemorySource.MEMORY, + start_line=1, + end_line=5, + text="Artificial intelligence is a technology that simulates human intelligence.", + hash=hashlib.md5(b"chunk1").hexdigest(), + embedding=None, # Will be populated later + metadata={"category": "AI", "importance": "high"}, + ), + MemoryChunk( + id=f"{id_prefix}chunk2_{base_hash}", + path=file_path, + source=MemorySource.MEMORY, + start_line=6, + end_line=10, + text="Machine learning is a subset of artificial intelligence that learns from data.", + hash=hashlib.md5(b"chunk2").hexdigest(), + embedding=None, + metadata={"category": "ML", "importance": "high"}, + ), + MemoryChunk( + id=f"{id_prefix}chunk3_{base_hash}", + path=file_path, + source=MemorySource.MEMORY, + start_line=11, + end_line=15, + text="Deep learning uses neural networks with multiple layers for complex tasks.", + hash=hashlib.md5(b"chunk3").hexdigest(), + embedding=None, + metadata={"category": "DL", "importance": "medium"}, + ), + ] + + @staticmethod + def create_session_chunks(file_path: str, prefix: str = "") -> List[MemoryChunk]: + """Create sample session chunks for testing. + + Args: + file_path: Path to the session file + prefix: Optional prefix for chunk_id to avoid conflicts + + Returns: + List[MemoryChunk]: List of session chunks + """ + id_prefix = f"{prefix}_" if prefix else "" + base_hash = hashlib.md5(file_path.encode()).hexdigest()[:8] + + return [ + MemoryChunk( + id=f"{id_prefix}session1_{base_hash}", + path=file_path, + source=MemorySource.SESSIONS, + start_line=1, + end_line=3, + text="User requested to analyze sales data for Q4 2024.", + hash=hashlib.md5(b"session1").hexdigest(), + embedding=None, + metadata={"session_id": "sess_001", "timestamp": "2024-12-01"}, + ), + MemoryChunk( + id=f"{id_prefix}session2_{base_hash}", + path=file_path, + source=MemorySource.SESSIONS, + start_line=4, + end_line=6, + text="Analyzed sales data and found a 15% increase in revenue.", + hash=hashlib.md5(b"session2").hexdigest(), + embedding=None, + metadata={"session_id": "sess_001", "timestamp": "2024-12-01"}, + ), + ] + + @staticmethod + def create_file_metadata(file_path: str, chunk_count: int = 0) -> FileMetadata: + """Create sample FileMetadata for testing. + + Args: + file_path: Path to the file + chunk_count: Number of chunks (optional) + + Returns: + FileMetadata: Sample file metadata + """ + content = f"Sample content for {file_path}" + return FileMetadata( + path=file_path, + hash=hashlib.md5(content.encode()).hexdigest(), + mtime_ms=time.time() * 1000, + size=len(content), + chunk_count=chunk_count, + ) + + +# ==================== Memory Store Factory ==================== + + +def get_store_type(store: BaseMemoryStore) -> str: + """Get the type identifier of a memory store instance. + + Args: + store: Memory store instance + + Returns: + str: Type identifier ("sqlite", etc.) + """ + if isinstance(store, SqliteMemoryStore): + return "sqlite" + else: + raise ValueError(f"Unknown memory store type: {type(store)}") + + +def create_memory_store(store_type: str) -> BaseMemoryStore: + """Create a memory store instance based on type. + + Args: + store_type: Type of memory store ("sqlite", etc.) + + Returns: + BaseMemoryStore: Initialized memory store instance + """ + config = TestConfig() + + # Initialize embedding model + embedding_model = OpenAIEmbeddingModel( + model_name=config.EMBEDDING_MODEL_NAME, + dimensions=config.EMBEDDING_DIMENSIONS, + ) + + if store_type == "sqlite": + return SqliteMemoryStore( + db_path=config.SQLITE_DB_PATH, + embedding_model=embedding_model, + vec_ext_path=config.SQLITE_VEC_EXT_PATH, + fts_enabled=config.SQLITE_FTS_ENABLED, + snippet_max_chars=config.SQLITE_SNIPPET_MAX_CHARS, + ) + else: + raise ValueError(f"Unknown store type: {store_type}") + + +# ==================== Test Functions ==================== + + +async def test_start_store(store: BaseMemoryStore, _store_name: str): + """Test store initialization.""" + logger.info("=" * 20 + " START STORE TEST " + "=" * 20) + + await store.start() + logger.info("✓ Store initialized successfully") + + # Verify tables created (SQLite specific) + if isinstance(store, SqliteMemoryStore): + cursor = store.conn.cursor() + cursor.execute( + "SELECT name FROM sqlite_master WHERE type='table' ORDER BY name", + ) + tables = [row[0] for row in cursor.fetchall()] + cursor.close() + + logger.info(f"Created tables: {tables}") + assert "files" in tables, "files table should exist" + assert "chunks" in tables, "chunks table should exist" + logger.info("✓ Required tables created") + + +async def test_upsert_file(store: BaseMemoryStore, _store_name: str) -> tuple[FileMetadata, List[MemoryChunk]]: + """Test file and chunks insertion.""" + logger.info("=" * 20 + " UPSERT FILE TEST " + "=" * 20) + + # Create sample data + file_path = "test_memory_file1.txt" + file_meta = SampleDataGenerator.create_file_metadata(file_path) + chunks = SampleDataGenerator.create_sample_chunks(file_path, prefix="test") + + # Generate embeddings for chunks + chunks = await store.get_chunk_embeddings(chunks) + file_meta.chunk_count = len(chunks) + + # Upsert file + await store.upsert_file(file_meta, MemorySource.MEMORY, chunks) + logger.info(f"✓ Upserted file: {file_path} with {len(chunks)} chunks") + + # Verify file exists + stored_meta = await store.get_file_metadata(file_path, MemorySource.MEMORY) + assert stored_meta is not None, "File should exist" + assert stored_meta.hash == file_meta.hash, "File hash should match" + logger.info(f"✓ Verified file hash: {stored_meta.hash}") + + # Verify chunks + stored_chunks = await store.get_file_chunks(file_path, MemorySource.MEMORY) + assert len(stored_chunks) == len(chunks), f"Should have {len(chunks)} chunks" + logger.info(f"✓ Verified {len(stored_chunks)} chunks stored") + + return file_meta, chunks + + +async def test_upsert_multiple_sources(store: BaseMemoryStore, _store_name: str): + """Test upserting files from different sources.""" + logger.info("=" * 20 + " UPSERT MULTIPLE SOURCES TEST " + "=" * 20) + + # Create memory file + memory_path = "test_memory_file2.txt" + memory_meta = SampleDataGenerator.create_file_metadata(memory_path) + memory_chunks = SampleDataGenerator.create_sample_chunks(memory_path, prefix="mem") + memory_chunks = await store.get_chunk_embeddings(memory_chunks) + memory_meta.chunk_count = len(memory_chunks) + + await store.upsert_file(memory_meta, MemorySource.MEMORY, memory_chunks) + logger.info(f"✓ Upserted MEMORY file: {memory_path}") + + # Create sessions file + session_path = "test_session_file1.jsonl" + session_meta = SampleDataGenerator.create_file_metadata(session_path) + session_chunks = SampleDataGenerator.create_session_chunks(session_path, prefix="sess") + session_chunks = await store.get_chunk_embeddings(session_chunks) + session_meta.chunk_count = len(session_chunks) + + await store.upsert_file(session_meta, MemorySource.SESSIONS, session_chunks) + logger.info(f"✓ Upserted SESSIONS file: {session_path}") + + # List files by source + memory_files = await store.list_files(MemorySource.MEMORY) + session_files = await store.list_files(MemorySource.SESSIONS) + + logger.info(f"MEMORY files: {len(memory_files)}") + logger.info(f"SESSIONS files: {len(session_files)}") + + assert memory_path in memory_files, "Memory file should be listed" + assert session_path in session_files, "Session file should be listed" + logger.info("✓ Multiple sources test passed") + + +async def test_update_file(store: BaseMemoryStore, _store_name: str): + """Test updating an existing file.""" + logger.info("=" * 20 + " UPDATE FILE TEST " + "=" * 20) + + file_path = "test_memory_file1.txt" + + # Get original metadata + original_meta = await store.get_file_metadata(file_path, MemorySource.MEMORY) + assert original_meta is not None, "Original file should exist" + original_chunk_count = original_meta.chunk_count + logger.info(f"Original chunk count: {original_chunk_count}") + + # Update with new chunks + updated_meta = SampleDataGenerator.create_file_metadata(file_path) + updated_meta.hash = hashlib.md5(b"updated content").hexdigest() + updated_chunks = SampleDataGenerator.create_sample_chunks(file_path, prefix="updated") + + # Add one more chunk + updated_chunks.append( + MemoryChunk( + id=f"updated_chunk4_{hashlib.md5(file_path.encode()).hexdigest()[:8]}", + path=file_path, + source=MemorySource.MEMORY, + start_line=16, + end_line=20, + text="Natural language processing enables computers to understand human language.", + hash=hashlib.md5(b"chunk4").hexdigest(), + embedding=None, + metadata={"category": "NLP", "importance": "high"}, + ), + ) + + updated_chunks = await store.get_chunk_embeddings(updated_chunks) + updated_meta.chunk_count = len(updated_chunks) + + # Upsert (update) + await store.delete_file(file_path, MemorySource.MEMORY) + await store.upsert_file(updated_meta, MemorySource.MEMORY, updated_chunks) + logger.info(f"✓ Updated file with {len(updated_chunks)} chunks") + + # Verify update + new_meta = await store.get_file_metadata(file_path, MemorySource.MEMORY) + assert new_meta.hash == updated_meta.hash, "Hash should be updated" + assert new_meta.chunk_count == len(updated_chunks), "Chunk count should be updated" + logger.info(f"✓ Verified update: new chunk count = {new_meta.chunk_count}") + + +async def test_get_file_metadata(store: BaseMemoryStore, _store_name: str): + """Test retrieving file metadata.""" + logger.info("=" * 20 + " GET FILE METADATA TEST " + "=" * 20) + + file_path = "test_memory_file1.txt" + meta = await store.get_file_metadata(file_path, MemorySource.MEMORY) + + assert meta is not None, "Metadata should exist" + assert meta.hash is not None, "Hash should exist" + assert meta.mtime_ms > 0, "Modification time should be positive" + assert meta.size > 0, "Size should be positive" + assert meta.chunk_count is not None and meta.chunk_count > 0, "Should have chunks" + + logger.info(f"File metadata: hash={meta.hash[:8]}..., chunks={meta.chunk_count}, size={meta.size}") + logger.info("✓ Get file metadata test passed") + + +async def test_list_files(store: BaseMemoryStore, _store_name: str): + """Test listing files by source.""" + logger.info("=" * 20 + " LIST FILES TEST " + "=" * 20) + + memory_files = await store.list_files(MemorySource.MEMORY) + session_files = await store.list_files(MemorySource.SESSIONS) + + logger.info(f"MEMORY files ({len(memory_files)}):") + for f in memory_files: + logger.info(f" - {f}") + + logger.info(f"SESSIONS files ({len(session_files)}):") + for f in session_files: + logger.info(f" - {f}") + + assert len(memory_files) > 0, "Should have at least one memory file" + logger.info("✓ List files test passed") + + +async def test_get_file_chunks(store: BaseMemoryStore, _store_name: str): + """Test retrieving chunks for a file.""" + logger.info("=" * 20 + " GET FILE CHUNKS TEST " + "=" * 20) + + file_path = "test_memory_file1.txt" + chunks = await store.get_file_chunks(file_path, MemorySource.MEMORY) + + assert len(chunks) > 0, "Should have chunks" + logger.info(f"Retrieved {len(chunks)} chunks") + + for i, chunk in enumerate(chunks, 1): + logger.info(f" Chunk {i}: lines {chunk.start_line}-{chunk.end_line}, text={chunk.text[:50]}...") + assert chunk.id is not None, "Chunk should have ID" + assert chunk.path == file_path, "Chunk path should match" + assert chunk.source == MemorySource.MEMORY, "Chunk source should match" + assert chunk.embedding is not None, "Chunk should have embedding" + + logger.info("✓ Get file chunks test passed") + + +async def test_vector_search(store: BaseMemoryStore, _store_name: str): + """Test vector similarity search.""" + logger.info("=" * 20 + " VECTOR SEARCH TEST " + "=" * 20) + + # Check if vector search is available (SQLite-specific) + if isinstance(store, SqliteMemoryStore) and not store.vector_available: + logger.warning("⚠ Vector extension not available, skipping vector search tests") + logger.info(" Install sqlite-vec extension to enable vector search") + logger.info(" See: https://github.com/asg017/sqlite-vec") + return + + # Search for AI-related content + query = "What is artificial intelligence and machine learning?" + results = await store.vector_search(query, limit=5) + + logger.info(f"Vector search for: '{query}'") + logger.info(f"Found {len(results)} results") + + for i, result in enumerate(results, 1): + logger.info(f"\n Result {i}:") + logger.info(f" Source: {result.source.value}") + logger.info(f" Path: {result.path}") + logger.info(f" Lines: {result.start_line}-{result.end_line}") + logger.info(f" Score: {result.score:.6f}") + logger.info(f" Snippet: {result.snippet}") + if result.metadata: + logger.info(f" Metadata: {result.metadata}") + + assert len(results) > 0, "Should find results" + assert results[0].score > 0, "Should have positive scores" + logger.info("\n✓ Vector search test passed") + + +async def test_vector_search_with_source_filter(store: BaseMemoryStore, _store_name: str): + """Test vector search with source filtering.""" + logger.info("=" * 20 + " VECTOR SEARCH WITH SOURCE FILTER TEST " + "=" * 20) + + # Check if vector search is available (SQLite-specific) + if isinstance(store, SqliteMemoryStore) and not store.vector_available: + logger.warning("⚠ Vector extension not available, skipping vector search with source filter test") + return + + query = "sales data analysis" + + # Search in MEMORY source + memory_results = await store.vector_search( + query, + limit=5, + sources=[MemorySource.MEMORY], + ) + logger.info(f"\nMEMORY source results: {len(memory_results)}") + for i, result in enumerate(memory_results, 1): + logger.info(f" {i}. Score: {result.score:.6f} | {result.path}:{result.start_line}-{result.end_line}") + logger.info(f" Snippet: {result.snippet}") + if result.metadata: + logger.info(f" Metadata: {result.metadata}") + + # Search in SESSIONS source + session_results = await store.vector_search( + query, + limit=5, + sources=[MemorySource.SESSIONS], + ) + logger.info(f"\nSESSIONS source results: {len(session_results)}") + for i, result in enumerate(session_results, 1): + logger.info(f" {i}. Score: {result.score:.6f} | {result.path}:{result.start_line}-{result.end_line}") + logger.info(f" Snippet: {result.snippet}") + if result.metadata: + logger.info(f" Metadata: {result.metadata}") + + # Search in all sources + all_results = await store.vector_search(query, limit=10) + logger.info(f"\nAll sources results: {len(all_results)}") + for i, result in enumerate(all_results, 1): + logger.info( + f" {i}. [{result.source.value}] Score: {result.score:.6f} | " + f"{result.path}:{result.start_line}-{result.end_line}", + ) + + # Verify source filtering + for r in memory_results: + assert r.source == MemorySource.MEMORY, "Memory results should only be from MEMORY source" + + for r in session_results: + assert r.source == MemorySource.SESSIONS, "Session results should only be from SESSIONS source" + + logger.info("\n✓ Vector search with source filter test passed") + + +async def test_keyword_search(store: BaseMemoryStore, _store_name: str): + """Test full-text keyword search.""" + logger.info("=" * 20 + " KEYWORD SEARCH TEST " + "=" * 20) + + # Check if FTS is available + if isinstance(store, SqliteMemoryStore) and not store.fts_available: + logger.info("⊘ Skipped: FTS not available") + return + + query = "neural networks" + results = await store.keyword_search(query, limit=5) + + logger.info(f"Keyword search for: '{query}'") + logger.info(f"Found {len(results)} results") + + for i, result in enumerate(results, 1): + logger.info(f"\n Result {i}:") + logger.info(f" Source: {result.source.value}") + logger.info(f" Path: {result.path}") + logger.info(f" Lines: {result.start_line}-{result.end_line}") + logger.info(f" Score: {result.score:.6f}") + logger.info(f" Snippet: {result.snippet}") + if result.metadata: + logger.info(f" Metadata: {result.metadata}") + + if len(results) > 0: + assert results[0].score > 0, "Should have positive scores" + logger.info("\n✓ Keyword search test passed") + else: + logger.info("\n⊘ No results found (may be expected depending on data)") + + +async def test_keyword_search_with_source_filter(store: BaseMemoryStore, _store_name: str): + """Test keyword search with source filtering.""" + logger.info("=" * 20 + " KEYWORD SEARCH WITH SOURCE FILTER TEST " + "=" * 20) + + # Check if FTS is available + if isinstance(store, SqliteMemoryStore) and not store.fts_available: + logger.info("⊘ Skipped: FTS not available") + return + + query = "data" + + # Search in different sources + memory_results = await store.keyword_search( + query, + limit=5, + sources=[MemorySource.MEMORY], + ) + logger.info(f"\nMEMORY source results: {len(memory_results)}") + for i, result in enumerate(memory_results, 1): + logger.info(f" {i}. Score: {result.score:.6f} | {result.path}:{result.start_line}-{result.end_line}") + logger.info(f" Snippet: {result.snippet}") + if result.metadata: + logger.info(f" Metadata: {result.metadata}") + + session_results = await store.keyword_search( + query, + limit=5, + sources=[MemorySource.SESSIONS], + ) + logger.info(f"\nSESSIONS source results: {len(session_results)}") + for i, result in enumerate(session_results, 1): + logger.info(f" {i}. Score: {result.score:.6f} | {result.path}:{result.start_line}-{result.end_line}") + logger.info(f" Snippet: {result.snippet}") + if result.metadata: + logger.info(f" Metadata: {result.metadata}") + + # Verify source filtering + for r in memory_results: + assert r.source == MemorySource.MEMORY, "Memory results should only be from MEMORY source" + + for r in session_results: + assert r.source == MemorySource.SESSIONS, "Session results should only be from SESSIONS source" + + logger.info("\n✓ Keyword search with source filter test passed") + + +async def test_delete_file(store: BaseMemoryStore, _store_name: str): + """Test file deletion.""" + logger.info("=" * 20 + " DELETE FILE TEST " + "=" * 20) + + # Create a file to delete + delete_path = "test_delete_file.txt" + delete_meta = SampleDataGenerator.create_file_metadata(delete_path) + delete_chunks = SampleDataGenerator.create_sample_chunks(delete_path, prefix="del") + delete_chunks = await store.get_chunk_embeddings(delete_chunks) + delete_meta.chunk_count = len(delete_chunks) + + await store.upsert_file(delete_meta, MemorySource.MEMORY, delete_chunks) + logger.info(f"✓ Created file: {delete_path}") + + # Verify it exists + meta_before = await store.get_file_metadata(delete_path, MemorySource.MEMORY) + assert meta_before is not None, "File should exist before deletion" + + # Delete the file + await store.delete_file(delete_path, MemorySource.MEMORY) + logger.info(f"✓ Deleted file: {delete_path}") + + # Verify deletion + meta_after = await store.get_file_metadata(delete_path, MemorySource.MEMORY) + assert meta_after is None, "File should not exist after deletion" + + chunks_after = await store.get_file_chunks(delete_path, MemorySource.MEMORY) + assert len(chunks_after) == 0, "Chunks should be deleted" + logger.info("✓ Verified deletion") + + +async def test_batch_upsert(store: BaseMemoryStore, _store_name: str): + """Test batch file upsertion.""" + logger.info("=" * 20 + " BATCH UPSERT TEST " + "=" * 20) + + # Create multiple files + batch_size = 10 + for i in range(batch_size): + file_path = f"test_batch_file_{i}.txt" + file_meta = SampleDataGenerator.create_file_metadata(file_path) + chunks = SampleDataGenerator.create_sample_chunks(file_path, prefix=f"batch{i}") + chunks = await store.get_chunk_embeddings(chunks) + file_meta.chunk_count = len(chunks) + + await store.upsert_file(file_meta, MemorySource.MEMORY, chunks) + + logger.info(f"✓ Batch upserted {batch_size} files") + + # Verify all files exist + memory_files = await store.list_files(MemorySource.MEMORY) + batch_files = [f for f in memory_files if f.startswith("test_batch_file_")] + + assert len(batch_files) >= batch_size, f"Should have at least {batch_size} batch files" + logger.info(f"✓ Verified {len(batch_files)} batch files") + + +async def test_concurrent_searches(store: BaseMemoryStore, _store_name: str): + """Test concurrent search operations.""" + logger.info("=" * 20 + " CONCURRENT SEARCHES TEST " + "=" * 20) + + queries = [ + "artificial intelligence", + "machine learning algorithms", + "deep learning neural networks", + "natural language processing", + "data analysis techniques", + ] + + # Concurrent vector searches + search_tasks = [store.vector_search(q, limit=3) for q in queries] + results = await asyncio.gather(*search_tasks) + + logger.info(f"✓ Completed {len(results)} concurrent vector searches") + for i, (query, result) in enumerate(zip(queries, results), 1): + logger.info(f" Query {i}: '{query}' -> {len(result)} results") + + # Concurrent keyword searches (if available) + if isinstance(store, SqliteMemoryStore) and store.fts_available: + keyword_tasks = [store.keyword_search(q, limit=3) for q in queries] + keyword_results = await asyncio.gather(*keyword_tasks) + logger.info(f"✓ Completed {len(keyword_results)} concurrent keyword searches") + + logger.info("✓ Concurrent searches test passed") + + +async def test_edge_cases(store: BaseMemoryStore, _store_name: str): + """Test edge cases and boundary conditions.""" + logger.info("=" * 20 + " EDGE CASES TEST " + "=" * 20) + + # Test 1: Empty chunk text + edge_path1 = "test_edge_empty_chunk.txt" + edge_meta1 = SampleDataGenerator.create_file_metadata(edge_path1) + edge_chunks1 = [ + MemoryChunk( + id="edge_empty_chunk", + path=edge_path1, + source=MemorySource.MEMORY, + start_line=1, + end_line=1, + text="", + hash=hashlib.md5(b"").hexdigest(), + embedding=None, + ), + ] + + try: + edge_chunks1 = await store.get_chunk_embeddings(edge_chunks1) + await store.upsert_file(edge_meta1, MemorySource.MEMORY, edge_chunks1) + logger.info("✓ Handled empty chunk text") + except Exception as e: + logger.info(f"⊘ Empty chunk not supported: {e}") + + # Test 2: Very long chunk text + edge_path2 = "test_edge_long_chunk.txt" + edge_meta2 = SampleDataGenerator.create_file_metadata(edge_path2) + long_text = "A" * 10000 # 10k characters + edge_chunks2 = [ + MemoryChunk( + id="edge_long_chunk", + path=edge_path2, + source=MemorySource.MEMORY, + start_line=1, + end_line=100, + text=long_text, + hash=hashlib.md5(long_text.encode()).hexdigest(), + embedding=None, + ), + ] + + edge_chunks2 = await store.get_chunk_embeddings(edge_chunks2) + await store.upsert_file(edge_meta2, MemorySource.MEMORY, edge_chunks2) + retrieved = await store.get_file_chunks(edge_path2, MemorySource.MEMORY) + assert len(retrieved[0].text) == 10000, "Long text should be preserved" + logger.info("✓ Handled very long chunk text (10k chars)") + + # Test 3: Special characters in text + edge_path3 = "test_edge_special_chars.txt" + edge_meta3 = SampleDataGenerator.create_file_metadata(edge_path3) + special_text = "Special chars: @#$%^&*()[]{}|\\;:'\",.<>?/~`+=−×÷" + edge_chunks3 = [ + MemoryChunk( + id="edge_special_chars", + path=edge_path3, + source=MemorySource.MEMORY, + start_line=1, + end_line=1, + text=special_text, + hash=hashlib.md5(special_text.encode()).hexdigest(), + embedding=None, + ), + ] + + edge_chunks3 = await store.get_chunk_embeddings(edge_chunks3) + await store.upsert_file(edge_meta3, MemorySource.MEMORY, edge_chunks3) + retrieved = await store.get_file_chunks(edge_path3, MemorySource.MEMORY) + assert "@#$%^&*()" in retrieved[0].text, "Special chars should be preserved" + logger.info("✓ Handled special characters in text") + + # Test 4: Unicode and emoji + edge_path4 = "test_edge_unicode.txt" + edge_meta4 = SampleDataGenerator.create_file_metadata(edge_path4) + unicode_text = "Unicode test: 你好世界 🌍 مرحبا العالم Привет мир" + edge_chunks4 = [ + MemoryChunk( + id="edge_unicode", + path=edge_path4, + source=MemorySource.MEMORY, + start_line=1, + end_line=1, + text=unicode_text, + hash=hashlib.md5(unicode_text.encode()).hexdigest(), + embedding=None, + ), + ] + + edge_chunks4 = await store.get_chunk_embeddings(edge_chunks4) + await store.upsert_file(edge_meta4, MemorySource.MEMORY, edge_chunks4) + retrieved = await store.get_file_chunks(edge_path4, MemorySource.MEMORY) + assert "你好世界" in retrieved[0].text, "Unicode should be preserved" + assert "🌍" in retrieved[0].text, "Emoji should be preserved" + logger.info("✓ Handled unicode and emoji") + + # Test 5: Search with empty query + try: + results = await store.vector_search("", limit=5) + logger.info(f"✓ Empty query returned {len(results)} results") + except Exception as e: + logger.info(f"⊘ Empty query not supported: {e}") + + # Test 6: Very high limit + results = await store.vector_search("test", limit=1000) + logger.info(f"✓ High limit search returned {len(results)} results") + + # Test 7: Non-existent file + non_existent_meta = await store.get_file_metadata("non_existent_file.txt", MemorySource.MEMORY) + assert non_existent_meta is None, "Non-existent file should return None" + logger.info("✓ Non-existent file handled gracefully") + + logger.info("✓ Edge cases test passed") + + +async def test_clear_all(store: BaseMemoryStore, _store_name: str): + """Test clearing all data.""" + logger.info("=" * 20 + " CLEAR ALL TEST " + "=" * 20) + + # Verify we have data before clearing + files_before = await store.list_files(MemorySource.MEMORY) + logger.info(f"Files before clear: {len(files_before)}") + assert len(files_before) > 0, "Should have files before clearing" + + # Clear all data + await store.clear_all() + logger.info("✓ Cleared all data") + + # Verify all data is gone + memory_files = await store.list_files(MemorySource.MEMORY) + session_files = await store.list_files(MemorySource.SESSIONS) + + assert len(memory_files) == 0, "All memory files should be deleted" + assert len(session_files) == 0, "All session files should be deleted" + logger.info("✓ Verified all data cleared") + + # Verify we can still insert after clearing + test_path = "test_after_clear.txt" + test_meta = SampleDataGenerator.create_file_metadata(test_path) + test_chunks = SampleDataGenerator.create_sample_chunks(test_path) + test_chunks = await store.get_chunk_embeddings(test_chunks) + test_meta.chunk_count = len(test_chunks) + + await store.upsert_file(test_meta, MemorySource.MEMORY, test_chunks) + logger.info("✓ Can insert data after clearing") + + logger.info("✓ Clear all test passed") + + +# ==================== Test Runner ==================== + + +async def run_all_tests_for_store(store_type: str, store_name: str): + """Run all tests for a specific memory store type. + + Args: + store_type: Type of memory store ("sqlite", etc.) + store_name: Display name for the memory store + """ + logger.info(f"\n\n{'#' * 60}") + logger.info(f"# Running all tests for: {store_name}") + logger.info(f"{'#' * 60}") + + # Create memory store instance + store = create_memory_store(store_type) + + try: + # ========== Basic Tests ========== + logger.info(f"\n{'#' * 60}") + logger.info("# BASIC FUNCTIONALITY TESTS") + logger.info(f"{'#' * 60}") + + await test_start_store(store, store_name) + await test_upsert_file(store, store_name) + await test_upsert_multiple_sources(store, store_name) + await test_update_file(store, store_name) + await test_get_file_metadata(store, store_name) + await test_list_files(store, store_name) + await test_get_file_chunks(store, store_name) + + # ========== Search Tests ========== + logger.info(f"\n{'#' * 60}") + logger.info("# SEARCH FUNCTIONALITY TESTS") + logger.info(f"{'#' * 60}") + + await test_vector_search(store, store_name) + await test_vector_search_with_source_filter(store, store_name) + await test_keyword_search(store, store_name) + await test_keyword_search_with_source_filter(store, store_name) + + # ========== Advanced Tests ========== + logger.info(f"\n{'#' * 60}") + logger.info("# ADVANCED FUNCTIONALITY TESTS") + logger.info(f"{'#' * 60}") + + await test_delete_file(store, store_name) + await test_batch_upsert(store, store_name) + await test_concurrent_searches(store, store_name) + await test_edge_cases(store, store_name) + + # ========== Cleanup Test ========== + logger.info(f"\n{'#' * 60}") + logger.info("# CLEANUP TESTS") + logger.info(f"{'#' * 60}") + + await test_clear_all(store, store_name) + + logger.info(f"\n{'=' * 60}") + logger.info(f"✓ All tests passed for {store_name}!") + logger.info(f"{'=' * 60}") + + except Exception as e: + logger.error(f"Test failed: {e}") + raise + finally: + # Cleanup + await cleanup_store(store, store_type) + + +async def cleanup_store(store: BaseMemoryStore, store_type: str): + """Clean up test resources for a memory store. + + Args: + store: Memory store instance + store_type: Type of memory store ("sqlite", etc.) + """ + logger.info("=" * 20 + " CLEANUP " + "=" * 20) + + try: + # Close connections + await store.close() + logger.info("✓ Closed store connections") + + # Clean up local directory if SqliteMemoryStore + if store_type == "sqlite": + config = TestConfig() + db_dir = Path(config.SQLITE_DB_PATH).parent + if db_dir.exists(): + shutil.rmtree(db_dir) + logger.info(f"✓ Cleaned up directory: {db_dir}") + + logger.info("✓ Cleanup completed") + except Exception as e: + logger.error(f"Cleanup error: {e}") + + +# ==================== Main Entry Point ==================== + + +async def main(): + """Main entry point for running tests.""" + parser = argparse.ArgumentParser( + description="Run memory store tests", + formatter_class=argparse.RawDescriptionHelpFormatter, + epilog=""" +Examples: + python test_memory_store.py --sqlite # Test SqliteMemoryStore only + python test_memory_store.py --all # Test all memory stores + """, + ) + parser.add_argument( + "--sqlite", + action="store_true", + help="Test SqliteMemoryStore", + ) + parser.add_argument( + "--all", + action="store_true", + help="Run tests for all available memory stores", + ) + + args = parser.parse_args() + + # Determine which memory stores to test + stores_to_test = [] + + if args.all: + stores_to_test = [ + ("sqlite", "SqliteMemoryStore"), + ] + else: + # Build list based on individual flags + if args.sqlite: + stores_to_test.append(("sqlite", "SqliteMemoryStore")) + + if not stores_to_test: + # Default to all memory stores if no argument provided + stores_to_test = [ + ("sqlite", "SqliteMemoryStore"), + ] + print("No memory store specified, defaulting to test all memory stores") + print("Use --sqlite to test specific ones\n") + + # Run tests for each memory store + for store_type, store_name in stores_to_test: + try: + await run_all_tests_for_store(store_type, store_name) + except Exception as e: + logger.error(f"\n✗ FAILED: {store_name} tests failed with error:") + logger.error(f" {type(e).__name__}: {e}") + raise + + # Final summary + print(f"\n\n{'#' * 60}") + print("# TEST SUMMARY") + print(f"{'#' * 60}") + print(f"✓ All tests passed for {len(stores_to_test)} memory store(s):") + for _, store_name in stores_to_test: + print(f" - {store_name}") + print(f"{'#' * 60}\n") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/test_tool.py b/tests/test_tool.py index 9e9ccaff..6c23ad18 100644 --- a/tests/test_tool.py +++ b/tests/test_tool.py @@ -10,10 +10,8 @@ import asyncio from reme import ReMe -app = ReMe() - -def test_search(): +async def test_search(_app): """Test search tool operations. Tests DashscopeSearch, MockSearch, and TavilySearch operations @@ -32,11 +30,11 @@ def test_search(): print(f"Testing {op.__class__.__name__}") print("=" * 60) print(f"Query: {query}") - output = asyncio.run(op.call(query=query, service_context=app.service_context)) + output = await op.call(query=query, service_context=_app.service_context) print(f"Output:\n{output}") -def test_execute(): +async def test_execute(_app): """Test code and shell execution tool operations. Tests ExecuteCode and ExecuteShell operations with various scenarios @@ -53,7 +51,7 @@ def test_execute(): op = ExecuteCode() code_to_execute = "print('hello world')" print(f"Executing Python code: {code_to_execute}") - output = asyncio.run(op.call(code=code_to_execute)) + output = await op.call(code=code_to_execute) print(f"Output:\n{output}") # Test ExecuteCode with more complex code @@ -64,7 +62,7 @@ def test_execute(): op = ExecuteCode() code_to_execute = "result = sum(range(1, 11))\nprint(f'Sum of 1-10: {result}')" print(f"Executing Python code:\n{code_to_execute}") - output = asyncio.run(op.call(code=code_to_execute)) + output = await op.call(code=code_to_execute) print(f"Output:\n{output}") # Test ExecuteShell @@ -75,7 +73,7 @@ def test_execute(): op = ExecuteShell() command = "ls" print(f"Executing shell command: {command}") - output = asyncio.run(op.call(command=command)) + output = await op.call(command=command) print(f"Output:\n{output}") # Test ExecuteShell with echo @@ -86,7 +84,7 @@ def test_execute(): op = ExecuteShell() command = "echo 'Hello from shell!'" print(f"Executing shell command: {command}") - output = asyncio.run(op.call(command=command)) + output = await op.call(command=command) print(f"Output:\n{output}") # Test ExecuteCode with error (syntax error) @@ -97,7 +95,7 @@ def test_execute(): op = ExecuteCode() code_to_execute = "print('missing closing quote)" print(f"Executing Python code with syntax error:\n{code_to_execute}") - output = asyncio.run(op.call(code=code_to_execute)) + output = await op.call(code=code_to_execute) print(f"Output:\n{output}") # Test ExecuteCode with runtime error @@ -108,7 +106,7 @@ def test_execute(): op = ExecuteCode() code_to_execute = "x = 1 / 0" print(f"Executing Python code with runtime error:\n{code_to_execute}") - output = asyncio.run(op.call(code=code_to_execute)) + output = await op.call(code=code_to_execute) print(f"Output:\n{output}") # Test ExecuteCode with undefined variable @@ -119,7 +117,7 @@ def test_execute(): op = ExecuteCode() code_to_execute = "print(undefined_variable)" print(f"Executing Python code with undefined variable:\n{code_to_execute}") - output = asyncio.run(op.call(code=code_to_execute)) + output = await op.call(code=code_to_execute) print(f"Output:\n{output}") # Test ExecuteShell with invalid command @@ -130,7 +128,7 @@ def test_execute(): op = ExecuteShell() command = "this_command_does_not_exist" print(f"Executing invalid shell command: {command}") - output = asyncio.run(op.call(command=command)) + output = await op.call(command=command) print(f"Output:\n{output}") # Test ExecuteShell with command that returns non-zero exit code @@ -141,7 +139,7 @@ def test_execute(): op = ExecuteShell() command = "ls /nonexistent_directory_12345" print(f"Executing shell command that should fail: {command}") - output = asyncio.run(op.call(command=command)) + output = await op.call(command=command) print(f"Output:\n{output}") print("\n" + "=" * 60) @@ -149,7 +147,7 @@ def test_execute(): print("=" * 60) -def test_simple_chat(): +async def test_simple_chat(app): """Test simple chat operation. Tests the SimpleChat agent with a basic query to verify @@ -158,11 +156,11 @@ def test_simple_chat(): from reme.agent.chat import SimpleChat op = SimpleChat() - output = asyncio.run(op.call(query="你好", service_context=app.service_context)) + output = await op.call(query="你好", service_context=app.service_context) print(output) -async def test_stream_chat(): +async def test_stream_chat(app): """Test streaming chat operation. Tests the StreamChat agent with a query to verify it can @@ -189,8 +187,16 @@ async def test_stream_chat(): print(chunk, end="") +async def main(): + """Main entry point for running tool tests.""" + app = ReMe() + await app.start() + await test_search(app) + await test_execute(app) + await test_simple_chat(app) + await test_stream_chat(app) + await app.close() + + if __name__ == "__main__": - test_search() - # test_execute() - # test_simple_chat() - # asyncio.run(test_stream_chat()) + asyncio.run(main())