feat(memory): implement file watcher with delta and full sync strategies

This commit is contained in:
jinli.yl 2026-02-06 01:58:28 +08:00
parent 857cd52a8e
commit c5438c52fc
41 changed files with 4453 additions and 1258 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

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

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

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

@ -1,890 +0,0 @@
"""Memory Index Manager - Main coordination layer.
This module provides the main MemoryIndexManager class that coordinates
file watching, embedding generation, and search operations across memory files
and session transcripts.
"""
import asyncio
import json
import os
import re
from typing import Any, Callable
from loguru import logger
from pydantic import BaseModel, Field
from watchfiles import awatch
from .ingestion.chunking import chunk_markdown
from .memory_storage.sqlite_memory_store import SqliteMemoryStore
from .utils.hashing import hash_text
from ..enumeration import MemorySource
from ..schema import FileMetadata, MemorySearchResult
# Constants
SNIPPET_MAX_CHARS = 700
SESSION_DIRTY_DEBOUNCE_MS = 5000
EMBEDDING_BATCH_MAX_TOKENS = 8000
EMBEDDING_APPROX_CHARS_PER_TOKEN = 1
EMBEDDING_INDEX_CONCURRENCY = 4
EMBEDDING_RETRY_MAX_ATTEMPTS = 3
EMBEDDING_RETRY_BASE_DELAY_MS = 500
EMBEDDING_RETRY_MAX_DELAY_MS = 8000
BATCH_FAILURE_LIMIT = 2
SESSION_DELTA_READ_CHUNK_BYTES = 64 * 1024
EMBEDDING_QUERY_TIMEOUT_REMOTE_MS = 60_000
EMBEDDING_QUERY_TIMEOUT_LOCAL_MS = 5 * 60_000
EMBEDDING_BATCH_TIMEOUT_REMOTE_MS = 2 * 60_000
EMBEDDING_BATCH_TIMEOUT_LOCAL_MS = 10 * 60_000
class MemorySyncProgressUpdate(BaseModel):
"""Progress update for memory sync operations."""
completed: int = Field(default=..., description="Number of items completed")
total: int = Field(default=..., description="Total number of items to process")
label: str | None = Field(default=None, description="Optional label for the progress operation")
class MemorySyncProgressState(BaseModel):
"""Internal state for tracking sync progress."""
completed: int = Field(default=0, description="Number of items completed")
total: int = Field(default=0, description="Total number of items to process")
label: str | None = Field(default=None, description="Optional label for the progress operation")
report: Callable[[MemorySyncProgressUpdate], None] | None = Field(
default=None,
description="Callback function to report progress updates",
)
class SessionDelta(BaseModel):
"""Tracks incremental changes in session files."""
last_size: int = Field(default=0, description="Last known size of the session file")
pending_bytes: int = Field(default=0, description="Number of pending bytes to process")
pending_messages: int = Field(default=0, description="Number of pending messages to process")
class MemorySearchConfig(BaseModel):
"""Configuration for memory search operations."""
model: str = Field(default="default", description="Model name for embeddings")
sources: list[MemorySource] = Field(default_factory=lambda: [MemorySource.MEMORY], description="Sources to search")
extra_paths: list[str] = Field(default_factory=list, description="Additional paths to include in search")
store_path: str = Field(default="memory.db", description="Path to SQLite database file")
vector_enabled: bool = Field(default=True, description="Whether to enable vector search")
vector_extension_path: str | None = Field(default=None, description="Path to vector extension for SQLite")
fts_enabled: bool = Field(default=True, description="Whether to enable full-text search")
chunk_tokens: int = Field(default=300, description="Number of tokens per chunk")
chunk_overlap: int = Field(default=30, description="Number of overlapping tokens between chunks")
watch_enabled: bool = Field(default=True, description="Whether to enable file watching")
watch_debounce_ms: int = Field(default=1000, description="Debounce time for file watcher in milliseconds")
interval_minutes: int = Field(default=0, description="Interval between automatic syncs in minutes (0 to disable)")
sync_on_search: bool = Field(default=True, description="Whether to sync before search operations")
sync_on_session_start: bool = Field(default=True, description="Whether to sync when a session starts")
query_min_score: float = Field(default=0.3, description="Minimum relevance score for search results")
query_max_results: int = Field(default=10, description="Maximum number of search results to return")
hybrid_enabled: bool = Field(default=True, description="Whether to use hybrid vector + keyword search")
hybrid_vector_weight: float = Field(default=0.7, description="Weight for vector search in hybrid scoring")
hybrid_text_weight: float = Field(default=0.3, description="Weight for text search in hybrid scoring")
hybrid_candidate_multiplier: float = Field(
default=2.0,
description="Multiplier for number of candidates to consider in hybrid search",
)
session_delta_bytes: int = Field(
default=0,
description="Threshold for session sync based on bytes changed (0 for any change)",
)
session_delta_messages: int = Field(
default=5,
description="Threshold for session sync based on messages changed (0 for any change)",
)
# Global cache for manager instances
INDEX_CACHE: dict[str, "MemoryIndexManager"] = {}
class MemoryIndexManager:
"""Main memory index manager coordinating all memory operations."""
# ============================================================================
# Initialization and Lifecycle
# ============================================================================
def __init__(
self,
agent_id: str,
workspace_dir: str,
settings: MemorySearchConfig,
store: SqliteMemoryStore,
):
"""Initialize the memory index manager."""
self.agent_id = agent_id
self.workspace_dir = workspace_dir
self.settings = settings
self.store = store
# State tracking
self.sources = set(settings.sources)
self.closed = False
self.dirty = MemorySource.MEMORY in self.sources
self.sessions_dirty = False
self.sessions_dirty_files: set[str] = set()
self.session_pending_files: set[str] = set()
self.session_deltas: dict[str, SessionDelta] = {}
self.session_warm: set[str] = set()
# Sync control
self.syncing: asyncio.Task | None = None
self.watch_task: asyncio.Task | None = None
self.session_watch_task: asyncio.Task | None = None
self.interval_task: asyncio.Task | None = None
# Batch failure tracking
self.batch_failure_count = 0
self.batch_failure_last_error: str | None = None
self.batch_failure_lock = asyncio.Lock()
async def close(self) -> None:
"""Close the manager and release resources."""
if self.closed:
return
self.closed = True
# Cancel all background tasks
if self.watch_task:
self.watch_task.cancel()
if self.session_watch_task:
self.session_watch_task.cancel()
if self.interval_task:
self.interval_task.cancel()
await self.store.close()
# ============================================================================
# Public API Methods
# ============================================================================
async def warm_session(self, session_key: str | None = None):
"""Pre-sync memory before a session starts."""
if not self.settings.sync_on_session_start:
return
key = (session_key or "").strip()
if key and key in self.session_warm:
return
await self.sync(reason="session-start")
if key:
self.session_warm.add(key)
async def sync(
self,
reason: str | None = None,
force: bool = False,
progress: Callable[[MemorySyncProgressUpdate], None] | None = None,
):
"""Synchronize memory index with file system."""
if self.syncing:
await self.syncing
return
self.syncing = asyncio.create_task(self._run_sync(reason, force, progress))
try:
await self.syncing
finally:
self.syncing = None
async def search(
self,
query: str,
max_results: int | None = None,
min_score: float | None = None,
session_key: str | None = None,
) -> list[MemorySearchResult]:
"""Search indexed memory with hybrid vector + keyword search.
Args:
query: Search query text
max_results: Maximum number of results to return
min_score: Minimum relevance score threshold
session_key: Optional session key for warmup
Returns:
List of search results sorted by relevance
"""
await self.warm_session(session_key)
if self.settings.sync_on_search and (self.dirty or self.sessions_dirty):
try:
await self.sync(reason="search")
except Exception as err:
logger.warning(f"memory sync failed (search): {err}")
cleaned = query.strip()
if not cleaned:
return []
min_score = min_score if min_score is not None else self.settings.query_min_score
max_results = max_results if max_results is not None else self.settings.query_max_results
hybrid = self.settings.hybrid_enabled
candidates = min(200, max(1, int(max_results * self.settings.hybrid_candidate_multiplier)))
# Run keyword search if hybrid enabled
keyword_results = []
if hybrid:
keyword_results = await self._search_keyword(cleaned, candidates)
# Perform vector search
vector_results = await self._search_vector(cleaned, candidates)
if not hybrid:
return [r for r in vector_results if r.score >= min_score][:max_results]
merged = self._merge_hybrid_results(
vector=vector_results,
keyword=keyword_results,
vector_weight=self.settings.hybrid_vector_weight,
text_weight=self.settings.hybrid_text_weight,
)
return [r for r in merged if r.score >= min_score][:max_results]
async def read_file(
self,
rel_path: str,
from_line: int | None = None,
num_lines: int | None = None,
) -> dict[str, str]:
"""Read a memory file with optional line range.
Args:
rel_path: Relative path to file
from_line: Starting line number (1-indexed)
num_lines: Number of lines to read
Returns:
Dictionary with 'text' and 'path' keys
Raises:
ValueError: If path is invalid or not allowed
"""
raw_path = rel_path.strip()
assert raw_path, "path required"
abs_path = os.path.abspath(os.path.join(self.workspace_dir, raw_path))
rel_path_clean = os.path.relpath(abs_path, self.workspace_dir)
in_workspace = not rel_path_clean.startswith("..") and not os.path.isabs(rel_path_clean)
allowed = in_workspace and self._is_memory_path(rel_path_clean)
if not allowed and self.settings.extra_paths:
for extra in self.settings.extra_paths:
extra_abs = os.path.abspath(extra)
if abs_path.startswith(extra_abs):
allowed = True
break
if not allowed:
raise ValueError("path required")
if not abs_path.endswith(".md"):
raise ValueError("path required")
# Read file
with open(abs_path, "r", encoding="utf-8") as f:
content = f.read()
if from_line is None and num_lines is None:
return {"text": content, "path": rel_path_clean}
lines = content.split("\n")
start = max(1, from_line or 1)
count = max(1, num_lines or len(lines))
slice_lines = lines[start - 1 : start - 1 + count]
return {"text": "\n".join(slice_lines), "path": rel_path_clean}
# ============================================================================
# Sync Logic
# ============================================================================
async def _run_sync(
self,
reason: str | None,
force: bool,
progress_callback: Callable[[MemorySyncProgressUpdate], None] | None,
):
"""Execute sync operation."""
progress = MemorySyncProgressState()
if progress_callback:
progress.report = progress_callback
should_sync_memory = MemorySource.MEMORY in self.sources and (force or self.dirty)
should_sync_sessions = self._should_sync_sessions(reason, force)
if should_sync_memory:
await self._sync_memory_files(progress)
self.dirty = False
if should_sync_sessions:
await self._sync_session_files(progress)
self.sessions_dirty = False
self.sessions_dirty_files.clear()
elif len(self.sessions_dirty_files) > 0:
self.sessions_dirty = True
else:
self.sessions_dirty = False
def _should_sync_sessions(self, reason: str | None, force: bool) -> bool:
"""Check if session sync is needed."""
if MemorySource.SESSIONS not in self.sources:
return False
if force:
return True
if reason in ("session-start", "watch"):
return False
return self.sessions_dirty and len(self.sessions_dirty_files) > 0
async def _sync_memory_files(self, progress: MemorySyncProgressState):
"""Sync memory markdown files."""
files = self._list_memory_files()
logger.debug("memory sync: indexing memory files", files=len(files))
active_paths = {f.path for f in files}
if progress.report:
progress.total += len(files)
progress.report(
MemorySyncProgressUpdate(
completed=progress.completed,
total=progress.total,
label="Indexing memory files…",
),
)
tasks = []
for file_entry in files:
task = self._index_memory_file(file_entry, progress)
tasks.append(task)
await asyncio.gather(*tasks)
indexed = await self.store.list_files(MemorySource.MEMORY)
for stale_path in indexed:
if stale_path not in active_paths:
await self.store.delete_file(stale_path, MemorySource.MEMORY)
async def _sync_session_files(self, progress: MemorySyncProgressState):
"""Sync session transcript files."""
files = self._list_session_files()
logger.debug(
"memory sync: indexing session files",
files=len(files),
index_all=len(self.sessions_dirty_files) == 0,
dirty_files=len(self.sessions_dirty_files),
)
if progress.report:
progress.total += len(files)
progress.report(
MemorySyncProgressUpdate(
completed=progress.completed,
total=progress.total,
label="Indexing session files...",
),
)
active_paths = set()
tasks = []
for abs_path in files:
rel_path = self._session_path_for_file(abs_path)
active_paths.add(rel_path)
if len(self.sessions_dirty_files) == 0 or abs_path in self.sessions_dirty_files:
task = self._index_session_file(abs_path, progress)
tasks.append(task)
else:
if progress.report:
progress.completed += 1
progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
await asyncio.gather(*tasks)
indexed = await self.store.list_files(MemorySource.SESSIONS)
for stale_path in indexed:
if stale_path not in active_paths:
await self.store.delete_file(stale_path, MemorySource.SESSIONS)
# ============================================================================
# File Indexing
# ============================================================================
async def _index_memory_file(self, file_meta: FileMetadata, progress: MemorySyncProgressState):
"""Index a single memory file."""
existing_meta = await self.store.get_file_metadata(file_meta.path, MemorySource.MEMORY)
if existing_meta and existing_meta.hash == file_meta.hash:
if progress.report:
progress.completed += 1
progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
return
# Read and chunk file
with open(file_meta.abs_path, "r", encoding="utf-8") as f:
content = f.read()
chunks = chunk_markdown(
content,
file_meta.path,
MemorySource.MEMORY,
self.settings.chunk_tokens,
self.settings.chunk_overlap,
)
chunks = [c for c in chunks if c.text.strip()]
if chunks:
chunks = await self.store.get_chunk_embeddings(chunks)
await self.store.upsert_file(file_meta, MemorySource.MEMORY, chunks)
if progress.report:
progress.completed += 1
progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
async def _index_session_file(self, abs_path: str, progress: MemorySyncProgressState):
"""Index a single session transcript file."""
file_meta = self._build_session_file_meta(abs_path)
if not file_meta:
if progress.report:
progress.completed += 1
progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
return
existing_meta = await self.store.get_file_metadata(file_meta.path, MemorySource.SESSIONS)
if existing_meta and existing_meta.hash == file_meta.hash:
self._reset_session_delta(abs_path, file_meta.size)
if progress.report:
progress.completed += 1
progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
return
chunks = chunk_markdown(
file_meta.content,
file_meta.path,
MemorySource.SESSIONS,
self.settings.chunk_tokens,
self.settings.chunk_overlap,
)
chunks = [c for c in chunks if c.text.strip()]
if chunks:
chunks = await self.store.get_chunk_embeddings(chunks)
await self.store.upsert_file(file_meta, MemorySource.SESSIONS, chunks)
self._reset_session_delta(abs_path, file_meta.size)
if progress.report:
progress.completed += 1
progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
# ============================================================================
# File Listing and Building
# ============================================================================
def _list_memory_files(self) -> list[FileMetadata]:
"""List all memory markdown files."""
files = []
# Scan workspace
memory_paths = [
os.path.join(self.workspace_dir, "MEMORY.md"),
os.path.join(self.workspace_dir, "memory.md"),
os.path.join(self.workspace_dir, "memory"),
]
for base_path in memory_paths:
if os.path.isfile(base_path) and base_path.endswith(".md"):
files.append(self._build_file_entry(base_path))
elif os.path.isdir(base_path):
for root, _, filenames in os.walk(base_path):
for filename in filenames:
if filename.endswith(".md"):
abs_path = os.path.join(root, filename)
files.append(self._build_file_entry(abs_path))
# Extra paths
for extra in self.settings.extra_paths:
if os.path.isfile(extra) and extra.endswith(".md"):
files.append(self._build_file_entry(extra))
elif os.path.isdir(extra):
for root, _, filenames in os.walk(extra):
for filename in filenames:
if filename.endswith(".md"):
abs_path = os.path.join(root, filename)
files.append(self._build_file_entry(abs_path))
return files
def _list_session_files(self) -> list[str]:
"""List all session transcript files."""
sessions_dir = os.path.join(self.workspace_dir, "sessions", self.agent_id)
if not os.path.exists(sessions_dir):
return []
files = []
for filename in os.listdir(sessions_dir):
if filename.endswith(".jsonl"):
files.append(os.path.join(sessions_dir, filename))
return files
def _build_file_entry(self, abs_path: str) -> FileMetadata:
"""Build file entry metadata."""
stat = os.stat(abs_path)
with open(abs_path, "r", encoding="utf-8") as f:
content = f.read()
rel_path = os.path.relpath(abs_path, self.workspace_dir)
return FileMetadata(
hash=hash_text(content),
mtime_ms=stat.st_mtime * 1000,
size=stat.st_size,
path=rel_path.replace("\\", "/"),
abs_path=abs_path,
)
def _build_session_file_meta(self, abs_path: str) -> FileMetadata | None:
"""Build session file entry with parsed content. TODO 修改message解析逻辑"""
stat = os.stat(abs_path)
with open(abs_path, "r", encoding="utf-8") as f:
raw = f.read()
lines = raw.split("\n")
collected = []
for line in lines:
if not line.strip():
continue
try:
record = json.loads(line)
except json.JSONDecodeError:
continue
if record.get("type") != "message":
continue
message = record.get("message", {})
role = message.get("role")
if role not in ("user", "assistant"):
continue
text = self._extract_session_text(message.get("content"))
if not text:
continue
label = "User" if role == "user" else "Assistant"
collected.append(f"{label}: {text}")
content = "\n".join(collected)
rel_path = self._session_path_for_file(abs_path)
return FileMetadata(
hash=hash_text(content),
mtime_ms=stat.st_mtime * 1000,
size=stat.st_size,
path=rel_path,
abs_path=abs_path,
content=content,
)
# ============================================================================
# Session Processing Helpers
# ============================================================================
@staticmethod
def _session_path_for_file(abs_path: str) -> str:
"""Convert absolute session path to relative."""
return f"sessions/{os.path.basename(abs_path)}"
def _extract_session_text(self, content: Any) -> str | None:
"""Extract text from session message content."""
if isinstance(content, str):
normalized = self._normalize_session_text(content)
return normalized if normalized else None
if not isinstance(content, list):
return None
parts = []
for block in content:
if not isinstance(block, dict):
continue
if block.get("type") != "text":
continue
text = block.get("text")
if isinstance(text, str):
normalized = self._normalize_session_text(text)
if normalized:
parts.append(normalized)
return " ".join(parts) if parts else None
@staticmethod
def _normalize_session_text(text: str) -> str:
"""Normalize session text by collapsing whitespace."""
text = re.sub(r"\s*\n+\s*", " ", text)
text = re.sub(r"\s+", " ", text)
return text.strip()
# ============================================================================
# Session Delta Tracking
# ============================================================================
async def _process_session_delta_batch(self) -> None:
"""Process pending session file changes."""
if not self.session_pending_files:
return
pending = list(self.session_pending_files)
self.session_pending_files.clear()
should_sync = False
for session_file in pending:
delta = await self._update_session_delta(session_file)
if not delta:
continue
bytes_threshold = self.settings.session_delta_bytes
messages_threshold = self.settings.session_delta_messages
if bytes_threshold <= 0:
bytes_hit = delta["pending_bytes"] > 0
else:
bytes_hit = delta["pending_bytes"] >= bytes_threshold
if messages_threshold <= 0:
messages_hit = delta["pending_messages"] > 0
else:
messages_hit = delta["pending_messages"] >= messages_threshold
if not bytes_hit and not messages_hit:
continue
self.sessions_dirty_files.add(session_file)
self.sessions_dirty = True
should_sync = True
if should_sync:
try:
await self.sync(reason="session-delta")
except Exception as err:
logger.warning(f"memory sync failed (session-delta): {err}")
async def _update_session_delta(self, session_file: str) -> dict[str, int] | None:
"""Update delta tracking for a session file."""
try:
stat = os.stat(session_file)
size = stat.st_size
except OSError:
return None
state = self.session_deltas.get(session_file)
if not state:
state = SessionDelta()
self.session_deltas[session_file] = state
delta_bytes = max(0, size - state.last_size)
if delta_bytes == 0 and size == state.last_size:
return {
"delta_bytes": self.settings.session_delta_bytes,
"delta_messages": self.settings.session_delta_messages,
"pending_bytes": state.pending_bytes,
"pending_messages": state.pending_messages,
}
if size < state.last_size:
state.last_size = size
state.pending_bytes += size
if self.settings.session_delta_messages > 0:
state.pending_messages += await self._count_newlines(session_file, 0, size)
else:
state.pending_bytes += delta_bytes
if self.settings.session_delta_messages > 0:
state.pending_messages += await self._count_newlines(session_file, state.last_size, size)
state.last_size = size
return {
"delta_bytes": self.settings.session_delta_bytes,
"delta_messages": self.settings.session_delta_messages,
"pending_bytes": state.pending_bytes,
"pending_messages": state.pending_messages,
}
def _reset_session_delta(self, abs_path: str, size: int) -> None:
"""Reset delta tracking for a session file."""
state = self.session_deltas.get(abs_path)
if state:
state.last_size = size
state.pending_bytes = 0
state.pending_messages = 0
@staticmethod
async def _count_newlines(abs_path: str, start: int, end: int) -> int:
"""Count newlines in a file range."""
if end <= start:
return 0
count = 0
with open(abs_path, "rb") as f:
f.seek(start)
remaining = end - start
while remaining > 0:
chunk_size = min(SESSION_DELTA_READ_CHUNK_BYTES, remaining)
chunk = f.read(chunk_size)
if not chunk:
break
count += chunk.count(b"\n")
remaining -= len(chunk)
return count
# ============================================================================
# File Watchers
# ============================================================================
async def _start_watchers(self):
"""Start file watching and interval sync tasks."""
if self.settings.watch_enabled and MemorySource.MEMORY in self.sources:
self.watch_task = asyncio.create_task(self._watch_memory_files())
if MemorySource.SESSIONS in self.sources:
self.session_watch_task = asyncio.create_task(self._watch_session_files())
if self.settings.interval_minutes > 0:
self.interval_task = asyncio.create_task(self._interval_sync())
async def _watch_memory_files(self) -> None:
"""Watch memory files for changes."""
watch_paths = [
os.path.join(self.workspace_dir, "MEMORY.md"),
os.path.join(self.workspace_dir, "memory.md"),
os.path.join(self.workspace_dir, "memory"),
]
for extra in self.settings.extra_paths:
watch_paths.append(extra)
async for changes in awatch(*watch_paths, stop_event=None):
if self.closed:
break
for _, path in changes:
if path.endswith(".md"):
self.dirty = True
await asyncio.sleep(self.settings.watch_debounce_ms / 1000)
try:
await self.sync(reason="watch")
except Exception as e:
logger.exception(f"memory sync failed (watch): {e}")
async def _watch_session_files(self):
"""Watch session files for changes."""
sessions_dir = os.path.join(self.workspace_dir, "sessions", self.agent_id)
if not os.path.exists(sessions_dir):
return
async for changes in awatch(sessions_dir, stop_event=None):
if self.closed:
break
for _, path in changes:
if path.endswith(".jsonl"):
self.session_pending_files.add(path)
await asyncio.sleep(SESSION_DIRTY_DEBOUNCE_MS / 1000)
await self._process_session_delta_batch()
async def _interval_sync(self) -> None:
"""Periodically sync the index."""
while not self.closed:
await asyncio.sleep(self.settings.interval_minutes * 60)
if not self.closed:
try:
await self.sync(reason="interval")
except Exception as err:
logger.warning(f"memory sync failed (interval): {err}")
# ============================================================================
# Search Methods
# ============================================================================
async def _search_vector(self, query: str, limit: int) -> list[MemorySearchResult]:
"""Perform vector similarity search."""
return await self.store.vector_search(query, limit, sources=list(self.sources))
async def _search_keyword(self, query: str, limit: int) -> list[MemorySearchResult]:
"""Perform keyword/FTS search."""
if not self.settings.fts_enabled:
return []
return await self.store.keyword_search(query, limit, sources=list(self.sources))
@staticmethod
def _merge_hybrid_results(
vector: list[MemorySearchResult],
keyword: list[MemorySearchResult],
vector_weight: float,
text_weight: float,
) -> list[MemorySearchResult]:
"""Merge vector and keyword search results."""
merged: dict[str, MemorySearchResult] = {}
# Process vector results
for result in vector:
result.score = result.score * vector_weight
merged[result.merge_key] = result
# Process keyword results
for result in keyword:
key = result.merge_key
if key in merged:
merged[key].score += result.score * text_weight
else:
result.score = result.score * text_weight
merged[key] = result
results = list(merged.values())
results.sort(key=lambda r: r.score, reverse=True)
return results
# ============================================================================
# Utility Methods
# ============================================================================
@staticmethod
def _is_memory_path(rel_path: str) -> bool:
"""Check if path is a valid memory path."""
normalized = rel_path.replace("\\", "/")
if normalized in ("MEMORY.md", "memory.md"):
return True
if normalized.startswith("memory/") and normalized.endswith(".md"):
return True
return False

