mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-10 03:30:56 +00:00
feat(memory): implement file watcher with delta and full sync strategies
This commit is contained in:
parent
857cd52a8e
commit
c5438c52fc
41 changed files with 4453 additions and 1258 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/*
|
||||
|
|
@ -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].
|
||||
|
|
@ -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,8 +4,10 @@ Defines the abstract base class and standard API for all embedding model impleme
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import time
|
||||
from abc import ABC
|
||||
from collections import OrderedDict
|
||||
|
||||
from loguru import logger
|
||||
|
||||
|
|
@ -28,17 +30,35 @@ class BaseEmbeddingModel(ABC):
|
|||
max_retries: int = 3,
|
||||
raise_exception: bool = True,
|
||||
max_input_length: int = 8192,
|
||||
max_cache_size: int = 10000,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize model configuration and parameters."""
|
||||
"""Initialize model configuration and parameters.
|
||||
|
||||
Args:
|
||||
model_name: Name of the embedding model
|
||||
dimensions: Vector dimensions of the embeddings
|
||||
max_batch_size: Maximum batch size for embedding requests
|
||||
max_retries: Maximum number of retry attempts on failure
|
||||
raise_exception: Whether to raise exceptions on failure
|
||||
max_input_length: Maximum input text length
|
||||
max_cache_size: Maximum number of embeddings to cache in memory (LRU)
|
||||
**kwargs: Additional model-specific parameters
|
||||
"""
|
||||
self.model_name = model_name
|
||||
self.dimensions = dimensions
|
||||
self.max_batch_size = max_batch_size
|
||||
self.max_retries = max_retries
|
||||
self.raise_exception = raise_exception
|
||||
self.max_input_length = max_input_length
|
||||
self.max_cache_size = max_cache_size
|
||||
self.kwargs = kwargs
|
||||
|
||||
# Initialize LRU cache for embeddings
|
||||
self._embedding_cache: OrderedDict[str, list[float]] = OrderedDict()
|
||||
self._cache_hits = 0
|
||||
self._cache_misses = 0
|
||||
|
||||
def _truncate_text(self, text: str) -> str:
|
||||
"""Truncate text to max_input_length if it exceeds the limit."""
|
||||
if len(text) > self.max_input_length:
|
||||
|
|
@ -52,6 +72,76 @@ class BaseEmbeddingModel(ABC):
|
|||
"""Truncate a list of texts to max_input_length."""
|
||||
return [self._truncate_text(text) for text in texts]
|
||||
|
||||
def _get_cache_key(self, text: str) -> str:
|
||||
"""Generate a cache key by hashing the input text.
|
||||
|
||||
Args:
|
||||
text: Input text to hash
|
||||
|
||||
Returns:
|
||||
SHA256 hash of the text as hexadecimal string
|
||||
"""
|
||||
return hashlib.sha256(text.encode("utf-8")).hexdigest()
|
||||
|
||||
def _get_from_cache(self, text: str) -> list[float] | None:
|
||||
"""Retrieve embedding from cache if it exists.
|
||||
|
||||
Args:
|
||||
text: Input text to look up
|
||||
|
||||
Returns:
|
||||
Cached embedding vector or None if not found
|
||||
"""
|
||||
cache_key = self._get_cache_key(text)
|
||||
if cache_key in self._embedding_cache:
|
||||
# Move to end (most recently used)
|
||||
self._embedding_cache.move_to_end(cache_key)
|
||||
self._cache_hits += 1
|
||||
return self._embedding_cache[cache_key]
|
||||
self._cache_misses += 1
|
||||
return None
|
||||
|
||||
def _put_to_cache(self, text: str, embedding: list[float]) -> None:
|
||||
"""Store embedding in cache with LRU eviction.
|
||||
|
||||
Args:
|
||||
text: Input text used as cache key
|
||||
embedding: Embedding vector to cache
|
||||
"""
|
||||
if self.max_cache_size <= 0:
|
||||
return
|
||||
|
||||
cache_key = self._get_cache_key(text)
|
||||
|
||||
# Remove oldest entry if cache is full
|
||||
if len(self._embedding_cache) >= self.max_cache_size and cache_key not in self._embedding_cache:
|
||||
self._embedding_cache.popitem(last=False)
|
||||
|
||||
self._embedding_cache[cache_key] = embedding
|
||||
self._embedding_cache.move_to_end(cache_key)
|
||||
|
||||
def get_cache_stats(self) -> dict[str, int]:
|
||||
"""Get cache statistics.
|
||||
|
||||
Returns:
|
||||
Dictionary with cache size, hits, misses, and hit rate
|
||||
"""
|
||||
total_requests = self._cache_hits + self._cache_misses
|
||||
hit_rate = self._cache_hits / total_requests if total_requests > 0 else 0.0
|
||||
return {
|
||||
"cache_size": len(self._embedding_cache),
|
||||
"max_cache_size": self.max_cache_size,
|
||||
"cache_hits": self._cache_hits,
|
||||
"cache_misses": self._cache_misses,
|
||||
"hit_rate": hit_rate,
|
||||
}
|
||||
|
||||
def clear_cache(self) -> None:
|
||||
"""Clear the embedding cache and reset statistics."""
|
||||
self._embedding_cache.clear()
|
||||
self._cache_hits = 0
|
||||
self._cache_misses = 0
|
||||
|
||||
async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]:
|
||||
"""Internal async implementation for calling the embedding API with batch input."""
|
||||
|
||||
|
|
@ -61,10 +151,20 @@ class BaseEmbeddingModel(ABC):
|
|||
async def get_embedding(self, input_text: str, **kwargs) -> list[float]:
|
||||
"""Async get embedding for a single text with exponential backoff retries."""
|
||||
truncated_text = self._truncate_text(input_text)
|
||||
|
||||
# Check cache first
|
||||
cached_embedding = self._get_from_cache(truncated_text)
|
||||
if cached_embedding is not None:
|
||||
return cached_embedding
|
||||
|
||||
# Cache miss - compute embedding
|
||||
for i in range(self.max_retries):
|
||||
try:
|
||||
result = await self._get_embeddings([truncated_text], **kwargs)
|
||||
return result[0]
|
||||
embedding = result[0]
|
||||
# Store in cache
|
||||
self._put_to_cache(truncated_text, embedding)
|
||||
return embedding
|
||||
except Exception as e:
|
||||
logger.error(f"Model {self.model_name} failed: {e}")
|
||||
if i == self.max_retries - 1:
|
||||
|
|
@ -79,16 +179,36 @@ class BaseEmbeddingModel(ABC):
|
|||
# Truncate all input texts first
|
||||
truncated_texts = self._truncate_texts(input_text)
|
||||
|
||||
# Split into batches and process sequentially to respect rate limits
|
||||
results = []
|
||||
for i in range(0, len(truncated_texts), self.max_batch_size):
|
||||
batch = truncated_texts[i : i + self.max_batch_size]
|
||||
# Check cache for each text and separate cached vs uncached
|
||||
results: list[list[float] | None] = [None] * len(truncated_texts)
|
||||
texts_to_compute: list[tuple[int, str]] = [] # (original_index, text)
|
||||
|
||||
for idx, text in enumerate(truncated_texts):
|
||||
cached = self._get_from_cache(text)
|
||||
if cached is not None:
|
||||
results[idx] = cached
|
||||
else:
|
||||
texts_to_compute.append((idx, text))
|
||||
|
||||
# If all texts were cached, return early
|
||||
if not texts_to_compute:
|
||||
return [r for r in results if r is not None]
|
||||
|
||||
# Compute embeddings for uncached texts in batches
|
||||
uncached_texts = [text for _, text in texts_to_compute]
|
||||
for i in range(0, len(uncached_texts), self.max_batch_size):
|
||||
batch_texts = uncached_texts[i : i + self.max_batch_size]
|
||||
batch_indices = [idx for idx, _ in texts_to_compute[i : i + self.max_batch_size]]
|
||||
|
||||
# Process each batch with retry logic
|
||||
for retry in range(self.max_retries):
|
||||
try:
|
||||
batch_res = await self._get_embeddings(batch, **kwargs)
|
||||
if batch_res:
|
||||
results.extend(batch_res)
|
||||
batch_embeddings = await self._get_embeddings(batch_texts, **kwargs)
|
||||
if batch_embeddings:
|
||||
# Store results and cache them
|
||||
for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings):
|
||||
results[orig_idx] = embedding
|
||||
self._put_to_cache(text, embedding)
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Model {self.model_name} batch failed: {e}")
|
||||
|
|
@ -97,15 +217,26 @@ class BaseEmbeddingModel(ABC):
|
|||
raise
|
||||
else:
|
||||
await asyncio.sleep(retry + 1)
|
||||
return results
|
||||
|
||||
return [r for r in results if r is not None]
|
||||
|
||||
def get_embedding_sync(self, input_text: str, **kwargs) -> list[float]:
|
||||
"""Synchronous get embedding for a single text with retry logic."""
|
||||
truncated_text = self._truncate_text(input_text)
|
||||
|
||||
# Check cache first
|
||||
cached_embedding = self._get_from_cache(truncated_text)
|
||||
if cached_embedding is not None:
|
||||
return cached_embedding
|
||||
|
||||
# Cache miss - compute embedding
|
||||
for i in range(self.max_retries):
|
||||
try:
|
||||
result = self._get_embeddings_sync([truncated_text], **kwargs)
|
||||
return result[0]
|
||||
embedding = result[0]
|
||||
# Store in cache
|
||||
self._put_to_cache(truncated_text, embedding)
|
||||
return embedding
|
||||
except Exception as exc:
|
||||
logger.error(f"Model {self.model_name} failed: {exc}")
|
||||
if i == self.max_retries - 1:
|
||||
|
|
@ -120,15 +251,36 @@ class BaseEmbeddingModel(ABC):
|
|||
# Truncate all input texts first
|
||||
truncated_texts = self._truncate_texts(input_text)
|
||||
|
||||
results = []
|
||||
for i in range(0, len(truncated_texts), self.max_batch_size):
|
||||
batch = truncated_texts[i : i + self.max_batch_size]
|
||||
# Check cache for each text and separate cached vs uncached
|
||||
results: list[list[float] | None] = [None] * len(truncated_texts)
|
||||
texts_to_compute: list[tuple[int, str]] = [] # (original_index, text)
|
||||
|
||||
for idx, text in enumerate(truncated_texts):
|
||||
cached = self._get_from_cache(text)
|
||||
if cached is not None:
|
||||
results[idx] = cached
|
||||
else:
|
||||
texts_to_compute.append((idx, text))
|
||||
|
||||
# If all texts were cached, return early
|
||||
if not texts_to_compute:
|
||||
return [r for r in results if r is not None]
|
||||
|
||||
# Compute embeddings for uncached texts in batches
|
||||
uncached_texts = [text for _, text in texts_to_compute]
|
||||
for i in range(0, len(uncached_texts), self.max_batch_size):
|
||||
batch_texts = uncached_texts[i : i + self.max_batch_size]
|
||||
batch_indices = [idx for idx, _ in texts_to_compute[i : i + self.max_batch_size]]
|
||||
|
||||
# Process each batch with retry logic
|
||||
for retry in range(self.max_retries):
|
||||
try:
|
||||
batch_res = self._get_embeddings_sync(batch, **kwargs)
|
||||
if batch_res:
|
||||
results.extend(batch_res)
|
||||
batch_embeddings = self._get_embeddings_sync(batch_texts, **kwargs)
|
||||
if batch_embeddings:
|
||||
# Store results and cache them
|
||||
for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings):
|
||||
results[orig_idx] = embedding
|
||||
self._put_to_cache(text, embedding)
|
||||
break
|
||||
except Exception as exc:
|
||||
logger.error(f"Model {self.model_name} batch failed: {exc}")
|
||||
|
|
@ -137,7 +289,8 @@ class BaseEmbeddingModel(ABC):
|
|||
raise
|
||||
else:
|
||||
time.sleep(retry + 1)
|
||||
return results
|
||||
|
||||
return [r for r in results if r is not None]
|
||||
|
||||
async def get_node_embedding(self, node: VectorNode, **kwargs) -> VectorNode:
|
||||
"""Async generate and populate vector field for a single VectorNode object."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -1,890 +0,0 @@
|
|||
"""Memory Index Manager - Main coordination layer.
|
||||
|
||||
This module provides the main MemoryIndexManager class that coordinates
|
||||
file watching, embedding generation, and search operations across memory files
|
||||
and session transcripts.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from typing import Any, Callable
|
||||
|
||||
from loguru import logger
|
||||
from pydantic import BaseModel, Field
|
||||
from watchfiles import awatch
|
||||
|
||||
from .ingestion.chunking import chunk_markdown
|
||||
from .memory_storage.sqlite_memory_store import SqliteMemoryStore
|
||||
from .utils.hashing import hash_text
|
||||
from ..enumeration import MemorySource
|
||||
from ..schema import FileMetadata, MemorySearchResult
|
||||
|
||||
# Constants
|
||||
SNIPPET_MAX_CHARS = 700
|
||||
SESSION_DIRTY_DEBOUNCE_MS = 5000
|
||||
EMBEDDING_BATCH_MAX_TOKENS = 8000
|
||||
EMBEDDING_APPROX_CHARS_PER_TOKEN = 1
|
||||
EMBEDDING_INDEX_CONCURRENCY = 4
|
||||
EMBEDDING_RETRY_MAX_ATTEMPTS = 3
|
||||
EMBEDDING_RETRY_BASE_DELAY_MS = 500
|
||||
EMBEDDING_RETRY_MAX_DELAY_MS = 8000
|
||||
BATCH_FAILURE_LIMIT = 2
|
||||
SESSION_DELTA_READ_CHUNK_BYTES = 64 * 1024
|
||||
EMBEDDING_QUERY_TIMEOUT_REMOTE_MS = 60_000
|
||||
EMBEDDING_QUERY_TIMEOUT_LOCAL_MS = 5 * 60_000
|
||||
EMBEDDING_BATCH_TIMEOUT_REMOTE_MS = 2 * 60_000
|
||||
EMBEDDING_BATCH_TIMEOUT_LOCAL_MS = 10 * 60_000
|
||||
|
||||
|
||||
class MemorySyncProgressUpdate(BaseModel):
|
||||
"""Progress update for memory sync operations."""
|
||||
|
||||
completed: int = Field(default=..., description="Number of items completed")
|
||||
total: int = Field(default=..., description="Total number of items to process")
|
||||
label: str | None = Field(default=None, description="Optional label for the progress operation")
|
||||
|
||||
|
||||
class MemorySyncProgressState(BaseModel):
|
||||
"""Internal state for tracking sync progress."""
|
||||
|
||||
completed: int = Field(default=0, description="Number of items completed")
|
||||
total: int = Field(default=0, description="Total number of items to process")
|
||||
label: str | None = Field(default=None, description="Optional label for the progress operation")
|
||||
report: Callable[[MemorySyncProgressUpdate], None] | None = Field(
|
||||
default=None,
|
||||
description="Callback function to report progress updates",
|
||||
)
|
||||
|
||||
|
||||
class SessionDelta(BaseModel):
|
||||
"""Tracks incremental changes in session files."""
|
||||
|
||||
last_size: int = Field(default=0, description="Last known size of the session file")
|
||||
pending_bytes: int = Field(default=0, description="Number of pending bytes to process")
|
||||
pending_messages: int = Field(default=0, description="Number of pending messages to process")
|
||||
|
||||
|
||||
class MemorySearchConfig(BaseModel):
|
||||
"""Configuration for memory search operations."""
|
||||
|
||||
model: str = Field(default="default", description="Model name for embeddings")
|
||||
sources: list[MemorySource] = Field(default_factory=lambda: [MemorySource.MEMORY], description="Sources to search")
|
||||
extra_paths: list[str] = Field(default_factory=list, description="Additional paths to include in search")
|
||||
store_path: str = Field(default="memory.db", description="Path to SQLite database file")
|
||||
vector_enabled: bool = Field(default=True, description="Whether to enable vector search")
|
||||
vector_extension_path: str | None = Field(default=None, description="Path to vector extension for SQLite")
|
||||
fts_enabled: bool = Field(default=True, description="Whether to enable full-text search")
|
||||
chunk_tokens: int = Field(default=300, description="Number of tokens per chunk")
|
||||
chunk_overlap: int = Field(default=30, description="Number of overlapping tokens between chunks")
|
||||
watch_enabled: bool = Field(default=True, description="Whether to enable file watching")
|
||||
watch_debounce_ms: int = Field(default=1000, description="Debounce time for file watcher in milliseconds")
|
||||
interval_minutes: int = Field(default=0, description="Interval between automatic syncs in minutes (0 to disable)")
|
||||
sync_on_search: bool = Field(default=True, description="Whether to sync before search operations")
|
||||
sync_on_session_start: bool = Field(default=True, description="Whether to sync when a session starts")
|
||||
query_min_score: float = Field(default=0.3, description="Minimum relevance score for search results")
|
||||
query_max_results: int = Field(default=10, description="Maximum number of search results to return")
|
||||
hybrid_enabled: bool = Field(default=True, description="Whether to use hybrid vector + keyword search")
|
||||
hybrid_vector_weight: float = Field(default=0.7, description="Weight for vector search in hybrid scoring")
|
||||
hybrid_text_weight: float = Field(default=0.3, description="Weight for text search in hybrid scoring")
|
||||
hybrid_candidate_multiplier: float = Field(
|
||||
default=2.0,
|
||||
description="Multiplier for number of candidates to consider in hybrid search",
|
||||
)
|
||||
session_delta_bytes: int = Field(
|
||||
default=0,
|
||||
description="Threshold for session sync based on bytes changed (0 for any change)",
|
||||
)
|
||||
session_delta_messages: int = Field(
|
||||
default=5,
|
||||
description="Threshold for session sync based on messages changed (0 for any change)",
|
||||
)
|
||||
|
||||
|
||||
# Global cache for manager instances
|
||||
INDEX_CACHE: dict[str, "MemoryIndexManager"] = {}
|
||||
|
||||
|
||||
class MemoryIndexManager:
|
||||
"""Main memory index manager coordinating all memory operations."""
|
||||
|
||||
# ============================================================================
|
||||
# Initialization and Lifecycle
|
||||
# ============================================================================
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
agent_id: str,
|
||||
workspace_dir: str,
|
||||
settings: MemorySearchConfig,
|
||||
store: SqliteMemoryStore,
|
||||
):
|
||||
"""Initialize the memory index manager."""
|
||||
|
||||
self.agent_id = agent_id
|
||||
self.workspace_dir = workspace_dir
|
||||
self.settings = settings
|
||||
self.store = store
|
||||
|
||||
# State tracking
|
||||
self.sources = set(settings.sources)
|
||||
self.closed = False
|
||||
self.dirty = MemorySource.MEMORY in self.sources
|
||||
self.sessions_dirty = False
|
||||
self.sessions_dirty_files: set[str] = set()
|
||||
self.session_pending_files: set[str] = set()
|
||||
self.session_deltas: dict[str, SessionDelta] = {}
|
||||
self.session_warm: set[str] = set()
|
||||
|
||||
# Sync control
|
||||
self.syncing: asyncio.Task | None = None
|
||||
self.watch_task: asyncio.Task | None = None
|
||||
self.session_watch_task: asyncio.Task | None = None
|
||||
self.interval_task: asyncio.Task | None = None
|
||||
|
||||
# Batch failure tracking
|
||||
self.batch_failure_count = 0
|
||||
self.batch_failure_last_error: str | None = None
|
||||
self.batch_failure_lock = asyncio.Lock()
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close the manager and release resources."""
|
||||
if self.closed:
|
||||
return
|
||||
|
||||
self.closed = True
|
||||
|
||||
# Cancel all background tasks
|
||||
if self.watch_task:
|
||||
self.watch_task.cancel()
|
||||
if self.session_watch_task:
|
||||
self.session_watch_task.cancel()
|
||||
if self.interval_task:
|
||||
self.interval_task.cancel()
|
||||
|
||||
await self.store.close()
|
||||
|
||||
# ============================================================================
|
||||
# Public API Methods
|
||||
# ============================================================================
|
||||
|
||||
async def warm_session(self, session_key: str | None = None):
|
||||
"""Pre-sync memory before a session starts."""
|
||||
if not self.settings.sync_on_session_start:
|
||||
return
|
||||
|
||||
key = (session_key or "").strip()
|
||||
if key and key in self.session_warm:
|
||||
return
|
||||
|
||||
await self.sync(reason="session-start")
|
||||
|
||||
if key:
|
||||
self.session_warm.add(key)
|
||||
|
||||
async def sync(
|
||||
self,
|
||||
reason: str | None = None,
|
||||
force: bool = False,
|
||||
progress: Callable[[MemorySyncProgressUpdate], None] | None = None,
|
||||
):
|
||||
"""Synchronize memory index with file system."""
|
||||
if self.syncing:
|
||||
await self.syncing
|
||||
return
|
||||
|
||||
self.syncing = asyncio.create_task(self._run_sync(reason, force, progress))
|
||||
try:
|
||||
await self.syncing
|
||||
finally:
|
||||
self.syncing = None
|
||||
|
||||
async def search(
|
||||
self,
|
||||
query: str,
|
||||
max_results: int | None = None,
|
||||
min_score: float | None = None,
|
||||
session_key: str | None = None,
|
||||
) -> list[MemorySearchResult]:
|
||||
"""Search indexed memory with hybrid vector + keyword search.
|
||||
|
||||
Args:
|
||||
query: Search query text
|
||||
max_results: Maximum number of results to return
|
||||
min_score: Minimum relevance score threshold
|
||||
session_key: Optional session key for warmup
|
||||
|
||||
Returns:
|
||||
List of search results sorted by relevance
|
||||
"""
|
||||
await self.warm_session(session_key)
|
||||
|
||||
if self.settings.sync_on_search and (self.dirty or self.sessions_dirty):
|
||||
try:
|
||||
await self.sync(reason="search")
|
||||
except Exception as err:
|
||||
logger.warning(f"memory sync failed (search): {err}")
|
||||
|
||||
cleaned = query.strip()
|
||||
if not cleaned:
|
||||
return []
|
||||
|
||||
min_score = min_score if min_score is not None else self.settings.query_min_score
|
||||
max_results = max_results if max_results is not None else self.settings.query_max_results
|
||||
|
||||
hybrid = self.settings.hybrid_enabled
|
||||
candidates = min(200, max(1, int(max_results * self.settings.hybrid_candidate_multiplier)))
|
||||
|
||||
# Run keyword search if hybrid enabled
|
||||
keyword_results = []
|
||||
if hybrid:
|
||||
keyword_results = await self._search_keyword(cleaned, candidates)
|
||||
|
||||
# Perform vector search
|
||||
vector_results = await self._search_vector(cleaned, candidates)
|
||||
|
||||
if not hybrid:
|
||||
return [r for r in vector_results if r.score >= min_score][:max_results]
|
||||
|
||||
merged = self._merge_hybrid_results(
|
||||
vector=vector_results,
|
||||
keyword=keyword_results,
|
||||
vector_weight=self.settings.hybrid_vector_weight,
|
||||
text_weight=self.settings.hybrid_text_weight,
|
||||
)
|
||||
|
||||
return [r for r in merged if r.score >= min_score][:max_results]
|
||||
|
||||
async def read_file(
|
||||
self,
|
||||
rel_path: str,
|
||||
from_line: int | None = None,
|
||||
num_lines: int | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""Read a memory file with optional line range.
|
||||
|
||||
Args:
|
||||
rel_path: Relative path to file
|
||||
from_line: Starting line number (1-indexed)
|
||||
num_lines: Number of lines to read
|
||||
|
||||
Returns:
|
||||
Dictionary with 'text' and 'path' keys
|
||||
|
||||
Raises:
|
||||
ValueError: If path is invalid or not allowed
|
||||
"""
|
||||
raw_path = rel_path.strip()
|
||||
assert raw_path, "path required"
|
||||
abs_path = os.path.abspath(os.path.join(self.workspace_dir, raw_path))
|
||||
rel_path_clean = os.path.relpath(abs_path, self.workspace_dir)
|
||||
|
||||
in_workspace = not rel_path_clean.startswith("..") and not os.path.isabs(rel_path_clean)
|
||||
allowed = in_workspace and self._is_memory_path(rel_path_clean)
|
||||
|
||||
if not allowed and self.settings.extra_paths:
|
||||
for extra in self.settings.extra_paths:
|
||||
extra_abs = os.path.abspath(extra)
|
||||
if abs_path.startswith(extra_abs):
|
||||
allowed = True
|
||||
break
|
||||
|
||||
if not allowed:
|
||||
raise ValueError("path required")
|
||||
|
||||
if not abs_path.endswith(".md"):
|
||||
raise ValueError("path required")
|
||||
|
||||
# Read file
|
||||
with open(abs_path, "r", encoding="utf-8") as f:
|
||||
content = f.read()
|
||||
|
||||
if from_line is None and num_lines is None:
|
||||
return {"text": content, "path": rel_path_clean}
|
||||
|
||||
lines = content.split("\n")
|
||||
start = max(1, from_line or 1)
|
||||
count = max(1, num_lines or len(lines))
|
||||
slice_lines = lines[start - 1 : start - 1 + count]
|
||||
|
||||
return {"text": "\n".join(slice_lines), "path": rel_path_clean}
|
||||
|
||||
# ============================================================================
|
||||
# Sync Logic
|
||||
# ============================================================================
|
||||
|
||||
async def _run_sync(
|
||||
self,
|
||||
reason: str | None,
|
||||
force: bool,
|
||||
progress_callback: Callable[[MemorySyncProgressUpdate], None] | None,
|
||||
):
|
||||
"""Execute sync operation."""
|
||||
progress = MemorySyncProgressState()
|
||||
if progress_callback:
|
||||
progress.report = progress_callback
|
||||
|
||||
should_sync_memory = MemorySource.MEMORY in self.sources and (force or self.dirty)
|
||||
should_sync_sessions = self._should_sync_sessions(reason, force)
|
||||
|
||||
if should_sync_memory:
|
||||
await self._sync_memory_files(progress)
|
||||
self.dirty = False
|
||||
|
||||
if should_sync_sessions:
|
||||
await self._sync_session_files(progress)
|
||||
self.sessions_dirty = False
|
||||
self.sessions_dirty_files.clear()
|
||||
elif len(self.sessions_dirty_files) > 0:
|
||||
self.sessions_dirty = True
|
||||
else:
|
||||
self.sessions_dirty = False
|
||||
|
||||
def _should_sync_sessions(self, reason: str | None, force: bool) -> bool:
|
||||
"""Check if session sync is needed."""
|
||||
if MemorySource.SESSIONS not in self.sources:
|
||||
return False
|
||||
|
||||
if force:
|
||||
return True
|
||||
|
||||
if reason in ("session-start", "watch"):
|
||||
return False
|
||||
|
||||
return self.sessions_dirty and len(self.sessions_dirty_files) > 0
|
||||
|
||||
async def _sync_memory_files(self, progress: MemorySyncProgressState):
|
||||
"""Sync memory markdown files."""
|
||||
files = self._list_memory_files()
|
||||
logger.debug("memory sync: indexing memory files", files=len(files))
|
||||
|
||||
active_paths = {f.path for f in files}
|
||||
if progress.report:
|
||||
progress.total += len(files)
|
||||
progress.report(
|
||||
MemorySyncProgressUpdate(
|
||||
completed=progress.completed,
|
||||
total=progress.total,
|
||||
label="Indexing memory files…",
|
||||
),
|
||||
)
|
||||
|
||||
tasks = []
|
||||
for file_entry in files:
|
||||
task = self._index_memory_file(file_entry, progress)
|
||||
tasks.append(task)
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
indexed = await self.store.list_files(MemorySource.MEMORY)
|
||||
for stale_path in indexed:
|
||||
if stale_path not in active_paths:
|
||||
await self.store.delete_file(stale_path, MemorySource.MEMORY)
|
||||
|
||||
async def _sync_session_files(self, progress: MemorySyncProgressState):
|
||||
"""Sync session transcript files."""
|
||||
files = self._list_session_files()
|
||||
logger.debug(
|
||||
"memory sync: indexing session files",
|
||||
files=len(files),
|
||||
index_all=len(self.sessions_dirty_files) == 0,
|
||||
dirty_files=len(self.sessions_dirty_files),
|
||||
)
|
||||
|
||||
if progress.report:
|
||||
progress.total += len(files)
|
||||
progress.report(
|
||||
MemorySyncProgressUpdate(
|
||||
completed=progress.completed,
|
||||
total=progress.total,
|
||||
label="Indexing session files...",
|
||||
),
|
||||
)
|
||||
|
||||
active_paths = set()
|
||||
tasks = []
|
||||
|
||||
for abs_path in files:
|
||||
rel_path = self._session_path_for_file(abs_path)
|
||||
active_paths.add(rel_path)
|
||||
|
||||
if len(self.sessions_dirty_files) == 0 or abs_path in self.sessions_dirty_files:
|
||||
task = self._index_session_file(abs_path, progress)
|
||||
tasks.append(task)
|
||||
else:
|
||||
if progress.report:
|
||||
progress.completed += 1
|
||||
progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
indexed = await self.store.list_files(MemorySource.SESSIONS)
|
||||
for stale_path in indexed:
|
||||
if stale_path not in active_paths:
|
||||
await self.store.delete_file(stale_path, MemorySource.SESSIONS)
|
||||
|
||||
# ============================================================================
|
||||
# File Indexing
|
||||
# ============================================================================
|
||||
|
||||
async def _index_memory_file(self, file_meta: FileMetadata, progress: MemorySyncProgressState):
|
||||
"""Index a single memory file."""
|
||||
existing_meta = await self.store.get_file_metadata(file_meta.path, MemorySource.MEMORY)
|
||||
if existing_meta and existing_meta.hash == file_meta.hash:
|
||||
if progress.report:
|
||||
progress.completed += 1
|
||||
progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
|
||||
return
|
||||
|
||||
# Read and chunk file
|
||||
with open(file_meta.abs_path, "r", encoding="utf-8") as f:
|
||||
content = f.read()
|
||||
|
||||
chunks = chunk_markdown(
|
||||
content,
|
||||
file_meta.path,
|
||||
MemorySource.MEMORY,
|
||||
self.settings.chunk_tokens,
|
||||
self.settings.chunk_overlap,
|
||||
)
|
||||
|
||||
chunks = [c for c in chunks if c.text.strip()]
|
||||
|
||||
if chunks:
|
||||
chunks = await self.store.get_chunk_embeddings(chunks)
|
||||
|
||||
await self.store.upsert_file(file_meta, MemorySource.MEMORY, chunks)
|
||||
if progress.report:
|
||||
progress.completed += 1
|
||||
progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
|
||||
|
||||
async def _index_session_file(self, abs_path: str, progress: MemorySyncProgressState):
|
||||
"""Index a single session transcript file."""
|
||||
file_meta = self._build_session_file_meta(abs_path)
|
||||
if not file_meta:
|
||||
if progress.report:
|
||||
progress.completed += 1
|
||||
progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
|
||||
return
|
||||
|
||||
existing_meta = await self.store.get_file_metadata(file_meta.path, MemorySource.SESSIONS)
|
||||
if existing_meta and existing_meta.hash == file_meta.hash:
|
||||
self._reset_session_delta(abs_path, file_meta.size)
|
||||
if progress.report:
|
||||
progress.completed += 1
|
||||
progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
|
||||
return
|
||||
|
||||
chunks = chunk_markdown(
|
||||
file_meta.content,
|
||||
file_meta.path,
|
||||
MemorySource.SESSIONS,
|
||||
self.settings.chunk_tokens,
|
||||
self.settings.chunk_overlap,
|
||||
)
|
||||
|
||||
chunks = [c for c in chunks if c.text.strip()]
|
||||
|
||||
if chunks:
|
||||
chunks = await self.store.get_chunk_embeddings(chunks)
|
||||
|
||||
await self.store.upsert_file(file_meta, MemorySource.SESSIONS, chunks)
|
||||
self._reset_session_delta(abs_path, file_meta.size)
|
||||
|
||||
if progress.report:
|
||||
progress.completed += 1
|
||||
progress.report(MemorySyncProgressUpdate(completed=progress.completed, total=progress.total))
|
||||
|
||||
# ============================================================================
|
||||
# File Listing and Building
|
||||
# ============================================================================
|
||||
|
||||
def _list_memory_files(self) -> list[FileMetadata]:
|
||||
"""List all memory markdown files."""
|
||||
files = []
|
||||
|
||||
# Scan workspace
|
||||
memory_paths = [
|
||||
os.path.join(self.workspace_dir, "MEMORY.md"),
|
||||
os.path.join(self.workspace_dir, "memory.md"),
|
||||
os.path.join(self.workspace_dir, "memory"),
|
||||
]
|
||||
|
||||
for base_path in memory_paths:
|
||||
if os.path.isfile(base_path) and base_path.endswith(".md"):
|
||||
files.append(self._build_file_entry(base_path))
|
||||
elif os.path.isdir(base_path):
|
||||
for root, _, filenames in os.walk(base_path):
|
||||
for filename in filenames:
|
||||
if filename.endswith(".md"):
|
||||
abs_path = os.path.join(root, filename)
|
||||
files.append(self._build_file_entry(abs_path))
|
||||
|
||||
# Extra paths
|
||||
for extra in self.settings.extra_paths:
|
||||
if os.path.isfile(extra) and extra.endswith(".md"):
|
||||
files.append(self._build_file_entry(extra))
|
||||
elif os.path.isdir(extra):
|
||||
for root, _, filenames in os.walk(extra):
|
||||
for filename in filenames:
|
||||
if filename.endswith(".md"):
|
||||
abs_path = os.path.join(root, filename)
|
||||
files.append(self._build_file_entry(abs_path))
|
||||
|
||||
return files
|
||||
|
||||
def _list_session_files(self) -> list[str]:
|
||||
"""List all session transcript files."""
|
||||
sessions_dir = os.path.join(self.workspace_dir, "sessions", self.agent_id)
|
||||
if not os.path.exists(sessions_dir):
|
||||
return []
|
||||
|
||||
files = []
|
||||
for filename in os.listdir(sessions_dir):
|
||||
if filename.endswith(".jsonl"):
|
||||
files.append(os.path.join(sessions_dir, filename))
|
||||
|
||||
return files
|
||||
|
||||
def _build_file_entry(self, abs_path: str) -> FileMetadata:
|
||||
"""Build file entry metadata."""
|
||||
stat = os.stat(abs_path)
|
||||
with open(abs_path, "r", encoding="utf-8") as f:
|
||||
content = f.read()
|
||||
|
||||
rel_path = os.path.relpath(abs_path, self.workspace_dir)
|
||||
|
||||
return FileMetadata(
|
||||
hash=hash_text(content),
|
||||
mtime_ms=stat.st_mtime * 1000,
|
||||
size=stat.st_size,
|
||||
path=rel_path.replace("\\", "/"),
|
||||
abs_path=abs_path,
|
||||
)
|
||||
|
||||
def _build_session_file_meta(self, abs_path: str) -> FileMetadata | None:
|
||||
"""Build session file entry with parsed content. TODO 修改message解析逻辑"""
|
||||
stat = os.stat(abs_path)
|
||||
with open(abs_path, "r", encoding="utf-8") as f:
|
||||
raw = f.read()
|
||||
|
||||
lines = raw.split("\n")
|
||||
collected = []
|
||||
|
||||
for line in lines:
|
||||
if not line.strip():
|
||||
continue
|
||||
|
||||
try:
|
||||
record = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
|
||||
if record.get("type") != "message":
|
||||
continue
|
||||
|
||||
message = record.get("message", {})
|
||||
role = message.get("role")
|
||||
|
||||
if role not in ("user", "assistant"):
|
||||
continue
|
||||
|
||||
text = self._extract_session_text(message.get("content"))
|
||||
if not text:
|
||||
continue
|
||||
|
||||
label = "User" if role == "user" else "Assistant"
|
||||
collected.append(f"{label}: {text}")
|
||||
|
||||
content = "\n".join(collected)
|
||||
rel_path = self._session_path_for_file(abs_path)
|
||||
|
||||
return FileMetadata(
|
||||
hash=hash_text(content),
|
||||
mtime_ms=stat.st_mtime * 1000,
|
||||
size=stat.st_size,
|
||||
path=rel_path,
|
||||
abs_path=abs_path,
|
||||
content=content,
|
||||
)
|
||||
|
||||
# ============================================================================
|
||||
# Session Processing Helpers
|
||||
# ============================================================================
|
||||
|
||||
@staticmethod
|
||||
def _session_path_for_file(abs_path: str) -> str:
|
||||
"""Convert absolute session path to relative."""
|
||||
return f"sessions/{os.path.basename(abs_path)}"
|
||||
|
||||
def _extract_session_text(self, content: Any) -> str | None:
|
||||
"""Extract text from session message content."""
|
||||
if isinstance(content, str):
|
||||
normalized = self._normalize_session_text(content)
|
||||
return normalized if normalized else None
|
||||
|
||||
if not isinstance(content, list):
|
||||
return None
|
||||
|
||||
parts = []
|
||||
for block in content:
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
|
||||
if block.get("type") != "text":
|
||||
continue
|
||||
|
||||
text = block.get("text")
|
||||
if isinstance(text, str):
|
||||
normalized = self._normalize_session_text(text)
|
||||
if normalized:
|
||||
parts.append(normalized)
|
||||
|
||||
return " ".join(parts) if parts else None
|
||||
|
||||
@staticmethod
|
||||
def _normalize_session_text(text: str) -> str:
|
||||
"""Normalize session text by collapsing whitespace."""
|
||||
text = re.sub(r"\s*\n+\s*", " ", text)
|
||||
text = re.sub(r"\s+", " ", text)
|
||||
return text.strip()
|
||||
|
||||
# ============================================================================
|
||||
# Session Delta Tracking
|
||||
# ============================================================================
|
||||
|
||||
async def _process_session_delta_batch(self) -> None:
|
||||
"""Process pending session file changes."""
|
||||
if not self.session_pending_files:
|
||||
return
|
||||
|
||||
pending = list(self.session_pending_files)
|
||||
self.session_pending_files.clear()
|
||||
|
||||
should_sync = False
|
||||
for session_file in pending:
|
||||
delta = await self._update_session_delta(session_file)
|
||||
if not delta:
|
||||
continue
|
||||
|
||||
bytes_threshold = self.settings.session_delta_bytes
|
||||
messages_threshold = self.settings.session_delta_messages
|
||||
|
||||
if bytes_threshold <= 0:
|
||||
bytes_hit = delta["pending_bytes"] > 0
|
||||
else:
|
||||
bytes_hit = delta["pending_bytes"] >= bytes_threshold
|
||||
|
||||
if messages_threshold <= 0:
|
||||
messages_hit = delta["pending_messages"] > 0
|
||||
else:
|
||||
messages_hit = delta["pending_messages"] >= messages_threshold
|
||||
|
||||
if not bytes_hit and not messages_hit:
|
||||
continue
|
||||
|
||||
self.sessions_dirty_files.add(session_file)
|
||||
self.sessions_dirty = True
|
||||
should_sync = True
|
||||
|
||||
if should_sync:
|
||||
try:
|
||||
await self.sync(reason="session-delta")
|
||||
except Exception as err:
|
||||
logger.warning(f"memory sync failed (session-delta): {err}")
|
||||
|
||||
async def _update_session_delta(self, session_file: str) -> dict[str, int] | None:
|
||||
"""Update delta tracking for a session file."""
|
||||
try:
|
||||
stat = os.stat(session_file)
|
||||
size = stat.st_size
|
||||
except OSError:
|
||||
return None
|
||||
|
||||
state = self.session_deltas.get(session_file)
|
||||
if not state:
|
||||
state = SessionDelta()
|
||||
self.session_deltas[session_file] = state
|
||||
|
||||
delta_bytes = max(0, size - state.last_size)
|
||||
|
||||
if delta_bytes == 0 and size == state.last_size:
|
||||
return {
|
||||
"delta_bytes": self.settings.session_delta_bytes,
|
||||
"delta_messages": self.settings.session_delta_messages,
|
||||
"pending_bytes": state.pending_bytes,
|
||||
"pending_messages": state.pending_messages,
|
||||
}
|
||||
|
||||
if size < state.last_size:
|
||||
state.last_size = size
|
||||
state.pending_bytes += size
|
||||
if self.settings.session_delta_messages > 0:
|
||||
state.pending_messages += await self._count_newlines(session_file, 0, size)
|
||||
else:
|
||||
state.pending_bytes += delta_bytes
|
||||
if self.settings.session_delta_messages > 0:
|
||||
state.pending_messages += await self._count_newlines(session_file, state.last_size, size)
|
||||
state.last_size = size
|
||||
|
||||
return {
|
||||
"delta_bytes": self.settings.session_delta_bytes,
|
||||
"delta_messages": self.settings.session_delta_messages,
|
||||
"pending_bytes": state.pending_bytes,
|
||||
"pending_messages": state.pending_messages,
|
||||
}
|
||||
|
||||
def _reset_session_delta(self, abs_path: str, size: int) -> None:
|
||||
"""Reset delta tracking for a session file."""
|
||||
state = self.session_deltas.get(abs_path)
|
||||
if state:
|
||||
state.last_size = size
|
||||
state.pending_bytes = 0
|
||||
state.pending_messages = 0
|
||||
|
||||
@staticmethod
|
||||
async def _count_newlines(abs_path: str, start: int, end: int) -> int:
|
||||
"""Count newlines in a file range."""
|
||||
if end <= start:
|
||||
return 0
|
||||
|
||||
count = 0
|
||||
with open(abs_path, "rb") as f:
|
||||
f.seek(start)
|
||||
remaining = end - start
|
||||
|
||||
while remaining > 0:
|
||||
chunk_size = min(SESSION_DELTA_READ_CHUNK_BYTES, remaining)
|
||||
chunk = f.read(chunk_size)
|
||||
if not chunk:
|
||||
break
|
||||
|
||||
count += chunk.count(b"\n")
|
||||
remaining -= len(chunk)
|
||||
|
||||
return count
|
||||
|
||||
# ============================================================================
|
||||
# File Watchers
|
||||
# ============================================================================
|
||||
|
||||
async def _start_watchers(self):
|
||||
"""Start file watching and interval sync tasks."""
|
||||
if self.settings.watch_enabled and MemorySource.MEMORY in self.sources:
|
||||
self.watch_task = asyncio.create_task(self._watch_memory_files())
|
||||
|
||||
if MemorySource.SESSIONS in self.sources:
|
||||
self.session_watch_task = asyncio.create_task(self._watch_session_files())
|
||||
|
||||
if self.settings.interval_minutes > 0:
|
||||
self.interval_task = asyncio.create_task(self._interval_sync())
|
||||
|
||||
async def _watch_memory_files(self) -> None:
|
||||
"""Watch memory files for changes."""
|
||||
watch_paths = [
|
||||
os.path.join(self.workspace_dir, "MEMORY.md"),
|
||||
os.path.join(self.workspace_dir, "memory.md"),
|
||||
os.path.join(self.workspace_dir, "memory"),
|
||||
]
|
||||
|
||||
for extra in self.settings.extra_paths:
|
||||
watch_paths.append(extra)
|
||||
|
||||
async for changes in awatch(*watch_paths, stop_event=None):
|
||||
if self.closed:
|
||||
break
|
||||
|
||||
for _, path in changes:
|
||||
if path.endswith(".md"):
|
||||
self.dirty = True
|
||||
await asyncio.sleep(self.settings.watch_debounce_ms / 1000)
|
||||
try:
|
||||
await self.sync(reason="watch")
|
||||
except Exception as e:
|
||||
logger.exception(f"memory sync failed (watch): {e}")
|
||||
|
||||
async def _watch_session_files(self):
|
||||
"""Watch session files for changes."""
|
||||
sessions_dir = os.path.join(self.workspace_dir, "sessions", self.agent_id)
|
||||
if not os.path.exists(sessions_dir):
|
||||
return
|
||||
|
||||
async for changes in awatch(sessions_dir, stop_event=None):
|
||||
if self.closed:
|
||||
break
|
||||
|
||||
for _, path in changes:
|
||||
if path.endswith(".jsonl"):
|
||||
self.session_pending_files.add(path)
|
||||
|
||||
await asyncio.sleep(SESSION_DIRTY_DEBOUNCE_MS / 1000)
|
||||
await self._process_session_delta_batch()
|
||||
|
||||
async def _interval_sync(self) -> None:
|
||||
"""Periodically sync the index."""
|
||||
while not self.closed:
|
||||
await asyncio.sleep(self.settings.interval_minutes * 60)
|
||||
if not self.closed:
|
||||
try:
|
||||
await self.sync(reason="interval")
|
||||
except Exception as err:
|
||||
logger.warning(f"memory sync failed (interval): {err}")
|
||||
|
||||
# ============================================================================
|
||||
# Search Methods
|
||||
# ============================================================================
|
||||
|
||||
async def _search_vector(self, query: str, limit: int) -> list[MemorySearchResult]:
|
||||
"""Perform vector similarity search."""
|
||||
return await self.store.vector_search(query, limit, sources=list(self.sources))
|
||||
|
||||
async def _search_keyword(self, query: str, limit: int) -> list[MemorySearchResult]:
|
||||
"""Perform keyword/FTS search."""
|
||||
if not self.settings.fts_enabled:
|
||||
return []
|
||||
|
||||
return await self.store.keyword_search(query, limit, sources=list(self.sources))
|
||||
|
||||
@staticmethod
|
||||
def _merge_hybrid_results(
|
||||
vector: list[MemorySearchResult],
|
||||
keyword: list[MemorySearchResult],
|
||||
vector_weight: float,
|
||||
text_weight: float,
|
||||
) -> list[MemorySearchResult]:
|
||||
"""Merge vector and keyword search results."""
|
||||
merged: dict[str, MemorySearchResult] = {}
|
||||
|
||||
# Process vector results
|
||||
for result in vector:
|
||||
result.score = result.score * vector_weight
|
||||
merged[result.merge_key] = result
|
||||
|
||||
# Process keyword results
|
||||
for result in keyword:
|
||||
key = result.merge_key
|
||||
if key in merged:
|
||||
merged[key].score += result.score * text_weight
|
||||
else:
|
||||
result.score = result.score * text_weight
|
||||
merged[key] = result
|
||||
|
||||
results = list(merged.values())
|
||||
results.sort(key=lambda r: r.score, reverse=True)
|
||||
return results
|
||||
|
||||
# ============================================================================
|
||||
# Utility Methods
|
||||
# ============================================================================
|
||||
|
||||
@staticmethod
|
||||
def _is_memory_path(rel_path: str) -> bool:
|
||||
"""Check if path is a valid memory path."""
|
||||
normalized = rel_path.replace("\\", "/")
|
||||
|
||||
if normalized in ("MEMORY.md", "memory.md"):
|
||||
return True
|
||||
|
||||
if normalized.startswith("memory/") and normalized.endswith(".md"):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
|
@ -1,15 +0,0 @@
|
|||
"""Utility functions for hashing text content."""
|
||||
|
||||
import hashlib
|
||||
|
||||
|
||||
def hash_text(text: str) -> str:
|
||||
"""Generate SHA-256 hash of text content.
|
||||
|
||||
Args:
|
||||
text: Input text to hash
|
||||
|
||||
Returns:
|
||||
Hexadecimal representation of the SHA-256 hash
|
||||
"""
|
||||
return hashlib.sha256(text.encode("utf-8")).hexdigest()
|
||||
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)
|
||||
|
|
@ -2,17 +2,31 @@
|
|||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from ...embedding import BaseEmbeddingModel
|
||||
from ...enumeration import MemorySource
|
||||
from ...schema import FileMetadata, MemoryIndexMeta, MemoryChunk, MemorySearchResult
|
||||
from ..embedding import BaseEmbeddingModel
|
||||
from ..enumeration import MemorySource
|
||||
from ..schema import FileMetadata, MemoryChunk, MemorySearchResult
|
||||
|
||||
|
||||
class BaseMemoryStore(ABC):
|
||||
"""Abstract base class for memory storage backends."""
|
||||
|
||||
def __init__(self, embedding_model: BaseEmbeddingModel):
|
||||
def __init__(
|
||||
self,
|
||||
store_name: str,
|
||||
embedding_model: BaseEmbeddingModel,
|
||||
fts_enabled: bool = True,
|
||||
snippet_max_chars: int = 700,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize"""
|
||||
self.store_name: str = store_name
|
||||
self.embedding_model: BaseEmbeddingModel = embedding_model
|
||||
self.fts_enabled: bool = fts_enabled
|
||||
self.snippet_max_chars: int = snippet_max_chars
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
self.vector_available = False
|
||||
self.fts_available = False
|
||||
|
||||
@property
|
||||
def embedding_dim(self) -> int:
|
||||
|
|
@ -20,77 +34,21 @@ class BaseMemoryStore(ABC):
|
|||
return self.embedding_model.dimensions
|
||||
|
||||
async def get_embedding(self, query: str, **kwargs) -> list[float]:
|
||||
"""Get embedding for a single query string.
|
||||
|
||||
Args:
|
||||
query: Input text to generate embedding for
|
||||
**kwargs: Additional arguments passed to the embedding model
|
||||
|
||||
Returns:
|
||||
Embedding vector as a list of floats
|
||||
"""
|
||||
"""Get embedding for a single query string."""
|
||||
return await self.embedding_model.get_embedding(query, **kwargs)
|
||||
|
||||
async def get_embeddings(self, queries: list[str], **kwargs) -> list[list[float]]:
|
||||
"""Get embeddings for a batch of query strings.
|
||||
|
||||
Args:
|
||||
queries: List of input texts to generate embeddings for
|
||||
**kwargs: Additional arguments passed to the embedding model
|
||||
|
||||
Returns:
|
||||
List of embedding vectors, each as a list of floats
|
||||
"""
|
||||
"""Get embeddings for a batch of query strings."""
|
||||
return await self.embedding_model.get_embeddings(queries, **kwargs)
|
||||
|
||||
async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
|
||||
"""Generate and populate embedding field for a single MemoryChunk object.
|
||||
|
||||
Args:
|
||||
chunk: MemoryChunk object containing text to embed
|
||||
**kwargs: Additional arguments passed to the embedding model
|
||||
|
||||
Returns:
|
||||
The same MemoryChunk object with populated embedding field
|
||||
"""
|
||||
"""Generate and populate embedding field for a single MemoryChunk object."""
|
||||
return await self.embedding_model.get_chunk_embedding(chunk, **kwargs)
|
||||
|
||||
async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
|
||||
"""Generate and populate embedding fields for a batch of MemoryChunk objects.
|
||||
|
||||
Args:
|
||||
chunks: List of MemoryChunk objects containing text to embed
|
||||
**kwargs: Additional arguments passed to the embedding model
|
||||
|
||||
Returns:
|
||||
The same list of MemoryChunk objects with populated embedding fields
|
||||
"""
|
||||
"""Generate and populate embedding fields for a batch of MemoryChunk objects."""
|
||||
return await self.embedding_model.get_chunk_embeddings(chunks, **kwargs)
|
||||
|
||||
def get_chunk_embedding_sync(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
|
||||
"""Synchronously generate and populate embedding field for a single MemoryChunk object.
|
||||
|
||||
Args:
|
||||
chunk: MemoryChunk object containing text to embed
|
||||
**kwargs: Additional arguments passed to the embedding model
|
||||
|
||||
Returns:
|
||||
The same MemoryChunk object with populated embedding field
|
||||
"""
|
||||
return self.embedding_model.get_chunk_embedding_sync(chunk, **kwargs)
|
||||
|
||||
def get_chunk_embeddings_sync(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
|
||||
"""Synchronously generate embeddings for a batch of MemoryChunk objects.
|
||||
|
||||
Args:
|
||||
chunks: List of MemoryChunk objects containing text to embed
|
||||
**kwargs: Additional arguments passed to the embedding model
|
||||
|
||||
Returns:
|
||||
The same list of MemoryChunk objects with populated embedding fields
|
||||
"""
|
||||
return self.embedding_model.get_chunk_embeddings_sync(chunks, **kwargs)
|
||||
|
||||
@abstractmethod
|
||||
async def start(self):
|
||||
"""Initialize the storage backend."""
|
||||
|
|
@ -104,19 +62,23 @@ class BaseMemoryStore(ABC):
|
|||
"""Delete a file and all its chunks."""
|
||||
|
||||
@abstractmethod
|
||||
async def get_file_hash(self, path: str, source: MemorySource) -> str | None:
|
||||
"""Get the hash of an indexed file."""
|
||||
async def delete_file_chunks(self, path: str, chunk_ids: list[str]):
|
||||
"""Delete chunks for a file."""
|
||||
|
||||
@abstractmethod
|
||||
async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None:
|
||||
"""Get full file metadata with statistics."""
|
||||
async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource):
|
||||
"""Insert or update specific chunks without affecting other chunks."""
|
||||
|
||||
@abstractmethod
|
||||
async def list_files(self, source: MemorySource) -> list[str]:
|
||||
"""List all indexed file paths for a source."""
|
||||
|
||||
@abstractmethod
|
||||
async def get_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
|
||||
async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None:
|
||||
"""Get full file metadata with statistics."""
|
||||
|
||||
@abstractmethod
|
||||
async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
|
||||
"""Get all chunks for a file."""
|
||||
|
||||
@abstractmethod
|
||||
|
|
@ -155,14 +117,6 @@ class BaseMemoryStore(ABC):
|
|||
List of search results sorted by relevance
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def read_meta(self, key: str) -> MemoryIndexMeta | None:
|
||||
"""Read metadata value."""
|
||||
|
||||
@abstractmethod
|
||||
async def write_meta(self, key: str, value: MemoryIndexMeta | dict):
|
||||
"""Write metadata value."""
|
||||
|
||||
@abstractmethod
|
||||
async def clear_all(self):
|
||||
"""Clear all indexed data."""
|
||||
|
|
@ -9,9 +9,8 @@ from pathlib import Path
|
|||
from loguru import logger
|
||||
|
||||
from .base_memory_store import BaseMemoryStore
|
||||
from ...embedding import BaseEmbeddingModel
|
||||
from ...enumeration import MemorySource
|
||||
from ...schema import FileMetadata, MemoryIndexMeta, MemoryChunk, MemorySearchResult
|
||||
from ..enumeration import MemorySource
|
||||
from ..schema import FileMetadata, MemoryChunk, MemorySearchResult
|
||||
|
||||
|
||||
class SqliteMemoryStore(BaseMemoryStore):
|
||||
|
|
@ -28,25 +27,32 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
- Efficient chunk and file metadata management
|
||||
"""
|
||||
|
||||
VECTOR_TABLE = "chunks_vec"
|
||||
FTS_TABLE = "chunks_fts"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
db_path: str,
|
||||
embedding_model: BaseEmbeddingModel,
|
||||
vec_ext_path: str = "",
|
||||
fts_enabled: bool = True,
|
||||
snippet_max_chars: int = 700,
|
||||
):
|
||||
super().__init__(embedding_model=embedding_model)
|
||||
def __init__(self, db_path: str = ".reme/memory.db", vec_ext_path: str = "", **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.db_path = db_path
|
||||
self.vec_ext_path = vec_ext_path
|
||||
self.fts_enabled = fts_enabled
|
||||
self.snippet_max_chars = snippet_max_chars
|
||||
|
||||
self.conn: sqlite3.Connection | None = None
|
||||
self.vector_available = False
|
||||
self.fts_available = False
|
||||
|
||||
@property
|
||||
def vector_table_name(self) -> str:
|
||||
"""Get the name of the vector table for this store."""
|
||||
return f"chunks_vec_{self.store_name}"
|
||||
|
||||
@property
|
||||
def fts_table_name(self) -> str:
|
||||
"""Get the name of the FTS table for this store."""
|
||||
return f"chunks_fts_{self.store_name}"
|
||||
|
||||
@property
|
||||
def chunks_table_name(self) -> str:
|
||||
"""Get the name of the chunks table for this store."""
|
||||
return f"chunks_{self.store_name}"
|
||||
|
||||
@property
|
||||
def files_table_name(self) -> str:
|
||||
"""Get the name of the files table for this store."""
|
||||
return f"files_{self.store_name}"
|
||||
|
||||
@staticmethod
|
||||
def vector_to_blob(embedding: list[float]) -> bytes:
|
||||
|
|
@ -55,6 +61,9 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
|
||||
async def start(self) -> None:
|
||||
"""Initialize database and load extensions."""
|
||||
if self.conn is not None:
|
||||
return
|
||||
|
||||
Path(self.db_path).parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
self.conn = sqlite3.connect(self.db_path, check_same_thread=False)
|
||||
|
|
@ -68,16 +77,27 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
logger.info(f"Loaded sqlite-vec: {self.vec_ext_path}")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load sqlite-vec: {e}")
|
||||
|
||||
else:
|
||||
# Try common extension names
|
||||
for name in ["vec0", "sqlite_vec", "vector0"]:
|
||||
try:
|
||||
self.conn.load_extension(name)
|
||||
self.vector_available = True
|
||||
logger.info(f"Loaded sqlite-vec: {name}")
|
||||
break
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
import sqlite_vec
|
||||
|
||||
ext_path = sqlite_vec.loadable_path()
|
||||
self.conn.load_extension(ext_path)
|
||||
self.vector_available = True
|
||||
logger.info(f"Loaded sqlite-vec from package: {ext_path}")
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load sqlite-vec from package: {e}")
|
||||
# Fallback: try common extension names
|
||||
for name in ["vec0", "sqlite_vec", "vector0"]:
|
||||
try:
|
||||
self.conn.load_extension(name)
|
||||
self.vector_available = True
|
||||
logger.info(f"Loaded sqlite-vec: {name}")
|
||||
break
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
self.conn.enable_load_extension(False)
|
||||
await self._create_tables()
|
||||
|
|
@ -86,20 +106,10 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
"""Create database schema."""
|
||||
cursor = self.conn.cursor()
|
||||
|
||||
# Metadata
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS meta (
|
||||
key TEXT PRIMARY KEY,
|
||||
value TEXT
|
||||
)
|
||||
""",
|
||||
)
|
||||
|
||||
# Files
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS files (
|
||||
f"""
|
||||
CREATE TABLE IF NOT EXISTS {self.files_table_name} (
|
||||
path TEXT,
|
||||
source TEXT,
|
||||
hash TEXT,
|
||||
|
|
@ -112,8 +122,8 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
|
||||
# Chunks
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS chunks (
|
||||
f"""
|
||||
CREATE TABLE IF NOT EXISTS {self.chunks_table_name} (
|
||||
id TEXT PRIMARY KEY,
|
||||
path TEXT,
|
||||
source TEXT,
|
||||
|
|
@ -127,49 +137,34 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
""",
|
||||
)
|
||||
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE INDEX IF NOT EXISTS idx_chunks_path_source
|
||||
ON chunks(path, source)
|
||||
""",
|
||||
)
|
||||
|
||||
# Vector table (sqlite-vec)
|
||||
if self.vector_available:
|
||||
try:
|
||||
cursor.execute(
|
||||
f"""
|
||||
CREATE VIRTUAL TABLE IF NOT EXISTS {self.VECTOR_TABLE} USING vec0(
|
||||
id TEXT PRIMARY KEY,
|
||||
embedding FLOAT[{self.embedding_dim}]
|
||||
)
|
||||
""",
|
||||
cursor.execute(
|
||||
f"""
|
||||
CREATE VIRTUAL TABLE IF NOT EXISTS {self.vector_table_name} USING vec0(
|
||||
id TEXT PRIMARY KEY,
|
||||
embedding FLOAT[{self.embedding_dim}]
|
||||
)
|
||||
logger.info(f"Created vector table (dims={self.embedding_dim})")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to create vector table: {e}")
|
||||
self.vector_available = False
|
||||
""",
|
||||
)
|
||||
logger.info(f"Created vector table (dims={self.embedding_dim})")
|
||||
|
||||
# FTS table
|
||||
if self.fts_enabled:
|
||||
try:
|
||||
cursor.execute(
|
||||
f"""
|
||||
CREATE VIRTUAL TABLE IF NOT EXISTS {self.FTS_TABLE} USING fts5(
|
||||
text,
|
||||
id UNINDEXED,
|
||||
path UNINDEXED,
|
||||
source UNINDEXED,
|
||||
start_line UNINDEXED,
|
||||
end_line UNINDEXED
|
||||
)
|
||||
""",
|
||||
cursor.execute(
|
||||
f"""
|
||||
CREATE VIRTUAL TABLE IF NOT EXISTS {self.fts_table_name} USING fts5(
|
||||
text,
|
||||
id UNINDEXED,
|
||||
path UNINDEXED,
|
||||
source UNINDEXED,
|
||||
start_line UNINDEXED,
|
||||
end_line UNINDEXED
|
||||
)
|
||||
self.fts_available = True
|
||||
logger.info("Created FTS5 table")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to create FTS table: {e}")
|
||||
self.fts_available = False
|
||||
""",
|
||||
)
|
||||
self.fts_available = True
|
||||
logger.info("Created FTS5 table")
|
||||
|
||||
self.conn.commit()
|
||||
cursor.close()
|
||||
|
|
@ -180,12 +175,11 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
|
||||
try:
|
||||
cursor.execute("BEGIN")
|
||||
await self._delete_file_internal(cursor, file_meta.path, source)
|
||||
|
||||
# Insert file
|
||||
cursor.execute(
|
||||
"""
|
||||
INSERT OR REPLACE INTO files (path, source, hash, mtime, size)
|
||||
f"""
|
||||
INSERT OR REPLACE INTO {self.files_table_name} (path, source, hash, mtime, size)
|
||||
VALUES (?, ?, ?, ?, ?)
|
||||
""",
|
||||
(file_meta.path, source.value, file_meta.hash, file_meta.mtime_ms, file_meta.size),
|
||||
|
|
@ -195,8 +189,8 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
now = int(time.time() * 1000)
|
||||
for chunk in chunks:
|
||||
cursor.execute(
|
||||
"""
|
||||
INSERT INTO chunks (
|
||||
f"""
|
||||
INSERT OR REPLACE INTO {self.chunks_table_name} (
|
||||
id, path, source, start_line, end_line,
|
||||
hash, text, embedding, updated_at
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
|
|
@ -215,38 +209,33 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
)
|
||||
|
||||
# Insert vector
|
||||
if self.vector_available and chunk.embedding:
|
||||
try:
|
||||
cursor.execute(
|
||||
f"""
|
||||
INSERT INTO {self.VECTOR_TABLE} (id, embedding)
|
||||
VALUES (?, ?)
|
||||
""",
|
||||
(chunk.id, self.vector_to_blob(chunk.embedding)),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Vector insert failed: {e}")
|
||||
if self.vector_available:
|
||||
assert chunk.embedding, "Embedding is required for vector insert"
|
||||
cursor.execute(
|
||||
f"""
|
||||
INSERT OR REPLACE INTO {self.vector_table_name} (id, embedding)
|
||||
VALUES (?, ?)
|
||||
""",
|
||||
(chunk.id, self.vector_to_blob(chunk.embedding)),
|
||||
)
|
||||
|
||||
# Insert FTS
|
||||
if self.fts_available:
|
||||
try:
|
||||
cursor.execute(
|
||||
f"""
|
||||
INSERT INTO {self.FTS_TABLE} (
|
||||
text, id, path, source, start_line, end_line
|
||||
) VALUES (?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
chunk.text,
|
||||
chunk.id,
|
||||
file_meta.path,
|
||||
source.value,
|
||||
chunk.start_line,
|
||||
chunk.end_line,
|
||||
),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"FTS insert failed: {e}")
|
||||
cursor.execute(
|
||||
f"""
|
||||
INSERT OR REPLACE INTO {self.fts_table_name} (
|
||||
text, id, path, source, start_line, end_line
|
||||
) VALUES (?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
chunk.text,
|
||||
chunk.id,
|
||||
file_meta.path,
|
||||
source.value,
|
||||
chunk.start_line,
|
||||
chunk.end_line,
|
||||
),
|
||||
)
|
||||
|
||||
cursor.execute("COMMIT")
|
||||
except Exception:
|
||||
|
|
@ -255,12 +244,50 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
finally:
|
||||
cursor.close()
|
||||
|
||||
async def delete_file(self, path: str, source: MemorySource) -> None:
|
||||
async def delete_file(self, path: str, source: MemorySource):
|
||||
"""Delete file and all its chunks."""
|
||||
cursor = self.conn.cursor()
|
||||
try:
|
||||
cursor.execute("BEGIN")
|
||||
await self._delete_file_internal(cursor, path, source)
|
||||
|
||||
# Get chunk IDs for vector deletion
|
||||
cursor.execute(
|
||||
f"SELECT id FROM {self.chunks_table_name} WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
chunk_ids = [row[0] for row in cursor.fetchall()]
|
||||
|
||||
# Delete vectors
|
||||
if self.vector_available and chunk_ids:
|
||||
for chunk_id in chunk_ids:
|
||||
try:
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.vector_table_name} WHERE id = ?",
|
||||
(chunk_id,),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Vector delete failed: {e}")
|
||||
|
||||
# Delete FTS entries
|
||||
if self.fts_available:
|
||||
try:
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.fts_table_name} WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"FTS delete failed: {e}")
|
||||
|
||||
# Delete chunks and file
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.chunks_table_name} WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.files_table_name} WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
|
||||
cursor.execute("COMMIT")
|
||||
except Exception:
|
||||
cursor.execute("ROLLBACK")
|
||||
|
|
@ -268,62 +295,132 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
finally:
|
||||
cursor.close()
|
||||
|
||||
async def _delete_file_internal(self, cursor: sqlite3.Cursor, path: str, source: MemorySource):
|
||||
"""Internal delete helper."""
|
||||
# Get chunk IDs for vector deletion
|
||||
cursor.execute(
|
||||
"SELECT id FROM chunks WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
chunk_ids = [row[0] for row in cursor.fetchall()]
|
||||
async def delete_file_chunks(self, path: str, chunk_ids: list[str]):
|
||||
"""Delete specific chunks for a file."""
|
||||
if not chunk_ids:
|
||||
return
|
||||
|
||||
# Delete vectors
|
||||
if self.vector_available and chunk_ids:
|
||||
for chunk_id in chunk_ids:
|
||||
cursor = self.conn.cursor()
|
||||
try:
|
||||
cursor.execute("BEGIN")
|
||||
|
||||
# Delete vectors
|
||||
if self.vector_available:
|
||||
for chunk_id in chunk_ids:
|
||||
try:
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.vector_table_name} WHERE id = ?",
|
||||
(chunk_id,),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Vector delete failed for {chunk_id}: {e}")
|
||||
|
||||
# Delete FTS entries
|
||||
if self.fts_available:
|
||||
placeholders = ",".join("?" * len(chunk_ids))
|
||||
try:
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.VECTOR_TABLE} WHERE id = ?",
|
||||
(chunk_id,),
|
||||
f"DELETE FROM {self.fts_table_name} WHERE id IN ({placeholders})",
|
||||
chunk_ids,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Vector delete failed: {e}")
|
||||
logger.debug(f"FTS delete failed: {e}")
|
||||
|
||||
# Delete FTS entries
|
||||
if self.fts_available:
|
||||
try:
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.FTS_TABLE} WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"FTS delete failed: {e}")
|
||||
# Delete chunks
|
||||
placeholders = ",".join("?" * len(chunk_ids))
|
||||
cursor.execute(
|
||||
f"DELETE FROM {self.chunks_table_name} WHERE id IN ({placeholders})",
|
||||
chunk_ids,
|
||||
)
|
||||
|
||||
# Delete chunks and file
|
||||
cursor.execute(
|
||||
"DELETE FROM chunks WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
cursor.execute(
|
||||
"DELETE FROM files WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
cursor.execute("COMMIT")
|
||||
except Exception:
|
||||
cursor.execute("ROLLBACK")
|
||||
raise
|
||||
finally:
|
||||
cursor.close()
|
||||
|
||||
async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource):
|
||||
"""Insert or update specific chunks without affecting other chunks."""
|
||||
if not chunks:
|
||||
return
|
||||
|
||||
async def get_file_hash(self, path: str, source: MemorySource) -> str | None:
|
||||
"""Get file hash."""
|
||||
cursor = self.conn.cursor()
|
||||
cursor.execute(
|
||||
"SELECT hash FROM files WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
row = cursor.fetchone()
|
||||
try:
|
||||
cursor.execute("BEGIN")
|
||||
|
||||
now = int(time.time() * 1000)
|
||||
for chunk in chunks:
|
||||
# Insert/update chunk
|
||||
cursor.execute(
|
||||
f"""
|
||||
INSERT OR REPLACE INTO {self.chunks_table_name} (
|
||||
id, path, source, start_line, end_line,
|
||||
hash, text, embedding, updated_at
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
chunk.id,
|
||||
chunk.path,
|
||||
source.value,
|
||||
chunk.start_line,
|
||||
chunk.end_line,
|
||||
chunk.hash,
|
||||
chunk.text,
|
||||
json.dumps(chunk.embedding) if chunk.embedding else None,
|
||||
now,
|
||||
),
|
||||
)
|
||||
|
||||
# Insert/update vector
|
||||
if self.vector_available:
|
||||
assert chunk.embedding, "Embedding is required for vector insert"
|
||||
cursor.execute(
|
||||
f"""
|
||||
INSERT OR REPLACE INTO {self.vector_table_name} (id, embedding)
|
||||
VALUES (?, ?)
|
||||
""",
|
||||
(chunk.id, self.vector_to_blob(chunk.embedding)),
|
||||
)
|
||||
|
||||
# Insert/update FTS
|
||||
if self.fts_available:
|
||||
cursor.execute(
|
||||
f"""
|
||||
INSERT OR REPLACE INTO {self.fts_table_name} (
|
||||
text, id, path, source, start_line, end_line
|
||||
) VALUES (?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
chunk.text,
|
||||
chunk.id,
|
||||
chunk.path,
|
||||
source.value,
|
||||
chunk.start_line,
|
||||
chunk.end_line,
|
||||
),
|
||||
)
|
||||
|
||||
cursor.execute("COMMIT")
|
||||
except Exception:
|
||||
cursor.execute("ROLLBACK")
|
||||
raise
|
||||
finally:
|
||||
cursor.close()
|
||||
|
||||
async def list_files(self, source: MemorySource) -> list[str]:
|
||||
"""List all indexed files."""
|
||||
cursor = self.conn.cursor()
|
||||
cursor.execute(f"SELECT path FROM {self.files_table_name} WHERE source = ?", (source.value,))
|
||||
paths = [row[0] for row in cursor.fetchall()]
|
||||
cursor.close()
|
||||
return row[0] if row else None
|
||||
return paths
|
||||
|
||||
async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None:
|
||||
"""Get file metadata with chunk count."""
|
||||
cursor = self.conn.cursor()
|
||||
cursor.execute(
|
||||
"SELECT hash, mtime, size FROM files WHERE path = ? AND source = ?",
|
||||
f"SELECT hash, mtime, size FROM {self.files_table_name} WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
row = cursor.fetchone()
|
||||
|
|
@ -333,7 +430,7 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
|
||||
hash_val, mtime, size = row
|
||||
cursor.execute(
|
||||
"SELECT COUNT(*) FROM chunks WHERE path = ? AND source = ?",
|
||||
f"SELECT COUNT(*) FROM {self.chunks_table_name} WHERE path = ? AND source = ?",
|
||||
(path, source.value),
|
||||
)
|
||||
chunk_count = cursor.fetchone()[0]
|
||||
|
|
@ -343,24 +440,17 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
hash=hash_val,
|
||||
mtime_ms=mtime,
|
||||
size=size,
|
||||
path=path,
|
||||
chunk_count=chunk_count,
|
||||
)
|
||||
|
||||
async def list_files(self, source: MemorySource) -> list[str]:
|
||||
"""List all indexed files."""
|
||||
cursor = self.conn.cursor()
|
||||
cursor.execute("SELECT path FROM files WHERE source = ?", (source.value,))
|
||||
paths = [row[0] for row in cursor.fetchall()]
|
||||
cursor.close()
|
||||
return paths
|
||||
|
||||
async def get_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
|
||||
async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
|
||||
"""Get all chunks for a file."""
|
||||
cursor = self.conn.cursor()
|
||||
cursor.execute(
|
||||
"""
|
||||
f"""
|
||||
SELECT id, path, source, start_line, end_line, text, hash, embedding
|
||||
FROM chunks WHERE path = ? AND source = ?
|
||||
FROM {self.chunks_table_name} WHERE path = ? AND source = ?
|
||||
ORDER BY start_line
|
||||
""",
|
||||
(path, source.value),
|
||||
|
|
@ -418,23 +508,25 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
|
||||
try:
|
||||
query_blob = self.vector_to_blob(query_embedding)
|
||||
|
||||
# Correct SQLite-vec syntax for vector search with limit
|
||||
# vec0 requires 'k = ?' constraint for knn queries
|
||||
query_sql = f"""
|
||||
SELECT c.id, c.path, c.start_line, c.end_line, c.source, c.text, v.distance
|
||||
FROM {self.VECTOR_TABLE} v
|
||||
JOIN chunks c ON v.id = c.id
|
||||
FROM {self.vector_table_name} v
|
||||
JOIN {self.chunks_table_name} c ON v.id = c.id
|
||||
WHERE v.embedding MATCH ?
|
||||
AND k = ?
|
||||
"""
|
||||
query_params: list = [query_blob]
|
||||
query_params: list = [query_blob, limit]
|
||||
|
||||
# Add source filter if specified
|
||||
if source_filter:
|
||||
query_sql += source_filter
|
||||
query_params.extend(params)
|
||||
|
||||
# Order and limit results
|
||||
query_sql += " ORDER BY v.distance LIMIT ?"
|
||||
query_params.append(str(limit))
|
||||
# Order by distance (k constraint already limits results)
|
||||
query_sql += " ORDER BY v.distance"
|
||||
|
||||
cursor.execute(query_sql, query_params)
|
||||
|
||||
|
|
@ -470,11 +562,21 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
if not self.fts_available:
|
||||
return []
|
||||
|
||||
# Build FTS5 query, escaping quotes
|
||||
cleaned = query.strip().replace('"', '""')
|
||||
# Build FTS5 query
|
||||
# Split query into tokens and join with OR for better recall
|
||||
# Individual words are automatically stemmed and matched by FTS5
|
||||
cleaned = query.strip()
|
||||
if not cleaned:
|
||||
return []
|
||||
fts_query = f'"{cleaned}"'
|
||||
|
||||
# Split into words and escape each
|
||||
words = cleaned.split()
|
||||
if not words:
|
||||
return []
|
||||
|
||||
# Use OR operator for better recall - match any of the query words
|
||||
escaped_words = [word.replace('"', '""') for word in words]
|
||||
fts_query = " OR ".join(escaped_words)
|
||||
|
||||
cursor = self.conn.cursor()
|
||||
source_filter = ""
|
||||
|
|
@ -490,7 +592,7 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
f"""
|
||||
SELECT fts.id, fts.path, fts.start_line, fts.end_line,
|
||||
fts.source, fts.text, rank
|
||||
FROM {self.FTS_TABLE} fts
|
||||
FROM {self.fts_table_name} fts
|
||||
WHERE fts.text MATCH ?{source_filter}
|
||||
ORDER BY rank
|
||||
LIMIT ?
|
||||
|
|
@ -521,52 +623,20 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
finally:
|
||||
cursor.close()
|
||||
|
||||
async def read_meta(self, key: str) -> MemoryIndexMeta | None:
|
||||
"""Read metadata value."""
|
||||
cursor = self.conn.cursor()
|
||||
cursor.execute("SELECT value FROM meta WHERE key = ?", (key,))
|
||||
row = cursor.fetchone()
|
||||
cursor.close()
|
||||
|
||||
if not row:
|
||||
return None
|
||||
|
||||
return MemoryIndexMeta(**json.loads(row[0]))
|
||||
|
||||
async def write_meta(self, key: str, value: MemoryIndexMeta | dict) -> None:
|
||||
"""Write metadata value."""
|
||||
data = value.model_dump() if isinstance(value, MemoryIndexMeta) else value
|
||||
cursor = self.conn.cursor()
|
||||
cursor.execute(
|
||||
"""
|
||||
INSERT OR REPLACE INTO meta (key, value)
|
||||
VALUES (?, ?)
|
||||
""",
|
||||
(key, json.dumps(data)),
|
||||
)
|
||||
self.conn.commit()
|
||||
cursor.close()
|
||||
|
||||
async def clear_all(self):
|
||||
"""Clear all indexed data."""
|
||||
cursor = self.conn.cursor()
|
||||
cursor.execute("BEGIN")
|
||||
|
||||
try:
|
||||
cursor.execute("DELETE FROM files")
|
||||
cursor.execute("DELETE FROM chunks")
|
||||
cursor.execute(f"DELETE FROM {self.files_table_name}")
|
||||
cursor.execute(f"DELETE FROM {self.chunks_table_name}")
|
||||
|
||||
if self.vector_available:
|
||||
try:
|
||||
cursor.execute(f"DELETE FROM {self.VECTOR_TABLE}")
|
||||
except Exception as e:
|
||||
logger.debug(f"Vector clear failed: {e}")
|
||||
cursor.execute(f"DELETE FROM {self.vector_table_name}")
|
||||
|
||||
if self.fts_available:
|
||||
try:
|
||||
cursor.execute(f"DELETE FROM {self.FTS_TABLE}")
|
||||
except Exception as e:
|
||||
logger.debug(f"FTS clear failed: {e}")
|
||||
cursor.execute(f"DELETE FROM {self.fts_table_name}")
|
||||
|
||||
cursor.execute("COMMIT")
|
||||
except Exception:
|
||||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -6,15 +6,10 @@ from pydantic import BaseModel, Field
|
|||
class FileMetadata(BaseModel):
|
||||
"""File metadata with optional extended fields for various use cases."""
|
||||
|
||||
# Core fields (always required)
|
||||
hash: str = Field(default=..., description="Hash of the file content")
|
||||
mtime_ms: float = Field(default=..., description="Last modification time in milliseconds")
|
||||
size: int = Field(default=..., description="File size in bytes")
|
||||
|
||||
# Extended fields for session files
|
||||
path: str | None = Field(default=None, description="Relative path to the session file")
|
||||
abs_path: str | None = Field(default=None, description="Absolute path to the session file")
|
||||
content: str | None = Field(default=None, description="Parsed content from the session file")
|
||||
|
||||
# Extended fields for statistics
|
||||
chunk_count: int | None = Field(default=None, description="Number of chunks in the file")
|
||||
metadata: dict = Field(default_factory=dict, description="Additional metadata")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -1,19 +1,17 @@
|
|||
"""Chunking logic for Markdown files."""
|
||||
|
||||
from typing import List, Dict, Any
|
||||
|
||||
from ..utils.hashing import hash_text
|
||||
from ...enumeration import MemorySource
|
||||
from ...schema import MemoryChunk
|
||||
from .common_utils import hash_text
|
||||
from ..enumeration import MemorySource
|
||||
from ..schema import MemoryChunk
|
||||
|
||||
|
||||
def chunk_markdown(
|
||||
text: str,
|
||||
path: str,
|
||||
source: MemorySource,
|
||||
chunk_tokens: int = 300,
|
||||
overlap: int = 30,
|
||||
) -> List[MemoryChunk]:
|
||||
chunk_tokens: int,
|
||||
overlap: int,
|
||||
) -> list[MemoryChunk]:
|
||||
"""
|
||||
Markdown chunking logic implemented based on the TypeScript version.
|
||||
|
||||
|
|
@ -35,10 +33,10 @@ def chunk_markdown(
|
|||
max_chars = max(32, chunk_tokens * 4)
|
||||
overlap_chars = max(0, overlap * 4)
|
||||
|
||||
chunks: List[MemoryChunk] = []
|
||||
chunks: list[MemoryChunk] = []
|
||||
|
||||
# Currently building chunk
|
||||
current: List[Dict[str, Any]] = [] # [{'line': str, 'line_no': int}]
|
||||
current: list[dict] = [] # [{'line': str, 'line_no': int}]
|
||||
current_chars = 0
|
||||
|
||||
def flush():
|
||||
|
|
@ -83,8 +81,8 @@ def chunk_markdown(
|
|||
kept = []
|
||||
|
||||
# Collect lines from the end until reaching overlap size
|
||||
for i in range(len(current) - 1, -1, -1):
|
||||
entry = current[i]
|
||||
for j in range(len(current) - 1, -1, -1):
|
||||
entry = current[j]
|
||||
if not entry:
|
||||
continue
|
||||
|
||||
|
|
@ -123,4 +121,4 @@ def chunk_markdown(
|
|||
# Process the final chunk
|
||||
flush()
|
||||
|
||||
return chunks
|
||||
return [c for c in chunks if c.text.strip()]
|
||||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -58,7 +58,7 @@ class ReMe(Application):
|
|||
target_user_names: list[str] | None = None,
|
||||
target_task_names: list[str] | None = None,
|
||||
target_tool_names: list[str] | None = None,
|
||||
profile_dir: str = "reme_profile",
|
||||
profile_dir: str = ".reme/profile",
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize ReMe with config.
|
||||
|
|
|
|||
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
|
||||
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())
|
||||
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