Merge pull request #100 from agentscope-ai/dev_0205

Dev 0205
This commit is contained in:
jinliyl 2026-02-06 12:17:15 +08:00 • committed by GitHub
commit 90cd9ddea8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
57 changed files with 7322 additions and 70 deletions

4
.gitignore vendored
View file

@ -39,4 +39,6 @@ chroma_vector_store/*
bench_results/*
meta_memory/*
*.sqlite3
**/data/*.json
**/data/*.json
*.db
memories/*

View file

@ -78,7 +78,7 @@ repos:
--disable=C3001,
--disable=R1702,
--disable=R0912,
--max-statements=75,
--max-statements=120,
--max-line-length=120,
]
- repo: https://github.com/regebro/pyroma

View file

@ -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/*

View file

@ -0,0 +1,9 @@
"""File system agents for memory management."""
from .fs_compactor import FsCompactor
from .fs_summarizer import FsSummarizer
__all__ = [
"FsSummarizer",
"FsCompactor",
]

View file

@ -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"<conversation>\n{conversation_text}\n</conversation>\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,
}

View file

@ -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
<previous-summary> tags.
<previous-summary>
{previous_summary}
</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>标签中提供。
<previous-summary>
{previous_summary}
</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>
{conversation_text}
</conversation>
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>
{conversation_text}
</conversation>
这是一个过长而无法保留的回合的前缀。后缀(最近的工作)已保留。
总结前缀以为保留的后缀提供上下文:
## 原始请求
[用户在此回合中要求了什么?]
## 早期进展
- [在前缀中做出的关键决策和完成的工作]
## 后缀上下文
- [理解保留的最近工作所需的信息]
保持简洁。专注于理解保留后缀所需的内容。

View file

@ -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

View file

@ -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].

View file

@ -66,7 +66,7 @@ class BaseMemoryAgent(BaseReact, metaclass=ABCMeta):
lines = []
for memory_target, memory_type in self.memory_target_type_mapping.items():
line = {
"agent": f"Agent managing {memory_type} memories for {memory_target}",
"agent": f"Agent managing {memory_type.value} memories for {memory_target}",
"memory_target": memory_target,
}
lines.append(json.dumps(line, ensure_ascii=False))

View file

@ -68,25 +68,27 @@ class ReMeRetriever(BaseMemoryAgent):
async def execute(self):
result = await super().execute()
tools: list[BaseTool] = result["tools"]
delegate_task_tool = tools[0]
agents: list[BaseMemoryAgent] = delegate_task_tool.response.metadata["agents"]
answer = []
success = True
messages = []
tools = []
tools_result = []
retrieved_nodes = []
for agent in agents:
answer.append(agent.response.answer)
success = success and agent.response.success
messages.extend(agent.response.metadata["messages"])
tools.extend(agent.response.metadata["tools"])
retrieved_nodes.extend(agent.response.metadata["retrieved_nodes"])
if tools:
delegate_task_tool = tools[0]
agents: list[BaseMemoryAgent] = delegate_task_tool.response.metadata["agents"]
for agent in agents:
answer.append(agent.response.answer)
success = success and agent.response.success
messages.extend(agent.response.metadata["messages"])
tools_result.extend(agent.response.metadata["tools"])
retrieved_nodes.extend(agent.response.metadata["retrieved_nodes"])
return {
"answer": "\n".join(answer),
"success": True,
"messages": messages,
"tools": tools,
"tools": tools_result,
"retrieved_nodes": retrieved_nodes,
}

View file

@ -74,22 +74,25 @@ class ReMeSummarizer(BaseMemoryAgent):
async def execute(self):
result = await super().execute()
tools: list[BaseTool] = result["tools"]
delegate_task_tool = tools[0]
agents: list[BaseMemoryAgent] = delegate_task_tool.response.metadata["agents"]
success = True
messages = []
tools = []
tools_result = []
memory_nodes = []
for agent in agents:
success = success and agent.response.success
messages.extend(agent.response.metadata["messages"])
tools.extend(agent.response.metadata["tools"])
memory_nodes.extend(agent.response.metadata["memory_nodes"])
if tools:
delegate_task_tool = tools[0]
agents: list[BaseMemoryAgent] = delegate_task_tool.response.metadata["agents"]
for agent in agents:
success = success and agent.response.success
messages.extend(agent.response.metadata["messages"])
tools_result.extend(agent.response.metadata["tools"])
memory_nodes.extend(agent.response.metadata["memory_nodes"])
return {
"answer": memory_nodes,
"success": True,
"messages": messages,
"tools": tools,
"tools": tools_result,
}

View file

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

View file

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

View file

@ -4,12 +4,15 @@ Defines the abstract base class and standard API for all embedding model impleme
"""
import asyncio
import hashlib
import time
from abc import ABC
from collections import OrderedDict
from loguru import logger
from ..schema import VectorNode
from ..schema.memory_chunk import MemoryChunk
class BaseEmbeddingModel(ABC):
@ -27,17 +30,35 @@ class BaseEmbeddingModel(ABC):
max_retries: int = 3,
raise_exception: bool = True,
max_input_length: int = 8192,
max_cache_size: int = 10000,
**kwargs,
):
"""Initialize model configuration and parameters."""
"""Initialize model configuration and parameters.
Args:
model_name: Name of the embedding model
dimensions: Vector dimensions of the embeddings
max_batch_size: Maximum batch size for embedding requests
max_retries: Maximum number of retry attempts on failure
raise_exception: Whether to raise exceptions on failure
max_input_length: Maximum input text length
max_cache_size: Maximum number of embeddings to cache in memory (LRU)
**kwargs: Additional model-specific parameters
"""
self.model_name = model_name
self.dimensions = dimensions
self.max_batch_size = max_batch_size
self.max_retries = max_retries
self.raise_exception = raise_exception
self.max_input_length = max_input_length
self.max_cache_size = max_cache_size
self.kwargs = kwargs
# Initialize LRU cache for embeddings
self._embedding_cache: OrderedDict[str, list[float]] = OrderedDict()
self._cache_hits = 0
self._cache_misses = 0
def _truncate_text(self, text: str) -> str:
"""Truncate text to max_input_length if it exceeds the limit."""
if len(text) > self.max_input_length:
@ -51,6 +72,76 @@ class BaseEmbeddingModel(ABC):
"""Truncate a list of texts to max_input_length."""
return [self._truncate_text(text) for text in texts]
def _get_cache_key(self, text: str) -> str:
"""Generate a cache key by hashing the input text.
Args:
text: Input text to hash
Returns:
SHA256 hash of the text as hexadecimal string
"""
return hashlib.sha256(text.encode("utf-8")).hexdigest()
def _get_from_cache(self, text: str) -> list[float] | None:
"""Retrieve embedding from cache if it exists.
Args:
text: Input text to look up
Returns:
Cached embedding vector or None if not found
"""
cache_key = self._get_cache_key(text)
if cache_key in self._embedding_cache:
# Move to end (most recently used)
self._embedding_cache.move_to_end(cache_key)
self._cache_hits += 1
return self._embedding_cache[cache_key]
self._cache_misses += 1
return None
def _put_to_cache(self, text: str, embedding: list[float]) -> None:
"""Store embedding in cache with LRU eviction.
Args:
text: Input text used as cache key
embedding: Embedding vector to cache
"""
if self.max_cache_size <= 0:
return
cache_key = self._get_cache_key(text)
# Remove oldest entry if cache is full
if len(self._embedding_cache) >= self.max_cache_size and cache_key not in self._embedding_cache:
self._embedding_cache.popitem(last=False)
self._embedding_cache[cache_key] = embedding
self._embedding_cache.move_to_end(cache_key)
def get_cache_stats(self) -> dict[str, int]:
"""Get cache statistics.
Returns:
Dictionary with cache size, hits, misses, and hit rate
"""
total_requests = self._cache_hits + self._cache_misses
hit_rate = self._cache_hits / total_requests if total_requests > 0 else 0.0
return {
"cache_size": len(self._embedding_cache),
"max_cache_size": self.max_cache_size,
"cache_hits": self._cache_hits,
"cache_misses": self._cache_misses,
"hit_rate": hit_rate,
}
def clear_cache(self) -> None:
"""Clear the embedding cache and reset statistics."""
self._embedding_cache.clear()
self._cache_hits = 0
self._cache_misses = 0
async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Internal async implementation for calling the embedding API with batch input."""
@ -60,10 +151,20 @@ class BaseEmbeddingModel(ABC):
async def get_embedding(self, input_text: str, **kwargs) -> list[float]:
"""Async get embedding for a single text with exponential backoff retries."""
truncated_text = self._truncate_text(input_text)
# Check cache first
cached_embedding = self._get_from_cache(truncated_text)
if cached_embedding is not None:
return cached_embedding
# Cache miss - compute embedding
for i in range(self.max_retries):
try:
result = await self._get_embeddings([truncated_text], **kwargs)
return result[0]
embedding = result[0]
# Store in cache
self._put_to_cache(truncated_text, embedding)
return embedding
except Exception as e:
logger.error(f"Model {self.model_name} failed: {e}")
if i == self.max_retries - 1:
@ -78,16 +179,36 @@ class BaseEmbeddingModel(ABC):
# Truncate all input texts first
truncated_texts = self._truncate_texts(input_text)
# Split into batches and process sequentially to respect rate limits
results = []
for i in range(0, len(truncated_texts), self.max_batch_size):
batch = truncated_texts[i : i + self.max_batch_size]
# Check cache for each text and separate cached vs uncached
results: list[list[float] | None] = [None] * len(truncated_texts)
texts_to_compute: list[tuple[int, str]] = [] # (original_index, text)
for idx, text in enumerate(truncated_texts):
cached = self._get_from_cache(text)
if cached is not None:
results[idx] = cached
else:
texts_to_compute.append((idx, text))
# If all texts were cached, return early
if not texts_to_compute:
return [r for r in results if r is not None]
# Compute embeddings for uncached texts in batches
uncached_texts = [text for _, text in texts_to_compute]
for i in range(0, len(uncached_texts), self.max_batch_size):
batch_texts = uncached_texts[i : i + self.max_batch_size]
batch_indices = [idx for idx, _ in texts_to_compute[i : i + self.max_batch_size]]
# Process each batch with retry logic
for retry in range(self.max_retries):
try:
batch_res = await self._get_embeddings(batch, **kwargs)
if batch_res:
results.extend(batch_res)
batch_embeddings = await self._get_embeddings(batch_texts, **kwargs)
if batch_embeddings:
# Store results and cache them
for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings):
results[orig_idx] = embedding
self._put_to_cache(text, embedding)
break
except Exception as e:
logger.error(f"Model {self.model_name} batch failed: {e}")
@ -96,15 +217,26 @@ class BaseEmbeddingModel(ABC):
raise
else:
await asyncio.sleep(retry + 1)
return results
return [r for r in results if r is not None]
def get_embedding_sync(self, input_text: str, **kwargs) -> list[float]:
"""Synchronous get embedding for a single text with retry logic."""
truncated_text = self._truncate_text(input_text)
# Check cache first
cached_embedding = self._get_from_cache(truncated_text)
if cached_embedding is not None:
return cached_embedding
# Cache miss - compute embedding
for i in range(self.max_retries):
try:
result = self._get_embeddings_sync([truncated_text], **kwargs)
return result[0]
embedding = result[0]
# Store in cache
self._put_to_cache(truncated_text, embedding)
return embedding
except Exception as exc:
logger.error(f"Model {self.model_name} failed: {exc}")
if i == self.max_retries - 1:
@ -119,15 +251,36 @@ class BaseEmbeddingModel(ABC):
# Truncate all input texts first
truncated_texts = self._truncate_texts(input_text)
results = []
for i in range(0, len(truncated_texts), self.max_batch_size):
batch = truncated_texts[i : i + self.max_batch_size]
# Check cache for each text and separate cached vs uncached
results: list[list[float] | None] = [None] * len(truncated_texts)
texts_to_compute: list[tuple[int, str]] = [] # (original_index, text)
for idx, text in enumerate(truncated_texts):
cached = self._get_from_cache(text)
if cached is not None:
results[idx] = cached
else:
texts_to_compute.append((idx, text))
# If all texts were cached, return early
if not texts_to_compute:
return [r for r in results if r is not None]
# Compute embeddings for uncached texts in batches
uncached_texts = [text for _, text in texts_to_compute]
for i in range(0, len(uncached_texts), self.max_batch_size):
batch_texts = uncached_texts[i : i + self.max_batch_size]
batch_indices = [idx for idx, _ in texts_to_compute[i : i + self.max_batch_size]]
# Process each batch with retry logic
for retry in range(self.max_retries):
try:
batch_res = self._get_embeddings_sync(batch, **kwargs)
if batch_res:
results.extend(batch_res)
batch_embeddings = self._get_embeddings_sync(batch_texts, **kwargs)
if batch_embeddings:
# Store results and cache them
for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings):
results[orig_idx] = embedding
self._put_to_cache(text, embedding)
break
except Exception as exc:
logger.error(f"Model {self.model_name} batch failed: {exc}")
@ -136,7 +289,8 @@ class BaseEmbeddingModel(ABC):
raise
else:
time.sleep(retry + 1)
return results
return [r for r in results if r is not None]
async def get_node_embedding(self, node: VectorNode, **kwargs) -> VectorNode:
"""Async generate and populate vector field for a single VectorNode object."""
@ -172,6 +326,72 @@ class BaseEmbeddingModel(ABC):
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(nodes)} nodes")
return nodes
async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
"""Async generate and populate embedding field for a single MemoryChunk object.
Args:
chunk: MemoryChunk object containing text to embed
**kwargs: Additional arguments passed to the embedding model
Returns:
The same MemoryChunk object with populated embedding field
"""
chunk.embedding = await self.get_embedding(chunk.text, **kwargs)
return chunk
async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
"""Async generate and populate embedding fields for a batch of MemoryChunk objects.
Args:
chunks: List of MemoryChunk objects containing text to embed
**kwargs: Additional arguments passed to the embedding model
Returns:
The same list of MemoryChunk objects with populated embedding fields
"""
texts = [chunk.text for chunk in chunks]
embeddings: list[list[float]] = await self.get_embeddings(texts, **kwargs)
if len(embeddings) == len(chunks):
for chunk, vec in zip(chunks, embeddings):
chunk.embedding = vec
else:
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(chunks)} chunks")
return chunks
def get_chunk_embedding_sync(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
"""Synchronously generate and populate embedding field for a single MemoryChunk object.
Args:
chunk: MemoryChunk object containing text to embed
**kwargs: Additional arguments passed to the embedding model
Returns:
The same MemoryChunk object with populated embedding field
"""
chunk.embedding = self.get_embedding_sync(chunk.text, **kwargs)
return chunk
def get_chunk_embeddings_sync(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
"""Synchronously generate embeddings for a batch of MemoryChunk objects.
Args:
chunks: List of MemoryChunk objects containing text to embed
**kwargs: Additional arguments passed to the embedding model
Returns:
The same list of MemoryChunk objects with populated embedding fields
"""
texts = [chunk.text for chunk in chunks]
embeddings: list[list[float]] = self.get_embeddings_sync(texts, **kwargs)
if len(embeddings) == len(chunks):
for chunk, vec in zip(chunks, embeddings):
chunk.embedding = vec
else:
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(chunks)} chunks")
return chunks
def close_sync(self):
"""Synchronously release resources and close connections."""

View file

@ -3,6 +3,7 @@
from .chunk_enum import ChunkEnum
from .http_enum import HttpEnum
from .json_schema_enum import JsonSchemaEnum
from .memory_source import MemorySource
from .memory_type import MemoryType
from .registry_enum import RegistryEnum
from .role import Role
@ -11,6 +12,7 @@ __all__ = [
"ChunkEnum",
"HttpEnum",
"JsonSchemaEnum",
"MemorySource",
"MemoryType",
"RegistryEnum",
"Role",

View file

@ -0,0 +1,11 @@
"""Memory source types."""
from enum import Enum
class MemorySource(str, Enum):
"""Source of memory data."""
MEMORY = "memory"
SESSIONS = "sessions"

View file

@ -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"

View file

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

View file

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

View file

@ -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

View file

@ -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

View file

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

View file

@ -0,0 +1,126 @@
"""Base storage interface for memory manager."""
from abc import ABC, abstractmethod
from ..embedding import BaseEmbeddingModel
from ..enumeration import MemorySource
from ..schema import FileMetadata, MemoryChunk, MemorySearchResult
class BaseMemoryStore(ABC):
"""Abstract base class for memory storage backends."""
def __init__(
self,
store_name: str,
embedding_model: BaseEmbeddingModel,
fts_enabled: bool = True,
snippet_max_chars: int = 700,
**kwargs,
):
"""Initialize"""
self.store_name: str = store_name
self.embedding_model: BaseEmbeddingModel = embedding_model
self.fts_enabled: bool = fts_enabled
self.snippet_max_chars: int = snippet_max_chars
self.kwargs: dict = kwargs
self.vector_available = False
self.fts_available = False
@property
def embedding_dim(self) -> int:
"""Get the embedding model's dimensionality."""
return self.embedding_model.dimensions
async def get_embedding(self, query: str, **kwargs) -> list[float]:
"""Get embedding for a single query string."""
return await self.embedding_model.get_embedding(query, **kwargs)
async def get_embeddings(self, queries: list[str], **kwargs) -> list[list[float]]:
"""Get embeddings for a batch of query strings."""
return await self.embedding_model.get_embeddings(queries, **kwargs)
async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
"""Generate and populate embedding field for a single MemoryChunk object."""
return await self.embedding_model.get_chunk_embedding(chunk, **kwargs)
async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
"""Generate and populate embedding fields for a batch of MemoryChunk objects."""
return await self.embedding_model.get_chunk_embeddings(chunks, **kwargs)
@abstractmethod
async def start(self):
"""Initialize the storage backend."""
@abstractmethod
async def upsert_file(self, file_meta: FileMetadata, source: MemorySource, chunks: list[MemoryChunk]):
"""Insert or update a file and its chunks."""
@abstractmethod
async def delete_file(self, path: str, source: MemorySource):
"""Delete a file and all its chunks."""
@abstractmethod
async def delete_file_chunks(self, path: str, chunk_ids: list[str]):
"""Delete chunks for a file."""
@abstractmethod
async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource):
"""Insert or update specific chunks without affecting other chunks."""
@abstractmethod
async def list_files(self, source: MemorySource) -> list[str]:
"""List all indexed file paths for a source."""
@abstractmethod
async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None:
"""Get full file metadata with statistics."""
@abstractmethod
async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
"""Get all chunks for a file."""
@abstractmethod
async def vector_search(
self,
query: str,
limit: int,
sources: list[MemorySource] | None = None,
) -> list[MemorySearchResult]:
"""Perform vector similarity search.
Args:
query: Query embedding vector
limit: Maximum number of results
sources: Optional list of sources to filter
Returns:
List of search results sorted by similarity
"""
@abstractmethod
async def keyword_search(
self,
query: str,
limit: int,
sources: list[MemorySource] | None = None,
) -> list[MemorySearchResult]:
"""Perform keyword/full-text search.
Args:
query: Search query text
limit: Maximum number of results
sources: Optional list of sources to filter
Returns:
List of search results sorted by relevance
"""
@abstractmethod
async def clear_all(self):
"""Clear all indexed data."""
@abstractmethod
async def close(self):
"""Close storage and release resources."""

View file

@ -0,0 +1,652 @@
"""SQLite storage backend for memory index."""
import json
import sqlite3
import struct
import time
from pathlib import Path
from loguru import logger
from .base_memory_store import BaseMemoryStore
from ..enumeration import MemorySource
from ..schema import FileMetadata, MemoryChunk, MemorySearchResult
class SqliteMemoryStore(BaseMemoryStore):
"""SQLite memory storage with vector and full-text search.
Inherits embedding methods from BaseMemoryStore:
- get_chunk_embedding / get_chunk_embeddings (async)
- get_chunk_embedding_sync / get_chunk_embeddings_sync (sync)
- get_embedding / get_embeddings (async)
Provides SQLite-backed persistent storage with:
- Vector similarity search (via sqlite-vec extension)
- Full-text search (via FTS5)
- Efficient chunk and file metadata management
"""
def __init__(self, db_path: str = ".reme/memory.db", vec_ext_path: str = "", **kwargs):
super().__init__(**kwargs)
self.db_path = db_path
self.vec_ext_path = vec_ext_path
self.conn: sqlite3.Connection | None = None
@property
def vector_table_name(self) -> str:
"""Get the name of the vector table for this store."""
return f"chunks_vec_{self.store_name}"
@property
def fts_table_name(self) -> str:
"""Get the name of the FTS table for this store."""
return f"chunks_fts_{self.store_name}"
@property
def chunks_table_name(self) -> str:
"""Get the name of the chunks table for this store."""
return f"chunks_{self.store_name}"
@property
def files_table_name(self) -> str:
"""Get the name of the files table for this store."""
return f"files_{self.store_name}"
@staticmethod
def vector_to_blob(embedding: list[float]) -> bytes:
"""Convert vector to binary blob for sqlite-vec."""
return struct.pack(f"{len(embedding)}f", *embedding)
async def start(self) -> None:
"""Initialize database and load extensions."""
if self.conn is not None:
return
Path(self.db_path).parent.mkdir(parents=True, exist_ok=True)
self.conn = sqlite3.connect(self.db_path, check_same_thread=False)
self.conn.enable_load_extension(True)
# Load sqlite-vec extension
if self.vec_ext_path:
try:
self.conn.load_extension(self.vec_ext_path)
self.vector_available = True
logger.info(f"Loaded sqlite-vec: {self.vec_ext_path}")
except Exception as e:
logger.warning(f"Failed to load sqlite-vec: {e}")
else:
try:
import sqlite_vec
ext_path = sqlite_vec.loadable_path()
self.conn.load_extension(ext_path)
self.vector_available = True
logger.info(f"Loaded sqlite-vec from package: {ext_path}")
except Exception as e:
logger.warning(f"Failed to load sqlite-vec from package: {e}")
# Fallback: try common extension names
for name in ["vec0", "sqlite_vec", "vector0"]:
try:
self.conn.load_extension(name)
self.vector_available = True
logger.info(f"Loaded sqlite-vec: {name}")
break
except Exception:
pass
self.conn.enable_load_extension(False)
await self._create_tables()
async def _create_tables(self) -> None:
"""Create database schema."""
cursor = self.conn.cursor()
# Files
cursor.execute(
f"""
CREATE TABLE IF NOT EXISTS {self.files_table_name} (
path TEXT,
source TEXT,
hash TEXT,
mtime REAL,
size INTEGER,
PRIMARY KEY (path, source)
)
""",
)
# Chunks
cursor.execute(
f"""
CREATE TABLE IF NOT EXISTS {self.chunks_table_name} (
id TEXT PRIMARY KEY,
path TEXT,
source TEXT,
start_line INTEGER,
end_line INTEGER,
hash TEXT,
text TEXT,
embedding TEXT,
updated_at INTEGER
)
""",
)
# Vector table (sqlite-vec)
if self.vector_available:
cursor.execute(
f"""
CREATE VIRTUAL TABLE IF NOT EXISTS {self.vector_table_name} USING vec0(
id TEXT PRIMARY KEY,
embedding FLOAT[{self.embedding_dim}]
)
""",
)
logger.info(f"Created vector table (dims={self.embedding_dim})")
# FTS table
if self.fts_enabled:
cursor.execute(
f"""
CREATE VIRTUAL TABLE IF NOT EXISTS {self.fts_table_name} USING fts5(
text,
id UNINDEXED,
path UNINDEXED,
source UNINDEXED,
start_line UNINDEXED,
end_line UNINDEXED
)
""",
)
self.fts_available = True
logger.info("Created FTS5 table")
self.conn.commit()
cursor.close()
async def upsert_file(self, file_meta: FileMetadata, source: MemorySource, chunks: list[MemoryChunk]):
"""Insert or update file and its chunks."""
cursor = self.conn.cursor()
try:
cursor.execute("BEGIN")
# Insert file
cursor.execute(
f"""
INSERT OR REPLACE INTO {self.files_table_name} (path, source, hash, mtime, size)
VALUES (?, ?, ?, ?, ?)
""",
(file_meta.path, source.value, file_meta.hash, file_meta.mtime_ms, file_meta.size),
)
# Insert chunks
now = int(time.time() * 1000)
for chunk in chunks:
cursor.execute(
f"""
INSERT OR REPLACE INTO {self.chunks_table_name} (
id, path, source, start_line, end_line,
hash, text, embedding, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
chunk.id,
file_meta.path,
source.value,
chunk.start_line,
chunk.end_line,
chunk.hash,
chunk.text,
json.dumps(chunk.embedding) if chunk.embedding else None,
now,
),
)
# Insert vector
if self.vector_available:
assert chunk.embedding, "Embedding is required for vector insert"
cursor.execute(
f"""
INSERT OR REPLACE INTO {self.vector_table_name} (id, embedding)
VALUES (?, ?)
""",
(chunk.id, self.vector_to_blob(chunk.embedding)),
)
# Insert FTS
if self.fts_available:
cursor.execute(
f"""
INSERT OR REPLACE INTO {self.fts_table_name} (
text, id, path, source, start_line, end_line
) VALUES (?, ?, ?, ?, ?, ?)
""",
(
chunk.text,
chunk.id,
file_meta.path,
source.value,
chunk.start_line,
chunk.end_line,
),
)
cursor.execute("COMMIT")
except Exception:
cursor.execute("ROLLBACK")
raise
finally:
cursor.close()
async def delete_file(self, path: str, source: MemorySource):
"""Delete file and all its chunks."""
cursor = self.conn.cursor()
try:
cursor.execute("BEGIN")
# Get chunk IDs for vector deletion
cursor.execute(
f"SELECT id FROM {self.chunks_table_name} WHERE path = ? AND source = ?",
(path, source.value),
)
chunk_ids = [row[0] for row in cursor.fetchall()]
# Delete vectors
if self.vector_available and chunk_ids:
for chunk_id in chunk_ids:
try:
cursor.execute(
f"DELETE FROM {self.vector_table_name} WHERE id = ?",
(chunk_id,),
)
except Exception as e:
logger.debug(f"Vector delete failed: {e}")
# Delete FTS entries
if self.fts_available:
try:
cursor.execute(
f"DELETE FROM {self.fts_table_name} WHERE path = ? AND source = ?",
(path, source.value),
)
except Exception as e:
logger.debug(f"FTS delete failed: {e}")
# Delete chunks and file
cursor.execute(
f"DELETE FROM {self.chunks_table_name} WHERE path = ? AND source = ?",
(path, source.value),
)
cursor.execute(
f"DELETE FROM {self.files_table_name} WHERE path = ? AND source = ?",
(path, source.value),
)
cursor.execute("COMMIT")
except Exception:
cursor.execute("ROLLBACK")
raise
finally:
cursor.close()
async def delete_file_chunks(self, path: str, chunk_ids: list[str]):
"""Delete specific chunks for a file."""
if not chunk_ids:
return
cursor = self.conn.cursor()
try:
cursor.execute("BEGIN")
# Delete vectors
if self.vector_available:
for chunk_id in chunk_ids:
try:
cursor.execute(
f"DELETE FROM {self.vector_table_name} WHERE id = ?",
(chunk_id,),
)
except Exception as e:
logger.debug(f"Vector delete failed for {chunk_id}: {e}")
# Delete FTS entries
if self.fts_available:
placeholders = ",".join("?" * len(chunk_ids))
try:
cursor.execute(
f"DELETE FROM {self.fts_table_name} WHERE id IN ({placeholders})",
chunk_ids,
)
except Exception as e:
logger.debug(f"FTS delete failed: {e}")
# Delete chunks
placeholders = ",".join("?" * len(chunk_ids))
cursor.execute(
f"DELETE FROM {self.chunks_table_name} WHERE id IN ({placeholders})",
chunk_ids,
)
cursor.execute("COMMIT")
except Exception:
cursor.execute("ROLLBACK")
raise
finally:
cursor.close()
async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource):
"""Insert or update specific chunks without affecting other chunks."""
if not chunks:
return
cursor = self.conn.cursor()
try:
cursor.execute("BEGIN")
now = int(time.time() * 1000)
for chunk in chunks:
# Insert/update chunk
cursor.execute(
f"""
INSERT OR REPLACE INTO {self.chunks_table_name} (
id, path, source, start_line, end_line,
hash, text, embedding, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
chunk.id,
chunk.path,
source.value,
chunk.start_line,
chunk.end_line,
chunk.hash,
chunk.text,
json.dumps(chunk.embedding) if chunk.embedding else None,
now,
),
)
# Insert/update vector
if self.vector_available:
assert chunk.embedding, "Embedding is required for vector insert"
cursor.execute(
f"""
INSERT OR REPLACE INTO {self.vector_table_name} (id, embedding)
VALUES (?, ?)
""",
(chunk.id, self.vector_to_blob(chunk.embedding)),
)
# Insert/update FTS
if self.fts_available:
cursor.execute(
f"""
INSERT OR REPLACE INTO {self.fts_table_name} (
text, id, path, source, start_line, end_line
) VALUES (?, ?, ?, ?, ?, ?)
""",
(
chunk.text,
chunk.id,
chunk.path,
source.value,
chunk.start_line,
chunk.end_line,
),
)
cursor.execute("COMMIT")
except Exception:
cursor.execute("ROLLBACK")
raise
finally:
cursor.close()
async def list_files(self, source: MemorySource) -> list[str]:
"""List all indexed files."""
cursor = self.conn.cursor()
cursor.execute(f"SELECT path FROM {self.files_table_name} WHERE source = ?", (source.value,))
paths = [row[0] for row in cursor.fetchall()]
cursor.close()
return paths
async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None:
"""Get file metadata with chunk count."""
cursor = self.conn.cursor()
cursor.execute(
f"SELECT hash, mtime, size FROM {self.files_table_name} WHERE path = ? AND source = ?",
(path, source.value),
)
row = cursor.fetchone()
if not row:
cursor.close()
return None
hash_val, mtime, size = row
cursor.execute(
f"SELECT COUNT(*) FROM {self.chunks_table_name} WHERE path = ? AND source = ?",
(path, source.value),
)
chunk_count = cursor.fetchone()[0]
cursor.close()
return FileMetadata(
hash=hash_val,
mtime_ms=mtime,
size=size,
path=path,
chunk_count=chunk_count,
)
async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
"""Get all chunks for a file."""
cursor = self.conn.cursor()
cursor.execute(
f"""
SELECT id, path, source, start_line, end_line, text, hash, embedding
FROM {self.chunks_table_name} WHERE path = ? AND source = ?
ORDER BY start_line
""",
(path, source.value),
)
chunks = []
for row in cursor.fetchall():
chunk_id, path_val, source_val, start, end, text, hash_val, emb_str = row
# Parse embedding from JSON string
embedding = None
if emb_str:
try:
embedding = json.loads(emb_str)
except (json.JSONDecodeError, TypeError):
embedding = None
chunks.append(
MemoryChunk(
id=chunk_id,
path=path_val,
source=MemorySource(source_val),
start_line=start,
end_line=end,
text=text,
hash=hash_val,
embedding=embedding,
),
)
cursor.close()
return chunks
async def vector_search(
self,
query: str,
limit: int,
sources: list[MemorySource] | None = None,
) -> list[MemorySearchResult]:
"""Perform vector similarity search."""
if not self.vector_available or not query:
return []
# Get query embedding
query_embedding = await self.get_embedding(query)
if not query_embedding:
return []
cursor = self.conn.cursor()
source_filter = ""
params: list = []
if sources:
placeholders = ",".join("?" * len(sources))
source_filter = f" AND c.source IN ({placeholders})"
params = [s.value for s in sources]
try:
query_blob = self.vector_to_blob(query_embedding)
# Correct SQLite-vec syntax for vector search with limit
# vec0 requires 'k = ?' constraint for knn queries
query_sql = f"""
SELECT c.id, c.path, c.start_line, c.end_line, c.source, c.text, v.distance
FROM {self.vector_table_name} v
JOIN {self.chunks_table_name} c ON v.id = c.id
WHERE v.embedding MATCH ?
AND k = ?
"""
query_params: list = [query_blob, limit]
# Add source filter if specified
if source_filter:
query_sql += source_filter
query_params.extend(params)
# Order by distance (k constraint already limits results)
query_sql += " ORDER BY v.distance"
cursor.execute(query_sql, query_params)
results = []
for _, path, start, end, src, text, dist in cursor.fetchall():
score = max(0.0, 1.0 - dist)
snippet = text[: self.snippet_max_chars] if len(text) > self.snippet_max_chars else text
results.append(
MemorySearchResult(
path=path,
start_line=start,
end_line=end,
score=score,
snippet=snippet,
source=MemorySource(src),
),
)
return results
except Exception as e:
logger.error(f"Vector search failed: {e}")
return []
finally:
cursor.close()
async def keyword_search(
self,
query: str,
limit: int,
sources: list[MemorySource] | None = None,
) -> list[MemorySearchResult]:
"""Perform full-text search."""
if not self.fts_available:
return []
# Build FTS5 query
# Split query into tokens and join with OR for better recall
# Individual words are automatically stemmed and matched by FTS5
cleaned = query.strip()
if not cleaned:
return []
# Split into words and escape each
words = cleaned.split()
if not words:
return []
# Use OR operator for better recall - match any of the query words
escaped_words = [word.replace('"', '""') for word in words]
fts_query = " OR ".join(escaped_words)
cursor = self.conn.cursor()
source_filter = ""
params: list = [fts_query]
if sources:
placeholders = ",".join("?" * len(sources))
source_filter = f" AND fts.source IN ({placeholders})"
params.extend([s.value for s in sources])
params.append(limit)
try:
cursor.execute(
f"""
SELECT fts.id, fts.path, fts.start_line, fts.end_line,
fts.source, fts.text, rank
FROM {self.fts_table_name} fts
WHERE fts.text MATCH ?{source_filter}
ORDER BY rank
LIMIT ?
""",
params,
)
results = []
for _, path, start, end, src, text, rank in cursor.fetchall():
# Convert BM25 rank (negative) to 0-1 score (higher=better)
score = max(0.0, 1.0 / (1.0 + abs(rank)))
snippet = text[: self.snippet_max_chars] if len(text) > self.snippet_max_chars else text
results.append(
MemorySearchResult(
path=path,
start_line=start,
end_line=end,
score=score,
snippet=snippet,
source=MemorySource(src),
),
)
return results
except Exception as e:
logger.error(f"Keyword search failed: {e}")
return []
finally:
cursor.close()
async def clear_all(self):
"""Clear all indexed data."""
cursor = self.conn.cursor()
cursor.execute("BEGIN")
try:
cursor.execute(f"DELETE FROM {self.files_table_name}")
cursor.execute(f"DELETE FROM {self.chunks_table_name}")
if self.vector_available:
cursor.execute(f"DELETE FROM {self.vector_table_name}")
if self.fts_available:
cursor.execute(f"DELETE FROM {self.fts_table_name}")
cursor.execute("COMMIT")
except Exception:
cursor.execute("ROLLBACK")
raise
finally:
cursor.close()
async def close(self):
"""Close database connection."""
if self.conn:
self.conn.close()
self.conn = None

View file

@ -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."""

View file

@ -1,6 +1,10 @@
"""schema"""
from .file_metadata import FileMetadata
from .memory_chunk import MemoryChunk
from .memory_index_meta import MemoryIndexMeta
from .memory_node import MemoryNode
from .memory_search_result import MemorySearchResult
from .message import ContentBlock, Message, Trajectory
from .request import Request
from .response import Response
@ -17,16 +21,22 @@ from .service_config import (
)
from .stream_chunk import StreamChunk
from .tool_call import ToolAttr, ToolCall
from .truncation_result import TruncationResult
from .vector_node import VectorNode
__all__ = [
"MemoryNode",
"CmdConfig",
"ContentBlock",
"EmbeddingModelConfig",
"FileMetadata",
"FlowConfig",
"HttpConfig",
"LLMConfig",
"MCPConfig",
"MemoryChunk",
"MemoryIndexMeta",
"MemoryNode",
"MemorySearchResult",
"Message",
"Request",
"Response",
@ -36,7 +46,7 @@ __all__ = [
"Trajectory",
"ToolAttr",
"ToolCall",
"TruncationResult",
"VectorNode",
"VectorStoreConfig",
"CmdConfig",
]

View file

@ -0,0 +1,15 @@
"""File metadata schema."""
from pydantic import BaseModel, Field
class FileMetadata(BaseModel):
"""File metadata with optional extended fields for various use cases."""
hash: str = Field(default=..., description="Hash of the file content")
mtime_ms: float = Field(default=..., description="Last modification time in milliseconds")
size: int = Field(default=..., description="File size in bytes")
path: str | None = Field(default=None, description="Relative path to the session file")
content: str | None = Field(default=None, description="Parsed content from the session file")
chunk_count: int | None = Field(default=None, description="Number of chunks in the file")
metadata: dict = Field(default_factory=dict, description="Additional metadata")

View file

@ -0,0 +1,19 @@
"""Memory chunk schema."""
from pydantic import BaseModel, Field
from ..enumeration import MemorySource
class MemoryChunk(BaseModel):
"""A chunk of memory content with metadata."""
id: str = Field(..., description="Unique identifier for the chunk")
path: str = Field(..., description="File path relative to workspace")
source: MemorySource = Field(..., description="Source of the memory data")
start_line: int = Field(..., description="Starting line number in the source file")
end_line: int = Field(..., description="Ending line number in the source file")
text: str = Field(..., description="Text content of the chunk")
hash: str = Field(..., description="Hash of the chunk content")
embedding: list[float] | None = Field(default=None, description="Vector embedding of the chunk")
metadata: dict = Field(default_factory=dict, description="Additional metadata")

View file

@ -0,0 +1,14 @@
"""Memory index metadata schema."""
from typing import Optional
from pydantic import BaseModel, Field
class MemoryIndexMeta(BaseModel):
"""Metadata for memory index configuration."""
model: str = Field(..., description="Name of the embedding model")
chunk_tokens: int = Field(..., description="Maximum tokens per chunk")
chunk_overlap: int = Field(..., description="Number of overlapping tokens between chunks")
vector_dims: Optional[int] = Field(default=None, description="Vector embedding dimensions")

View file

@ -0,0 +1,24 @@
"""Memory search result schema."""
from typing import Any, Dict
from pydantic import BaseModel, Field
from ..enumeration import MemorySource
class MemorySearchResult(BaseModel):
"""Search result from memory index."""
path: str = Field(..., description="File path relative to workspace")
start_line: int = Field(..., description="Starting line number of the match")
end_line: int = Field(..., description="Ending line number of the match")
score: float = Field(..., description="Relevance score of the search result")
snippet: str = Field(..., description="Text snippet from the matched content")
source: MemorySource = Field(..., description="Source of the memory data")
metadata: Dict[str, Any] = Field(default_factory=dict, description="Additional metadata")
@property
def merge_key(self) -> str:
"""Merge key for the search result."""
return self.path + f":{self.start_line}:{self.end_line}"

View file

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

View file

@ -0,0 +1,35 @@
"""Truncation result schema for command output truncation."""
from typing import Literal
from pydantic import BaseModel, Field
class TruncationResult(BaseModel):
"""Result of output truncation operation.
Attributes:
content: The truncated content
truncated: Whether truncation occurred
total_lines: Total number of lines in original output
output_lines: Number of lines in truncated output
total_bytes: Total bytes in original output
output_bytes: Bytes in truncated output
truncated_by: What caused truncation ('lines' or 'bytes')
last_line_partial: Whether last line was partially truncated
"""
content: str = Field(description="The truncated content")
truncated: bool = Field(description="Whether truncation occurred")
total_lines: int = Field(description="Total number of lines in original output")
output_lines: int = Field(description="Number of lines in truncated output")
total_bytes: int = Field(description="Total bytes in original output")
output_bytes: int = Field(description="Bytes in truncated output")
truncated_by: Literal["lines", "bytes"] | None = Field(
default=None,
description="What caused truncation ('lines' or 'bytes')",
)
last_line_partial: bool = Field(
default=False,
description="Whether last line was partially truncated",
)

View file

@ -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",

View file

@ -0,0 +1,124 @@
"""Chunking logic for Markdown files."""
from .common_utils import hash_text
from ..enumeration import MemorySource
from ..schema import MemoryChunk
def chunk_markdown(
text: str,
path: str,
source: MemorySource,
chunk_tokens: int,
overlap: int,
) -> list[MemoryChunk]:
"""
Markdown chunking logic implemented based on the TypeScript version.
Args:
text: Input text
path: File path
source: Memory source
chunk_tokens: Maximum tokens per chunk
overlap: Overlap tokens between chunks
Returns:
List of MemoryChunk objects
"""
lines = text.split("\n")
if not lines:
return []
# Convert tokens to characters (~1 token = 4 chars)
max_chars = max(32, chunk_tokens * 4)
overlap_chars = max(0, overlap * 4)
chunks: list[MemoryChunk] = []
# Currently building chunk
current: list[dict] = [] # [{'line': str, 'line_no': int}]
current_chars = 0
def flush():
"""Add current chunk to results list"""
if not current:
return
first_entry = current[0]
last_entry = current[-1]
if not first_entry or not last_entry:
return
chunk_text = "\n".join([entry["line"] for entry in current])
start_line = first_entry["line_no"]
end_line = last_entry["line_no"]
chunk_hash = hash_text(chunk_text)
chunks.append(
MemoryChunk(
id=hash_text(f"{source}:{path}:{start_line}:{end_line}:{chunk_hash}:{len(chunks)}"),
path=path,
source=source,
start_line=start_line,
end_line=end_line,
text=chunk_text,
hash=chunk_hash,
),
)
def carry_overlap():
"""Keep overlapping part and clear the rest"""
nonlocal current, current_chars
if overlap_chars <= 0 or not current:
current = []
current_chars = 0
return
acc = 0
kept = []
# Collect lines from the end until reaching overlap size
for j in range(len(current) - 1, -1, -1):
entry = current[j]
if not entry:
continue
acc += len(entry["line"]) + 1 # +1 for newline
kept.insert(0, entry) # Insert at the beginning to maintain order
if acc >= overlap_chars:
break
current = kept
current_chars = sum(len(entry["line"]) + 1 for entry in kept)
for i, line in enumerate(lines):
line_no = i + 1
# Split long lines into multiple segments
segments = []
if not line: # Empty line
segments.append("")
else:
# If line is too long, split by maximum character count
for start in range(0, len(line), max_chars):
segments.append(line[start : start + max_chars])
for segment in segments:
line_size = len(segment) + 1 # +1 for newline
# If adding current segment would exceed the limit, flush current chunk
if current_chars + line_size > max_chars and current:
flush()
carry_overlap()
current.append({"line": segment, "line_no": line_no})
current_chars += line_size
# Process the final chunk
flush()
return [c for c in chunks if c.text.strip()]

View file

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

View file

@ -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}")

View file

@ -20,7 +20,7 @@ from .agent.memory import (
)
from .config import ReMeConfigParser
from .core import Application
from .core.enumeration import MemoryType
from .core.enumeration import MemoryType, Role
from .core.schema import Message, MemoryNode
from .tool.memory import (
RetrieveMemory,
@ -59,7 +59,7 @@ class ReMe(Application):
target_user_names: list[str] | None = None,
target_task_names: list[str] | None = None,
target_tool_names: list[str] | None = None,
profile_dir: str = "reme_profile",
profile_dir: str = ".reme/profile",
**kwargs,
):
"""Initialize ReMe with config.
@ -83,11 +83,14 @@ class ReMe(Application):
Example:
```python
reme = await ReMe(...).start()
reme = ReMe(...)
await reme.start()
# reme = await ReMe.create(...) # both ok
await reme.summarize_memory(...)
await reme.retrieve_memory(...)
await reme.close()
```
"""
@ -308,7 +311,8 @@ class ReMe(Application):
if user_name:
if isinstance(user_name, str):
for message in format_messages:
message.name = user_name
if message.role is Role.USER:
message.name = user_name
self._add_meta_memory(MemoryType.PERSONAL, user_name)
memory_targets.append(user_name)
elif isinstance(user_name, list):

24
reme/tool/fs/__init__.py Normal file
View file

@ -0,0 +1,24 @@
"""File system tools."""
from .bash_tool import BashTool
from .edit_tool import EditTool
from .find_tool import FindTool
from .grep_tool import GrepTool
from .ls_tool import LsTool
from .read_tool import ReadTool
from .write_tool import WriteTool
from ...core import R
__all__ = [
"BashTool",
"EditTool",
"FindTool",
"GrepTool",
"LsTool",
"ReadTool",
"WriteTool",
]
for name in __all__:
tool_class = globals()[name]
R.op.register(tool_class)

191
reme/tool/fs/bash_tool.py Normal file
View file

@ -0,0 +1,191 @@
"""Bash command execution tool with production-grade features.
This module provides a production-grade tool for executing bash commands with:
- Smart output truncation (keeps last N lines/bytes to prevent memory issues)
- Process tree termination (prevents orphan processes)
"""
import asyncio
import os
import platform
import signal
from pathlib import Path
from .truncate import DEFAULT_MAX_BYTES, DEFAULT_MAX_LINES, truncate_tail
from ...core.op import BaseTool
from ...core.schema import ToolCall, TruncationResult
def get_shell_config() -> tuple[str, list[str]]:
"""Get the appropriate shell and arguments for the current platform.
Returns:
Tuple of (shell_path, args) for subprocess execution
"""
system = platform.system()
if system == "Windows":
# Use PowerShell on Windows
return "powershell.exe", ["-Command"]
else:
# Use bash on Unix-like systems
shell = os.environ.get("SHELL", "/bin/bash")
return shell, ["-c"]
def kill_process_tree(pid: int) -> None:
"""Kill a process and all its children.
Args:
pid: Process ID to kill
"""
try:
if platform.system() == "Windows":
# Windows: use taskkill
os.system(f"taskkill /F /T /PID {pid}")
else:
# Unix: kill process group
try:
os.killpg(os.getpgid(pid), signal.SIGTERM)
except ProcessLookupError:
pass # Process already dead
except Exception:
pass # Best effort
class BashTool(BaseTool):
"""Production-grade tool for executing bash commands.
Features:
- Smart output truncation (preserves last N lines or M bytes)
- Kills entire process tree on timeout (prevents orphan processes)
"""
def __init__(self, cwd: str | None = None, command_prefix: str | None = None):
"""Initialize bash tool.
Args:
cwd: Working directory (defaults to current directory)
command_prefix: Optional prefix prepended to every command
"""
super().__init__()
self.cwd = cwd or os.getcwd()
self.command_prefix = command_prefix
def _build_tool_call(self) -> ToolCall:
max_kb = DEFAULT_MAX_BYTES // 1024
return ToolCall(
**{
"description": (
f"Execute a bash command in the current working directory. "
f"Returns stdout and stderr. Output is truncated to last "
f"{DEFAULT_MAX_LINES} lines or {max_kb}KB (whichever is hit first). "
f"Optionally provide a timeout in seconds."
),
"parameters": {
"type": "object",
"properties": {
"command": {
"type": "string",
"description": "Bash command to execute",
},
"timeout": {
"type": "number",
"description": "Timeout in seconds (optional, no default timeout)",
},
},
"required": ["command"],
},
},
)
async def execute(self) -> str:
"""Execute the bash command with production-grade features."""
command: str = self.context.command
timeout: float | None = self.context.get("timeout", None)
# Apply command prefix if configured
if self.command_prefix:
command = f"{self.command_prefix}\n{command}"
# Verify working directory exists
if not Path(self.cwd).exists():
raise FileNotFoundError(
f"Working directory does not exist: {self.cwd}\n" f"Cannot execute bash commands.",
)
# Get shell configuration
shell, shell_args = get_shell_config()
# Start process
try:
process = await asyncio.create_subprocess_exec(
shell,
*shell_args,
command,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
cwd=self.cwd,
# Create process group for clean termination
preexec_fn=os.setpgrp if platform.system() != "Windows" else None,
)
except Exception as e:
raise RuntimeError(f"Failed to start process: {e}") from e
# Execute command with optional timeout
try:
if timeout and timeout > 0:
try:
stdout, stderr = await asyncio.wait_for(
process.communicate(),
timeout=timeout,
)
except asyncio.TimeoutError as e:
# Kill process tree on timeout
if process.pid:
kill_process_tree(process.pid)
try:
await asyncio.wait_for(process.wait(), timeout=1.0)
except asyncio.TimeoutError:
process.kill()
raise TimeoutError(f"Command timed out after {timeout} seconds") from e
else:
stdout, stderr = await process.communicate()
except TimeoutError as e:
raise RuntimeError(str(e)) from e
# Decode output
full_output = stdout.decode("utf-8", errors="ignore")
if stderr:
stderr_text = stderr.decode("utf-8", errors="ignore")
if full_output:
full_output += "\n"
full_output += stderr_text
# Apply tail truncation_result to prevent memory issues
truncation_result: TruncationResult = truncate_tail(full_output)
output_text = truncation_result.content or "(no output)"
# Build truncation_result notice if needed
if truncation_result.truncated:
start_line = truncation_result.total_lines - truncation_result.output_lines + 1
end_line = truncation_result.total_lines
if truncation_result.truncated_by == "lines":
output_text += (
f"\n\n[Output truncated: showing lines {start_line}-{end_line} "
f"of {truncation_result.total_lines} total lines]"
)
else:
max_kb = DEFAULT_MAX_BYTES // 1024
output_text += (
f"\n\n[Output truncated: showing lines {start_line}-{end_line} "
f"of {truncation_result.total_lines} ({max_kb}KB limit reached)]"
)
# Handle non-zero exit code
if process.returncode != 0:
output_text += f"\n\nCommand exited with code {process.returncode}"
raise RuntimeError(output_text)
return output_text

164
reme/tool/fs/edit_diff.py Normal file
View file

@ -0,0 +1,164 @@
"""Diff utilities for edit tool."""
import re
from dataclasses import dataclass
from difflib import unified_diff
def detect_line_ending(content: str) -> str:
"""Detect line ending style (CRLF or LF)."""
crlf_idx = content.find("\r\n")
lf_idx = content.find("\n")
if lf_idx == -1:
return "\n"
if crlf_idx == -1:
return "\n"
return "\r\n" if crlf_idx < lf_idx else "\n"
def normalize_to_lf(text: str) -> str:
"""Normalize line endings to LF."""
return text.replace("\r\n", "\n").replace("\r", "\n")
def restore_line_endings(text: str, ending: str) -> str:
"""Restore original line endings."""
return text.replace("\n", ending) if ending == "\r\n" else text
def normalize_for_fuzzy_match(text: str) -> str:
"""Normalize text for fuzzy matching: strip trailing whitespace, normalize quotes/dashes."""
lines = text.split("\n")
normalized = "\n".join(line.rstrip() for line in lines)
# Smart quotes → ASCII
normalized = re.sub(r"[\u2018\u2019\u201A\u201B]", "'", normalized)
normalized = re.sub(r"[\u201C\u201D\u201E\u201F]", '"', normalized)
# Dashes → hyphen
normalized = re.sub(r"[\u2010\u2011\u2012\u2013\u2014\u2015\u2212]", "-", normalized)
# Special spaces → regular space
normalized = re.sub(r"[\u00A0\u2002-\u200A\u202F\u205F\u3000]", " ", normalized)
return normalized
@dataclass
class FuzzyMatchResult:
"""Result of fuzzy text matching."""
found: bool
index: int
match_length: int
used_fuzzy_match: bool
content_for_replacement: str
def fuzzy_find_text(content: str, old_text: str) -> FuzzyMatchResult:
"""Find old_text in content, trying exact match first, then fuzzy match."""
# Try exact match
exact_index = content.find(old_text)
if exact_index != -1:
return FuzzyMatchResult(
found=True,
index=exact_index,
match_length=len(old_text),
used_fuzzy_match=False,
content_for_replacement=content,
)
# Try fuzzy match
fuzzy_content = normalize_for_fuzzy_match(content)
fuzzy_old_text = normalize_for_fuzzy_match(old_text)
fuzzy_index = fuzzy_content.find(fuzzy_old_text)
if fuzzy_index == -1:
return FuzzyMatchResult(
found=False,
index=-1,
match_length=0,
used_fuzzy_match=False,
content_for_replacement=content,
)
return FuzzyMatchResult(
found=True,
index=fuzzy_index,
match_length=len(fuzzy_old_text),
used_fuzzy_match=True,
content_for_replacement=fuzzy_content,
)
def strip_bom(content: str) -> tuple[str, str]:
"""Strip UTF-8 BOM, return (bom, text_without_bom)."""
if content.startswith("\ufeff"):
return "\ufeff", content[1:]
return "", content
@dataclass
class DiffResult:
"""Result of diff generation."""
diff: str
first_changed_line: int | None
def generate_diff_string(old_content: str, new_content: str, context_lines: int = 4) -> DiffResult:
"""Generate unified diff with line numbers."""
old_lines = old_content.split("\n")
new_lines = new_content.split("\n")
# Use difflib to get the changes
diff_lines = list(
unified_diff(
old_lines,
new_lines,
lineterm="",
n=context_lines,
),
)
if not diff_lines:
return DiffResult(diff="", first_changed_line=None)
# Parse and format the diff
output = []
first_changed_line = None
max_line_num = max(len(old_lines), len(new_lines))
line_num_width = len(str(max_line_num))
old_line_num = 1
new_line_num = 1
for line in diff_lines[2:]: # Skip header lines
if line.startswith("@@"):
# Parse hunk header
match = re.match(r"@@ -(\d+),?\d* \+(\d+),?\d* @@", line)
if match:
old_line_num = int(match.group(1))
new_line_num = int(match.group(2))
continue
if line.startswith("+"):
if first_changed_line is None:
first_changed_line = new_line_num
line_num = str(new_line_num).rjust(line_num_width)
output.append(f"+{line_num} {line[1:]}")
new_line_num += 1
elif line.startswith("-"):
if first_changed_line is None:
first_changed_line = new_line_num
line_num = str(old_line_num).rjust(line_num_width)
output.append(f"-{line_num} {line[1:]}")
old_line_num += 1
else:
# Context line
line_num = str(old_line_num).rjust(line_num_width)
output.append(f" {line_num} {line[1:] if line.startswith(' ') else line}")
old_line_num += 1
new_line_num += 1
return DiffResult(diff="\n".join(output), first_changed_line=first_changed_line)

141
reme/tool/fs/edit_tool.py Normal file
View file

@ -0,0 +1,141 @@
"""File editing tool with exact text replacement."""
import os
from pathlib import Path
from .edit_diff import (
detect_line_ending,
fuzzy_find_text,
generate_diff_string,
normalize_for_fuzzy_match,
normalize_to_lf,
restore_line_endings,
strip_bom,
)
from ...core.op import BaseTool
from ...core.schema import ToolCall
class EditTool(BaseTool):
"""Edit a file by replacing exact text."""
def __init__(self, cwd: str | None = None):
"""Initialize edit tool.
Args:
cwd: Working directory (defaults to current directory)
"""
super().__init__()
self.cwd = cwd or os.getcwd()
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": (
"Edit a file by replacing exact text. The oldText must match exactly "
"(including whitespace). Use this for precise, surgical edits."
),
"parameters": {
"type": "object",
"properties": {
"path": {
"type": "string",
"description": "Path to the file to edit (relative or absolute)",
},
"oldText": {
"type": "string",
"description": "Exact text to find and replace (must match exactly)",
},
"newText": {
"type": "string",
"description": "New text to replace the old text with",
},
},
"required": ["path", "oldText", "newText"],
},
},
)
async def execute(self) -> str:
"""Execute the edit operation."""
path: str = self.context.path
old_text: str = self.context.oldText
new_text: str = self.context.newText
# Resolve path
if not os.path.isabs(path):
absolute_path = os.path.join(self.cwd, path)
else:
absolute_path = path
# Check file exists and is writable
path_obj = Path(absolute_path)
if not path_obj.exists():
raise FileNotFoundError(f"File not found: {path}")
if not os.access(absolute_path, os.R_OK | os.W_OK):
raise PermissionError(f"File not readable/writable: {path}")
# Read file
try:
with open(absolute_path, "r", encoding="utf-8") as f:
raw_content = f.read()
except Exception as e:
raise IOError(f"Failed to read file {path}: {e}") from e
# Strip BOM (LLM won't include invisible BOM in oldText)
bom, content = strip_bom(raw_content)
original_ending = detect_line_ending(content)
normalized_content = normalize_to_lf(content)
normalized_old_text = normalize_to_lf(old_text)
normalized_new_text = normalize_to_lf(new_text)
# Find old text using fuzzy matching
match_result = fuzzy_find_text(normalized_content, normalized_old_text)
if not match_result.found:
raise ValueError(
f"Could not find the exact text in {path}. The old text must match "
f"exactly including all whitespace and newlines.",
)
# Count occurrences for uniqueness check
fuzzy_content = normalize_for_fuzzy_match(normalized_content)
fuzzy_old_text = normalize_for_fuzzy_match(normalized_old_text)
occurrences = fuzzy_content.count(fuzzy_old_text)
if occurrences > 1:
raise ValueError(
f"Found {occurrences} occurrences of the text in {path}. "
f"The text must be unique. Please provide more context to make it unique.",
)
# Perform replacement
base_content = match_result.content_for_replacement
new_content = (
base_content[: match_result.index]
+ normalized_new_text
+ base_content[match_result.index + match_result.match_length :]
)
# Verify replacement changed something
if base_content == new_content:
raise ValueError(
f"No changes made to {path}. The replacement produced identical content. "
f"This might indicate an issue with special characters or the text not "
f"exist as expected.",
)
# Write file
final_content = bom + restore_line_endings(new_content, original_ending)
try:
with open(absolute_path, "w", encoding="utf-8") as f:
f.write(final_content)
except Exception as e:
raise IOError(f"Failed to write file {path}: {e}") from e
# Generate diff
diff_result = generate_diff_string(base_content, new_content)
return f"Successfully replaced text in {path}.\n\n{diff_result.diff}"

183
reme/tool/fs/find_tool.py Normal file
View file

@ -0,0 +1,183 @@
"""File search tool using glob patterns with gitignore support."""
import os
from pathlib import Path
from .truncate import FIND_MAX_BYTES, FIND_MAX_LINES, format_size, truncate_head
from ...core.op import BaseTool
from ...core.schema import ToolCall
class FindTool(BaseTool):
"""Search for files by glob pattern, respecting .gitignore."""
def __init__(self, cwd: str | None = None):
"""Initialize find tool.
Args:
cwd: Working directory (defaults to current directory)
"""
super().__init__()
self.cwd = cwd or os.getcwd()
def _build_tool_call(self) -> ToolCall:
max_kb = FIND_MAX_BYTES // 1024
return ToolCall(
**{
"description": (
f"Search for files by glob pattern. Returns matching file paths relative "
f"to the search directory. Respects .gitignore. Output is truncated to "
f"1000 results or {max_kb}KB (whichever is hit first)."
),
"parameters": {
"type": "object",
"properties": {
"pattern": {
"type": "string",
"description": "Glob pattern to match files, "
"e.g. '*.ts', '**/*.json', or 'src/**/*.spec.ts'",
},
"path": {
"type": "string",
"description": "Directory to search in (default: current directory)",
},
"limit": {
"type": "number",
"description": "Maximum number of results (default: 1000)",
},
},
"required": ["pattern"],
},
},
)
def _load_gitignore_patterns(self, search_path: Path) -> list[str]:
"""Load gitignore patterns from directory and subdirectories."""
patterns = ["**/node_modules/**", "**/.git/**"]
# Load root .gitignore
gitignore_path = search_path / ".gitignore"
if gitignore_path.exists():
try:
with open(gitignore_path, "r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if line and not line.startswith("#"):
patterns.append(line)
except Exception:
pass # Ignore errors
# Load nested .gitignore files
try:
for gitignore in search_path.rglob(".gitignore"):
if gitignore == gitignore_path:
continue
try:
with open(gitignore, "r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if line and not line.startswith("#"):
patterns.append(line)
except Exception:
pass # Ignore errors
except Exception:
pass # Ignore glob errors
return patterns
def _should_ignore(self, path: Path, ignore_patterns: list[str]) -> bool:
"""Check if path matches any ignore pattern."""
path_str = str(path)
for pattern in ignore_patterns:
# Simple pattern matching (not full gitignore spec)
if "**" in pattern:
# Recursive match
clean_pattern = pattern.replace("**/", "").replace("/**", "")
if clean_pattern in path_str:
return True
elif "*" in pattern:
# Wildcard match
from fnmatch import fnmatch
if fnmatch(path.name, pattern):
return True
elif pattern in path_str:
return True
return False
async def execute(self) -> str:
"""Execute file search."""
pattern: str = self.context.pattern
search_dir: str = self.context.get("path", ".")
limit: int = self.context.get("limit", 1000)
# Resolve search path
if not os.path.isabs(search_dir):
search_path = Path(self.cwd) / search_dir
else:
search_path = Path(search_dir)
# Check if directory exists
if not search_path.exists():
raise FileNotFoundError(f"Path not found: {search_dir}")
if not search_path.is_dir():
raise NotADirectoryError(f"Path is not a directory: {search_dir}")
# Load gitignore patterns
ignore_patterns = self._load_gitignore_patterns(search_path)
# Search for files
results = []
try:
for file_path in search_path.glob(pattern):
if len(results) >= limit:
break
# Skip if matches ignore patterns
if self._should_ignore(file_path, ignore_patterns):
continue
# Get relative path
try:
rel_path = file_path.relative_to(search_path)
# Add trailing slash for directories
if file_path.is_dir():
results.append(f"{rel_path}/")
else:
results.append(str(rel_path))
except ValueError:
# If relative_to fails, use the path as-is
results.append(str(file_path))
except Exception as e:
raise RuntimeError(f"Error searching for files: {e}") from e
# Handle no results
if not results:
return "No files found matching pattern"
# Sort results for consistency
results.sort()
# Apply limit and truncation
result_limit_reached = len(results) >= limit
raw_output = "\n".join(results)
truncation = truncate_head(raw_output, max_lines=FIND_MAX_LINES, max_bytes=FIND_MAX_BYTES)
output = truncation.content
notices = []
if result_limit_reached:
notices.append(
f"{limit} results limit reached. Use limit={limit * 2} for more, or refine pattern",
)
if truncation.truncated:
notices.append(f"{format_size(FIND_MAX_BYTES)} limit reached")
if notices:
output += f"\n\n[{'. '.join(notices)}]"
return output

View file

@ -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

View file

@ -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

276
reme/tool/fs/grep_tool.py Normal file
View file

@ -0,0 +1,276 @@
"""Grep tool for searching file contents using ripgrep.
This module provides a tool for searching file contents with:
- Pattern matching (regex or literal string)
- Smart output truncation (prevents memory issues)
- Context lines support
- Respects .gitignore
"""
import asyncio
import json
import os
import shutil
from pathlib import Path
from .truncate import (
DEFAULT_MAX_BYTES,
GREP_MAX_LINE_LENGTH,
format_size,
truncate_head,
truncate_line,
)
from ...core.op import BaseTool
from ...core.schema import ToolCall
# Default limits
DEFAULT_LIMIT = 100 # Maximum number of matches
class GrepTool(BaseTool):
"""Tool for searching file contents using ripgrep.
Features:
- Pattern matching with regex or literal string
- Context lines support
- Smart output truncation
- Respects .gitignore
"""
def __init__(self, cwd: str | None = None):
"""Initialize grep tool.
Args:
cwd: Working directory (defaults to current directory)
"""
super().__init__()
self.cwd = cwd or os.getcwd()
def _build_tool_call(self) -> ToolCall:
max_kb = DEFAULT_MAX_BYTES // 1024
return ToolCall(
**{
"description": (
f"Search file contents for a pattern. Returns matching lines with "
f"file paths and line numbers. Respects .gitignore. Output is "
f"truncated to {DEFAULT_LIMIT} matches or {max_kb}KB (whichever is "
f"hit first). Long lines are truncated to {GREP_MAX_LINE_LENGTH} chars."
),
"parameters": {
"type": "object",
"properties": {
"pattern": {
"type": "string",
"description": "Search pattern (regex or literal string)",
},
"path": {
"type": "string",
"description": "Directory or file to search (default: current directory)",
},
"glob": {
"type": "string",
"description": "Filter files by glob pattern, e.g. '*.ts' or '**/*.spec.ts'",
},
"ignoreCase": {
"type": "boolean",
"description": "Case-insensitive search (default: false)",
},
"literal": {
"type": "boolean",
"description": "Treat pattern as literal string instead of regex (default: false)",
},
"contextLines": {
"type": "number",
"description": "Number of lines to show before and after each match (default: 0)",
},
"limit": {
"type": "number",
"description": f"Maximum number of matches to return (default: {DEFAULT_LIMIT})",
},
},
"required": ["pattern"],
},
},
)
async def execute(self) -> str:
"""Execute the grep search."""
pattern: str = self.context.pattern
search_path: str = self.context.get("path", ".")
glob: str | None = self.context.get("glob", None)
ignore_case: bool = self.context.get("ignoreCase", False)
literal: bool = self.context.get("literal", False)
context_lines: int = self.context.get("contextLines", 0)
limit: int = self.context.get("limit", DEFAULT_LIMIT)
# Check if ripgrep is available
rg_path = shutil.which("rg")
if not rg_path:
raise RuntimeError(
"ripgrep (rg) is not available. Please install it:\n"
" macOS: brew install ripgrep\n"
" Ubuntu: apt-get install ripgrep\n"
" Other: https://github.com/BurntSushi/ripgrep",
)
# Resolve search path
if not os.path.isabs(search_path):
search_path = os.path.join(self.cwd, search_path)
# Check if path exists
if not Path(search_path).exists():
raise FileNotFoundError(f"Path not found: {search_path}")
is_directory = Path(search_path).is_dir()
effective_limit = max(1, limit)
# Build ripgrep arguments
args = [
rg_path,
"--json",
"--line-number",
"--color=never",
"--hidden",
]
if ignore_case:
args.append("--ignore-case")
if literal:
args.append("--fixed-strings")
if glob:
args.extend(["--glob", glob])
args.extend([pattern, search_path])
# Execute ripgrep
try:
process = await asyncio.create_subprocess_exec(
*args,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
cwd=self.cwd,
)
except Exception as e:
raise RuntimeError(f"Failed to run ripgrep: {e}") from e
stdout, stderr = await process.communicate()
# Parse JSON output
matches = []
match_count = 0
lines_truncated = False
for line in stdout.decode("utf-8", errors="ignore").splitlines():
if not line.strip() or match_count >= effective_limit:
break
try:
event = json.loads(line)
except json.JSONDecodeError:
continue
if event.get("type") == "match":
match_count += 1
data = event.get("data", {})
file_path = data.get("path", {}).get("text", "")
line_number = data.get("line_number", 0)
if file_path and line_number:
matches.append({"file_path": file_path, "line_number": line_number})
if match_count >= effective_limit:
break
# Check for errors
if process.returncode not in (0, 1) and match_count == 0:
error_msg = stderr.decode("utf-8", errors="ignore").strip()
if error_msg:
raise RuntimeError(error_msg)
raise RuntimeError(f"ripgrep exited with code {process.returncode}")
# No matches found
if match_count == 0:
return "No matches found"
# Format matches with context
output_lines = []
file_cache = {}
for match in matches:
file_path = match["file_path"]
line_number = match["line_number"]
# Read file if not cached
if file_path not in file_cache:
try:
with open(file_path, "r", encoding="utf-8", errors="ignore") as f:
file_cache[file_path] = f.read().replace("\r\n", "\n").replace("\r", "\n").split("\n")
except Exception:
file_cache[file_path] = []
lines = file_cache[file_path]
# Format relative path
if is_directory:
relative_path = os.path.relpath(file_path, search_path)
if not relative_path.startswith(".."):
display_path = relative_path.replace("\\", "/")
else:
display_path = os.path.basename(file_path)
else:
display_path = os.path.basename(file_path)
# Generate context block
if not lines:
output_lines.append(f"{display_path}:{line_number}: (unable to read file)")
continue
context_value = max(0, context_lines)
start = max(1, line_number - context_value) if context_value > 0 else line_number
end = min(len(lines), line_number + context_value) if context_value > 0 else line_number
for current in range(start, end + 1):
if current < 1 or current > len(lines):
continue
line_text = lines[current - 1]
is_match_line = current == line_number
# Truncate long lines
truncated_text, was_truncated = truncate_line(line_text)
if was_truncated:
lines_truncated = True
if is_match_line:
output_lines.append(f"{display_path}:{current}: {truncated_text}")
else:
output_lines.append(f"{display_path}-{current}- {truncated_text}")
# Apply byte truncation
raw_output = "\n".join(output_lines)
truncation = truncate_head(raw_output, max_lines=999999999)
output = truncation.content
notices = []
# Add notices
if match_count >= effective_limit:
notices.append(
f"{effective_limit} matches limit reached. "
f"Use limit={effective_limit * 2} for more, or refine pattern",
)
if truncation.truncated:
notices.append(f"{format_size(DEFAULT_MAX_BYTES)} limit reached")
if lines_truncated:
notices.append(
f"Some lines truncated to {GREP_MAX_LINE_LENGTH} chars. " f"Use read tool to see full lines",
)
if notices:
output += f"\n\n[{'. '.join(notices)}]"
return output

128
reme/tool/fs/ls_tool.py Normal file
View file

@ -0,0 +1,128 @@
"""Directory listing tool with truncation support."""
import os
from pathlib import Path
from .truncate import DEFAULT_MAX_BYTES, truncate_head
from ...core.op import BaseTool
from ...core.schema import ToolCall
DEFAULT_LIMIT = 500
class LsTool(BaseTool):
"""List directory contents with smart truncation.
Features:
- Returns entries sorted alphabetically (case-insensitive)
- Directory indicators ('/' suffix)
- Includes dotfiles
- Entry count limiting
- Byte truncation
"""
def __init__(self, cwd: str | None = None):
"""Initialize ls tool.
Args:
cwd: Working directory (defaults to current directory)
"""
super().__init__()
self.cwd = cwd or os.getcwd()
def _build_tool_call(self) -> ToolCall:
max_kb = DEFAULT_MAX_BYTES // 1024
return ToolCall(
**{
"description": (
f"List directory contents. Returns entries sorted alphabetically, "
f"with '/' suffix for directories. Includes dotfiles. Output is truncated "
f"to {DEFAULT_LIMIT} entries or {max_kb}KB (whichever is hit first)."
),
"parameters": {
"type": "object",
"properties": {
"path": {
"type": "string",
"description": "Directory to list (default: current directory)",
},
"limit": {
"type": "number",
"description": f"Maximum number of entries to return (default: {DEFAULT_LIMIT})",
},
},
"required": [],
},
},
)
async def execute(self) -> str:
"""List directory contents with production-grade features."""
path: str | None = self.context.get("path", None)
limit: int | None = self.context.get("limit", None)
# Resolve directory path
dir_path = Path(self.cwd) / (path or ".")
dir_path = dir_path.resolve()
effective_limit = limit if limit is not None else DEFAULT_LIMIT
# Check if path exists
if not dir_path.exists():
raise FileNotFoundError(f"Path not found: {dir_path}")
# Check if path is a directory
if not dir_path.is_dir():
raise NotADirectoryError(f"Not a directory: {dir_path}")
# Read directory entries
try:
entries = list(dir_path.iterdir())
except Exception as e:
raise PermissionError(f"Cannot read directory: {e}") from e
# Sort alphabetically (case-insensitive)
entries.sort(key=lambda e: e.name.lower())
# Format entries with directory indicators
results: list[str] = []
entry_limit_reached = False
for entry in entries:
if len(results) >= effective_limit:
entry_limit_reached = True
break
try:
# Add '/' suffix for directories
suffix = "/" if entry.is_dir() else ""
results.append(entry.name + suffix)
except Exception:
# Skip entries we can't stat
continue
# Handle empty directory
if len(results) == 0:
return "(empty directory)"
# Apply byte truncation
raw_output = "\n".join(results)
truncation_result = truncate_head(raw_output, max_lines=float("inf"))
output_text = truncation_result.content
# Build notices
notices: list[str] = []
if entry_limit_reached:
notices.append(
f"{effective_limit} entries limit reached. Use limit={effective_limit * 2} for more",
)
if truncation_result.truncated:
max_kb = DEFAULT_MAX_BYTES // 1024
notices.append(f"{max_kb}KB limit reached")
if notices:
output_text += f"\n\n[{'. '.join(notices)}]"
return output_text

219
reme/tool/fs/read_tool.py Normal file
View file

@ -0,0 +1,219 @@
"""Read file tool with smart truncation and image support.
Features:
- Reads text files with offset/limit support
- Detects and handles image files (jpg, png, gif, webp)
- Smart truncation to prevent memory issues
"""
import os
from pathlib import Path
from .truncate import DEFAULT_MAX_BYTES, DEFAULT_MAX_LINES, format_size, truncate_head
from ...core.op import BaseTool
from ...core.schema import ToolCall
# Supported image extensions
IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".gif", ".webp"}
def is_image_file(path: str) -> bool:
"""Check if file is a supported image type.
Args:
path: File path to check
Returns:
True if file is a supported image
"""
return Path(path).suffix.lower() in IMAGE_EXTENSIONS
class ReadTool(BaseTool):
"""Read file contents with smart truncation.
Features:
- Supports text files and images (jpg, png, gif, webp)
- Smart truncation for large files
- Offset/limit for reading specific portions
"""
def __init__(self, cwd: str | None = None):
"""Initialize read tool.
Args:
cwd: Working directory (defaults to current directory)
"""
super().__init__()
self.cwd = cwd or os.getcwd()
def _build_tool_call(self) -> ToolCall:
max_kb = DEFAULT_MAX_BYTES // 1024
return ToolCall(
**{
"description": (
f"Read the contents of a file. Supports text files and images "
f"(jpg, png, gif, webp). Images are sent as attachments. For text files, "
f"output is truncated to {DEFAULT_MAX_LINES} lines or {max_kb}KB "
f"(whichever is hit first). Use offset/limit for large files. "
f"When you need the full file, continue with offset until complete."
),
"parameters": {
"type": "object",
"properties": {
"path": {
"type": "string",
"description": "Path to the file to read (relative or absolute)",
},
"offset": {
"type": "number",
"description": "Line number to start reading from (1-indexed)",
},
"limit": {
"type": "number",
"description": "Maximum number of lines to read",
},
},
"required": ["path"],
},
},
)
async def execute(self) -> str:
"""Execute the read operation."""
path: str = self.context.path
offset: int | None = self.context.get("offset", None)
limit: int | None = self.context.get("limit", None)
# Resolve path
if not os.path.isabs(path):
absolute_path = os.path.join(self.cwd, path)
else:
absolute_path = path
absolute_path = os.path.normpath(absolute_path)
# Check file exists and is readable
if not os.path.exists(absolute_path):
raise ValueError(f"File not found: {path}")
if not os.path.isfile(absolute_path):
raise ValueError(f"Not a file: {path}")
if not os.access(absolute_path, os.R_OK):
raise ValueError(f"File not readable: {path}")
# Check if image
if is_image_file(absolute_path):
return await self._read_image(absolute_path, path)
else:
return await self._read_text(absolute_path, path, offset, limit)
@staticmethod
async def _read_image(absolute_path: str, display_path: str) -> str:
"""Read and return image file information.
Args:
absolute_path: Absolute path to image
display_path: Path to display to user
Returns:
Image information text
"""
# Get file size
file_size = os.path.getsize(absolute_path)
file_ext = Path(absolute_path).suffix.lower()
# For Python tools, we typically can't return image data directly to LLM
# So we return a descriptive message
return (
f"Read image file [{file_ext}]\n"
f"Path: {display_path}\n"
f"Size: {format_size(file_size)}\n"
f"Note: Image content cannot be displayed in text format. "
f"Use bash tool or other methods to process the image."
)
@staticmethod
async def _read_text(
absolute_path: str,
_display_path: str,
offset: int | None,
limit: int | None,
) -> str:
"""Read text file with smart truncation.
Args:
absolute_path: Absolute path to file
_display_path: Path to display to user
offset: Starting line (1-indexed)
limit: Maximum lines to read
Returns:
File contents with truncation notices
"""
# Read file
try:
with open(absolute_path, "r", encoding="utf-8") as f:
content = f.read()
except UnicodeDecodeError:
# Try with error handling for binary files
with open(absolute_path, "r", encoding="utf-8", errors="ignore") as f:
content = f.read()
all_lines = content.split("\n")
total_file_lines = len(all_lines)
# Apply offset if specified (convert 1-indexed to 0-indexed)
start_line = max(0, (offset - 1)) if offset else 0
start_line_display = start_line + 1
# Check offset bounds
if start_line >= len(all_lines):
raise IndexError(
f"Offset {offset} is beyond end of file ({len(all_lines)} lines total)",
)
# Apply user limit if specified
if limit is not None:
end_line = min(start_line + limit, len(all_lines))
selected_content = "\n".join(all_lines[start_line:end_line])
user_limited_lines = end_line - start_line
else:
selected_content = "\n".join(all_lines[start_line:])
user_limited_lines = None
# Apply truncation
truncation = truncate_head(selected_content)
# Build output with truncation notices
if truncation.truncated:
# Truncation occurred
end_line_display = start_line_display + truncation.output_lines - 1
next_offset = end_line_display + 1
output_text = truncation.content
if truncation.truncated_by == "lines":
output_text += (
f"\n\n[Showing lines {start_line_display}-{end_line_display} "
f"of {total_file_lines}. Use offset={next_offset} to continue.]"
)
else:
max_kb = DEFAULT_MAX_BYTES // 1024
output_text += (
f"\n\n[Showing lines {start_line_display}-{end_line_display} "
f"of {total_file_lines} ({max_kb}KB limit). "
f"Use offset={next_offset} to continue.]"
)
elif user_limited_lines is not None and start_line + user_limited_lines < len(all_lines):
# User limit exceeded, but no truncation
remaining = len(all_lines) - (start_line + user_limited_lines)
next_offset = start_line + user_limited_lines + 1
output_text = truncation.content
output_text += f"\n\n[{remaining} more lines in file. " f"Use offset={next_offset} to continue.]"
else:
# No truncation or user limit exceeded
output_text = truncation.content
return output_text

209
reme/tool/fs/truncate.py Normal file
View file

@ -0,0 +1,209 @@
"""fs utils"""
from typing import Literal
from ...core.schema import TruncationResult
# Default limits for output truncation
DEFAULT_MAX_LINES = 1000 # Maximum lines to keep for tail truncation
DEFAULT_MAX_BYTES = 30 * 1024 # Maximum bytes to keep (30KB)
# Find tool limits
FIND_MAX_LINES = 2000 # Maximum lines for find output
FIND_MAX_BYTES = 50 * 1024 # 50KB for find output
# Grep tool limits
GREP_MAX_LINE_LENGTH = 500 # Maximum line length for grep output
def format_size(num_bytes: int) -> str:
"""Format byte size in human-readable format.
Args:
num_bytes: Number of bytes
Returns:
Formatted string (e.g., "1.5KB", "2.3MB")
"""
if num_bytes < 1024:
return f"{num_bytes}B"
elif num_bytes < 1024 * 1024:
return f"{num_bytes / 1024:.1f}KB"
else:
return f"{num_bytes / (1024 * 1024):.1f}MB"
def truncate_line(text: str, max_length: int = GREP_MAX_LINE_LENGTH) -> tuple[str, bool]:
"""Truncate a single line if it exceeds max length.
Args:
text: Line text
max_length: Maximum line length
Returns:
Tuple of (truncated_text, was_truncated)
"""
if len(text) <= max_length:
return text, False
return text[:max_length] + "...", True
def truncate_tail(
text: str,
max_lines: int = DEFAULT_MAX_LINES,
max_bytes: int = DEFAULT_MAX_BYTES,
) -> TruncationResult:
"""Truncate text to keep only the tail (last portion).
Keeps the last N lines or M bytes, whichever is hit first.
This is useful for command outputs where the end is most relevant.
Args:
text: The text to truncate
max_lines: Maximum number of lines to keep
max_bytes: Maximum bytes to keep
Returns:
TruncationResult with truncated content and metadata
"""
if not text:
return TruncationResult(
content="",
truncated=False,
total_lines=0,
output_lines=0,
total_bytes=0,
output_bytes=0,
)
total_bytes = len(text.encode("utf-8"))
lines = text.split("\n")
total_lines = len(lines)
# Check if we need to truncate
if total_lines <= max_lines and total_bytes <= max_bytes:
return TruncationResult(
content=text,
truncated=False,
total_lines=total_lines,
output_lines=total_lines,
total_bytes=total_bytes,
output_bytes=total_bytes,
)
# Keep last N lines
kept_lines = lines[-max_lines:] if total_lines > max_lines else lines
truncated_by: Literal["lines", "bytes"] = "lines" if total_lines > max_lines else "bytes"
# Check byte limit on kept lines
kept_text = "\n".join(kept_lines)
kept_bytes = len(kept_text.encode("utf-8"))
# If still over byte limit, truncate further
last_line_partial = False
if kept_bytes > max_bytes:
truncated_by = "bytes"
# Keep truncating from the start until under byte limit
while kept_lines and len("\n".join(kept_lines).encode("utf-8")) > max_bytes:
kept_lines.pop(0)
# If still over (single line > max_bytes), truncate the line itself
if kept_lines and len("\n".join(kept_lines).encode("utf-8")) > max_bytes:
last_line = kept_lines[-1]
# Binary search to find how much of last line fits
encoded = last_line.encode("utf-8")
if len(encoded) > max_bytes:
last_line_partial = True
# Take last max_bytes of the line
kept_lines[-1] = encoded[-max_bytes:].decode("utf-8", errors="ignore")
kept_text = "\n".join(kept_lines)
kept_bytes = len(kept_text.encode("utf-8"))
return TruncationResult(
content=kept_text,
truncated=True,
total_lines=total_lines,
output_lines=len(kept_lines),
total_bytes=total_bytes,
output_bytes=kept_bytes,
truncated_by=truncated_by,
last_line_partial=last_line_partial,
)
def truncate_head(
text: str,
max_lines: int = FIND_MAX_LINES,
max_bytes: int = FIND_MAX_BYTES,
) -> TruncationResult:
"""Truncate text to keep only the head (first portion).
Keeps the first N lines or M bytes, whichever is hit first.
Suitable for file reads where you want to see the beginning.
Args:
text: The text to truncate
max_lines: Maximum number of lines to keep
max_bytes: Maximum bytes to keep
Returns:
TruncationResult with truncated content and metadata
"""
if not text:
return TruncationResult(
content="",
truncated=False,
total_lines=0,
output_lines=0,
total_bytes=0,
output_bytes=0,
)
total_bytes = len(text.encode("utf-8"))
lines = text.split("\n")
total_lines = len(lines)
# Check if no truncation needed
if total_lines <= max_lines and total_bytes <= max_bytes:
return TruncationResult(
content=text,
truncated=False,
total_lines=total_lines,
output_lines=total_lines,
total_bytes=total_bytes,
output_bytes=total_bytes,
)
# Collect complete lines that fit
kept_lines = []
kept_bytes = 0
truncated_by: Literal["lines", "bytes"] = "lines"
for i, line in enumerate(lines):
if i >= max_lines:
truncated_by = "lines"
break
# Calculate bytes for this line (+1 for newline except first line)
line_bytes = len(line.encode("utf-8")) + (1 if i > 0 else 0)
if kept_bytes + line_bytes > max_bytes:
truncated_by = "bytes"
break
kept_lines.append(line)
kept_bytes += line_bytes
kept_text = "\n".join(kept_lines)
final_bytes = len(kept_text.encode("utf-8"))
return TruncationResult(
content=kept_text,
truncated=True,
total_lines=total_lines,
output_lines=len(kept_lines),
total_bytes=total_bytes,
output_bytes=final_bytes,
truncated_by=truncated_by,
)

View file

@ -0,0 +1,81 @@
"""Write tool for creating and overwriting files.
This module provides a tool for writing content to files with:
- Automatic parent directory creation
- File overwriting (creates if doesn't exist, overwrites if exists)
- Path resolution (relative to working directory)
"""
import os
from ...core.op import BaseTool
from ...core.schema import ToolCall
class WriteTool(BaseTool):
"""Tool for writing content to files.
Features:
- Creates file if it doesn't exist, overwrites if it does
- Automatically creates parent directories
- Supports both relative and absolute paths
"""
def __init__(self, cwd: str | None = None):
"""Initialize write tool.
Args:
cwd: Working directory (defaults to current directory)
"""
super().__init__()
self.cwd = cwd or os.getcwd()
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": (
"Write content to a file. Creates the file if it doesn't exist, "
"overwrites if it does. Automatically creates parent directories."
),
"parameters": {
"type": "object",
"properties": {
"path": {
"type": "string",
"description": "Path to the file to write (relative or absolute)",
},
"content": {
"type": "string",
"description": "Content to write to the file",
},
},
"required": ["path", "content"],
},
},
)
async def execute(self) -> str:
"""Execute the write operation."""
path: str = self.context.path
content: str = self.context.content
# Resolve path to absolute
if not os.path.isabs(path):
absolute_path = os.path.join(self.cwd, path)
else:
absolute_path = path
absolute_path = os.path.normpath(absolute_path)
# Create parent directories if needed
parent_dir = os.path.dirname(absolute_path)
if parent_dir:
os.makedirs(parent_dir, exist_ok=True)
# Write the file
with open(absolute_path, "w", encoding="utf-8") as f:
f.write(content)
# Return success message
content_bytes = len(content.encode("utf-8"))
return f"Successfully wrote {content_bytes} bytes to {path}"

