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/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/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 e43da12b..4bd4cbe3 100644
--- a/reme/core/embedding/base_embedding_model.py
+++ b/reme/core/embedding/base_embedding_model.py
@@ -4,8 +4,10 @@ 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
@@ -28,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:
@@ -52,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."""
@@ -61,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:
@@ -79,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}")
@@ -97,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:
@@ -120,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}")
@@ -137,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."""
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_manager/__init__.py b/reme/core/memory_manager/__init__.py
deleted file mode 100644
index e69de29b..00000000
diff --git a/reme/core/memory_manager/ingestion/__init__.py b/reme/core/memory_manager/ingestion/__init__.py
deleted file mode 100644
index e69de29b..00000000
diff --git a/reme/core/memory_manager/manager.py b/reme/core/memory_manager/manager.py
deleted file mode 100644
index f8eccada..00000000
--- a/reme/core/memory_manager/manager.py
+++ /dev/null
@@ -1,890 +0,0 @@
-"""Memory Index Manager - Main coordination layer.
-
-This module provides the main MemoryIndexManager class that coordinates
-file watching, embedding generation, and search operations across memory files
-and session transcripts.
-"""
-
-import asyncio
-import json
-import os
-import re
-from typing import Any, Callable
-
-from loguru import logger
-from pydantic import BaseModel, Field
-from watchfiles import awatch
-
-from .ingestion.chunking import chunk_markdown
-from .memory_storage.sqlite_memory_store import SqliteMemoryStore
-from .utils.hashing import hash_text
-from ..enumeration import MemorySource
-from ..schema import FileMetadata, MemorySearchResult
-
-# Constants
-SNIPPET_MAX_CHARS = 700
-SESSION_DIRTY_DEBOUNCE_MS = 5000
-EMBEDDING_BATCH_MAX_TOKENS = 8000
-EMBEDDING_APPROX_CHARS_PER_TOKEN = 1
-EMBEDDING_INDEX_CONCURRENCY = 4
-EMBEDDING_RETRY_MAX_ATTEMPTS = 3
-EMBEDDING_RETRY_BASE_DELAY_MS = 500
-EMBEDDING_RETRY_MAX_DELAY_MS = 8000
-BATCH_FAILURE_LIMIT = 2
-SESSION_DELTA_READ_CHUNK_BYTES = 64 * 1024
-EMBEDDING_QUERY_TIMEOUT_REMOTE_MS = 60_000
-EMBEDDING_QUERY_TIMEOUT_LOCAL_MS = 5 * 60_000
-EMBEDDING_BATCH_TIMEOUT_REMOTE_MS = 2 * 60_000
-EMBEDDING_BATCH_TIMEOUT_LOCAL_MS = 10 * 60_000
-
-
-class MemorySyncProgressUpdate(BaseModel):
- """Progress update for memory sync operations."""
-
- completed: int = Field(default=..., description="Number of items completed")
- total: int = Field(default=..., description="Total number of items to process")
- label: str | None = Field(default=None, description="Optional label for the progress operation")
-
-
-class MemorySyncProgressState(BaseModel):
- """Internal state for tracking sync progress."""
-
- completed: int = Field(default=0, description="Number of items completed")
- total: int = Field(default=0, description="Total number of items to process")
- label: str | None = Field(default=None, description="Optional label for the progress operation")
- report: Callable[[MemorySyncProgressUpdate], None] | None = Field(
- default=None,
- description="Callback function to report progress updates",
- )
-
-
-class SessionDelta(BaseModel):
- """Tracks incremental changes in session files."""
-
- last_size: int = Field(default=0, description="Last known size of the session file")
- pending_bytes: int = Field(default=0, description="Number of pending bytes to process")
- pending_messages: int = Field(default=0, description="Number of pending messages to process")
-
-
-class MemorySearchConfig(BaseModel):
- """Configuration for memory search operations."""
-
- model: str = Field(default="default", description="Model name for embeddings")
- sources: list[MemorySource] = Field(default_factory=lambda: [MemorySource.MEMORY], description="Sources to search")
- extra_paths: list[str] = Field(default_factory=list, description="Additional paths to include in search")
- store_path: str = Field(default="memory.db", description="Path to SQLite database file")
- vector_enabled: bool = Field(default=True, description="Whether to enable vector search")
- vector_extension_path: str | None = Field(default=None, description="Path to vector extension for SQLite")
- fts_enabled: bool = Field(default=True, description="Whether to enable full-text search")
- chunk_tokens: int = Field(default=300, description="Number of tokens per chunk")
- chunk_overlap: int = Field(default=30, description="Number of overlapping tokens between chunks")
- watch_enabled: bool = Field(default=True, description="Whether to enable file watching")
- watch_debounce_ms: int = Field(default=1000, description="Debounce time for file watcher in milliseconds")
- interval_minutes: int = Field(default=0, description="Interval between automatic syncs in minutes (0 to disable)")
- sync_on_search: bool = Field(default=True, description="Whether to sync before search operations")
- sync_on_session_start: bool = Field(default=True, description="Whether to sync when a session starts")
- query_min_score: float = Field(default=0.3, description="Minimum relevance score for search results")
- query_max_results: int = Field(default=10, description="Maximum number of search results to return")
- hybrid_enabled: bool = Field(default=True, description="Whether to use hybrid vector + keyword search")
- hybrid_vector_weight: float = Field(default=0.7, description="Weight for vector search in hybrid scoring")
- hybrid_text_weight: float = Field(default=0.3, description="Weight for text search in hybrid scoring")
- hybrid_candidate_multiplier: float = Field(
- default=2.0,
- description="Multiplier for number of candidates to consider in hybrid search",
- )
- session_delta_bytes: int = Field(
- default=0,
- description="Threshold for session sync based on bytes changed (0 for any change)",
- )
- session_delta_messages: int = Field(
- default=5,
- description="Threshold for session sync based on messages changed (0 for any change)",
- )
-
-
-# Global cache for manager instances
-INDEX_CACHE: dict[str, "MemoryIndexManager"] = {}
-
-
-class MemoryIndexManager:
- """Main memory index manager coordinating all memory operations."""
-
- # ============================================================================
- # Initialization and Lifecycle
- # ============================================================================
-
- def __init__(
- self,
- agent_id: str,
- workspace_dir: str,
- settings: MemorySearchConfig,
- store: SqliteMemoryStore,
- ):
- """Initialize the memory index manager."""
-
- self.agent_id = agent_id
- self.workspace_dir = workspace_dir
- self.settings = settings
- self.store = store
-
- # State tracking
- self.sources = set(settings.sources)
- self.closed = False
- self.dirty = MemorySource.MEMORY in self.sources
- self.sessions_dirty = False
- self.sessions_dirty_files: set[str] = set()
- self.session_pending_files: set[str] = set()
- self.session_deltas: dict[str, SessionDelta] = {}
- self.session_warm: set[str] = set()
-
- # Sync control
- self.syncing: asyncio.Task | None = None
- self.watch_task: asyncio.Task | None = None
- self.session_watch_task: asyncio.Task | None = None
- self.interval_task: asyncio.Task | None = None
-
- # Batch failure tracking
- self.batch_failure_count = 0
- self.batch_failure_last_error: str | None = None
- self.batch_failure_lock = asyncio.Lock()
-
- async def close(self) -> None:
- """Close the manager and release resources."""
- if self.closed:
- return
-
- self.closed = True
-
- # Cancel all background tasks
- if self.watch_task:
- self.watch_task.cancel()
- if self.session_watch_task:
- self.session_watch_task.cancel()
- if self.interval_task:
- self.interval_task.cancel()
-
- await self.store.close()
-
- # ============================================================================
- # Public API Methods
- # ============================================================================
-
- async def warm_session(self, session_key: str | None = None):
- """Pre-sync memory before a session starts."""
- if not self.settings.sync_on_session_start:
- return
-
- key = (session_key or "").strip()
- if key and key in self.session_warm:
- return
-
- await self.sync(reason="session-start")
-
- if key:
- self.session_warm.add(key)
-
- async def sync(
- self,
- reason: str | None = None,
- force: bool = False,
- progress: Callable[[MemorySyncProgressUpdate], None] | None = None,
- ):
- """Synchronize memory index with file system."""
- if self.syncing:
- await self.syncing
- return
-
- self.syncing = asyncio.create_task(self._run_sync(reason, force, progress))
- try:
- await self.syncing
- finally:
- self.syncing = None
-
- async def search(
- self,
- query: str,
- max_results: int | None = None,
- min_score: float | None = None,
- session_key: str | None = None,
- ) -> list[MemorySearchResult]:
- """Search indexed memory with hybrid vector + keyword search.
-
- Args:
- query: Search query text
- max_results: Maximum number of results to return
- min_score: Minimum relevance score threshold
- session_key: Optional session key for warmup
-
- Returns:
- List of search results sorted by relevance
- """
- await self.warm_session(session_key)
-
- if self.settings.sync_on_search and (self.dirty or self.sessions_dirty):
- try:
- await self.sync(reason="search")
- except Exception as err:
- logger.warning(f"memory sync failed (search): {err}")
-
- cleaned = query.strip()
- if not cleaned:
- return []
-
- min_score = min_score if min_score is not None else self.settings.query_min_score
- max_results = max_results if max_results is not None else self.settings.query_max_results
-
- hybrid = self.settings.hybrid_enabled
- candidates = min(200, max(1, int(max_results * self.settings.hybrid_candidate_multiplier)))
-
- # Run keyword search if hybrid enabled
- keyword_results = []
- if hybrid:
- keyword_results = await self._search_keyword(cleaned, candidates)
-
- # Perform vector search
- vector_results = await self._search_vector(cleaned, candidates)
-
- if not hybrid:
- return [r for r in vector_results if r.score >= min_score][:max_results]
-
- merged = self._merge_hybrid_results(
- vector=vector_results,
- keyword=keyword_results,
- vector_weight=self.settings.hybrid_vector_weight,
- text_weight=self.settings.hybrid_text_weight,
- )
-
- return [r for r in merged if r.score >= min_score][:max_results]
-
- async def read_file(
- self,
- rel_path: str,
- from_line: int | None = None,
- num_lines: int | None = None,
- ) -> dict[str, str]:
- """Read a memory file with optional line range.
-
- Args:
- rel_path: Relative path to file
- from_line: Starting line number (1-indexed)
- num_lines: Number of lines to read
-
- Returns:
- Dictionary with 'text' and 'path' keys
-
- Raises:
- ValueError: If path is invalid or not allowed
- """
- raw_path = rel_path.strip()
- assert raw_path, "path required"
- abs_path = os.path.abspath(os.path.join(self.workspace_dir, raw_path))
- rel_path_clean = os.path.relpath(abs_path, self.workspace_dir)
-
- in_workspace = not rel_path_clean.startswith("..") and not os.path.isabs(rel_path_clean)
- allowed = in_workspace and self._is_memory_path(rel_path_clean)
-
- if not allowed and self.settings.extra_paths:
- for extra in self.settings.extra_paths:
- extra_abs = os.path.abspath(extra)
- if abs_path.startswith(extra_abs):
- allowed = True
- break
-
- if not allowed:
- raise ValueError("path required")
-
- if not abs_path.endswith(".md"):
- raise ValueError("path required")
-
- # Read file
- with open(abs_path, "r", encoding="utf-8") as f:
- content = f.read()
-
- if from_line is None and num_lines is None:
- return {"text": content, "path": rel_path_clean}
-
- lines = content.split("\n")
- start = max(1, from_line or 1)
- count = max(1, num_lines or len(lines))
- slice_lines = lines[start - 1 : start - 1 + count]
-
- return {"text": "\n".join(slice_lines), "path": rel_path_clean}
-
- # ============================================================================
- # Sync Logic
- # ============================================================================
-
- async def _run_sync(
- self,
- reason: str | None,
- force: bool,
- progress_callback: Callable[[MemorySyncProgressUpdate], None] | None,
- ):
- """Execute sync operation."""
- progress = MemorySyncProgressState()
- if progress_callback:
- progress.report = progress_callback
-
- should_sync_memory = MemorySource.MEMORY in self.sources and (force or self.dirty)
- should_sync_sessions = self._should_sync_sessions(reason, force)
-
- if should_sync_memory:
- await self._sync_memory_files(progress)
- self.dirty = False
-
- if should_sync_sessions:
- await self._sync_session_files(progress)
- self.sessions_dirty = False
- self.sessions_dirty_files.clear()
- elif len(self.sessions_dirty_files) > 0:
- self.sessions_dirty = True
- else:
- self.sessions_dirty = False
-
- def _should_sync_sessions(self, reason: str | None, force: bool) -> bool:
- """Check if session sync is needed."""
- if MemorySource.SESSIONS not in self.sources:
- return False
-
- if force:
- return True
-
- if reason in ("session-start", "watch"):
- return False
-
- return self.sessions_dirty and len(self.sessions_dirty_files) > 0
-
- async def _sync_memory_files(self, progress: MemorySyncProgressState):
- """Sync memory markdown files."""
- files = self._list_memory_files()
- logger.debug("memory sync: indexing memory files", files=len(files))
-
- active_paths = {f.path for f in files}
- if progress.report:
- progress.total += len(files)
- progress.report(
- MemorySyncProgressUpdate(
- completed=progress.completed,
- total=progress.total,
- label="Indexing memory files…",
- ),
- )
-
- tasks = []
- for file_entry in files:
- task = self._index_memory_file(file_entry, progress)
- tasks.append(task)
- await asyncio.gather(*tasks)
-
- indexed = await self.store.list_files(MemorySource.MEMORY)
- for stale_path in indexed:
- if stale_path not in active_paths:
- await self.store.delete_file(stale_path, MemorySource.MEMORY)
-
- async def _sync_session_files(self, progress: MemorySyncProgressState):
- """Sync session transcript files."""
- files = self._list_session_files()
- logger.debug(
- "memory sync: indexing session files",
- files=len(files),
- index_all=len(self.sessions_dirty_files) == 0,
- dirty_files=len(self.sessions_dirty_files),
- )
-
- if progress.report:
- progress.total += len(files)
- progress.report(
- MemorySyncProgressUpdate(
- completed=progress.completed,
- total=progress.total,
- label="Indexing session files...",
- ),
- )
-
- active_paths = set()
- tasks = []
-
- for abs_path in files:
- rel_path = self._session_path_for_file(abs_path)
- active_paths.add(rel_path)
-
- if len(self.sessions_dirty_files) == 0 or abs_path in self.sessions_dirty_files:
- task = self._index_session_file(abs_path, progress)
- tasks.append(task)
- else:
- if progress.report:
- progress.completed += 1
- progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
- await asyncio.gather(*tasks)
-
- indexed = await self.store.list_files(MemorySource.SESSIONS)
- for stale_path in indexed:
- if stale_path not in active_paths:
- await self.store.delete_file(stale_path, MemorySource.SESSIONS)
-
- # ============================================================================
- # File Indexing
- # ============================================================================
-
- async def _index_memory_file(self, file_meta: FileMetadata, progress: MemorySyncProgressState):
- """Index a single memory file."""
- existing_meta = await self.store.get_file_metadata(file_meta.path, MemorySource.MEMORY)
- if existing_meta and existing_meta.hash == file_meta.hash:
- if progress.report:
- progress.completed += 1
- progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
- return
-
- # Read and chunk file
- with open(file_meta.abs_path, "r", encoding="utf-8") as f:
- content = f.read()
-
- chunks = chunk_markdown(
- content,
- file_meta.path,
- MemorySource.MEMORY,
- self.settings.chunk_tokens,
- self.settings.chunk_overlap,
- )
-
- chunks = [c for c in chunks if c.text.strip()]
-
- if chunks:
- chunks = await self.store.get_chunk_embeddings(chunks)
-
- await self.store.upsert_file(file_meta, MemorySource.MEMORY, chunks)
- if progress.report:
- progress.completed += 1
- progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
-
- async def _index_session_file(self, abs_path: str, progress: MemorySyncProgressState):
- """Index a single session transcript file."""
- file_meta = self._build_session_file_meta(abs_path)
- if not file_meta:
- if progress.report:
- progress.completed += 1
- progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
- return
-
- existing_meta = await self.store.get_file_metadata(file_meta.path, MemorySource.SESSIONS)
- if existing_meta and existing_meta.hash == file_meta.hash:
- self._reset_session_delta(abs_path, file_meta.size)
- if progress.report:
- progress.completed += 1
- progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
- return
-
- chunks = chunk_markdown(
- file_meta.content,
- file_meta.path,
- MemorySource.SESSIONS,
- self.settings.chunk_tokens,
- self.settings.chunk_overlap,
- )
-
- chunks = [c for c in chunks if c.text.strip()]
-
- if chunks:
- chunks = await self.store.get_chunk_embeddings(chunks)
-
- await self.store.upsert_file(file_meta, MemorySource.SESSIONS, chunks)
- self._reset_session_delta(abs_path, file_meta.size)
-
- if progress.report:
- progress.completed += 1
- progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
-
- # ============================================================================
- # File Listing and Building
- # ============================================================================
-
- def _list_memory_files(self) -> list[FileMetadata]:
- """List all memory markdown files."""
- files = []
-
- # Scan workspace
- memory_paths = [
- os.path.join(self.workspace_dir, "MEMORY.md"),
- os.path.join(self.workspace_dir, "memory.md"),
- os.path.join(self.workspace_dir, "memory"),
- ]
-
- for base_path in memory_paths:
- if os.path.isfile(base_path) and base_path.endswith(".md"):
- files.append(self._build_file_entry(base_path))
- elif os.path.isdir(base_path):
- for root, _, filenames in os.walk(base_path):
- for filename in filenames:
- if filename.endswith(".md"):
- abs_path = os.path.join(root, filename)
- files.append(self._build_file_entry(abs_path))
-
- # Extra paths
- for extra in self.settings.extra_paths:
- if os.path.isfile(extra) and extra.endswith(".md"):
- files.append(self._build_file_entry(extra))
- elif os.path.isdir(extra):
- for root, _, filenames in os.walk(extra):
- for filename in filenames:
- if filename.endswith(".md"):
- abs_path = os.path.join(root, filename)
- files.append(self._build_file_entry(abs_path))
-
- return files
-
- def _list_session_files(self) -> list[str]:
- """List all session transcript files."""
- sessions_dir = os.path.join(self.workspace_dir, "sessions", self.agent_id)
- if not os.path.exists(sessions_dir):
- return []
-
- files = []
- for filename in os.listdir(sessions_dir):
- if filename.endswith(".jsonl"):
- files.append(os.path.join(sessions_dir, filename))
-
- return files
-
- def _build_file_entry(self, abs_path: str) -> FileMetadata:
- """Build file entry metadata."""
- stat = os.stat(abs_path)
- with open(abs_path, "r", encoding="utf-8") as f:
- content = f.read()
-
- rel_path = os.path.relpath(abs_path, self.workspace_dir)
-
- return FileMetadata(
- hash=hash_text(content),
- mtime_ms=stat.st_mtime * 1000,
- size=stat.st_size,
- path=rel_path.replace("\\", "/"),
- abs_path=abs_path,
- )
-
- def _build_session_file_meta(self, abs_path: str) -> FileMetadata | None:
- """Build session file entry with parsed content. TODO 修改message解析逻辑"""
- stat = os.stat(abs_path)
- with open(abs_path, "r", encoding="utf-8") as f:
- raw = f.read()
-
- lines = raw.split("\n")
- collected = []
-
- for line in lines:
- if not line.strip():
- continue
-
- try:
- record = json.loads(line)
- except json.JSONDecodeError:
- continue
-
- if record.get("type") != "message":
- continue
-
- message = record.get("message", {})
- role = message.get("role")
-
- if role not in ("user", "assistant"):
- continue
-
- text = self._extract_session_text(message.get("content"))
- if not text:
- continue
-
- label = "User" if role == "user" else "Assistant"
- collected.append(f"{label}: {text}")
-
- content = "\n".join(collected)
- rel_path = self._session_path_for_file(abs_path)
-
- return FileMetadata(
- hash=hash_text(content),
- mtime_ms=stat.st_mtime * 1000,
- size=stat.st_size,
- path=rel_path,
- abs_path=abs_path,
- content=content,
- )
-
- # ============================================================================
- # Session Processing Helpers
- # ============================================================================
-
- @staticmethod
- def _session_path_for_file(abs_path: str) -> str:
- """Convert absolute session path to relative."""
- return f"sessions/{os.path.basename(abs_path)}"
-
- def _extract_session_text(self, content: Any) -> str | None:
- """Extract text from session message content."""
- if isinstance(content, str):
- normalized = self._normalize_session_text(content)
- return normalized if normalized else None
-
- if not isinstance(content, list):
- return None
-
- parts = []
- for block in content:
- if not isinstance(block, dict):
- continue
-
- if block.get("type") != "text":
- continue
-
- text = block.get("text")
- if isinstance(text, str):
- normalized = self._normalize_session_text(text)
- if normalized:
- parts.append(normalized)
-
- return " ".join(parts) if parts else None
-
- @staticmethod
- def _normalize_session_text(text: str) -> str:
- """Normalize session text by collapsing whitespace."""
- text = re.sub(r"\s*\n+\s*", " ", text)
- text = re.sub(r"\s+", " ", text)
- return text.strip()
-
- # ============================================================================
- # Session Delta Tracking
- # ============================================================================
-
- async def _process_session_delta_batch(self) -> None:
- """Process pending session file changes."""
- if not self.session_pending_files:
- return
-
- pending = list(self.session_pending_files)
- self.session_pending_files.clear()
-
- should_sync = False
- for session_file in pending:
- delta = await self._update_session_delta(session_file)
- if not delta:
- continue
-
- bytes_threshold = self.settings.session_delta_bytes
- messages_threshold = self.settings.session_delta_messages
-
- if bytes_threshold <= 0:
- bytes_hit = delta["pending_bytes"] > 0
- else:
- bytes_hit = delta["pending_bytes"] >= bytes_threshold
-
- if messages_threshold <= 0:
- messages_hit = delta["pending_messages"] > 0
- else:
- messages_hit = delta["pending_messages"] >= messages_threshold
-
- if not bytes_hit and not messages_hit:
- continue
-
- self.sessions_dirty_files.add(session_file)
- self.sessions_dirty = True
- should_sync = True
-
- if should_sync:
- try:
- await self.sync(reason="session-delta")
- except Exception as err:
- logger.warning(f"memory sync failed (session-delta): {err}")
-
- async def _update_session_delta(self, session_file: str) -> dict[str, int] | None:
- """Update delta tracking for a session file."""
- try:
- stat = os.stat(session_file)
- size = stat.st_size
- except OSError:
- return None
-
- state = self.session_deltas.get(session_file)
- if not state:
- state = SessionDelta()
- self.session_deltas[session_file] = state
-
- delta_bytes = max(0, size - state.last_size)
-
- if delta_bytes == 0 and size == state.last_size:
- return {
- "delta_bytes": self.settings.session_delta_bytes,
- "delta_messages": self.settings.session_delta_messages,
- "pending_bytes": state.pending_bytes,
- "pending_messages": state.pending_messages,
- }
-
- if size < state.last_size:
- state.last_size = size
- state.pending_bytes += size
- if self.settings.session_delta_messages > 0:
- state.pending_messages += await self._count_newlines(session_file, 0, size)
- else:
- state.pending_bytes += delta_bytes
- if self.settings.session_delta_messages > 0:
- state.pending_messages += await self._count_newlines(session_file, state.last_size, size)
- state.last_size = size
-
- return {
- "delta_bytes": self.settings.session_delta_bytes,
- "delta_messages": self.settings.session_delta_messages,
- "pending_bytes": state.pending_bytes,
- "pending_messages": state.pending_messages,
- }
-
- def _reset_session_delta(self, abs_path: str, size: int) -> None:
- """Reset delta tracking for a session file."""
- state = self.session_deltas.get(abs_path)
- if state:
- state.last_size = size
- state.pending_bytes = 0
- state.pending_messages = 0
-
- @staticmethod
- async def _count_newlines(abs_path: str, start: int, end: int) -> int:
- """Count newlines in a file range."""
- if end <= start:
- return 0
-
- count = 0
- with open(abs_path, "rb") as f:
- f.seek(start)
- remaining = end - start
-
- while remaining > 0:
- chunk_size = min(SESSION_DELTA_READ_CHUNK_BYTES, remaining)
- chunk = f.read(chunk_size)
- if not chunk:
- break
-
- count += chunk.count(b"\n")
- remaining -= len(chunk)
-
- return count
-
- # ============================================================================
- # File Watchers
- # ============================================================================
-
- async def _start_watchers(self):
- """Start file watching and interval sync tasks."""
- if self.settings.watch_enabled and MemorySource.MEMORY in self.sources:
- self.watch_task = asyncio.create_task(self._watch_memory_files())
-
- if MemorySource.SESSIONS in self.sources:
- self.session_watch_task = asyncio.create_task(self._watch_session_files())
-
- if self.settings.interval_minutes > 0:
- self.interval_task = asyncio.create_task(self._interval_sync())
-
- async def _watch_memory_files(self) -> None:
- """Watch memory files for changes."""
- watch_paths = [
- os.path.join(self.workspace_dir, "MEMORY.md"),
- os.path.join(self.workspace_dir, "memory.md"),
- os.path.join(self.workspace_dir, "memory"),
- ]
-
- for extra in self.settings.extra_paths:
- watch_paths.append(extra)
-
- async for changes in awatch(*watch_paths, stop_event=None):
- if self.closed:
- break
-
- for _, path in changes:
- if path.endswith(".md"):
- self.dirty = True
- await asyncio.sleep(self.settings.watch_debounce_ms / 1000)
- try:
- await self.sync(reason="watch")
- except Exception as e:
- logger.exception(f"memory sync failed (watch): {e}")
-
- async def _watch_session_files(self):
- """Watch session files for changes."""
- sessions_dir = os.path.join(self.workspace_dir, "sessions", self.agent_id)
- if not os.path.exists(sessions_dir):
- return
-
- async for changes in awatch(sessions_dir, stop_event=None):
- if self.closed:
- break
-
- for _, path in changes:
- if path.endswith(".jsonl"):
- self.session_pending_files.add(path)
-
- await asyncio.sleep(SESSION_DIRTY_DEBOUNCE_MS / 1000)
- await self._process_session_delta_batch()
-
- async def _interval_sync(self) -> None:
- """Periodically sync the index."""
- while not self.closed:
- await asyncio.sleep(self.settings.interval_minutes * 60)
- if not self.closed:
- try:
- await self.sync(reason="interval")
- except Exception as err:
- logger.warning(f"memory sync failed (interval): {err}")
-
- # ============================================================================
- # Search Methods
- # ============================================================================
-
- async def _search_vector(self, query: str, limit: int) -> list[MemorySearchResult]:
- """Perform vector similarity search."""
- return await self.store.vector_search(query, limit, sources=list(self.sources))
-
- async def _search_keyword(self, query: str, limit: int) -> list[MemorySearchResult]:
- """Perform keyword/FTS search."""
- if not self.settings.fts_enabled:
- return []
-
- return await self.store.keyword_search(query, limit, sources=list(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."""
- 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
-
- results = list(merged.values())
- results.sort(key=lambda r: r.score, reverse=True)
- return results
-
- # ============================================================================
- # Utility Methods
- # ============================================================================
-
- @staticmethod
- def _is_memory_path(rel_path: str) -> bool:
- """Check if path is a valid memory path."""
- normalized = rel_path.replace("\\", "/")
-
- if normalized in ("MEMORY.md", "memory.md"):
- return True
-
- if normalized.startswith("memory/") and normalized.endswith(".md"):
- return True
-
- return False
diff --git a/reme/core/memory_manager/memory_storage/__init__.py b/reme/core/memory_manager/memory_storage/__init__.py
deleted file mode 100644
index e69de29b..00000000
diff --git a/reme/core/memory_manager/utils/__init__.py b/reme/core/memory_manager/utils/__init__.py
deleted file mode 100644
index e69de29b..00000000
diff --git a/reme/core/memory_manager/utils/hashing.py b/reme/core/memory_manager/utils/hashing.py
deleted file mode 100644
index 7894d69b..00000000
--- a/reme/core/memory_manager/utils/hashing.py
+++ /dev/null
@@ -1,15 +0,0 @@
-"""Utility functions for hashing text content."""
-
-import hashlib
-
-
-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/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_manager/memory_storage/base_memory_store.py b/reme/core/memory_storage/base_memory_store.py
similarity index 52%
rename from reme/core/memory_manager/memory_storage/base_memory_store.py
rename to reme/core/memory_storage/base_memory_store.py
index d17b8fa5..ed8f59e7 100644
--- a/reme/core/memory_manager/memory_storage/base_memory_store.py
+++ b/reme/core/memory_storage/base_memory_store.py
@@ -2,17 +2,31 @@
from abc import ABC, abstractmethod
-from ...embedding import BaseEmbeddingModel
-from ...enumeration import MemorySource
-from ...schema import FileMetadata, MemoryIndexMeta, MemoryChunk, MemorySearchResult
+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, embedding_model: BaseEmbeddingModel):
+ 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:
@@ -20,77 +34,21 @@ class BaseMemoryStore(ABC):
return self.embedding_model.dimensions
async def get_embedding(self, query: str, **kwargs) -> list[float]:
- """Get embedding for a single query string.
-
- Args:
- query: Input text to generate embedding for
- **kwargs: Additional arguments passed to the embedding model
-
- Returns:
- Embedding vector as a list of floats
- """
+ """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.
-
- Args:
- queries: List of input texts to generate embeddings for
- **kwargs: Additional arguments passed to the embedding model
-
- Returns:
- List of embedding vectors, each as a list of floats
- """
+ """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.
-
- 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
- """
+ """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.
-
- 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
- """
+ """Generate and populate embedding fields for a batch of MemoryChunk objects."""
return await self.embedding_model.get_chunk_embeddings(chunks, **kwargs)
- 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
- """
- return self.embedding_model.get_chunk_embedding_sync(chunk, **kwargs)
-
- 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
- """
- return self.embedding_model.get_chunk_embeddings_sync(chunks, **kwargs)
-
@abstractmethod
async def start(self):
"""Initialize the storage backend."""
@@ -104,19 +62,23 @@ class BaseMemoryStore(ABC):
"""Delete a file and all its chunks."""
@abstractmethod
- async def get_file_hash(self, path: str, source: MemorySource) -> str | None:
- """Get the hash of an indexed file."""
+ async def delete_file_chunks(self, path: str, chunk_ids: list[str]):
+ """Delete chunks for a file."""
@abstractmethod
- async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None:
- """Get full file metadata with statistics."""
+ 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_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
+ 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
@@ -155,14 +117,6 @@ class BaseMemoryStore(ABC):
List of search results sorted by relevance
"""
- @abstractmethod
- async def read_meta(self, key: str) -> MemoryIndexMeta | None:
- """Read metadata value."""
-
- @abstractmethod
- async def write_meta(self, key: str, value: MemoryIndexMeta | dict):
- """Write metadata value."""
-
@abstractmethod
async def clear_all(self):
"""Clear all indexed data."""
diff --git a/reme/core/memory_manager/memory_storage/sqlite_memory_store.py b/reme/core/memory_storage/sqlite_memory_store.py
similarity index 53%
rename from reme/core/memory_manager/memory_storage/sqlite_memory_store.py
rename to reme/core/memory_storage/sqlite_memory_store.py
index 0d6b47fe..756c2b3b 100644
--- a/reme/core/memory_manager/memory_storage/sqlite_memory_store.py
+++ b/reme/core/memory_storage/sqlite_memory_store.py
@@ -9,9 +9,8 @@ from pathlib import Path
from loguru import logger
from .base_memory_store import BaseMemoryStore
-from ...embedding import BaseEmbeddingModel
-from ...enumeration import MemorySource
-from ...schema import FileMetadata, MemoryIndexMeta, MemoryChunk, MemorySearchResult
+from ..enumeration import MemorySource
+from ..schema import FileMetadata, MemoryChunk, MemorySearchResult
class SqliteMemoryStore(BaseMemoryStore):
@@ -28,25 +27,32 @@ class SqliteMemoryStore(BaseMemoryStore):
- Efficient chunk and file metadata management
"""
- VECTOR_TABLE = "chunks_vec"
- FTS_TABLE = "chunks_fts"
-
- def __init__(
- self,
- db_path: str,
- embedding_model: BaseEmbeddingModel,
- vec_ext_path: str = "",
- fts_enabled: bool = True,
- snippet_max_chars: int = 700,
- ):
- super().__init__(embedding_model=embedding_model)
+ 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.fts_enabled = fts_enabled
- self.snippet_max_chars = snippet_max_chars
+
self.conn: sqlite3.Connection | None = None
- self.vector_available = False
- self.fts_available = False
+
+ @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:
@@ -55,6 +61,9 @@ class SqliteMemoryStore(BaseMemoryStore):
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)
@@ -68,16 +77,27 @@ class SqliteMemoryStore(BaseMemoryStore):
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 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
+ 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()
@@ -86,20 +106,10 @@ class SqliteMemoryStore(BaseMemoryStore):
"""Create database schema."""
cursor = self.conn.cursor()
- # Metadata
- cursor.execute(
- """
- CREATE TABLE IF NOT EXISTS meta (
- key TEXT PRIMARY KEY,
- value TEXT
- )
- """,
- )
-
# Files
cursor.execute(
- """
- CREATE TABLE IF NOT EXISTS files (
+ f"""
+ CREATE TABLE IF NOT EXISTS {self.files_table_name} (
path TEXT,
source TEXT,
hash TEXT,
@@ -112,8 +122,8 @@ class SqliteMemoryStore(BaseMemoryStore):
# Chunks
cursor.execute(
- """
- CREATE TABLE IF NOT EXISTS chunks (
+ f"""
+ CREATE TABLE IF NOT EXISTS {self.chunks_table_name} (
id TEXT PRIMARY KEY,
path TEXT,
source TEXT,
@@ -127,49 +137,34 @@ class SqliteMemoryStore(BaseMemoryStore):
""",
)
- cursor.execute(
- """
- CREATE INDEX IF NOT EXISTS idx_chunks_path_source
- ON chunks(path, source)
- """,
- )
-
# Vector table (sqlite-vec)
if self.vector_available:
- try:
- cursor.execute(
- f"""
- CREATE VIRTUAL TABLE IF NOT EXISTS {self.VECTOR_TABLE} USING vec0(
- id TEXT PRIMARY KEY,
- embedding FLOAT[{self.embedding_dim}]
- )
- """,
+ 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})")
- except Exception as e:
- logger.warning(f"Failed to create vector table: {e}")
- self.vector_available = False
+ """,
+ )
+ logger.info(f"Created vector table (dims={self.embedding_dim})")
# FTS table
if self.fts_enabled:
- try:
- cursor.execute(
- f"""
- CREATE VIRTUAL TABLE IF NOT EXISTS {self.FTS_TABLE} USING fts5(
- text,
- id UNINDEXED,
- path UNINDEXED,
- source UNINDEXED,
- start_line UNINDEXED,
- end_line UNINDEXED
- )
- """,
+ 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")
- except Exception as e:
- logger.warning(f"Failed to create FTS table: {e}")
- self.fts_available = False
+ """,
+ )
+ self.fts_available = True
+ logger.info("Created FTS5 table")
self.conn.commit()
cursor.close()
@@ -180,12 +175,11 @@ class SqliteMemoryStore(BaseMemoryStore):
try:
cursor.execute("BEGIN")
- await self._delete_file_internal(cursor, file_meta.path, source)
# Insert file
cursor.execute(
- """
- INSERT OR REPLACE INTO files (path, source, hash, mtime, size)
+ 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),
@@ -195,8 +189,8 @@ class SqliteMemoryStore(BaseMemoryStore):
now = int(time.time() * 1000)
for chunk in chunks:
cursor.execute(
- """
- INSERT INTO chunks (
+ f"""
+ INSERT OR REPLACE INTO {self.chunks_table_name} (
id, path, source, start_line, end_line,
hash, text, embedding, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
@@ -215,38 +209,33 @@ class SqliteMemoryStore(BaseMemoryStore):
)
# Insert vector
- if self.vector_available and chunk.embedding:
- try:
- cursor.execute(
- f"""
- INSERT INTO {self.VECTOR_TABLE} (id, embedding)
- VALUES (?, ?)
- """,
- (chunk.id, self.vector_to_blob(chunk.embedding)),
- )
- except Exception as e:
- logger.debug(f"Vector insert failed: {e}")
+ 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:
- try:
- cursor.execute(
- f"""
- INSERT INTO {self.FTS_TABLE} (
- text, id, path, source, start_line, end_line
- ) VALUES (?, ?, ?, ?, ?, ?)
- """,
- (
- chunk.text,
- chunk.id,
- file_meta.path,
- source.value,
- chunk.start_line,
- chunk.end_line,
- ),
- )
- except Exception as e:
- logger.debug(f"FTS insert failed: {e}")
+ 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:
@@ -255,12 +244,50 @@ class SqliteMemoryStore(BaseMemoryStore):
finally:
cursor.close()
- async def delete_file(self, path: str, source: MemorySource) -> None:
+ async def delete_file(self, path: str, source: MemorySource):
"""Delete file and all its chunks."""
cursor = self.conn.cursor()
try:
cursor.execute("BEGIN")
- await self._delete_file_internal(cursor, path, source)
+
+ # 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")
@@ -268,62 +295,132 @@ class SqliteMemoryStore(BaseMemoryStore):
finally:
cursor.close()
- async def _delete_file_internal(self, cursor: sqlite3.Cursor, path: str, source: MemorySource):
- """Internal delete helper."""
- # Get chunk IDs for vector deletion
- cursor.execute(
- "SELECT id FROM chunks WHERE path = ? AND source = ?",
- (path, source.value),
- )
- chunk_ids = [row[0] for row in cursor.fetchall()]
+ async def delete_file_chunks(self, path: str, chunk_ids: list[str]):
+ """Delete specific chunks for a file."""
+ if not chunk_ids:
+ return
- # Delete vectors
- if self.vector_available and chunk_ids:
- for chunk_id in chunk_ids:
+ 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.VECTOR_TABLE} WHERE id = ?",
- (chunk_id,),
+ f"DELETE FROM {self.fts_table_name} WHERE id IN ({placeholders})",
+ chunk_ids,
)
except Exception as e:
- logger.debug(f"Vector delete failed: {e}")
+ logger.debug(f"FTS delete failed: {e}")
- # Delete FTS entries
- if self.fts_available:
- try:
- cursor.execute(
- f"DELETE FROM {self.FTS_TABLE} WHERE path = ? AND source = ?",
- (path, source.value),
- )
- 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,
+ )
- # Delete chunks and file
- cursor.execute(
- "DELETE FROM chunks WHERE path = ? AND source = ?",
- (path, source.value),
- )
- cursor.execute(
- "DELETE FROM files WHERE path = ? AND source = ?",
- (path, source.value),
- )
+ 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
- async def get_file_hash(self, path: str, source: MemorySource) -> str | None:
- """Get file hash."""
cursor = self.conn.cursor()
- cursor.execute(
- "SELECT hash FROM files WHERE path = ? AND source = ?",
- (path, source.value),
- )
- row = cursor.fetchone()
+ 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 row[0] if row else None
+ 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(
- "SELECT hash, mtime, size FROM files WHERE path = ? AND source = ?",
+ f"SELECT hash, mtime, size FROM {self.files_table_name} WHERE path = ? AND source = ?",
(path, source.value),
)
row = cursor.fetchone()
@@ -333,7 +430,7 @@ class SqliteMemoryStore(BaseMemoryStore):
hash_val, mtime, size = row
cursor.execute(
- "SELECT COUNT(*) FROM chunks WHERE path = ? AND source = ?",
+ f"SELECT COUNT(*) FROM {self.chunks_table_name} WHERE path = ? AND source = ?",
(path, source.value),
)
chunk_count = cursor.fetchone()[0]
@@ -343,24 +440,17 @@ class SqliteMemoryStore(BaseMemoryStore):
hash=hash_val,
mtime_ms=mtime,
size=size,
+ path=path,
chunk_count=chunk_count,
)
- async def list_files(self, source: MemorySource) -> list[str]:
- """List all indexed files."""
- cursor = self.conn.cursor()
- cursor.execute("SELECT path FROM files WHERE source = ?", (source.value,))
- paths = [row[0] for row in cursor.fetchall()]
- cursor.close()
- return paths
-
- async def get_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
+ 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 chunks WHERE path = ? AND source = ?
+ FROM {self.chunks_table_name} WHERE path = ? AND source = ?
ORDER BY start_line
""",
(path, source.value),
@@ -418,23 +508,25 @@ class SqliteMemoryStore(BaseMemoryStore):
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} v
- JOIN chunks c ON v.id = c.id
+ 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]
+ query_params: list = [query_blob, limit]
# Add source filter if specified
if source_filter:
query_sql += source_filter
query_params.extend(params)
- # Order and limit results
- query_sql += " ORDER BY v.distance LIMIT ?"
- query_params.append(str(limit))
+ # Order by distance (k constraint already limits results)
+ query_sql += " ORDER BY v.distance"
cursor.execute(query_sql, query_params)
@@ -470,11 +562,21 @@ class SqliteMemoryStore(BaseMemoryStore):
if not self.fts_available:
return []
- # Build FTS5 query, escaping quotes
- cleaned = query.strip().replace('"', '""')
+ # 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 []
- fts_query = f'"{cleaned}"'
+
+ # 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 = ""
@@ -490,7 +592,7 @@ class SqliteMemoryStore(BaseMemoryStore):
f"""
SELECT fts.id, fts.path, fts.start_line, fts.end_line,
fts.source, fts.text, rank
- FROM {self.FTS_TABLE} fts
+ FROM {self.fts_table_name} fts
WHERE fts.text MATCH ?{source_filter}
ORDER BY rank
LIMIT ?
@@ -521,52 +623,20 @@ class SqliteMemoryStore(BaseMemoryStore):
finally:
cursor.close()
- async def read_meta(self, key: str) -> MemoryIndexMeta | None:
- """Read metadata value."""
- cursor = self.conn.cursor()
- cursor.execute("SELECT value FROM meta WHERE key = ?", (key,))
- row = cursor.fetchone()
- cursor.close()
-
- if not row:
- return None
-
- return MemoryIndexMeta(**json.loads(row[0]))
-
- async def write_meta(self, key: str, value: MemoryIndexMeta | dict) -> None:
- """Write metadata value."""
- data = value.model_dump() if isinstance(value, MemoryIndexMeta) else value
- cursor = self.conn.cursor()
- cursor.execute(
- """
- INSERT OR REPLACE INTO meta (key, value)
- VALUES (?, ?)
- """,
- (key, json.dumps(data)),
- )
- self.conn.commit()
- cursor.close()
-
async def clear_all(self):
"""Clear all indexed data."""
cursor = self.conn.cursor()
cursor.execute("BEGIN")
try:
- cursor.execute("DELETE FROM files")
- cursor.execute("DELETE FROM chunks")
+ cursor.execute(f"DELETE FROM {self.files_table_name}")
+ cursor.execute(f"DELETE FROM {self.chunks_table_name}")
if self.vector_available:
- try:
- cursor.execute(f"DELETE FROM {self.VECTOR_TABLE}")
- except Exception as e:
- logger.debug(f"Vector clear failed: {e}")
+ cursor.execute(f"DELETE FROM {self.vector_table_name}")
if self.fts_available:
- try:
- cursor.execute(f"DELETE FROM {self.FTS_TABLE}")
- except Exception as e:
- logger.debug(f"FTS clear failed: {e}")
+ cursor.execute(f"DELETE FROM {self.fts_table_name}")
cursor.execute("COMMIT")
except Exception:
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/file_metadata.py b/reme/core/schema/file_metadata.py
index 053f1e62..672e9c27 100644
--- a/reme/core/schema/file_metadata.py
+++ b/reme/core/schema/file_metadata.py
@@ -6,15 +6,10 @@ from pydantic import BaseModel, Field
class FileMetadata(BaseModel):
"""File metadata with optional extended fields for various use cases."""
- # Core fields (always required)
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")
-
- # Extended fields for session files
path: str | None = Field(default=None, description="Relative path to the session file")
- abs_path: str | None = Field(default=None, description="Absolute path to the session file")
content: str | None = Field(default=None, description="Parsed content from the session file")
-
- # Extended fields for statistics
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/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/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/memory_manager/ingestion/chunking.py b/reme/core/utils/chunking_utils.py
similarity index 87%
rename from reme/core/memory_manager/ingestion/chunking.py
rename to reme/core/utils/chunking_utils.py
index 9a159cf0..e1511200 100644
--- a/reme/core/memory_manager/ingestion/chunking.py
+++ b/reme/core/utils/chunking_utils.py
@@ -1,19 +1,17 @@
"""Chunking logic for Markdown files."""
-from typing import List, Dict, Any
-
-from ..utils.hashing import hash_text
-from ...enumeration import MemorySource
-from ...schema import MemoryChunk
+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 = 300,
- overlap: int = 30,
-) -> List[MemoryChunk]:
+ chunk_tokens: int,
+ overlap: int,
+) -> list[MemoryChunk]:
"""
Markdown chunking logic implemented based on the TypeScript version.
@@ -35,10 +33,10 @@ def chunk_markdown(
max_chars = max(32, chunk_tokens * 4)
overlap_chars = max(0, overlap * 4)
- chunks: List[MemoryChunk] = []
+ chunks: list[MemoryChunk] = []
# Currently building chunk
- current: List[Dict[str, Any]] = [] # [{'line': str, 'line_no': int}]
+ current: list[dict] = [] # [{'line': str, 'line_no': int}]
current_chars = 0
def flush():
@@ -83,8 +81,8 @@ def chunk_markdown(
kept = []
# Collect lines from the end until reaching overlap size
- for i in range(len(current) - 1, -1, -1):
- entry = current[i]
+ for j in range(len(current) - 1, -1, -1):
+ entry = current[j]
if not entry:
continue
@@ -123,4 +121,4 @@ def chunk_markdown(
# Process the final chunk
flush()
- return chunks
+ 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 060019b7..81050cbd 100644
--- a/reme/reme.py
+++ b/reme/reme.py
@@ -58,7 +58,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.
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/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_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())