mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
feat(cli): add ReMeCli class with interactive chat functionality
This commit is contained in:
parent
bba509465c
commit
eae6c63d15
16 changed files with 682 additions and 514 deletions
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"<conversation>\n{conversation}\n</conversation>\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(
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -5,7 +5,6 @@ on any change, ensuring complete synchronization.
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from loguru import logger
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
31
reme/reme.py
31
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."""
|
||||
|
|
|
|||
199
reme/reme_cli.py
Normal file
199
reme/reme_cli.py
Normal file
|
|
@ -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()
|
||||
221
reme/reme_fs.py
221
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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue