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())