mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
commit
90cd9ddea8
57 changed files with 7322 additions and 70 deletions
4
.gitignore
vendored
4
.gitignore
vendored
|
|
@ -39,4 +39,6 @@ chroma_vector_store/*
|
|||
bench_results/*
|
||||
meta_memory/*
|
||||
*.sqlite3
|
||||
**/data/*.json
|
||||
**/data/*.json
|
||||
*.db
|
||||
memories/*
|
||||
|
|
@ -78,7 +78,7 @@ repos:
|
|||
--disable=C3001,
|
||||
--disable=R1702,
|
||||
--disable=R0912,
|
||||
--max-statements=75,
|
||||
--max-statements=120,
|
||||
--max-line-length=120,
|
||||
]
|
||||
- repo: https://github.com/regebro/pyroma
|
||||
|
|
|
|||
|
|
@ -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/*
|
||||
|
|
|
|||
9
reme/agent/fs/__init__.py
Normal file
9
reme/agent/fs/__init__.py
Normal file
|
|
@ -0,0 +1,9 @@
|
|||
"""File system agents for memory management."""
|
||||
|
||||
from .fs_compactor import FsCompactor
|
||||
from .fs_summarizer import FsSummarizer
|
||||
|
||||
__all__ = [
|
||||
"FsSummarizer",
|
||||
"FsCompactor",
|
||||
]
|
||||
245
reme/agent/fs/fs_compactor.py
Normal file
245
reme/agent/fs/fs_compactor.py
Normal 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,
|
||||
}
|
||||
212
reme/agent/fs/fs_compactor.yaml
Normal file
212
reme/agent/fs/fs_compactor.yaml
Normal 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>
|
||||
|
||||
这是一个过长而无法保留的回合的前缀。后缀(最近的工作)已保留。
|
||||
|
||||
总结前缀以为保留的后缀提供上下文:
|
||||
|
||||
## 原始请求
|
||||
[用户在此回合中要求了什么?]
|
||||
|
||||
## 早期进展
|
||||
- [在前缀中做出的关键决策和完成的工作]
|
||||
|
||||
## 后缀上下文
|
||||
- [理解保留的最近工作所需的信息]
|
||||
|
||||
保持简洁。专注于理解保留后缀所需的内容。
|
||||
81
reme/agent/fs/fs_summarizer.py
Normal file
81
reme/agent/fs/fs_summarizer.py
Normal 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
|
||||
15
reme/agent/fs/fs_summarizer.yaml
Normal file
15
reme/agent/fs/fs_summarizer.yaml
Normal 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].
|
||||
|
|
@ -66,7 +66,7 @@ class BaseMemoryAgent(BaseReact, metaclass=ABCMeta):
|
|||
lines = []
|
||||
for memory_target, memory_type in self.memory_target_type_mapping.items():
|
||||
line = {
|
||||
"agent": f"Agent managing {memory_type} memories for {memory_target}",
|
||||
"agent": f"Agent managing {memory_type.value} memories for {memory_target}",
|
||||
"memory_target": memory_target,
|
||||
}
|
||||
lines.append(json.dumps(line, ensure_ascii=False))
|
||||
|
|
|
|||
|
|
@ -68,25 +68,27 @@ class ReMeRetriever(BaseMemoryAgent):
|
|||
async def execute(self):
|
||||
result = await super().execute()
|
||||
tools: list[BaseTool] = result["tools"]
|
||||
delegate_task_tool = tools[0]
|
||||
agents: list[BaseMemoryAgent] = delegate_task_tool.response.metadata["agents"]
|
||||
|
||||
answer = []
|
||||
success = True
|
||||
messages = []
|
||||
tools = []
|
||||
tools_result = []
|
||||
retrieved_nodes = []
|
||||
for agent in agents:
|
||||
answer.append(agent.response.answer)
|
||||
success = success and agent.response.success
|
||||
messages.extend(agent.response.metadata["messages"])
|
||||
tools.extend(agent.response.metadata["tools"])
|
||||
retrieved_nodes.extend(agent.response.metadata["retrieved_nodes"])
|
||||
|
||||
if tools:
|
||||
delegate_task_tool = tools[0]
|
||||
agents: list[BaseMemoryAgent] = delegate_task_tool.response.metadata["agents"]
|
||||
for agent in agents:
|
||||
answer.append(agent.response.answer)
|
||||
success = success and agent.response.success
|
||||
messages.extend(agent.response.metadata["messages"])
|
||||
tools_result.extend(agent.response.metadata["tools"])
|
||||
retrieved_nodes.extend(agent.response.metadata["retrieved_nodes"])
|
||||
|
||||
return {
|
||||
"answer": "\n".join(answer),
|
||||
"success": True,
|
||||
"messages": messages,
|
||||
"tools": tools,
|
||||
"tools": tools_result,
|
||||
"retrieved_nodes": retrieved_nodes,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -74,22 +74,25 @@ class ReMeSummarizer(BaseMemoryAgent):
|
|||
async def execute(self):
|
||||
result = await super().execute()
|
||||
tools: list[BaseTool] = result["tools"]
|
||||
delegate_task_tool = tools[0]
|
||||
agents: list[BaseMemoryAgent] = delegate_task_tool.response.metadata["agents"]
|
||||
|
||||
success = True
|
||||
messages = []
|
||||
tools = []
|
||||
tools_result = []
|
||||
memory_nodes = []
|
||||
for agent in agents:
|
||||
success = success and agent.response.success
|
||||
messages.extend(agent.response.metadata["messages"])
|
||||
tools.extend(agent.response.metadata["tools"])
|
||||
memory_nodes.extend(agent.response.metadata["memory_nodes"])
|
||||
|
||||
if tools:
|
||||
delegate_task_tool = tools[0]
|
||||
agents: list[BaseMemoryAgent] = delegate_task_tool.response.metadata["agents"]
|
||||
|
||||
for agent in agents:
|
||||
success = success and agent.response.success
|
||||
messages.extend(agent.response.metadata["messages"])
|
||||
tools_result.extend(agent.response.metadata["tools"])
|
||||
memory_nodes.extend(agent.response.metadata["memory_nodes"])
|
||||
|
||||
return {
|
||||
"answer": memory_nodes,
|
||||
"success": True,
|
||||
"messages": messages,
|
||||
"tools": tools,
|
||||
"tools": tools_result,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -4,12 +4,15 @@ Defines the abstract base class and standard API for all embedding model impleme
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import time
|
||||
from abc import ABC
|
||||
from collections import OrderedDict
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from ..schema import VectorNode
|
||||
from ..schema.memory_chunk import MemoryChunk
|
||||
|
||||
|
||||
class BaseEmbeddingModel(ABC):
|
||||
|
|
@ -27,17 +30,35 @@ class BaseEmbeddingModel(ABC):
|
|||
max_retries: int = 3,
|
||||
raise_exception: bool = True,
|
||||
max_input_length: int = 8192,
|
||||
max_cache_size: int = 10000,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize model configuration and parameters."""
|
||||
"""Initialize model configuration and parameters.
|
||||
|
||||
Args:
|
||||
model_name: Name of the embedding model
|
||||
dimensions: Vector dimensions of the embeddings
|
||||
max_batch_size: Maximum batch size for embedding requests
|
||||
max_retries: Maximum number of retry attempts on failure
|
||||
raise_exception: Whether to raise exceptions on failure
|
||||
max_input_length: Maximum input text length
|
||||
max_cache_size: Maximum number of embeddings to cache in memory (LRU)
|
||||
**kwargs: Additional model-specific parameters
|
||||
"""
|
||||
self.model_name = model_name
|
||||
self.dimensions = dimensions
|
||||
self.max_batch_size = max_batch_size
|
||||
self.max_retries = max_retries
|
||||
self.raise_exception = raise_exception
|
||||
self.max_input_length = max_input_length
|
||||
self.max_cache_size = max_cache_size
|
||||
self.kwargs = kwargs
|
||||
|
||||
# Initialize LRU cache for embeddings
|
||||
self._embedding_cache: OrderedDict[str, list[float]] = OrderedDict()
|
||||
self._cache_hits = 0
|
||||
self._cache_misses = 0
|
||||
|
||||
def _truncate_text(self, text: str) -> str:
|
||||
"""Truncate text to max_input_length if it exceeds the limit."""
|
||||
if len(text) > self.max_input_length:
|
||||
|
|
@ -51,6 +72,76 @@ class BaseEmbeddingModel(ABC):
|
|||
"""Truncate a list of texts to max_input_length."""
|
||||
return [self._truncate_text(text) for text in texts]
|
||||
|
||||
def _get_cache_key(self, text: str) -> str:
|
||||
"""Generate a cache key by hashing the input text.
|
||||
|
||||
Args:
|
||||
text: Input text to hash
|
||||
|
||||
Returns:
|
||||
SHA256 hash of the text as hexadecimal string
|
||||
"""
|
||||
return hashlib.sha256(text.encode("utf-8")).hexdigest()
|
||||
|
||||
def _get_from_cache(self, text: str) -> list[float] | None:
|
||||
"""Retrieve embedding from cache if it exists.
|
||||
|
||||
Args:
|
||||
text: Input text to look up
|
||||
|
||||
Returns:
|
||||
Cached embedding vector or None if not found
|
||||
"""
|
||||
cache_key = self._get_cache_key(text)
|
||||
if cache_key in self._embedding_cache:
|
||||
# Move to end (most recently used)
|
||||
self._embedding_cache.move_to_end(cache_key)
|
||||
self._cache_hits += 1
|
||||
return self._embedding_cache[cache_key]
|
||||
self._cache_misses += 1
|
||||
return None
|
||||
|
||||
def _put_to_cache(self, text: str, embedding: list[float]) -> None:
|
||||
"""Store embedding in cache with LRU eviction.
|
||||
|
||||
Args:
|
||||
text: Input text used as cache key
|
||||
embedding: Embedding vector to cache
|
||||
"""
|
||||
if self.max_cache_size <= 0:
|
||||
return
|
||||
|
||||
cache_key = self._get_cache_key(text)
|
||||
|
||||
# Remove oldest entry if cache is full
|
||||
if len(self._embedding_cache) >= self.max_cache_size and cache_key not in self._embedding_cache:
|
||||
self._embedding_cache.popitem(last=False)
|
||||
|
||||
self._embedding_cache[cache_key] = embedding
|
||||
self._embedding_cache.move_to_end(cache_key)
|
||||
|
||||
def get_cache_stats(self) -> dict[str, int]:
|
||||
"""Get cache statistics.
|
||||
|
||||
Returns:
|
||||
Dictionary with cache size, hits, misses, and hit rate
|
||||
"""
|
||||
total_requests = self._cache_hits + self._cache_misses
|
||||
hit_rate = self._cache_hits / total_requests if total_requests > 0 else 0.0
|
||||
return {
|
||||
"cache_size": len(self._embedding_cache),
|
||||
"max_cache_size": self.max_cache_size,
|
||||
"cache_hits": self._cache_hits,
|
||||
"cache_misses": self._cache_misses,
|
||||
"hit_rate": hit_rate,
|
||||
}
|
||||
|
||||
def clear_cache(self) -> None:
|
||||
"""Clear the embedding cache and reset statistics."""
|
||||
self._embedding_cache.clear()
|
||||
self._cache_hits = 0
|
||||
self._cache_misses = 0
|
||||
|
||||
async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]:
|
||||
"""Internal async implementation for calling the embedding API with batch input."""
|
||||
|
||||
|
|
@ -60,10 +151,20 @@ class BaseEmbeddingModel(ABC):
|
|||
async def get_embedding(self, input_text: str, **kwargs) -> list[float]:
|
||||
"""Async get embedding for a single text with exponential backoff retries."""
|
||||
truncated_text = self._truncate_text(input_text)
|
||||
|
||||
# Check cache first
|
||||
cached_embedding = self._get_from_cache(truncated_text)
|
||||
if cached_embedding is not None:
|
||||
return cached_embedding
|
||||
|
||||
# Cache miss - compute embedding
|
||||
for i in range(self.max_retries):
|
||||
try:
|
||||
result = await self._get_embeddings([truncated_text], **kwargs)
|
||||
return result[0]
|
||||
embedding = result[0]
|
||||
# Store in cache
|
||||
self._put_to_cache(truncated_text, embedding)
|
||||
return embedding
|
||||
except Exception as e:
|
||||
logger.error(f"Model {self.model_name} failed: {e}")
|
||||
if i == self.max_retries - 1:
|
||||
|
|
@ -78,16 +179,36 @@ class BaseEmbeddingModel(ABC):
|
|||
# Truncate all input texts first
|
||||
truncated_texts = self._truncate_texts(input_text)
|
||||
|
||||
# Split into batches and process sequentially to respect rate limits
|
||||
results = []
|
||||
for i in range(0, len(truncated_texts), self.max_batch_size):
|
||||
batch = truncated_texts[i : i + self.max_batch_size]
|
||||
# Check cache for each text and separate cached vs uncached
|
||||
results: list[list[float] | None] = [None] * len(truncated_texts)
|
||||
texts_to_compute: list[tuple[int, str]] = [] # (original_index, text)
|
||||
|
||||
for idx, text in enumerate(truncated_texts):
|
||||
cached = self._get_from_cache(text)
|
||||
if cached is not None:
|
||||
results[idx] = cached
|
||||
else:
|
||||
texts_to_compute.append((idx, text))
|
||||
|
||||
# If all texts were cached, return early
|
||||
if not texts_to_compute:
|
||||
return [r for r in results if r is not None]
|
||||
|
||||
# Compute embeddings for uncached texts in batches
|
||||
uncached_texts = [text for _, text in texts_to_compute]
|
||||
for i in range(0, len(uncached_texts), self.max_batch_size):
|
||||
batch_texts = uncached_texts[i : i + self.max_batch_size]
|
||||
batch_indices = [idx for idx, _ in texts_to_compute[i : i + self.max_batch_size]]
|
||||
|
||||
# Process each batch with retry logic
|
||||
for retry in range(self.max_retries):
|
||||
try:
|
||||
batch_res = await self._get_embeddings(batch, **kwargs)
|
||||
if batch_res:
|
||||
results.extend(batch_res)
|
||||
batch_embeddings = await self._get_embeddings(batch_texts, **kwargs)
|
||||
if batch_embeddings:
|
||||
# Store results and cache them
|
||||
for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings):
|
||||
results[orig_idx] = embedding
|
||||
self._put_to_cache(text, embedding)
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Model {self.model_name} batch failed: {e}")
|
||||
|
|
@ -96,15 +217,26 @@ class BaseEmbeddingModel(ABC):
|
|||
raise
|
||||
else:
|
||||
await asyncio.sleep(retry + 1)
|
||||
return results
|
||||
|
||||
return [r for r in results if r is not None]
|
||||
|
||||
def get_embedding_sync(self, input_text: str, **kwargs) -> list[float]:
|
||||
"""Synchronous get embedding for a single text with retry logic."""
|
||||
truncated_text = self._truncate_text(input_text)
|
||||
|
||||
# Check cache first
|
||||
cached_embedding = self._get_from_cache(truncated_text)
|
||||
if cached_embedding is not None:
|
||||
return cached_embedding
|
||||
|
||||
# Cache miss - compute embedding
|
||||
for i in range(self.max_retries):
|
||||
try:
|
||||
result = self._get_embeddings_sync([truncated_text], **kwargs)
|
||||
return result[0]
|
||||
embedding = result[0]
|
||||
# Store in cache
|
||||
self._put_to_cache(truncated_text, embedding)
|
||||
return embedding
|
||||
except Exception as exc:
|
||||
logger.error(f"Model {self.model_name} failed: {exc}")
|
||||
if i == self.max_retries - 1:
|
||||
|
|
@ -119,15 +251,36 @@ class BaseEmbeddingModel(ABC):
|
|||
# Truncate all input texts first
|
||||
truncated_texts = self._truncate_texts(input_text)
|
||||
|
||||
results = []
|
||||
for i in range(0, len(truncated_texts), self.max_batch_size):
|
||||
batch = truncated_texts[i : i + self.max_batch_size]
|
||||
# Check cache for each text and separate cached vs uncached
|
||||
results: list[list[float] | None] = [None] * len(truncated_texts)
|
||||
texts_to_compute: list[tuple[int, str]] = [] # (original_index, text)
|
||||
|
||||
for idx, text in enumerate(truncated_texts):
|
||||
cached = self._get_from_cache(text)
|
||||
if cached is not None:
|
||||
results[idx] = cached
|
||||
else:
|
||||
texts_to_compute.append((idx, text))
|
||||
|
||||
# If all texts were cached, return early
|
||||
if not texts_to_compute:
|
||||
return [r for r in results if r is not None]
|
||||
|
||||
# Compute embeddings for uncached texts in batches
|
||||
uncached_texts = [text for _, text in texts_to_compute]
|
||||
for i in range(0, len(uncached_texts), self.max_batch_size):
|
||||
batch_texts = uncached_texts[i : i + self.max_batch_size]
|
||||
batch_indices = [idx for idx, _ in texts_to_compute[i : i + self.max_batch_size]]
|
||||
|
||||
# Process each batch with retry logic
|
||||
for retry in range(self.max_retries):
|
||||
try:
|
||||
batch_res = self._get_embeddings_sync(batch, **kwargs)
|
||||
if batch_res:
|
||||
results.extend(batch_res)
|
||||
batch_embeddings = self._get_embeddings_sync(batch_texts, **kwargs)
|
||||
if batch_embeddings:
|
||||
# Store results and cache them
|
||||
for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings):
|
||||
results[orig_idx] = embedding
|
||||
self._put_to_cache(text, embedding)
|
||||
break
|
||||
except Exception as exc:
|
||||
logger.error(f"Model {self.model_name} batch failed: {exc}")
|
||||
|
|
@ -136,7 +289,8 @@ class BaseEmbeddingModel(ABC):
|
|||
raise
|
||||
else:
|
||||
time.sleep(retry + 1)
|
||||
return results
|
||||
|
||||
return [r for r in results if r is not None]
|
||||
|
||||
async def get_node_embedding(self, node: VectorNode, **kwargs) -> VectorNode:
|
||||
"""Async generate and populate vector field for a single VectorNode object."""
|
||||
|
|
@ -172,6 +326,72 @@ class BaseEmbeddingModel(ABC):
|
|||
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(nodes)} nodes")
|
||||
return nodes
|
||||
|
||||
async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
|
||||
"""Async generate and populate embedding field for a single MemoryChunk object.
|
||||
|
||||
Args:
|
||||
chunk: MemoryChunk object containing text to embed
|
||||
**kwargs: Additional arguments passed to the embedding model
|
||||
|
||||
Returns:
|
||||
The same MemoryChunk object with populated embedding field
|
||||
"""
|
||||
chunk.embedding = await self.get_embedding(chunk.text, **kwargs)
|
||||
return chunk
|
||||
|
||||
async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
|
||||
"""Async generate and populate embedding fields for a batch of MemoryChunk objects.
|
||||
|
||||
Args:
|
||||
chunks: List of MemoryChunk objects containing text to embed
|
||||
**kwargs: Additional arguments passed to the embedding model
|
||||
|
||||
Returns:
|
||||
The same list of MemoryChunk objects with populated embedding fields
|
||||
"""
|
||||
texts = [chunk.text for chunk in chunks]
|
||||
embeddings: list[list[float]] = await self.get_embeddings(texts, **kwargs)
|
||||
|
||||
if len(embeddings) == len(chunks):
|
||||
for chunk, vec in zip(chunks, embeddings):
|
||||
chunk.embedding = vec
|
||||
else:
|
||||
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(chunks)} chunks")
|
||||
return chunks
|
||||
|
||||
def get_chunk_embedding_sync(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
|
||||
"""Synchronously generate and populate embedding field for a single MemoryChunk object.
|
||||
|
||||
Args:
|
||||
chunk: MemoryChunk object containing text to embed
|
||||
**kwargs: Additional arguments passed to the embedding model
|
||||
|
||||
Returns:
|
||||
The same MemoryChunk object with populated embedding field
|
||||
"""
|
||||
chunk.embedding = self.get_embedding_sync(chunk.text, **kwargs)
|
||||
return chunk
|
||||
|
||||
def get_chunk_embeddings_sync(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
|
||||
"""Synchronously generate embeddings for a batch of MemoryChunk objects.
|
||||
|
||||
Args:
|
||||
chunks: List of MemoryChunk objects containing text to embed
|
||||
**kwargs: Additional arguments passed to the embedding model
|
||||
|
||||
Returns:
|
||||
The same list of MemoryChunk objects with populated embedding fields
|
||||
"""
|
||||
texts = [chunk.text for chunk in chunks]
|
||||
embeddings: list[list[float]] = self.get_embeddings_sync(texts, **kwargs)
|
||||
|
||||
if len(embeddings) == len(chunks):
|
||||
for chunk, vec in zip(chunks, embeddings):
|
||||
chunk.embedding = vec
|
||||
else:
|
||||
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(chunks)} chunks")
|
||||
return chunks
|
||||
|
||||
def close_sync(self):
|
||||
"""Synchronously release resources and close connections."""
|
||||
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@
|
|||
from .chunk_enum import ChunkEnum
|
||||
from .http_enum import HttpEnum
|
||||
from .json_schema_enum import JsonSchemaEnum
|
||||
from .memory_source import MemorySource
|
||||
from .memory_type import MemoryType
|
||||
from .registry_enum import RegistryEnum
|
||||
from .role import Role
|
||||
|
|
@ -11,6 +12,7 @@ __all__ = [
|
|||
"ChunkEnum",
|
||||
"HttpEnum",
|
||||
"JsonSchemaEnum",
|
||||
"MemorySource",
|
||||
"MemoryType",
|
||||
"RegistryEnum",
|
||||
"Role",
|
||||
|
|
|
|||
11
reme/core/enumeration/memory_source.py
Normal file
11
reme/core/enumeration/memory_source.py
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
"""Memory source types."""
|
||||
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class MemorySource(str, Enum):
|
||||
"""Source of memory data."""
|
||||
|
||||
MEMORY = "memory"
|
||||
|
||||
SESSIONS = "sessions"
|
||||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
19
reme/core/file_watcher/__init__.py
Normal file
19
reme/core/file_watcher/__init__.py
Normal 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)
|
||||
126
reme/core/file_watcher/base_file_watcher.py
Normal file
126
reme/core/file_watcher/base_file_watcher.py
Normal 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()
|
||||
277
reme/core/file_watcher/delta_file_watcher.py
Normal file
277
reme/core/file_watcher/delta_file_watcher.py
Normal 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
|
||||
76
reme/core/file_watcher/full_file_watcher.py
Normal file
76
reme/core/file_watcher/full_file_watcher.py
Normal 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
|
||||
16
reme/core/memory_storage/__init__.py
Normal file
16
reme/core/memory_storage/__init__.py
Normal 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)
|
||||
126
reme/core/memory_storage/base_memory_store.py
Normal file
126
reme/core/memory_storage/base_memory_store.py
Normal file
|
|
@ -0,0 +1,126 @@
|
|||
"""Base storage interface for memory manager."""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from ..embedding import BaseEmbeddingModel
|
||||
from ..enumeration import MemorySource
|
||||
from ..schema import FileMetadata, MemoryChunk, MemorySearchResult
|
||||
|
||||
|
||||
class BaseMemoryStore(ABC):
|
||||
"""Abstract base class for memory storage backends."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
store_name: str,
|
||||
embedding_model: BaseEmbeddingModel,
|
||||
fts_enabled: bool = True,
|
||||
snippet_max_chars: int = 700,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize"""
|
||||
self.store_name: str = store_name
|
||||
self.embedding_model: BaseEmbeddingModel = embedding_model
|
||||
self.fts_enabled: bool = fts_enabled
|
||||
self.snippet_max_chars: int = snippet_max_chars
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
self.vector_available = False
|
||||
self.fts_available = False
|
||||
|
||||
@property
|
||||
def embedding_dim(self) -> int:
|
||||
"""Get the embedding model's dimensionality."""
|
||||
return self.embedding_model.dimensions
|
||||
|
||||
async def get_embedding(self, query: str, **kwargs) -> list[float]:
|
||||
"""Get embedding for a single query string."""
|
||||
return await self.embedding_model.get_embedding(query, **kwargs)
|
||||
|
||||
async def get_embeddings(self, queries: list[str], **kwargs) -> list[list[float]]:
|
||||
"""Get embeddings for a batch of query strings."""
|
||||
return await self.embedding_model.get_embeddings(queries, **kwargs)
|
||||
|
||||
async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
|
||||
"""Generate and populate embedding field for a single MemoryChunk object."""
|
||||
return await self.embedding_model.get_chunk_embedding(chunk, **kwargs)
|
||||
|
||||
async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
|
||||
"""Generate and populate embedding fields for a batch of MemoryChunk objects."""
|
||||
return await self.embedding_model.get_chunk_embeddings(chunks, **kwargs)
|
||||
|
||||
@abstractmethod
|
||||
async def start(self):
|
||||
"""Initialize the storage backend."""
|
||||
|
||||
@abstractmethod
|
||||
async def upsert_file(self, file_meta: FileMetadata, source: MemorySource, chunks: list[MemoryChunk]):
|
||||
"""Insert or update a file and its chunks."""
|
||||
|
||||
@abstractmethod
|
||||
async def delete_file(self, path: str, source: MemorySource):
|
||||
"""Delete a file and all its chunks."""
|
||||
|
||||
@abstractmethod
|
||||
async def delete_file_chunks(self, path: str, chunk_ids: list[str]):
|
||||
"""Delete chunks for a file."""
|
||||
|
||||
@abstractmethod
|
||||
async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource):
|
||||
"""Insert or update specific chunks without affecting other chunks."""
|
||||
|
||||
@abstractmethod
|
||||
async def list_files(self, source: MemorySource) -> list[str]:
|
||||
"""List all indexed file paths for a source."""
|
||||
|
||||
@abstractmethod
|
||||
async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None:
|
||||
"""Get full file metadata with statistics."""
|
||||
|
||||
@abstractmethod
|
||||
async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
|
||||
"""Get all chunks for a file."""
|
||||
|
||||
@abstractmethod
|
||||
async def vector_search(
|
||||
self,
|
||||
query: str,
|
||||
limit: int,
|
||||
sources: list[MemorySource] | None = None,
|
||||
) -> list[MemorySearchResult]:
|
||||
"""Perform vector similarity search.
|
||||
|
||||
Args:
|
||||
query: Query embedding vector
|
||||
limit: Maximum number of results
|
||||
sources: Optional list of sources to filter
|
||||
|
||||
Returns:
|
||||
List of search results sorted by similarity
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def keyword_search(
|
||||
self,
|
||||
query: str,
|
||||
limit: int,
|
||||
sources: list[MemorySource] | None = None,
|
||||
) -> list[MemorySearchResult]:
|
||||
"""Perform keyword/full-text search.
|
||||
|
||||
Args:
|
||||
query: Search query text
|
||||
limit: Maximum number of results
|
||||
sources: Optional list of sources to filter
|
||||
|
||||
Returns:
|
||||
List of search results sorted by relevance
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def clear_all(self):
|
||||
"""Clear all indexed data."""
|
||||
|
||||
@abstractmethod
|
||||
async def close(self):
|
||||
"""Close storage and release resources."""
|
||||
652
reme/core/memory_storage/sqlite_memory_store.py
Normal file
652
reme/core/memory_storage/sqlite_memory_store.py
Normal file
|
|
@ -0,0 +1,652 @@
|
|||
"""SQLite storage backend for memory index."""
|
||||
|
||||
import json
|
||||
import sqlite3
|
||||
import struct
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from .base_memory_store import BaseMemoryStore
|
||||
from ..enumeration import MemorySource
|
||||
from ..schema import FileMetadata, MemoryChunk, MemorySearchResult
|
||||
|
||||
|
||||
class SqliteMemoryStore(BaseMemoryStore):
|
||||
"""SQLite memory storage with vector and full-text search.
|
||||
|
||||
Inherits embedding methods from BaseMemoryStore:
|
||||
- get_chunk_embedding / get_chunk_embeddings (async)
|
||||
- get_chunk_embedding_sync / get_chunk_embeddings_sync (sync)
|
||||
- get_embedding / get_embeddings (async)
|
||||
|
||||
Provides SQLite-backed persistent storage with:
|
||||
- Vector similarity search (via sqlite-vec extension)
|
||||
- Full-text search (via FTS5)
|
||||
- Efficient chunk and file metadata management
|
||||
"""
|
||||
|
||||
def __init__(self, db_path: str = ".reme/memory.db", vec_ext_path: str = "", **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.db_path = db_path
|
||||
self.vec_ext_path = vec_ext_path
|
||||
|
||||
self.conn: sqlite3.Connection | None = None
|
||||
|
||||
@property
|
||||
def vector_table_name(self) -> str:
|
||||
"""Get the name of the vector table for this store."""
|
||||
return f"chunks_vec_{self.store_name}"
|
||||
|
||||
@property
|
||||
def fts_table_name(self) -> str:
|
||||
"""Get the name of the FTS table for this store."""
|
||||
return f"chunks_fts_{self.store_name}"
|
||||
|
||||
@property
|
||||
def chunks_table_name(self) -> str:
|
||||
"""Get the name of the chunks table for this store."""
|
||||
return f"chunks_{self.store_name}"
|
||||
|
||||
@property
|
||||
def files_table_name(self) -> str:
|
||||
"""Get the name of the files table for this store."""
|
||||
return f"files_{self.store_name}"
|
||||
|
||||
@staticmethod
|
||||
def vector_to_blob(embedding: list[float]) -> bytes:
|
||||
"""Convert vector to binary blob for sqlite-vec."""
|
||||
return struct.pack(f"{len(embedding)}f", *embedding)
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Initialize database and load extensions."""
|
||||
if self.conn is not None:
|
||||
return
|
||||
|
||||
Path(self.db_path).parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
self.conn = sqlite3.connect(self.db_path, check_same_thread=False)
|
||||
self.conn.enable_load_extension(True)
|
||||
|
||||
# Load sqlite-vec extension
|
||||
if self.vec_ext_path:
|
||||
try:
|
||||
self.conn.load_extension(self.vec_ext_path)
|
||||
self.vector_available = True
|
||||
logger.info(f"Loaded sqlite-vec: {self.vec_ext_path}")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load sqlite-vec: {e}")
|
||||
|
||||
else:
|
||||
try:
|
||||
import sqlite_vec
|
||||
|
||||
ext_path = sqlite_vec.loadable_path()
|
||||
self.conn.load_extension(ext_path)
|
||||
self.vector_available = True
|
||||
logger.info(f"Loaded sqlite-vec from package: {ext_path}")
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load sqlite-vec from package: {e}")
|
||||
# Fallback: try common extension names
|
||||
for name in ["vec0", "sqlite_vec", "vector0"]:
|
||||
try:
|
||||
self.conn.load_extension(name)
|
||||
self.vector_available = True
|
||||
logger.info(f"Loaded sqlite-vec: {name}")
|
||||
break
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
self.conn.enable_load_extension(False)
|
||||
await self._create_tables()
|
||||
|
||||
async def _create_tables(self) -> None:
|
||||
"""Create database schema."""
|
||||
cursor = self.conn.cursor()
|
||||
|
||||
# Files
|
||||
cursor.execute(
|
||||
f"""
|
||||
CREATE TABLE IF NOT EXISTS {self.files_table_name} (
|
||||
path TEXT,
|
||||
source TEXT,
|
||||
hash TEXT,
|
||||
mtime REAL,
|
||||
size INTEGER,
|
||||
PRIMARY KEY (path, source)
|
||||
)
|
||||
""",
|
||||
)
|
||||
|
||||
# Chunks
|
||||
cursor.execute(
|
||||
f"""
|
||||
CREATE TABLE IF NOT EXISTS {self.chunks_table_name} (
|
||||
id TEXT PRIMARY KEY,
|
||||
path TEXT,
|
||||
source TEXT,
|
||||
start_line INTEGER,
|
||||
end_line INTEGER,
|
||||
hash TEXT,
|
||||
text TEXT,
|
||||
embedding TEXT,
|
||||
updated_at INTEGER
|
||||
)
|
||||
""",
|
||||
)
|
||||
|
||||
# Vector table (sqlite-vec)
|
||||
if self.vector_available:
|
||||
cursor.execute(
|
||||
f"""
|
||||
CREATE VIRTUAL TABLE IF NOT EXISTS {self.vector_table_name} USING vec0(
|
||||
id TEXT PRIMARY KEY,
|
||||
embedding FLOAT[{self.embedding_dim}]
|
||||
)
|
||||
""",
|
||||
)
|
||||
logger.info(f"Created vector table (dims={self.embedding_dim})")
|
||||
|
||||
# FTS table
|
||||
if self.fts_enabled:
|
||||
cursor.execute(
|
||||
f"""
|
||||
CREATE VIRTUAL TABLE IF NOT EXISTS {self.fts_table_name} USING fts5(
|
||||
text,
|
||||
id UNINDEXED,
|
||||
path UNINDEXED,
|
||||
source UNINDEXED,
|
||||
start_line UNINDEXED,
|
||||
end_line UNINDEXED
|
||||
)
|
||||
""",
|
||||
)
|
||||
self.fts_available = True
|
||||
logger.info("Created FTS5 table")
|
||||
|
||||
self.conn.commit()
|
||||
cursor.close()
|
||||
|
||||
async def upsert_file(self, file_meta: FileMetadata, source: MemorySource, chunks: list[MemoryChunk]):
|
||||
"""Insert or update file and its chunks."""
|
||||
cursor = self.conn.cursor()
|
||||
|
||||
try:
|
||||
cursor.execute("BEGIN")
|
||||
|
||||
# Insert file
|
||||
cursor.execute(
|
||||
f"""
|
||||
INSERT OR REPLACE INTO {self.files_table_name} (path, source, hash, mtime, size)
|
||||
VALUES (?, ?, ?, ?, ?)
|
||||
""",
|
||||
(file_meta.path, source.value, file_meta.hash, file_meta.mtime_ms, file_meta.size),
|
||||
)
|
||||
|
||||
# Insert chunks
|
||||
now = int(time.time() * 1000)
|
||||
for chunk in chunks:
|
||||
cursor.execute(
|
||||
f"""
|
||||
INSERT OR REPLACE INTO {self.chunks_table_name} (
|
||||
id, path, source, start_line, end_line,
|
||||
hash, text, embedding, updated_at
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
chunk.id,
|
||||
file_meta.path,
|
||||
source.value,
|
||||
chunk.start_line,
|
||||
chunk.end_line,
|
||||
chunk.hash,
|
||||
chunk.text,
|
||||
json.dumps(chunk.embedding) if chunk.embedding else None,
|
||||
now,
|
||||
),
|
||||
)
|
||||
|
||||
# Insert vector
|
||||
if self.vector_available:
|
||||
assert chunk.embedding, "Embedding is required for vector insert"
|
||||
cursor.execute(
|
||||
f"""
|
||||
INSERT OR REPLACE INTO {self.vector_table_name} (id, embedding)
|
||||
VALUES (?, ?)
|
||||
""",
|
||||
(chunk.id, self.vector_to_blob(chunk.embedding)),
|
||||
)
|
||||
|
||||
# Insert FTS
|
||||
if self.fts_available:
|
||||
cursor.execute(
|
||||
f"""
|
||||
INSERT OR REPLACE INTO {self.fts_table_name} (
|
||||
text, id, path, source, start_line, end_line
|
||||
) VALUES (?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
chunk.text,
|
||||
chunk.id,
|
||||
file_meta.path,
|
||||
source.value,
|
||||
chunk.start_line,
|
||||
chunk.end_line,
|
||||
),
|
||||
)
|
||||
|
||||
cursor.execute("COMMIT")
|
||||
except Exception:
|
||||
cursor.execute("ROLLBACK")
|
||||
raise
|
||||
finally:
|
||||
cursor.close()
|
||||
|
||||
async def delete_file(self, path: str, source: MemorySource):
|
||||
"""Delete file and all its chunks."""
|
||||
cursor = self.conn.cursor()
|
||||
try:
|
||||
cursor.execute("BEGIN")
|
||||
|
||||
# Get chunk IDs for vector deletion
|
||||
cursor.execute(
|
||||
f"SELECT id FROM {self.chunks_table_name} WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
chunk_ids = [row[0] for row in cursor.fetchall()]
|
||||
|
||||
# Delete vectors
|
||||
if self.vector_available and chunk_ids:
|
||||
for chunk_id in chunk_ids:
|
||||
try:
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.vector_table_name} WHERE id = ?",
|
||||
(chunk_id,),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Vector delete failed: {e}")
|
||||
|
||||
# Delete FTS entries
|
||||
if self.fts_available:
|
||||
try:
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.fts_table_name} WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"FTS delete failed: {e}")
|
||||
|
||||
# Delete chunks and file
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.chunks_table_name} WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.files_table_name} WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
|
||||
cursor.execute("COMMIT")
|
||||
except Exception:
|
||||
cursor.execute("ROLLBACK")
|
||||
raise
|
||||
finally:
|
||||
cursor.close()
|
||||
|
||||
async def delete_file_chunks(self, path: str, chunk_ids: list[str]):
|
||||
"""Delete specific chunks for a file."""
|
||||
if not chunk_ids:
|
||||
return
|
||||
|
||||
cursor = self.conn.cursor()
|
||||
try:
|
||||
cursor.execute("BEGIN")
|
||||
|
||||
# Delete vectors
|
||||
if self.vector_available:
|
||||
for chunk_id in chunk_ids:
|
||||
try:
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.vector_table_name} WHERE id = ?",
|
||||
(chunk_id,),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Vector delete failed for {chunk_id}: {e}")
|
||||
|
||||
# Delete FTS entries
|
||||
if self.fts_available:
|
||||
placeholders = ",".join("?" * len(chunk_ids))
|
||||
try:
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.fts_table_name} WHERE id IN ({placeholders})",
|
||||
chunk_ids,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"FTS delete failed: {e}")
|
||||
|
||||
# Delete chunks
|
||||
placeholders = ",".join("?" * len(chunk_ids))
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.chunks_table_name} WHERE id IN ({placeholders})",
|
||||
chunk_ids,
|
||||
)
|
||||
|
||||
cursor.execute("COMMIT")
|
||||
except Exception:
|
||||
cursor.execute("ROLLBACK")
|
||||
raise
|
||||
finally:
|
||||
cursor.close()
|
||||
|
||||
async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource):
|
||||
"""Insert or update specific chunks without affecting other chunks."""
|
||||
if not chunks:
|
||||
return
|
||||
|
||||
cursor = self.conn.cursor()
|
||||
try:
|
||||
cursor.execute("BEGIN")
|
||||
|
||||
now = int(time.time() * 1000)
|
||||
for chunk in chunks:
|
||||
# Insert/update chunk
|
||||
cursor.execute(
|
||||
f"""
|
||||
INSERT OR REPLACE INTO {self.chunks_table_name} (
|
||||
id, path, source, start_line, end_line,
|
||||
hash, text, embedding, updated_at
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
chunk.id,
|
||||
chunk.path,
|
||||
source.value,
|
||||
chunk.start_line,
|
||||
chunk.end_line,
|
||||
chunk.hash,
|
||||
chunk.text,
|
||||
json.dumps(chunk.embedding) if chunk.embedding else None,
|
||||
now,
|
||||
),
|
||||
)
|
||||
|
||||
# Insert/update vector
|
||||
if self.vector_available:
|
||||
assert chunk.embedding, "Embedding is required for vector insert"
|
||||
cursor.execute(
|
||||
f"""
|
||||
INSERT OR REPLACE INTO {self.vector_table_name} (id, embedding)
|
||||
VALUES (?, ?)
|
||||
""",
|
||||
(chunk.id, self.vector_to_blob(chunk.embedding)),
|
||||
)
|
||||
|
||||
# Insert/update FTS
|
||||
if self.fts_available:
|
||||
cursor.execute(
|
||||
f"""
|
||||
INSERT OR REPLACE INTO {self.fts_table_name} (
|
||||
text, id, path, source, start_line, end_line
|
||||
) VALUES (?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
chunk.text,
|
||||
chunk.id,
|
||||
chunk.path,
|
||||
source.value,
|
||||
chunk.start_line,
|
||||
chunk.end_line,
|
||||
),
|
||||
)
|
||||
|
||||
cursor.execute("COMMIT")
|
||||
except Exception:
|
||||
cursor.execute("ROLLBACK")
|
||||
raise
|
||||
finally:
|
||||
cursor.close()
|
||||
|
||||
async def list_files(self, source: MemorySource) -> list[str]:
|
||||
"""List all indexed files."""
|
||||
cursor = self.conn.cursor()
|
||||
cursor.execute(f"SELECT path FROM {self.files_table_name} WHERE source = ?", (source.value,))
|
||||
paths = [row[0] for row in cursor.fetchall()]
|
||||
cursor.close()
|
||||
return paths
|
||||
|
||||
async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None:
|
||||
"""Get file metadata with chunk count."""
|
||||
cursor = self.conn.cursor()
|
||||
cursor.execute(
|
||||
f"SELECT hash, mtime, size FROM {self.files_table_name} WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
row = cursor.fetchone()
|
||||
if not row:
|
||||
cursor.close()
|
||||
return None
|
||||
|
||||
hash_val, mtime, size = row
|
||||
cursor.execute(
|
||||
f"SELECT COUNT(*) FROM {self.chunks_table_name} WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
chunk_count = cursor.fetchone()[0]
|
||||
cursor.close()
|
||||
|
||||
return FileMetadata(
|
||||
hash=hash_val,
|
||||
mtime_ms=mtime,
|
||||
size=size,
|
||||
path=path,
|
||||
chunk_count=chunk_count,
|
||||
)
|
||||
|
||||
async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
|
||||
"""Get all chunks for a file."""
|
||||
cursor = self.conn.cursor()
|
||||
cursor.execute(
|
||||
f"""
|
||||
SELECT id, path, source, start_line, end_line, text, hash, embedding
|
||||
FROM {self.chunks_table_name} WHERE path = ? AND source = ?
|
||||
ORDER BY start_line
|
||||
""",
|
||||
(path, source.value),
|
||||
)
|
||||
|
||||
chunks = []
|
||||
for row in cursor.fetchall():
|
||||
chunk_id, path_val, source_val, start, end, text, hash_val, emb_str = row
|
||||
# Parse embedding from JSON string
|
||||
embedding = None
|
||||
if emb_str:
|
||||
try:
|
||||
embedding = json.loads(emb_str)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
embedding = None
|
||||
|
||||
chunks.append(
|
||||
MemoryChunk(
|
||||
id=chunk_id,
|
||||
path=path_val,
|
||||
source=MemorySource(source_val),
|
||||
start_line=start,
|
||||
end_line=end,
|
||||
text=text,
|
||||
hash=hash_val,
|
||||
embedding=embedding,
|
||||
),
|
||||
)
|
||||
|
||||
cursor.close()
|
||||
return chunks
|
||||
|
||||
async def vector_search(
|
||||
self,
|
||||
query: str,
|
||||
limit: int,
|
||||
sources: list[MemorySource] | None = None,
|
||||
) -> list[MemorySearchResult]:
|
||||
"""Perform vector similarity search."""
|
||||
if not self.vector_available or not query:
|
||||
return []
|
||||
|
||||
# Get query embedding
|
||||
query_embedding = await self.get_embedding(query)
|
||||
if not query_embedding:
|
||||
return []
|
||||
|
||||
cursor = self.conn.cursor()
|
||||
source_filter = ""
|
||||
params: list = []
|
||||
if sources:
|
||||
placeholders = ",".join("?" * len(sources))
|
||||
source_filter = f" AND c.source IN ({placeholders})"
|
||||
params = [s.value for s in sources]
|
||||
|
||||
try:
|
||||
query_blob = self.vector_to_blob(query_embedding)
|
||||
|
||||
# Correct SQLite-vec syntax for vector search with limit
|
||||
# vec0 requires 'k = ?' constraint for knn queries
|
||||
query_sql = f"""
|
||||
SELECT c.id, c.path, c.start_line, c.end_line, c.source, c.text, v.distance
|
||||
FROM {self.vector_table_name} v
|
||||
JOIN {self.chunks_table_name} c ON v.id = c.id
|
||||
WHERE v.embedding MATCH ?
|
||||
AND k = ?
|
||||
"""
|
||||
query_params: list = [query_blob, limit]
|
||||
|
||||
# Add source filter if specified
|
||||
if source_filter:
|
||||
query_sql += source_filter
|
||||
query_params.extend(params)
|
||||
|
||||
# Order by distance (k constraint already limits results)
|
||||
query_sql += " ORDER BY v.distance"
|
||||
|
||||
cursor.execute(query_sql, query_params)
|
||||
|
||||
results = []
|
||||
for _, path, start, end, src, text, dist in cursor.fetchall():
|
||||
score = max(0.0, 1.0 - dist)
|
||||
snippet = text[: self.snippet_max_chars] if len(text) > self.snippet_max_chars else text
|
||||
results.append(
|
||||
MemorySearchResult(
|
||||
path=path,
|
||||
start_line=start,
|
||||
end_line=end,
|
||||
score=score,
|
||||
snippet=snippet,
|
||||
source=MemorySource(src),
|
||||
),
|
||||
)
|
||||
|
||||
return results
|
||||
except Exception as e:
|
||||
logger.error(f"Vector search failed: {e}")
|
||||
return []
|
||||
finally:
|
||||
cursor.close()
|
||||
|
||||
async def keyword_search(
|
||||
self,
|
||||
query: str,
|
||||
limit: int,
|
||||
sources: list[MemorySource] | None = None,
|
||||
) -> list[MemorySearchResult]:
|
||||
"""Perform full-text search."""
|
||||
if not self.fts_available:
|
||||
return []
|
||||
|
||||
# Build FTS5 query
|
||||
# Split query into tokens and join with OR for better recall
|
||||
# Individual words are automatically stemmed and matched by FTS5
|
||||
cleaned = query.strip()
|
||||
if not cleaned:
|
||||
return []
|
||||
|
||||
# Split into words and escape each
|
||||
words = cleaned.split()
|
||||
if not words:
|
||||
return []
|
||||
|
||||
# Use OR operator for better recall - match any of the query words
|
||||
escaped_words = [word.replace('"', '""') for word in words]
|
||||
fts_query = " OR ".join(escaped_words)
|
||||
|
||||
cursor = self.conn.cursor()
|
||||
source_filter = ""
|
||||
params: list = [fts_query]
|
||||
if sources:
|
||||
placeholders = ",".join("?" * len(sources))
|
||||
source_filter = f" AND fts.source IN ({placeholders})"
|
||||
params.extend([s.value for s in sources])
|
||||
params.append(limit)
|
||||
|
||||
try:
|
||||
cursor.execute(
|
||||
f"""
|
||||
SELECT fts.id, fts.path, fts.start_line, fts.end_line,
|
||||
fts.source, fts.text, rank
|
||||
FROM {self.fts_table_name} fts
|
||||
WHERE fts.text MATCH ?{source_filter}
|
||||
ORDER BY rank
|
||||
LIMIT ?
|
||||
""",
|
||||
params,
|
||||
)
|
||||
|
||||
results = []
|
||||
for _, path, start, end, src, text, rank in cursor.fetchall():
|
||||
# Convert BM25 rank (negative) to 0-1 score (higher=better)
|
||||
score = max(0.0, 1.0 / (1.0 + abs(rank)))
|
||||
snippet = text[: self.snippet_max_chars] if len(text) > self.snippet_max_chars else text
|
||||
results.append(
|
||||
MemorySearchResult(
|
||||
path=path,
|
||||
start_line=start,
|
||||
end_line=end,
|
||||
score=score,
|
||||
snippet=snippet,
|
||||
source=MemorySource(src),
|
||||
),
|
||||
)
|
||||
|
||||
return results
|
||||
except Exception as e:
|
||||
logger.error(f"Keyword search failed: {e}")
|
||||
return []
|
||||
finally:
|
||||
cursor.close()
|
||||
|
||||
async def clear_all(self):
|
||||
"""Clear all indexed data."""
|
||||
cursor = self.conn.cursor()
|
||||
cursor.execute("BEGIN")
|
||||
|
||||
try:
|
||||
cursor.execute(f"DELETE FROM {self.files_table_name}")
|
||||
cursor.execute(f"DELETE FROM {self.chunks_table_name}")
|
||||
|
||||
if self.vector_available:
|
||||
cursor.execute(f"DELETE FROM {self.vector_table_name}")
|
||||
|
||||
if self.fts_available:
|
||||
cursor.execute(f"DELETE FROM {self.fts_table_name}")
|
||||
|
||||
cursor.execute("COMMIT")
|
||||
except Exception:
|
||||
cursor.execute("ROLLBACK")
|
||||
raise
|
||||
finally:
|
||||
cursor.close()
|
||||
|
||||
async def close(self):
|
||||
"""Close database connection."""
|
||||
if self.conn:
|
||||
self.conn.close()
|
||||
self.conn = None
|
||||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -1,6 +1,10 @@
|
|||
"""schema"""
|
||||
|
||||
from .file_metadata import FileMetadata
|
||||
from .memory_chunk import MemoryChunk
|
||||
from .memory_index_meta import MemoryIndexMeta
|
||||
from .memory_node import MemoryNode
|
||||
from .memory_search_result import MemorySearchResult
|
||||
from .message import ContentBlock, Message, Trajectory
|
||||
from .request import Request
|
||||
from .response import Response
|
||||
|
|
@ -17,16 +21,22 @@ from .service_config import (
|
|||
)
|
||||
from .stream_chunk import StreamChunk
|
||||
from .tool_call import ToolAttr, ToolCall
|
||||
from .truncation_result import TruncationResult
|
||||
from .vector_node import VectorNode
|
||||
|
||||
__all__ = [
|
||||
"MemoryNode",
|
||||
"CmdConfig",
|
||||
"ContentBlock",
|
||||
"EmbeddingModelConfig",
|
||||
"FileMetadata",
|
||||
"FlowConfig",
|
||||
"HttpConfig",
|
||||
"LLMConfig",
|
||||
"MCPConfig",
|
||||
"MemoryChunk",
|
||||
"MemoryIndexMeta",
|
||||
"MemoryNode",
|
||||
"MemorySearchResult",
|
||||
"Message",
|
||||
"Request",
|
||||
"Response",
|
||||
|
|
@ -36,7 +46,7 @@ __all__ = [
|
|||
"Trajectory",
|
||||
"ToolAttr",
|
||||
"ToolCall",
|
||||
"TruncationResult",
|
||||
"VectorNode",
|
||||
"VectorStoreConfig",
|
||||
"CmdConfig",
|
||||
]
|
||||
|
|
|
|||
15
reme/core/schema/file_metadata.py
Normal file
15
reme/core/schema/file_metadata.py
Normal file
|
|
@ -0,0 +1,15 @@
|
|||
"""File metadata schema."""
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class FileMetadata(BaseModel):
|
||||
"""File metadata with optional extended fields for various use cases."""
|
||||
|
||||
hash: str = Field(default=..., description="Hash of the file content")
|
||||
mtime_ms: float = Field(default=..., description="Last modification time in milliseconds")
|
||||
size: int = Field(default=..., description="File size in bytes")
|
||||
path: str | None = Field(default=None, description="Relative path to the session file")
|
||||
content: str | None = Field(default=None, description="Parsed content from the session file")
|
||||
chunk_count: int | None = Field(default=None, description="Number of chunks in the file")
|
||||
metadata: dict = Field(default_factory=dict, description="Additional metadata")
|
||||
19
reme/core/schema/memory_chunk.py
Normal file
19
reme/core/schema/memory_chunk.py
Normal file
|
|
@ -0,0 +1,19 @@
|
|||
"""Memory chunk schema."""
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from ..enumeration import MemorySource
|
||||
|
||||
|
||||
class MemoryChunk(BaseModel):
|
||||
"""A chunk of memory content with metadata."""
|
||||
|
||||
id: str = Field(..., description="Unique identifier for the chunk")
|
||||
path: str = Field(..., description="File path relative to workspace")
|
||||
source: MemorySource = Field(..., description="Source of the memory data")
|
||||
start_line: int = Field(..., description="Starting line number in the source file")
|
||||
end_line: int = Field(..., description="Ending line number in the source file")
|
||||
text: str = Field(..., description="Text content of the chunk")
|
||||
hash: str = Field(..., description="Hash of the chunk content")
|
||||
embedding: list[float] | None = Field(default=None, description="Vector embedding of the chunk")
|
||||
metadata: dict = Field(default_factory=dict, description="Additional metadata")
|
||||
14
reme/core/schema/memory_index_meta.py
Normal file
14
reme/core/schema/memory_index_meta.py
Normal file
|
|
@ -0,0 +1,14 @@
|
|||
"""Memory index metadata schema."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class MemoryIndexMeta(BaseModel):
|
||||
"""Metadata for memory index configuration."""
|
||||
|
||||
model: str = Field(..., description="Name of the embedding model")
|
||||
chunk_tokens: int = Field(..., description="Maximum tokens per chunk")
|
||||
chunk_overlap: int = Field(..., description="Number of overlapping tokens between chunks")
|
||||
vector_dims: Optional[int] = Field(default=None, description="Vector embedding dimensions")
|
||||
24
reme/core/schema/memory_search_result.py
Normal file
24
reme/core/schema/memory_search_result.py
Normal file
|
|
@ -0,0 +1,24 @@
|
|||
"""Memory search result schema."""
|
||||
|
||||
from typing import Any, Dict
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from ..enumeration import MemorySource
|
||||
|
||||
|
||||
class MemorySearchResult(BaseModel):
|
||||
"""Search result from memory index."""
|
||||
|
||||
path: str = Field(..., description="File path relative to workspace")
|
||||
start_line: int = Field(..., description="Starting line number of the match")
|
||||
end_line: int = Field(..., description="Ending line number of the match")
|
||||
score: float = Field(..., description="Relevance score of the search result")
|
||||
snippet: str = Field(..., description="Text snippet from the matched content")
|
||||
source: MemorySource = Field(..., description="Source of the memory data")
|
||||
metadata: Dict[str, Any] = Field(default_factory=dict, description="Additional metadata")
|
||||
|
||||
@property
|
||||
def merge_key(self) -> str:
|
||||
"""Merge key for the search result."""
|
||||
return self.path + f":{self.start_line}:{self.end_line}"
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
35
reme/core/schema/truncation_result.py
Normal file
35
reme/core/schema/truncation_result.py
Normal file
|
|
@ -0,0 +1,35 @@
|
|||
"""Truncation result schema for command output truncation."""
|
||||
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class TruncationResult(BaseModel):
|
||||
"""Result of output truncation operation.
|
||||
|
||||
Attributes:
|
||||
content: The truncated content
|
||||
truncated: Whether truncation occurred
|
||||
total_lines: Total number of lines in original output
|
||||
output_lines: Number of lines in truncated output
|
||||
total_bytes: Total bytes in original output
|
||||
output_bytes: Bytes in truncated output
|
||||
truncated_by: What caused truncation ('lines' or 'bytes')
|
||||
last_line_partial: Whether last line was partially truncated
|
||||
"""
|
||||
|
||||
content: str = Field(description="The truncated content")
|
||||
truncated: bool = Field(description="Whether truncation occurred")
|
||||
total_lines: int = Field(description="Total number of lines in original output")
|
||||
output_lines: int = Field(description="Number of lines in truncated output")
|
||||
total_bytes: int = Field(description="Total bytes in original output")
|
||||
output_bytes: int = Field(description="Bytes in truncated output")
|
||||
truncated_by: Literal["lines", "bytes"] | None = Field(
|
||||
default=None,
|
||||
description="What caused truncation ('lines' or 'bytes')",
|
||||
)
|
||||
last_line_partial: bool = Field(
|
||||
default=False,
|
||||
description="Whether last line was partially truncated",
|
||||
)
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
124
reme/core/utils/chunking_utils.py
Normal file
124
reme/core/utils/chunking_utils.py
Normal file
|
|
@ -0,0 +1,124 @@
|
|||
"""Chunking logic for Markdown files."""
|
||||
|
||||
from .common_utils import hash_text
|
||||
from ..enumeration import MemorySource
|
||||
from ..schema import MemoryChunk
|
||||
|
||||
|
||||
def chunk_markdown(
|
||||
text: str,
|
||||
path: str,
|
||||
source: MemorySource,
|
||||
chunk_tokens: int,
|
||||
overlap: int,
|
||||
) -> list[MemoryChunk]:
|
||||
"""
|
||||
Markdown chunking logic implemented based on the TypeScript version.
|
||||
|
||||
Args:
|
||||
text: Input text
|
||||
path: File path
|
||||
source: Memory source
|
||||
chunk_tokens: Maximum tokens per chunk
|
||||
overlap: Overlap tokens between chunks
|
||||
|
||||
Returns:
|
||||
List of MemoryChunk objects
|
||||
"""
|
||||
lines = text.split("\n")
|
||||
if not lines:
|
||||
return []
|
||||
|
||||
# Convert tokens to characters (~1 token = 4 chars)
|
||||
max_chars = max(32, chunk_tokens * 4)
|
||||
overlap_chars = max(0, overlap * 4)
|
||||
|
||||
chunks: list[MemoryChunk] = []
|
||||
|
||||
# Currently building chunk
|
||||
current: list[dict] = [] # [{'line': str, 'line_no': int}]
|
||||
current_chars = 0
|
||||
|
||||
def flush():
|
||||
"""Add current chunk to results list"""
|
||||
if not current:
|
||||
return
|
||||
|
||||
first_entry = current[0]
|
||||
last_entry = current[-1]
|
||||
|
||||
if not first_entry or not last_entry:
|
||||
return
|
||||
|
||||
chunk_text = "\n".join([entry["line"] for entry in current])
|
||||
start_line = first_entry["line_no"]
|
||||
end_line = last_entry["line_no"]
|
||||
|
||||
chunk_hash = hash_text(chunk_text)
|
||||
|
||||
chunks.append(
|
||||
MemoryChunk(
|
||||
id=hash_text(f"{source}:{path}:{start_line}:{end_line}:{chunk_hash}:{len(chunks)}"),
|
||||
path=path,
|
||||
source=source,
|
||||
start_line=start_line,
|
||||
end_line=end_line,
|
||||
text=chunk_text,
|
||||
hash=chunk_hash,
|
||||
),
|
||||
)
|
||||
|
||||
def carry_overlap():
|
||||
"""Keep overlapping part and clear the rest"""
|
||||
nonlocal current, current_chars
|
||||
|
||||
if overlap_chars <= 0 or not current:
|
||||
current = []
|
||||
current_chars = 0
|
||||
return
|
||||
|
||||
acc = 0
|
||||
kept = []
|
||||
|
||||
# Collect lines from the end until reaching overlap size
|
||||
for j in range(len(current) - 1, -1, -1):
|
||||
entry = current[j]
|
||||
if not entry:
|
||||
continue
|
||||
|
||||
acc += len(entry["line"]) + 1 # +1 for newline
|
||||
kept.insert(0, entry) # Insert at the beginning to maintain order
|
||||
|
||||
if acc >= overlap_chars:
|
||||
break
|
||||
|
||||
current = kept
|
||||
current_chars = sum(len(entry["line"]) + 1 for entry in kept)
|
||||
|
||||
for i, line in enumerate(lines):
|
||||
line_no = i + 1
|
||||
|
||||
# Split long lines into multiple segments
|
||||
segments = []
|
||||
if not line: # Empty line
|
||||
segments.append("")
|
||||
else:
|
||||
# If line is too long, split by maximum character count
|
||||
for start in range(0, len(line), max_chars):
|
||||
segments.append(line[start : start + max_chars])
|
||||
|
||||
for segment in segments:
|
||||
line_size = len(segment) + 1 # +1 for newline
|
||||
|
||||
# If adding current segment would exceed the limit, flush current chunk
|
||||
if current_chars + line_size > max_chars and current:
|
||||
flush()
|
||||
carry_overlap()
|
||||
|
||||
current.append({"line": segment, "line_no": line_no})
|
||||
current_chars += line_size
|
||||
|
||||
# Process the final chunk
|
||||
flush()
|
||||
|
||||
return [c for c in chunks if c.text.strip()]
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
12
reme/reme.py
12
reme/reme.py
|
|
@ -20,7 +20,7 @@ from .agent.memory import (
|
|||
)
|
||||
from .config import ReMeConfigParser
|
||||
from .core import Application
|
||||
from .core.enumeration import MemoryType
|
||||
from .core.enumeration import MemoryType, Role
|
||||
from .core.schema import Message, MemoryNode
|
||||
from .tool.memory import (
|
||||
RetrieveMemory,
|
||||
|
|
@ -59,7 +59,7 @@ class ReMe(Application):
|
|||
target_user_names: list[str] | None = None,
|
||||
target_task_names: list[str] | None = None,
|
||||
target_tool_names: list[str] | None = None,
|
||||
profile_dir: str = "reme_profile",
|
||||
profile_dir: str = ".reme/profile",
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize ReMe with config.
|
||||
|
|
@ -83,11 +83,14 @@ class ReMe(Application):
|
|||
|
||||
Example:
|
||||
```python
|
||||
reme = await ReMe(...).start()
|
||||
reme = ReMe(...)
|
||||
await reme.start()
|
||||
# reme = await ReMe.create(...) # both ok
|
||||
|
||||
await reme.summarize_memory(...)
|
||||
await reme.retrieve_memory(...)
|
||||
|
||||
await reme.close()
|
||||
```
|
||||
|
||||
"""
|
||||
|
|
@ -308,7 +311,8 @@ class ReMe(Application):
|
|||
if user_name:
|
||||
if isinstance(user_name, str):
|
||||
for message in format_messages:
|
||||
message.name = user_name
|
||||
if message.role is Role.USER:
|
||||
message.name = user_name
|
||||
self._add_meta_memory(MemoryType.PERSONAL, user_name)
|
||||
memory_targets.append(user_name)
|
||||
elif isinstance(user_name, list):
|
||||
|
|
|
|||
24
reme/tool/fs/__init__.py
Normal file
24
reme/tool/fs/__init__.py
Normal file
|
|
@ -0,0 +1,24 @@
|
|||
"""File system tools."""
|
||||
|
||||
from .bash_tool import BashTool
|
||||
from .edit_tool import EditTool
|
||||
from .find_tool import FindTool
|
||||
from .grep_tool import GrepTool
|
||||
from .ls_tool import LsTool
|
||||
from .read_tool import ReadTool
|
||||
from .write_tool import WriteTool
|
||||
from ...core import R
|
||||
|
||||
__all__ = [
|
||||
"BashTool",
|
||||
"EditTool",
|
||||
"FindTool",
|
||||
"GrepTool",
|
||||
"LsTool",
|
||||
"ReadTool",
|
||||
"WriteTool",
|
||||
]
|
||||
|
||||
for name in __all__:
|
||||
tool_class = globals()[name]
|
||||
R.op.register(tool_class)
|
||||
191
reme/tool/fs/bash_tool.py
Normal file
191
reme/tool/fs/bash_tool.py
Normal file
|
|
@ -0,0 +1,191 @@
|
|||
"""Bash command execution tool with production-grade features.
|
||||
|
||||
This module provides a production-grade tool for executing bash commands with:
|
||||
- Smart output truncation (keeps last N lines/bytes to prevent memory issues)
|
||||
- Process tree termination (prevents orphan processes)
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import platform
|
||||
import signal
|
||||
from pathlib import Path
|
||||
|
||||
from .truncate import DEFAULT_MAX_BYTES, DEFAULT_MAX_LINES, truncate_tail
|
||||
from ...core.op import BaseTool
|
||||
from ...core.schema import ToolCall, TruncationResult
|
||||
|
||||
|
||||
def get_shell_config() -> tuple[str, list[str]]:
|
||||
"""Get the appropriate shell and arguments for the current platform.
|
||||
|
||||
Returns:
|
||||
Tuple of (shell_path, args) for subprocess execution
|
||||
"""
|
||||
system = platform.system()
|
||||
|
||||
if system == "Windows":
|
||||
# Use PowerShell on Windows
|
||||
return "powershell.exe", ["-Command"]
|
||||
else:
|
||||
# Use bash on Unix-like systems
|
||||
shell = os.environ.get("SHELL", "/bin/bash")
|
||||
return shell, ["-c"]
|
||||
|
||||
|
||||
def kill_process_tree(pid: int) -> None:
|
||||
"""Kill a process and all its children.
|
||||
|
||||
Args:
|
||||
pid: Process ID to kill
|
||||
"""
|
||||
try:
|
||||
if platform.system() == "Windows":
|
||||
# Windows: use taskkill
|
||||
os.system(f"taskkill /F /T /PID {pid}")
|
||||
else:
|
||||
# Unix: kill process group
|
||||
try:
|
||||
os.killpg(os.getpgid(pid), signal.SIGTERM)
|
||||
except ProcessLookupError:
|
||||
pass # Process already dead
|
||||
except Exception:
|
||||
pass # Best effort
|
||||
|
||||
|
||||
class BashTool(BaseTool):
|
||||
"""Production-grade tool for executing bash commands.
|
||||
|
||||
Features:
|
||||
- Smart output truncation (preserves last N lines or M bytes)
|
||||
- Kills entire process tree on timeout (prevents orphan processes)
|
||||
"""
|
||||
|
||||
def __init__(self, cwd: str | None = None, command_prefix: str | None = None):
|
||||
"""Initialize bash tool.
|
||||
|
||||
Args:
|
||||
cwd: Working directory (defaults to current directory)
|
||||
command_prefix: Optional prefix prepended to every command
|
||||
"""
|
||||
super().__init__()
|
||||
self.cwd = cwd or os.getcwd()
|
||||
self.command_prefix = command_prefix
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
max_kb = DEFAULT_MAX_BYTES // 1024
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": (
|
||||
f"Execute a bash command in the current working directory. "
|
||||
f"Returns stdout and stderr. Output is truncated to last "
|
||||
f"{DEFAULT_MAX_LINES} lines or {max_kb}KB (whichever is hit first). "
|
||||
f"Optionally provide a timeout in seconds."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"command": {
|
||||
"type": "string",
|
||||
"description": "Bash command to execute",
|
||||
},
|
||||
"timeout": {
|
||||
"type": "number",
|
||||
"description": "Timeout in seconds (optional, no default timeout)",
|
||||
},
|
||||
},
|
||||
"required": ["command"],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self) -> str:
|
||||
"""Execute the bash command with production-grade features."""
|
||||
command: str = self.context.command
|
||||
timeout: float | None = self.context.get("timeout", None)
|
||||
|
||||
# Apply command prefix if configured
|
||||
if self.command_prefix:
|
||||
command = f"{self.command_prefix}\n{command}"
|
||||
|
||||
# Verify working directory exists
|
||||
if not Path(self.cwd).exists():
|
||||
raise FileNotFoundError(
|
||||
f"Working directory does not exist: {self.cwd}\n" f"Cannot execute bash commands.",
|
||||
)
|
||||
|
||||
# Get shell configuration
|
||||
shell, shell_args = get_shell_config()
|
||||
|
||||
# Start process
|
||||
try:
|
||||
process = await asyncio.create_subprocess_exec(
|
||||
shell,
|
||||
*shell_args,
|
||||
command,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
cwd=self.cwd,
|
||||
# Create process group for clean termination
|
||||
preexec_fn=os.setpgrp if platform.system() != "Windows" else None,
|
||||
)
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Failed to start process: {e}") from e
|
||||
|
||||
# Execute command with optional timeout
|
||||
try:
|
||||
if timeout and timeout > 0:
|
||||
try:
|
||||
stdout, stderr = await asyncio.wait_for(
|
||||
process.communicate(),
|
||||
timeout=timeout,
|
||||
)
|
||||
except asyncio.TimeoutError as e:
|
||||
# Kill process tree on timeout
|
||||
if process.pid:
|
||||
kill_process_tree(process.pid)
|
||||
try:
|
||||
await asyncio.wait_for(process.wait(), timeout=1.0)
|
||||
except asyncio.TimeoutError:
|
||||
process.kill()
|
||||
raise TimeoutError(f"Command timed out after {timeout} seconds") from e
|
||||
else:
|
||||
stdout, stderr = await process.communicate()
|
||||
except TimeoutError as e:
|
||||
raise RuntimeError(str(e)) from e
|
||||
|
||||
# Decode output
|
||||
full_output = stdout.decode("utf-8", errors="ignore")
|
||||
if stderr:
|
||||
stderr_text = stderr.decode("utf-8", errors="ignore")
|
||||
if full_output:
|
||||
full_output += "\n"
|
||||
full_output += stderr_text
|
||||
|
||||
# Apply tail truncation_result to prevent memory issues
|
||||
truncation_result: TruncationResult = truncate_tail(full_output)
|
||||
output_text = truncation_result.content or "(no output)"
|
||||
|
||||
# Build truncation_result notice if needed
|
||||
if truncation_result.truncated:
|
||||
start_line = truncation_result.total_lines - truncation_result.output_lines + 1
|
||||
end_line = truncation_result.total_lines
|
||||
|
||||
if truncation_result.truncated_by == "lines":
|
||||
output_text += (
|
||||
f"\n\n[Output truncated: showing lines {start_line}-{end_line} "
|
||||
f"of {truncation_result.total_lines} total lines]"
|
||||
)
|
||||
else:
|
||||
max_kb = DEFAULT_MAX_BYTES // 1024
|
||||
output_text += (
|
||||
f"\n\n[Output truncated: showing lines {start_line}-{end_line} "
|
||||
f"of {truncation_result.total_lines} ({max_kb}KB limit reached)]"
|
||||
)
|
||||
|
||||
# Handle non-zero exit code
|
||||
if process.returncode != 0:
|
||||
output_text += f"\n\nCommand exited with code {process.returncode}"
|
||||
raise RuntimeError(output_text)
|
||||
|
||||
return output_text
|
||||
164
reme/tool/fs/edit_diff.py
Normal file
164
reme/tool/fs/edit_diff.py
Normal file
|
|
@ -0,0 +1,164 @@
|
|||
"""Diff utilities for edit tool."""
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from difflib import unified_diff
|
||||
|
||||
|
||||
def detect_line_ending(content: str) -> str:
|
||||
"""Detect line ending style (CRLF or LF)."""
|
||||
crlf_idx = content.find("\r\n")
|
||||
lf_idx = content.find("\n")
|
||||
if lf_idx == -1:
|
||||
return "\n"
|
||||
if crlf_idx == -1:
|
||||
return "\n"
|
||||
return "\r\n" if crlf_idx < lf_idx else "\n"
|
||||
|
||||
|
||||
def normalize_to_lf(text: str) -> str:
|
||||
"""Normalize line endings to LF."""
|
||||
return text.replace("\r\n", "\n").replace("\r", "\n")
|
||||
|
||||
|
||||
def restore_line_endings(text: str, ending: str) -> str:
|
||||
"""Restore original line endings."""
|
||||
return text.replace("\n", ending) if ending == "\r\n" else text
|
||||
|
||||
|
||||
def normalize_for_fuzzy_match(text: str) -> str:
|
||||
"""Normalize text for fuzzy matching: strip trailing whitespace, normalize quotes/dashes."""
|
||||
lines = text.split("\n")
|
||||
normalized = "\n".join(line.rstrip() for line in lines)
|
||||
|
||||
# Smart quotes → ASCII
|
||||
normalized = re.sub(r"[\u2018\u2019\u201A\u201B]", "'", normalized)
|
||||
normalized = re.sub(r"[\u201C\u201D\u201E\u201F]", '"', normalized)
|
||||
|
||||
# Dashes → hyphen
|
||||
normalized = re.sub(r"[\u2010\u2011\u2012\u2013\u2014\u2015\u2212]", "-", normalized)
|
||||
|
||||
# Special spaces → regular space
|
||||
normalized = re.sub(r"[\u00A0\u2002-\u200A\u202F\u205F\u3000]", " ", normalized)
|
||||
|
||||
return normalized
|
||||
|
||||
|
||||
@dataclass
|
||||
class FuzzyMatchResult:
|
||||
"""Result of fuzzy text matching."""
|
||||
|
||||
found: bool
|
||||
index: int
|
||||
match_length: int
|
||||
used_fuzzy_match: bool
|
||||
content_for_replacement: str
|
||||
|
||||
|
||||
def fuzzy_find_text(content: str, old_text: str) -> FuzzyMatchResult:
|
||||
"""Find old_text in content, trying exact match first, then fuzzy match."""
|
||||
# Try exact match
|
||||
exact_index = content.find(old_text)
|
||||
if exact_index != -1:
|
||||
return FuzzyMatchResult(
|
||||
found=True,
|
||||
index=exact_index,
|
||||
match_length=len(old_text),
|
||||
used_fuzzy_match=False,
|
||||
content_for_replacement=content,
|
||||
)
|
||||
|
||||
# Try fuzzy match
|
||||
fuzzy_content = normalize_for_fuzzy_match(content)
|
||||
fuzzy_old_text = normalize_for_fuzzy_match(old_text)
|
||||
fuzzy_index = fuzzy_content.find(fuzzy_old_text)
|
||||
|
||||
if fuzzy_index == -1:
|
||||
return FuzzyMatchResult(
|
||||
found=False,
|
||||
index=-1,
|
||||
match_length=0,
|
||||
used_fuzzy_match=False,
|
||||
content_for_replacement=content,
|
||||
)
|
||||
|
||||
return FuzzyMatchResult(
|
||||
found=True,
|
||||
index=fuzzy_index,
|
||||
match_length=len(fuzzy_old_text),
|
||||
used_fuzzy_match=True,
|
||||
content_for_replacement=fuzzy_content,
|
||||
)
|
||||
|
||||
|
||||
def strip_bom(content: str) -> tuple[str, str]:
|
||||
"""Strip UTF-8 BOM, return (bom, text_without_bom)."""
|
||||
if content.startswith("\ufeff"):
|
||||
return "\ufeff", content[1:]
|
||||
return "", content
|
||||
|
||||
|
||||
@dataclass
|
||||
class DiffResult:
|
||||
"""Result of diff generation."""
|
||||
|
||||
diff: str
|
||||
first_changed_line: int | None
|
||||
|
||||
|
||||
def generate_diff_string(old_content: str, new_content: str, context_lines: int = 4) -> DiffResult:
|
||||
"""Generate unified diff with line numbers."""
|
||||
old_lines = old_content.split("\n")
|
||||
new_lines = new_content.split("\n")
|
||||
|
||||
# Use difflib to get the changes
|
||||
diff_lines = list(
|
||||
unified_diff(
|
||||
old_lines,
|
||||
new_lines,
|
||||
lineterm="",
|
||||
n=context_lines,
|
||||
),
|
||||
)
|
||||
|
||||
if not diff_lines:
|
||||
return DiffResult(diff="", first_changed_line=None)
|
||||
|
||||
# Parse and format the diff
|
||||
output = []
|
||||
first_changed_line = None
|
||||
max_line_num = max(len(old_lines), len(new_lines))
|
||||
line_num_width = len(str(max_line_num))
|
||||
|
||||
old_line_num = 1
|
||||
new_line_num = 1
|
||||
|
||||
for line in diff_lines[2:]: # Skip header lines
|
||||
if line.startswith("@@"):
|
||||
# Parse hunk header
|
||||
match = re.match(r"@@ -(\d+),?\d* \+(\d+),?\d* @@", line)
|
||||
if match:
|
||||
old_line_num = int(match.group(1))
|
||||
new_line_num = int(match.group(2))
|
||||
continue
|
||||
|
||||
if line.startswith("+"):
|
||||
if first_changed_line is None:
|
||||
first_changed_line = new_line_num
|
||||
line_num = str(new_line_num).rjust(line_num_width)
|
||||
output.append(f"+{line_num} {line[1:]}")
|
||||
new_line_num += 1
|
||||
elif line.startswith("-"):
|
||||
if first_changed_line is None:
|
||||
first_changed_line = new_line_num
|
||||
line_num = str(old_line_num).rjust(line_num_width)
|
||||
output.append(f"-{line_num} {line[1:]}")
|
||||
old_line_num += 1
|
||||
else:
|
||||
# Context line
|
||||
line_num = str(old_line_num).rjust(line_num_width)
|
||||
output.append(f" {line_num} {line[1:] if line.startswith(' ') else line}")
|
||||
old_line_num += 1
|
||||
new_line_num += 1
|
||||
|
||||
return DiffResult(diff="\n".join(output), first_changed_line=first_changed_line)
|
||||
141
reme/tool/fs/edit_tool.py
Normal file
141
reme/tool/fs/edit_tool.py
Normal file
|
|
@ -0,0 +1,141 @@
|
|||
"""File editing tool with exact text replacement."""
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from .edit_diff import (
|
||||
detect_line_ending,
|
||||
fuzzy_find_text,
|
||||
generate_diff_string,
|
||||
normalize_for_fuzzy_match,
|
||||
normalize_to_lf,
|
||||
restore_line_endings,
|
||||
strip_bom,
|
||||
)
|
||||
from ...core.op import BaseTool
|
||||
from ...core.schema import ToolCall
|
||||
|
||||
|
||||
class EditTool(BaseTool):
|
||||
"""Edit a file by replacing exact text."""
|
||||
|
||||
def __init__(self, cwd: str | None = None):
|
||||
"""Initialize edit tool.
|
||||
|
||||
Args:
|
||||
cwd: Working directory (defaults to current directory)
|
||||
"""
|
||||
super().__init__()
|
||||
self.cwd = cwd or os.getcwd()
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": (
|
||||
"Edit a file by replacing exact text. The oldText must match exactly "
|
||||
"(including whitespace). Use this for precise, surgical edits."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"path": {
|
||||
"type": "string",
|
||||
"description": "Path to the file to edit (relative or absolute)",
|
||||
},
|
||||
"oldText": {
|
||||
"type": "string",
|
||||
"description": "Exact text to find and replace (must match exactly)",
|
||||
},
|
||||
"newText": {
|
||||
"type": "string",
|
||||
"description": "New text to replace the old text with",
|
||||
},
|
||||
},
|
||||
"required": ["path", "oldText", "newText"],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self) -> str:
|
||||
"""Execute the edit operation."""
|
||||
path: str = self.context.path
|
||||
old_text: str = self.context.oldText
|
||||
new_text: str = self.context.newText
|
||||
|
||||
# Resolve path
|
||||
if not os.path.isabs(path):
|
||||
absolute_path = os.path.join(self.cwd, path)
|
||||
else:
|
||||
absolute_path = path
|
||||
|
||||
# Check file exists and is writable
|
||||
path_obj = Path(absolute_path)
|
||||
if not path_obj.exists():
|
||||
raise FileNotFoundError(f"File not found: {path}")
|
||||
|
||||
if not os.access(absolute_path, os.R_OK | os.W_OK):
|
||||
raise PermissionError(f"File not readable/writable: {path}")
|
||||
|
||||
# Read file
|
||||
try:
|
||||
with open(absolute_path, "r", encoding="utf-8") as f:
|
||||
raw_content = f.read()
|
||||
except Exception as e:
|
||||
raise IOError(f"Failed to read file {path}: {e}") from e
|
||||
|
||||
# Strip BOM (LLM won't include invisible BOM in oldText)
|
||||
bom, content = strip_bom(raw_content)
|
||||
|
||||
original_ending = detect_line_ending(content)
|
||||
normalized_content = normalize_to_lf(content)
|
||||
normalized_old_text = normalize_to_lf(old_text)
|
||||
normalized_new_text = normalize_to_lf(new_text)
|
||||
|
||||
# Find old text using fuzzy matching
|
||||
match_result = fuzzy_find_text(normalized_content, normalized_old_text)
|
||||
|
||||
if not match_result.found:
|
||||
raise ValueError(
|
||||
f"Could not find the exact text in {path}. The old text must match "
|
||||
f"exactly including all whitespace and newlines.",
|
||||
)
|
||||
|
||||
# Count occurrences for uniqueness check
|
||||
fuzzy_content = normalize_for_fuzzy_match(normalized_content)
|
||||
fuzzy_old_text = normalize_for_fuzzy_match(normalized_old_text)
|
||||
occurrences = fuzzy_content.count(fuzzy_old_text)
|
||||
|
||||
if occurrences > 1:
|
||||
raise ValueError(
|
||||
f"Found {occurrences} occurrences of the text in {path}. "
|
||||
f"The text must be unique. Please provide more context to make it unique.",
|
||||
)
|
||||
|
||||
# Perform replacement
|
||||
base_content = match_result.content_for_replacement
|
||||
new_content = (
|
||||
base_content[: match_result.index]
|
||||
+ normalized_new_text
|
||||
+ base_content[match_result.index + match_result.match_length :]
|
||||
)
|
||||
|
||||
# Verify replacement changed something
|
||||
if base_content == new_content:
|
||||
raise ValueError(
|
||||
f"No changes made to {path}. The replacement produced identical content. "
|
||||
f"This might indicate an issue with special characters or the text not "
|
||||
f"exist as expected.",
|
||||
)
|
||||
|
||||
# Write file
|
||||
final_content = bom + restore_line_endings(new_content, original_ending)
|
||||
try:
|
||||
with open(absolute_path, "w", encoding="utf-8") as f:
|
||||
f.write(final_content)
|
||||
except Exception as e:
|
||||
raise IOError(f"Failed to write file {path}: {e}") from e
|
||||
|
||||
# Generate diff
|
||||
diff_result = generate_diff_string(base_content, new_content)
|
||||
|
||||
return f"Successfully replaced text in {path}.\n\n{diff_result.diff}"
|
||||
183
reme/tool/fs/find_tool.py
Normal file
183
reme/tool/fs/find_tool.py
Normal file
|
|
@ -0,0 +1,183 @@
|
|||
"""File search tool using glob patterns with gitignore support."""
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from .truncate import FIND_MAX_BYTES, FIND_MAX_LINES, format_size, truncate_head
|
||||
from ...core.op import BaseTool
|
||||
from ...core.schema import ToolCall
|
||||
|
||||
|
||||
class FindTool(BaseTool):
|
||||
"""Search for files by glob pattern, respecting .gitignore."""
|
||||
|
||||
def __init__(self, cwd: str | None = None):
|
||||
"""Initialize find tool.
|
||||
|
||||
Args:
|
||||
cwd: Working directory (defaults to current directory)
|
||||
"""
|
||||
super().__init__()
|
||||
self.cwd = cwd or os.getcwd()
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
max_kb = FIND_MAX_BYTES // 1024
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": (
|
||||
f"Search for files by glob pattern. Returns matching file paths relative "
|
||||
f"to the search directory. Respects .gitignore. Output is truncated to "
|
||||
f"1000 results or {max_kb}KB (whichever is hit first)."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"pattern": {
|
||||
"type": "string",
|
||||
"description": "Glob pattern to match files, "
|
||||
"e.g. '*.ts', '**/*.json', or 'src/**/*.spec.ts'",
|
||||
},
|
||||
"path": {
|
||||
"type": "string",
|
||||
"description": "Directory to search in (default: current directory)",
|
||||
},
|
||||
"limit": {
|
||||
"type": "number",
|
||||
"description": "Maximum number of results (default: 1000)",
|
||||
},
|
||||
},
|
||||
"required": ["pattern"],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
def _load_gitignore_patterns(self, search_path: Path) -> list[str]:
|
||||
"""Load gitignore patterns from directory and subdirectories."""
|
||||
patterns = ["**/node_modules/**", "**/.git/**"]
|
||||
|
||||
# Load root .gitignore
|
||||
gitignore_path = search_path / ".gitignore"
|
||||
if gitignore_path.exists():
|
||||
try:
|
||||
with open(gitignore_path, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if line and not line.startswith("#"):
|
||||
patterns.append(line)
|
||||
except Exception:
|
||||
pass # Ignore errors
|
||||
|
||||
# Load nested .gitignore files
|
||||
try:
|
||||
for gitignore in search_path.rglob(".gitignore"):
|
||||
if gitignore == gitignore_path:
|
||||
continue
|
||||
try:
|
||||
with open(gitignore, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if line and not line.startswith("#"):
|
||||
patterns.append(line)
|
||||
except Exception:
|
||||
pass # Ignore errors
|
||||
except Exception:
|
||||
pass # Ignore glob errors
|
||||
|
||||
return patterns
|
||||
|
||||
def _should_ignore(self, path: Path, ignore_patterns: list[str]) -> bool:
|
||||
"""Check if path matches any ignore pattern."""
|
||||
path_str = str(path)
|
||||
|
||||
for pattern in ignore_patterns:
|
||||
# Simple pattern matching (not full gitignore spec)
|
||||
if "**" in pattern:
|
||||
# Recursive match
|
||||
clean_pattern = pattern.replace("**/", "").replace("/**", "")
|
||||
if clean_pattern in path_str:
|
||||
return True
|
||||
elif "*" in pattern:
|
||||
# Wildcard match
|
||||
from fnmatch import fnmatch
|
||||
|
||||
if fnmatch(path.name, pattern):
|
||||
return True
|
||||
elif pattern in path_str:
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
async def execute(self) -> str:
|
||||
"""Execute file search."""
|
||||
pattern: str = self.context.pattern
|
||||
search_dir: str = self.context.get("path", ".")
|
||||
limit: int = self.context.get("limit", 1000)
|
||||
|
||||
# Resolve search path
|
||||
if not os.path.isabs(search_dir):
|
||||
search_path = Path(self.cwd) / search_dir
|
||||
else:
|
||||
search_path = Path(search_dir)
|
||||
|
||||
# Check if directory exists
|
||||
if not search_path.exists():
|
||||
raise FileNotFoundError(f"Path not found: {search_dir}")
|
||||
|
||||
if not search_path.is_dir():
|
||||
raise NotADirectoryError(f"Path is not a directory: {search_dir}")
|
||||
|
||||
# Load gitignore patterns
|
||||
ignore_patterns = self._load_gitignore_patterns(search_path)
|
||||
|
||||
# Search for files
|
||||
results = []
|
||||
try:
|
||||
for file_path in search_path.glob(pattern):
|
||||
if len(results) >= limit:
|
||||
break
|
||||
|
||||
# Skip if matches ignore patterns
|
||||
if self._should_ignore(file_path, ignore_patterns):
|
||||
continue
|
||||
|
||||
# Get relative path
|
||||
try:
|
||||
rel_path = file_path.relative_to(search_path)
|
||||
# Add trailing slash for directories
|
||||
if file_path.is_dir():
|
||||
results.append(f"{rel_path}/")
|
||||
else:
|
||||
results.append(str(rel_path))
|
||||
except ValueError:
|
||||
# If relative_to fails, use the path as-is
|
||||
results.append(str(file_path))
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Error searching for files: {e}") from e
|
||||
|
||||
# Handle no results
|
||||
if not results:
|
||||
return "No files found matching pattern"
|
||||
|
||||
# Sort results for consistency
|
||||
results.sort()
|
||||
|
||||
# Apply limit and truncation
|
||||
result_limit_reached = len(results) >= limit
|
||||
raw_output = "\n".join(results)
|
||||
truncation = truncate_head(raw_output, max_lines=FIND_MAX_LINES, max_bytes=FIND_MAX_BYTES)
|
||||
|
||||
output = truncation.content
|
||||
notices = []
|
||||
|
||||
if result_limit_reached:
|
||||
notices.append(
|
||||
f"{limit} results limit reached. Use limit={limit * 2} for more, or refine pattern",
|
||||
)
|
||||
|
||||
if truncation.truncated:
|
||||
notices.append(f"{format_size(FIND_MAX_BYTES)} limit reached")
|
||||
|
||||
if notices:
|
||||
output += f"\n\n[{'. '.join(notices)}]"
|
||||
|
||||
return output
|
||||
79
reme/tool/fs/fs_memory_get.py
Normal file
79
reme/tool/fs/fs_memory_get.py
Normal 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
|
||||
134
reme/tool/fs/fs_memory_search.py
Normal file
134
reme/tool/fs/fs_memory_search.py
Normal file
|
|
@ -0,0 +1,134 @@
|
|||
"""Memory search tool for semantic search in memory files."""
|
||||
|
||||
import json
|
||||
|
||||
from reme.core.enumeration import MemorySource
|
||||
from reme.core.op import BaseTool
|
||||
from reme.core.schema import MemorySearchResult, ToolCall
|
||||
|
||||
|
||||
class FsMemorySearch(BaseTool):
|
||||
"""Semantically search MEMORY.md and memory files."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
sources: list[MemorySource] | None = None,
|
||||
min_score: float = 0.1,
|
||||
max_results: int = 20,
|
||||
hybrid_enabled: bool = True,
|
||||
hybrid_vector_weight: float = 0.7,
|
||||
hybrid_text_weight: float = 0.3,
|
||||
hybrid_candidate_multiplier: float = 3.0,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize memory search tool."""
|
||||
kwargs.setdefault("name", "memory_search")
|
||||
super().__init__(**kwargs)
|
||||
self.sources = sources or [MemorySource.MEMORY]
|
||||
self.min_score = min_score
|
||||
self.max_results = max_results
|
||||
self.hybrid_enabled = hybrid_enabled
|
||||
self.hybrid_vector_weight = hybrid_vector_weight
|
||||
self.hybrid_text_weight = hybrid_text_weight
|
||||
self.hybrid_candidate_multiplier = hybrid_candidate_multiplier
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": (
|
||||
"Mandatory recall step: semantically search MEMORY.md + memory/*.md "
|
||||
"(and optional session transcripts) before answering questions about "
|
||||
"prior work, decisions, dates, people, preferences, or todos; returns "
|
||||
"top snippets with path + lines."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "The semantic search query to find relevant memory snippets",
|
||||
},
|
||||
"maxResults": {
|
||||
"type": "integer",
|
||||
"description": "Maximum number of search results to return (optional)",
|
||||
},
|
||||
"minScore": {
|
||||
"type": "number",
|
||||
"description": "Minimum similarity score threshold for results (optional)",
|
||||
},
|
||||
},
|
||||
"required": ["query"],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self) -> str:
|
||||
"""Execute the memory search operation."""
|
||||
query: str = self.context.query.strip()
|
||||
min_score = self.context.get("minScore", self.min_score)
|
||||
max_results = self.context.get("maxResults", self.max_results)
|
||||
candidates = min(200, max(1, int(max_results * self.hybrid_candidate_multiplier)))
|
||||
|
||||
# Perform hybrid search (vector + keyword)
|
||||
if self.hybrid_enabled:
|
||||
keyword_results = []
|
||||
if self.memory_store.fts_enabled:
|
||||
keyword_results = await self._search_keyword(query, candidates)
|
||||
vector_results = await self._search_vector(query, candidates)
|
||||
|
||||
if not keyword_results:
|
||||
results = [r for r in vector_results if r.score >= min_score][:max_results]
|
||||
elif not vector_results:
|
||||
results = [r for r in keyword_results if r.score >= min_score][:max_results]
|
||||
else:
|
||||
merged = self._merge_hybrid_results(
|
||||
vector=vector_results,
|
||||
keyword=keyword_results,
|
||||
vector_weight=self.hybrid_vector_weight,
|
||||
text_weight=self.hybrid_text_weight,
|
||||
)
|
||||
results = [r for r in merged if r.score >= min_score][:max_results]
|
||||
else:
|
||||
vector_results = await self._search_vector(query, candidates)
|
||||
results = [r for r in vector_results if r.score >= min_score][:max_results]
|
||||
|
||||
return json.dumps([result.model_dump(exclude_none=True) for result in results], indent=2, ensure_ascii=False)
|
||||
|
||||
async def _search_vector(self, query: str, limit: int) -> list[MemorySearchResult]:
|
||||
"""Perform vector similarity search."""
|
||||
return await self.memory_store.vector_search(query, limit, sources=self.sources)
|
||||
|
||||
async def _search_keyword(self, query: str, limit: int) -> list[MemorySearchResult]:
|
||||
"""Perform keyword/FTS search."""
|
||||
if not self.memory_store.fts_enabled:
|
||||
return []
|
||||
return await self.memory_store.keyword_search(query, limit, sources=self.sources)
|
||||
|
||||
@staticmethod
|
||||
def _merge_hybrid_results(
|
||||
vector: list[MemorySearchResult],
|
||||
keyword: list[MemorySearchResult],
|
||||
vector_weight: float,
|
||||
text_weight: float,
|
||||
) -> list[MemorySearchResult]:
|
||||
"""Merge vector and keyword search results with weighted scoring."""
|
||||
merged: dict[str, MemorySearchResult] = {}
|
||||
|
||||
# Process vector results
|
||||
for result in vector:
|
||||
result.score = result.score * vector_weight
|
||||
merged[result.merge_key] = result
|
||||
|
||||
# Process keyword results
|
||||
for result in keyword:
|
||||
key = result.merge_key
|
||||
if key in merged:
|
||||
merged[key].score += result.score * text_weight
|
||||
else:
|
||||
result.score = result.score * text_weight
|
||||
merged[key] = result
|
||||
|
||||
# Sort by score and return
|
||||
results = list(merged.values())
|
||||
results.sort(key=lambda r: r.score, reverse=True)
|
||||
return results
|
||||
276
reme/tool/fs/grep_tool.py
Normal file
276
reme/tool/fs/grep_tool.py
Normal file
|
|
@ -0,0 +1,276 @@
|
|||
"""Grep tool for searching file contents using ripgrep.
|
||||
|
||||
This module provides a tool for searching file contents with:
|
||||
- Pattern matching (regex or literal string)
|
||||
- Smart output truncation (prevents memory issues)
|
||||
- Context lines support
|
||||
- Respects .gitignore
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
|
||||
from .truncate import (
|
||||
DEFAULT_MAX_BYTES,
|
||||
GREP_MAX_LINE_LENGTH,
|
||||
format_size,
|
||||
truncate_head,
|
||||
truncate_line,
|
||||
)
|
||||
from ...core.op import BaseTool
|
||||
from ...core.schema import ToolCall
|
||||
|
||||
# Default limits
|
||||
DEFAULT_LIMIT = 100 # Maximum number of matches
|
||||
|
||||
|
||||
class GrepTool(BaseTool):
|
||||
"""Tool for searching file contents using ripgrep.
|
||||
|
||||
Features:
|
||||
- Pattern matching with regex or literal string
|
||||
- Context lines support
|
||||
- Smart output truncation
|
||||
- Respects .gitignore
|
||||
"""
|
||||
|
||||
def __init__(self, cwd: str | None = None):
|
||||
"""Initialize grep tool.
|
||||
|
||||
Args:
|
||||
cwd: Working directory (defaults to current directory)
|
||||
"""
|
||||
super().__init__()
|
||||
self.cwd = cwd or os.getcwd()
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
max_kb = DEFAULT_MAX_BYTES // 1024
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": (
|
||||
f"Search file contents for a pattern. Returns matching lines with "
|
||||
f"file paths and line numbers. Respects .gitignore. Output is "
|
||||
f"truncated to {DEFAULT_LIMIT} matches or {max_kb}KB (whichever is "
|
||||
f"hit first). Long lines are truncated to {GREP_MAX_LINE_LENGTH} chars."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"pattern": {
|
||||
"type": "string",
|
||||
"description": "Search pattern (regex or literal string)",
|
||||
},
|
||||
"path": {
|
||||
"type": "string",
|
||||
"description": "Directory or file to search (default: current directory)",
|
||||
},
|
||||
"glob": {
|
||||
"type": "string",
|
||||
"description": "Filter files by glob pattern, e.g. '*.ts' or '**/*.spec.ts'",
|
||||
},
|
||||
"ignoreCase": {
|
||||
"type": "boolean",
|
||||
"description": "Case-insensitive search (default: false)",
|
||||
},
|
||||
"literal": {
|
||||
"type": "boolean",
|
||||
"description": "Treat pattern as literal string instead of regex (default: false)",
|
||||
},
|
||||
"contextLines": {
|
||||
"type": "number",
|
||||
"description": "Number of lines to show before and after each match (default: 0)",
|
||||
},
|
||||
"limit": {
|
||||
"type": "number",
|
||||
"description": f"Maximum number of matches to return (default: {DEFAULT_LIMIT})",
|
||||
},
|
||||
},
|
||||
"required": ["pattern"],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self) -> str:
|
||||
"""Execute the grep search."""
|
||||
pattern: str = self.context.pattern
|
||||
search_path: str = self.context.get("path", ".")
|
||||
glob: str | None = self.context.get("glob", None)
|
||||
ignore_case: bool = self.context.get("ignoreCase", False)
|
||||
literal: bool = self.context.get("literal", False)
|
||||
context_lines: int = self.context.get("contextLines", 0)
|
||||
limit: int = self.context.get("limit", DEFAULT_LIMIT)
|
||||
|
||||
# Check if ripgrep is available
|
||||
rg_path = shutil.which("rg")
|
||||
if not rg_path:
|
||||
raise RuntimeError(
|
||||
"ripgrep (rg) is not available. Please install it:\n"
|
||||
" macOS: brew install ripgrep\n"
|
||||
" Ubuntu: apt-get install ripgrep\n"
|
||||
" Other: https://github.com/BurntSushi/ripgrep",
|
||||
)
|
||||
|
||||
# Resolve search path
|
||||
if not os.path.isabs(search_path):
|
||||
search_path = os.path.join(self.cwd, search_path)
|
||||
|
||||
# Check if path exists
|
||||
if not Path(search_path).exists():
|
||||
raise FileNotFoundError(f"Path not found: {search_path}")
|
||||
|
||||
is_directory = Path(search_path).is_dir()
|
||||
effective_limit = max(1, limit)
|
||||
|
||||
# Build ripgrep arguments
|
||||
args = [
|
||||
rg_path,
|
||||
"--json",
|
||||
"--line-number",
|
||||
"--color=never",
|
||||
"--hidden",
|
||||
]
|
||||
|
||||
if ignore_case:
|
||||
args.append("--ignore-case")
|
||||
|
||||
if literal:
|
||||
args.append("--fixed-strings")
|
||||
|
||||
if glob:
|
||||
args.extend(["--glob", glob])
|
||||
|
||||
args.extend([pattern, search_path])
|
||||
|
||||
# Execute ripgrep
|
||||
try:
|
||||
process = await asyncio.create_subprocess_exec(
|
||||
*args,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
cwd=self.cwd,
|
||||
)
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Failed to run ripgrep: {e}") from e
|
||||
|
||||
stdout, stderr = await process.communicate()
|
||||
|
||||
# Parse JSON output
|
||||
matches = []
|
||||
match_count = 0
|
||||
lines_truncated = False
|
||||
|
||||
for line in stdout.decode("utf-8", errors="ignore").splitlines():
|
||||
if not line.strip() or match_count >= effective_limit:
|
||||
break
|
||||
|
||||
try:
|
||||
event = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
|
||||
if event.get("type") == "match":
|
||||
match_count += 1
|
||||
data = event.get("data", {})
|
||||
file_path = data.get("path", {}).get("text", "")
|
||||
line_number = data.get("line_number", 0)
|
||||
|
||||
if file_path and line_number:
|
||||
matches.append({"file_path": file_path, "line_number": line_number})
|
||||
|
||||
if match_count >= effective_limit:
|
||||
break
|
||||
|
||||
# Check for errors
|
||||
if process.returncode not in (0, 1) and match_count == 0:
|
||||
error_msg = stderr.decode("utf-8", errors="ignore").strip()
|
||||
if error_msg:
|
||||
raise RuntimeError(error_msg)
|
||||
raise RuntimeError(f"ripgrep exited with code {process.returncode}")
|
||||
|
||||
# No matches found
|
||||
if match_count == 0:
|
||||
return "No matches found"
|
||||
|
||||
# Format matches with context
|
||||
output_lines = []
|
||||
file_cache = {}
|
||||
|
||||
for match in matches:
|
||||
file_path = match["file_path"]
|
||||
line_number = match["line_number"]
|
||||
|
||||
# Read file if not cached
|
||||
if file_path not in file_cache:
|
||||
try:
|
||||
with open(file_path, "r", encoding="utf-8", errors="ignore") as f:
|
||||
file_cache[file_path] = f.read().replace("\r\n", "\n").replace("\r", "\n").split("\n")
|
||||
except Exception:
|
||||
file_cache[file_path] = []
|
||||
|
||||
lines = file_cache[file_path]
|
||||
|
||||
# Format relative path
|
||||
if is_directory:
|
||||
relative_path = os.path.relpath(file_path, search_path)
|
||||
if not relative_path.startswith(".."):
|
||||
display_path = relative_path.replace("\\", "/")
|
||||
else:
|
||||
display_path = os.path.basename(file_path)
|
||||
else:
|
||||
display_path = os.path.basename(file_path)
|
||||
|
||||
# Generate context block
|
||||
if not lines:
|
||||
output_lines.append(f"{display_path}:{line_number}: (unable to read file)")
|
||||
continue
|
||||
|
||||
context_value = max(0, context_lines)
|
||||
start = max(1, line_number - context_value) if context_value > 0 else line_number
|
||||
end = min(len(lines), line_number + context_value) if context_value > 0 else line_number
|
||||
|
||||
for current in range(start, end + 1):
|
||||
if current < 1 or current > len(lines):
|
||||
continue
|
||||
|
||||
line_text = lines[current - 1]
|
||||
is_match_line = current == line_number
|
||||
|
||||
# Truncate long lines
|
||||
truncated_text, was_truncated = truncate_line(line_text)
|
||||
if was_truncated:
|
||||
lines_truncated = True
|
||||
|
||||
if is_match_line:
|
||||
output_lines.append(f"{display_path}:{current}: {truncated_text}")
|
||||
else:
|
||||
output_lines.append(f"{display_path}-{current}- {truncated_text}")
|
||||
|
||||
# Apply byte truncation
|
||||
raw_output = "\n".join(output_lines)
|
||||
truncation = truncate_head(raw_output, max_lines=999999999)
|
||||
|
||||
output = truncation.content
|
||||
notices = []
|
||||
|
||||
# Add notices
|
||||
if match_count >= effective_limit:
|
||||
notices.append(
|
||||
f"{effective_limit} matches limit reached. "
|
||||
f"Use limit={effective_limit * 2} for more, or refine pattern",
|
||||
)
|
||||
|
||||
if truncation.truncated:
|
||||
notices.append(f"{format_size(DEFAULT_MAX_BYTES)} limit reached")
|
||||
|
||||
if lines_truncated:
|
||||
notices.append(
|
||||
f"Some lines truncated to {GREP_MAX_LINE_LENGTH} chars. " f"Use read tool to see full lines",
|
||||
)
|
||||
|
||||
if notices:
|
||||
output += f"\n\n[{'. '.join(notices)}]"
|
||||
|
||||
return output
|
||||
128
reme/tool/fs/ls_tool.py
Normal file
128
reme/tool/fs/ls_tool.py
Normal file
|
|
@ -0,0 +1,128 @@
|
|||
"""Directory listing tool with truncation support."""
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from .truncate import DEFAULT_MAX_BYTES, truncate_head
|
||||
from ...core.op import BaseTool
|
||||
from ...core.schema import ToolCall
|
||||
|
||||
DEFAULT_LIMIT = 500
|
||||
|
||||
|
||||
class LsTool(BaseTool):
|
||||
"""List directory contents with smart truncation.
|
||||
|
||||
Features:
|
||||
- Returns entries sorted alphabetically (case-insensitive)
|
||||
- Directory indicators ('/' suffix)
|
||||
- Includes dotfiles
|
||||
- Entry count limiting
|
||||
- Byte truncation
|
||||
"""
|
||||
|
||||
def __init__(self, cwd: str | None = None):
|
||||
"""Initialize ls tool.
|
||||
|
||||
Args:
|
||||
cwd: Working directory (defaults to current directory)
|
||||
"""
|
||||
super().__init__()
|
||||
self.cwd = cwd or os.getcwd()
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
max_kb = DEFAULT_MAX_BYTES // 1024
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": (
|
||||
f"List directory contents. Returns entries sorted alphabetically, "
|
||||
f"with '/' suffix for directories. Includes dotfiles. Output is truncated "
|
||||
f"to {DEFAULT_LIMIT} entries or {max_kb}KB (whichever is hit first)."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"path": {
|
||||
"type": "string",
|
||||
"description": "Directory to list (default: current directory)",
|
||||
},
|
||||
"limit": {
|
||||
"type": "number",
|
||||
"description": f"Maximum number of entries to return (default: {DEFAULT_LIMIT})",
|
||||
},
|
||||
},
|
||||
"required": [],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self) -> str:
|
||||
"""List directory contents with production-grade features."""
|
||||
path: str | None = self.context.get("path", None)
|
||||
limit: int | None = self.context.get("limit", None)
|
||||
|
||||
# Resolve directory path
|
||||
dir_path = Path(self.cwd) / (path or ".")
|
||||
dir_path = dir_path.resolve()
|
||||
effective_limit = limit if limit is not None else DEFAULT_LIMIT
|
||||
|
||||
# Check if path exists
|
||||
if not dir_path.exists():
|
||||
raise FileNotFoundError(f"Path not found: {dir_path}")
|
||||
|
||||
# Check if path is a directory
|
||||
if not dir_path.is_dir():
|
||||
raise NotADirectoryError(f"Not a directory: {dir_path}")
|
||||
|
||||
# Read directory entries
|
||||
try:
|
||||
entries = list(dir_path.iterdir())
|
||||
except Exception as e:
|
||||
raise PermissionError(f"Cannot read directory: {e}") from e
|
||||
|
||||
# Sort alphabetically (case-insensitive)
|
||||
entries.sort(key=lambda e: e.name.lower())
|
||||
|
||||
# Format entries with directory indicators
|
||||
results: list[str] = []
|
||||
entry_limit_reached = False
|
||||
|
||||
for entry in entries:
|
||||
if len(results) >= effective_limit:
|
||||
entry_limit_reached = True
|
||||
break
|
||||
|
||||
try:
|
||||
# Add '/' suffix for directories
|
||||
suffix = "/" if entry.is_dir() else ""
|
||||
results.append(entry.name + suffix)
|
||||
except Exception:
|
||||
# Skip entries we can't stat
|
||||
continue
|
||||
|
||||
# Handle empty directory
|
||||
if len(results) == 0:
|
||||
return "(empty directory)"
|
||||
|
||||
# Apply byte truncation
|
||||
raw_output = "\n".join(results)
|
||||
truncation_result = truncate_head(raw_output, max_lines=float("inf"))
|
||||
|
||||
output_text = truncation_result.content
|
||||
|
||||
# Build notices
|
||||
notices: list[str] = []
|
||||
|
||||
if entry_limit_reached:
|
||||
notices.append(
|
||||
f"{effective_limit} entries limit reached. Use limit={effective_limit * 2} for more",
|
||||
)
|
||||
|
||||
if truncation_result.truncated:
|
||||
max_kb = DEFAULT_MAX_BYTES // 1024
|
||||
notices.append(f"{max_kb}KB limit reached")
|
||||
|
||||
if notices:
|
||||
output_text += f"\n\n[{'. '.join(notices)}]"
|
||||
|
||||
return output_text
|
||||
219
reme/tool/fs/read_tool.py
Normal file
219
reme/tool/fs/read_tool.py
Normal file
|
|
@ -0,0 +1,219 @@
|
|||
"""Read file tool with smart truncation and image support.
|
||||
|
||||
Features:
|
||||
- Reads text files with offset/limit support
|
||||
- Detects and handles image files (jpg, png, gif, webp)
|
||||
- Smart truncation to prevent memory issues
|
||||
"""
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from .truncate import DEFAULT_MAX_BYTES, DEFAULT_MAX_LINES, format_size, truncate_head
|
||||
from ...core.op import BaseTool
|
||||
from ...core.schema import ToolCall
|
||||
|
||||
# Supported image extensions
|
||||
IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".gif", ".webp"}
|
||||
|
||||
|
||||
def is_image_file(path: str) -> bool:
|
||||
"""Check if file is a supported image type.
|
||||
|
||||
Args:
|
||||
path: File path to check
|
||||
|
||||
Returns:
|
||||
True if file is a supported image
|
||||
"""
|
||||
return Path(path).suffix.lower() in IMAGE_EXTENSIONS
|
||||
|
||||
|
||||
class ReadTool(BaseTool):
|
||||
"""Read file contents with smart truncation.
|
||||
|
||||
Features:
|
||||
- Supports text files and images (jpg, png, gif, webp)
|
||||
- Smart truncation for large files
|
||||
- Offset/limit for reading specific portions
|
||||
"""
|
||||
|
||||
def __init__(self, cwd: str | None = None):
|
||||
"""Initialize read tool.
|
||||
|
||||
Args:
|
||||
cwd: Working directory (defaults to current directory)
|
||||
"""
|
||||
super().__init__()
|
||||
self.cwd = cwd or os.getcwd()
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
max_kb = DEFAULT_MAX_BYTES // 1024
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": (
|
||||
f"Read the contents of a file. Supports text files and images "
|
||||
f"(jpg, png, gif, webp). Images are sent as attachments. For text files, "
|
||||
f"output is truncated to {DEFAULT_MAX_LINES} lines or {max_kb}KB "
|
||||
f"(whichever is hit first). Use offset/limit for large files. "
|
||||
f"When you need the full file, continue with offset until complete."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"path": {
|
||||
"type": "string",
|
||||
"description": "Path to the file to read (relative or absolute)",
|
||||
},
|
||||
"offset": {
|
||||
"type": "number",
|
||||
"description": "Line number to start reading from (1-indexed)",
|
||||
},
|
||||
"limit": {
|
||||
"type": "number",
|
||||
"description": "Maximum number of lines to read",
|
||||
},
|
||||
},
|
||||
"required": ["path"],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self) -> str:
|
||||
"""Execute the read operation."""
|
||||
path: str = self.context.path
|
||||
offset: int | None = self.context.get("offset", None)
|
||||
limit: int | None = self.context.get("limit", None)
|
||||
|
||||
# Resolve path
|
||||
if not os.path.isabs(path):
|
||||
absolute_path = os.path.join(self.cwd, path)
|
||||
else:
|
||||
absolute_path = path
|
||||
absolute_path = os.path.normpath(absolute_path)
|
||||
|
||||
# Check file exists and is readable
|
||||
if not os.path.exists(absolute_path):
|
||||
raise ValueError(f"File not found: {path}")
|
||||
|
||||
if not os.path.isfile(absolute_path):
|
||||
raise ValueError(f"Not a file: {path}")
|
||||
|
||||
if not os.access(absolute_path, os.R_OK):
|
||||
raise ValueError(f"File not readable: {path}")
|
||||
|
||||
# Check if image
|
||||
if is_image_file(absolute_path):
|
||||
return await self._read_image(absolute_path, path)
|
||||
else:
|
||||
return await self._read_text(absolute_path, path, offset, limit)
|
||||
|
||||
@staticmethod
|
||||
async def _read_image(absolute_path: str, display_path: str) -> str:
|
||||
"""Read and return image file information.
|
||||
|
||||
Args:
|
||||
absolute_path: Absolute path to image
|
||||
display_path: Path to display to user
|
||||
|
||||
Returns:
|
||||
Image information text
|
||||
"""
|
||||
# Get file size
|
||||
file_size = os.path.getsize(absolute_path)
|
||||
file_ext = Path(absolute_path).suffix.lower()
|
||||
|
||||
# For Python tools, we typically can't return image data directly to LLM
|
||||
# So we return a descriptive message
|
||||
return (
|
||||
f"Read image file [{file_ext}]\n"
|
||||
f"Path: {display_path}\n"
|
||||
f"Size: {format_size(file_size)}\n"
|
||||
f"Note: Image content cannot be displayed in text format. "
|
||||
f"Use bash tool or other methods to process the image."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _read_text(
|
||||
absolute_path: str,
|
||||
_display_path: str,
|
||||
offset: int | None,
|
||||
limit: int | None,
|
||||
) -> str:
|
||||
"""Read text file with smart truncation.
|
||||
|
||||
Args:
|
||||
absolute_path: Absolute path to file
|
||||
_display_path: Path to display to user
|
||||
offset: Starting line (1-indexed)
|
||||
limit: Maximum lines to read
|
||||
|
||||
Returns:
|
||||
File contents with truncation notices
|
||||
"""
|
||||
# Read file
|
||||
try:
|
||||
with open(absolute_path, "r", encoding="utf-8") as f:
|
||||
content = f.read()
|
||||
except UnicodeDecodeError:
|
||||
# Try with error handling for binary files
|
||||
with open(absolute_path, "r", encoding="utf-8", errors="ignore") as f:
|
||||
content = f.read()
|
||||
|
||||
all_lines = content.split("\n")
|
||||
total_file_lines = len(all_lines)
|
||||
|
||||
# Apply offset if specified (convert 1-indexed to 0-indexed)
|
||||
start_line = max(0, (offset - 1)) if offset else 0
|
||||
start_line_display = start_line + 1
|
||||
|
||||
# Check offset bounds
|
||||
if start_line >= len(all_lines):
|
||||
raise IndexError(
|
||||
f"Offset {offset} is beyond end of file ({len(all_lines)} lines total)",
|
||||
)
|
||||
|
||||
# Apply user limit if specified
|
||||
if limit is not None:
|
||||
end_line = min(start_line + limit, len(all_lines))
|
||||
selected_content = "\n".join(all_lines[start_line:end_line])
|
||||
user_limited_lines = end_line - start_line
|
||||
else:
|
||||
selected_content = "\n".join(all_lines[start_line:])
|
||||
user_limited_lines = None
|
||||
|
||||
# Apply truncation
|
||||
truncation = truncate_head(selected_content)
|
||||
|
||||
# Build output with truncation notices
|
||||
if truncation.truncated:
|
||||
# Truncation occurred
|
||||
end_line_display = start_line_display + truncation.output_lines - 1
|
||||
next_offset = end_line_display + 1
|
||||
|
||||
output_text = truncation.content
|
||||
|
||||
if truncation.truncated_by == "lines":
|
||||
output_text += (
|
||||
f"\n\n[Showing lines {start_line_display}-{end_line_display} "
|
||||
f"of {total_file_lines}. Use offset={next_offset} to continue.]"
|
||||
)
|
||||
else:
|
||||
max_kb = DEFAULT_MAX_BYTES // 1024
|
||||
output_text += (
|
||||
f"\n\n[Showing lines {start_line_display}-{end_line_display} "
|
||||
f"of {total_file_lines} ({max_kb}KB limit). "
|
||||
f"Use offset={next_offset} to continue.]"
|
||||
)
|
||||
elif user_limited_lines is not None and start_line + user_limited_lines < len(all_lines):
|
||||
# User limit exceeded, but no truncation
|
||||
remaining = len(all_lines) - (start_line + user_limited_lines)
|
||||
next_offset = start_line + user_limited_lines + 1
|
||||
|
||||
output_text = truncation.content
|
||||
output_text += f"\n\n[{remaining} more lines in file. " f"Use offset={next_offset} to continue.]"
|
||||
else:
|
||||
# No truncation or user limit exceeded
|
||||
output_text = truncation.content
|
||||
|
||||
return output_text
|
||||
209
reme/tool/fs/truncate.py
Normal file
209
reme/tool/fs/truncate.py
Normal file
|
|
@ -0,0 +1,209 @@
|
|||
"""fs utils"""
|
||||
|
||||
from typing import Literal
|
||||
|
||||
from ...core.schema import TruncationResult
|
||||
|
||||
# Default limits for output truncation
|
||||
DEFAULT_MAX_LINES = 1000 # Maximum lines to keep for tail truncation
|
||||
DEFAULT_MAX_BYTES = 30 * 1024 # Maximum bytes to keep (30KB)
|
||||
|
||||
# Find tool limits
|
||||
FIND_MAX_LINES = 2000 # Maximum lines for find output
|
||||
FIND_MAX_BYTES = 50 * 1024 # 50KB for find output
|
||||
|
||||
# Grep tool limits
|
||||
GREP_MAX_LINE_LENGTH = 500 # Maximum line length for grep output
|
||||
|
||||
|
||||
def format_size(num_bytes: int) -> str:
|
||||
"""Format byte size in human-readable format.
|
||||
|
||||
Args:
|
||||
num_bytes: Number of bytes
|
||||
|
||||
Returns:
|
||||
Formatted string (e.g., "1.5KB", "2.3MB")
|
||||
"""
|
||||
if num_bytes < 1024:
|
||||
return f"{num_bytes}B"
|
||||
elif num_bytes < 1024 * 1024:
|
||||
return f"{num_bytes / 1024:.1f}KB"
|
||||
else:
|
||||
return f"{num_bytes / (1024 * 1024):.1f}MB"
|
||||
|
||||
|
||||
def truncate_line(text: str, max_length: int = GREP_MAX_LINE_LENGTH) -> tuple[str, bool]:
|
||||
"""Truncate a single line if it exceeds max length.
|
||||
|
||||
Args:
|
||||
text: Line text
|
||||
max_length: Maximum line length
|
||||
|
||||
Returns:
|
||||
Tuple of (truncated_text, was_truncated)
|
||||
"""
|
||||
if len(text) <= max_length:
|
||||
return text, False
|
||||
return text[:max_length] + "...", True
|
||||
|
||||
|
||||
def truncate_tail(
|
||||
text: str,
|
||||
max_lines: int = DEFAULT_MAX_LINES,
|
||||
max_bytes: int = DEFAULT_MAX_BYTES,
|
||||
) -> TruncationResult:
|
||||
"""Truncate text to keep only the tail (last portion).
|
||||
|
||||
Keeps the last N lines or M bytes, whichever is hit first.
|
||||
This is useful for command outputs where the end is most relevant.
|
||||
|
||||
Args:
|
||||
text: The text to truncate
|
||||
max_lines: Maximum number of lines to keep
|
||||
max_bytes: Maximum bytes to keep
|
||||
|
||||
Returns:
|
||||
TruncationResult with truncated content and metadata
|
||||
"""
|
||||
if not text:
|
||||
return TruncationResult(
|
||||
content="",
|
||||
truncated=False,
|
||||
total_lines=0,
|
||||
output_lines=0,
|
||||
total_bytes=0,
|
||||
output_bytes=0,
|
||||
)
|
||||
|
||||
total_bytes = len(text.encode("utf-8"))
|
||||
lines = text.split("\n")
|
||||
total_lines = len(lines)
|
||||
|
||||
# Check if we need to truncate
|
||||
if total_lines <= max_lines and total_bytes <= max_bytes:
|
||||
return TruncationResult(
|
||||
content=text,
|
||||
truncated=False,
|
||||
total_lines=total_lines,
|
||||
output_lines=total_lines,
|
||||
total_bytes=total_bytes,
|
||||
output_bytes=total_bytes,
|
||||
)
|
||||
|
||||
# Keep last N lines
|
||||
kept_lines = lines[-max_lines:] if total_lines > max_lines else lines
|
||||
truncated_by: Literal["lines", "bytes"] = "lines" if total_lines > max_lines else "bytes"
|
||||
|
||||
# Check byte limit on kept lines
|
||||
kept_text = "\n".join(kept_lines)
|
||||
kept_bytes = len(kept_text.encode("utf-8"))
|
||||
|
||||
# If still over byte limit, truncate further
|
||||
last_line_partial = False
|
||||
if kept_bytes > max_bytes:
|
||||
truncated_by = "bytes"
|
||||
# Keep truncating from the start until under byte limit
|
||||
while kept_lines and len("\n".join(kept_lines).encode("utf-8")) > max_bytes:
|
||||
kept_lines.pop(0)
|
||||
|
||||
# If still over (single line > max_bytes), truncate the line itself
|
||||
if kept_lines and len("\n".join(kept_lines).encode("utf-8")) > max_bytes:
|
||||
last_line = kept_lines[-1]
|
||||
# Binary search to find how much of last line fits
|
||||
encoded = last_line.encode("utf-8")
|
||||
if len(encoded) > max_bytes:
|
||||
last_line_partial = True
|
||||
# Take last max_bytes of the line
|
||||
kept_lines[-1] = encoded[-max_bytes:].decode("utf-8", errors="ignore")
|
||||
|
||||
kept_text = "\n".join(kept_lines)
|
||||
kept_bytes = len(kept_text.encode("utf-8"))
|
||||
|
||||
return TruncationResult(
|
||||
content=kept_text,
|
||||
truncated=True,
|
||||
total_lines=total_lines,
|
||||
output_lines=len(kept_lines),
|
||||
total_bytes=total_bytes,
|
||||
output_bytes=kept_bytes,
|
||||
truncated_by=truncated_by,
|
||||
last_line_partial=last_line_partial,
|
||||
)
|
||||
|
||||
|
||||
def truncate_head(
|
||||
text: str,
|
||||
max_lines: int = FIND_MAX_LINES,
|
||||
max_bytes: int = FIND_MAX_BYTES,
|
||||
) -> TruncationResult:
|
||||
"""Truncate text to keep only the head (first portion).
|
||||
|
||||
Keeps the first N lines or M bytes, whichever is hit first.
|
||||
Suitable for file reads where you want to see the beginning.
|
||||
|
||||
Args:
|
||||
text: The text to truncate
|
||||
max_lines: Maximum number of lines to keep
|
||||
max_bytes: Maximum bytes to keep
|
||||
|
||||
Returns:
|
||||
TruncationResult with truncated content and metadata
|
||||
"""
|
||||
if not text:
|
||||
return TruncationResult(
|
||||
content="",
|
||||
truncated=False,
|
||||
total_lines=0,
|
||||
output_lines=0,
|
||||
total_bytes=0,
|
||||
output_bytes=0,
|
||||
)
|
||||
|
||||
total_bytes = len(text.encode("utf-8"))
|
||||
lines = text.split("\n")
|
||||
total_lines = len(lines)
|
||||
|
||||
# Check if no truncation needed
|
||||
if total_lines <= max_lines and total_bytes <= max_bytes:
|
||||
return TruncationResult(
|
||||
content=text,
|
||||
truncated=False,
|
||||
total_lines=total_lines,
|
||||
output_lines=total_lines,
|
||||
total_bytes=total_bytes,
|
||||
output_bytes=total_bytes,
|
||||
)
|
||||
|
||||
# Collect complete lines that fit
|
||||
kept_lines = []
|
||||
kept_bytes = 0
|
||||
truncated_by: Literal["lines", "bytes"] = "lines"
|
||||
|
||||
for i, line in enumerate(lines):
|
||||
if i >= max_lines:
|
||||
truncated_by = "lines"
|
||||
break
|
||||
|
||||
# Calculate bytes for this line (+1 for newline except first line)
|
||||
line_bytes = len(line.encode("utf-8")) + (1 if i > 0 else 0)
|
||||
|
||||
if kept_bytes + line_bytes > max_bytes:
|
||||
truncated_by = "bytes"
|
||||
break
|
||||
|
||||
kept_lines.append(line)
|
||||
kept_bytes += line_bytes
|
||||
|
||||
kept_text = "\n".join(kept_lines)
|
||||
final_bytes = len(kept_text.encode("utf-8"))
|
||||
|
||||
return TruncationResult(
|
||||
content=kept_text,
|
||||
truncated=True,
|
||||
total_lines=total_lines,
|
||||
output_lines=len(kept_lines),
|
||||
total_bytes=total_bytes,
|
||||
output_bytes=final_bytes,
|
||||
truncated_by=truncated_by,
|
||||
)
|
||||
81
reme/tool/fs/write_tool.py
Normal file
81
reme/tool/fs/write_tool.py
Normal file
|
|
@ -0,0 +1,81 @@
|
|||
"""Write tool for creating and overwriting files.
|
||||
|
||||
This module provides a tool for writing content to files with:
|
||||
- Automatic parent directory creation
|
||||
- File overwriting (creates if doesn't exist, overwrites if exists)
|
||||
- Path resolution (relative to working directory)
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
from ...core.op import BaseTool
|
||||
from ...core.schema import ToolCall
|
||||
|
||||
|
||||
class WriteTool(BaseTool):
|
||||
"""Tool for writing content to files.
|
||||
|
||||
Features:
|
||||
- Creates file if it doesn't exist, overwrites if it does
|
||||
- Automatically creates parent directories
|
||||
- Supports both relative and absolute paths
|
||||
"""
|
||||
|
||||
def __init__(self, cwd: str | None = None):
|
||||
"""Initialize write tool.
|
||||
|
||||
Args:
|
||||
cwd: Working directory (defaults to current directory)
|
||||
"""
|
||||
super().__init__()
|
||||
self.cwd = cwd or os.getcwd()
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": (
|
||||
"Write content to a file. Creates the file if it doesn't exist, "
|
||||
"overwrites if it does. Automatically creates parent directories."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"path": {
|
||||
"type": "string",
|
||||
"description": "Path to the file to write (relative or absolute)",
|
||||
},
|
||||
"content": {
|
||||
"type": "string",
|
||||
"description": "Content to write to the file",
|
||||
},
|
||||
},
|
||||
"required": ["path", "content"],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self) -> str:
|
||||
"""Execute the write operation."""
|
||||
path: str = self.context.path
|
||||
content: str = self.context.content
|
||||
|
||||
# Resolve path to absolute
|
||||
if not os.path.isabs(path):
|
||||
absolute_path = os.path.join(self.cwd, path)
|
||||
else:
|
||||
absolute_path = path
|
||||
|
||||
absolute_path = os.path.normpath(absolute_path)
|
||||
|
||||
# Create parent directories if needed
|
||||
parent_dir = os.path.dirname(absolute_path)
|
||||
if parent_dir:
|
||||
os.makedirs(parent_dir, exist_ok=True)
|
||||
|
||||
# Write the file
|
||||
with open(absolute_path, "w", encoding="utf-8") as f:
|
||||
f.write(content)
|
||||
|
||||
# Return success message
|
||||
content_bytes = len(content.encode("utf-8"))
|
||||
return f"Successfully wrote {content_bytes} bytes to {path}"
|
||||
400
tests/demo_memory_search.py
Normal file
400
tests/demo_memory_search.py
Normal 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())
|
||||
209
tests/test_cache_memory_usage.py
Normal file
209
tests/test_cache_memory_usage.py
Normal 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")
|
||||
268
tests/test_chunking_utils.py
Normal file
268
tests/test_chunking_utils.py
Normal 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"])
|
||||
394
tests/test_embedding_cache.py
Normal file
394
tests/test_embedding_cache.py
Normal 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())
|
||||
445
tests/test_file_system_tool.py
Normal file
445
tests/test_file_system_tool.py
Normal file
|
|
@ -0,0 +1,445 @@
|
|||
"""Tests for file system tools including bash, edit, find, grep, ls, read, and write tools."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
async def test_bash_tool():
|
||||
"""Test BashTool."""
|
||||
from reme.tool.fs import BashTool
|
||||
|
||||
print("=== Testing BashTool ===")
|
||||
bash_tool = BashTool()
|
||||
result = await bash_tool.call(command="echo 'Hello World'")
|
||||
print(f"Result: {result}")
|
||||
assert "Hello World" in result
|
||||
print("✓ BashTool test passed\n")
|
||||
|
||||
|
||||
async def test_edit_tool():
|
||||
"""Test EditTool."""
|
||||
from reme.tool.fs import EditTool
|
||||
|
||||
print("=== Testing EditTool ===")
|
||||
|
||||
# Create temp file
|
||||
with tempfile.NamedTemporaryFile(mode="w", suffix=".txt", delete=False) as f:
|
||||
temp_path = f.name
|
||||
f.write("Hello World\nThis is a test\nGoodbye World\n")
|
||||
|
||||
try:
|
||||
# Test edit
|
||||
edit_tool = EditTool()
|
||||
result = await edit_tool.call(
|
||||
path=temp_path,
|
||||
oldText="This is a test",
|
||||
newText="This is an updated test",
|
||||
)
|
||||
print(f"Result: {result}")
|
||||
|
||||
# Verify content
|
||||
with open(temp_path, "r", encoding="utf-8") as f:
|
||||
content = f.read()
|
||||
|
||||
assert "This is an updated test" in content
|
||||
assert "This is a test" not in content
|
||||
print("✓ EditTool test passed\n")
|
||||
|
||||
# Test error: file not found
|
||||
print("=== Testing file not found error ===")
|
||||
result = await edit_tool.call(
|
||||
path="/nonexistent/file.txt",
|
||||
oldText="test",
|
||||
newText="new",
|
||||
)
|
||||
print(f"Expected error result: {result}")
|
||||
assert "failed" in result and "File not found" in result
|
||||
print("✓ File not found error test passed\n")
|
||||
|
||||
# Test error: text not found
|
||||
print("=== Testing text not found error ===")
|
||||
result = await edit_tool.call(
|
||||
path=temp_path,
|
||||
oldText="nonexistent text",
|
||||
newText="new",
|
||||
)
|
||||
print(f"Expected error result: {result}")
|
||||
assert "failed" in result and "Could not find" in result
|
||||
print("✓ Text not found error test passed\n")
|
||||
|
||||
finally:
|
||||
# Cleanup
|
||||
if os.path.exists(temp_path):
|
||||
os.unlink(temp_path)
|
||||
|
||||
|
||||
async def test_find_tool():
|
||||
"""Test FindTool."""
|
||||
from reme.tool.fs import FindTool
|
||||
|
||||
print("=== Testing FindTool ===")
|
||||
|
||||
# Create temp directory with test files
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
temp_path = Path(temp_dir)
|
||||
|
||||
# Create some test files
|
||||
(temp_path / "test1.txt").write_text("test file 1")
|
||||
(temp_path / "test2.txt").write_text("test file 2")
|
||||
(temp_path / "readme.md").write_text("readme")
|
||||
|
||||
# Create subdirectory with files
|
||||
sub_dir = temp_path / "subdir"
|
||||
sub_dir.mkdir()
|
||||
(sub_dir / "test3.txt").write_text("test file 3")
|
||||
(sub_dir / "config.json").write_text("{}")
|
||||
|
||||
# Create .gitignore to ignore certain files
|
||||
(temp_path / ".gitignore").write_text("*.md\n")
|
||||
|
||||
# Test: find all txt files
|
||||
find_tool = FindTool(cwd=str(temp_path))
|
||||
result = await find_tool.call(pattern="*.txt")
|
||||
print(f"Find *.txt result:\n{result}")
|
||||
assert "test1.txt" in result
|
||||
assert "test2.txt" in result
|
||||
assert "readme.md" not in result # Should be ignored by .gitignore
|
||||
print("✓ Find *.txt test passed\n")
|
||||
|
||||
# Test: find with recursive pattern
|
||||
result = await find_tool.call(pattern="**/*.txt")
|
||||
print(f"Find **/*.txt result:\n{result}")
|
||||
assert "test1.txt" in result
|
||||
assert "subdir/test3.txt" in result or "test3.txt" in result
|
||||
print("✓ Find **/*.txt test passed\n")
|
||||
|
||||
# Test: find with no matches
|
||||
result = await find_tool.call(pattern="*.nonexistent")
|
||||
print(f"Find *.nonexistent result:\n{result}")
|
||||
assert "No files found" in result
|
||||
print("✓ No matches test passed\n")
|
||||
|
||||
# Test: error - directory not found
|
||||
print("=== Testing directory not found error ===")
|
||||
result = await find_tool.call(pattern="*.txt", path="/nonexistent/dir")
|
||||
print(f"Expected error result: {result}")
|
||||
assert "failed" in result and "Path not found" in result
|
||||
print("✓ Directory not found error test passed\n")
|
||||
|
||||
|
||||
async def test_grep_tool():
|
||||
"""Test GrepTool."""
|
||||
from reme.tool.fs import GrepTool
|
||||
|
||||
print("=== Testing GrepTool ===")
|
||||
|
||||
# Create temp directory with test files
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
temp_path = Path(temp_dir)
|
||||
|
||||
# Create test files with searchable content
|
||||
(temp_path / "file1.txt").write_text("Hello World\nThis is a test\nGoodbye World\n")
|
||||
(temp_path / "file2.txt").write_text("Another test file\nWith multiple lines\nHello again\n")
|
||||
(temp_path / "script.py").write_text("def hello():\n print('Hello')\n return True\n")
|
||||
|
||||
# Create subdirectory with files
|
||||
sub_dir = temp_path / "subdir"
|
||||
sub_dir.mkdir()
|
||||
(sub_dir / "nested.txt").write_text("Nested file content\nWith hello keyword\n")
|
||||
|
||||
# Test: search for pattern
|
||||
grep_tool = GrepTool(cwd=str(temp_path))
|
||||
result = await grep_tool.call(pattern="Hello", path=str(temp_path))
|
||||
print(f"Search 'Hello' result:\n{result}")
|
||||
assert "file1.txt" in result
|
||||
assert "Hello World" in result or "Hello" in result
|
||||
print("✓ Basic search test passed\n")
|
||||
|
||||
# Test: case-insensitive search
|
||||
result = await grep_tool.call(pattern="hello", path=str(temp_path), ignoreCase=True)
|
||||
print(f"Case-insensitive search result:\n{result}")
|
||||
assert "file1.txt" in result or "Hello" in result.lower()
|
||||
print("✓ Case-insensitive search test passed\n")
|
||||
|
||||
# Test: literal string search
|
||||
result = await grep_tool.call(pattern="Hello()", path=str(temp_path), literal=True)
|
||||
print(f"Literal search result:\n{result}")
|
||||
# Should not find regex interpretation
|
||||
print("✓ Literal search test passed\n")
|
||||
|
||||
# Test: glob filter
|
||||
result = await grep_tool.call(pattern="Hello", path=str(temp_path), glob="*.txt")
|
||||
print(f"Glob filter *.txt result:\n{result}")
|
||||
assert "file1.txt" in result or "file2.txt" in result
|
||||
assert ".py" not in result # Python files should be excluded
|
||||
print("✓ Glob filter test passed\n")
|
||||
|
||||
# Test: context lines
|
||||
result = await grep_tool.call(pattern="test", path=str(temp_path), contextLines=1)
|
||||
print(f"Context lines result:\n{result}")
|
||||
# Should include lines before and after matches
|
||||
print("✓ Context lines test passed\n")
|
||||
|
||||
# Test: limit matches
|
||||
result = await grep_tool.call(pattern="Hello", path=str(temp_path), limit=1)
|
||||
print(f"Limit to 1 match result:\n{result}")
|
||||
assert "limit reached" in result or result.count(":") >= 1
|
||||
print("✓ Limit test passed\n")
|
||||
|
||||
# Test: no matches
|
||||
result = await grep_tool.call(pattern="nonexistent_pattern_xyz", path=str(temp_path))
|
||||
print(f"No matches result:\n{result}")
|
||||
assert "No matches found" in result
|
||||
print("✓ No matches test passed\n")
|
||||
|
||||
# Test: error - path not found
|
||||
print("=== Testing path not found error ===")
|
||||
try:
|
||||
result = await grep_tool.call(pattern="test", path="/nonexistent/path")
|
||||
assert "failed" in result and "not found" in result
|
||||
except Exception as e:
|
||||
assert "not found" in str(e).lower()
|
||||
print("✓ Path not found error test passed\n")
|
||||
|
||||
|
||||
async def test_ls_tool():
|
||||
"""Test LsTool."""
|
||||
from reme.tool.fs import LsTool
|
||||
|
||||
print("=== Testing LsTool ===")
|
||||
|
||||
# Create temp directory with test files
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
temp_path = Path(temp_dir)
|
||||
|
||||
# Create test files and directories
|
||||
(temp_path / "file1.txt").write_text("test file 1")
|
||||
(temp_path / "file2.py").write_text("test file 2")
|
||||
(temp_path / ".hidden").write_text("hidden file")
|
||||
(temp_path / "README.md").write_text("readme")
|
||||
|
||||
# Create subdirectories
|
||||
(temp_path / "subdir1").mkdir()
|
||||
(temp_path / "subdir2").mkdir()
|
||||
|
||||
# Test: list current directory
|
||||
ls_tool = LsTool(cwd=str(temp_path))
|
||||
result = await ls_tool.call()
|
||||
print(f"List directory result:\n{result}")
|
||||
assert ".hidden" in result # Includes dotfiles
|
||||
assert "file1.txt" in result
|
||||
assert "file2.py" in result
|
||||
assert "subdir1/" in result # Directories have '/' suffix
|
||||
assert "subdir2/" in result
|
||||
print("✓ Basic ls test passed\n")
|
||||
|
||||
# Test: list specific path
|
||||
result = await ls_tool.call(path=".")
|
||||
print(f"List current directory result:\n{result}")
|
||||
assert "file1.txt" in result
|
||||
print("✓ Specific path test passed\n")
|
||||
|
||||
# Test: empty directory
|
||||
empty_dir = temp_path / "empty"
|
||||
empty_dir.mkdir()
|
||||
result = await ls_tool.call(path="empty")
|
||||
print(f"Empty directory result:\n{result}")
|
||||
assert "(empty directory)" in result
|
||||
print("✓ Empty directory test passed\n")
|
||||
|
||||
# Test: entry limit
|
||||
# Create many files
|
||||
for i in range(10):
|
||||
(temp_path / f"file{i:03d}.txt").write_text(f"file {i}")
|
||||
|
||||
result = await ls_tool.call(limit=5)
|
||||
print(f"Limited entries result:\n{result}")
|
||||
assert "entries limit reached" in result
|
||||
assert "limit=10" in result # Should suggest doubling the limit
|
||||
print("✓ Entry limit test passed\n")
|
||||
|
||||
# Test: error - path not found
|
||||
print("=== Testing path not found error ===")
|
||||
result = await ls_tool.call(path="/nonexistent/path")
|
||||
print(f"Expected error result: {result}")
|
||||
assert "failed" in result and "Path not found" in result
|
||||
print("✓ Path not found error test passed\n")
|
||||
|
||||
# Test: error - not a directory
|
||||
print("=== Testing not a directory error ===")
|
||||
result = await ls_tool.call(path="file1.txt")
|
||||
print(f"Expected error result: {result}")
|
||||
assert "failed" in result and "Not a directory" in result
|
||||
print("✓ Not a directory error test passed\n")
|
||||
|
||||
|
||||
async def test_read_tool():
|
||||
"""Test ReadTool."""
|
||||
from reme.tool.fs import ReadTool
|
||||
|
||||
print("=== Testing ReadTool ===")
|
||||
|
||||
# Create temp directory with test files
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
temp_path = Path(temp_dir)
|
||||
|
||||
# Create test text file
|
||||
test_file = temp_path / "test.txt"
|
||||
test_content = "\n".join([f"Line {i}" for i in range(1, 101)]) # 100 lines
|
||||
test_file.write_text(test_content)
|
||||
|
||||
# Create test image file
|
||||
image_file = temp_path / "test.jpg"
|
||||
image_file.write_bytes(b"\xff\xd8\xff\xe0") # Minimal JPEG header
|
||||
|
||||
# Test: read full file
|
||||
read_tool = ReadTool(cwd=str(temp_path))
|
||||
result = await read_tool.call(path="test.txt")
|
||||
print(f"Read full file result:\n{result[:200]}...")
|
||||
assert "Line 1" in result
|
||||
assert "Line 100" in result
|
||||
print("✓ Read full file test passed\n")
|
||||
|
||||
# Test: read with offset
|
||||
result = await read_tool.call(path="test.txt", offset=50)
|
||||
print(f"Read with offset=50 result:\n{result[:200]}...")
|
||||
# Check that we start from Line 50 (should be first line of content)
|
||||
assert result.startswith("Line 50"), f"Should start with 'Line 50', got: {result[:50]}"
|
||||
assert "Line 100" in result
|
||||
print("✓ Read with offset test passed\n")
|
||||
|
||||
# Test: read with limit
|
||||
result = await read_tool.call(path="test.txt", limit=10)
|
||||
print(f"Read with limit=10 result:\n{result}")
|
||||
assert "Line 1" in result
|
||||
assert "Line 10" in result or "more lines in file" in result
|
||||
assert "Line 50" not in result
|
||||
print("✓ Read with limit test passed\n")
|
||||
|
||||
# Test: read with offset and limit
|
||||
result = await read_tool.call(path="test.txt", offset=20, limit=5)
|
||||
print(f"Read with offset=20, limit=5 result:\n{result}")
|
||||
assert "Line 20" in result
|
||||
assert "Line 24" in result or "more lines" in result
|
||||
print("✓ Read with offset and limit test passed\n")
|
||||
|
||||
# Test: read image file
|
||||
result = await read_tool.call(path="test.jpg")
|
||||
print(f"Read image result:\n{result}")
|
||||
assert "image file" in result.lower() or ".jpg" in result.lower()
|
||||
print("✓ Read image test passed\n")
|
||||
|
||||
# Test: offset beyond file
|
||||
print("=== Testing offset beyond file error ===")
|
||||
result = await read_tool.call(path="test.txt", offset=200)
|
||||
print(f"Expected error result: {result}")
|
||||
assert "failed" in result and ("beyond end of file" in result or "offset" in result.lower())
|
||||
print("✓ Offset beyond file error test passed\n")
|
||||
|
||||
# Test: file not found
|
||||
print("=== Testing file not found error ===")
|
||||
result = await read_tool.call(path="nonexistent.txt")
|
||||
print(f"Expected error result: {result}")
|
||||
assert "failed" in result and "not found" in result.lower()
|
||||
print("✓ File not found error test passed\n")
|
||||
|
||||
# Test: read directory (should fail)
|
||||
print("=== Testing read directory error ===")
|
||||
sub_dir = temp_path / "subdir"
|
||||
sub_dir.mkdir()
|
||||
result = await read_tool.call(path="subdir")
|
||||
print(f"Expected error result: {result}")
|
||||
assert "failed" in result and ("Not a file" in result or "directory" in result.lower())
|
||||
print("✓ Read directory error test passed\n")
|
||||
|
||||
|
||||
async def test_write_tool():
|
||||
"""Test WriteTool."""
|
||||
from reme.tool.fs import WriteTool
|
||||
|
||||
print("=== Testing WriteTool ===")
|
||||
|
||||
# Create temp directory
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
temp_path = Path(temp_dir)
|
||||
|
||||
# Test: write new file
|
||||
write_tool = WriteTool(cwd=str(temp_path))
|
||||
test_content = "Hello World\nThis is a test file\n"
|
||||
result = await write_tool.call(path="test.txt", content=test_content)
|
||||
print(f"Write result: {result}")
|
||||
assert "Successfully wrote" in result
|
||||
assert "test.txt" in result
|
||||
|
||||
# Verify file was created
|
||||
test_file = temp_path / "test.txt"
|
||||
assert test_file.exists()
|
||||
assert test_file.read_text() == test_content
|
||||
print("✓ Write new file test passed\n")
|
||||
|
||||
# Test: overwrite existing file
|
||||
new_content = "Updated content\n"
|
||||
result = await write_tool.call(path="test.txt", content=new_content)
|
||||
print(f"Overwrite result: {result}")
|
||||
assert "Successfully wrote" in result
|
||||
|
||||
# Verify file was overwritten
|
||||
assert test_file.read_text() == new_content
|
||||
assert test_content not in test_file.read_text()
|
||||
print("✓ Overwrite existing file test passed\n")
|
||||
|
||||
# Test: create file with parent directories
|
||||
nested_path = "subdir1/subdir2/nested.txt"
|
||||
nested_content = "Nested file content"
|
||||
result = await write_tool.call(path=nested_path, content=nested_content)
|
||||
print(f"Create with parents result: {result}")
|
||||
assert "Successfully wrote" in result
|
||||
|
||||
# Verify nested file was created
|
||||
nested_file = temp_path / "subdir1" / "subdir2" / "nested.txt"
|
||||
assert nested_file.exists()
|
||||
assert nested_file.read_text() == nested_content
|
||||
print("✓ Create file with parent directories test passed\n")
|
||||
|
||||
# Test: write empty file
|
||||
result = await write_tool.call(path="empty.txt", content="")
|
||||
print(f"Write empty file result: {result}")
|
||||
assert "Successfully wrote" in result
|
||||
assert "0 bytes" in result
|
||||
|
||||
# Verify empty file
|
||||
empty_file = temp_path / "empty.txt"
|
||||
assert empty_file.exists()
|
||||
assert empty_file.read_text() == ""
|
||||
print("✓ Write empty file test passed\n")
|
||||
|
||||
# Test: write file with absolute path
|
||||
abs_path = str(temp_path / "absolute.txt")
|
||||
abs_content = "Absolute path content"
|
||||
result = await write_tool.call(path=abs_path, content=abs_content)
|
||||
print(f"Write absolute path result: {result}")
|
||||
assert "Successfully wrote" in result
|
||||
|
||||
# Verify absolute path file
|
||||
abs_file = Path(abs_path)
|
||||
assert abs_file.exists()
|
||||
assert abs_file.read_text(encoding="utf-8") == abs_content
|
||||
print("✓ Write absolute path test passed\n")
|
||||
|
||||
|
||||
async def main():
|
||||
"""Run all file system tool tests."""
|
||||
await test_bash_tool()
|
||||
await test_edit_tool()
|
||||
await test_find_tool()
|
||||
await test_grep_tool()
|
||||
await test_ls_tool()
|
||||
await test_read_tool()
|
||||
await test_write_tool()
|
||||
print("=== All tests passed! ===")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
323
tests/test_fs_agent.py
Normal file
323
tests/test_fs_agent.py
Normal 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
979
tests/test_memory_store.py
Normal 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())
|
||||
|
|
@ -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())
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue