From eae6c63d159469b0f57cfbdae7bf63730626f54b Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 11 Feb 2026 17:23:15 +0800 Subject: [PATCH] feat(cli): add ReMeCli class with interactive chat functionality --- pyproject.toml | 2 +- reme/__init__.py | 4 +- reme/agent/chat/fs_cli.py | 10 +- reme/agent/fs/fs_summarizer.py | 18 +- reme/config/fs.yaml | 4 +- reme/core/application.py | 2 + reme/core/context/service_context.py | 83 ++- reme/core/file_watcher/full_file_watcher.py | 1 - reme/core/memory_store/base_memory_store.py | 10 +- reme/core/memory_store/sqlite_memory_store.py | 568 +++++++++++------- reme/core/schema/service_config.py | 10 +- reme/reme.py | 31 +- reme/reme_cli.py | 199 ++++++ reme/reme_fs.py | 221 +------ reme/tool/fs/fs_memory_search.py | 27 +- tests/test_fs_memory_search.py | 6 +- 16 files changed, 682 insertions(+), 514 deletions(-) create mode 100644 reme/reme_cli.py diff --git a/pyproject.toml b/pyproject.toml index df1d919f..028f4fa8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -109,7 +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" +remecli = "reme.reme_cli:main" [tool.pytest.ini_options] asyncio_default_fixture_loop_scope = "function" diff --git a/reme/__init__.py b/reme/__init__.py index de22976c..383140c2 100644 --- a/reme/__init__.py +++ b/reme/__init__.py @@ -6,6 +6,7 @@ from . import core from . import tool from . import workflow from .reme import ReMe +from .reme_cli import ReMeCli from .reme_fs import ReMeFs __all__ = [ @@ -15,10 +16,11 @@ __all__ = [ "tool", "workflow", "ReMe", + "ReMeCli", "ReMeFs", ] -__version__ = "0.3.0.0a4" +__version__ = "0.3.0.0a5" """ diff --git a/reme/agent/chat/fs_cli.py b/reme/agent/chat/fs_cli.py index e7794f59..671062e8 100644 --- a/reme/agent/chat/fs_cli.py +++ b/reme/agent/chat/fs_cli.py @@ -20,10 +20,6 @@ class FsCli(BaseReactStream): 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) @@ -32,10 +28,6 @@ class FsCli(BaseReactStream): 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 = "" @@ -168,7 +160,7 @@ class FsCli(BaseReactStream): 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)}") + logger.info(f"[{self.__class__.__name__}] msg[{i}] role={role} {message.simple_dump(as_dict=False)}") t_tools, messages, success = await self.react(messages, self.tools) diff --git a/reme/agent/fs/fs_summarizer.py b/reme/agent/fs/fs_summarizer.py index d1406462..91ff0510 100644 --- a/reme/agent/fs/fs_summarizer.py +++ b/reme/agent/fs/fs_summarizer.py @@ -29,13 +29,29 @@ class FsSummarizer(BaseReact): role=Role.USER, content=self.prompt_format( "user_message_default", - conversation=format_messages(messages, add_index=False), working_dir=self.working_dir, date=date_str, memory_dir=self.memory_dir, ), ), ) + + elif self.version == "v1": + conversation = format_messages(messages, add_index=False) + messages = [ + Message( + role=Role.USER, + content=f"\n{conversation}\n\n" + + self.prompt_format( + "user_message_default", + conversation=format_messages(messages, add_index=False), + working_dir=self.working_dir, + date=date_str, + memory_dir=self.memory_dir, + ), + ), + ] + else: messages.extend( [ diff --git a/reme/config/fs.yaml b/reme/config/fs.yaml index 2bcf0516..daa59c6a 100644 --- a/reme/config/fs.yaml +++ b/reme/config/fs.yaml @@ -18,13 +18,15 @@ embedding_models: memory_stores: default: backend: sqlite - store_name: test_hybrid + store_name: reme embedding_model: default fts_enabled: true + vector_enabled: false file_watchers: default: backend: full + memory_store: default watch_paths: [".reme", ".reme/memory"] suffix_filters: [".md"] recursive: false diff --git a/reme/core/application.py b/reme/core/application.py index e380a21f..31c5de8e 100644 --- a/reme/core/application.py +++ b/reme/core/application.py @@ -24,6 +24,7 @@ class Application: llm_base_url: str | None = None, embedding_api_key: str | None = None, embedding_base_url: str | None = None, + working_dir: str | None = None, config_path: str | None = None, enable_logo: bool = True, log_to_console: bool = True, @@ -44,6 +45,7 @@ class Application: embedding_base_url=embedding_base_url, service_config=None, parser=parser, + working_dir=working_dir, config_path=config_path, enable_logo=enable_logo, log_to_console=log_to_console, diff --git a/reme/core/context/service_context.py b/reme/core/context/service_context.py index 29214776..9fee8a31 100644 --- a/reme/core/context/service_context.py +++ b/reme/core/context/service_context.py @@ -2,6 +2,7 @@ import os from concurrent.futures import ThreadPoolExecutor +from pathlib import Path from typing import TYPE_CHECKING from loguru import logger @@ -34,6 +35,7 @@ class ServiceContext(BaseContext): embedding_base_url: str | None = None, service_config: ServiceConfig | None = None, parser: type[PydanticConfigParser] | None = None, + working_dir: str | None = None, config_path: str | None = None, enable_logo: bool = True, log_to_console: bool = True, @@ -74,13 +76,23 @@ class ServiceContext(BaseContext): self._update_section_config(kwargs, "memory_stores", **default_memory_store_config) 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 + + kwargs.update( + { + "enable_logo": enable_logo, + "log_to_console": log_to_console, + "working_dir": working_dir, + }, + ) logger.info(f"update with args: {input_args} kwargs: {kwargs}") service_config = parser.parse_args(*input_args, **kwargs) self.service_config: ServiceConfig = service_config init_logger(log_to_console=self.service_config.log_to_console) + logger.info(f"ReMe Config: {service_config.model_dump_json()}") + + if self.service_config.working_dir: + Path(self.service_config.working_dir).mkdir(parents=True, exist_ok=True) if self.service_config.enable_logo: print_logo(service_config=self.service_config) @@ -147,41 +159,58 @@ class ServiceContext(BaseContext): async def start(self): """Start the service context by initializing all configured components.""" for name, config in self.service_config.llms.items(): - self.llms[name] = R.llms[config.backend](model_name=config.model_name, **config.model_extra) + if config.backend not in R.llms: + logger.warning(f"LLM backend {config.backend} is not supported.") + else: + self.llms[name] = R.llms[config.backend](model_name=config.model_name, **config.model_extra) for name, config in self.service_config.embedding_models.items(): - self.embedding_models[name] = R.embedding_models[config.backend]( - model_name=config.model_name, - **config.model_extra, - ) + if config.backend not in R.embedding_models: + logger.warning(f"Embedding model backend {config.backend} is not supported.") + else: + self.embedding_models[name] = R.embedding_models[config.backend]( + model_name=config.model_name, + **config.model_extra, + ) for name, config in self.service_config.token_counters.items(): - self.token_counters[name] = R.token_counters[config.backend]( - model_name=config.model_name, - **config.model_extra, - ) + if config.backend not in R.token_counters: + logger.warning(f"Token counter backend {config.backend} is not supported.") + else: + self.token_counters[name] = R.token_counters[config.backend]( + model_name=config.model_name, + **config.model_extra, + ) for name, config in self.service_config.vector_stores.items(): - # 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) + if config.backend not in R.vector_stores: + logger.warning(f"Vector store backend {config.backend} is not supported.") + else: + # 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(): - # 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() + if config.backend not in R.memory_stores: + logger.warning(f"Memory store backend {config.backend} is not supported.") + else: + # 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(): - # 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 config.backend not in R.file_watchers: + logger.warning(f"File watcher backend {config.backend} is not supported.") + else: + 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: await self.prepare_mcp_servers() diff --git a/reme/core/file_watcher/full_file_watcher.py b/reme/core/file_watcher/full_file_watcher.py index a172a4ed..1e26e936 100644 --- a/reme/core/file_watcher/full_file_watcher.py +++ b/reme/core/file_watcher/full_file_watcher.py @@ -5,7 +5,6 @@ on any change, ensuring complete synchronization. """ import asyncio -import os from pathlib import Path from loguru import logger diff --git a/reme/core/memory_store/base_memory_store.py b/reme/core/memory_store/base_memory_store.py index 5d028e75..069fe820 100644 --- a/reme/core/memory_store/base_memory_store.py +++ b/reme/core/memory_store/base_memory_store.py @@ -15,6 +15,7 @@ class BaseMemoryStore(ABC): self, store_name: str, embedding_model: BaseEmbeddingModel, + vector_enabled: bool = False, fts_enabled: bool = True, **kwargs, ): @@ -23,14 +24,17 @@ class BaseMemoryStore(ABC): # Only allow alphanumeric characters and underscores if not re.match(r"^[a-zA-Z0-9_]+$", store_name): raise ValueError(f"Invalid '{store_name}'. Only alphanumeric characters and underscores are allowed.") + + # Ensure at least one search method is enabled + if not vector_enabled and not fts_enabled: + raise ValueError("At least one of vector_enabled or fts_enabled must be True.") + self.store_name: str = store_name self.embedding_model: BaseEmbeddingModel = embedding_model + self.vector_enabled: bool = vector_enabled self.fts_enabled: bool = fts_enabled self.kwargs: dict = kwargs - self.vector_available = False - self.fts_available = False - @property def embedding_dim(self) -> int: """Get the embedding model's dimensionality.""" diff --git a/reme/core/memory_store/sqlite_memory_store.py b/reme/core/memory_store/sqlite_memory_store.py index 73afead2..e9a8a6ba 100644 --- a/reme/core/memory_store/sqlite_memory_store.py +++ b/reme/core/memory_store/sqlite_memory_store.py @@ -67,108 +67,114 @@ class SqliteMemoryStore(BaseMemoryStore): Path(self.db_path).parent.mkdir(parents=True, exist_ok=True) self.conn = sqlite3.connect(self.db_path, check_same_thread=False) - self.conn.enable_load_extension(True) - # Load sqlite-vec extension - if self.vec_ext_path: - try: - self.conn.load_extension(self.vec_ext_path) - self.vector_available = True - logger.info(f"Loaded sqlite-vec: {self.vec_ext_path}") - except Exception as e: - logger.warning(f"Failed to load sqlite-vec: {e}") + # Only load sqlite-vec extension if vector search is enabled + if self.vector_enabled: + self.conn.enable_load_extension(True) + # Load sqlite-vec extension + if self.vec_ext_path: + try: + self.conn.load_extension(self.vec_ext_path) + logger.info(f"Loaded sqlite-vec: {self.vec_ext_path}") + except Exception as e: + logger.warning(f"Failed to load sqlite-vec: {e}") + + else: + try: + import sqlite_vec + + ext_path = sqlite_vec.loadable_path() + self.conn.load_extension(ext_path) + logger.info(f"Loaded sqlite-vec from package: {ext_path}") + + except Exception as e: + logger.warning(f"Failed to load sqlite-vec from package: {e}") + # Fallback: try common extension names + for name in ["vec0", "sqlite_vec", "vector0"]: + try: + self.conn.load_extension(name) + logger.info(f"Loaded sqlite-vec: {name}") + break + except Exception: + pass + + self.conn.enable_load_extension(False) else: - try: - import sqlite_vec + logger.info("Vector search disabled, skipping sqlite-vec extension loading") - ext_path = sqlite_vec.loadable_path() - self.conn.load_extension(ext_path) - self.vector_available = True - logger.info(f"Loaded sqlite-vec from package: {ext_path}") - - except Exception as e: - logger.warning(f"Failed to load sqlite-vec from package: {e}") - # Fallback: try common extension names - for name in ["vec0", "sqlite_vec", "vector0"]: - try: - self.conn.load_extension(name) - self.vector_available = True - logger.info(f"Loaded sqlite-vec: {name}") - break - except Exception: - pass - - self.conn.enable_load_extension(False) await self._create_tables() async def _create_tables(self) -> None: """Create database schema.""" cursor = self.conn.cursor() - - # Files - cursor.execute( - f""" - CREATE TABLE IF NOT EXISTS {self.files_table_name} ( - path TEXT, - source TEXT, - hash TEXT, - mtime REAL, - size INTEGER, - PRIMARY KEY (path, source) - ) - """, - ) - - # Chunks - cursor.execute( - f""" - CREATE TABLE IF NOT EXISTS {self.chunks_table_name} ( - id TEXT PRIMARY KEY, - path TEXT, - source TEXT, - start_line INTEGER, - end_line INTEGER, - hash TEXT, - text TEXT, - embedding TEXT, - updated_at INTEGER - ) - """, - ) - - # Vector table (sqlite-vec) - if self.vector_available: + try: + # Files cursor.execute( f""" - CREATE VIRTUAL TABLE IF NOT EXISTS {self.vector_table_name} USING vec0( + CREATE TABLE IF NOT EXISTS {self.files_table_name} ( + path TEXT, + source TEXT, + hash TEXT, + mtime REAL, + size INTEGER, + PRIMARY KEY (path, source) + ) + """, + ) + + # Chunks + cursor.execute( + f""" + CREATE TABLE IF NOT EXISTS {self.chunks_table_name} ( id TEXT PRIMARY KEY, - embedding FLOAT[{self.embedding_dim}] + path TEXT, + source TEXT, + start_line INTEGER, + end_line INTEGER, + hash TEXT, + text TEXT, + embedding TEXT, + updated_at INTEGER ) """, ) - logger.info(f"Created vector table (dims={self.embedding_dim})") - # FTS table - if self.fts_enabled: - cursor.execute( - f""" - CREATE VIRTUAL TABLE IF NOT EXISTS {self.fts_table_name} USING fts5( - text, - id UNINDEXED, - path UNINDEXED, - source UNINDEXED, - start_line UNINDEXED, - end_line UNINDEXED, - tokenize='trigram' + # Vector table (sqlite-vec) + if self.vector_enabled: + cursor.execute( + f""" + CREATE VIRTUAL TABLE IF NOT EXISTS {self.vector_table_name} USING vec0( + id TEXT PRIMARY KEY, + embedding FLOAT[{self.embedding_dim}] + ) + """, ) - """, - ) - self.fts_available = True - logger.info("Created FTS5 table with trigram tokenizer") + logger.info(f"Created vector table (dims={self.embedding_dim})") - self.conn.commit() - cursor.close() + # FTS table + if self.fts_enabled: + cursor.execute( + f""" + CREATE VIRTUAL TABLE IF NOT EXISTS {self.fts_table_name} USING fts5( + text, + id UNINDEXED, + path UNINDEXED, + source UNINDEXED, + start_line UNINDEXED, + end_line UNINDEXED, + tokenize='trigram' + ) + """, + ) + logger.info("Created FTS5 table with trigram tokenizer") + + self.conn.commit() + except Exception as e: + logger.error(f"Failed to create tables: {e}") + raise + finally: + cursor.close() async def upsert_file(self, file_meta: FileMetadata, source: MemorySource, chunks: list[MemoryChunk]): """Insert or update file and its chunks.""" @@ -210,21 +216,25 @@ class SqliteMemoryStore(BaseMemoryStore): ) # Insert vector (vec0 doesn't support OR REPLACE, use DELETE + INSERT) - if self.vector_available: - assert chunk.embedding, "Embedding is required for vector insert" - # Delete existing vector first - cursor.execute( - f"DELETE FROM {self.vector_table_name} WHERE id = ?", - (chunk.id,), - ) - # Then insert new vector - cursor.execute( - f"INSERT INTO {self.vector_table_name} (id, embedding) VALUES (?, ?)", - (chunk.id, self.vector_to_blob(chunk.embedding)), - ) + if self.vector_enabled: + if not chunk.embedding: + logger.warning( + f"Chunk {chunk.id} missing embedding for vector insert, skipping vector indexing", + ) + else: + # Delete existing vector first + cursor.execute( + f"DELETE FROM {self.vector_table_name} WHERE id = ?", + (chunk.id,), + ) + # Then insert new vector + cursor.execute( + f"INSERT INTO {self.vector_table_name} (id, embedding) VALUES (?, ?)", + (chunk.id, self.vector_to_blob(chunk.embedding)), + ) # Insert FTS - if self.fts_available: + if self.fts_enabled: cursor.execute( f""" INSERT OR REPLACE INTO {self.fts_table_name} ( @@ -242,8 +252,9 @@ class SqliteMemoryStore(BaseMemoryStore): ) cursor.execute("COMMIT") - except Exception: + except Exception as e: cursor.execute("ROLLBACK") + logger.error(f"Failed to upsert file {file_meta.path}: {e}") raise finally: cursor.close() @@ -262,25 +273,19 @@ class SqliteMemoryStore(BaseMemoryStore): chunk_ids = [row[0] for row in cursor.fetchall()] # Delete vectors - if self.vector_available and chunk_ids: + if self.vector_enabled and chunk_ids: for chunk_id in chunk_ids: - try: - cursor.execute( - f"DELETE FROM {self.vector_table_name} WHERE id = ?", - (chunk_id,), - ) - except Exception as e: - logger.debug(f"Vector delete failed: {e}") + cursor.execute( + f"DELETE FROM {self.vector_table_name} WHERE id = ?", + (chunk_id,), + ) # Delete FTS entries - if self.fts_available: - try: - cursor.execute( - f"DELETE FROM {self.fts_table_name} WHERE path = ? AND source = ?", - (path, source.value), - ) - except Exception as e: - logger.debug(f"FTS delete failed: {e}") + if self.fts_enabled: + cursor.execute( + f"DELETE FROM {self.fts_table_name} WHERE path = ? AND source = ?", + (path, source.value), + ) # Delete chunks and file cursor.execute( @@ -293,8 +298,9 @@ class SqliteMemoryStore(BaseMemoryStore): ) cursor.execute("COMMIT") - except Exception: + except Exception as e: cursor.execute("ROLLBACK") + logger.error(f"Failed to delete file {path}: {e}") raise finally: cursor.close() @@ -309,26 +315,20 @@ class SqliteMemoryStore(BaseMemoryStore): cursor.execute("BEGIN") # Delete vectors - if self.vector_available: + if self.vector_enabled: for chunk_id in chunk_ids: - try: - cursor.execute( - f"DELETE FROM {self.vector_table_name} WHERE id = ?", - (chunk_id,), - ) - except Exception as e: - logger.debug(f"Vector delete failed for {chunk_id}: {e}") + cursor.execute( + f"DELETE FROM {self.vector_table_name} WHERE id = ?", + (chunk_id,), + ) # Delete FTS entries - if self.fts_available: + if self.fts_enabled: placeholders = ",".join("?" * len(chunk_ids)) - try: - cursor.execute( - f"DELETE FROM {self.fts_table_name} WHERE id IN ({placeholders})", - chunk_ids, - ) - except Exception as e: - logger.debug(f"FTS delete failed: {e}") + cursor.execute( + f"DELETE FROM {self.fts_table_name} WHERE id IN ({placeholders})", + chunk_ids, + ) # Delete chunks placeholders = ",".join("?" * len(chunk_ids)) @@ -338,8 +338,9 @@ class SqliteMemoryStore(BaseMemoryStore): ) cursor.execute("COMMIT") - except Exception: + except Exception as e: cursor.execute("ROLLBACK") + logger.error(f"Failed to delete chunks for {path}: {e}") raise finally: cursor.close() @@ -377,21 +378,25 @@ class SqliteMemoryStore(BaseMemoryStore): ) # Insert/update vector (vec0 doesn't support OR REPLACE, use DELETE + INSERT) - if self.vector_available: - assert chunk.embedding, "Embedding is required for vector insert" - # Delete existing vector first - cursor.execute( - f"DELETE FROM {self.vector_table_name} WHERE id = ?", - (chunk.id,), - ) - # Then insert new vector - cursor.execute( - f"INSERT INTO {self.vector_table_name} (id, embedding) VALUES (?, ?)", - (chunk.id, self.vector_to_blob(chunk.embedding)), - ) + if self.vector_enabled: + if not chunk.embedding: + logger.warning( + f"Chunk {chunk.id} missing embedding for vector insert, skipping vector indexing", + ) + else: + # Delete existing vector first + cursor.execute( + f"DELETE FROM {self.vector_table_name} WHERE id = ?", + (chunk.id,), + ) + # Then insert new vector + cursor.execute( + f"INSERT INTO {self.vector_table_name} (id, embedding) VALUES (?, ?)", + (chunk.id, self.vector_to_blob(chunk.embedding)), + ) # Insert/update FTS - if self.fts_available: + if self.fts_enabled: cursor.execute( f""" INSERT OR REPLACE INTO {self.fts_table_name} ( @@ -409,8 +414,9 @@ class SqliteMemoryStore(BaseMemoryStore): ) cursor.execute("COMMIT") - except Exception: + except Exception as e: cursor.execute("ROLLBACK") + logger.error(f"Failed to upsert chunks: {e}") raise finally: cursor.close() @@ -418,77 +424,91 @@ class SqliteMemoryStore(BaseMemoryStore): async def list_files(self, source: MemorySource) -> list[str]: """List all indexed files.""" cursor = self.conn.cursor() - cursor.execute(f"SELECT path FROM {self.files_table_name} WHERE source = ?", (source.value,)) - paths = [row[0] for row in cursor.fetchall()] - cursor.close() - return paths + try: + cursor.execute(f"SELECT path FROM {self.files_table_name} WHERE source = ?", (source.value,)) + paths = [row[0] for row in cursor.fetchall()] + return paths + except Exception as e: + logger.error(f"Failed to list files: {e}") + raise + finally: + cursor.close() async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None: """Get file metadata with chunk count.""" cursor = self.conn.cursor() - cursor.execute( - f"SELECT hash, mtime, size FROM {self.files_table_name} WHERE path = ? AND source = ?", - (path, source.value), - ) - row = cursor.fetchone() - if not row: + try: + cursor.execute( + f"SELECT hash, mtime, size FROM {self.files_table_name} WHERE path = ? AND source = ?", + (path, source.value), + ) + row = cursor.fetchone() + if not row: + return None + + hash_val, mtime, size = row + cursor.execute( + f"SELECT COUNT(*) FROM {self.chunks_table_name} WHERE path = ? AND source = ?", + (path, source.value), + ) + chunk_count = cursor.fetchone()[0] + + return FileMetadata( + hash=hash_val, + mtime_ms=mtime, + size=size, + path=path, + chunk_count=chunk_count, + ) + except Exception as e: + logger.error(f"Failed to get file metadata for {path}: {e}") + raise + finally: cursor.close() - return None - - hash_val, mtime, size = row - cursor.execute( - f"SELECT COUNT(*) FROM {self.chunks_table_name} WHERE path = ? AND source = ?", - (path, source.value), - ) - chunk_count = cursor.fetchone()[0] - cursor.close() - - return FileMetadata( - hash=hash_val, - mtime_ms=mtime, - size=size, - path=path, - chunk_count=chunk_count, - ) async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]: """Get all chunks for a file.""" cursor = self.conn.cursor() - cursor.execute( - f""" - SELECT id, path, source, start_line, end_line, text, hash, embedding - FROM {self.chunks_table_name} WHERE path = ? AND source = ? - ORDER BY start_line - """, - (path, source.value), - ) - - chunks = [] - for row in cursor.fetchall(): - chunk_id, path_val, source_val, start, end, text, hash_val, emb_str = row - # Parse embedding from JSON string - embedding = None - if emb_str: - try: - embedding = json.loads(emb_str) - except (json.JSONDecodeError, TypeError): - embedding = None - - chunks.append( - MemoryChunk( - id=chunk_id, - path=path_val, - source=MemorySource(source_val), - start_line=start, - end_line=end, - text=text, - hash=hash_val, - embedding=embedding, - ), + try: + cursor.execute( + f""" + SELECT id, path, source, start_line, end_line, text, hash, embedding + FROM {self.chunks_table_name} WHERE path = ? AND source = ? + ORDER BY start_line + """, + (path, source.value), ) - cursor.close() - return chunks + chunks = [] + for row in cursor.fetchall(): + chunk_id, path_val, source_val, start, end, text, hash_val, emb_str = row + # Parse embedding from JSON string + embedding = None + if emb_str: + try: + embedding = json.loads(emb_str) + except (json.JSONDecodeError, TypeError): + embedding = None + + chunks.append( + MemoryChunk( + id=chunk_id, + path=path_val, + source=MemorySource(source_val), + start_line=start, + end_line=end, + text=text, + hash=hash_val, + embedding=embedding, + ), + ) + + return chunks + except Exception as e: + logger.error(f"Failed to get file chunks for {path}: {e}") + raise + finally: + cursor.close() async def vector_search( self, @@ -497,7 +517,7 @@ class SqliteMemoryStore(BaseMemoryStore): sources: list[MemorySource] | None = None, ) -> list[MemorySearchResult]: """Perform vector similarity search.""" - if not self.vector_available or not query: + if not self.vector_enabled or not query: return [] # Get query embedding @@ -556,6 +576,7 @@ class SqliteMemoryStore(BaseMemoryStore): ), ) + results.sort(key=lambda r: r.score, reverse=True) return results except Exception as e: logger.error(f"Vector search failed: {e}") @@ -563,7 +584,8 @@ class SqliteMemoryStore(BaseMemoryStore): finally: cursor.close() - def _sanitize_fts_query(self, query: str) -> str: + @staticmethod + def _sanitize_fts_query(query: str) -> str: """Sanitize query string for FTS5 search. Removes or escapes special characters that have special meaning in FTS5: @@ -643,22 +665,41 @@ class SqliteMemoryStore(BaseMemoryStore): limit: int, sources: list[MemorySource] | None = None, ) -> list[MemorySearchResult]: - """Perform full-text search.""" - if not self.fts_available: + """Perform keyword search. + + Strategy: + - FTS5 trigram (fast path): used when ALL terms >= 3 chars (trigram minimum). + - LIKE (universal fallback): used when any term < 3 chars, covering CJK + short words, single/double-char queries, and mixed-length queries. + """ + if not self.fts_enabled: return [] - # Sanitize and prepare query cleaned = self._sanitize_fts_query(query) if not cleaned: return [] - # Split into words and escape double quotes for FTS5 phrase matching words = cleaned.split() if not words: return [] - # Use OR operator for better recall - match any of the query words - escaped_words = [word.replace('"', '""') for word in words] + # FTS5 trigram requires every term >= 3 characters + if all(len(w) >= 3 for w in words): + results = await self._fts_trigram_search(words, limit, sources) + if results: + return results + + # Universal fallback: LIKE-based substring search + return await self._like_search(cleaned, words, limit, sources) + + async def _fts_trigram_search( + self, + words: list[str], + limit: int, + sources: list[MemorySource] | None = None, + ) -> list[MemorySearchResult]: + """FTS5 trigram search. All terms must be >= 3 characters.""" + escaped_words = [w.replace('"', '""') for w in words] fts_query = " OR ".join(escaped_words) cursor = self.conn.cursor() @@ -685,24 +726,102 @@ class SqliteMemoryStore(BaseMemoryStore): results = [] for _, path, start, end, src, text, rank in cursor.fetchall(): - # Convert BM25 rank (negative) to 0-1 score (higher=better) score = max(0.0, 1.0 / (1.0 + abs(rank))) - snippet = text results.append( MemorySearchResult( path=path, start_line=start, end_line=end, score=score, - snippet=snippet, + snippet=text, source=MemorySource(src), raw_metric=rank, ), ) - + results.sort(key=lambda r: r.score, reverse=True) return results except Exception as e: - logger.error(f"Keyword search failed: {e}") + logger.error(f"FTS trigram search failed: {e}") + return [] + finally: + cursor.close() + + async def _like_search( + self, + phrase: str, + words: list[str], + limit: int, + sources: list[MemorySource] | None = None, + ) -> list[MemorySearchResult]: + """LIKE-based substring search with Python-side relevance scoring. + + Handles any term length and all languages (CJK, Latin, etc.). + Scores results by: word-match ratio + full-phrase bonus. + """ + cursor = self.conn.cursor() + try: + # Build OR conditions: match any individual word + like_clauses = [] + params: list = [] + for word in words: + like_clauses.append("c.text LIKE ?") + params.append(f"%{word}%") + + where_clause = " OR ".join(like_clauses) + + source_filter = "" + if sources: + placeholders = ",".join("?" * len(sources)) + source_filter = f" AND c.source IN ({placeholders})" + params.extend([s.value for s in sources]) + + # Fetch extra candidates for re-ranking in Python + fetch_limit = min(limit * 3, 200) + params.append(fetch_limit) + + cursor.execute( + f""" + SELECT c.id, c.path, c.start_line, c.end_line, c.source, c.text + FROM {self.chunks_table_name} c + WHERE ({where_clause}){source_filter} + LIMIT ? + """, + params, + ) + + results = [] + phrase_lower = phrase.lower() + words_lower = [w.lower() for w in words] + n_words = len(words) + + for _, path, start, end, src, text in cursor.fetchall(): + text_lower = text.lower() + + # Base score: proportion of query words found in text + match_count = sum(1 for w in words_lower if w in text_lower) + base_score = match_count / n_words + + # Bonus: full phrase appears as contiguous substring + phrase_bonus = 0.2 if n_words > 1 and phrase_lower in text_lower else 0.0 + + score = min(1.0, base_score * 0.8 + phrase_bonus) + + results.append( + MemorySearchResult( + path=path, + start_line=start, + end_line=end, + score=score, + snippet=text, + source=MemorySource(src), + ), + ) + + # Sort by score descending, return top `limit` + results.sort(key=lambda r: r.score, reverse=True) + return results[:limit] + except Exception as e: + logger.error(f"LIKE search failed: {e}") return [] finally: cursor.close() @@ -710,21 +829,22 @@ class SqliteMemoryStore(BaseMemoryStore): async def clear_all(self): """Clear all indexed data.""" cursor = self.conn.cursor() - cursor.execute("BEGIN") - try: + cursor.execute("BEGIN") + cursor.execute(f"DELETE FROM {self.files_table_name}") cursor.execute(f"DELETE FROM {self.chunks_table_name}") - if self.vector_available: + if self.vector_enabled: cursor.execute(f"DELETE FROM {self.vector_table_name}") - if self.fts_available: + if self.fts_enabled: cursor.execute(f"DELETE FROM {self.fts_table_name}") cursor.execute("COMMIT") - except Exception: + except Exception as e: cursor.execute("ROLLBACK") + logger.error(f"Failed to clear all data: {e}") raise finally: cursor.close() diff --git a/reme/core/schema/service_config.py b/reme/core/schema/service_config.py index 7152dc40..766d1fd3 100644 --- a/reme/core/schema/service_config.py +++ b/reme/core/schema/service_config.py @@ -85,7 +85,6 @@ class MemoryStoreConfig(BaseModel): backend: str = Field(default="sqlite") store_name: str = Field(default="reme") embedding_model: str = Field(default="default") - fts_enabled: bool = Field(default=True) class TokenCounterConfig(BaseModel): @@ -103,14 +102,8 @@ class FileWatcherConfig(BaseModel): model_config = ConfigDict(extra="allow") backend: str = Field(default="") + memory_store: str = Field(default="") watch_paths: list[str] = Field(default_factory=list) - suffix_filters: list[str] = Field(default_factory=list) - recursive: bool = Field(default=False) - debounce: int = Field(default=500) - chunk_tokens: int = Field(default=1000) - chunk_overlap: int = Field(default=100) - memory_store: str = Field(default="default") - scan_on_start: bool = Field(default=True) class ServiceConfig(BaseModel): @@ -120,6 +113,7 @@ class ServiceConfig(BaseModel): backend: str = Field(default="") app_name: str = Field(default=os.getenv("APP_NAME", "ReMe")) + working_dir: str | None = Field(default=None) enable_logo: bool = Field(default=True) language: str = Field(default="") thread_pool_max_workers: int = Field(default=16) diff --git a/reme/reme.py b/reme/reme.py index 7b013010..cab3a2c4 100644 --- a/reme/reme.py +++ b/reme/reme.py @@ -51,6 +51,7 @@ class ReMe(Application): llm_base_url: str | None = None, embedding_api_key: str | None = None, embedding_base_url: str | None = None, + working_dir: str = ".reme", config_path: str = "default", enable_logo: bool = True, log_to_console: bool = True, @@ -61,42 +62,18 @@ class ReMe(Application): target_user_names: list[str] | None = None, target_task_names: list[str] | None = None, target_tool_names: list[str] | None = None, - profile_dir: str = ".reme/profile", **kwargs, ): """Initialize ReMe with config. - Args: - *args: Arguments passed to Application - llm_api_key: API key for LLM provider - llm_base_url: API base for LLM provider - embedding_api_key: API key for embedding provider - embedding_base_url: 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 - default_token_counter_config: Token counter configuration - target_user_names: List of user names for personal memory - target_task_names: List of task names for procedural memory - target_tool_names: List of tool names for tool memory - profile_dir: Directory for profile storage - **kwargs: Additional keyword arguments passed to Application - Example: ```python reme = ReMe(...) await reme.start() - # reme = await ReMe.create(...) # both ok - await reme.summarize_memory(...) await reme.retrieve_memory(...) - await reme.close() ``` - """ super().__init__( *args, @@ -104,6 +81,7 @@ class ReMe(Application): llm_base_url=llm_base_url, embedding_api_key=embedding_api_key, embedding_base_url=embedding_base_url, + working_dir=working_dir, config_path=config_path, enable_logo=enable_logo, log_to_console=log_to_console, @@ -131,7 +109,10 @@ class ReMe(Application): memory_target_type_mapping[name] = MemoryType.TOOL self.service_context.memory_target_type_mapping = memory_target_type_mapping - self.profile_dir: str = profile_dir + + profile_path = Path(self.service_context.service_config.working_dir) / "profile" + profile_path.mkdir(parents=True, exist_ok=True) + self.profile_dir: str = str(profile_path) def _add_meta_memory(self, memory_type: str | MemoryType, memory_target: str): """Register or validate a memory target with the given memory type.""" diff --git a/reme/reme_cli.py b/reme/reme_cli.py new file mode 100644 index 00000000..476cdf2e --- /dev/null +++ b/reme/reme_cli.py @@ -0,0 +1,199 @@ +"""ReMe File System""" + +import asyncio +import sys +from typing import AsyncGenerator + +from prompt_toolkit import PromptSession + +from .agent.chat import FsCli +from .core.enumeration import ChunkEnum +from .core.schema import StreamChunk +from .core.utils import execute_stream_task +from .reme_fs import ReMeFs +from .tool.fs import ( + BashTool, + EditTool, + FsMemorySearch, + LsTool, + ReadTool, + WriteTool, +) +from .tool.gallery import ExecuteCode +from .tool.search import DashscopeSearch + + +class ReMeCli(ReMeFs): + """ReMe Cli""" + + def __init__(self, *args, **kwargs): + """Initialize ReMe with config.""" + super().__init__(*args, **kwargs) + self.commands = { + "/new": "Create a new conversation.", + "/compact": "Compact messages into a summary.", + "/exit": "Exit the application.", + "/clear": "Clear the history.", + "/help": "Show help.", + } + + async def chat_with_remy(self, tool_result_max_size: int = 100, language: str = "zh", **kwargs): + """Interactive CLI chat with Remy using simple streaming output.""" + fs_cli = FsCli( + tools=[ + FsMemorySearch( + hybrid_vector_weight=self.hybrid_vector_weight, + hybrid_candidate_multiplier=self.hybrid_candidate_multiplier, + ), + BashTool(cwd=self.working_dir), + LsTool(cwd=self.working_dir), + ReadTool(cwd=self.working_dir), + EditTool(cwd=self.working_dir), + WriteTool(cwd=self.working_dir), + ExecuteCode(), + DashscopeSearch(), + ], + context_window_tokens=self.context_window_tokens, + reserve_tokens=self.reserve_tokens, + keep_recent_tokens=self.keep_recent_tokens, + working_dir=self.working_dir, + language=language, + **kwargs, + ) + session = PromptSession() + + # Print welcome banner + print("\n========================================") + print(" Welcome to Remy Chat!") + 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: ") + user_input = user_input.strip() + if not user_input: + continue + + # Handle commands + if user_input == "/exit": + break + + if user_input == "/new": + result = await fs_cli.reset() + print(f"{result}\nConversation reset\n") + continue + + if user_input == "/compact": + result = await fs_cli.compact(force_compact=True) + print(f"{result}\nHistory compacted.\n") + continue + + if user_input == "/clear": + fs_cli.messages.clear() + print("History cleared.\n") + continue + + if user_input == "/help": + print("\nCommands:") + for command, description in self.commands.items(): + print(f" {command}: {description}") + 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 -> {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 ReMeCli(*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() diff --git a/reme/reme_fs.py b/reme/reme_fs.py index 65f06c84..33c335d6 100644 --- a/reme/reme_fs.py +++ b/reme/reme_fs.py @@ -1,19 +1,11 @@ """ReMe File System""" -import asyncio -import sys from pathlib import Path -from typing import AsyncGenerator -from prompt_toolkit import PromptSession - -from .agent.chat import FsCli from .agent.fs import FsCompactor, FsContextChecker, FsSummarizer from .config import ReMeConfigParser from .core import Application -from .core.enumeration import ChunkEnum -from .core.schema import Message, StreamChunk -from .core.utils import execute_stream_task +from .core.schema import Message from .tool.fs import ( BashTool, EditTool, @@ -23,8 +15,6 @@ from .tool.fs import ( ReadTool, WriteTool, ) -from .tool.gallery import ExecuteCode -from .tool.search import DashscopeSearch class ReMeFs(Application): @@ -46,6 +36,8 @@ class ReMeFs(Application): default_embedding_model_name: str | None = None, default_embedding_model_config: dict | None = None, default_store_name: str = "reme", + vector_enabled: bool = False, + fts_enabled: bool = True, default_memory_store_config: dict | None = None, token_counter_backend: str = "base", default_token_counter_config: dict | None = None, @@ -60,9 +52,7 @@ class ReMeFs(Application): 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, ): @@ -82,10 +72,20 @@ class ReMeFs(Application): default_embedding_model_config["model_name"] = default_embedding_model_name default_memory_store_config = default_memory_store_config or {} - default_memory_store_config["store_name"] = default_store_name + default_memory_store_config.update( + { + "store_name": default_store_name, + "vector_enabled": vector_enabled, + "fts_enabled": fts_enabled, + }, + ) default_token_counter_config = default_token_counter_config or {} - default_token_counter_config["backend"] = token_counter_backend + default_token_counter_config.update( + { + "backend": token_counter_backend, + }, + ) default_file_watcher_config = default_file_watcher_config or {} default_file_watcher_config.update( @@ -126,20 +126,9 @@ class ReMeFs(Application): 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 - # Commands - self.commands = { - "/new": "Create a new conversation.", - "/compact": "Compact messages into a summary.", - "/exit": "Exit the application.", - "/clear": "Clear the history.", - "/help": "Show help.", - } - async def context_check(self, messages: list[Message | dict]) -> dict: """Check if messages exceed context limits.""" checker = FsContextChecker( @@ -166,7 +155,14 @@ class ReMeFs(Application): service_context=self.service_context, ) - async def summary(self, messages: list[Message | dict], date: str, language: str = "zh", **kwargs): + async def summary( + self, + messages: list[Message | dict], + date: str, + version: str = "default", + language: str = "zh", + **kwargs, + ): """Generate a summary of the given messages.""" summarizer = FsSummarizer( tools=[ @@ -178,6 +174,7 @@ class ReMeFs(Application): ], working_dir=self.working_dir, language=language, + version=version, **kwargs, ) return await summarizer.call(messages=messages, date=date, service_context=self.service_context) @@ -197,9 +194,7 @@ class ReMeFs(Application): 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( @@ -234,171 +229,3 @@ class ReMeFs(Application): ) 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, language: str = "zh", **kwargs): - """Interactive CLI chat with Remy using simple streaming output.""" - fs_cli = FsCli( - working_dir=self.working_dir, - tools=[ - 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, - ), - BashTool(cwd=self.working_dir), - LsTool(cwd=self.working_dir), - ReadTool(cwd=self.working_dir), - EditTool(cwd=self.working_dir), - WriteTool(cwd=self.working_dir), - ExecuteCode(), - DashscopeSearch(), - ], - 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, - language=language, - **kwargs, - ) - session = PromptSession() - - # Print welcome banner - print("\n========================================") - print(" Welcome to Remy Chat!") - 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: ") - user_input = user_input.strip() - if not user_input: - continue - - # Handle commands - if user_input == "/exit": - break - - if user_input == "/new": - result = await fs_cli.reset() - print(f"{result}\nConversation reset\n") - continue - - if user_input == "/compact": - result = await fs_cli.compact(force_compact=True) - print(f"{result}\nHistory compacted.\n") - continue - - if user_input == "/clear": - fs_cli.messages.clear() - print("History cleared.\n") - continue - - if user_input == "/help": - print("\nCommands:") - for command, description in self.commands.items(): - print(f" {command}: {description}") - 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 -> {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() diff --git a/reme/tool/fs/fs_memory_search.py b/reme/tool/fs/fs_memory_search.py index fb749a1d..012dc0d8 100644 --- a/reme/tool/fs/fs_memory_search.py +++ b/reme/tool/fs/fs_memory_search.py @@ -17,21 +17,20 @@ class FsMemorySearch(BaseFsTool): sources: list[MemorySource] | None = None, min_score: float = 0.1, max_results: int = 5, - hybrid_enabled: bool = True, hybrid_vector_weight: float = 0.7, - hybrid_text_weight: float = 0.3, hybrid_candidate_multiplier: float = 3.0, **kwargs, ): """Initialize memory search tool.""" + assert ( + 0.0 <= hybrid_vector_weight <= 1.0 + ), f"hybrid_vector_weight must be between 0 and 1, got {hybrid_vector_weight}" kwargs.setdefault("name", "memory_search") super().__init__(**kwargs) self.sources = sources or [MemorySource.MEMORY] self.min_score = min_score self.max_results = max_results - self.hybrid_enabled = hybrid_enabled self.hybrid_vector_weight = hybrid_vector_weight - self.hybrid_text_weight = hybrid_text_weight self.hybrid_candidate_multiplier = hybrid_candidate_multiplier def _build_tool_call(self) -> ToolCall: @@ -71,11 +70,12 @@ class FsMemorySearch(BaseFsTool): max_results = self.context.get("max_results", self.max_results) candidates = min(200, max(1, int(max_results * self.hybrid_candidate_multiplier))) - # Perform hybrid search (vector + keyword) - if self.hybrid_enabled: - keyword_results = [] - if self.memory_store.fts_enabled: - keyword_results = await self._search_keyword(query, candidates) + vector_enabled = self.memory_store.vector_enabled + fts_enabled = self.memory_store.fts_enabled + + # Perform search based on enabled backends + if vector_enabled and fts_enabled: + keyword_results = await self._search_keyword(query, candidates) vector_results = await self._search_vector(query, candidates) # Log original vector results @@ -99,7 +99,7 @@ class FsMemorySearch(BaseFsTool): vector=vector_results, keyword=keyword_results, vector_weight=self.hybrid_vector_weight, - text_weight=self.hybrid_text_weight, + text_weight=1.0 - self.hybrid_vector_weight, ) # Log merged results @@ -109,9 +109,14 @@ class FsMemorySearch(BaseFsTool): logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") results = [r for r in merged if r.score >= min_score][:max_results] - else: + elif vector_enabled: vector_results = await self._search_vector(query, candidates) results = [r for r in vector_results if r.score >= min_score][:max_results] + elif fts_enabled: + keyword_results = await self._search_keyword(query, candidates) + results = [r for r in keyword_results if r.score >= min_score][:max_results] + else: + results = [] return json.dumps([result.model_dump(exclude_none=True) for result in results], indent=2, ensure_ascii=False) diff --git a/tests/test_fs_memory_search.py b/tests/test_fs_memory_search.py index 81ae44aa..d2a6256f 100644 --- a/tests/test_fs_memory_search.py +++ b/tests/test_fs_memory_search.py @@ -611,9 +611,7 @@ async def test_memory_search_hybrid_mode(): "fts_enabled": True, }, search_params={ - "hybrid_enabled": True, "hybrid_vector_weight": 0.7, - "hybrid_text_weight": 0.3, }, ) await reme_fs_hybrid.start() @@ -641,7 +639,7 @@ async def test_memory_search_hybrid_mode(): "embedding_model": "default", "fts_enabled": True, }, - search_params={"hybrid_enabled": False}, + search_params={}, ) await reme_fs_vector.start() result_json_vector = await reme_fs_vector.memory_search( @@ -674,9 +672,7 @@ async def test_memory_search_hybrid_mode(): "fts_enabled": True, }, search_params={ - "hybrid_enabled": True, "hybrid_vector_weight": vec_weight, - "hybrid_text_weight": text_weight, }, ) await reme_fs_weights.start()