400
tests/demo_memory_search.py Normal file
View file

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

View file

@ -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")

View file

@ -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"])

View file

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

View file

@ -0,0 +1,445 @@
"""Tests for file system tools including bash, edit, find, grep, ls, read, and write tools."""
import asyncio
import os
import tempfile
from pathlib import Path
async def test_bash_tool():
"""Test BashTool."""
from reme.tool.fs import BashTool
print("=== Testing BashTool ===")
bash_tool = BashTool()
result = await bash_tool.call(command="echo 'Hello World'")
print(f"Result: {result}")
assert "Hello World" in result
print("✓ BashTool test passed\n")
async def test_edit_tool():
"""Test EditTool."""
from reme.tool.fs import EditTool
print("=== Testing EditTool ===")
# Create temp file
with tempfile.NamedTemporaryFile(mode="w", suffix=".txt", delete=False) as f:
temp_path = f.name
f.write("Hello World\nThis is a test\nGoodbye World\n")
try:
# Test edit
edit_tool = EditTool()
result = await edit_tool.call(
path=temp_path,
oldText="This is a test",
newText="This is an updated test",
)
print(f"Result: {result}")
# Verify content
with open(temp_path, "r", encoding="utf-8") as f:
content = f.read()
assert "This is an updated test" in content
assert "This is a test" not in content
print("✓ EditTool test passed\n")
# Test error: file not found
print("=== Testing file not found error ===")
result = await edit_tool.call(
path="/nonexistent/file.txt",
oldText="test",
newText="new",
)
print(f"Expected error result: {result}")
assert "failed" in result and "File not found" in result
print("✓ File not found error test passed\n")
# Test error: text not found
print("=== Testing text not found error ===")
result = await edit_tool.call(
path=temp_path,
oldText="nonexistent text",
newText="new",
)
print(f"Expected error result: {result}")
assert "failed" in result and "Could not find" in result
print("✓ Text not found error test passed\n")
finally:
# Cleanup
if os.path.exists(temp_path):
os.unlink(temp_path)
async def test_find_tool():
"""Test FindTool."""
from reme.tool.fs import FindTool
print("=== Testing FindTool ===")
# Create temp directory with test files
with tempfile.TemporaryDirectory() as temp_dir:
temp_path = Path(temp_dir)
# Create some test files
(temp_path / "test1.txt").write_text("test file 1")
(temp_path / "test2.txt").write_text("test file 2")
(temp_path / "readme.md").write_text("readme")
# Create subdirectory with files
sub_dir = temp_path / "subdir"
sub_dir.mkdir()
(sub_dir / "test3.txt").write_text("test file 3")
(sub_dir / "config.json").write_text("{}")
# Create .gitignore to ignore certain files
(temp_path / ".gitignore").write_text("*.md\n")
# Test: find all txt files
find_tool = FindTool(cwd=str(temp_path))
result = await find_tool.call(pattern="*.txt")
print(f"Find *.txt result:\n{result}")
assert "test1.txt" in result
assert "test2.txt" in result
assert "readme.md" not in result # Should be ignored by .gitignore
print("✓ Find *.txt test passed\n")
# Test: find with recursive pattern
result = await find_tool.call(pattern="**/*.txt")
print(f"Find **/*.txt result:\n{result}")
assert "test1.txt" in result
assert "subdir/test3.txt" in result or "test3.txt" in result
print("✓ Find **/*.txt test passed\n")
# Test: find with no matches
result = await find_tool.call(pattern="*.nonexistent")
print(f"Find *.nonexistent result:\n{result}")
assert "No files found" in result
print("✓ No matches test passed\n")
# Test: error - directory not found
print("=== Testing directory not found error ===")
result = await find_tool.call(pattern="*.txt", path="/nonexistent/dir")
print(f"Expected error result: {result}")
assert "failed" in result and "Path not found" in result
print("✓ Directory not found error test passed\n")
async def test_grep_tool():
"""Test GrepTool."""
from reme.tool.fs import GrepTool
print("=== Testing GrepTool ===")
# Create temp directory with test files
with tempfile.TemporaryDirectory() as temp_dir:
temp_path = Path(temp_dir)
# Create test files with searchable content
(temp_path / "file1.txt").write_text("Hello World\nThis is a test\nGoodbye World\n")
(temp_path / "file2.txt").write_text("Another test file\nWith multiple lines\nHello again\n")
(temp_path / "script.py").write_text("def hello():\n print('Hello')\n return True\n")
# Create subdirectory with files
sub_dir = temp_path / "subdir"
sub_dir.mkdir()
(sub_dir / "nested.txt").write_text("Nested file content\nWith hello keyword\n")
# Test: search for pattern
grep_tool = GrepTool(cwd=str(temp_path))
result = await grep_tool.call(pattern="Hello", path=str(temp_path))
print(f"Search 'Hello' result:\n{result}")
assert "file1.txt" in result
assert "Hello World" in result or "Hello" in result
print("✓ Basic search test passed\n")
# Test: case-insensitive search
result = await grep_tool.call(pattern="hello", path=str(temp_path), ignoreCase=True)
print(f"Case-insensitive search result:\n{result}")
assert "file1.txt" in result or "Hello" in result.lower()
print("✓ Case-insensitive search test passed\n")
# Test: literal string search
result = await grep_tool.call(pattern="Hello()", path=str(temp_path), literal=True)
print(f"Literal search result:\n{result}")
# Should not find regex interpretation
print("✓ Literal search test passed\n")
# Test: glob filter
result = await grep_tool.call(pattern="Hello", path=str(temp_path), glob="*.txt")
print(f"Glob filter *.txt result:\n{result}")
assert "file1.txt" in result or "file2.txt" in result
assert ".py" not in result # Python files should be excluded
print("✓ Glob filter test passed\n")
# Test: context lines
result = await grep_tool.call(pattern="test", path=str(temp_path), contextLines=1)
print(f"Context lines result:\n{result}")
# Should include lines before and after matches
print("✓ Context lines test passed\n")
# Test: limit matches
result = await grep_tool.call(pattern="Hello", path=str(temp_path), limit=1)
print(f"Limit to 1 match result:\n{result}")
assert "limit reached" in result or result.count(":") >= 1
print("✓ Limit test passed\n")
# Test: no matches
result = await grep_tool.call(pattern="nonexistent_pattern_xyz", path=str(temp_path))
print(f"No matches result:\n{result}")
assert "No matches found" in result
print("✓ No matches test passed\n")
# Test: error - path not found
print("=== Testing path not found error ===")
try:
result = await grep_tool.call(pattern="test", path="/nonexistent/path")
assert "failed" in result and "not found" in result
except Exception as e:
assert "not found" in str(e).lower()
print("✓ Path not found error test passed\n")
async def test_ls_tool():
"""Test LsTool."""
from reme.tool.fs import LsTool
print("=== Testing LsTool ===")
# Create temp directory with test files
with tempfile.TemporaryDirectory() as temp_dir:
temp_path = Path(temp_dir)
# Create test files and directories
(temp_path / "file1.txt").write_text("test file 1")
(temp_path / "file2.py").write_text("test file 2")
(temp_path / ".hidden").write_text("hidden file")
(temp_path / "README.md").write_text("readme")
# Create subdirectories
(temp_path / "subdir1").mkdir()
(temp_path / "subdir2").mkdir()
# Test: list current directory
ls_tool = LsTool(cwd=str(temp_path))
result = await ls_tool.call()
print(f"List directory result:\n{result}")
assert ".hidden" in result # Includes dotfiles
assert "file1.txt" in result
assert "file2.py" in result
assert "subdir1/" in result # Directories have '/' suffix
assert "subdir2/" in result
print("✓ Basic ls test passed\n")
# Test: list specific path
result = await ls_tool.call(path=".")
print(f"List current directory result:\n{result}")
assert "file1.txt" in result
print("✓ Specific path test passed\n")
# Test: empty directory
empty_dir = temp_path / "empty"
empty_dir.mkdir()
result = await ls_tool.call(path="empty")
print(f"Empty directory result:\n{result}")
assert "(empty directory)" in result
print("✓ Empty directory test passed\n")
# Test: entry limit
# Create many files
for i in range(10):
(temp_path / f"file{i:03d}.txt").write_text(f"file {i}")
result = await ls_tool.call(limit=5)
print(f"Limited entries result:\n{result}")
assert "entries limit reached" in result
assert "limit=10" in result # Should suggest doubling the limit
print("✓ Entry limit test passed\n")
# Test: error - path not found
print("=== Testing path not found error ===")
result = await ls_tool.call(path="/nonexistent/path")
print(f"Expected error result: {result}")
assert "failed" in result and "Path not found" in result
print("✓ Path not found error test passed\n")
# Test: error - not a directory
print("=== Testing not a directory error ===")
result = await ls_tool.call(path="file1.txt")
print(f"Expected error result: {result}")
assert "failed" in result and "Not a directory" in result
print("✓ Not a directory error test passed\n")
async def test_read_tool():
"""Test ReadTool."""
from reme.tool.fs import ReadTool
print("=== Testing ReadTool ===")
# Create temp directory with test files
with tempfile.TemporaryDirectory() as temp_dir:
temp_path = Path(temp_dir)
# Create test text file
test_file = temp_path / "test.txt"
test_content = "\n".join([f"Line {i}" for i in range(1, 101)]) # 100 lines
test_file.write_text(test_content)
# Create test image file
image_file = temp_path / "test.jpg"
image_file.write_bytes(b"\xff\xd8\xff\xe0") # Minimal JPEG header
# Test: read full file
read_tool = ReadTool(cwd=str(temp_path))
result = await read_tool.call(path="test.txt")
print(f"Read full file result:\n{result[:200]}...")
assert "Line 1" in result
assert "Line 100" in result
print("✓ Read full file test passed\n")
# Test: read with offset
result = await read_tool.call(path="test.txt", offset=50)
print(f"Read with offset=50 result:\n{result[:200]}...")
# Check that we start from Line 50 (should be first line of content)
assert result.startswith("Line 50"), f"Should start with 'Line 50', got: {result[:50]}"
assert "Line 100" in result
print("✓ Read with offset test passed\n")
# Test: read with limit
result = await read_tool.call(path="test.txt", limit=10)
print(f"Read with limit=10 result:\n{result}")
assert "Line 1" in result
assert "Line 10" in result or "more lines in file" in result
assert "Line 50" not in result
print("✓ Read with limit test passed\n")
# Test: read with offset and limit
result = await read_tool.call(path="test.txt", offset=20, limit=5)
print(f"Read with offset=20, limit=5 result:\n{result}")
assert "Line 20" in result
assert "Line 24" in result or "more lines" in result
print("✓ Read with offset and limit test passed\n")
# Test: read image file
result = await read_tool.call(path="test.jpg")
print(f"Read image result:\n{result}")
assert "image file" in result.lower() or ".jpg" in result.lower()
print("✓ Read image test passed\n")
# Test: offset beyond file
print("=== Testing offset beyond file error ===")
result = await read_tool.call(path="test.txt", offset=200)
print(f"Expected error result: {result}")
assert "failed" in result and ("beyond end of file" in result or "offset" in result.lower())
print("✓ Offset beyond file error test passed\n")
# Test: file not found
print("=== Testing file not found error ===")
result = await read_tool.call(path="nonexistent.txt")
print(f"Expected error result: {result}")
assert "failed" in result and "not found" in result.lower()
print("✓ File not found error test passed\n")
# Test: read directory (should fail)
print("=== Testing read directory error ===")
sub_dir = temp_path / "subdir"
sub_dir.mkdir()
result = await read_tool.call(path="subdir")
print(f"Expected error result: {result}")
assert "failed" in result and ("Not a file" in result or "directory" in result.lower())
print("✓ Read directory error test passed\n")
async def test_write_tool():
"""Test WriteTool."""
from reme.tool.fs import WriteTool
print("=== Testing WriteTool ===")
# Create temp directory
with tempfile.TemporaryDirectory() as temp_dir:
temp_path = Path(temp_dir)
# Test: write new file
write_tool = WriteTool(cwd=str(temp_path))
test_content = "Hello World\nThis is a test file\n"
result = await write_tool.call(path="test.txt", content=test_content)
print(f"Write result: {result}")
assert "Successfully wrote" in result
assert "test.txt" in result
# Verify file was created
test_file = temp_path / "test.txt"
assert test_file.exists()
assert test_file.read_text() == test_content
print("✓ Write new file test passed\n")
# Test: overwrite existing file
new_content = "Updated content\n"
result = await write_tool.call(path="test.txt", content=new_content)
print(f"Overwrite result: {result}")
assert "Successfully wrote" in result
# Verify file was overwritten
assert test_file.read_text() == new_content
assert test_content not in test_file.read_text()
print("✓ Overwrite existing file test passed\n")
# Test: create file with parent directories
nested_path = "subdir1/subdir2/nested.txt"
nested_content = "Nested file content"
result = await write_tool.call(path=nested_path, content=nested_content)
print(f"Create with parents result: {result}")
assert "Successfully wrote" in result
# Verify nested file was created
nested_file = temp_path / "subdir1" / "subdir2" / "nested.txt"
assert nested_file.exists()
assert nested_file.read_text() == nested_content
print("✓ Create file with parent directories test passed\n")
# Test: write empty file
result = await write_tool.call(path="empty.txt", content="")
print(f"Write empty file result: {result}")
assert "Successfully wrote" in result
assert "0 bytes" in result
# Verify empty file
empty_file = temp_path / "empty.txt"
assert empty_file.exists()
assert empty_file.read_text() == ""
print("✓ Write empty file test passed\n")
# Test: write file with absolute path
abs_path = str(temp_path / "absolute.txt")
abs_content = "Absolute path content"
result = await write_tool.call(path=abs_path, content=abs_content)
print(f"Write absolute path result: {result}")
assert "Successfully wrote" in result
# Verify absolute path file
abs_file = Path(abs_path)
assert abs_file.exists()
assert abs_file.read_text(encoding="utf-8") == abs_content
print("✓ Write absolute path test passed\n")
async def main():
"""Run all file system tool tests."""
await test_bash_tool()
await test_edit_tool()
await test_find_tool()
await test_grep_tool()
await test_ls_tool()
await test_read_tool()
await test_write_tool()
print("=== All tests passed! ===")
if __name__ == "__main__":
asyncio.run(main())

323
tests/test_fs_agent.py Normal file
View file

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

979
tests/test_memory_store.py Normal file
View file

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

View file

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