View file

@ -1,15 +0,0 @@
"""Utility functions for hashing text content."""
import hashlib
def hash_text(text: str) -> str:
"""Generate SHA-256 hash of text content.
Args:
text: Input text to hash
Returns:
Hexadecimal representation of the SHA-256 hash
"""
return hashlib.sha256(text.encode("utf-8")).hexdigest()

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

@ -2,17 +2,31 @@
from abc import ABC, abstractmethod
from ...embedding import BaseEmbeddingModel
from ...enumeration import MemorySource
from ...schema import FileMetadata, MemoryIndexMeta, MemoryChunk, MemorySearchResult
from ..embedding import BaseEmbeddingModel
from ..enumeration import MemorySource
from ..schema import FileMetadata, MemoryChunk, MemorySearchResult
class BaseMemoryStore(ABC):
"""Abstract base class for memory storage backends."""
def __init__(self, embedding_model: BaseEmbeddingModel):
def __init__(
self,
store_name: str,
embedding_model: BaseEmbeddingModel,
fts_enabled: bool = True,
snippet_max_chars: int = 700,
**kwargs,
):
"""Initialize"""
self.store_name: str = store_name
self.embedding_model: BaseEmbeddingModel = embedding_model
self.fts_enabled: bool = fts_enabled
self.snippet_max_chars: int = snippet_max_chars
self.kwargs: dict = kwargs
self.vector_available = False
self.fts_available = False
@property
def embedding_dim(self) -> int:
@ -20,77 +34,21 @@ class BaseMemoryStore(ABC):
return self.embedding_model.dimensions
async def get_embedding(self, query: str, **kwargs) -> list[float]:
"""Get embedding for a single query string.
Args:
query: Input text to generate embedding for
**kwargs: Additional arguments passed to the embedding model
Returns:
Embedding vector as a list of floats
"""
"""Get embedding for a single query string."""
return await self.embedding_model.get_embedding(query, **kwargs)
async def get_embeddings(self, queries: list[str], **kwargs) -> list[list[float]]:
"""Get embeddings for a batch of query strings.
Args:
queries: List of input texts to generate embeddings for
**kwargs: Additional arguments passed to the embedding model
Returns:
List of embedding vectors, each as a list of floats
"""
"""Get embeddings for a batch of query strings."""
return await self.embedding_model.get_embeddings(queries, **kwargs)
async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
"""Generate and populate embedding field for a single MemoryChunk object.
Args:
chunk: MemoryChunk object containing text to embed
**kwargs: Additional arguments passed to the embedding model
Returns:
The same MemoryChunk object with populated embedding field
"""
"""Generate and populate embedding field for a single MemoryChunk object."""
return await self.embedding_model.get_chunk_embedding(chunk, **kwargs)
async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
"""Generate and populate embedding fields for a batch of MemoryChunk objects.
Args:
chunks: List of MemoryChunk objects containing text to embed
**kwargs: Additional arguments passed to the embedding model
Returns:
The same list of MemoryChunk objects with populated embedding fields
"""
"""Generate and populate embedding fields for a batch of MemoryChunk objects."""
return await self.embedding_model.get_chunk_embeddings(chunks, **kwargs)
def get_chunk_embedding_sync(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
"""Synchronously generate and populate embedding field for a single MemoryChunk object.
Args:
chunk: MemoryChunk object containing text to embed
**kwargs: Additional arguments passed to the embedding model
Returns:
The same MemoryChunk object with populated embedding field
"""
return self.embedding_model.get_chunk_embedding_sync(chunk, **kwargs)
def get_chunk_embeddings_sync(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
"""Synchronously generate embeddings for a batch of MemoryChunk objects.
Args:
chunks: List of MemoryChunk objects containing text to embed
**kwargs: Additional arguments passed to the embedding model
Returns:
The same list of MemoryChunk objects with populated embedding fields
"""
return self.embedding_model.get_chunk_embeddings_sync(chunks, **kwargs)
@abstractmethod
async def start(self):
"""Initialize the storage backend."""
@ -104,19 +62,23 @@ class BaseMemoryStore(ABC):
"""Delete a file and all its chunks."""
@abstractmethod
async def get_file_hash(self, path: str, source: MemorySource) -> str | None:
"""Get the hash of an indexed file."""
async def delete_file_chunks(self, path: str, chunk_ids: list[str]):
"""Delete chunks for a file."""
@abstractmethod
async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None:
"""Get full file metadata with statistics."""
async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource):
"""Insert or update specific chunks without affecting other chunks."""
@abstractmethod
async def list_files(self, source: MemorySource) -> list[str]:
"""List all indexed file paths for a source."""
@abstractmethod
async def get_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None:
"""Get full file metadata with statistics."""
@abstractmethod
async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
"""Get all chunks for a file."""
@abstractmethod
@ -155,14 +117,6 @@ class BaseMemoryStore(ABC):
List of search results sorted by relevance
"""
@abstractmethod
async def read_meta(self, key: str) -> MemoryIndexMeta | None:
"""Read metadata value."""
@abstractmethod
async def write_meta(self, key: str, value: MemoryIndexMeta | dict):
"""Write metadata value."""
@abstractmethod
async def clear_all(self):
"""Clear all indexed data."""

View file

@ -9,9 +9,8 @@ from pathlib import Path
from loguru import logger
from .base_memory_store import BaseMemoryStore
from ...embedding import BaseEmbeddingModel
from ...enumeration import MemorySource
from ...schema import FileMetadata, MemoryIndexMeta, MemoryChunk, MemorySearchResult
from ..enumeration import MemorySource
from ..schema import FileMetadata, MemoryChunk, MemorySearchResult
class SqliteMemoryStore(BaseMemoryStore):
@ -28,25 +27,32 @@ class SqliteMemoryStore(BaseMemoryStore):
- Efficient chunk and file metadata management
"""
VECTOR_TABLE = "chunks_vec"
FTS_TABLE = "chunks_fts"
def __init__(
self,
db_path: str,
embedding_model: BaseEmbeddingModel,
vec_ext_path: str = "",
fts_enabled: bool = True,
snippet_max_chars: int = 700,
):
super().__init__(embedding_model=embedding_model)
def __init__(self, db_path: str = ".reme/memory.db", vec_ext_path: str = "", **kwargs):
super().__init__(**kwargs)
self.db_path = db_path
self.vec_ext_path = vec_ext_path
self.fts_enabled = fts_enabled
self.snippet_max_chars = snippet_max_chars
self.conn: sqlite3.Connection | None = None
self.vector_available = False
self.fts_available = False
@property
def vector_table_name(self) -> str:
"""Get the name of the vector table for this store."""
return f"chunks_vec_{self.store_name}"
@property
def fts_table_name(self) -> str:
"""Get the name of the FTS table for this store."""
return f"chunks_fts_{self.store_name}"
@property
def chunks_table_name(self) -> str:
"""Get the name of the chunks table for this store."""
return f"chunks_{self.store_name}"
@property
def files_table_name(self) -> str:
"""Get the name of the files table for this store."""
return f"files_{self.store_name}"
@staticmethod
def vector_to_blob(embedding: list[float]) -> bytes:
@ -55,6 +61,9 @@ class SqliteMemoryStore(BaseMemoryStore):
async def start(self) -> None:
"""Initialize database and load extensions."""
if self.conn is not None:
return
Path(self.db_path).parent.mkdir(parents=True, exist_ok=True)
self.conn = sqlite3.connect(self.db_path, check_same_thread=False)
@ -68,16 +77,27 @@ class SqliteMemoryStore(BaseMemoryStore):
logger.info(f"Loaded sqlite-vec: {self.vec_ext_path}")
except Exception as e:
logger.warning(f"Failed to load sqlite-vec: {e}")
else:
# Try common extension names
for name in ["vec0", "sqlite_vec", "vector0"]:
try:
self.conn.load_extension(name)
self.vector_available = True
logger.info(f"Loaded sqlite-vec: {name}")
break
except Exception:
pass
try:
import sqlite_vec
ext_path = sqlite_vec.loadable_path()
self.conn.load_extension(ext_path)
self.vector_available = True
logger.info(f"Loaded sqlite-vec from package: {ext_path}")
except Exception as e:
logger.warning(f"Failed to load sqlite-vec from package: {e}")
# Fallback: try common extension names
for name in ["vec0", "sqlite_vec", "vector0"]:
try:
self.conn.load_extension(name)
self.vector_available = True
logger.info(f"Loaded sqlite-vec: {name}")
break
except Exception:
pass
self.conn.enable_load_extension(False)
await self._create_tables()
@ -86,20 +106,10 @@ class SqliteMemoryStore(BaseMemoryStore):
"""Create database schema."""
cursor = self.conn.cursor()
# Metadata
cursor.execute(
"""
CREATE TABLE IF NOT EXISTS meta (
key TEXT PRIMARY KEY,
value TEXT
)
""",
)
# Files
cursor.execute(
"""
CREATE TABLE IF NOT EXISTS files (
f"""
CREATE TABLE IF NOT EXISTS {self.files_table_name} (
path TEXT,
source TEXT,
hash TEXT,
@ -112,8 +122,8 @@ class SqliteMemoryStore(BaseMemoryStore):
# Chunks
cursor.execute(
"""
CREATE TABLE IF NOT EXISTS chunks (
f"""
CREATE TABLE IF NOT EXISTS {self.chunks_table_name} (
id TEXT PRIMARY KEY,
path TEXT,
source TEXT,
@ -127,49 +137,34 @@ class SqliteMemoryStore(BaseMemoryStore):
""",
)
cursor.execute(
"""
CREATE INDEX IF NOT EXISTS idx_chunks_path_source
ON chunks(path, source)
""",
)
# Vector table (sqlite-vec)
if self.vector_available:
try:
cursor.execute(
f"""
CREATE VIRTUAL TABLE IF NOT EXISTS {self.VECTOR_TABLE} USING vec0(
id TEXT PRIMARY KEY,
embedding FLOAT[{self.embedding_dim}]
)
""",
cursor.execute(
f"""
CREATE VIRTUAL TABLE IF NOT EXISTS {self.vector_table_name} USING vec0(
id TEXT PRIMARY KEY,
embedding FLOAT[{self.embedding_dim}]
)
logger.info(f"Created vector table (dims={self.embedding_dim})")
except Exception as e:
logger.warning(f"Failed to create vector table: {e}")
self.vector_available = False
""",
)
logger.info(f"Created vector table (dims={self.embedding_dim})")
# FTS table
if self.fts_enabled:
try:
cursor.execute(
f"""
CREATE VIRTUAL TABLE IF NOT EXISTS {self.FTS_TABLE} USING fts5(
text,
id UNINDEXED,
path UNINDEXED,
source UNINDEXED,
start_line UNINDEXED,
end_line UNINDEXED
)
""",
cursor.execute(
f"""
CREATE VIRTUAL TABLE IF NOT EXISTS {self.fts_table_name} USING fts5(
text,
id UNINDEXED,
path UNINDEXED,
source UNINDEXED,
start_line UNINDEXED,
end_line UNINDEXED
)
self.fts_available = True
logger.info("Created FTS5 table")
except Exception as e:
logger.warning(f"Failed to create FTS table: {e}")
self.fts_available = False
""",
)
self.fts_available = True
logger.info("Created FTS5 table")
self.conn.commit()
cursor.close()
@ -180,12 +175,11 @@ class SqliteMemoryStore(BaseMemoryStore):
try:
cursor.execute("BEGIN")
await self._delete_file_internal(cursor, file_meta.path, source)
# Insert file
cursor.execute(
"""
INSERT OR REPLACE INTO files (path, source, hash, mtime, size)
f"""
INSERT OR REPLACE INTO {self.files_table_name} (path, source, hash, mtime, size)
VALUES (?, ?, ?, ?, ?)
""",
(file_meta.path, source.value, file_meta.hash, file_meta.mtime_ms, file_meta.size),
@ -195,8 +189,8 @@ class SqliteMemoryStore(BaseMemoryStore):
now = int(time.time() * 1000)
for chunk in chunks:
cursor.execute(
"""
INSERT INTO chunks (
f"""
INSERT OR REPLACE INTO {self.chunks_table_name} (
id, path, source, start_line, end_line,
hash, text, embedding, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
@ -215,38 +209,33 @@ class SqliteMemoryStore(BaseMemoryStore):
)
# Insert vector
if self.vector_available and chunk.embedding:
try:
cursor.execute(
f"""
INSERT INTO {self.VECTOR_TABLE} (id, embedding)
VALUES (?, ?)
""",
(chunk.id, self.vector_to_blob(chunk.embedding)),
)
except Exception as e:
logger.debug(f"Vector insert failed: {e}")
if self.vector_available:
assert chunk.embedding, "Embedding is required for vector insert"
cursor.execute(
f"""
INSERT OR REPLACE INTO {self.vector_table_name} (id, embedding)
VALUES (?, ?)
""",
(chunk.id, self.vector_to_blob(chunk.embedding)),
)
# Insert FTS
if self.fts_available:
try:
cursor.execute(
f"""
INSERT INTO {self.FTS_TABLE} (
text, id, path, source, start_line, end_line
) VALUES (?, ?, ?, ?, ?, ?)
""",
(
chunk.text,
chunk.id,
file_meta.path,
source.value,
chunk.start_line,
chunk.end_line,
),
)
except Exception as e:
logger.debug(f"FTS insert failed: {e}")
cursor.execute(
f"""
INSERT OR REPLACE INTO {self.fts_table_name} (
text, id, path, source, start_line, end_line
) VALUES (?, ?, ?, ?, ?, ?)
""",
(
chunk.text,
chunk.id,
file_meta.path,
source.value,
chunk.start_line,
chunk.end_line,
),
)
cursor.execute("COMMIT")
except Exception:
@ -255,12 +244,50 @@ class SqliteMemoryStore(BaseMemoryStore):
finally:
cursor.close()
async def delete_file(self, path: str, source: MemorySource) -> None:
async def delete_file(self, path: str, source: MemorySource):
"""Delete file and all its chunks."""
cursor = self.conn.cursor()
try:
cursor.execute("BEGIN")
await self._delete_file_internal(cursor, path, source)
# Get chunk IDs for vector deletion
cursor.execute(
f"SELECT id FROM {self.chunks_table_name} WHERE path = ? AND source = ?",
(path, source.value),
)
chunk_ids = [row[0] for row in cursor.fetchall()]
# Delete vectors
if self.vector_available and chunk_ids:
for chunk_id in chunk_ids:
try:
cursor.execute(
f"DELETE FROM {self.vector_table_name} WHERE id = ?",
(chunk_id,),
)
except Exception as e:
logger.debug(f"Vector delete failed: {e}")
# Delete FTS entries
if self.fts_available:
try:
cursor.execute(
f"DELETE FROM {self.fts_table_name} WHERE path = ? AND source = ?",
(path, source.value),
)
except Exception as e:
logger.debug(f"FTS delete failed: {e}")
# Delete chunks and file
cursor.execute(
f"DELETE FROM {self.chunks_table_name} WHERE path = ? AND source = ?",
(path, source.value),
)
cursor.execute(
f"DELETE FROM {self.files_table_name} WHERE path = ? AND source = ?",
(path, source.value),
)
cursor.execute("COMMIT")
except Exception:
cursor.execute("ROLLBACK")
@ -268,62 +295,132 @@ class SqliteMemoryStore(BaseMemoryStore):
finally:
cursor.close()
async def _delete_file_internal(self, cursor: sqlite3.Cursor, path: str, source: MemorySource):
"""Internal delete helper."""
# Get chunk IDs for vector deletion
cursor.execute(
"SELECT id FROM chunks WHERE path = ? AND source = ?",
(path, source.value),
)
chunk_ids = [row[0] for row in cursor.fetchall()]
async def delete_file_chunks(self, path: str, chunk_ids: list[str]):
"""Delete specific chunks for a file."""
if not chunk_ids:
return
# Delete vectors
if self.vector_available and chunk_ids:
for chunk_id in chunk_ids:
cursor = self.conn.cursor()
try:
cursor.execute("BEGIN")
# Delete vectors
if self.vector_available:
for chunk_id in chunk_ids:
try:
cursor.execute(
f"DELETE FROM {self.vector_table_name} WHERE id = ?",
(chunk_id,),
)
except Exception as e:
logger.debug(f"Vector delete failed for {chunk_id}: {e}")
# Delete FTS entries
if self.fts_available:
placeholders = ",".join("?" * len(chunk_ids))
try:
cursor.execute(
f"DELETE FROM {self.VECTOR_TABLE} WHERE id = ?",
(chunk_id,),
f"DELETE FROM {self.fts_table_name} WHERE id IN ({placeholders})",
chunk_ids,
)
except Exception as e:
logger.debug(f"Vector delete failed: {e}")
logger.debug(f"FTS delete failed: {e}")
# Delete FTS entries
if self.fts_available:
try:
cursor.execute(
f"DELETE FROM {self.FTS_TABLE} WHERE path = ? AND source = ?",
(path, source.value),
)
except Exception as e:
logger.debug(f"FTS delete failed: {e}")
# Delete chunks
placeholders = ",".join("?" * len(chunk_ids))
cursor.execute(
f"DELETE FROM {self.chunks_table_name} WHERE id IN ({placeholders})",
chunk_ids,
)
# Delete chunks and file
cursor.execute(
"DELETE FROM chunks WHERE path = ? AND source = ?",
(path, source.value),
)
cursor.execute(
"DELETE FROM files WHERE path = ? AND source = ?",
(path, source.value),
)
cursor.execute("COMMIT")
except Exception:
cursor.execute("ROLLBACK")
raise
finally:
cursor.close()
async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource):
"""Insert or update specific chunks without affecting other chunks."""
if not chunks:
return
async def get_file_hash(self, path: str, source: MemorySource) -> str | None:
"""Get file hash."""
cursor = self.conn.cursor()
cursor.execute(
"SELECT hash FROM files WHERE path = ? AND source = ?",
(path, source.value),
)
row = cursor.fetchone()
try:
cursor.execute("BEGIN")
now = int(time.time() * 1000)
for chunk in chunks:
# Insert/update chunk
cursor.execute(
f"""
INSERT OR REPLACE INTO {self.chunks_table_name} (
id, path, source, start_line, end_line,
hash, text, embedding, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
chunk.id,
chunk.path,
source.value,
chunk.start_line,
chunk.end_line,
chunk.hash,
chunk.text,
json.dumps(chunk.embedding) if chunk.embedding else None,
now,
),
)
# Insert/update vector
if self.vector_available:
assert chunk.embedding, "Embedding is required for vector insert"
cursor.execute(
f"""
INSERT OR REPLACE INTO {self.vector_table_name} (id, embedding)
VALUES (?, ?)
""",
(chunk.id, self.vector_to_blob(chunk.embedding)),
)
# Insert/update FTS
if self.fts_available:
cursor.execute(
f"""
INSERT OR REPLACE INTO {self.fts_table_name} (
text, id, path, source, start_line, end_line
) VALUES (?, ?, ?, ?, ?, ?)
""",
(
chunk.text,
chunk.id,
chunk.path,
source.value,
chunk.start_line,
chunk.end_line,
),
)
cursor.execute("COMMIT")
except Exception:
cursor.execute("ROLLBACK")
raise
finally:
cursor.close()
async def list_files(self, source: MemorySource) -> list[str]:
"""List all indexed files."""
cursor = self.conn.cursor()
cursor.execute(f"SELECT path FROM {self.files_table_name} WHERE source = ?", (source.value,))
paths = [row[0] for row in cursor.fetchall()]
cursor.close()
return row[0] if row else None
return paths
async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None:
"""Get file metadata with chunk count."""
cursor = self.conn.cursor()
cursor.execute(
"SELECT hash, mtime, size FROM files WHERE path = ? AND source = ?",
f"SELECT hash, mtime, size FROM {self.files_table_name} WHERE path = ? AND source = ?",
(path, source.value),
)
row = cursor.fetchone()
@ -333,7 +430,7 @@ class SqliteMemoryStore(BaseMemoryStore):
hash_val, mtime, size = row
cursor.execute(
"SELECT COUNT(*) FROM chunks WHERE path = ? AND source = ?",
f"SELECT COUNT(*) FROM {self.chunks_table_name} WHERE path = ? AND source = ?",
(path, source.value),
)
chunk_count = cursor.fetchone()[0]
@ -343,24 +440,17 @@ class SqliteMemoryStore(BaseMemoryStore):
hash=hash_val,
mtime_ms=mtime,
size=size,
path=path,
chunk_count=chunk_count,
)
async def list_files(self, source: MemorySource) -> list[str]:
"""List all indexed files."""
cursor = self.conn.cursor()
cursor.execute("SELECT path FROM files WHERE source = ?", (source.value,))
paths = [row[0] for row in cursor.fetchall()]
cursor.close()
return paths
async def get_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
"""Get all chunks for a file."""
cursor = self.conn.cursor()
cursor.execute(
"""
f"""
SELECT id, path, source, start_line, end_line, text, hash, embedding
FROM chunks WHERE path = ? AND source = ?
FROM {self.chunks_table_name} WHERE path = ? AND source = ?
ORDER BY start_line
""",
(path, source.value),
@ -418,23 +508,25 @@ class SqliteMemoryStore(BaseMemoryStore):
try:
query_blob = self.vector_to_blob(query_embedding)
# Correct SQLite-vec syntax for vector search with limit
# vec0 requires 'k = ?' constraint for knn queries
query_sql = f"""
SELECT c.id, c.path, c.start_line, c.end_line, c.source, c.text, v.distance
FROM {self.VECTOR_TABLE} v
JOIN chunks c ON v.id = c.id
FROM {self.vector_table_name} v
JOIN {self.chunks_table_name} c ON v.id = c.id
WHERE v.embedding MATCH ?
AND k = ?
"""
query_params: list = [query_blob]
query_params: list = [query_blob, limit]
# Add source filter if specified
if source_filter:
query_sql += source_filter
query_params.extend(params)
# Order and limit results
query_sql += " ORDER BY v.distance LIMIT ?"
query_params.append(str(limit))
# Order by distance (k constraint already limits results)
query_sql += " ORDER BY v.distance"
cursor.execute(query_sql, query_params)
@ -470,11 +562,21 @@ class SqliteMemoryStore(BaseMemoryStore):
if not self.fts_available:
return []
# Build FTS5 query, escaping quotes
cleaned = query.strip().replace('"', '""')
# Build FTS5 query
# Split query into tokens and join with OR for better recall
# Individual words are automatically stemmed and matched by FTS5
cleaned = query.strip()
if not cleaned:
return []
fts_query = f'"{cleaned}"'
# Split into words and escape each
words = cleaned.split()
if not words:
return []
# Use OR operator for better recall - match any of the query words
escaped_words = [word.replace('"', '""') for word in words]
fts_query = " OR ".join(escaped_words)
cursor = self.conn.cursor()
source_filter = ""
@ -490,7 +592,7 @@ class SqliteMemoryStore(BaseMemoryStore):
f"""
SELECT fts.id, fts.path, fts.start_line, fts.end_line,
fts.source, fts.text, rank
FROM {self.FTS_TABLE} fts
FROM {self.fts_table_name} fts
WHERE fts.text MATCH ?{source_filter}
ORDER BY rank
LIMIT ?
@ -521,52 +623,20 @@ class SqliteMemoryStore(BaseMemoryStore):
finally:
cursor.close()
async def read_meta(self, key: str) -> MemoryIndexMeta | None:
"""Read metadata value."""
cursor = self.conn.cursor()
cursor.execute("SELECT value FROM meta WHERE key = ?", (key,))
row = cursor.fetchone()
cursor.close()
if not row:
return None
return MemoryIndexMeta(**json.loads(row[0]))
async def write_meta(self, key: str, value: MemoryIndexMeta | dict) -> None:
"""Write metadata value."""
data = value.model_dump() if isinstance(value, MemoryIndexMeta) else value
cursor = self.conn.cursor()
cursor.execute(
"""
INSERT OR REPLACE INTO meta (key, value)
VALUES (?, ?)
""",
(key, json.dumps(data)),
)
self.conn.commit()
cursor.close()
async def clear_all(self):
"""Clear all indexed data."""
cursor = self.conn.cursor()
cursor.execute("BEGIN")
try:
cursor.execute("DELETE FROM files")
cursor.execute("DELETE FROM chunks")
cursor.execute(f"DELETE FROM {self.files_table_name}")
cursor.execute(f"DELETE FROM {self.chunks_table_name}")
if self.vector_available:
try:
cursor.execute(f"DELETE FROM {self.VECTOR_TABLE}")
except Exception as e:
logger.debug(f"Vector clear failed: {e}")
cursor.execute(f"DELETE FROM {self.vector_table_name}")
if self.fts_available:
try:
cursor.execute(f"DELETE FROM {self.FTS_TABLE}")
except Exception as e:
logger.debug(f"FTS clear failed: {e}")
cursor.execute(f"DELETE FROM {self.fts_table_name}")
cursor.execute("COMMIT")
except Exception:

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

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

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

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

@ -1,19 +1,17 @@
"""Chunking logic for Markdown files."""
from typing import List, Dict, Any
from ..utils.hashing import hash_text
from ...enumeration import MemorySource
from ...schema import MemoryChunk
from .common_utils import hash_text
from ..enumeration import MemorySource
from ..schema import MemoryChunk
def chunk_markdown(
text: str,
path: str,
source: MemorySource,
chunk_tokens: int = 300,
overlap: int = 30,
) -> List[MemoryChunk]:
chunk_tokens: int,
overlap: int,
) -> list[MemoryChunk]:
"""
Markdown chunking logic implemented based on the TypeScript version.
@ -35,10 +33,10 @@ def chunk_markdown(
max_chars = max(32, chunk_tokens * 4)
overlap_chars = max(0, overlap * 4)
chunks: List[MemoryChunk] = []
chunks: list[MemoryChunk] = []
# Currently building chunk
current: List[Dict[str, Any]] = [] # [{'line': str, 'line_no': int}]
current: list[dict] = [] # [{'line': str, 'line_no': int}]
current_chars = 0
def flush():
@ -83,8 +81,8 @@ def chunk_markdown(
kept = []
# Collect lines from the end until reaching overlap size
for i in range(len(current) - 1, -1, -1):
entry = current[i]
for j in range(len(current) - 1, -1, -1):
entry = current[j]
if not entry:
continue
@ -123,4 +121,4 @@ def chunk_markdown(
# Process the final chunk
flush()
return chunks
return [c for c in chunks if c.text.strip()]

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

@ -58,7 +58,7 @@ class ReMe(Application):
target_user_names: list[str] | None = None,
target_task_names: list[str] | None = None,
target_tool_names: list[str] | None = None,
profile_dir: str = "reme_profile",
profile_dir: str = ".reme/profile",
**kwargs,
):
"""Initialize ReMe with config.

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

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

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