mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
commit
05e87afbd2
46 changed files with 3808 additions and 692 deletions
|
|
@ -30,14 +30,42 @@ classifiers = [
|
|||
"Typing :: Typed",
|
||||
]
|
||||
|
||||
keywords = ["llm", "memory", "experience", "memoryscope", "ai", "mcp", "http"]
|
||||
keywords = ["llm", "memory", "experience", "memoryscope", "ai", "mcp", "http", "reme", "personal"]
|
||||
|
||||
dependencies = [
|
||||
"flowllm[reme]>=0.2.0.10",
|
||||
"sqlite-vec>=0.1.6",
|
||||
"prompt_toolkit>=3.0.52",
|
||||
"rich>=14.2.0",
|
||||
"asyncpg>=0.31.0",
|
||||
"chromadb>=1.3.5",
|
||||
"dashscope>=1.25.1",
|
||||
"elasticsearch>=9.2.0",
|
||||
"fastapi>=0.121.3",
|
||||
"fastmcp>=2.14.1",
|
||||
"httpx>=0.28.1",
|
||||
"litellm>=1.80.0",
|
||||
"loguru>=0.7.3",
|
||||
"mcp>=1.25.0",
|
||||
"numpy>=2.2.6",
|
||||
"openai>=2.8.1",
|
||||
"pandas>=2.3.3",
|
||||
"pydantic>=2.12.4",
|
||||
"qdrant-client>=1.16.0",
|
||||
"tavily-python>=0.7.13",
|
||||
"tiktoken>=0.12.0",
|
||||
"tqdm>=4.67.1",
|
||||
"transformers>=4.57.3",
|
||||
"uvicorn>=0.40.0",
|
||||
"watchfiles>=1.1.1",
|
||||
"pyyaml>=6.0.3",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
ray = [
|
||||
"ray",
|
||||
]
|
||||
|
||||
dev = [
|
||||
"jupyter-book",
|
||||
"ghp-import",
|
||||
|
|
@ -48,12 +76,8 @@ dev = [
|
|||
"pre-commit",
|
||||
]
|
||||
|
||||
token = [
|
||||
"flowllm[token]>=0.2.0.10"
|
||||
]
|
||||
|
||||
full = [
|
||||
"reme_ai[dev,token]"
|
||||
"reme_ai[dev,ray]"
|
||||
]
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
|
|
@ -85,6 +109,7 @@ Repository = "https://github.com/agentscope-ai/ReMe"
|
|||
[project.scripts]
|
||||
reme = "reme_ai.main:main"
|
||||
reme2 = "reme.reme:main"
|
||||
remefs = "reme.reme_fs:main"
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
asyncio_default_fixture_loop_scope = "function"
|
||||
|
|
|
|||
|
|
@ -19,3 +19,10 @@ __all__ = [
|
|||
]
|
||||
|
||||
__version__ = "0.3.0.0a1"
|
||||
|
||||
|
||||
"""
|
||||
conda create -n fl_test2 python=3.10
|
||||
conda activate fl_test2
|
||||
conda env remove -n fl_test2
|
||||
"""
|
||||
|
|
@ -1,13 +1,16 @@
|
|||
"""chat agent"""
|
||||
|
||||
from .fs_cli import FsCli
|
||||
from .simple_chat import SimpleChat
|
||||
from .stream_chat import StreamChat
|
||||
from ...core import R
|
||||
|
||||
__all__ = [
|
||||
"FsCli",
|
||||
"StreamChat",
|
||||
"SimpleChat",
|
||||
]
|
||||
|
||||
R.ops.register(FsCli)
|
||||
R.ops.register(SimpleChat)
|
||||
R.ops.register(StreamChat)
|
||||
|
|
|
|||
174
reme/agent/chat/fs_cli.py
Normal file
174
reme/agent/chat/fs_cli.py
Normal file
|
|
@ -0,0 +1,174 @@
|
|||
"""FsCli system prompt"""
|
||||
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
from ...core.enumeration import Role, ChunkEnum
|
||||
from ...core.op import BaseReactStream
|
||||
from ...core.schema import Message, StreamChunk
|
||||
|
||||
|
||||
class FsCli(BaseReactStream):
|
||||
"""FsCli agent with system prompt."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
working_dir: str,
|
||||
context_window_tokens: int = 128000,
|
||||
reserve_tokens: int = 36000,
|
||||
keep_recent_tokens: int = 20000,
|
||||
hybrid_enabled: bool = True,
|
||||
hybrid_vector_weight: float = 0.7,
|
||||
hybrid_text_weight: float = 0.3,
|
||||
hybrid_candidate_multiplier: float = 3.0,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
self.working_dir: str = working_dir
|
||||
Path(self.working_dir).mkdir(parents=True, exist_ok=True)
|
||||
self.context_window_tokens: int = context_window_tokens
|
||||
self.reserve_tokens: int = reserve_tokens
|
||||
self.keep_recent_tokens: int = keep_recent_tokens
|
||||
self.hybrid_enabled: bool = hybrid_enabled
|
||||
self.hybrid_vector_weight: float = hybrid_vector_weight
|
||||
self.hybrid_text_weight: float = hybrid_text_weight
|
||||
self.hybrid_candidate_multiplier: float = hybrid_candidate_multiplier
|
||||
|
||||
self.messages: list[Message] = []
|
||||
self.previous_summary: str = ""
|
||||
|
||||
async def reset(self) -> str:
|
||||
"""Reset conversation history using summary.
|
||||
|
||||
Summarizes current messages to memory files and clears history.
|
||||
"""
|
||||
if not self.messages:
|
||||
self.messages.clear()
|
||||
self.previous_summary = ""
|
||||
return "No history to reset."
|
||||
|
||||
# Import required modules
|
||||
from ..fs import FsSummarizer
|
||||
|
||||
# Summarize current conversation and save to memory files
|
||||
current_date = datetime.now().strftime("%Y-%m-%d")
|
||||
summarizer = FsSummarizer(tools=self.tools, working_dir=self.working_dir)
|
||||
|
||||
result = await summarizer.call(messages=self.messages, date=current_date, service_context=self.service_context)
|
||||
self.messages.clear()
|
||||
self.previous_summary = ""
|
||||
return f"History saved to memory files and reset. Result: {result.get('answer', 'Done')}"
|
||||
|
||||
async def context_check(self) -> dict:
|
||||
"""Check if messages exceed token limits."""
|
||||
# Import required modules
|
||||
from ..fs import FsContextChecker
|
||||
|
||||
# Step 1: Check and find cut point
|
||||
checker = FsContextChecker(
|
||||
context_window_tokens=self.context_window_tokens,
|
||||
reserve_tokens=self.reserve_tokens,
|
||||
keep_recent_tokens=self.keep_recent_tokens,
|
||||
)
|
||||
return await checker.call(messages=self.messages, service_context=self.service_context)
|
||||
|
||||
async def compact(self, force_compact: bool = False) -> str:
|
||||
"""Compact history then reset.
|
||||
|
||||
First compacts messages if they exceed token limits (generating a summary),
|
||||
then calls reset_history to save to files and clear.
|
||||
|
||||
Args:
|
||||
force_compact: If True, force compaction of all messages into summary
|
||||
|
||||
Returns:
|
||||
str: Summary of compaction result
|
||||
"""
|
||||
if not self.messages:
|
||||
return "No history to compact."
|
||||
|
||||
# Import required modules
|
||||
from ..fs import FsCompactor
|
||||
|
||||
# Step 1: Check and find cut point
|
||||
cut_result = await self.context_check()
|
||||
tokens_before = cut_result.get("token_count", 0)
|
||||
|
||||
if force_compact:
|
||||
# Force compact: summarize all messages, leave only summary
|
||||
messages_to_summarize = self.messages
|
||||
turn_prefix_messages = []
|
||||
left_messages = []
|
||||
elif not cut_result.get("needs_compaction", False):
|
||||
# No compaction needed
|
||||
return "History is within token limits, no compaction needed."
|
||||
else:
|
||||
# Normal compaction: use cut point result
|
||||
messages_to_summarize = cut_result.get("messages_to_summarize", [])
|
||||
turn_prefix_messages = cut_result.get("turn_prefix_messages", [])
|
||||
left_messages = cut_result.get("left_messages", [])
|
||||
|
||||
# Step 2: Generate summary via Compactor
|
||||
compactor = FsCompactor()
|
||||
summary_content = await compactor.call(
|
||||
messages_to_summarize=messages_to_summarize,
|
||||
turn_prefix_messages=turn_prefix_messages,
|
||||
previous_summary=self.previous_summary,
|
||||
service_context=self.service_context,
|
||||
)
|
||||
|
||||
# Step 3: Assemble final messages
|
||||
summary_message = Message(role=Role.USER, content=summary_content)
|
||||
self.messages = [summary_message] + left_messages
|
||||
self.previous_summary = summary_content
|
||||
|
||||
# Step 4: Call reset_history to save and clear
|
||||
reset_result = await self.reset()
|
||||
return f"History compacted from {tokens_before} tokens. {reset_result}"
|
||||
|
||||
async def build_messages(self) -> list[Message]:
|
||||
"""Build system prompt message."""
|
||||
current_time = datetime.now().strftime("%Y-%m-%d %H:%M:%S %A")
|
||||
|
||||
system_prompt = self.prompt_format(
|
||||
"system_prompt",
|
||||
workspace_dir=self.working_dir,
|
||||
current_time=current_time,
|
||||
has_previous_summary=bool(self.previous_summary),
|
||||
previous_summary=self.previous_summary or "",
|
||||
)
|
||||
|
||||
return [
|
||||
Message(role=Role.SYSTEM, content=system_prompt),
|
||||
*self.messages,
|
||||
Message(role=Role.USER, content=self.context.query),
|
||||
]
|
||||
|
||||
async def execute(self):
|
||||
"""Execute the agent."""
|
||||
messages = await self.build_messages()
|
||||
|
||||
t_tools, messages, success = await self.react(messages, self.tools)
|
||||
|
||||
# Update self.messages: react() returns [SYSTEM, ...history...],
|
||||
# so we remove the first SYSTEM message
|
||||
self.messages = messages[1:]
|
||||
|
||||
# Emit final done signal
|
||||
await self.context.add_stream_chunk(
|
||||
StreamChunk(
|
||||
chunk_type=ChunkEnum.DONE,
|
||||
chunk="",
|
||||
metadata={
|
||||
"success": success,
|
||||
"total_steps": len(t_tools),
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
return {
|
||||
"answer": messages[-1].content if success else "",
|
||||
"success": success,
|
||||
"messages": messages,
|
||||
"tools": t_tools,
|
||||
}
|
||||
63
reme/agent/chat/fs_cli.yaml
Normal file
63
reme/agent/chat/fs_cli.yaml
Normal file
|
|
@ -0,0 +1,63 @@
|
|||
system_prompt: |
|
||||
You are a personal assistant named Remy.
|
||||
|
||||
## Current Time
|
||||
{current_time}
|
||||
|
||||
[has_previous_summary]## Previous Conversation Summary
|
||||
[has_previous_summary]<previous-summary>
|
||||
[has_previous_summary]{previous_summary}
|
||||
[has_previous_summary]</previous-summary>
|
||||
[has_previous_summary]
|
||||
[has_previous_summary]The above is a summary of our previous conversation. Use it as context to maintain continuity.
|
||||
|
||||
## Memory System
|
||||
You wake up fresh each session. These files provide continuity:
|
||||
|
||||
### 📝 Daily Notes: `memory/YYYY-MM-DD.md`
|
||||
- Raw logs of what happened today
|
||||
- Create `memory/` directory if needed
|
||||
- Write events, conversations, tasks, decisions as they happen
|
||||
- Capture what matters
|
||||
|
||||
### 🧠 Long-Term Memory: `MEMORY.md`
|
||||
- Your curated memories, like a human's long-term memory
|
||||
- The distilled essence, not raw logs
|
||||
- Contains: significant events, thoughts, decisions, opinions, lessons learned
|
||||
- Maintenance: periodically review daily files and promote important context here
|
||||
|
||||
### 🔍 Memory Recall
|
||||
Before answering questions about prior work, decisions, dates, people, preferences, or todos:
|
||||
1. Run `memory_search` on MEMORY.md + memory/*.md
|
||||
2. Use `memory_get` to pull only the needed lines
|
||||
|
||||
### 💾 Write It Down - No "Mental Notes"!
|
||||
- **Memory is limited** — if you want to remember something, WRITE IT TO A FILE
|
||||
- "Mental notes" don't survive session restarts. Files do.
|
||||
- When someone says "remember this" → update `memory/YYYY-MM-DD.md` or MEMORY.md
|
||||
- When you learn a lesson → update `memory/YYYY-MM-DD.md` or MEMORY.md
|
||||
- When you make a mistake → update `memory/YYYY-MM-DD.md` or MEMORY.md, so future-you doesn't repeat it
|
||||
- **Text > Brain** 📝
|
||||
|
||||
## Behavior Guidelines
|
||||
|
||||
### 😊 React Like a Human
|
||||
**React when:**
|
||||
- You appreciate something but don't need to reply (👍, ❤️, 🙌)
|
||||
- Something made you laugh (😂, 💀)
|
||||
- You find it interesting or thought-provoking (🤔, 💡)
|
||||
- You want to acknowledge without interrupting the flow
|
||||
- It's a simple yes/no or approval situation (✅, 👀)
|
||||
|
||||
**Why:** Reactions are lightweight social signals. Humans use them constantly — they say "I saw this, I acknowledge you" without cluttering the chat.
|
||||
|
||||
**Don't overdo it:** One reaction per message max. Pick the one that fits best.
|
||||
|
||||
### 🛡️ Safety Rules
|
||||
- Don't exfiltrate private data. Ever.
|
||||
- Don't run destructive commands without asking
|
||||
- Prefer `trash` over `rm` (recoverable beats gone forever)
|
||||
- When in doubt, ask
|
||||
|
||||
## Continuous Improvement
|
||||
This is a starting point. Add your own conventions, style, and rules as you figure out what works.
|
||||
|
|
@ -1,9 +1,11 @@
|
|||
"""File system agents for memory management."""
|
||||
|
||||
from .fs_compactor import FsCompactor
|
||||
from .fs_context_checker import FsContextChecker
|
||||
from .fs_summarizer import FsSummarizer
|
||||
|
||||
__all__ = [
|
||||
"FsSummarizer",
|
||||
"FsCompactor",
|
||||
"FsContextChecker",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -2,165 +2,62 @@
|
|||
|
||||
from loguru import logger
|
||||
|
||||
from ...core.enumeration import Role, MemoryType
|
||||
from ...core.op import BaseReact
|
||||
from ...core.schema import CutPointResult, Message
|
||||
from ...core.enumeration import Role
|
||||
from ...core.op import BaseOp
|
||||
from ...core.schema import Message
|
||||
from ...core.utils import format_messages
|
||||
|
||||
|
||||
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
|
||||
class FsCompactor(BaseOp):
|
||||
"""Generate summaries for conversation history compaction."""
|
||||
|
||||
@staticmethod
|
||||
def _normalize_messages(messages: list[Message | dict]) -> list[Message]:
|
||||
"""Convert dict messages to Message objects."""
|
||||
return [Message(**m) if isinstance(m, dict) else m for m in messages]
|
||||
|
||||
@staticmethod
|
||||
def _is_user_message(message: Message) -> bool:
|
||||
"""Check if message is user role."""
|
||||
return message.role is Role.USER
|
||||
|
||||
def _find_turn_start_index(self, messages: list[Message], entry_index: int) -> int:
|
||||
"""Find user message that starts the turn. Returns -1 if not found."""
|
||||
if not messages or entry_index < 0 or entry_index >= len(messages):
|
||||
return -1
|
||||
|
||||
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]) -> CutPointResult:
|
||||
"""
|
||||
Find cut point with split turn detection.
|
||||
|
||||
Split turn: User → Assistant → [CUT] → Assistant → User
|
||||
Clean cut: User → [CUT] → Assistant → User
|
||||
"""
|
||||
if not messages:
|
||||
return CutPointResult()
|
||||
|
||||
accumulated_tokens = 0
|
||||
cut_index = 0
|
||||
|
||||
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
|
||||
logger.debug(f"Cut point at index {cut_index}, {accumulated_tokens} tokens")
|
||||
break
|
||||
|
||||
if cut_index == 0:
|
||||
return CutPointResult(left_messages=messages)
|
||||
|
||||
cut_message = messages[cut_index]
|
||||
is_user_cut = self._is_user_message(cut_message)
|
||||
|
||||
if is_user_cut:
|
||||
return CutPointResult(
|
||||
messages_to_summarize=messages[:cut_index],
|
||||
left_messages=messages[cut_index:],
|
||||
cut_index=cut_index,
|
||||
)
|
||||
|
||||
turn_start_index = self._find_turn_start_index(messages, cut_index)
|
||||
|
||||
if turn_start_index == -1:
|
||||
logger.warning("Split turn detected but no turn start found, treating as clean cut")
|
||||
return CutPointResult(
|
||||
messages_to_summarize=messages[:cut_index],
|
||||
left_messages=messages[cut_index:],
|
||||
cut_index=cut_index,
|
||||
)
|
||||
|
||||
return CutPointResult(
|
||||
messages_to_summarize=messages[:turn_start_index],
|
||||
turn_prefix_messages=messages[turn_start_index:cut_index],
|
||||
left_messages=messages[cut_index:],
|
||||
is_split_turn=True,
|
||||
cut_index=cut_index,
|
||||
)
|
||||
|
||||
async def _generate_summary(self, prompt_messages: list[Message]) -> str:
|
||||
"""Generate summary via LLM. Returns empty string if no messages."""
|
||||
if not prompt_messages:
|
||||
return ""
|
||||
|
||||
try:
|
||||
assistant_message = await self.llm.chat(prompt_messages)
|
||||
return assistant_message.content if assistant_message.content else ""
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to generate summary: {e}")
|
||||
raise RuntimeError(f"Summarization failed: {e}") from e
|
||||
assistant_message = await self.llm.chat(prompt_messages)
|
||||
return assistant_message.content
|
||||
|
||||
@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 format_messages(
|
||||
messages=messages,
|
||||
add_index=False,
|
||||
add_time=False,
|
||||
use_name=True,
|
||||
add_reasoning=False,
|
||||
add_tools=True,
|
||||
strip_markdown_headers=False,
|
||||
)
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
def build_messages_s1(self) -> list[Message]:
|
||||
def _build_history_prompt(self, messages_to_summarize: list[Message], previous_summary: str = "") -> list[Message]:
|
||||
"""Build prompt for main history summary."""
|
||||
messages = self._normalize_messages(self.context.messages)
|
||||
cut_result = self._find_cut_point(messages)
|
||||
|
||||
self.context.is_split_turn = cut_result.is_split_turn
|
||||
self.context.turn_prefix_messages = cut_result.turn_prefix_messages
|
||||
self.context.left_messages = cut_result.left_messages
|
||||
|
||||
if not cut_result.messages_to_summarize:
|
||||
logger.info("No messages to summarize")
|
||||
if not 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)
|
||||
if previous_summary:
|
||||
user_prompt = self.prompt_format("update_user_message", previous_summary=previous_summary)
|
||||
else:
|
||||
user_prompt = self.get_prompt("initial_user_message")
|
||||
conversation_text = self._serialize_conversation(cut_result.messages_to_summarize)
|
||||
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]:
|
||||
def _build_turn_prefix_prompt(self, turn_prefix_messages: list[Message]) -> list[Message]:
|
||||
"""Build prompt for turn prefix summary (split turn only)."""
|
||||
if not self.context.turn_prefix_messages:
|
||||
if not turn_prefix_messages:
|
||||
return []
|
||||
|
||||
system_prompt = self.get_prompt("system_prompt")
|
||||
conversation_text = self._serialize_conversation(self.context.turn_prefix_messages)
|
||||
conversation_text = self._serialize_conversation(turn_prefix_messages)
|
||||
turn_prefix_prompt = self.prompt_format("turn_prefix_summarization", conversation_text=conversation_text)
|
||||
|
||||
return [
|
||||
|
|
@ -168,58 +65,38 @@ class FsCompactor(BaseReact):
|
|||
Message(role=Role.USER, content=turn_prefix_prompt),
|
||||
]
|
||||
|
||||
async def execute(self):
|
||||
async def execute(self) -> str:
|
||||
"""
|
||||
Execute compaction if needed.
|
||||
Generate summary for conversation history.
|
||||
|
||||
Returns: [summary_message, ...left_messages] if compacted, else original messages.
|
||||
Expects context to have:
|
||||
- messages_to_summarize: list[Message] (required)
|
||||
- turn_prefix_messages: list[Message] (optional, for split turn)
|
||||
- previous_summary: str (optional, for incremental summarization)
|
||||
|
||||
Returns:
|
||||
str: Generated summary text formatted with compaction_summary_format.
|
||||
Returns empty string if no messages to summarize.
|
||||
"""
|
||||
original_messages = self._normalize_messages(self.context.messages)
|
||||
token_count: int = self.token_counter.count_token(original_messages)
|
||||
threshold = self.context_window_tokens - self.reserve_tokens
|
||||
messages_to_summarize = self.context.get("messages_to_summarize", [])
|
||||
turn_prefix_messages = self.context.get("turn_prefix_messages", [])
|
||||
previous_summary = self.context.get("previous_summary", "")
|
||||
|
||||
if token_count < threshold:
|
||||
logger.info(f"Token count {token_count} below threshold ({threshold}), skipping compaction")
|
||||
return {
|
||||
"compacted": False,
|
||||
"tokens_before": token_count,
|
||||
"is_split_turn": False,
|
||||
"messages": original_messages,
|
||||
}
|
||||
|
||||
logger.info(f"Starting compaction, token count: {token_count}, threshold: {threshold}")
|
||||
|
||||
history_prompt_messages = self.build_messages_s1()
|
||||
|
||||
if not history_prompt_messages and not self.context.get("is_split_turn"):
|
||||
logger.warning("No messages to summarize and not a split turn, returning original messages")
|
||||
return {
|
||||
"compacted": False,
|
||||
"tokens_before": token_count,
|
||||
"is_split_turn": False,
|
||||
"messages": original_messages,
|
||||
}
|
||||
|
||||
history_summary = await self._generate_summary(history_prompt_messages) if history_prompt_messages else ""
|
||||
|
||||
if self.context.is_split_turn and self.context.turn_prefix_messages:
|
||||
logger.info("Split turn detected, generating turn prefix summary")
|
||||
turn_prefix_prompt_messages = self.build_messages_s2()
|
||||
turn_prefix_summary = await self._generate_summary(turn_prefix_prompt_messages)
|
||||
summary = f"{history_summary}\n\n---\n\n**Turn Context (split turn):**\n\n{turn_prefix_summary}"
|
||||
messages_to_summarize = self._normalize_messages(messages_to_summarize)
|
||||
if messages_to_summarize:
|
||||
history_prompt_messages = self._build_history_prompt(messages_to_summarize, previous_summary)
|
||||
history_summary = "**Turn Context**:\n\n" + await self._generate_summary(history_prompt_messages)
|
||||
else:
|
||||
summary = history_summary
|
||||
history_summary = ""
|
||||
|
||||
logger.info(f"Compaction complete, summary length: {len(summary)}, split_turn: {self.context.is_split_turn}")
|
||||
turn_prefix_messages = self._normalize_messages(turn_prefix_messages)
|
||||
if turn_prefix_messages:
|
||||
turn_prefix_prompt_messages = self._build_turn_prefix_prompt(turn_prefix_messages)
|
||||
turn_prefix_summary = "**Turn Context**:\n\n" + await self._generate_summary(turn_prefix_prompt_messages)
|
||||
else:
|
||||
turn_prefix_summary = ""
|
||||
|
||||
summary = "\n\n---".join([history_summary, turn_prefix_summary])
|
||||
summary_content = self.prompt_format("compaction_summary_format", summary=summary)
|
||||
summary_message = Message(role=Role.USER, content=summary_content)
|
||||
left_messages = self.context.get("left_messages", [])
|
||||
final_messages = [summary_message] + left_messages
|
||||
|
||||
return {
|
||||
"compacted": True,
|
||||
"tokens_before": token_count,
|
||||
"is_split_turn": self.context.is_split_turn,
|
||||
"messages": final_messages,
|
||||
}
|
||||
logger.info(f"Generated summary: {summary}")
|
||||
return summary_content
|
||||
|
|
|
|||
161
reme/agent/fs/fs_context_checker.py
Normal file
161
reme/agent/fs/fs_context_checker.py
Normal file
|
|
@ -0,0 +1,161 @@
|
|||
"""Context window limit checker for reactive agents."""
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from ...core.enumeration import Role
|
||||
from ...core.op import BaseReact
|
||||
from ...core.schema import CutPointResult, Message
|
||||
|
||||
|
||||
class FsContextChecker(BaseReact):
|
||||
"""Check if context exceeds token limits and find cut point for compaction."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
context_window_tokens: int = 128000,
|
||||
reserve_tokens: int = 36000,
|
||||
keep_recent_tokens: int = 20000,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Initialize context checker.
|
||||
|
||||
Args:
|
||||
context_window_tokens: Total context window size.
|
||||
reserve_tokens: Tokens to reserve for output and overhead.
|
||||
keep_recent_tokens: Tokens to keep in recent messages.
|
||||
**kwargs: Additional BaseReact arguments.
|
||||
"""
|
||||
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 _normalize_messages(messages: list[Message | dict]) -> list[Message]:
|
||||
"""Convert dict messages to Message objects."""
|
||||
return [Message(**m) if isinstance(m, dict) else m for m in messages]
|
||||
|
||||
@staticmethod
|
||||
def _is_user_message(message: Message) -> bool:
|
||||
"""Check if message is user role."""
|
||||
return message.role is Role.USER
|
||||
|
||||
def _find_turn_start_index(self, messages: list[Message], entry_index: int) -> int:
|
||||
"""Find user message that starts the turn. Returns -1 if not found."""
|
||||
if not messages or entry_index < 0 or entry_index >= len(messages):
|
||||
return -1
|
||||
|
||||
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],
|
||||
token_count: int,
|
||||
threshold: int,
|
||||
) -> CutPointResult:
|
||||
"""
|
||||
Find cut point with split turn detection.
|
||||
|
||||
Split turn: User → Assistant → [CUT] → Assistant → User
|
||||
Clean cut: User → [CUT] → Assistant → User
|
||||
"""
|
||||
if not messages:
|
||||
return CutPointResult(
|
||||
needs_compaction=False,
|
||||
token_count=token_count,
|
||||
threshold=threshold,
|
||||
)
|
||||
|
||||
accumulated_tokens = 0
|
||||
cut_index = 0
|
||||
|
||||
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
|
||||
logger.debug(f"Cut point at index {cut_index}, {accumulated_tokens} tokens")
|
||||
break
|
||||
|
||||
if cut_index == 0:
|
||||
return CutPointResult(
|
||||
left_messages=messages,
|
||||
needs_compaction=True,
|
||||
token_count=token_count,
|
||||
threshold=threshold,
|
||||
accumulated_tokens=accumulated_tokens,
|
||||
)
|
||||
|
||||
cut_message = messages[cut_index]
|
||||
is_user_cut = self._is_user_message(cut_message)
|
||||
|
||||
if is_user_cut:
|
||||
return CutPointResult(
|
||||
messages_to_summarize=messages[:cut_index],
|
||||
left_messages=messages[cut_index:],
|
||||
cut_index=cut_index,
|
||||
needs_compaction=True,
|
||||
token_count=token_count,
|
||||
threshold=threshold,
|
||||
accumulated_tokens=accumulated_tokens,
|
||||
)
|
||||
|
||||
turn_start_index = self._find_turn_start_index(messages, cut_index)
|
||||
|
||||
if turn_start_index == -1:
|
||||
logger.warning("Split turn detected but no turn start found, treating as clean cut")
|
||||
return CutPointResult(
|
||||
messages_to_summarize=messages[:cut_index],
|
||||
left_messages=messages[cut_index:],
|
||||
cut_index=cut_index,
|
||||
needs_compaction=True,
|
||||
token_count=token_count,
|
||||
threshold=threshold,
|
||||
accumulated_tokens=accumulated_tokens,
|
||||
)
|
||||
|
||||
return CutPointResult(
|
||||
messages_to_summarize=messages[:turn_start_index],
|
||||
turn_prefix_messages=messages[turn_start_index:cut_index],
|
||||
left_messages=messages[cut_index:],
|
||||
is_split_turn=True,
|
||||
cut_index=cut_index,
|
||||
needs_compaction=True,
|
||||
token_count=token_count,
|
||||
threshold=threshold,
|
||||
accumulated_tokens=accumulated_tokens,
|
||||
)
|
||||
|
||||
async def execute(self):
|
||||
"""
|
||||
Execute context check and find cut point.
|
||||
|
||||
Returns:
|
||||
dict: CutPointResult.model_dump() with cut point information.
|
||||
"""
|
||||
messages = self.context.messages
|
||||
normalized_messages = self._normalize_messages(messages)
|
||||
token_count: int = self.token_counter.count_token(normalized_messages)
|
||||
threshold = self.context_window_tokens - self.reserve_tokens
|
||||
|
||||
needs_compaction = token_count >= threshold
|
||||
|
||||
if not needs_compaction:
|
||||
logger.info(f"Token count {token_count} below threshold ({threshold}), no compaction needed")
|
||||
cut_result = CutPointResult(
|
||||
needs_compaction=False,
|
||||
token_count=token_count,
|
||||
threshold=threshold,
|
||||
left_messages=normalized_messages,
|
||||
)
|
||||
return cut_result.model_dump()
|
||||
|
||||
logger.info(f"Compaction needed, token count: {token_count}, threshold: {threshold}")
|
||||
cut_result = self._find_cut_point(normalized_messages, token_count, threshold)
|
||||
return cut_result.model_dump()
|
||||
|
|
@ -4,7 +4,7 @@ import datetime
|
|||
|
||||
from loguru import logger
|
||||
|
||||
from ...core.enumeration import Role, MemoryType
|
||||
from ...core.enumeration import Role
|
||||
from ...core.op import BaseReact
|
||||
from ...core.schema import Message
|
||||
|
||||
|
|
@ -12,14 +12,7 @@ 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 = "memory",
|
||||
version: str = "default",
|
||||
**kwargs,
|
||||
):
|
||||
def __init__(self, memory_dir: str = "memory", version: str = "default", **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.memory_dir: str = memory_dir
|
||||
self.version: str = version
|
||||
|
|
@ -40,16 +33,18 @@ class FsSummarizer(BaseReact):
|
|||
),
|
||||
)
|
||||
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",
|
||||
date=date_str,
|
||||
memory_dir=self.memory_dir,
|
||||
messages.extend(
|
||||
[
|
||||
Message(role=Role.SYSTEM, content=self.get_prompt("system_prompt")),
|
||||
Message(
|
||||
role=Role.USER,
|
||||
content=self.prompt_format(
|
||||
"user_message",
|
||||
date=date_str,
|
||||
memory_dir=self.memory_dir,
|
||||
),
|
||||
),
|
||||
),
|
||||
],
|
||||
)
|
||||
return messages
|
||||
|
||||
|
|
|
|||
|
|
@ -17,7 +17,23 @@ user_message_v2: |
|
|||
1. Check if {memory_dir}/ exists; if not, create it via bash
|
||||
2. Check if {memory_dir}/YYYY-MM-DD.md exists (use actual date)
|
||||
3. If file is NEW: Write memories directly (be concise)
|
||||
4. If file EXISTS: Read it first, then UPDATE with new memories (keep concise, merge/deduplicate)
|
||||
4. If file EXISTS:
|
||||
a) Read the existing file content
|
||||
b) Compare conversation history with existing content
|
||||
c) Identify NEW/UPDATED information not yet captured
|
||||
d) Use edit_tool to add/update only the new information (preserve existing content)
|
||||
e) If conversation contains NO new information, skip writing
|
||||
5. If NO valuable information to store: Reply with reason and [SILENT]
|
||||
|
||||
IMPORTANT for updates:
|
||||
- Only add information that is NOT already in the file
|
||||
- Preserve all existing entries
|
||||
- Merge duplicate information intelligently
|
||||
- Use edit_tool for surgical updates, not write_tool (which overwrites)
|
||||
|
||||
Example of what counts as NEW information:
|
||||
- Existing: "Alice: Software engineer"
|
||||
- Conversation: "Alice loves Python and AI projects"
|
||||
- Action: ADD "Enjoys Python programming and AI project work" to Alice's entry
|
||||
|
||||
Store durable memories. Keep entries concise and well-organized.
|
||||
|
|
@ -21,6 +21,7 @@ llms:
|
|||
default:
|
||||
backend: openai
|
||||
model_name: qwen3-30b-a3b-instruct-2507
|
||||
# model_name: qwen3-30b-a3b-thinking-2507
|
||||
request_interval: 1
|
||||
# temperature: 0.0001
|
||||
|
||||
|
|
|
|||
40
reme/config/fs.yaml
Normal file
40
reme/config/fs.yaml
Normal file
|
|
@ -0,0 +1,40 @@
|
|||
backend: cmd
|
||||
|
||||
llms:
|
||||
default:
|
||||
backend: openai
|
||||
model_name: qwen3-30b-a3b-instruct-2507
|
||||
# model_name: qwen3-30b-a3b-thinking-2507
|
||||
request_interval: 1
|
||||
# temperature: 0.0001
|
||||
|
||||
embedding_models:
|
||||
default:
|
||||
backend: openai
|
||||
model_name: text-embedding-v4
|
||||
dimensions: 1024
|
||||
|
||||
memory_stores:
|
||||
default:
|
||||
backend: sqlite
|
||||
store_name: test_hybrid
|
||||
embedding_model: default
|
||||
fts_enabled: true
|
||||
snippet_max_chars: 700
|
||||
|
||||
file_watchers:
|
||||
default:
|
||||
backend: full
|
||||
watch_paths: [".reme", ".reme/memory"]
|
||||
suffix_filters: [".md"]
|
||||
recursive: false
|
||||
scan_on_start: true
|
||||
|
||||
token_counters:
|
||||
default:
|
||||
backend: base
|
||||
|
||||
hf:
|
||||
backend: hf
|
||||
model_name: Qwen/Qwen3-Coder-30B-A3B-Instruct
|
||||
use_mirror: true
|
||||
|
|
@ -24,7 +24,9 @@ class Application:
|
|||
llm_api_base: str | None = None,
|
||||
embedding_api_key: str | None = None,
|
||||
embedding_api_base: str | None = None,
|
||||
config_path: str | None = None,
|
||||
enable_logo: bool = True,
|
||||
log_to_console: bool = True,
|
||||
parser: type[PydanticConfigParser] | None = None,
|
||||
default_llm_config: dict | None = None,
|
||||
default_embedding_model_config: dict | None = None,
|
||||
|
|
@ -42,8 +44,9 @@ class Application:
|
|||
embedding_api_base=embedding_api_base,
|
||||
service_config=None,
|
||||
parser=parser,
|
||||
config_path=None,
|
||||
config_path=config_path,
|
||||
enable_logo=enable_logo,
|
||||
log_to_console=log_to_console,
|
||||
default_llm_config=default_llm_config,
|
||||
default_embedding_model_config=default_embedding_model_config,
|
||||
default_vector_store_config=default_vector_store_config,
|
||||
|
|
@ -136,7 +139,7 @@ class Application:
|
|||
stream_queue=stream_queue,
|
||||
task=task,
|
||||
task_name=name,
|
||||
as_bytes=False,
|
||||
output_format="str",
|
||||
):
|
||||
yield chunk
|
||||
|
||||
|
|
|
|||
|
|
@ -36,6 +36,7 @@ class ServiceContext(BaseContext):
|
|||
parser: type[PydanticConfigParser] | None = None,
|
||||
config_path: str | None = None,
|
||||
enable_logo: bool = True,
|
||||
log_to_console: bool = True,
|
||||
default_llm_config: dict | None = None,
|
||||
default_embedding_model_config: dict | None = None,
|
||||
default_vector_store_config: dict | None = None,
|
||||
|
|
@ -74,13 +75,12 @@ class ServiceContext(BaseContext):
|
|||
if default_file_watcher_config:
|
||||
self._update_section_config(kwargs, "file_watchers", **default_file_watcher_config)
|
||||
kwargs["enable_logo"] = enable_logo
|
||||
kwargs["log_to_console"] = log_to_console
|
||||
logger.info(f"update with args: {input_args} kwargs: {kwargs}")
|
||||
service_config = parser.parse_args(*input_args, **kwargs)
|
||||
|
||||
self.service_config: ServiceConfig = service_config
|
||||
|
||||
if self.service_config.init_logger:
|
||||
init_logger()
|
||||
init_logger(log_to_console=self.service_config.log_to_console)
|
||||
|
||||
if self.service_config.enable_logo:
|
||||
print_logo(service_config=self.service_config)
|
||||
|
|
@ -162,35 +162,25 @@ class ServiceContext(BaseContext):
|
|||
)
|
||||
|
||||
for name, config in self.service_config.vector_stores.items():
|
||||
self.vector_stores[name] = R.vector_stores[config.backend](
|
||||
collection_name=config.collection_name,
|
||||
embedding_model=self.embedding_models[config.embedding_model],
|
||||
thread_pool=self.thread_pool,
|
||||
**config.model_extra,
|
||||
)
|
||||
# Extract config dict and replace special fields with actual instances
|
||||
config_dict = config.model_dump(exclude={"backend", "embedding_model"})
|
||||
config_dict["embedding_model"] = self.embedding_models[config.embedding_model]
|
||||
config_dict["thread_pool"] = self.thread_pool
|
||||
self.vector_stores[name] = R.vector_stores[config.backend](**config_dict)
|
||||
await self.vector_stores[name].create_collection(config.collection_name)
|
||||
|
||||
for name, config in self.service_config.memory_stores.items():
|
||||
self.memory_stores[name] = R.memory_stores[config.backend](
|
||||
store_name=config.store_name,
|
||||
embedding_model=self.embedding_models[config.embedding_model],
|
||||
fts_enabled=config.fts_enabled,
|
||||
snippet_max_chars=config.snippet_max_chars,
|
||||
**config.model_extra,
|
||||
)
|
||||
# Extract config dict and replace embedding_model string with actual instance
|
||||
config_dict = config.model_dump(exclude={"backend", "embedding_model"})
|
||||
config_dict["embedding_model"] = self.embedding_models[config.embedding_model]
|
||||
self.memory_stores[name] = R.memory_stores[config.backend](**config_dict)
|
||||
await self.memory_stores[name].start()
|
||||
|
||||
for name, config in self.service_config.file_watchers.items():
|
||||
self.file_watchers[name] = R.file_watchers[config.backend](
|
||||
watch_paths=config.watch_paths,
|
||||
suffix_filters=config.suffix_filters,
|
||||
recursive=config.recursive,
|
||||
debounce=config.debounce,
|
||||
chunk_tokens=config.chunk_tokens,
|
||||
chunk_overlap=config.chunk_overlap,
|
||||
memory_store=self.memory_stores[config.memory_store],
|
||||
**config.model_extra,
|
||||
)
|
||||
# Extract config dict and replace memory_store string with actual instance
|
||||
config_dict = config.model_dump(exclude={"backend", "memory_store"})
|
||||
config_dict["memory_store"] = self.memory_stores[config.memory_store]
|
||||
self.file_watchers[name] = R.file_watchers[config.backend](**config_dict)
|
||||
await self.file_watchers[name].start()
|
||||
|
||||
if self.service_config.mcp_servers:
|
||||
|
|
@ -229,6 +219,9 @@ class ServiceContext(BaseContext):
|
|||
for _, memory_store in self.memory_stores.items():
|
||||
await memory_store.close()
|
||||
|
||||
for _, file_watcher in self.file_watchers.items():
|
||||
await file_watcher.close()
|
||||
|
||||
for _, llm in self.llms.items():
|
||||
await llm.close()
|
||||
|
||||
|
|
|
|||
|
|
@ -21,5 +21,11 @@ class ChunkEnum(str, Enum):
|
|||
# Error messages or exception details
|
||||
ERROR = "error"
|
||||
|
||||
# Signal indicating the start of a new ReAct step
|
||||
STEP_START = "step_start"
|
||||
|
||||
# Tool execution result
|
||||
TOOL_RESULT = "tool_result"
|
||||
|
||||
# Final signal indicating the completion of the stream
|
||||
DONE = "done"
|
||||
|
|
|
|||
|
|
@ -6,11 +6,13 @@ that monitor file system changes and trigger callbacks.
|
|||
|
||||
import asyncio
|
||||
from collections.abc import Coroutine
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable
|
||||
|
||||
from loguru import logger
|
||||
from watchfiles import awatch, Change
|
||||
|
||||
from ..enumeration import MemorySource
|
||||
from ..memory_store import BaseMemoryStore
|
||||
|
||||
|
||||
|
|
@ -32,10 +34,24 @@ class BaseFileWatcher:
|
|||
chunk_overlap: int = 80,
|
||||
memory_store: BaseMemoryStore | None = None,
|
||||
callback: Callable[[set[tuple[Change, str]]], None | Coroutine[Any, Any, None]] | None = None,
|
||||
scan_on_start: bool = False,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Initialize the file watcher"""
|
||||
Initialize the file watcher
|
||||
|
||||
Args:
|
||||
watch_paths: Paths to watch for changes
|
||||
suffix_filters: File suffix filters (e.g., ['.py', '.txt'])
|
||||
recursive: Whether to watch directories recursively
|
||||
debounce: Debounce time in milliseconds
|
||||
chunk_tokens: Token size for chunking
|
||||
chunk_overlap: Overlap size for chunks
|
||||
memory_store: Memory store instance
|
||||
callback: Callback function for changes
|
||||
scan_on_start: If True, scan existing files on start and trigger on_changes with Change.added
|
||||
**kwargs: Additional keyword arguments
|
||||
"""
|
||||
self.watch_paths: list[str] = [watch_paths] if isinstance(watch_paths, str) else watch_paths
|
||||
self.suffix_filters: list[str] = suffix_filters or []
|
||||
self.recursive: bool = recursive
|
||||
|
|
@ -44,6 +60,7 @@ class BaseFileWatcher:
|
|||
self.chunk_overlap: int = chunk_overlap
|
||||
self.memory_store: BaseMemoryStore = memory_store
|
||||
self.callback = callback
|
||||
self.scan_on_start: bool = scan_on_start
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
self._stop_event = asyncio.Event()
|
||||
|
|
@ -56,6 +73,11 @@ class BaseFileWatcher:
|
|||
return
|
||||
|
||||
self._running = True
|
||||
|
||||
# Scan existing files if requested
|
||||
if self.scan_on_start:
|
||||
await self._scan_existing_files()
|
||||
|
||||
self._watch_task = asyncio.create_task(self._watch_loop())
|
||||
logger.info(f"Started watching: {self.watch_paths}")
|
||||
|
||||
|
|
@ -83,23 +105,69 @@ class BaseFileWatcher:
|
|||
|
||||
return False
|
||||
|
||||
async def _scan_existing_files(self):
|
||||
"""Scan existing files matching watch criteria and trigger on_changes with Change.added"""
|
||||
existing_files: set[tuple[Change, str]] = set()
|
||||
|
||||
for watch_path_str in self.watch_paths:
|
||||
watch_path = Path(watch_path_str)
|
||||
|
||||
if not watch_path.exists():
|
||||
logger.warning(f"Watch path does not exist: {watch_path}")
|
||||
continue
|
||||
|
||||
if watch_path.is_file():
|
||||
# Single file
|
||||
if self.watch_filter(Change.added, str(watch_path)):
|
||||
existing_files.add((Change.added, str(watch_path)))
|
||||
elif watch_path.is_dir():
|
||||
# Directory
|
||||
if self.recursive:
|
||||
# Recursive scan
|
||||
for file_path in watch_path.rglob("*"):
|
||||
if file_path.is_file() and self.watch_filter(Change.added, str(file_path)):
|
||||
existing_files.add((Change.added, str(file_path)))
|
||||
else:
|
||||
# Non-recursive scan (only immediate children)
|
||||
for file_path in watch_path.iterdir():
|
||||
if file_path.is_file() and self.watch_filter(Change.added, str(file_path)):
|
||||
existing_files.add((Change.added, str(file_path)))
|
||||
|
||||
if existing_files:
|
||||
logger.info(f"Found {len(existing_files)} existing files to process")
|
||||
await self.on_changes(existing_files)
|
||||
else:
|
||||
logger.info("No existing files found matching watch criteria")
|
||||
|
||||
files: list[str] = await self.memory_store.list_files(MemorySource.MEMORY)
|
||||
for file_path in files:
|
||||
chunks = await self.memory_store.get_file_chunks(file_path, MemorySource.MEMORY)
|
||||
logger.info(f"Found existing file: {file_path}, {len(chunks)} chunks")
|
||||
|
||||
async def _watch_loop(self):
|
||||
"""Core monitoring loop"""
|
||||
if not self.watch_paths:
|
||||
logger.warning("No watch paths specified")
|
||||
return
|
||||
|
||||
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
|
||||
try:
|
||||
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)
|
||||
await self.on_changes(changes)
|
||||
except FileNotFoundError as e:
|
||||
# Watch path was deleted, this is expected during cleanup
|
||||
logger.debug(f"Watch path no longer exists: {e}")
|
||||
except Exception as e:
|
||||
# Log other exceptions but don't crash
|
||||
logger.error(f"Error in watch loop: {e}", exc_info=True)
|
||||
|
||||
async def _on_changes(self, changes: set[tuple[Change, str]]):
|
||||
"""Callback method to handle file changes"""
|
||||
|
|
@ -112,6 +180,7 @@ class BaseFileWatcher:
|
|||
await result
|
||||
else:
|
||||
await self._on_changes(changes)
|
||||
logger.info(f"[{self.__class__.__name__}] on_changes: {changes}")
|
||||
|
||||
def is_running(self) -> bool:
|
||||
"""Check if the watcher is running"""
|
||||
|
|
|
|||
|
|
@ -61,11 +61,16 @@ class FullFileWatcher(BaseFileWatcher):
|
|||
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)
|
||||
logger.info(f"delete_file {file_meta.path}")
|
||||
|
||||
await self.memory_store.upsert_file(file_meta, MemorySource.MEMORY, chunks)
|
||||
logger.info(f"Upserted {file_meta.chunk_count} chunks for {file_meta.path}")
|
||||
|
||||
elif change_type == Change.deleted:
|
||||
await self.memory_store.delete_file(path, MemorySource.MEMORY)
|
||||
logger.info(f"Deleted {path}")
|
||||
|
||||
else:
|
||||
logger.warning(f"Unknown change type: {change_type}")
|
||||
|
|
|
|||
|
|
@ -99,7 +99,6 @@ class BaseLLM(ABC):
|
|||
stream_kwargs: dict,
|
||||
) -> AsyncGenerator[StreamChunk, None]:
|
||||
"""Async generator for streaming response chunks."""
|
||||
raise NotImplementedError
|
||||
|
||||
def _stream_chat_sync(
|
||||
self,
|
||||
|
|
@ -108,7 +107,6 @@ class BaseLLM(ABC):
|
|||
stream_kwargs: dict | None = None,
|
||||
) -> Generator[StreamChunk, None, None]:
|
||||
"""Sync generator for streaming response chunks."""
|
||||
raise NotImplementedError
|
||||
|
||||
async def stream_chat(
|
||||
self,
|
||||
|
|
@ -117,7 +115,7 @@ class BaseLLM(ABC):
|
|||
model_name: str | None = None,
|
||||
**kwargs,
|
||||
) -> AsyncGenerator[StreamChunk, None]:
|
||||
"""Stream chat completions with retries."""
|
||||
"""Stream chat completions with retries and return final message."""
|
||||
if self.request_interval > 0:
|
||||
async with self._request_lock:
|
||||
current_time = time.time()
|
||||
|
|
@ -143,7 +141,8 @@ class BaseLLM(ABC):
|
|||
try:
|
||||
async for chunk in self._stream_chat(messages=messages, tools=tools, stream_kwargs=stream_kwargs):
|
||||
yield chunk
|
||||
return
|
||||
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f"Stream chat error (model={self.model_name}): {e.args}")
|
||||
|
|
@ -152,7 +151,7 @@ class BaseLLM(ABC):
|
|||
if self.raise_exception:
|
||||
raise e
|
||||
yield StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e))
|
||||
return
|
||||
break
|
||||
|
||||
yield StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e))
|
||||
await asyncio.sleep(i + 1)
|
||||
|
|
@ -170,7 +169,7 @@ class BaseLLM(ABC):
|
|||
for i in range(self.max_retries):
|
||||
try:
|
||||
yield from self._stream_chat_sync(messages=messages, tools=tools, stream_kwargs=stream_kwargs)
|
||||
return
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f"Stream chat sync error (model={self.model_name}): {e.args}")
|
||||
|
|
@ -179,7 +178,7 @@ class BaseLLM(ABC):
|
|||
if self.raise_exception:
|
||||
raise e
|
||||
yield StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e))
|
||||
return
|
||||
break
|
||||
|
||||
yield StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e))
|
||||
time.sleep(i + 1)
|
||||
|
|
|
|||
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)
|
||||
|
|
@ -558,6 +558,38 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
finally:
|
||||
cursor.close()
|
||||
|
||||
def _sanitize_fts_query(self, query: str) -> str:
|
||||
"""Sanitize query string for FTS5 search.
|
||||
|
||||
Removes or escapes special characters that have special meaning in FTS5:
|
||||
- * (prefix match)
|
||||
- ? (not used in FTS5, but can cause issues)
|
||||
- " (phrase search, needs escaping)
|
||||
- : (column filter)
|
||||
- ^ (start of line anchor, not standard FTS5)
|
||||
- Other special chars that may interfere
|
||||
|
||||
Args:
|
||||
query: Raw query string
|
||||
|
||||
Returns:
|
||||
Sanitized query string safe for FTS5
|
||||
"""
|
||||
if not query:
|
||||
return ""
|
||||
|
||||
# Remove FTS5 special characters that we don't want users to use
|
||||
# Keep only alphanumeric, spaces, and some safe punctuation
|
||||
special_chars = ["*", "?", ":", "^", "(", ")", "[", "]", "{", "}"]
|
||||
cleaned = query
|
||||
for char in special_chars:
|
||||
cleaned = cleaned.replace(char, " ")
|
||||
|
||||
# Normalize whitespace
|
||||
cleaned = " ".join(cleaned.split())
|
||||
|
||||
return cleaned
|
||||
|
||||
async def keyword_search(
|
||||
self,
|
||||
query: str,
|
||||
|
|
@ -568,14 +600,12 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
if not self.fts_available:
|
||||
return []
|
||||
|
||||
# Build FTS5 query
|
||||
# Split query into tokens and join with OR for better recall
|
||||
# Individual words are automatically stemmed and matched by FTS5
|
||||
cleaned = query.strip()
|
||||
# Sanitize and prepare query
|
||||
cleaned = self._sanitize_fts_query(query)
|
||||
if not cleaned:
|
||||
return []
|
||||
|
||||
# Split into words and escape each
|
||||
# Split into words and escape double quotes for FTS5 phrase matching
|
||||
words = cleaned.split()
|
||||
if not words:
|
||||
return []
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@
|
|||
from .base_op import BaseOp
|
||||
from .base_ray_op import BaseRayOp
|
||||
from .base_react import BaseReact
|
||||
from .base_react_stream import BaseReactStream
|
||||
from .base_tool import BaseTool
|
||||
from .mcp_tool import MCPTool
|
||||
from .parallel_op import ParallelOp
|
||||
|
|
@ -13,6 +14,7 @@ __all__ = [
|
|||
"BaseOp",
|
||||
"BaseRayOp",
|
||||
"BaseReact",
|
||||
"BaseReactStream",
|
||||
"BaseTool",
|
||||
"MCPTool",
|
||||
"ParallelOp",
|
||||
|
|
|
|||
|
|
@ -50,7 +50,7 @@ class BaseOp(metaclass=ABCMeta):
|
|||
sub_ops: dict[str, "BaseOp"] | list["BaseOp"] | Optional["BaseOp"] = None,
|
||||
input_mapping: dict[str, str] | None = None,
|
||||
output_mapping: dict[str, str] | None = None,
|
||||
enable_sync_thread_pool: bool = True,
|
||||
enable_parallel: bool = False,
|
||||
max_retries: int = 1,
|
||||
raise_exception: bool = False,
|
||||
**kwargs,
|
||||
|
|
@ -76,7 +76,7 @@ class BaseOp(metaclass=ABCMeta):
|
|||
|
||||
self.input_mapping = input_mapping
|
||||
self.output_mapping = output_mapping
|
||||
self.enable_sync_thread_pool = enable_sync_thread_pool
|
||||
self.enable_parallel = enable_parallel # Control whether to execute tasks in parallel
|
||||
self.max_retries = max(1, max_retries)
|
||||
self.raise_exception = raise_exception
|
||||
self.op_params = kwargs
|
||||
|
|
@ -233,7 +233,7 @@ class BaseOp(metaclass=ABCMeta):
|
|||
|
||||
def submit_sync_task(self, fn: Callable, *args, **kwargs) -> "BaseOp":
|
||||
"""Submit a task to the thread pool or local queue."""
|
||||
if self.enable_sync_thread_pool:
|
||||
if self.enable_parallel:
|
||||
task = self.service_context.thread_pool.submit(fn, *args, **kwargs)
|
||||
else:
|
||||
task = (fn, args, kwargs)
|
||||
|
|
@ -250,7 +250,7 @@ class BaseOp(metaclass=ABCMeta):
|
|||
"""Wait for all pending sync tasks and return flattened results."""
|
||||
results = []
|
||||
for task in tqdm(self._pending_tasks, desc=task_desc or self.name):
|
||||
if self.enable_sync_thread_pool:
|
||||
if self.enable_parallel:
|
||||
result = task.result()
|
||||
else:
|
||||
result = task[0](*task[1], **task[2])
|
||||
|
|
@ -264,7 +264,20 @@ class BaseOp(metaclass=ABCMeta):
|
|||
|
||||
async def join_async_tasks(self, return_exceptions: bool = True) -> list:
|
||||
"""Wait for all pending async tasks and aggregate results."""
|
||||
raw_results = await asyncio.gather(*self._pending_tasks, return_exceptions=return_exceptions)
|
||||
if self.enable_parallel:
|
||||
raw_results = await asyncio.gather(*self._pending_tasks, return_exceptions=return_exceptions)
|
||||
else:
|
||||
raw_results = []
|
||||
for task in self._pending_tasks:
|
||||
try:
|
||||
result = await task
|
||||
raw_results.append(result)
|
||||
except Exception as e:
|
||||
if return_exceptions:
|
||||
raw_results.append(e)
|
||||
else:
|
||||
raise
|
||||
|
||||
results = []
|
||||
for result in raw_results:
|
||||
if isinstance(result, Exception):
|
||||
|
|
|
|||
|
|
@ -71,7 +71,7 @@ class BaseReact(BaseOp):
|
|||
assistant_message: Message = await self.llm.chat(messages=messages, tools=tool_calls, **kwargs)
|
||||
messages.append(assistant_message)
|
||||
assistant_content: str = assistant_message.simple_dump(as_dict=False)
|
||||
logger.info(f"[{self.__class__.__name__} {stage or ''} step{step + 1}] assistant={assistant_content}")
|
||||
logger.info(f"[{self.__class__.__name__} {stage or ''} step{step}] assistant={assistant_content}")
|
||||
|
||||
# Determine if tools should be called
|
||||
should_act = bool(assistant_message.tool_calls)
|
||||
|
|
@ -95,7 +95,7 @@ class BaseReact(BaseOp):
|
|||
# Create tool name to tool instance mapping
|
||||
tool_dict = {t.tool_call.name: t for t in tools}
|
||||
for j, tool_call in enumerate(assistant_message.tool_calls):
|
||||
prefix: str = f"[{self.__class__.__name__} {stage or ''} step{step + 1}.{j}]"
|
||||
prefix: str = f"[{self.__class__.__name__} {stage or ''} step{step}.{j}]"
|
||||
if tool_call.name not in tool_dict:
|
||||
logger.warning(f"{prefix} unknown tool_call={tool_call.name}")
|
||||
continue
|
||||
|
|
@ -125,7 +125,7 @@ class BaseReact(BaseOp):
|
|||
tool_call_id=tool.tool_call.id,
|
||||
),
|
||||
)
|
||||
prefix: str = f"[{self.__class__.__name__} {stage or ''} step{step + 1}.{j}]"
|
||||
prefix: str = f"[{self.__class__.__name__} {stage or ''} step{step}.{j}]"
|
||||
logger.info(f"{prefix} join tool={tool.name} result={tool.response.answer}")
|
||||
return tool_list, tool_messages
|
||||
|
||||
|
|
@ -153,7 +153,7 @@ class BaseReact(BaseOp):
|
|||
"""Execute the ReAct agent and return final results."""
|
||||
# Log available tools
|
||||
for i, tool in enumerate(self.tools):
|
||||
logger.info(f"[{self.__class__.__name__}] tool_call[{i}]={tool.tool_call.simple_input_dump(as_dict=False)}")
|
||||
logger.info(f"[{self.__class__.__name__}] {i}.tool_call={tool.tool_call.simple_input_dump(as_dict=False)}")
|
||||
|
||||
# Build and log initial messages
|
||||
messages = await self.build_messages()
|
||||
|
|
@ -163,8 +163,13 @@ class BaseReact(BaseOp):
|
|||
|
||||
# Run ReAct loop
|
||||
t_tools, messages, success = await self.react(messages, self.tools)
|
||||
|
||||
# Get the last assistant message as the final answer
|
||||
assistant_messages = [m for m in messages if m.role == Role.ASSISTANT]
|
||||
answer = assistant_messages[-1].content if assistant_messages else ""
|
||||
|
||||
return {
|
||||
"answer": messages[-1].content if success else "",
|
||||
"answer": answer,
|
||||
"success": success,
|
||||
"messages": messages,
|
||||
"tools": t_tools,
|
||||
|
|
|
|||
249
reme/core/op/base_react_stream.py
Normal file
249
reme/core/op/base_react_stream.py
Normal file
|
|
@ -0,0 +1,249 @@
|
|||
"""Base memory agent for handling memory operations with tool-based reasoning."""
|
||||
|
||||
import asyncio
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from ..enumeration import Role, ChunkEnum
|
||||
from ..op import BaseOp
|
||||
from ..schema import Message, StreamChunk
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from . import BaseTool
|
||||
|
||||
|
||||
class BaseReactStream(BaseOp):
|
||||
"""ReAct agent that performs reasoning and acting cycles with tools."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tools: list["BaseTool"],
|
||||
tool_call_interval: float = 0,
|
||||
max_steps: int = 10,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize ReAct agent with tools and execution parameters."""
|
||||
kwargs["sub_ops"] = tools or []
|
||||
super().__init__(**kwargs)
|
||||
# Filter only BaseTool instances from sub_ops
|
||||
from . import BaseTool
|
||||
|
||||
self.sub_ops: list[BaseTool] = [t for t in self.sub_ops if isinstance(t, BaseTool)]
|
||||
self.tool_call_interval: float = tool_call_interval
|
||||
self.max_steps: int = max_steps
|
||||
|
||||
@property
|
||||
def tools(self) -> list["BaseTool"]:
|
||||
"""Return available tools for the agent."""
|
||||
return self.sub_ops
|
||||
|
||||
def pop_tool(self, name: str) -> "BaseTool | None":
|
||||
"""Remove and return a tool from self.tools by name."""
|
||||
for i, tool in enumerate(self.sub_ops):
|
||||
if tool.tool_call.name == name:
|
||||
return self.sub_ops.pop(i)
|
||||
return None
|
||||
|
||||
async def build_messages(self) -> list[Message]:
|
||||
"""Build initial message list from context query or messages."""
|
||||
if self.context.get("query"):
|
||||
messages = [Message(role=Role.USER, content=self.context.query)]
|
||||
elif self.context.get("messages"):
|
||||
messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages]
|
||||
else:
|
||||
raise ValueError("input must have either `query` or `messages`")
|
||||
return messages
|
||||
|
||||
async def _reasoning_step(
|
||||
self,
|
||||
messages: list[Message],
|
||||
tools: list["BaseTool"],
|
||||
step: int,
|
||||
stage: str = "",
|
||||
**kwargs,
|
||||
) -> tuple[Message, bool]:
|
||||
"""Execute one reasoning step where LLM decides whether to use tools."""
|
||||
tool_calls = [t.tool_call for t in tools]
|
||||
|
||||
start_chunk = StreamChunk(chunk_type=ChunkEnum.STEP_START, metadata={"step": step, "stage": stage})
|
||||
await self.context.add_stream_chunk(start_chunk)
|
||||
|
||||
# State for accumulating message content from stream
|
||||
state = {
|
||||
"reasoning_content": "",
|
||||
"content": "",
|
||||
"tool_calls": [],
|
||||
}
|
||||
|
||||
async for stream_chunk in self.llm.stream_chat(messages=messages, tools=tool_calls, **kwargs): # noqa
|
||||
if stream_chunk.chunk_type in [ChunkEnum.ANSWER, ChunkEnum.THINK, ChunkEnum.ERROR]:
|
||||
await self.context.add_stream_chunk(stream_chunk)
|
||||
|
||||
# Accumulate content based on chunk type
|
||||
if stream_chunk.chunk_type is ChunkEnum.THINK:
|
||||
state["reasoning_content"] += stream_chunk.chunk
|
||||
|
||||
elif stream_chunk.chunk_type is ChunkEnum.ANSWER:
|
||||
state["content"] += stream_chunk.chunk
|
||||
|
||||
elif stream_chunk.chunk_type is ChunkEnum.TOOL:
|
||||
state["tool_calls"].append(stream_chunk.chunk)
|
||||
|
||||
# Build the final assistant message from accumulated state
|
||||
assistant_message = Message(role=Role.ASSISTANT, **state)
|
||||
messages.append(assistant_message)
|
||||
logger.info(
|
||||
f"[{self.__class__.__name__} {stage or ''} step{step}] "
|
||||
f"assistant={assistant_message.simple_dump(as_dict=False)}",
|
||||
)
|
||||
|
||||
should_act = bool(assistant_message.tool_calls)
|
||||
return assistant_message, should_act
|
||||
|
||||
async def _acting_step(
|
||||
self,
|
||||
assistant_message: Message,
|
||||
tools: list["BaseTool"],
|
||||
step: int,
|
||||
stage: str = "",
|
||||
**kwargs,
|
||||
) -> tuple[list["BaseTool"], list[Message]]:
|
||||
"""Execute tool calls serially and collect results with streaming output."""
|
||||
tool_list: list["BaseTool"] = []
|
||||
tool_messages: list[Message] = []
|
||||
|
||||
if not assistant_message.tool_calls:
|
||||
return tool_list, tool_messages
|
||||
|
||||
# Create tool name to tool instance mapping
|
||||
tool_dict = {t.tool_call.name: t for t in tools}
|
||||
|
||||
# Execute tools serially for better streaming experience
|
||||
for j, tool_call in enumerate(assistant_message.tool_calls):
|
||||
prefix: str = f"[{self.__class__.__name__} {stage or ''} step{step}.{j}]"
|
||||
if tool_call.name not in tool_dict:
|
||||
logger.warning(f"{prefix} unknown tool_call={tool_call.name}")
|
||||
# Emit error chunk for unknown tool
|
||||
await self.context.add_stream_chunk(
|
||||
StreamChunk(
|
||||
chunk_type=ChunkEnum.ERROR,
|
||||
chunk=f"Unknown tool: {tool_call.name}",
|
||||
metadata={"step": step, "tool_index": j, "tool_name": tool_call.name},
|
||||
),
|
||||
)
|
||||
continue
|
||||
|
||||
logger.info(f"{prefix} submit tool_call[{tool_call.name}] arguments={tool_call.arguments}")
|
||||
|
||||
# Emit tool execution start signal
|
||||
await self.context.add_stream_chunk(
|
||||
StreamChunk(
|
||||
chunk_type=ChunkEnum.TOOL,
|
||||
chunk=f"Executing tool: {tool_call.name} {tool_call.arguments}",
|
||||
metadata={
|
||||
"step": step,
|
||||
"tool_index": j,
|
||||
"tool_name": tool_call.name,
|
||||
"arguments": tool_call.arguments,
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
# Create independent tool copy with unique ID
|
||||
tool_copy: BaseTool = tool_dict[tool_call.name].copy()
|
||||
tool_copy.tool_call.id = tool_call.id
|
||||
tool_list.append(tool_copy)
|
||||
|
||||
# Create isolated kwargs for each tool call to avoid parameter conflicts
|
||||
tool_kwargs = {**kwargs, **tool_call.argument_dict}
|
||||
|
||||
# Execute tool serially (wait for completion before next tool)
|
||||
await tool_copy.call(service_context=self.service_context, **tool_kwargs)
|
||||
|
||||
# Get tool result immediately after execution
|
||||
tool_result = tool_copy.response.answer
|
||||
tool_messages.append(
|
||||
Message(
|
||||
role=Role.TOOL,
|
||||
content=tool_result,
|
||||
tool_call_id=tool_copy.tool_call.id,
|
||||
),
|
||||
)
|
||||
logger.info(f"{prefix} tool={tool_copy.name} result={tool_result}")
|
||||
|
||||
await self.context.add_stream_chunk(
|
||||
StreamChunk(
|
||||
chunk_type=ChunkEnum.TOOL_RESULT,
|
||||
chunk=tool_result,
|
||||
metadata={
|
||||
"step": step,
|
||||
"tool_index": j,
|
||||
"tool_name": tool_copy.name,
|
||||
"tool_call_id": tool_copy.tool_call.id,
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
# Optional interval between tool calls
|
||||
if self.tool_call_interval > 0 and j < len(assistant_message.tool_calls) - 1:
|
||||
await asyncio.sleep(self.tool_call_interval)
|
||||
|
||||
return tool_list, tool_messages
|
||||
|
||||
async def react(self, messages: list[Message], tools: list["BaseTool"], stage: str = ""):
|
||||
"""Run ReAct loop alternating between reasoning and acting until completion."""
|
||||
success: bool = False
|
||||
used_tools: list[BaseTool] = []
|
||||
for step in range(self.max_steps):
|
||||
# Reasoning: LLM decides next action
|
||||
assistant_message, should_act = await self._reasoning_step(messages, tools, step=step, stage=stage)
|
||||
|
||||
if not should_act:
|
||||
# No tools requested, task complete
|
||||
success = True
|
||||
break
|
||||
|
||||
# Acting: execute tools and collect results
|
||||
t_tools, tool_messages = await self._acting_step(assistant_message, tools, step=step, stage=stage)
|
||||
used_tools.extend(t_tools)
|
||||
messages.extend(tool_messages)
|
||||
|
||||
return used_tools, messages, success
|
||||
|
||||
async def execute(self):
|
||||
"""Execute the ReAct agent with streaming output and return final results."""
|
||||
for i, tool in enumerate(self.tools):
|
||||
logger.info(f"[{self.__class__.__name__}] {i}.tool_call={tool.tool_call.simple_input_dump(as_dict=False)}")
|
||||
|
||||
# Build and log initial messages
|
||||
messages = await self.build_messages()
|
||||
for i, message in enumerate(messages):
|
||||
role = message.name or message.role
|
||||
logger.info(f"[{self.__class__.__name__}] role={role} {message.simple_dump(as_dict=False)}")
|
||||
|
||||
# Run ReAct loop with streaming
|
||||
t_tools, messages, success = await self.react(messages, self.tools)
|
||||
|
||||
# Emit final done signal
|
||||
await self.context.add_stream_chunk(
|
||||
StreamChunk(
|
||||
chunk_type=ChunkEnum.DONE,
|
||||
chunk="",
|
||||
metadata={
|
||||
"success": success,
|
||||
"total_steps": len(t_tools),
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
# Get the last assistant message as the final answer
|
||||
assistant_messages = [m for m in messages if m.role == Role.ASSISTANT]
|
||||
answer = assistant_messages[-1].content if assistant_messages else ""
|
||||
|
||||
return {
|
||||
"answer": answer,
|
||||
"success": success,
|
||||
"messages": messages,
|
||||
"tools": t_tools,
|
||||
}
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
"""schema"""
|
||||
|
||||
from .compaction_result import CutPointResult
|
||||
from .cut_point_result import CutPointResult
|
||||
from .file_metadata import FileMetadata
|
||||
from .memory_chunk import MemoryChunk
|
||||
from .memory_node import MemoryNode
|
||||
|
|
@ -25,9 +25,9 @@ from .truncation_result import TruncationResult
|
|||
from .vector_node import VectorNode
|
||||
|
||||
__all__ = [
|
||||
"CutPointResult",
|
||||
"CmdConfig",
|
||||
"ContentBlock",
|
||||
"CutPointResult",
|
||||
"EmbeddingModelConfig",
|
||||
"FileMetadata",
|
||||
"FlowConfig",
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
"""Compaction result schemas for context window management."""
|
||||
"""Cut point result schemas for context window management."""
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
|
@ -11,5 +11,11 @@ class CutPointResult(BaseModel):
|
|||
messages_to_summarize: list[Message] = Field(default_factory=list, description="Complete turns before cut point")
|
||||
turn_prefix_messages: list[Message] = Field(default_factory=list, description="Turn prefix if split turn")
|
||||
left_messages: list[Message] = Field(default_factory=list, description="Messages to keep from cut point onwards")
|
||||
|
||||
is_split_turn: bool = Field(default=False, description="Whether cut point is mid-turn")
|
||||
cut_index: int = Field(default=0, description="Index of cut point in original message list")
|
||||
|
||||
needs_compaction: bool = Field(default=False, description="Whether compaction is actually needed")
|
||||
token_count: int = Field(default=0, description="Total token count of original messages")
|
||||
threshold: int = Field(default=0, description="Token threshold that triggers compaction")
|
||||
accumulated_tokens: int = Field(default=0, description="Tokens accumulated when finding cut point")
|
||||
|
|
@ -103,6 +103,7 @@ class FileWatcherConfig(BaseModel):
|
|||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
backend: str = Field(default="")
|
||||
watch_paths: list[str] = Field(default_factory=list)
|
||||
suffix_filters: list[str] = Field(default_factory=list)
|
||||
recursive: bool = Field(default=False)
|
||||
|
|
@ -110,6 +111,7 @@ class FileWatcherConfig(BaseModel):
|
|||
chunk_tokens: int = Field(default=400)
|
||||
chunk_overlap: int = Field(default=80)
|
||||
memory_store: str = Field(default="default")
|
||||
scan_on_start: bool = Field(default=True)
|
||||
|
||||
|
||||
class ServiceConfig(BaseModel):
|
||||
|
|
@ -123,7 +125,7 @@ class ServiceConfig(BaseModel):
|
|||
language: str = Field(default="")
|
||||
thread_pool_max_workers: int = Field(default=16)
|
||||
ray_max_workers: int = Field(default=-1)
|
||||
init_logger: bool = Field(default=True)
|
||||
log_to_console: bool = Field(default=True)
|
||||
disabled_flows: list[str] = Field(default_factory=list)
|
||||
enabled_flows: list[str] = Field(default_factory=list)
|
||||
mcp_servers: dict[str, dict] = Field(default_factory=dict)
|
||||
|
|
|
|||
|
|
@ -66,7 +66,7 @@ class HttpService(BaseService):
|
|||
stream_queue=stream_queue,
|
||||
task=task,
|
||||
task_name=tool_call.name,
|
||||
as_bytes=True,
|
||||
output_format="bytes",
|
||||
):
|
||||
yield chunk
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""utils"""
|
||||
|
||||
from .agentscope_utils import convert_dashscope_to_agentscope
|
||||
from .cache_handler import CacheHandler
|
||||
from .case_converter import snake_to_camel, camel_to_snake
|
||||
from .chunking_utils import chunk_markdown
|
||||
|
|
@ -17,6 +18,7 @@ from .singleton import singleton
|
|||
from .time import timer, get_now_time
|
||||
|
||||
__all__ = [
|
||||
"convert_dashscope_to_agentscope",
|
||||
"CacheHandler",
|
||||
"snake_to_camel",
|
||||
"camel_to_snake",
|
||||
|
|
|
|||
389
reme/core/utils/agentscope_utils.py
Normal file
389
reme/core/utils/agentscope_utils.py
Normal file
|
|
@ -0,0 +1,389 @@
|
|||
# -*- coding: utf-8 -*-
|
||||
"""Utilities for converting between DashScope format and AgentScope Msg format."""
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agentscope.message import Msg
|
||||
|
||||
|
||||
class DashScopeToAgentScopeConverter:
|
||||
"""Converter for DashScope format to AgentScope Msg format."""
|
||||
|
||||
def __init__(self, default_name: str = "assistant") -> None:
|
||||
"""Initialize the converter.
|
||||
|
||||
Args:
|
||||
default_name: Default name for assistant messages when not specified.
|
||||
"""
|
||||
self.default_name = default_name
|
||||
|
||||
def convert_message(
|
||||
self,
|
||||
dashscope_msg: dict[str, Any],
|
||||
name: str | None = None,
|
||||
) -> "Msg":
|
||||
"""Convert a single DashScope format message to AgentScope Msg.
|
||||
|
||||
Args:
|
||||
dashscope_msg: DashScope format message dictionary containing
|
||||
'role', 'content', and optionally 'tool_calls', 'reasoning_content',
|
||||
'tool_call_id', 'name'.
|
||||
name: Override name for the message. If None, uses the name from
|
||||
dashscope_msg or default_name.
|
||||
|
||||
Returns:
|
||||
AgentScope Msg object.
|
||||
|
||||
Examples:
|
||||
>>> converter = DashScopeToAgentScopeConverter()
|
||||
>>> # Plain text message
|
||||
>>> ds_msg = {"role": "assistant", "content": "Hello!"}
|
||||
>>> msg = converter.convert_message(ds_msg)
|
||||
>>> # Multimodal message with content blocks
|
||||
>>> ds_msg = {
|
||||
... "role": "user",
|
||||
... "content": [
|
||||
... {"text": "What's in this image?"},
|
||||
... {"image": "https://example.com/image.jpg"}
|
||||
... ]
|
||||
... }
|
||||
>>> msg = converter.convert_message(ds_msg)
|
||||
>>> # Tool call message
|
||||
>>> ds_msg = {
|
||||
... "role": "assistant",
|
||||
... "content": "",
|
||||
... "tool_calls": [{
|
||||
... "id": "call_123",
|
||||
... "type": "function",
|
||||
... "function": {
|
||||
... "name": "get_weather",
|
||||
... "arguments": '{"city": "Beijing"}'
|
||||
... }
|
||||
... }]
|
||||
... }
|
||||
>>> msg = converter.convert_message(ds_msg)
|
||||
"""
|
||||
from agentscope.message import Msg
|
||||
|
||||
role = dashscope_msg.get("role", "assistant")
|
||||
if role not in ["user", "assistant", "system"]:
|
||||
# Map 'tool' role to 'user' since tool results are inputs to assistant
|
||||
if role == "tool":
|
||||
role = "user"
|
||||
else:
|
||||
role = "assistant"
|
||||
|
||||
# Determine message name
|
||||
msg_name = name or dashscope_msg.get("name") or (self.default_name if role == "assistant" else role)
|
||||
|
||||
# Handle tool result messages (role="tool")
|
||||
if dashscope_msg.get("role") == "tool":
|
||||
content_blocks = self._convert_tool_result_to_blocks(dashscope_msg)
|
||||
return Msg(
|
||||
name=msg_name,
|
||||
content=content_blocks,
|
||||
role=role,
|
||||
)
|
||||
|
||||
# Extract content
|
||||
raw_content = dashscope_msg.get("content", "")
|
||||
tool_calls = dashscope_msg.get("tool_calls", [])
|
||||
reasoning_content = dashscope_msg.get("reasoning_content", "")
|
||||
|
||||
# Check if we need ContentBlocks or plain string
|
||||
_ = self._has_multimodal_content(raw_content)
|
||||
has_tools = len(tool_calls) > 0
|
||||
has_reasoning = bool(reasoning_content)
|
||||
|
||||
# If only plain text without tools/reasoning/multimodal, use string content
|
||||
if isinstance(raw_content, str) and not has_tools and not has_reasoning:
|
||||
return Msg(
|
||||
name=msg_name,
|
||||
content=raw_content or "",
|
||||
role=role,
|
||||
)
|
||||
|
||||
# Otherwise, build ContentBlock list
|
||||
content_blocks = []
|
||||
|
||||
# Add reasoning content (thinking block)
|
||||
if has_reasoning:
|
||||
from agentscope.message import ThinkingBlock
|
||||
|
||||
content_blocks.append(
|
||||
ThinkingBlock(
|
||||
type="thinking",
|
||||
thinking=reasoning_content,
|
||||
),
|
||||
)
|
||||
|
||||
# Convert content to blocks
|
||||
content_blocks.extend(self._convert_content_to_blocks(raw_content))
|
||||
|
||||
# Convert tool calls to blocks
|
||||
if has_tools:
|
||||
content_blocks.extend(self._convert_tool_calls_to_blocks(tool_calls))
|
||||
|
||||
# If we have no blocks but expected to have content, return empty string
|
||||
if not content_blocks:
|
||||
return Msg(
|
||||
name=msg_name,
|
||||
content="",
|
||||
role=role,
|
||||
)
|
||||
|
||||
return Msg(
|
||||
name=msg_name,
|
||||
content=content_blocks,
|
||||
role=role,
|
||||
)
|
||||
|
||||
def convert_messages(
|
||||
self,
|
||||
dashscope_msgs: list[dict[str, Any]],
|
||||
) -> list["Msg"]:
|
||||
"""Convert a list of DashScope format messages to AgentScope Msgs.
|
||||
|
||||
Args:
|
||||
dashscope_msgs: List of DashScope format message dictionaries.
|
||||
|
||||
Returns:
|
||||
List of AgentScope Msg objects.
|
||||
"""
|
||||
return [self.convert_message(msg) for msg in dashscope_msgs]
|
||||
|
||||
def _has_multimodal_content(self, content: Any) -> bool:
|
||||
"""Check if content contains multimodal data.
|
||||
|
||||
Args:
|
||||
content: Content to check (string or list of content blocks).
|
||||
|
||||
Returns:
|
||||
True if content contains images, audio, or video.
|
||||
"""
|
||||
if not isinstance(content, list):
|
||||
return False
|
||||
|
||||
for item in content:
|
||||
if isinstance(item, dict):
|
||||
item_type = item.get("type", "")
|
||||
if item_type in ["image", "audio", "video", "image_url"]:
|
||||
return True
|
||||
# Check for keys that indicate media
|
||||
if any(key in item for key in ["image", "audio", "video", "image_url"]):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
def _convert_content_to_blocks(
|
||||
self,
|
||||
content: str | list[dict[str, Any]],
|
||||
) -> list[Any]:
|
||||
"""Convert DashScope content to AgentScope content blocks.
|
||||
|
||||
Args:
|
||||
content: DashScope content (string or list of content items).
|
||||
|
||||
Returns:
|
||||
List of AgentScope content blocks.
|
||||
"""
|
||||
from agentscope.message import (
|
||||
AudioBlock,
|
||||
ImageBlock,
|
||||
TextBlock,
|
||||
URLSource,
|
||||
VideoBlock,
|
||||
)
|
||||
|
||||
blocks = []
|
||||
|
||||
if isinstance(content, str):
|
||||
if content:
|
||||
blocks.append(
|
||||
TextBlock(
|
||||
type="text",
|
||||
text=content,
|
||||
),
|
||||
)
|
||||
elif isinstance(content, list):
|
||||
for item in content:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
|
||||
# Handle text blocks
|
||||
if "text" in item:
|
||||
text = item["text"]
|
||||
if text:
|
||||
blocks.append(
|
||||
TextBlock(
|
||||
type="text",
|
||||
text=text,
|
||||
),
|
||||
)
|
||||
|
||||
# Handle image blocks
|
||||
elif "image" in item or item.get("type") == "image":
|
||||
url = item.get("image", "")
|
||||
blocks.append(
|
||||
ImageBlock(
|
||||
type="image",
|
||||
source=URLSource(
|
||||
type="url",
|
||||
url=url,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
# Handle image_url format (OpenAI style)
|
||||
elif "image_url" in item or item.get("type") == "image_url":
|
||||
image_url = item.get("image_url", {})
|
||||
if isinstance(image_url, dict):
|
||||
url = image_url.get("url", "")
|
||||
else:
|
||||
url = str(image_url)
|
||||
|
||||
blocks.append(
|
||||
ImageBlock(
|
||||
type="image",
|
||||
source=URLSource(
|
||||
type="url",
|
||||
url=url,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
# Handle audio blocks
|
||||
elif "audio" in item or item.get("type") == "audio":
|
||||
url = item.get("audio", "")
|
||||
blocks.append(
|
||||
AudioBlock(
|
||||
type="audio",
|
||||
source=URLSource(
|
||||
type="url",
|
||||
url=url,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
# Handle video blocks
|
||||
elif "video" in item or item.get("type") == "video":
|
||||
video_data = item.get("video", "")
|
||||
# Video can be a URL string or list of frame URLs
|
||||
if isinstance(video_data, list):
|
||||
# Use first frame as URL for now
|
||||
url = video_data[0] if video_data else ""
|
||||
else:
|
||||
url = str(video_data)
|
||||
|
||||
blocks.append(
|
||||
VideoBlock(
|
||||
type="video",
|
||||
source=URLSource(
|
||||
type="url",
|
||||
url=url,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
return blocks
|
||||
|
||||
def _convert_tool_calls_to_blocks(
|
||||
self,
|
||||
tool_calls: list[dict[str, Any]],
|
||||
) -> list[Any]:
|
||||
"""Convert DashScope tool_calls to AgentScope ToolUseBlocks.
|
||||
|
||||
Args:
|
||||
tool_calls: List of DashScope tool call dictionaries.
|
||||
|
||||
Returns:
|
||||
List of AgentScope ToolUseBlock objects.
|
||||
"""
|
||||
from agentscope.message import ToolUseBlock
|
||||
|
||||
blocks = []
|
||||
|
||||
for tool_call in tool_calls:
|
||||
tool_id = tool_call.get("id", "")
|
||||
function = tool_call.get("function", {})
|
||||
name = function.get("name", "")
|
||||
arguments_str = function.get("arguments", "{}")
|
||||
|
||||
# Parse arguments JSON string to dict
|
||||
try:
|
||||
arguments = json.loads(arguments_str)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
arguments = {}
|
||||
|
||||
blocks.append(
|
||||
ToolUseBlock(
|
||||
type="tool_use",
|
||||
id=tool_id,
|
||||
name=name,
|
||||
input=arguments,
|
||||
),
|
||||
)
|
||||
|
||||
return blocks
|
||||
|
||||
def _convert_tool_result_to_blocks(
|
||||
self,
|
||||
dashscope_msg: dict[str, Any],
|
||||
) -> list[Any]:
|
||||
"""Convert DashScope tool result message to AgentScope ToolResultBlock.
|
||||
|
||||
Args:
|
||||
dashscope_msg: DashScope tool result message with role="tool".
|
||||
|
||||
Returns:
|
||||
List containing a single ToolResultBlock.
|
||||
"""
|
||||
from agentscope.message import ToolResultBlock
|
||||
|
||||
tool_call_id = dashscope_msg.get("tool_call_id", "")
|
||||
content = dashscope_msg.get("content", "")
|
||||
name = dashscope_msg.get("name", "")
|
||||
|
||||
# Tool result content should be plain text
|
||||
return [
|
||||
ToolResultBlock(
|
||||
type="tool_result",
|
||||
id=tool_call_id,
|
||||
name=name,
|
||||
output=content if content else "",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def convert_dashscope_to_agentscope(
|
||||
dashscope_msg: dict[str, Any] | list[dict[str, Any]],
|
||||
name: str | None = None,
|
||||
default_name: str = "assistant",
|
||||
) -> "Msg | list[Msg]":
|
||||
"""Convenience function to convert DashScope format to AgentScope Msg.
|
||||
|
||||
Args:
|
||||
dashscope_msg: Single message dict or list of message dicts in DashScope format.
|
||||
name: Override name for the message(s).
|
||||
default_name: Default name for assistant messages.
|
||||
|
||||
Returns:
|
||||
Single Msg object or list of Msg objects.
|
||||
|
||||
Examples:
|
||||
>>> # Single message
|
||||
>>> msg = convert_dashscope_to_agentscope({"role": "assistant", "content": "Hi"})
|
||||
>>> # Multiple messages
|
||||
>>> msgs = convert_dashscope_to_agentscope([
|
||||
... {"role": "user", "content": "Hello"},
|
||||
... {"role": "assistant", "content": "Hi there!"}
|
||||
... ])
|
||||
"""
|
||||
converter = DashScopeToAgentScopeConverter(default_name=default_name)
|
||||
|
||||
if isinstance(dashscope_msg, list):
|
||||
return converter.convert_messages(dashscope_msg)
|
||||
else:
|
||||
return converter.convert_message(dashscope_msg, name=name)
|
||||
|
|
@ -3,7 +3,7 @@
|
|||
import asyncio
|
||||
import hashlib
|
||||
from collections.abc import AsyncGenerator, Coroutine
|
||||
from typing import Any
|
||||
from typing import Any, Literal
|
||||
|
||||
import numpy as np
|
||||
from loguru import logger
|
||||
|
|
@ -31,8 +31,8 @@ async def execute_stream_task(
|
|||
stream_queue: asyncio.Queue,
|
||||
task: asyncio.Task,
|
||||
task_name: str | None = None,
|
||||
as_bytes: bool = False,
|
||||
) -> AsyncGenerator[str | bytes, None]:
|
||||
output_format: Literal["str", "bytes", "chunk"] = "str",
|
||||
) -> AsyncGenerator[str | bytes | StreamChunk, None]:
|
||||
"""
|
||||
Core stream flow execution logic.
|
||||
|
||||
|
|
@ -43,46 +43,83 @@ async def execute_stream_task(
|
|||
stream_queue: Queue to receive StreamChunk objects from
|
||||
task: Background task executing the flow
|
||||
task_name: Optional flow name for logging purposes
|
||||
as_bytes: If True, yield bytes for HTTP responses; if False, yield strings
|
||||
output_format: Output format control
|
||||
- "str": SSE-formatted string (default)
|
||||
- "bytes": SSE-formatted bytes for HTTP responses
|
||||
- "chunk": Raw StreamChunk objects
|
||||
|
||||
Yields:
|
||||
SSE-formatted data chunks (either str or bytes based on as_bytes)
|
||||
"""
|
||||
done_msg = b"data:[DONE]\n\n" if as_bytes else "data:[DONE]\n\n"
|
||||
- str: SSE-formatted data when output_format="str"
|
||||
- bytes: SSE-formatted data when output_format="bytes"
|
||||
- StreamChunk: Raw chunk objects when output_format="chunk"
|
||||
|
||||
Raises:
|
||||
Exception: Re-raises any exception from the background task
|
||||
"""
|
||||
try:
|
||||
while True:
|
||||
# Wait for next chunk or check if task failed
|
||||
get_chunk = asyncio.create_task(stream_queue.get())
|
||||
done, _ = await asyncio.wait({get_chunk, task}, return_when=asyncio.FIRST_COMPLETED)
|
||||
done, _pending = await asyncio.wait({get_chunk, task}, return_when=asyncio.FIRST_COMPLETED)
|
||||
|
||||
if get_chunk in done:
|
||||
# Priority 1: Check if main task finished (may have exception)
|
||||
if task in done:
|
||||
# Task finished - check for exceptions first
|
||||
exc = task.exception()
|
||||
if exc:
|
||||
log_msg = f"Task error in {task_name}: {exc}" if task_name else f"Task error: {exc}"
|
||||
logger.exception(log_msg)
|
||||
raise exc
|
||||
|
||||
# Task completed successfully - drain remaining chunks if any
|
||||
if get_chunk in done:
|
||||
chunk: StreamChunk = get_chunk.result()
|
||||
if output_format == "chunk":
|
||||
yield chunk
|
||||
if chunk.done:
|
||||
break
|
||||
else:
|
||||
if chunk.done:
|
||||
yield b"data:[DONE]\n\n" if output_format == "bytes" else "data:[DONE]\n\n"
|
||||
break
|
||||
data = f"data:{chunk.model_dump_json()}\n\n"
|
||||
yield data.encode() if output_format == "bytes" else data
|
||||
else:
|
||||
# No more chunks, task completed
|
||||
get_chunk.cancel()
|
||||
if output_format == "chunk":
|
||||
yield StreamChunk(chunk_type=ChunkEnum.DONE, chunk="", done=True)
|
||||
else:
|
||||
yield b"data:[DONE]\n\n" if output_format == "bytes" else "data:[DONE]\n\n"
|
||||
break
|
||||
|
||||
elif get_chunk in done:
|
||||
# Got a chunk from the queue (task still running)
|
||||
chunk: StreamChunk = get_chunk.result()
|
||||
|
||||
# Handle raw chunk mode
|
||||
if output_format == "chunk":
|
||||
yield chunk
|
||||
if chunk.done:
|
||||
break
|
||||
continue
|
||||
|
||||
# Handle SSE format mode (str or bytes)
|
||||
if chunk.done:
|
||||
yield done_msg
|
||||
yield b"data:[DONE]\n\n" if output_format == "bytes" else "data:[DONE]\n\n"
|
||||
break
|
||||
|
||||
data = f"data:{chunk.model_dump_json()}\n\n"
|
||||
yield data.encode() if as_bytes else data
|
||||
else:
|
||||
# Task finished unexpectedly or raised exception
|
||||
await task
|
||||
yield done_msg
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
log_msg = f"Stream error in {task_name}: {e}" if task_name else f"Stream error: {e}"
|
||||
logger.exception(log_msg)
|
||||
|
||||
err = StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e), done=True)
|
||||
err_data = f"data:{err.model_dump_json()}\n\n"
|
||||
yield err_data.encode() if as_bytes else err_data
|
||||
yield done_msg
|
||||
yield data.encode() if output_format == "bytes" else data
|
||||
|
||||
finally:
|
||||
# Ensure task is cancelled if still running to avoid resource leaks
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
|
||||
def hash_text(text: str) -> str:
|
||||
|
|
|
|||
|
|
@ -12,6 +12,11 @@ from ..schema import Message, Trajectory, MemoryNode
|
|||
def format_messages(
|
||||
messages: list[Message | dict],
|
||||
add_index: bool = True,
|
||||
add_time: bool = True,
|
||||
use_name: bool = True,
|
||||
add_reasoning: bool = True,
|
||||
add_tools: bool = True,
|
||||
strip_markdown_headers: bool = True,
|
||||
enable_system: bool = False,
|
||||
) -> str:
|
||||
"""Formats a list of messages into a single string, optionally filtering system roles."""
|
||||
|
|
@ -25,11 +30,11 @@ def format_messages(
|
|||
formatted_lines.append(
|
||||
message.format_message(
|
||||
index=i if add_index else None,
|
||||
add_time=True,
|
||||
use_name=True,
|
||||
add_reasoning=True,
|
||||
add_tools=True,
|
||||
strip_markdown_headers=True,
|
||||
add_time=add_time,
|
||||
use_name=use_name,
|
||||
add_reasoning=add_reasoning,
|
||||
add_tools=add_tools,
|
||||
strip_markdown_headers=strip_markdown_headers,
|
||||
),
|
||||
)
|
||||
return "\n".join(formatted_lines)
|
||||
|
|
|
|||
|
|
@ -5,8 +5,14 @@ import sys
|
|||
from datetime import datetime
|
||||
|
||||
|
||||
def init_logger(log_dir: str = "logs", level: str = "INFO") -> None:
|
||||
"""Initialize the logger with both file and console handlers."""
|
||||
def init_logger(log_dir: str = "logs", level: str = "INFO", log_to_console: bool = True) -> None:
|
||||
"""Initialize the logger with both file and console handlers.
|
||||
|
||||
Args:
|
||||
log_dir: Directory path for log files
|
||||
level: Logging level (DEBUG, INFO, WARNING, ERROR, CRITICAL)
|
||||
log_to_console: Whether to print logs to console/screen
|
||||
"""
|
||||
from loguru import logger
|
||||
|
||||
# Remove default handler to avoid duplicate logs
|
||||
|
|
@ -31,10 +37,11 @@ def init_logger(log_dir: str = "logs", level: str = "INFO") -> None:
|
|||
format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {file}:{line} | {function} | {message}",
|
||||
)
|
||||
|
||||
# Configure colorized standard output logging
|
||||
logger.add(
|
||||
sink=sys.stdout,
|
||||
level=level,
|
||||
format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {file}:{line} | {function} | {message}",
|
||||
colorize=True,
|
||||
)
|
||||
# Configure colorized standard output logging if enabled
|
||||
if log_to_console:
|
||||
logger.add(
|
||||
sink=sys.stdout,
|
||||
level=level,
|
||||
format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {file}:{line} | {function} | {message}",
|
||||
colorize=True,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -51,7 +51,9 @@ class ReMe(Application):
|
|||
llm_api_base: str | None = None,
|
||||
embedding_api_key: str | None = None,
|
||||
embedding_api_base: str | None = None,
|
||||
config_path: str = "default",
|
||||
enable_logo: bool = True,
|
||||
log_to_console: bool = True,
|
||||
default_llm_config: dict | None = None,
|
||||
default_embedding_model_config: dict | None = None,
|
||||
default_vector_store_config: dict | None = None,
|
||||
|
|
@ -70,7 +72,9 @@ class ReMe(Application):
|
|||
llm_api_base: API base for LLM provider
|
||||
embedding_api_key: API key for embedding provider
|
||||
embedding_api_base: API base for embedding provider
|
||||
config_path: Path to config file
|
||||
enable_logo: Enable logo
|
||||
log_to_console: Log to console
|
||||
default_llm_config: LLM configuration
|
||||
default_embedding_model_config: Embedding model configuration
|
||||
default_vector_store_config: Vector store configuration
|
||||
|
|
@ -100,7 +104,9 @@ class ReMe(Application):
|
|||
llm_api_base=llm_api_base,
|
||||
embedding_api_key=embedding_api_key,
|
||||
embedding_api_base=embedding_api_base,
|
||||
config_path=config_path,
|
||||
enable_logo=enable_logo,
|
||||
log_to_console=log_to_console,
|
||||
parser=ReMeConfigParser,
|
||||
default_llm_config=default_llm_config,
|
||||
default_embedding_model_config=default_embedding_model_config,
|
||||
|
|
|
|||
336
reme/reme_fs.py
336
reme/reme_fs.py
|
|
@ -1,13 +1,20 @@
|
|||
"""ReMe File System"""
|
||||
|
||||
import asyncio
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import AsyncGenerator
|
||||
|
||||
from .agent.fs import FsCompactor, FsSummarizer
|
||||
from prompt_toolkit import PromptSession
|
||||
|
||||
from reme.core.utils import execute_stream_task
|
||||
from .agent.chat import FsCli
|
||||
from .agent.fs import FsCompactor, FsContextChecker, FsSummarizer
|
||||
from .config import ReMeConfigParser
|
||||
from .core import Application
|
||||
from .core.enumeration import MemorySource
|
||||
from .core.enumeration import ChunkEnum
|
||||
from .core.op import BaseTool
|
||||
from .core.schema import Message
|
||||
from .core.schema import Message, StreamChunk
|
||||
from .tool.fs import (
|
||||
BashTool,
|
||||
EditTool,
|
||||
|
|
@ -31,13 +38,22 @@ class ReMeFs(Application):
|
|||
llm_api_base: str | None = None,
|
||||
embedding_api_key: str | None = None,
|
||||
embedding_api_base: str | None = None,
|
||||
config_path: str = "fs",
|
||||
enable_logo: bool = True,
|
||||
log_to_console: bool = True,
|
||||
default_llm_config: dict | None = None,
|
||||
default_embedding_model_config: dict | None = None,
|
||||
default_memory_store_config: dict | None = None,
|
||||
default_token_counter_config: dict | None = None,
|
||||
default_file_watcher_config: dict | None = None,
|
||||
working_dir: str = ".reme",
|
||||
context_window_tokens: int = 128000,
|
||||
reserve_tokens: int = 36000,
|
||||
keep_recent_tokens: int = 20000,
|
||||
hybrid_enabled: bool = True,
|
||||
hybrid_vector_weight: float = 0.7,
|
||||
hybrid_text_weight: float = 0.3,
|
||||
hybrid_candidate_multiplier: float = 3.0,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize ReMe with config."""
|
||||
|
|
@ -47,7 +63,9 @@ class ReMeFs(Application):
|
|||
llm_api_base=llm_api_base,
|
||||
embedding_api_key=embedding_api_key,
|
||||
embedding_api_base=embedding_api_base,
|
||||
config_path=config_path,
|
||||
enable_logo=enable_logo,
|
||||
log_to_console=log_to_console,
|
||||
parser=ReMeConfigParser,
|
||||
default_llm_config=default_llm_config,
|
||||
default_embedding_model_config=default_embedding_model_config,
|
||||
|
|
@ -56,9 +74,25 @@ class ReMeFs(Application):
|
|||
default_file_watcher_config=default_file_watcher_config,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
self.working_dir: str = working_dir
|
||||
Path(self.working_dir).mkdir(parents=True, exist_ok=True)
|
||||
self.context_window_tokens: int = context_window_tokens
|
||||
self.reserve_tokens: int = reserve_tokens
|
||||
self.keep_recent_tokens: int = keep_recent_tokens
|
||||
self.hybrid_enabled: bool = hybrid_enabled
|
||||
self.hybrid_vector_weight: float = hybrid_vector_weight
|
||||
self.hybrid_text_weight: float = hybrid_text_weight
|
||||
self.hybrid_candidate_multiplier: float = hybrid_candidate_multiplier
|
||||
|
||||
# Setup file system tools
|
||||
self.fs_tools: list[BaseTool] = [
|
||||
FsMemorySearch(
|
||||
hybrid_enabled=hybrid_enabled,
|
||||
hybrid_vector_weight=hybrid_vector_weight,
|
||||
hybrid_text_weight=hybrid_text_weight,
|
||||
hybrid_candidate_multiplier=hybrid_candidate_multiplier,
|
||||
),
|
||||
FsMemoryGet(cwd=self.working_dir),
|
||||
BashTool(cwd=self.working_dir),
|
||||
EditTool(cwd=self.working_dir),
|
||||
FindTool(cwd=self.working_dir),
|
||||
|
|
@ -67,67 +101,82 @@ class ReMeFs(Application):
|
|||
ReadTool(cwd=self.working_dir),
|
||||
WriteTool(cwd=self.working_dir),
|
||||
]
|
||||
self.working_path: Path = Path(self.working_dir)
|
||||
self.working_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Commands
|
||||
self.commands = [
|
||||
"/new",
|
||||
"/compact",
|
||||
"/exit",
|
||||
"/help",
|
||||
"/clear",
|
||||
]
|
||||
|
||||
async def context_check(self, messages: list[Message | dict]) -> dict:
|
||||
"""Check if messages exceed context limits."""
|
||||
checker = FsContextChecker(
|
||||
context_window_tokens=self.context_window_tokens,
|
||||
reserve_tokens=self.reserve_tokens,
|
||||
keep_recent_tokens=self.keep_recent_tokens,
|
||||
)
|
||||
return await checker.call(messages=messages, service_context=self.service_context)
|
||||
|
||||
async def compact(
|
||||
self,
|
||||
messages: list[Message | dict],
|
||||
context_window_tokens: int = 128000,
|
||||
reserve_tokens: int = 36000,
|
||||
keep_recent_tokens: int = 20000,
|
||||
):
|
||||
"""Compact messages."""
|
||||
messages = [Message(**message) if isinstance(message, dict) else message for message in messages]
|
||||
compactor = FsCompactor(
|
||||
context_window_tokens=context_window_tokens,
|
||||
reserve_tokens=reserve_tokens,
|
||||
keep_recent_tokens=keep_recent_tokens,
|
||||
messages_to_summarize: list[Message | dict] = None,
|
||||
turn_prefix_messages: list[Message | dict] = None,
|
||||
previous_summary: str = "",
|
||||
) -> str:
|
||||
"""Compact messages into a summary.
|
||||
|
||||
Args:
|
||||
messages_to_summarize: Messages to summarize
|
||||
turn_prefix_messages: Messages to prepend to each turn
|
||||
previous_summary: Previous summary to build upon
|
||||
|
||||
Returns:
|
||||
Compaction result from FsCompactor
|
||||
"""
|
||||
compactor = FsCompactor()
|
||||
return await compactor.call(
|
||||
messages_to_summarize=messages_to_summarize or [],
|
||||
turn_prefix_messages=turn_prefix_messages or [],
|
||||
previous_summary=previous_summary,
|
||||
service_context=self.service_context,
|
||||
)
|
||||
|
||||
return await compactor.call(messages=messages, service_context=self.service_context)
|
||||
async def summary(self, messages: list[Message | dict], date: str):
|
||||
"""Generate a summary of the given messages.
|
||||
|
||||
async def summary(
|
||||
self,
|
||||
messages: list[Message | dict],
|
||||
date: str,
|
||||
version: str = "default",
|
||||
context_window_tokens: int = 128000,
|
||||
reserve_tokens: int = 32000,
|
||||
soft_threshold_tokens: int = 4000,
|
||||
):
|
||||
"""Summarize messages."""
|
||||
messages = [Message(**message) if isinstance(message, dict) else message for message in messages]
|
||||
summarizer = FsSummarizer(
|
||||
tools=self.fs_tools,
|
||||
version=version,
|
||||
context_window_tokens=context_window_tokens,
|
||||
reserve_tokens=reserve_tokens,
|
||||
soft_threshold_tokens=soft_threshold_tokens,
|
||||
)
|
||||
Args:
|
||||
messages: Messages to summarize
|
||||
date: Date of the conversation
|
||||
|
||||
Returns:
|
||||
Summary of the given messages
|
||||
"""
|
||||
summarizer = FsSummarizer(tools=self.fs_tools, working_dir=self.working_dir)
|
||||
return await summarizer.call(messages=messages, date=date, service_context=self.service_context)
|
||||
|
||||
async def memory_search(
|
||||
self,
|
||||
query: str,
|
||||
max_results: int = 20,
|
||||
min_score: float = 0.1,
|
||||
sources: list[MemorySource] | None = None,
|
||||
hybrid_enabled: bool = True,
|
||||
hybrid_vector_weight: float = 0.7,
|
||||
hybrid_text_weight: float = 0.3,
|
||||
hybrid_candidate_multiplier: float = 3.0,
|
||||
) -> str:
|
||||
"""Semantically search memory files."""
|
||||
search_tool = FsMemorySearch(
|
||||
sources=sources,
|
||||
hybrid_enabled=hybrid_enabled,
|
||||
hybrid_vector_weight=hybrid_vector_weight,
|
||||
hybrid_text_weight=hybrid_text_weight,
|
||||
hybrid_candidate_multiplier=hybrid_candidate_multiplier,
|
||||
)
|
||||
async def memory_search(self, query: str, max_results: int = 10, min_score: float = 0.3) -> str:
|
||||
"""
|
||||
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.
|
||||
|
||||
Args:
|
||||
query: The semantic search query to find relevant memory snippets
|
||||
max_results: Maximum number of search results to return (optional), default is 10
|
||||
min_score: Minimum similarity score threshold for results (optional), default is 0.3
|
||||
|
||||
Returns:
|
||||
Search results as formatted string
|
||||
"""
|
||||
search_tool = FsMemorySearch(
|
||||
hybrid_enabled=self.hybrid_enabled,
|
||||
hybrid_vector_weight=self.hybrid_vector_weight,
|
||||
hybrid_text_weight=self.hybrid_text_weight,
|
||||
hybrid_candidate_multiplier=self.hybrid_candidate_multiplier,
|
||||
)
|
||||
return await search_tool.call(
|
||||
query=query,
|
||||
max_results=max_results,
|
||||
|
|
@ -136,6 +185,179 @@ class ReMeFs(Application):
|
|||
)
|
||||
|
||||
async def memory_get(self, path: str, offset: int | None = None, limit: int | None = None) -> str:
|
||||
"""Read specific snippets from memory files."""
|
||||
get_tool = FsMemoryGet(workspace_dir=self.working_dir)
|
||||
"""
|
||||
Safe snippet read from MEMORY.md, memory/*.md with optional offset/limit;
|
||||
use after memory_search to pull only the needed lines and keep context small.
|
||||
|
||||
Args:
|
||||
path: Path to the memory file to read (relative or absolute)
|
||||
offset: Starting line number (1-indexed, optional)
|
||||
limit: Number of lines to read from the starting line (optional)
|
||||
|
||||
Returns:
|
||||
Memory file content as string
|
||||
"""
|
||||
get_tool = FsMemoryGet(cwd=self.working_dir)
|
||||
return await get_tool.call(path=path, offset=offset, limit=limit, service_context=self.service_context)
|
||||
|
||||
async def needs_compaction(self, messages: list[Message | dict]) -> bool:
|
||||
"""Check if messages need compaction based on context window limits."""
|
||||
messages = [Message(**message) if isinstance(message, dict) else message for message in messages]
|
||||
checker = FsContextChecker(
|
||||
context_window_tokens=self.context_window_tokens,
|
||||
reserve_tokens=self.reserve_tokens,
|
||||
)
|
||||
result = await checker.call(messages=messages, service_context=self.service_context)
|
||||
return result["needs_compaction"]
|
||||
|
||||
async def chat_with_remy(self, tool_result_max_size: int = 100):
|
||||
"""Interactive CLI chat with Remy using simple streaming output."""
|
||||
fs_cli = FsCli(
|
||||
working_dir=self.working_dir,
|
||||
tools=self.fs_tools,
|
||||
context_window_tokens=self.context_window_tokens,
|
||||
reserve_tokens=self.reserve_tokens,
|
||||
keep_recent_tokens=self.keep_recent_tokens,
|
||||
hybrid_enabled=self.hybrid_enabled,
|
||||
hybrid_vector_weight=self.hybrid_vector_weight,
|
||||
hybrid_text_weight=self.hybrid_text_weight,
|
||||
hybrid_candidate_multiplier=self.hybrid_candidate_multiplier,
|
||||
tool_result_max_size=tool_result_max_size,
|
||||
)
|
||||
session = PromptSession()
|
||||
|
||||
# Print welcome banner
|
||||
print("\n========================================")
|
||||
print(" Welcome to Remy Chat!")
|
||||
print(" Type /exit to quit, /new to start fresh.")
|
||||
print("========================================\n")
|
||||
|
||||
async def chat(q: str) -> AsyncGenerator[StreamChunk, None]:
|
||||
"""Execute chat query and yield streaming chunks."""
|
||||
stream_queue = asyncio.Queue()
|
||||
task = asyncio.create_task(
|
||||
fs_cli.call(
|
||||
query=q,
|
||||
stream_queue=stream_queue,
|
||||
service_context=self.service_context,
|
||||
),
|
||||
)
|
||||
async for _chunk in execute_stream_task(
|
||||
stream_queue=stream_queue,
|
||||
task=task,
|
||||
task_name="cli",
|
||||
output_format="chunk",
|
||||
):
|
||||
yield _chunk
|
||||
|
||||
while True:
|
||||
try:
|
||||
# Get user input (async)
|
||||
user_input = await session.prompt_async("You: ", default="")
|
||||
if not user_input.strip():
|
||||
continue
|
||||
|
||||
# Handle commands
|
||||
if user_input.strip() == "/exit":
|
||||
break
|
||||
|
||||
if user_input.strip() == "/new":
|
||||
result = await fs_cli.reset()
|
||||
print(f"{result}\nConversation reset\n")
|
||||
continue
|
||||
|
||||
if user_input.strip() == "/compact":
|
||||
result = await fs_cli.compact()
|
||||
print(f"{result}\nHistory compacted.\n")
|
||||
continue
|
||||
|
||||
if user_input.strip() == "/clear":
|
||||
fs_cli.messages.clear()
|
||||
print("History cleared.\n")
|
||||
continue
|
||||
|
||||
if user_input.strip() == "/help":
|
||||
print("\nCommands:")
|
||||
for command in self.commands:
|
||||
print(f" {command}")
|
||||
continue
|
||||
|
||||
# Stream processing state
|
||||
in_thinking = False
|
||||
in_answer = False
|
||||
|
||||
try:
|
||||
async for chunk in chat(user_input):
|
||||
if chunk.chunk_type == ChunkEnum.THINK:
|
||||
if not in_thinking:
|
||||
print("\033[90mThinking: ", end="", flush=True)
|
||||
in_thinking = True
|
||||
print(chunk.chunk, end="", flush=True)
|
||||
|
||||
elif chunk.chunk_type == ChunkEnum.ANSWER:
|
||||
if in_thinking:
|
||||
print("\033[0m") # reset color after thinking
|
||||
in_thinking = False
|
||||
if not in_answer:
|
||||
print("\nRemy: ", end="", flush=True)
|
||||
in_answer = True
|
||||
print(chunk.chunk, end="", flush=True)
|
||||
|
||||
elif chunk.chunk_type == ChunkEnum.TOOL:
|
||||
if in_thinking:
|
||||
print("\033[0m") # reset color after thinking
|
||||
in_thinking = False
|
||||
print(f"\033[36m -> Tool: {chunk.chunk}\033[0m")
|
||||
|
||||
elif chunk.chunk_type == ChunkEnum.TOOL_RESULT:
|
||||
tool_name = chunk.metadata.get("tool_name", "unknown")
|
||||
result = chunk.chunk
|
||||
if len(result) > tool_result_max_size:
|
||||
result = result[:tool_result_max_size] + f"... ({len(chunk.chunk)} chars total)"
|
||||
print(f"\033[36m Tool result for {tool_name}: {result.strip()}\033[0m")
|
||||
|
||||
elif chunk.chunk_type == ChunkEnum.ERROR:
|
||||
print(f"\n\033[91m[ERROR] {chunk.chunk}\033[0m")
|
||||
# Also log the full error metadata if available
|
||||
if chunk.metadata:
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
elif chunk.chunk_type == ChunkEnum.DONE:
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
print(f"\nStream error: {e}")
|
||||
|
||||
# End current streaming line
|
||||
print("\n")
|
||||
print("----------------------------------------\n")
|
||||
|
||||
except EOFError:
|
||||
break
|
||||
except KeyboardInterrupt:
|
||||
print("\nInterrupted.")
|
||||
break
|
||||
except Exception as e:
|
||||
print(f"Error: {e}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
print("\nGoodbye!\n")
|
||||
|
||||
|
||||
async def async_main():
|
||||
"""Main function for testing the ReMeFs CLI."""
|
||||
async with ReMeFs(*sys.argv[1:], log_to_console=False) as reme:
|
||||
await reme.chat_with_remy()
|
||||
|
||||
|
||||
def main():
|
||||
"""Main function for testing the ReMeFs CLI."""
|
||||
asyncio.run(async_main())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
|
|||
|
|
@ -10,11 +10,11 @@ from .base_fs_tool import BaseFsTool
|
|||
class FsMemoryGet(BaseFsTool):
|
||||
"""Read specific snippets from memory files."""
|
||||
|
||||
def __init__(self, workspace_dir: str | None = None, **kwargs):
|
||||
def __init__(self, cwd: str | None = None, **kwargs):
|
||||
"""Initialize memory get tool."""
|
||||
kwargs.setdefault("name", "memory_get")
|
||||
super().__init__(**kwargs)
|
||||
self.workspace_dir = workspace_dir or os.getcwd()
|
||||
self.cwd = cwd or os.getcwd()
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
return ToolCall(
|
||||
|
|
@ -53,7 +53,7 @@ class FsMemoryGet(BaseFsTool):
|
|||
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))
|
||||
abs_path = os.path.abspath(os.path.join(self.cwd, raw_path))
|
||||
assert abs_path.lower().endswith(".md")
|
||||
|
||||
# Check file exists, is not a symlink, and is a regular file
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@
|
|||
|
||||
import json
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from reme.core.enumeration import MemorySource
|
||||
from reme.core.schema import MemorySearchResult, ToolCall
|
||||
from .base_fs_tool import BaseFsTool
|
||||
|
|
@ -76,6 +78,18 @@ class FsMemorySearch(BaseFsTool):
|
|||
keyword_results = await self._search_keyword(query, candidates)
|
||||
vector_results = await self._search_vector(query, candidates)
|
||||
|
||||
# Log original vector results
|
||||
logger.debug("\n=== Vector Search Results ===")
|
||||
for i, r in enumerate(vector_results[:10], 1):
|
||||
snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet
|
||||
logger.debug(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}")
|
||||
|
||||
# Log original keyword results
|
||||
logger.debug("\n=== Keyword Search Results ===")
|
||||
for i, r in enumerate(keyword_results[:10], 1):
|
||||
snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet
|
||||
logger.debug(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}")
|
||||
|
||||
if not keyword_results:
|
||||
results = [r for r in vector_results if r.score >= min_score][:max_results]
|
||||
elif not vector_results:
|
||||
|
|
@ -87,6 +101,13 @@ class FsMemorySearch(BaseFsTool):
|
|||
vector_weight=self.hybrid_vector_weight,
|
||||
text_weight=self.hybrid_text_weight,
|
||||
)
|
||||
|
||||
# Log merged results
|
||||
logger.debug("\n=== Merged Hybrid Results ===")
|
||||
for i, r in enumerate(merged[:10], 1):
|
||||
snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet
|
||||
logger.debug(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}")
|
||||
|
||||
results = [r for r in merged if r.score >= min_score][:max_results]
|
||||
else:
|
||||
vector_results = await self._search_vector(query, candidates)
|
||||
|
|
|
|||
383
tests/test_agentscope_converter.py
Normal file
383
tests/test_agentscope_converter.py
Normal file
|
|
@ -0,0 +1,383 @@
|
|||
"""Test cases for DashScope to AgentScope message conversion."""
|
||||
|
||||
import json
|
||||
|
||||
|
||||
def test_plain_text_list_conversion():
|
||||
"""Test converting a long list of plain text DashScope messages to AgentScope Msgs."""
|
||||
from reme.core.utils.agentscope_utils import convert_dashscope_to_agentscope
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("TEST 1: Plain Text List Conversion (List[Dict] -> List[Msg])")
|
||||
print("=" * 80)
|
||||
|
||||
# Long conversation with plain text messages
|
||||
dashscope_msgs = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "你是一个专业的AI助手,擅长回答各种问题。",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "你好!请问你能帮我做什么?",
|
||||
"name": "用户A",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "你好!我可以帮你回答问题、提供建议、进行对话等。有什么我可以帮助你的吗?",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "我想了解一下今天北京的天气情况。",
|
||||
"name": "用户A",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "好的,让我帮你查询一下北京的天气。",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_weather_001",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"arguments": '{"city": "北京", "date": "今天"}',
|
||||
},
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_weather_001",
|
||||
"name": "get_weather",
|
||||
"content": "北京今天天气:晴转多云,气温15-25°C,风力3-4级,空气质量良好,适合户外活动。",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "根据天气查询结果,北京今天的天气情况如下:\n- 天气:晴转多云\n- 气温:15-25°C\n- 风力:3-4级\n- 空气质量:良好\n\n今天天气不错,适合户外活动哦!",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "太好了!那你能推荐一些户外活动吗?",
|
||||
"name": "用户A",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": (
|
||||
"当然可以!根据今天的天气情况,我推荐以下几个户外活动:\n\n"
|
||||
"1. 公园散步或慢跑\n2. 骑自行车游览城市\n3. 去郊外爬山\n"
|
||||
"4. 在户外咖啡厅享受阳光\n5. 拍摄城市风景照片\n\n你对哪个活动比较感兴趣呢?"
|
||||
),
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "爬山听起来不错!你能推荐几个北京周边的爬山地点吗?",
|
||||
"name": "用户A",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"reasoning_content": "用户想要北京周边的爬山地点推荐。我应该推荐一些知名且适合休闲爬山的地方,考虑交通便利性和难度适中。",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "北京周边有很多适合爬山的好去处,这里给你推荐几个:\n\n**初级难度:**\n1. 香山公园 - 红叶季节尤其美丽\n"
|
||||
"2. 景山公园 - 可以俯瞰故宫全景\n\n**中级难度:**\n3. 八达岭长城 - 著名的世界文化遗产\n"
|
||||
"4. 慕田峪长城 - 相对人少,风景优美\n\n**进阶难度:**\n5. 妙峰山 - 自然风光秀丽\n"
|
||||
"6. 百花山 - 植被丰富,空气清新\n\n建议提前查看开放时间和门票信息,准备好登山装备和充足的水。祝你爬山愉快!",
|
||||
},
|
||||
]
|
||||
|
||||
print(f"\n[Input] DashScope messages: {len(dashscope_msgs)} messages")
|
||||
print(json.dumps(dashscope_msgs, ensure_ascii=False, indent=2))
|
||||
|
||||
# Convert to AgentScope Msgs
|
||||
msgs = convert_dashscope_to_agentscope(dashscope_msgs)
|
||||
|
||||
print(f"\n[Output] AgentScope Msgs: {len(msgs)} messages")
|
||||
print("=" * 80)
|
||||
|
||||
for i, msg in enumerate(msgs):
|
||||
print(f"\n【Message {i+1}/{len(msgs)}】")
|
||||
print(f" name: {msg.name}")
|
||||
print(f" role: {msg.role}")
|
||||
print(f" content type: {type(msg.content).__name__}")
|
||||
print(f" timestamp: {msg.timestamp}")
|
||||
|
||||
if isinstance(msg.content, str):
|
||||
content_preview = msg.content[:100] + "..." if len(msg.content) > 100 else msg.content
|
||||
print(f" content: {content_preview}")
|
||||
elif isinstance(msg.content, list):
|
||||
print(f" content blocks: {len(msg.content)} blocks")
|
||||
for j, block in enumerate(msg.content):
|
||||
block_type = block.get("type")
|
||||
print(f" [{j}] type={block_type}", end="")
|
||||
if block_type == "text":
|
||||
text = block.get("text", "")
|
||||
text_preview = text[:60] + "..." if len(text) > 60 else text
|
||||
print(f", text='{text_preview}'")
|
||||
elif block_type == "tool_use":
|
||||
print(f", name={block.get('name')}, id={block.get('id')}, input={block.get('input')}")
|
||||
elif block_type == "tool_result":
|
||||
output = block.get("output", "")
|
||||
output_preview = output[:60] + "..." if len(output) > 60 else output
|
||||
print(f", name={block.get('name')}, id={block.get('id')}, output='{output_preview}'")
|
||||
elif block_type == "thinking":
|
||||
thinking = block.get("thinking", "")
|
||||
thinking_preview = thinking[:60] + "..." if len(thinking) > 60 else thinking
|
||||
print(f", thinking='{thinking_preview}'")
|
||||
else:
|
||||
print()
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("✓ Plain Text List Conversion Test Completed")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
|
||||
def test_multimodal_list_conversion():
|
||||
"""Test converting a long list of multimodal DashScope messages to AgentScope Msgs."""
|
||||
from reme.core.utils.agentscope_utils import convert_dashscope_to_agentscope
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("TEST 2: Multimodal List Conversion (List[Dict] -> List[Msg])")
|
||||
print("=" * 80)
|
||||
|
||||
# Long conversation with multimodal content
|
||||
dashscope_msgs = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "你是一个视觉分析助手,可以分析图片、视频和音频内容。",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"text": "你好!我想让你帮我分析几张照片。"},
|
||||
],
|
||||
"name": "摄影师",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "你好!我很乐意帮你分析照片。请上传你想分析的照片。",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"text": "首先,这是我拍的一张风景照,你觉得构图怎么样?"},
|
||||
{
|
||||
"image": "https://img.alicdn.com/imgextra/i1/O1CN01gDEY8M1W114Hi3XcN_"
|
||||
"!!6000000002727-0-tps-1024-406.jpg",
|
||||
},
|
||||
],
|
||||
"name": "摄影师",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"text": "这张风景照的构图很不错!主要优点包括:\n\n1. 采用了经典的三分法构图\n2. 前景、中景、远景层次分明\n"
|
||||
"3. 色彩饱和度适中,视觉效果舒适\n4. 光线运用得当,明暗对比自然\n\n"
|
||||
"如果要改进的话,可以考虑稍微调整一下地平线的位置。",
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"text": "太感谢了!那这两张照片呢?我想对比一下:"},
|
||||
{"text": "\n第一张:"},
|
||||
{"image": "https://example.com/photo1_sunrise.jpg"},
|
||||
{"text": "\n第二张:"},
|
||||
{"image": "https://example.com/photo2_sunset.jpg"},
|
||||
{"text": "\n它们分别是日出和日落时拍摄的,你觉得哪张效果更好?"},
|
||||
],
|
||||
"name": "摄影师",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"text": "让我对比分析一下这两张照片:\n\n**日出照片(第一张):**\n- 光线柔和,色调偏冷\n"
|
||||
"- 天空呈现淡蓝到橙黄的渐变\n- 画面整体清新明快\n- 适合表现希望和新生的主题\n\n"
|
||||
"**日落照片(第二张):**\n- 光线温暖,色调偏暖\n- 天空呈现金黄到橙红的渐变\n"
|
||||
"- 画面更有戏剧性和情绪感染力\n"
|
||||
"- 适合表现浪漫和感性的主题\n\n"
|
||||
"两张照片各有特色,难分伯仲。如果是为了表现宁静和希望,推荐日出;如果想营造温馨浪漫的氛围,日落会更好。",
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"text": "太专业了!我还拍了一段延时摄影视频,能帮我看看吗?"},
|
||||
{
|
||||
"video": [
|
||||
"https://example.com/timelapse/frame001.jpg",
|
||||
"https://example.com/timelapse/frame002.jpg",
|
||||
"https://example.com/timelapse/frame003.jpg",
|
||||
"https://example.com/timelapse/frame004.jpg",
|
||||
"https://example.com/timelapse/frame005.jpg",
|
||||
],
|
||||
},
|
||||
],
|
||||
"name": "摄影师",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"text": (
|
||||
"这段延时摄影很棒!从视频帧可以看出:\n\n"
|
||||
"1. **节奏控制**:时间间隔掌握得很好,云朵移动流畅自然\n"
|
||||
"2. **曝光稳定**:各帧之间曝光一致,没有明显的闪烁\n"
|
||||
"3. **画面稳定**:机位固定得很稳,没有抖动\n"
|
||||
"4. **动态元素**:云朵的流动和光线的变化增加了画面的生动性\n\n"
|
||||
"建议:\n- 可以尝试更长的拍摄时间,捕捉完整的天色变化\n"
|
||||
"- 考虑加入前景元素(如建筑或树木)增加空间感"
|
||||
),
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"text": "最后还有一段录音,是我在海边录的环境音,你能听听吗?"},
|
||||
{"audio": "https://example.com/audio/beach_ambience.mp3"},
|
||||
],
|
||||
"name": "摄影师",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_audio_001",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "analyze_audio",
|
||||
"arguments": (
|
||||
'{"audio_url": "https://example.com/audio/beach_ambience.mp3", '
|
||||
'"analysis_type": "环境音分析"}'
|
||||
),
|
||||
},
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_audio_001",
|
||||
"name": "analyze_audio",
|
||||
"content": (
|
||||
"音频分析结果:\n- 音频时长:3分24秒\n- 主要声音:海浪拍岸声、海鸥鸣叫、轻微的风声\n"
|
||||
"- 音质:清晰,背景噪音较少\n- 立体声效果:良好,空间感强\n"
|
||||
"- 音量平衡:适中,没有过载或削波\n- 频率分布:低频(海浪)和高频(海鸥)分布均衡"
|
||||
),
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"text": (
|
||||
"你的海边环境音录制得非常专业!根据音频分析:\n\n**技术质量:**\n"
|
||||
"✓ 音质清晰,噪音控制得很好\n✓ 立体声效果出色,能感受到空间感\n✓ 音量平衡适中,听感舒适\n\n"
|
||||
"**内容表现:**\n✓ 成功捕捉了海浪、海鸥、风声等多层次声音\n"
|
||||
"✓ 声音元素丰富但不杂乱\n✓ 很好地营造了宁静放松的海边氛围\n\n**应用建议:**\n"
|
||||
"- 适合用作冥想或放松音乐的背景\n- 可以配合你的海边照片/视频使用\n"
|
||||
"- 建议保留原始文件,方便后期调音\n\n"
|
||||
"总的来说,你在摄影和录音方面都展现了很高的专业水平!"
|
||||
),
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"text": "非常感谢你详细的分析和建议!这对我帮助很大。"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://example.com/thank_you.jpg"},
|
||||
},
|
||||
{"text": "这是我做的一张感谢卡片,送给你!"},
|
||||
],
|
||||
"name": "摄影师",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"text": "谢谢你精美的感谢卡片!很高兴能帮到你。\n\n你的作品都很出色,继续保持这份对摄影和创作的热情!如果以后还有作品想分析或讨论,随时欢迎找我。\n\n祝你创作顺利!📸✨",
|
||||
},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
print(f"\n[Input] DashScope messages: {len(dashscope_msgs)} messages")
|
||||
print(json.dumps(dashscope_msgs, ensure_ascii=False, indent=2))
|
||||
|
||||
# Convert to AgentScope Msgs
|
||||
msgs = convert_dashscope_to_agentscope(dashscope_msgs)
|
||||
|
||||
print(f"\n[Output] AgentScope Msgs: {len(msgs)} messages")
|
||||
print("=" * 80)
|
||||
|
||||
for i, msg in enumerate(msgs):
|
||||
print(f"\n【Message {i+1}/{len(msgs)}】")
|
||||
print(f" name: {msg.name}")
|
||||
print(f" role: {msg.role}")
|
||||
print(f" content type: {type(msg.content).__name__}")
|
||||
print(f" timestamp: {msg.timestamp}")
|
||||
|
||||
if isinstance(msg.content, str):
|
||||
content_preview = msg.content[:100] + "..." if len(msg.content) > 100 else msg.content
|
||||
print(f" content: {content_preview}")
|
||||
elif isinstance(msg.content, list):
|
||||
print(f" content blocks: {len(msg.content)} blocks")
|
||||
for j, block in enumerate(msg.content):
|
||||
block_type = block.get("type")
|
||||
print(f" [{j}] type={block_type}", end="")
|
||||
|
||||
if block_type == "text":
|
||||
text = block.get("text", "")
|
||||
text_preview = text[:50] + "..." if len(text) > 50 else text
|
||||
print(f", text='{text_preview}'")
|
||||
elif block_type == "image":
|
||||
source = block.get("source", {})
|
||||
url = source.get("url", "")
|
||||
url_preview = url[:50] + "..." if len(url) > 50 else url
|
||||
print(f", url='{url_preview}'")
|
||||
elif block_type == "video":
|
||||
source = block.get("source", {})
|
||||
url = source.get("url", "")
|
||||
url_preview = url[:50] + "..." if len(url) > 50 else url
|
||||
print(f", url='{url_preview}'")
|
||||
elif block_type == "audio":
|
||||
source = block.get("source", {})
|
||||
url = source.get("url", "")
|
||||
url_preview = url[:50] + "..." if len(url) > 50 else url
|
||||
print(f", url='{url_preview}'")
|
||||
elif block_type == "tool_use":
|
||||
print(f", name={block.get('name')}, id={block.get('id')}")
|
||||
print(f" input={json.dumps(block.get('input'), ensure_ascii=False)}")
|
||||
elif block_type == "tool_result":
|
||||
output = block.get("output", "")
|
||||
output_preview = output[:50] + "..." if len(output) > 50 else output
|
||||
print(f", name={block.get('name')}, id={block.get('id')}")
|
||||
print(f" output='{output_preview}'")
|
||||
elif block_type == "thinking":
|
||||
thinking = block.get("thinking", "")
|
||||
thinking_preview = thinking[:50] + "..." if len(thinking) > 50 else thinking
|
||||
print(f", thinking='{thinking_preview}'")
|
||||
else:
|
||||
print()
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("✓ Multimodal List Conversion Test Completed")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run both tests
|
||||
test_plain_text_list_conversion()
|
||||
test_multimodal_list_conversion()
|
||||
|
||||
print("\n" + "🎉" * 40)
|
||||
print("All tests completed successfully!")
|
||||
print("🎉" * 40 + "\n")
|
||||
|
|
@ -1,301 +0,0 @@
|
|||
"""Tests for ReMeFs compact interface.
|
||||
|
||||
This module tests the compact() method of ReMeFs class which provides
|
||||
a high-level interface for conversation compaction.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
||||
from reme import ReMeFs
|
||||
from reme.core.enumeration import Role
|
||||
from reme.core.schema import Message
|
||||
|
||||
|
||||
def print_messages(messages: list[Message], title: str = "Messages", max_content_len: int = 150):
|
||||
"""Print messages with their role and content.
|
||||
|
||||
Args:
|
||||
messages: List of messages to print
|
||||
title: Title for the message list
|
||||
max_content_len: Maximum content length to display (truncate if longer)
|
||||
"""
|
||||
print(f"\n{title}: (count: {len(messages)})")
|
||||
print("-" * 80)
|
||||
for i, msg in enumerate(messages):
|
||||
content = str(msg.content)
|
||||
if len(content) > max_content_len:
|
||||
content = content[:max_content_len] + "..."
|
||||
print(f" [{i}] {msg.role.value:10s}: {content}")
|
||||
print("-" * 80)
|
||||
|
||||
|
||||
def create_test_messages(num_messages: int = 10) -> list[Message]:
|
||||
"""Create a list of test messages.
|
||||
|
||||
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:
|
||||
messages.append(
|
||||
Message(
|
||||
role=Role.USER,
|
||||
content=f"User message {i}: Can you help me with task {i}?",
|
||||
),
|
||||
)
|
||||
else:
|
||||
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,
|
||||
),
|
||||
)
|
||||
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...""",
|
||||
),
|
||||
]
|
||||
|
||||
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_compact_below_threshold():
|
||||
"""Test compact() when messages are below threshold.
|
||||
|
||||
Expects: compacted=False, returns original messages
|
||||
"""
|
||||
print("\n" + "=" * 80)
|
||||
print("TEST 1: Compact - Below Threshold (No Compaction)")
|
||||
print("=" * 80)
|
||||
|
||||
reme_fs = ReMeFs(enable_logo=False, vector_store=None)
|
||||
await reme_fs.start()
|
||||
|
||||
messages = create_test_messages(num_messages=4)
|
||||
print_messages(messages, "INPUT MESSAGES", max_content_len=80)
|
||||
|
||||
print("\nParameters:")
|
||||
print(" context_window_tokens: 5000")
|
||||
print(" reserve_tokens: 2000 (threshold = 3000)")
|
||||
print(" keep_recent_tokens: 1000")
|
||||
|
||||
result = await reme_fs.compact(
|
||||
messages=messages,
|
||||
context_window_tokens=5000,
|
||||
reserve_tokens=2000,
|
||||
keep_recent_tokens=1000,
|
||||
)
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print("RESULT:")
|
||||
print(f" compacted: {result.get('compacted')}")
|
||||
print(f" tokens_before: {result.get('tokens_before')}")
|
||||
print(f" is_split_turn: {result.get('is_split_turn')}")
|
||||
|
||||
result_messages = result.get("messages", [])
|
||||
print_messages(result_messages, "OUTPUT MESSAGES", max_content_len=80)
|
||||
|
||||
assert result.get("compacted") is False, "Should not compact below threshold"
|
||||
assert len(result_messages) == len(messages), "Should return all original messages"
|
||||
print("\n✓ TEST PASSED: No compaction below threshold\n")
|
||||
|
||||
await reme_fs.close()
|
||||
|
||||
|
||||
async def test_compact_above_threshold():
|
||||
"""Test compact() when messages exceed threshold.
|
||||
|
||||
Expects: compacted=True, returns summary + left_messages
|
||||
"""
|
||||
print("\n" + "=" * 80)
|
||||
print("TEST 2: Compact - Above Threshold (With Compaction & LLM Summary)")
|
||||
print("=" * 80)
|
||||
|
||||
reme_fs = ReMeFs(enable_logo=False, vector_store=None)
|
||||
await reme_fs.start()
|
||||
|
||||
messages = create_test_messages(num_messages=12)
|
||||
print_messages(messages, "INPUT MESSAGES", max_content_len=60)
|
||||
|
||||
print("\nParameters:")
|
||||
print(" context_window_tokens: 3000")
|
||||
print(" reserve_tokens: 1500 (threshold = 1500)")
|
||||
print(" keep_recent_tokens: 500 (keep only recent messages)")
|
||||
|
||||
result = await reme_fs.compact(
|
||||
messages=messages,
|
||||
context_window_tokens=3000,
|
||||
reserve_tokens=1500,
|
||||
keep_recent_tokens=500,
|
||||
)
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print("RESULT:")
|
||||
print(f" compacted: {result.get('compacted')}")
|
||||
print(f" tokens_before: {result.get('tokens_before')}")
|
||||
print(f" is_split_turn: {result.get('is_split_turn')}")
|
||||
|
||||
result_messages = result.get("messages", [])
|
||||
if result.get("compacted") and result_messages:
|
||||
has_summary = "<summary>" in str(result_messages[0].content)
|
||||
print(f"\n *** First message contains summary: {has_summary}")
|
||||
|
||||
print_messages(result_messages, "OUTPUT MESSAGES (Summary + Recent)", max_content_len=1500)
|
||||
|
||||
assert result.get("compacted") is True, "Should compact above threshold"
|
||||
assert len(result_messages) < len(messages), "Should reduce message count"
|
||||
print("\n✓ TEST PASSED: Compaction triggered and summary generated\n")
|
||||
|
||||
await reme_fs.close()
|
||||
|
||||
|
||||
async def test_compact_split_turn_scenario():
|
||||
"""Test compact() with split turn scenario.
|
||||
|
||||
Expects: is_split_turn=True when cut point is mid-turn
|
||||
"""
|
||||
print("\n" + "=" * 80)
|
||||
print("TEST 3: Compact - Split Turn Scenario (Cut in Middle of Assistant Response)")
|
||||
print("=" * 80)
|
||||
|
||||
reme_fs = ReMeFs(enable_logo=False, vector_store=None)
|
||||
await reme_fs.start()
|
||||
|
||||
messages = []
|
||||
|
||||
# Add initial conversation
|
||||
for i in range(3):
|
||||
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 multi-part assistant response
|
||||
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. " * 50,
|
||||
),
|
||||
)
|
||||
messages.append(
|
||||
Message(
|
||||
role=Role.ASSISTANT,
|
||||
content="This is the continuation of the response. " * 50,
|
||||
),
|
||||
)
|
||||
messages.append(
|
||||
Message(
|
||||
role=Role.ASSISTANT,
|
||||
content="And here's the final part with the conclusion. " * 30,
|
||||
),
|
||||
)
|
||||
|
||||
print_messages(messages, "INPUT MESSAGES", max_content_len=80)
|
||||
|
||||
print("\nParameters:")
|
||||
print(" context_window_tokens: 3000")
|
||||
print(" reserve_tokens: 1000 (threshold = 2000)")
|
||||
print(" keep_recent_tokens: 800 (should cut in middle of assistant responses)")
|
||||
|
||||
result = await reme_fs.compact(
|
||||
messages=messages,
|
||||
context_window_tokens=3000,
|
||||
reserve_tokens=1000,
|
||||
keep_recent_tokens=800,
|
||||
)
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print("RESULT:")
|
||||
print(f" compacted: {result.get('compacted')}")
|
||||
print(f" tokens_before: {result.get('tokens_before')}")
|
||||
print(f" is_split_turn: {result.get('is_split_turn')} *** (should be True)")
|
||||
|
||||
result_messages = result.get("messages", [])
|
||||
print_messages(result_messages, "OUTPUT MESSAGES (Summary with Turn Context + Recent)", max_content_len=150)
|
||||
|
||||
if result.get("is_split_turn"):
|
||||
print("\n✓ TEST PASSED: Split turn correctly detected and handled\n")
|
||||
else:
|
||||
print("\n⚠ WARNING: Split turn not detected (parameters may need adjustment)\n")
|
||||
|
||||
await reme_fs.close()
|
||||
|
||||
|
||||
async def main():
|
||||
"""Run core compact interface tests."""
|
||||
print("\n" + "=" * 80)
|
||||
print("ReMeFs Compact Interface - Core Test Suite")
|
||||
print("=" * 80)
|
||||
print("\nThis test suite demonstrates the three key scenarios of conversation compaction:")
|
||||
print(" 1. Below threshold - no compaction needed")
|
||||
print(" 2. Above threshold - full compaction with LLM summary")
|
||||
print(" 3. Split turn - cut point falls in middle of assistant response")
|
||||
print("=" * 80)
|
||||
|
||||
# Test 1: No compaction (below threshold)
|
||||
await test_compact_below_threshold()
|
||||
|
||||
# Test 2: Full compaction (requires LLM)
|
||||
await test_compact_above_threshold()
|
||||
|
||||
# Test 3: Split turn compaction (requires LLM)
|
||||
await test_compact_split_turn_scenario()
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("All basic tests completed!")
|
||||
print("=" * 80)
|
||||
print("\nNote: Tests requiring LLM calls are commented out.")
|
||||
print("Uncomment them in the main() function to run with actual LLM.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
678
tests/test_fs_compactor.py
Normal file
678
tests/test_fs_compactor.py
Normal file
|
|
@ -0,0 +1,678 @@
|
|||
"""Tests for FsCompactor - conversation history summarization.
|
||||
|
||||
This module tests the summary generation logic of FsCompactor class,
|
||||
which creates compact summaries of conversation history using LLM.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
||||
from reme import ReMeFs
|
||||
from reme.core.enumeration import Role
|
||||
from reme.core.schema import Message
|
||||
|
||||
|
||||
def print_messages(messages: list[Message], title: str = "Messages", max_content_len: int = 150):
|
||||
"""Print messages with their role and content.
|
||||
|
||||
Args:
|
||||
messages: List of messages to print
|
||||
title: Title for the message list
|
||||
max_content_len: Maximum content length to display (truncate if longer)
|
||||
"""
|
||||
print(f"\n{title}: (count: {len(messages)})")
|
||||
print("-" * 80)
|
||||
for i, msg in enumerate(messages):
|
||||
content = str(msg.content)
|
||||
if len(content) > max_content_len:
|
||||
content = content[:max_content_len] + "..."
|
||||
print(f" [{i}] {msg.role.value:10s}: {content}")
|
||||
print("-" * 80)
|
||||
|
||||
|
||||
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...""",
|
||||
),
|
||||
]
|
||||
|
||||
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
|
||||
|
||||
|
||||
def create_realistic_personal_conversation() -> list[Message]:
|
||||
"""Create a realistic conversation with personal information for testing compaction."""
|
||||
messages = [
|
||||
Message(
|
||||
role=Role.USER,
|
||||
content="Hi! I'm planning a trip to Japan next month. I need help organizing my itinerary.",
|
||||
),
|
||||
Message(
|
||||
role=Role.ASSISTANT,
|
||||
content="""\
|
||||
Great! I'd be happy to help you plan your Japan trip. To give you the best recommendations, could you tell me:
|
||||
1. Which cities are you planning to visit?
|
||||
2. How many days will you be there?
|
||||
3. What are your interests (food, culture, nature, shopping, etc.)?
|
||||
4. What's your approximate budget per day?""",
|
||||
),
|
||||
Message(
|
||||
role=Role.USER,
|
||||
content="""\
|
||||
I'm planning to visit Tokyo, Kyoto, and Osaka. The trip will be 14 days, from March 15th to March 29th.
|
||||
My name is Sarah Chen, and I'm 28 years old. I'm traveling with my boyfriend David Miller who is 30.
|
||||
We're both really into Japanese food - sushi, ramen, tempura, you name it! We also love
|
||||
temples, gardens, and want to experience traditional Japanese culture.
|
||||
Our budget is around $200 per day per person, not including accommodation which we've already booked.""",
|
||||
),
|
||||
Message(
|
||||
role=Role.ASSISTANT,
|
||||
content="""Perfect! With 14 days and your interests, you'll have an amazing time.
|
||||
Here's a suggested breakdown:
|
||||
|
||||
**Tokyo (5 days - March 15-19)**
|
||||
- Day 1: Arrive, check into hotel, explore Shibuya and Harajuku
|
||||
- Day 2: Visit Senso-ji Temple in Asakusa, explore Ueno Park
|
||||
- Day 3: Tsukiji Outer Market for breakfast, Imperial Palace, Ginza shopping
|
||||
- Day 4: Day trip to Nikko (UNESCO World Heritage temples)
|
||||
- Day 5: TeamLab Borderless, Odaiba, evening in Roppongi
|
||||
|
||||
**Kyoto (6 days - March 20-25)**
|
||||
- Day 6: Travel to Kyoto, Fushimi Inari Shrine (iconic red torii gates)
|
||||
- Day 7: Arashiyama Bamboo Grove, Tenryu-ji Temple, Monkey Park
|
||||
- Day 8: Kinkaku-ji (Golden Pavilion), Ryoan-ji rock garden
|
||||
- Day 9: Gion district, traditional tea ceremony, Kiyomizu-dera Temple
|
||||
- Day 10: Day trip to Nara (deer park, Todai-ji Temple)
|
||||
- Day 11: Philosopher's Path, Ginkaku-ji (Silver Pavilion), Heian Shrine
|
||||
|
||||
**Osaka (3 days - March 26-28)**
|
||||
- Day 12: Travel to Osaka, Osaka Castle, Dotonbori food street
|
||||
- Day 13: Kuromon Market, Shinsaibashi shopping, Universal Studios Japan option
|
||||
- Day 14: Sumiyoshi Taisha Shrine, last-minute shopping, departure prep
|
||||
|
||||
Would you like me to suggest specific restaurants or any other details?""",
|
||||
),
|
||||
Message(
|
||||
role=Role.USER,
|
||||
content="""\
|
||||
This looks fantastic! Yes, please suggest restaurants. We're especially interested in
|
||||
authentic local places, not touristy spots.
|
||||
Also, I should mention that David has a shellfish allergy, so we need to be careful
|
||||
about that.
|
||||
My email is sarah.chen@email.com if you want to send me a detailed PDF later.
|
||||
My phone is +1-555-0123-4567.
|
||||
We're staying at the Park Hyatt Tokyo (5 days), traditional ryokan in Kyoto called
|
||||
Yoshikawa Inn (6 days), and Hotel Granvia Osaka (3 days).""",
|
||||
),
|
||||
Message(
|
||||
role=Role.ASSISTANT,
|
||||
content="""Excellent accommodation choices! The Yoshikawa Inn is particularly special
|
||||
- they have an amazing kaiseki restaurant.
|
||||
|
||||
**Restaurant Recommendations:**
|
||||
|
||||
**Tokyo:**
|
||||
1. **Ichiran Ramen** (Shibuya) - Private booth ramen experience, customize your bowl. """
|
||||
"""Shellfish-free option available.
|
||||
2. **Sushi Dai** (Toyosu Market) - Early morning sushi breakfast, tell them about the
|
||||
shellfish allergy """
|
||||
"""("ebi to kai rui arerugī")
|
||||
3. **Tempura Kondo** (Ginza) - Michelin-starred, vegetable-focused tempura
|
||||
4. **Narisawa** (Aoyama) - If you want a splurge, this 2-Michelin-star restaurant has
|
||||
innovative Japanese cuisine
|
||||
5. **Afuri Ramen** (Harajuku) - Yuzu-based light ramen, very refreshing
|
||||
|
||||
**Kyoto:**
|
||||
1. **Yoshikawa Inn's own restaurant** - Their kaiseki is exceptional, """
|
||||
"""inform them of the allergy when booking
|
||||
2. **Omen** (near Ginkaku-ji) - Handmade udon in a traditional house
|
||||
3. **Kikunoi Roan** - Kaiseki experience, they're very accommodating with allergies
|
||||
4. **Katsukura** (Sanjo) - Best tonkatsu (pork cutlet) in Kyoto
|
||||
5. **Nishiki Market** - Walk through for street food, but David should be cautious with mystery items
|
||||
|
||||
**Osaka:**
|
||||
1. **Harukoma Sushi** (Tenjinbashisuji) - Conveyor belt sushi, excellent quality, affordable
|
||||
2. **Fukutaro** (Dotonbori) - Okonomiyaki specialists, no shellfish in standard recipes
|
||||
3. **Matsusakagyu Yakiniku M** (Hozenji Yokocho) - High-end yakiniku (grilled beef)
|
||||
4. **Ichiran Ramen** (Dotonbori branch) - Consistent quality
|
||||
5. **Kuromon Market** - Various stalls, ask vendors about ingredients
|
||||
|
||||
**Important Allergy Notes:**
|
||||
- Learn this phrase: "Ebi to kai rui arerugī ga arimasu"
|
||||
(I have a shellfish allergy)
|
||||
- Many broths use shellfish stock, always ask
|
||||
- Carry allergy cards in Japanese
|
||||
- Download Google Translate offline for emergency communication
|
||||
|
||||
For your trip planning, I'll note: Sarah Chen (sarah.chen@email.com, +1-555-0123-4567)
|
||||
and David Miller, March 15-29 Japan trip.""",
|
||||
),
|
||||
Message(
|
||||
role=Role.USER,
|
||||
content="""\
|
||||
Thank you so much! A few more questions:
|
||||
1. What's the best way to get from Tokyo to Kyoto? We heard about the JR Pass.
|
||||
2. Should we rent a pocket WiFi or buy a SIM card?
|
||||
3. Any cultural etiquette we should be aware of?
|
||||
4. David wants to know if there are good bars or nightlife spots. He loves whisky.
|
||||
""",
|
||||
),
|
||||
Message(
|
||||
role=Role.ASSISTANT,
|
||||
content="""Great questions! Let me address each:
|
||||
|
||||
**1. Transportation:**
|
||||
Absolutely get the JR Pass! For your 14-day trip, get the 14-day pass (around ¥47,250/$315 per person).
|
||||
- Tokyo to Kyoto: Shinkansen (bullet train), 2.5 hours, covered by JR Pass
|
||||
- Kyoto to Osaka: 30 minutes by train, also covered
|
||||
- The pass also covers JR trains in cities and the Nikko trip
|
||||
- Order it before leaving the US, pick up at airport upon arrival
|
||||
|
||||
**2. Internet:**
|
||||
I recommend renting pocket WiFi over SIM card because:
|
||||
- You can share it (both Sarah and David can connect)
|
||||
- No need to swap your phone SIM
|
||||
- Rent from companies like Japan Wireless or Ninja WiFi
|
||||
- Reserve online, pick up at Narita/Haneda airport
|
||||
- Cost: around $8-10/day for unlimited data
|
||||
- Return at airport when you leave
|
||||
|
||||
**3. Cultural Etiquette:**
|
||||
- Take off shoes when entering homes, some restaurants, temples (you'll see slippers at entrance)
|
||||
- Don't tip - it can be considered insulting
|
||||
- At temples: bow before entering, don't take photos inside unless permitted
|
||||
- On trains: silent mode on phones, no phone calls
|
||||
- Eating: slurping noodles is okay (shows appreciation), don't stick chopsticks upright in rice
|
||||
- Onsen (if your ryokan has one): shower completely before entering communal bath, no swimwear
|
||||
- Walking: don't eat while walking, finish at the stall or sit down
|
||||
- Trash: carry your trash with you, public bins are rare
|
||||
|
||||
**4. Bars & Whisky for David:**
|
||||
|
||||
**Tokyo:**
|
||||
1. **Zoetrope** (Shinjuku) - 300+ types of Japanese whisky, cozy atmosphere
|
||||
2. **Bar Benfiddich** (Shinjuku) - Award-winning cocktail bar, house-made bitters
|
||||
3. **Mori Bar** (Ginza) - Classic Japanese bar, extensive whisky collection
|
||||
4. **Golden Gai** (Shinjuku) - Alley with 200+ tiny bars, unique atmosphere
|
||||
|
||||
**Kyoto:**
|
||||
1. **Sake Bar Yoramu** - Also has excellent whisky selection
|
||||
2. **Rocking Bar Ing** - Relaxed atmosphere, good whisky menu
|
||||
3. **Bar K6** (Kawaramachi) - Intimate whisky bar
|
||||
|
||||
**Osaka:**
|
||||
1. **Bar Nayuta** (Namba) - World-class whisky bar
|
||||
2. **Moonshine Karaoke Bar** (Dotonbori) - If you want karaoke with drinks
|
||||
3. **Bar Americana** - Extensive bourbon and whisky selection
|
||||
|
||||
**Whisky Shopping:**
|
||||
- **Tanakaya** (Tokyo) - Huge selection of Japanese whisky
|
||||
- **Liquors Hasegawa** (Tokyo Station) - Great for last-minute bottles before flying home
|
||||
- Note: Japanese whisky (Yamazaki, Hibiki, Hakushu) can be expensive even in Japan due to global demand
|
||||
|
||||
Do you need help with anything else for your trip? Train routes, more specific timing, or """
|
||||
"""budget breakdown?""",
|
||||
),
|
||||
Message(
|
||||
role=Role.USER,
|
||||
content="""\
|
||||
This is incredibly helpful! One more thing - we want to attend a traditional tea
|
||||
ceremony in Kyoto. Can you recommend where to do this?
|
||||
Also, what should we pack? I know March can be unpredictable weather-wise.
|
||||
Oh, and my birthday is on March 22nd - any special restaurant recommendation for
|
||||
that evening? It'll be our 3rd anniversary too!
|
||||
David's credit card is Visa ending in 4892, and mine is Mastercard ending in 7651
|
||||
- will these work everywhere in Japan?""",
|
||||
),
|
||||
Message(
|
||||
role=Role.ASSISTANT,
|
||||
content="""Wonderful questions! And happy early birthday & anniversary!
|
||||
|
||||
**Tea Ceremony Experiences in Kyoto:**
|
||||
|
||||
1. **Camellia Tea Ceremony** (Higashiyama) - Highly recommended!
|
||||
- English-speaking host
|
||||
- Includes kimono wearing experience
|
||||
- Small groups (max 6 people)
|
||||
- 2-hour experience, around ¥6,000 per person
|
||||
- Book online in advance: camelliatearoom.com
|
||||
|
||||
2. **En Tea Ceremony** (near Kiyomizu-dera)
|
||||
- Traditional machiya (townhouse) setting
|
||||
- Very authentic, less touristy
|
||||
- Private ceremony option available
|
||||
- Around ¥5,500 per person
|
||||
|
||||
3. **Wak Japan** (Gion area)
|
||||
- Combines tea ceremony with flower arrangement or calligraphy
|
||||
- Good for couples
|
||||
- Around ¥8,000 per person for combined experience
|
||||
|
||||
I recommend booking Day 9 (March 22nd) morning for the tea ceremony, then evening for your special dinner!
|
||||
|
||||
**March Weather & Packing:**
|
||||
March in Japan: transitioning from winter to spring
|
||||
- Temperature: 8-15°C (46-59°F)
|
||||
- Cherry blossoms might just start blooming late March (you might catch early bloomers!)
|
||||
|
||||
**Pack:**
|
||||
- Layering clothes: light sweater, cardigan, light jacket
|
||||
- One warmer jacket for evenings
|
||||
- Comfortable walking shoes (you'll walk 15,000+ steps daily)
|
||||
- Umbrella (March has occasional rain)
|
||||
- Slip-on shoes (easier for temple visits)
|
||||
- Nice outfit for fancy restaurants
|
||||
- Power adapter (Japan uses Type A plugs, 100V)
|
||||
- Portable charger for phones
|
||||
- Small day backpack
|
||||
|
||||
**Birthday & Anniversary Dinner - March 22nd:**
|
||||
|
||||
For such a special occasion in Kyoto, I highly recommend:
|
||||
|
||||
**Kikunoi Honten** (Main Branch) - 3 Michelin Stars
|
||||
- Ultimate kaiseki experience
|
||||
- Beautiful traditional setting with garden views
|
||||
- Multi-course seasonal menu
|
||||
- Reserve 1-2 months in advance
|
||||
- Budget: ¥25,000-40,000 per person (worth it for anniversary!)
|
||||
- Inform them of: shellfish allergy, birthday, anniversary
|
||||
- They'll often add special touches for celebrations
|
||||
|
||||
**Alternative (still amazing):**
|
||||
**Hyotei** - 3 Michelin Stars
|
||||
- Historic (400 years old)
|
||||
- Beautiful garden setting
|
||||
- Breakfast kaiseki is famous, but dinner is exceptional
|
||||
- Similar price range
|
||||
|
||||
**More Accessible but Still Special:**
|
||||
**Gion Karyo**
|
||||
- 1 Michelin Star
|
||||
- Beautiful Gion location
|
||||
- Around ¥15,000-20,000 per person
|
||||
- Easier to book
|
||||
|
||||
**Birthday Cake:**
|
||||
Order a birthday cake from **Malebranche** (Kyoto) - they do beautiful matcha cakes.
|
||||
Ask your ryokan to keep it refrigerated and present it after you return from dinner!
|
||||
|
||||
**Credit Cards:**
|
||||
- Major hotels and restaurants: Visa & Mastercard widely accepted
|
||||
(Sarah's Mastercard 7651 and David's Visa 4892 should work fine)
|
||||
- However, Japan is still heavily CASH-based
|
||||
- Bring cash for:
|
||||
- Small restaurants
|
||||
- Temples/shrines
|
||||
- Street food
|
||||
- Some shops
|
||||
- Taxis
|
||||
- Vending machines
|
||||
|
||||
**ATM Tips:**
|
||||
- 7-Eleven ATMs work with foreign cards (open 24/7)
|
||||
- Japan Post Bank ATMs also reliable
|
||||
- Withdraw larger amounts to minimize fees
|
||||
- Keep ¥10,000-20,000 ($70-140) cash daily
|
||||
- Your banks: notify them of Japan travel to avoid card blocks
|
||||
|
||||
**Pre-Trip Checklist:**
|
||||
✅ JR Pass ordered
|
||||
✅ Pocket WiFi reserved
|
||||
✅ Kikunoi dinner reservation (book NOW for March 22nd)
|
||||
✅ Tea ceremony booking
|
||||
✅ Notify banks: Sarah (MC 7651), David (Visa 4892)
|
||||
✅ Travel insurance
|
||||
✅ Download offline maps
|
||||
✅ Book airport transfers
|
||||
✅ Birthday cake ordered
|
||||
|
||||
Would you like me to help with a day-by-day detailed schedule with specific times and addresses?""",
|
||||
),
|
||||
Message(
|
||||
role=Role.USER,
|
||||
content="""\
|
||||
Yes! A detailed daily schedule would be amazing. Can you include specific addresses and opening hours?
|
||||
Also, we're thinking of proposing to our friends to join us - my best friend Jessica Wang and her """
|
||||
"""husband Tom. What would change if we're traveling as a group of 4?""",
|
||||
),
|
||||
Message(
|
||||
role=Role.ASSISTANT,
|
||||
content="""\
|
||||
Absolutely! I'll create a detailed schedule. Traveling as a group of 4 (Sarah & David + Jessica & """
|
||||
"""Tom) actually has some advantages:
|
||||
|
||||
**Group of 4 Benefits:**
|
||||
✅ Split pocket WiFi cost (1 device for 4 people)
|
||||
✅ Share taxi costs (taxis fit 4 people)
|
||||
✅ Some restaurants have minimum 4-person set menus
|
||||
✅ Private tea ceremony for your group
|
||||
✅ Better for group photos!
|
||||
|
||||
**Considerations:**
|
||||
- Book restaurants for 4 people
|
||||
- Some tiny bars in Golden Gai might not fit all
|
||||
- Reserve 2 rooms/apartments when needed
|
||||
- Coordinate meeting points if you split up
|
||||
|
||||
**DETAILED 14-DAY SCHEDULE WITH ADDRESSES:**
|
||||
|
||||
**DAY 1 - March 15 (Friday) - TOKYO ARRIVAL**
|
||||
|
||||
*Morning/Afternoon:*
|
||||
- Arrive Narita/Haneda Airport
|
||||
- Pick up: JR Pass, Pocket WiFi
|
||||
- Exchange yen at airport (recommend ¥50,000+ per person)
|
||||
- Take train to hotel: Park Hyatt Tokyo
|
||||
📍 3-7-1-2 Nishishinjuku, Shinjuku-ku, Tokyo 163-1055
|
||||
🚇 Shinjuku Station → Oedo Line to Tochomae Station (5 min walk)
|
||||
|
||||
*Evening (6:00 PM - 9:00 PM):*
|
||||
- Check in, rest, freshen up
|
||||
- Dinner: **Omoide Yokocho** (Memory Lane)
|
||||
📍 1 Chome Nishishinjuku, Shinjuku-ku, Tokyo
|
||||
🕒 Open till midnight
|
||||
💴 ¥2,000-3,000/person
|
||||
- Narrow alley with small yakitori stands
|
||||
- Cash only, very local atmosphere
|
||||
- Ask about shellfish ("kai rui") in skewers
|
||||
|
||||
*Night:*
|
||||
- Walk around Shinjuku, see the night lights
|
||||
- Convenience store snacks (7-Eleven/Family Mart)
|
||||
- Early sleep (jet lag)
|
||||
|
||||
---
|
||||
|
||||
**DAY 2 - March 16 (Saturday) - ASAKUSA & UENO**
|
||||
|
||||
*Morning (9:00 AM - 12:00 PM):*
|
||||
- Breakfast at hotel or nearby bakery
|
||||
- 🚇 Train to Asakusa (30 min from Shinjuku)
|
||||
|
||||
- **Senso-ji Temple**
|
||||
📍 2-3-1 Asakusa, Taito-ku, Tokyo
|
||||
🕒 6:00 AM - 5:00 PM (grounds always open)
|
||||
💴 Free
|
||||
- Arrive by 9:30 AM to avoid crowds
|
||||
- Walk through Kaminarimon Gate, Nakamise Shopping Street
|
||||
- Draw fortune (omikuji) - ¥100
|
||||
- Visit main hall, incense burner
|
||||
|
||||
*Lunch (12:00 PM):*
|
||||
- **Daikokuya Tempura**
|
||||
📍 1-38-10 Asakusa, Taito-ku, Tokyo
|
||||
🕒 11:00 AM - 8:30 PM (closed Mon)
|
||||
💴 ¥2,000-3,000/person
|
||||
- Famous tendon (tempura rice bowl)
|
||||
- Mention shellfish allergy to David's order
|
||||
|
||||
*Afternoon (1:30 PM - 5:00 PM):*
|
||||
- Walk to Ueno (15 min) or train (2 stops)
|
||||
|
||||
- **Ueno Park**
|
||||
📍 Uenokoen, Taito-ku, Tokyo
|
||||
🕒 5:00 AM - 11:00 PM
|
||||
💴 Free (museums extra)
|
||||
- Cherry blossom trees (might see early bloomers!)
|
||||
- Visit **Tokyo National Museum** if interested
|
||||
🕒 9:30 AM - 5:00 PM (closed Mon)
|
||||
💴 ¥1,000/person
|
||||
|
||||
- **Ameya-Yokocho Market**
|
||||
📍 4 Chome Ueno, Taito-ku, Tokyo
|
||||
- Shopping street, bargain clothes, snacks
|
||||
|
||||
*Dinner (6:30 PM):*
|
||||
- **Ichiran Ramen Ueno**
|
||||
📍 6-11-11 Ueno, Taito-ku, Tokyo
|
||||
🕒 24 hours
|
||||
💴 ¥1,000-1,500/person
|
||||
- Individual booth experience
|
||||
- Order via vending machine (English available)
|
||||
- Customize your ramen
|
||||
|
||||
*Night:*
|
||||
- Return to Shinjuku
|
||||
- Optional: **Zoetrope Whisky Bar** for David & Tom
|
||||
📍 Sankoubldg. 3F, 1-7-10 Nishi-Shinjuku, Shinjuku-ku
|
||||
🕒 6:00 PM - 12:00 AM (closed Sun)
|
||||
💴 ¥1,500-3,000/drink
|
||||
|
||||
---
|
||||
|
||||
**DAY 3 - March 17 (Sunday) - TSUKIJI, GINZA, IMPERIAL PALACE**
|
||||
|
||||
*Early Morning (5:30 AM - 8:00 AM):*
|
||||
- Wake up early!
|
||||
- **Tsukiji Outer Market**
|
||||
📍 4 Chome Tsukiji, Chuo-ku, Tokyo
|
||||
🚇 Tsukijishijo Station (Oedo Line)
|
||||
🕒 Most stalls: 5:00 AM - 2:00 PM
|
||||
|
||||
- Breakfast at **Sushi Dai** (or Daiwa Sushi)
|
||||
📍 Inside Toyosu Market (new location)
|
||||
🕒 5:30 AM - 1:30 PM
|
||||
💴 ¥3,500-5,000/person
|
||||
⚠️ Expect 1-2 hour wait, go early!
|
||||
- Tell chef about David's shellfish allergy
|
||||
- Omakase sushi breakfast
|
||||
|
||||
*Late Morning (9:00 AM - 12:00 PM):*
|
||||
- **Imperial Palace East Gardens**
|
||||
📍 1-1 Chiyoda, Chiyoda-ku, Tokyo
|
||||
🕒 9:00 AM - 4:30 PM (closed Mon, Fri)
|
||||
💴 Free
|
||||
- Beautiful gardens, historic site
|
||||
- 1-1.5 hour visit
|
||||
|
||||
*Lunch (12:30 PM):*
|
||||
- **Ginza**
|
||||
📍 Ginza, Chuo-ku, Tokyo
|
||||
- Many options for lunch
|
||||
|
||||
- **Tempura Kondo** (if you can get reservation)
|
||||
📍 Sakaguchi Bldg. 9F, 5-5-13 Ginza, Chuo-ku
|
||||
🕒 Lunch 12:00-2:00 PM, Dinner 5:30-9:00 PM (closed Sun)
|
||||
💴 Lunch ¥8,000-12,000/person
|
||||
- Reserve online or call: +81-3-5568-0923
|
||||
|
||||
- **Backup: Ginza Kagari Ramen**
|
||||
📍 Ginza, Chuo-ku (search exact location)
|
||||
💴 ¥1,200/person
|
||||
- Creamy chicken ramen, no shellfish
|
||||
|
||||
*Afternoon (2:00 PM - 6:00 PM):*
|
||||
- **Ginza Shopping**
|
||||
- UNIQLO flagship (12 floors)
|
||||
- Mitsukoshi Department Store
|
||||
- Dover Street Market (avant-garde fashion)
|
||||
- MUJI flagship
|
||||
- Window shop luxury brands
|
||||
|
||||
*Dinner (6:30 PM):*
|
||||
- **Afuri Ramen**
|
||||
📍 1-1-7 Ebisu, Shibuya-ku, Tokyo (Ebisu location)
|
||||
🕒 11:00 AM - 11:00 PM
|
||||
💴 ¥1,200/person
|
||||
- Yuzu-salt ramen, light and refreshing
|
||||
|
||||
*Night:*
|
||||
- Train to Shibuya for evening walk
|
||||
- See Shibuya Crossing at night
|
||||
- Return hotel
|
||||
|
||||
---
|
||||
|
||||
This is getting quite long! Should I continue with the rest of the days (Days 4-14)? I can also send """
|
||||
"""you this as a Google Doc or PDF if that's easier. Just need to confirm - are Jessica and Tom """
|
||||
"""definitely joining, or still maybe?
|
||||
|
||||
Also, does anyone have other dietary restrictions besides David's shellfish allergy? And what are """
|
||||
"""your hotel/ryokan confirmations - should I include check-in/check-out timing?""",
|
||||
),
|
||||
]
|
||||
|
||||
return messages
|
||||
|
||||
|
||||
async def test_full_compact_with_summary():
|
||||
"""Test complete compaction flow with LLM summary generation.
|
||||
|
||||
This is a complex integration test that exercises the full compaction pipeline:
|
||||
1. Create a long conversation that exceeds token threshold
|
||||
2. Context checker finds cut point (may include split turn detection)
|
||||
3. Compactor generates summary for messages to summarize
|
||||
4. Compactor handles turn prefix if split turn detected
|
||||
5. Final output contains summary + recent messages
|
||||
|
||||
Expects:
|
||||
- compacted=True
|
||||
- Summary message generated with proper format
|
||||
- Split turn handling if applicable
|
||||
- Reduced message count
|
||||
- Token count within limits
|
||||
"""
|
||||
print("\n" + "=" * 80)
|
||||
print("TEST: Full Compaction with LLM Summary Generation")
|
||||
print("=" * 80)
|
||||
|
||||
reme_fs = ReMeFs(
|
||||
enable_logo=False,
|
||||
vector_store=None,
|
||||
compact_params={
|
||||
"context_window_tokens": 3000,
|
||||
"reserve_tokens": 1500,
|
||||
"keep_recent_tokens": 500,
|
||||
},
|
||||
)
|
||||
await reme_fs.start()
|
||||
|
||||
messages = create_long_conversation()
|
||||
print_messages(messages, "INPUT MESSAGES (Long Conversation)", max_content_len=60)
|
||||
|
||||
print("\nParameters:")
|
||||
print(" context_window_tokens: 3000")
|
||||
print(" reserve_tokens: 1500 (threshold = 1500)")
|
||||
print(" keep_recent_tokens: 500")
|
||||
print("\nExpectations:")
|
||||
print(" - Token count exceeds threshold")
|
||||
print(" - Context checker finds cut point")
|
||||
print(" - Compactor generates summary via LLM")
|
||||
print(" - May detect split turn scenario")
|
||||
print(" - Returns summary + recent messages")
|
||||
|
||||
# Execute full compact flow
|
||||
result = await reme_fs.compact(messages_to_summarize=messages)
|
||||
|
||||
print(f"\n{'=' * 80}")
|
||||
print("RESULT:")
|
||||
print(f" compacted: {result}")
|
||||
await reme_fs.close()
|
||||
|
||||
|
||||
async def test_realistic_personal_conversation_compact():
|
||||
"""Test compaction with realistic personal conversation and return summary string.
|
||||
|
||||
This test:
|
||||
1. Creates a realistic conversation with personal details
|
||||
2. Runs compaction to generate a summary
|
||||
3. Returns the summary as a string
|
||||
4. Validates the compaction result
|
||||
"""
|
||||
print("\n" + "=" * 80)
|
||||
print("TEST: Realistic Personal Conversation Compaction")
|
||||
print("=" * 80)
|
||||
|
||||
reme_fs = ReMeFs(
|
||||
enable_logo=False,
|
||||
vector_store=None,
|
||||
compact_params={
|
||||
"context_window_tokens": 4000,
|
||||
"reserve_tokens": 2000,
|
||||
"keep_recent_tokens": 800,
|
||||
},
|
||||
)
|
||||
await reme_fs.start()
|
||||
|
||||
messages = create_realistic_personal_conversation()
|
||||
print_messages(messages, "INPUT: Realistic Personal Conversation", max_content_len=100)
|
||||
|
||||
print(f"\n{'=' * 80}")
|
||||
print("COMPACTING CONVERSATION...")
|
||||
print(f"{'=' * 80}")
|
||||
|
||||
# Execute compaction
|
||||
result = await reme_fs.compact(messages_to_summarize=messages)
|
||||
|
||||
print(f"\n{'=' * 80}")
|
||||
print("COMPACTION RESULT:")
|
||||
print(f"{'=' * 80}")
|
||||
print(f" compacted: {result}")
|
||||
await reme_fs.close()
|
||||
|
||||
|
||||
async def main():
|
||||
"""Run compactor tests."""
|
||||
print("\n" + "=" * 80)
|
||||
print("FsCompactor - Summary Generation Test Suite")
|
||||
print("=" * 80)
|
||||
print("\nThis test suite validates the LLM-based summarization:")
|
||||
print(" - Full compaction flow (context check + summary generation)")
|
||||
print(" - Summary format and structure")
|
||||
print(" - Split turn handling")
|
||||
print(" - Message preservation")
|
||||
print(" - Realistic personal conversation compaction")
|
||||
print("=" * 80)
|
||||
print("\nNote: This test requires LLM access and may take some time.")
|
||||
print("=" * 80)
|
||||
|
||||
# Run the comprehensive compaction test
|
||||
await test_full_compact_with_summary()
|
||||
|
||||
# Run the realistic personal conversation test
|
||||
await test_realistic_personal_conversation_compact()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
291
tests/test_fs_context_checker.py
Normal file
291
tests/test_fs_context_checker.py
Normal file
|
|
@ -0,0 +1,291 @@
|
|||
"""Tests for FsContextChecker - context window limit checking and cut point finding.
|
||||
|
||||
This module tests the cut point finding logic of FsContextChecker class,
|
||||
which determines where to split conversation history when token limits are exceeded.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
||||
from reme import ReMeFs
|
||||
from reme.core.enumeration import Role
|
||||
from reme.core.schema import Message
|
||||
|
||||
|
||||
def print_messages(messages: list[Message], title: str = "Messages", max_content_len: int = 150):
|
||||
"""Print messages with their role and content.
|
||||
|
||||
Args:
|
||||
messages: List of messages to print
|
||||
title: Title for the message list
|
||||
max_content_len: Maximum content length to display (truncate if longer)
|
||||
"""
|
||||
print(f"\n{title}: (count: {len(messages)})")
|
||||
print("-" * 80)
|
||||
for i, msg in enumerate(messages):
|
||||
content = str(msg.content)
|
||||
if len(content) > max_content_len:
|
||||
content = content[:max_content_len] + "..."
|
||||
print(f" [{i}] {msg.role.value:10s}: {content}")
|
||||
print("-" * 80)
|
||||
|
||||
|
||||
def create_test_messages(num_messages: int = 10) -> list[Message]:
|
||||
"""Create a list of test messages.
|
||||
|
||||
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:
|
||||
messages.append(
|
||||
Message(
|
||||
role=Role.USER,
|
||||
content=f"User message {i}: Can you help me with task {i}?",
|
||||
),
|
||||
)
|
||||
else:
|
||||
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,
|
||||
),
|
||||
)
|
||||
return messages
|
||||
|
||||
|
||||
async def test_no_compaction_needed():
|
||||
"""Test 1: Below threshold - no compaction needed.
|
||||
|
||||
Expects: needs_compaction=False, returns original messages
|
||||
"""
|
||||
print("\n" + "=" * 80)
|
||||
print("TEST 1: Below Threshold - No Cut Point Needed")
|
||||
print("=" * 80)
|
||||
|
||||
reme_fs = ReMeFs(
|
||||
"vector_stores={}", # Override config to disable vector stores
|
||||
enable_logo=False,
|
||||
context_window_tokens=5000,
|
||||
reserve_tokens=2000,
|
||||
keep_recent_tokens=1000,
|
||||
)
|
||||
await reme_fs.start()
|
||||
|
||||
messages = create_test_messages(num_messages=4)
|
||||
print_messages(messages, "INPUT MESSAGES", max_content_len=80)
|
||||
|
||||
print("\nParameters:")
|
||||
print(" context_window_tokens: 5000")
|
||||
print(" reserve_tokens: 2000 (threshold = 3000)")
|
||||
print(" keep_recent_tokens: 1000")
|
||||
|
||||
# Use the new context_check method
|
||||
result = await reme_fs.context_check(messages)
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print("RESULT:")
|
||||
print(f" needs_compaction: {result.get('needs_compaction')}")
|
||||
print(f" token_count: {result.get('token_count')}")
|
||||
print(f" threshold: {result.get('threshold')}")
|
||||
print(f" cut_index: {result.get('cut_index')}")
|
||||
print(f" is_split_turn: {result.get('is_split_turn')}")
|
||||
|
||||
assert result.get("needs_compaction") is False, "Should not need compaction below threshold"
|
||||
assert result.get("left_messages") is not None, "Should return all messages in left_messages"
|
||||
print("\n✓ TEST PASSED: No cut point needed below threshold\n")
|
||||
|
||||
await reme_fs.close()
|
||||
|
||||
|
||||
async def test_compaction_needed_above_threshold():
|
||||
"""Test 2: Compaction needed when exceeding threshold.
|
||||
|
||||
When messages exceed threshold, compaction should be triggered.
|
||||
The cut point location depends on token estimation.
|
||||
"""
|
||||
print("\n" + "=" * 80)
|
||||
print("TEST 2: Compaction Needed Above Threshold")
|
||||
print("=" * 80)
|
||||
|
||||
reme_fs = ReMeFs(
|
||||
"vector_stores={}", # Override config to disable vector stores
|
||||
enable_logo=False,
|
||||
context_window_tokens=1500,
|
||||
reserve_tokens=700, # threshold = 800 (below 892 tokens)
|
||||
keep_recent_tokens=220, # Increased to hit next user message (index 40)
|
||||
)
|
||||
await reme_fs.start()
|
||||
|
||||
# Create simple, short messages with uniform size for predictable cutting
|
||||
messages = []
|
||||
for i in range(50): # More messages to exceed threshold
|
||||
if i % 2 == 0:
|
||||
messages.append(Message(role=Role.USER, content=f"Question {i}?"))
|
||||
else:
|
||||
messages.append(Message(role=Role.ASSISTANT, content=f"Answer {i}: " + "details " * 15)) # Longer assistant
|
||||
|
||||
print_messages(messages, "INPUT MESSAGES", max_content_len=40)
|
||||
|
||||
print("\nParameters:")
|
||||
print(" context_window_tokens: 1500")
|
||||
print(" reserve_tokens: 700 (threshold = 800)")
|
||||
print(" keep_recent_tokens: 220 (should cut at a user message)")
|
||||
|
||||
# Use the new context_check method
|
||||
result = await reme_fs.context_check(messages)
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print("RESULT:")
|
||||
print(f" needs_compaction: {result.get('needs_compaction')}")
|
||||
print(f" token_count: {result.get('token_count')}")
|
||||
print(f" threshold: {result.get('threshold')}")
|
||||
print(f" cut_index: {result.get('cut_index')}")
|
||||
print(f" is_split_turn: {result.get('is_split_turn')}")
|
||||
print(f" accumulated_tokens: {result.get('accumulated_tokens')}")
|
||||
|
||||
messages_to_summarize = result.get("messages_to_summarize", [])
|
||||
left_messages = result.get("left_messages", [])
|
||||
print(f"\n Messages to summarize: {len(messages_to_summarize)}")
|
||||
print(f" Left messages: {len(left_messages)}")
|
||||
|
||||
# Print cut message role for debugging
|
||||
if result.get("cut_index") is not None:
|
||||
cut_idx = result.get("cut_index")
|
||||
if cut_idx < len(messages):
|
||||
print(f" Cut message role: {messages[cut_idx].role.value}")
|
||||
|
||||
assert result.get("needs_compaction") is True, "Should need compaction"
|
||||
# Note: Due to token estimation variability, may or may not be a split turn
|
||||
# The important part is that compaction is triggered
|
||||
assert len(messages_to_summarize) > 0, "Should have messages to summarize"
|
||||
assert len(left_messages) > 0, "Should have left messages"
|
||||
print(f"\n Detected split_turn: {result.get('is_split_turn')}")
|
||||
print("\n✓ TEST PASSED: Compaction triggered when exceeding threshold\n")
|
||||
|
||||
await reme_fs.close()
|
||||
|
||||
|
||||
async def test_split_turn_scenario():
|
||||
"""Test 3: Split turn - cut point in middle of assistant response.
|
||||
|
||||
When cut point lands on an assistant message, we need to find the turn start
|
||||
and handle turn prefix separately.
|
||||
Expects: is_split_turn=True, has turn_prefix_messages
|
||||
"""
|
||||
print("\n" + "=" * 80)
|
||||
print("TEST 3: Split Turn - Cut in Middle of Assistant Response")
|
||||
print("=" * 80)
|
||||
|
||||
reme_fs = ReMeFs(
|
||||
"vector_stores={}", # Override config to disable vector stores
|
||||
enable_logo=False,
|
||||
context_window_tokens=2000,
|
||||
reserve_tokens=300,
|
||||
keep_recent_tokens=600,
|
||||
)
|
||||
await reme_fs.start()
|
||||
|
||||
messages = []
|
||||
|
||||
# Add initial conversation
|
||||
for i in range(3):
|
||||
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 multi-part assistant response
|
||||
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. " * 50,
|
||||
),
|
||||
)
|
||||
messages.append(
|
||||
Message(
|
||||
role=Role.ASSISTANT,
|
||||
content="This is the continuation of the response. " * 50,
|
||||
),
|
||||
)
|
||||
messages.append(
|
||||
Message(
|
||||
role=Role.ASSISTANT,
|
||||
content="And here's the final part with the conclusion. " * 30,
|
||||
),
|
||||
)
|
||||
|
||||
print_messages(messages, "INPUT MESSAGES", max_content_len=80)
|
||||
|
||||
print("\nParameters:")
|
||||
print(" context_window_tokens: 2000")
|
||||
print(" reserve_tokens: 300 (threshold = 1700)")
|
||||
print(" keep_recent_tokens: 600 (should cut in middle of assistant responses)")
|
||||
|
||||
# Use the new context_check method
|
||||
result = await reme_fs.context_check(messages)
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print("RESULT:")
|
||||
print(f" needs_compaction: {result.get('needs_compaction')}")
|
||||
print(f" token_count: {result.get('token_count')}")
|
||||
print(f" threshold: {result.get('threshold')}")
|
||||
print(f" cut_index: {result.get('cut_index')}")
|
||||
print(f" is_split_turn: {result.get('is_split_turn')} *** (should be True)")
|
||||
print(f" accumulated_tokens: {result.get('accumulated_tokens')}")
|
||||
|
||||
messages_to_summarize = result.get("messages_to_summarize", [])
|
||||
turn_prefix_messages = result.get("turn_prefix_messages", [])
|
||||
left_messages = result.get("left_messages", [])
|
||||
print(f"\n Messages to summarize: {len(messages_to_summarize)}")
|
||||
print(f" Turn prefix messages: {len(turn_prefix_messages)}")
|
||||
print(f" Left messages: {len(left_messages)}")
|
||||
|
||||
if turn_prefix_messages:
|
||||
print("\n Turn prefix messages detail:")
|
||||
for i, msg in enumerate(turn_prefix_messages):
|
||||
role = msg["role"] if isinstance(msg, dict) else msg.role.value
|
||||
content = msg["content"] if isinstance(msg, dict) else msg.content
|
||||
print(f" [{i}] {role}: {str(content)[:60]}...")
|
||||
|
||||
assert result.get("needs_compaction") is True, "Should need compaction"
|
||||
assert result.get("is_split_turn") is True, "Should detect split turn"
|
||||
assert len(turn_prefix_messages) > 0, "Should have turn prefix messages"
|
||||
assert len(messages_to_summarize) > 0, "Should have messages to summarize"
|
||||
assert len(left_messages) > 0, "Should have left messages"
|
||||
|
||||
print("\n✓ TEST PASSED: Split turn correctly detected and cut point found\n")
|
||||
|
||||
await reme_fs.close()
|
||||
|
||||
|
||||
async def main():
|
||||
"""Run context checker tests."""
|
||||
print("\n" + "=" * 80)
|
||||
print("FsContextChecker - Cut Point Finding Test Suite")
|
||||
print("=" * 80)
|
||||
print("\nThis test suite validates the cut point finding logic:")
|
||||
print(" 1. Below threshold - no compaction needed")
|
||||
print(" 2. Above threshold - compaction triggered")
|
||||
print(" 3. Split turn - cut point in middle of assistant response")
|
||||
print("=" * 80)
|
||||
|
||||
# Test 1: No compaction needed
|
||||
await test_no_compaction_needed()
|
||||
|
||||
# Test 2: Compaction triggered above threshold
|
||||
await test_compaction_needed_above_threshold()
|
||||
|
||||
# Test 3: Split turn detection
|
||||
await test_split_turn_scenario()
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("All context checker tests completed!")
|
||||
print("=" * 80)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
523
tests/test_fs_file_watch_integration.py
Normal file
523
tests/test_fs_file_watch_integration.py
Normal file
|
|
@ -0,0 +1,523 @@
|
|||
"""Integration test for ReMeFs file watching with memory_search and memory_get.
|
||||
|
||||
This test demonstrates the complete workflow:
|
||||
1. Create markdown files with personal information in test_reme folder
|
||||
2. Initialize ReMeFs with file watching enabled
|
||||
3. Start file watching to automatically index files into the database
|
||||
4. Use memory_search and memory_get to retrieve the indexed content
|
||||
5. Modify the markdown files
|
||||
6. Verify that modified content is properly indexed and retrievable
|
||||
|
||||
This validates the full pipeline:
|
||||
- File creation → File watcher → Database indexing
|
||||
- Search and retrieval functionality
|
||||
- File modification → Re-indexing → Updated search results
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
|
||||
from reme import ReMeFs
|
||||
|
||||
|
||||
# ==================== Test Configuration ====================
|
||||
|
||||
|
||||
class TestConfig:
|
||||
"""Test configuration settings."""
|
||||
|
||||
WORKING_DIR = "test_reme"
|
||||
MEMORY_SUBDIR = "memory"
|
||||
|
||||
|
||||
# ==================== Helper Functions ====================
|
||||
|
||||
|
||||
def create_test_markdown_files(base_dir: str):
|
||||
"""Create test markdown files with personal information.
|
||||
|
||||
Args:
|
||||
base_dir: Base directory to create test files in
|
||||
"""
|
||||
base_path = Path(base_dir)
|
||||
base_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
memory_path = base_path / TestConfig.MEMORY_SUBDIR
|
||||
memory_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Create personal profile markdown
|
||||
profile_file = memory_path / "profile.md"
|
||||
profile_content = """# Personal Profile
|
||||
|
||||
## Basic Information
|
||||
My name is Zhang Wei (张伟). I am a 32-year-old software engineer living in Beijing, China.
|
||||
I work at ByteDance as a senior backend engineer.
|
||||
|
||||
## Professional Skills
|
||||
- Programming Languages: Python, Go, Java
|
||||
- Specialization: Distributed systems and microservices architecture
|
||||
- Experience: 8 years in software development
|
||||
|
||||
## Education
|
||||
- Master's degree in Computer Science from Tsinghua University (2014)
|
||||
- Focus on machine learning and data mining
|
||||
"""
|
||||
profile_file.write_text(profile_content, encoding="utf-8")
|
||||
print(f"✓ Created: {profile_file}")
|
||||
|
||||
# Create hobbies and interests markdown
|
||||
hobbies_file = memory_path / "hobbies.md"
|
||||
hobbies_content = """# Hobbies and Interests
|
||||
|
||||
## Technical Interests
|
||||
I am passionate about cloud computing and containerization technologies.
|
||||
Recently, I've been exploring Kubernetes and service mesh architectures.
|
||||
|
||||
## Personal Hobbies
|
||||
- Reading: Love science fiction novels, especially works by Liu Cixin
|
||||
- Sports: Play basketball every weekend with friends
|
||||
- Travel: Visited 15 provinces in China, planning to visit Japan next year
|
||||
|
||||
## Learning Goals
|
||||
- Deep dive into distributed tracing systems
|
||||
- Learn more about database internals
|
||||
- Improve English communication skills
|
||||
"""
|
||||
hobbies_file.write_text(hobbies_content, encoding="utf-8")
|
||||
print(f"✓ Created: {hobbies_file}")
|
||||
|
||||
# Create work projects markdown
|
||||
projects_file = memory_path / "projects.md"
|
||||
projects_content = """# Work Projects
|
||||
|
||||
## Current Projects
|
||||
|
||||
### Project Alpha (2024-present)
|
||||
Building a high-performance message queue system to handle 1M+ QPS.
|
||||
Using Go and Redis for the core infrastructure.
|
||||
|
||||
### Project Beta (2023-2024)
|
||||
Developed a distributed configuration management system.
|
||||
Integrated with Kubernetes for dynamic config updates.
|
||||
|
||||
## Past Experience
|
||||
- Led the migration of monolithic services to microservices (2021-2023)
|
||||
- Built automated deployment pipelines using Jenkins and GitLab CI (2020-2021)
|
||||
|
||||
## Technical Challenges Solved
|
||||
- Resolved race conditions in concurrent data processing
|
||||
- Optimized database queries reducing response time by 60%
|
||||
"""
|
||||
projects_file.write_text(projects_content, encoding="utf-8")
|
||||
print(f"✓ Created: {projects_file}")
|
||||
|
||||
return [profile_file, hobbies_file, projects_file]
|
||||
|
||||
|
||||
def modify_test_markdown_files(base_dir: str):
|
||||
"""Modify the test markdown files with updated information.
|
||||
|
||||
Args:
|
||||
base_dir: Base directory containing test files
|
||||
"""
|
||||
base_path = Path(base_dir)
|
||||
memory_path = base_path / TestConfig.MEMORY_SUBDIR
|
||||
|
||||
# Modify profile - update job title and add new skill
|
||||
profile_file = memory_path / "profile.md"
|
||||
profile_content = """# Personal Profile
|
||||
|
||||
## Basic Information
|
||||
My name is Zhang Wei (张伟). I am a 32-year-old software engineer living in Beijing, China.
|
||||
I work at ByteDance as a **principal engineer** and tech lead.
|
||||
|
||||
## Professional Skills
|
||||
- Programming Languages: Python, Go, Java, Rust
|
||||
- Specialization: Distributed systems, microservices, and cloud-native architectures
|
||||
- Experience: 8 years in software development
|
||||
- **New**: Expert in observability and monitoring systems
|
||||
|
||||
## Education
|
||||
- Master's degree in Computer Science from Tsinghua University (2014)
|
||||
- Focus on machine learning and data mining
|
||||
"""
|
||||
profile_file.write_text(profile_content, encoding="utf-8")
|
||||
print(f"✓ Modified: {profile_file}")
|
||||
|
||||
# Modify hobbies - add new hobby
|
||||
hobbies_file = memory_path / "hobbies.md"
|
||||
hobbies_content = """# Hobbies and Interests
|
||||
|
||||
## Technical Interests
|
||||
I am passionate about cloud computing and containerization technologies.
|
||||
Recently, I've been exploring Kubernetes, service mesh, and eBPF technologies.
|
||||
|
||||
## Personal Hobbies
|
||||
- Reading: Love science fiction novels, especially works by Liu Cixin
|
||||
- Sports: Play basketball every weekend with friends
|
||||
- Travel: Visited 15 provinces in China, planning to visit Japan next year
|
||||
- **New**: Photography - Recently bought a Sony A7 III camera
|
||||
|
||||
## Learning Goals
|
||||
- Deep dive into distributed tracing and eBPF
|
||||
- Learn more about database internals and query optimization
|
||||
- Improve English communication skills
|
||||
- **New**: Master advanced photography techniques
|
||||
"""
|
||||
hobbies_file.write_text(hobbies_content, encoding="utf-8")
|
||||
print(f"✓ Modified: {hobbies_file}")
|
||||
|
||||
# Modify projects - add new project
|
||||
projects_file = memory_path / "projects.md"
|
||||
projects_content = """# Work Projects
|
||||
|
||||
## Current Projects
|
||||
|
||||
### Project Gamma (2024-present) **NEW**
|
||||
Leading the development of an observability platform using OpenTelemetry.
|
||||
Integrating metrics, traces, and logs into a unified dashboard.
|
||||
|
||||
### Project Alpha (2024-present)
|
||||
Building a high-performance message queue system to handle 1M+ QPS.
|
||||
Using Go and Redis for the core infrastructure.
|
||||
**Update**: Successfully deployed to production, handling 2M+ QPS now.
|
||||
|
||||
### Project Beta (2023-2024)
|
||||
Developed a distributed configuration management system.
|
||||
Integrated with Kubernetes for dynamic config updates.
|
||||
|
||||
## Past Experience
|
||||
- Led the migration of monolithic services to microservices (2021-2023)
|
||||
- Built automated deployment pipelines using Jenkins and GitLab CI (2020-2021)
|
||||
|
||||
## Technical Challenges Solved
|
||||
- Resolved race conditions in concurrent data processing
|
||||
- Optimized database queries reducing response time by 60%
|
||||
- **New**: Implemented distributed tracing reducing MTTR by 40%
|
||||
"""
|
||||
projects_file.write_text(projects_content, encoding="utf-8")
|
||||
print(f"✓ Modified: {projects_file}")
|
||||
|
||||
|
||||
def print_separator(title: str):
|
||||
"""Print a formatted separator line."""
|
||||
print(f"\n{'=' * 80}")
|
||||
print(f" {title}")
|
||||
print(f"{'=' * 80}\n")
|
||||
|
||||
|
||||
def print_search_results(results: list[dict], query: str, context: str):
|
||||
"""Pretty print search results.
|
||||
|
||||
Args:
|
||||
results: List of search results
|
||||
query: The search query
|
||||
context: Context description (e.g., "BEFORE MODIFICATION")
|
||||
"""
|
||||
print(f"\n{'-' * 80}")
|
||||
print(f"Search Results - {context}")
|
||||
print(f"Query: '{query}'")
|
||||
print(f"Found: {len(results)} results")
|
||||
print(f"{'-' * 80}")
|
||||
|
||||
for i, result in enumerate(results, 1):
|
||||
print(f"\n[{i}] Path: {result.get('path', 'N/A')}")
|
||||
print(f" Lines: {result.get('start_line', '?')}-{result.get('end_line', '?')}")
|
||||
print(f" Score: {result.get('score', 0):.4f}")
|
||||
snippet = result.get("snippet", result.get("text", ""))
|
||||
if len(snippet) > 200:
|
||||
snippet = snippet[:200] + "..."
|
||||
print(f" Snippet: {snippet}")
|
||||
|
||||
print(f"{'-' * 80}\n")
|
||||
|
||||
|
||||
def print_get_result(content: str, path: str, context: str):
|
||||
"""Pretty print memory_get result.
|
||||
|
||||
Args:
|
||||
content: Content retrieved from memory_get
|
||||
path: File path
|
||||
context: Context description
|
||||
"""
|
||||
print(f"\n{'-' * 80}")
|
||||
print(f"Memory Get Result - {context}")
|
||||
print(f"Path: {path}")
|
||||
print(f"Content length: {len(content)} chars, {len(content.split(chr(10)))} lines")
|
||||
print(f"{'-' * 80}")
|
||||
print(content[:500] + ("..." if len(content) > 500 else ""))
|
||||
print(f"{'-' * 80}\n")
|
||||
|
||||
|
||||
# ==================== Test Functions ====================
|
||||
|
||||
|
||||
async def test_file_watch_integration():
|
||||
"""Complete integration test for file watching with search and get.
|
||||
|
||||
This test validates:
|
||||
1. File creation and automatic indexing via file watcher
|
||||
2. Search functionality returns correct results
|
||||
3. Get functionality retrieves correct content
|
||||
4. File modification triggers re-indexing
|
||||
5. Updated content is properly searchable and retrievable
|
||||
"""
|
||||
print_separator("FILE WATCH INTEGRATION TEST - START")
|
||||
|
||||
# Clean up any existing test directory
|
||||
test_dir = Path(TestConfig.WORKING_DIR)
|
||||
if test_dir.exists():
|
||||
shutil.rmtree(test_dir)
|
||||
print(f"✓ Cleaned up existing test directory: {test_dir}")
|
||||
|
||||
# ==================== STEP 1: Create Test Files ====================
|
||||
print_separator("STEP 1: Creating Test Files")
|
||||
|
||||
test_files = create_test_markdown_files(TestConfig.WORKING_DIR)
|
||||
print(f"\n✓ Created {len(test_files)} markdown files in {TestConfig.WORKING_DIR}")
|
||||
|
||||
# ==================== STEP 2: Initialize ReMeFs ====================
|
||||
print_separator("STEP 2: Initializing ReMeFs with File Watching")
|
||||
|
||||
reme_fs = ReMeFs(
|
||||
enable_logo=False,
|
||||
working_dir=TestConfig.WORKING_DIR,
|
||||
default_memory_store_config={
|
||||
"backend": "sqlite",
|
||||
"store_name": "test_integration",
|
||||
"embedding_model": "default",
|
||||
"fts_enabled": True,
|
||||
"snippet_max_chars": 700,
|
||||
},
|
||||
default_file_watcher_config={
|
||||
"backend": "full",
|
||||
"watch_paths": [TestConfig.WORKING_DIR, f"{TestConfig.WORKING_DIR}/memory"],
|
||||
"suffix_filters": [".md"],
|
||||
"recursive": False,
|
||||
"scan_on_start": True,
|
||||
},
|
||||
)
|
||||
|
||||
print("✓ ReMeFs instance created")
|
||||
print(f" Working directory: {TestConfig.WORKING_DIR}")
|
||||
print(f" Watch paths: {TestConfig.WORKING_DIR}, {TestConfig.WORKING_DIR}/memory")
|
||||
print(" File filters: .md files")
|
||||
|
||||
# ==================== STEP 3: Start File Watching ====================
|
||||
print_separator("STEP 3: Starting File Watcher")
|
||||
|
||||
await reme_fs.start()
|
||||
print("✓ File watcher started")
|
||||
print(" Files will be automatically indexed into the database")
|
||||
|
||||
# Give file watcher time to process files
|
||||
print("\nWaiting 3 seconds for file watcher to index files...")
|
||||
await asyncio.sleep(3)
|
||||
print("✓ File watcher should have processed all files")
|
||||
|
||||
# ==================== STEP 4: Search Initial Content ====================
|
||||
print_separator("STEP 4: Searching Initial Content")
|
||||
|
||||
queries_initial = [
|
||||
"What programming languages does Zhang Wei know?",
|
||||
"What are Zhang Wei's hobbies?",
|
||||
"What projects is Zhang Wei working on?",
|
||||
]
|
||||
|
||||
results_before = {}
|
||||
|
||||
for query in queries_initial:
|
||||
print(f"\n📍 Searching: '{query}'")
|
||||
result_json = await reme_fs.memory_search(
|
||||
query=query,
|
||||
max_results=3,
|
||||
min_score=0.0,
|
||||
)
|
||||
results = json.loads(result_json)
|
||||
results_before[query] = results
|
||||
print_search_results(results, query, "BEFORE MODIFICATION")
|
||||
|
||||
assert len(results) > 0, f"Should find results for query: {query}"
|
||||
print(f"✓ Found {len(results)} results")
|
||||
|
||||
# ==================== STEP 5: Get Specific Content ====================
|
||||
print_separator("STEP 5: Getting Specific Content with memory_get")
|
||||
|
||||
# Try to get content from profile.md
|
||||
profile_path = f"{TestConfig.MEMORY_SUBDIR}/profile.md"
|
||||
print(f"\n📍 Getting content from: {profile_path}")
|
||||
|
||||
profile_content_before = await reme_fs.memory_get(
|
||||
path=profile_path,
|
||||
offset=1,
|
||||
limit=10,
|
||||
)
|
||||
print_get_result(profile_content_before, profile_path, "BEFORE MODIFICATION")
|
||||
|
||||
assert "Zhang Wei" in profile_content_before, "Should contain Zhang Wei"
|
||||
assert "software engineer" in profile_content_before, "Should contain job title"
|
||||
print("✓ Content retrieved successfully")
|
||||
|
||||
# Get full hobbies.md content
|
||||
hobbies_path = f"{TestConfig.MEMORY_SUBDIR}/hobbies.md"
|
||||
print(f"\n📍 Getting full content from: {hobbies_path}")
|
||||
|
||||
hobbies_content_before = await reme_fs.memory_get(path=hobbies_path)
|
||||
print_get_result(hobbies_content_before, hobbies_path, "BEFORE MODIFICATION")
|
||||
|
||||
assert "basketball" in hobbies_content_before, "Should contain hobbies"
|
||||
print("✓ Full content retrieved successfully")
|
||||
|
||||
# ==================== STEP 6: Modify Files ====================
|
||||
print_separator("STEP 6: Modifying Test Files")
|
||||
|
||||
print("Modifying markdown files with updated information...")
|
||||
modify_test_markdown_files(TestConfig.WORKING_DIR)
|
||||
|
||||
# Give file watcher time to detect and re-index changes
|
||||
print("\nWaiting 3 seconds for file watcher to detect and re-index changes...")
|
||||
await asyncio.sleep(3)
|
||||
print("✓ File watcher should have re-indexed modified files")
|
||||
|
||||
# ==================== STEP 7: Search Modified Content ====================
|
||||
print_separator("STEP 7: Searching Modified Content")
|
||||
|
||||
queries_modified = [
|
||||
"What is Zhang Wei's current job title?",
|
||||
"Does Zhang Wei have any new hobbies?",
|
||||
"What new projects is Zhang Wei working on?",
|
||||
"What expertise does Zhang Wei have in observability?",
|
||||
]
|
||||
|
||||
results_after = {}
|
||||
|
||||
for query in queries_modified:
|
||||
print(f"\n📍 Searching: '{query}'")
|
||||
result_json = await reme_fs.memory_search(
|
||||
query=query,
|
||||
max_results=3,
|
||||
min_score=0.0,
|
||||
)
|
||||
results = json.loads(result_json)
|
||||
results_after[query] = results
|
||||
print_search_results(results, query, "AFTER MODIFICATION")
|
||||
|
||||
assert len(results) > 0, f"Should find results for query: {query}"
|
||||
print(f"✓ Found {len(results)} results")
|
||||
|
||||
# ==================== STEP 8: Get Modified Content ====================
|
||||
print_separator("STEP 8: Getting Modified Content")
|
||||
|
||||
# Get updated profile content
|
||||
print(f"\n📍 Getting updated content from: {profile_path}")
|
||||
profile_content_after = await reme_fs.memory_get(
|
||||
path=profile_path,
|
||||
offset=1,
|
||||
limit=10,
|
||||
)
|
||||
print_get_result(profile_content_after, profile_path, "AFTER MODIFICATION")
|
||||
|
||||
assert "principal engineer" in profile_content_after, "Should contain updated job title"
|
||||
assert "Rust" in profile_content_after, "Should contain new programming language"
|
||||
print("✓ Updated profile content retrieved successfully")
|
||||
|
||||
# Get updated hobbies content
|
||||
print(f"\n📍 Getting updated content from: {hobbies_path}")
|
||||
hobbies_content_after = await reme_fs.memory_get(path=hobbies_path)
|
||||
print_get_result(hobbies_content_after, hobbies_path, "AFTER MODIFICATION")
|
||||
|
||||
assert "Photography" in hobbies_content_after, "Should contain new hobby"
|
||||
assert "Sony A7 III" in hobbies_content_after, "Should contain camera info"
|
||||
print("✓ Updated hobbies content retrieved successfully")
|
||||
|
||||
# Get updated projects content
|
||||
projects_path = f"{TestConfig.MEMORY_SUBDIR}/projects.md"
|
||||
print(f"\n📍 Getting updated content from: {projects_path}")
|
||||
projects_content_after = await reme_fs.memory_get(path=projects_path)
|
||||
print_get_result(projects_content_after, projects_path, "AFTER MODIFICATION")
|
||||
|
||||
assert "Project Gamma" in projects_content_after, "Should contain new project"
|
||||
assert "OpenTelemetry" in projects_content_after, "Should contain new technology"
|
||||
print("✓ Updated projects content retrieved successfully")
|
||||
|
||||
# ==================== STEP 9: Verify Changes ====================
|
||||
print_separator("STEP 9: Verifying Content Changes")
|
||||
|
||||
print("\n📍 Comparing BEFORE vs AFTER content:")
|
||||
|
||||
# Verify profile changes
|
||||
print("\n1. Profile.md changes:")
|
||||
print(f" Before: Contains 'software engineer' = {('software engineer' in profile_content_before.lower())}")
|
||||
print(f" After: Contains 'principal engineer' = {('principal engineer' in profile_content_after.lower())}")
|
||||
print(f" After: Contains 'Rust' = {('rust' in profile_content_after.lower())}")
|
||||
|
||||
# Verify hobbies changes
|
||||
print("\n2. Hobbies.md changes:")
|
||||
print(f" Before: Contains 'Photography' = {('photography' in hobbies_content_before.lower())}")
|
||||
print(f" After: Contains 'Photography' = {('photography' in hobbies_content_after.lower())}")
|
||||
print(f" After: Contains 'Sony A7 III' = {('sony' in hobbies_content_after.lower())}")
|
||||
|
||||
# Verify projects changes
|
||||
print("\n3. Projects.md changes:")
|
||||
print(f" After: Contains 'Project Gamma' = {('Project Gamma' in projects_content_after)}")
|
||||
print(f" After: Contains 'OpenTelemetry' = {('OpenTelemetry' in projects_content_after)}")
|
||||
|
||||
print("\n✓ All content changes verified successfully")
|
||||
|
||||
# ==================== STEP 10: Cleanup ====================
|
||||
print_separator("STEP 10: Cleanup")
|
||||
|
||||
await reme_fs.close()
|
||||
print("✓ ReMeFs closed")
|
||||
|
||||
# Clean up test directory
|
||||
if test_dir.exists():
|
||||
shutil.rmtree(test_dir)
|
||||
print(f"✓ Removed test directory: {test_dir}")
|
||||
else:
|
||||
print(f"⚠️ Test directory does not exist: {test_dir}")
|
||||
|
||||
print("\n✓ All test data cleaned up")
|
||||
|
||||
print_separator("FILE WATCH INTEGRATION TEST - COMPLETED SUCCESSFULLY")
|
||||
|
||||
|
||||
# ==================== Main Entry Point ====================
|
||||
|
||||
|
||||
async def main():
|
||||
"""Run the file watch integration test."""
|
||||
print("\n" + "=" * 80)
|
||||
print(" ReMeFs File Watch Integration Test")
|
||||
print("=" * 80)
|
||||
print("\nThis test validates the complete file watching workflow:")
|
||||
print(" 1. Create markdown files with personal information")
|
||||
print(" 2. Initialize ReMeFs and start file watching")
|
||||
print(" 3. Verify automatic indexing into database")
|
||||
print(" 4. Search and retrieve initial content")
|
||||
print(" 5. Modify files and verify re-indexing")
|
||||
print(" 6. Search and retrieve modified content")
|
||||
print(" 7. Compare before/after results")
|
||||
print("=" * 80)
|
||||
|
||||
try:
|
||||
await test_file_watch_integration()
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print(" ✓ All tests passed successfully!")
|
||||
print("=" * 80)
|
||||
|
||||
except Exception as e:
|
||||
print("\n" + "=" * 80)
|
||||
print(f" ✗ Test failed with error: {e}")
|
||||
print("=" * 80)
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
raise
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
|
|
@ -340,13 +340,6 @@ async def test_memory_search_with_source_filter():
|
|||
reme_fs = ReMeFs(
|
||||
enable_logo=False,
|
||||
working_dir=TestConfig.WORKING_DIR,
|
||||
default_memory_store_config={
|
||||
"backend": "sqlite",
|
||||
"store_name": "test_source_filter",
|
||||
"embedding_model": "default",
|
||||
"fts_enabled": True,
|
||||
"snippet_max_chars": 700,
|
||||
},
|
||||
)
|
||||
await reme_fs.start()
|
||||
|
||||
|
|
@ -382,11 +375,19 @@ async def test_memory_search_with_source_filter():
|
|||
|
||||
# Search only MEMORY source
|
||||
print(f"\n--- Searching MEMORY source for: '{query}' ---")
|
||||
result_json = await reme_fs.memory_search(
|
||||
# Create a new instance with MEMORY source filter
|
||||
reme_fs_memory = ReMeFs(
|
||||
enable_logo=False,
|
||||
working_dir=TestConfig.WORKING_DIR,
|
||||
search_params={"sources": [MemorySource.MEMORY]},
|
||||
)
|
||||
await reme_fs_memory.start()
|
||||
result_json = await reme_fs_memory.memory_search(
|
||||
query=query,
|
||||
max_results=5,
|
||||
sources=[MemorySource.MEMORY],
|
||||
)
|
||||
await reme_fs_memory.close()
|
||||
|
||||
import json
|
||||
|
||||
memory_results = json.loads(result_json)
|
||||
|
|
@ -396,11 +397,18 @@ async def test_memory_search_with_source_filter():
|
|||
|
||||
# Search only SESSIONS source
|
||||
print(f"\n--- Searching SESSIONS source for: '{query}' ---")
|
||||
result_json = await reme_fs.memory_search(
|
||||
# Create a new instance with SESSIONS source filter
|
||||
reme_fs_sessions = ReMeFs(
|
||||
enable_logo=False,
|
||||
working_dir=TestConfig.WORKING_DIR,
|
||||
search_params={"sources": [MemorySource.SESSIONS]},
|
||||
)
|
||||
await reme_fs_sessions.start()
|
||||
result_json = await reme_fs_sessions.memory_search(
|
||||
query=query,
|
||||
max_results=5,
|
||||
sources=[MemorySource.SESSIONS],
|
||||
)
|
||||
await reme_fs_sessions.close()
|
||||
session_results = json.loads(result_json)
|
||||
print(f"Found {len(session_results)} results in SESSIONS source")
|
||||
for result in session_results:
|
||||
|
|
@ -597,13 +605,29 @@ async def test_memory_search_hybrid_mode():
|
|||
|
||||
# Test with hybrid enabled
|
||||
print(f"\n--- Hybrid search (enabled) for: '{query}' ---")
|
||||
result_json_hybrid = await reme_fs.memory_search(
|
||||
# Create instance with hybrid enabled
|
||||
reme_fs_hybrid = ReMeFs(
|
||||
enable_logo=False,
|
||||
working_dir=TestConfig.WORKING_DIR,
|
||||
default_memory_store_config={
|
||||
"backend": "sqlite",
|
||||
"store_name": "test_hybrid",
|
||||
"embedding_model": "default",
|
||||
"fts_enabled": True,
|
||||
"snippet_max_chars": 700,
|
||||
},
|
||||
search_params={
|
||||
"hybrid_enabled": True,
|
||||
"hybrid_vector_weight": 0.7,
|
||||
"hybrid_text_weight": 0.3,
|
||||
},
|
||||
)
|
||||
await reme_fs_hybrid.start()
|
||||
result_json_hybrid = await reme_fs_hybrid.memory_search(
|
||||
query=query,
|
||||
max_results=5,
|
||||
hybrid_enabled=True,
|
||||
hybrid_vector_weight=0.7,
|
||||
hybrid_text_weight=0.3,
|
||||
)
|
||||
await reme_fs_hybrid.close()
|
||||
|
||||
import json
|
||||
|
||||
|
|
@ -613,11 +637,25 @@ async def test_memory_search_hybrid_mode():
|
|||
|
||||
# Test with hybrid disabled (vector only)
|
||||
print(f"\n--- Vector-only search for: '{query}' ---")
|
||||
result_json_vector = await reme_fs.memory_search(
|
||||
# Create instance with hybrid disabled
|
||||
reme_fs_vector = ReMeFs(
|
||||
enable_logo=False,
|
||||
working_dir=TestConfig.WORKING_DIR,
|
||||
default_memory_store_config={
|
||||
"backend": "sqlite",
|
||||
"store_name": "test_hybrid",
|
||||
"embedding_model": "default",
|
||||
"fts_enabled": True,
|
||||
"snippet_max_chars": 700,
|
||||
},
|
||||
search_params={"hybrid_enabled": False},
|
||||
)
|
||||
await reme_fs_vector.start()
|
||||
result_json_vector = await reme_fs_vector.memory_search(
|
||||
query=query,
|
||||
max_results=5,
|
||||
hybrid_enabled=False,
|
||||
)
|
||||
await reme_fs_vector.close()
|
||||
|
||||
vector_results = json.loads(result_json_vector)
|
||||
print(f"Vector search found {len(vector_results)} results")
|
||||
|
|
@ -632,13 +670,29 @@ async def test_memory_search_hybrid_mode():
|
|||
]
|
||||
|
||||
for vec_weight, text_weight in weight_configs:
|
||||
result_json = await reme_fs.memory_search(
|
||||
# Create instance with specific weights
|
||||
reme_fs_weights = ReMeFs(
|
||||
enable_logo=False,
|
||||
working_dir=TestConfig.WORKING_DIR,
|
||||
default_memory_store_config={
|
||||
"backend": "sqlite",
|
||||
"store_name": "test_hybrid",
|
||||
"embedding_model": "default",
|
||||
"fts_enabled": True,
|
||||
"snippet_max_chars": 700,
|
||||
},
|
||||
search_params={
|
||||
"hybrid_enabled": True,
|
||||
"hybrid_vector_weight": vec_weight,
|
||||
"hybrid_text_weight": text_weight,
|
||||
},
|
||||
)
|
||||
await reme_fs_weights.start()
|
||||
result_json = await reme_fs_weights.memory_search(
|
||||
query=query,
|
||||
max_results=5,
|
||||
hybrid_enabled=True,
|
||||
hybrid_vector_weight=vec_weight,
|
||||
hybrid_text_weight=text_weight,
|
||||
)
|
||||
await reme_fs_weights.close()
|
||||
results = json.loads(result_json)
|
||||
print(f" Vector:{vec_weight}/Text:{text_weight} -> {len(results)} results")
|
||||
|
||||
|
|
|
|||
|
|
@ -155,7 +155,6 @@ async def test_summary_personal_info_storage():
|
|||
|
||||
result = await reme_fs.summary(
|
||||
messages=messages,
|
||||
version="default",
|
||||
date="2023-09-01",
|
||||
)
|
||||
|
||||
|
|
@ -191,7 +190,6 @@ async def test_summary_detailed_profile():
|
|||
|
||||
result = await reme_fs.summary(
|
||||
messages=messages,
|
||||
version="default",
|
||||
date="2023-10-01",
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ class TestConfig:
|
|||
"""Configuration for test execution."""
|
||||
|
||||
# SqliteMemoryStore settings
|
||||
NAME = "test"
|
||||
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
|
||||
|
|
@ -205,6 +206,7 @@ def create_memory_store(store_type: str) -> BaseMemoryStore:
|
|||
|
||||
if store_type == "sqlite":
|
||||
return SqliteMemoryStore(
|
||||
store_name=config.NAME,
|
||||
db_path=config.SQLITE_DB_PATH,
|
||||
embedding_model=embedding_model,
|
||||
vec_ext_path=config.SQLITE_VEC_EXT_PATH,
|
||||
|
|
@ -235,8 +237,8 @@ async def test_start_store(store: BaseMemoryStore, _store_name: str):
|
|||
cursor.close()
|
||||
|
||||
logger.info(f"Created tables: {tables}")
|
||||
assert "files" in tables, "files table should exist"
|
||||
assert "chunks" in tables, "chunks table should exist"
|
||||
assert store.files_table_name in tables, f"{store.files_table_name} table should exist"
|
||||
assert store.chunks_table_name in tables, f"{store.chunks_table_name} table should exist"
|
||||
logger.info("✓ Required tables created")
|
||||
|
||||
|
||||
|
|
@ -577,6 +579,42 @@ async def test_keyword_search_with_source_filter(store: BaseMemoryStore, _store_
|
|||
logger.info("\n✓ Keyword search with source filter test passed")
|
||||
|
||||
|
||||
async def test_keyword_search_special_chars(store: BaseMemoryStore, _store_name: str):
|
||||
"""Test keyword search with special characters like ?, *, etc."""
|
||||
logger.info("=" * 20 + " KEYWORD SEARCH SPECIAL CHARS TEST " + "=" * 20)
|
||||
|
||||
# Check if FTS is available
|
||||
if isinstance(store, SqliteMemoryStore) and not store.fts_available:
|
||||
logger.info("⊘ Skipped: FTS not available")
|
||||
return
|
||||
|
||||
# Test various queries with special characters
|
||||
test_queries = [
|
||||
"What is the status?",
|
||||
"How does it work?",
|
||||
"Why is this important?",
|
||||
"data?",
|
||||
"test*",
|
||||
"query with ? marks",
|
||||
]
|
||||
|
||||
for query in test_queries:
|
||||
logger.info(f"\nTesting query: '{query}'")
|
||||
try:
|
||||
results = await store.keyword_search(query, limit=3)
|
||||
logger.info(f"✓ Query succeeded, found {len(results)} results")
|
||||
if results:
|
||||
for i, result in enumerate(results[:2], 1): # Show first 2 results
|
||||
logger.info(
|
||||
f" {i}. {result.path}:{result.start_line}-{result.end_line} (score: {result.score:.4f})",
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"✗ Query failed: {e}")
|
||||
raise
|
||||
|
||||
logger.info("\n✓ Keyword search with special characters test passed")
|
||||
|
||||
|
||||
async def test_delete_file(store: BaseMemoryStore, _store_name: str):
|
||||
"""Test file deletion."""
|
||||
logger.info("=" * 20 + " DELETE FILE TEST " + "=" * 20)
|
||||
|
|
@ -852,6 +890,7 @@ async def run_all_tests_for_store(store_type: str, store_name: str):
|
|||
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)
|
||||
await test_keyword_search_special_chars(store, store_name)
|
||||
|
||||
# ========== Advanced Tests ==========
|
||||
logger.info(f"\n{'#' * 60}")
|
||||
|
|
|
|||
|
|
@ -182,7 +182,7 @@ async def test_stream_chat(app):
|
|||
stream_queue=context.stream_queue,
|
||||
task=asyncio.create_task(task()),
|
||||
task_name="test_stream_chat",
|
||||
as_bytes=False,
|
||||
output_format="str",
|
||||
):
|
||||
print(chunk, end="")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue