From 9983a9854dcaff39fad5171666c04347dd1e41da Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 9 Apr 2026 10:38:04 +0800 Subject: [PATCH] init --- reme_cli/application.py | 16 +- reme_cli/component/__init__.py | 5 +- reme_cli/component/application_context.py | 24 + reme_cli/component/as_llm/__init__.py | 31 +- .../component/as_llm_formatter/__init__.py | 32 +- reme_cli/component/base_component.py | 78 +- reme_cli/component/embedding/__init__.py | 10 + .../embedding/base_embedding_model.py | 407 ++++++++ .../embedding/openai_embedding_model.py | 56 + reme_cli/component/file_store/__init__.py | 23 - .../component/file_store/base_file_store.py | 227 ---- .../component/file_store/chroma_file_store.py | 633 ------------ .../component/file_store/local_file_store.py | 461 --------- .../component/file_store/sqlite_file_store.py | 978 ------------------ reme_cli/component/file_watcher/__init__.py | 19 - .../file_watcher/base_file_watcher.py | 240 ----- .../file_watcher/delta_file_watcher.py | 280 ----- .../file_watcher/full_file_watcher.py | 79 -- reme_cli/component/flow/__init__.py | 14 + reme_cli/component/flow/base_flow.py | 208 ++++ reme_cli/component/flow/cmd_flow.py | 18 + reme_cli/component/flow/expression_flow.py | 37 + reme_cli/{ => component}/op/__init__.py | 0 reme_cli/{ => component}/op/base_op.py | 0 reme_cli/component/registry_factory.py | 50 - reme_cli/enumeration/__init__.py | 2 + reme_cli/enumeration/json_schema_enum.py | 41 + reme_cli/schema/__init__.py | 7 + reme_cli/schema/application_config.py | 20 + reme_cli/schema/base_node.py | 10 + reme_cli/schema/service_config.py | 144 --- reme_cli/schema/tool_call.py | 196 ++++ reme_cli/utils/__init__.py | 4 + reme_cli/utils/logger_utils.py | 50 +- reme_cli/utils/pydantic_config_parser.py | 194 ++++ 35 files changed, 1414 insertions(+), 3180 deletions(-) create mode 100644 reme_cli/component/application_context.py create mode 100644 reme_cli/component/embedding/__init__.py create mode 100644 reme_cli/component/embedding/base_embedding_model.py create mode 100644 reme_cli/component/embedding/openai_embedding_model.py delete mode 100644 reme_cli/component/file_store/__init__.py delete mode 100644 reme_cli/component/file_store/base_file_store.py delete mode 100644 reme_cli/component/file_store/chroma_file_store.py delete mode 100644 reme_cli/component/file_store/local_file_store.py delete mode 100644 reme_cli/component/file_store/sqlite_file_store.py delete mode 100644 reme_cli/component/file_watcher/__init__.py delete mode 100644 reme_cli/component/file_watcher/base_file_watcher.py delete mode 100644 reme_cli/component/file_watcher/delta_file_watcher.py delete mode 100644 reme_cli/component/file_watcher/full_file_watcher.py create mode 100644 reme_cli/component/flow/__init__.py create mode 100644 reme_cli/component/flow/base_flow.py create mode 100644 reme_cli/component/flow/cmd_flow.py create mode 100644 reme_cli/component/flow/expression_flow.py rename reme_cli/{ => component}/op/__init__.py (100%) rename reme_cli/{ => component}/op/base_op.py (100%) delete mode 100644 reme_cli/component/registry_factory.py create mode 100644 reme_cli/enumeration/json_schema_enum.py create mode 100644 reme_cli/schema/application_config.py create mode 100644 reme_cli/schema/base_node.py delete mode 100644 reme_cli/schema/service_config.py create mode 100644 reme_cli/schema/tool_call.py create mode 100644 reme_cli/utils/pydantic_config_parser.py diff --git a/reme_cli/application.py b/reme_cli/application.py index a7522647..640c0c7c 100644 --- a/reme_cli/application.py +++ b/reme_cli/application.py @@ -1,20 +1,16 @@ -from reme_cli.component import BaseComponent +from .application_context import ApplicationContext +from .component import BaseComponent class Application(BaseComponent): """Application component for managing the main application.""" - def __init__(self) -> None: + def __init__(self, *args, config: str = "", **kwargs) -> None: super().__init__() - ... + self.context = ApplicationContext(*args, config=config, **kwargs) - async def start(self) -> None: + async def _start(self, app_context: ApplicationContext | None = None) -> None: """Start the application.""" - # 初始化llm formater - # - pass - async def close(self) -> None: + async def _close(self) -> None: """Close the application.""" - pass - diff --git a/reme_cli/component/__init__.py b/reme_cli/component/__init__.py index 695c24f7..e2b6373d 100644 --- a/reme_cli/component/__init__.py +++ b/reme_cli/component/__init__.py @@ -1,5 +1,8 @@ from .base_component import BaseComponent +from .component_registry import ComponentRegistry, R __all__ = [ "BaseComponent", -] \ No newline at end of file + "ComponentRegistry", + "R", +] diff --git a/reme_cli/component/application_context.py b/reme_cli/component/application_context.py new file mode 100644 index 00000000..c9f08b05 --- /dev/null +++ b/reme_cli/component/application_context.py @@ -0,0 +1,24 @@ +from typing import TYPE_CHECKING + +from ..enumeration import ComponentEnum +from ..schema import ApplicationConfig +from ..utils import PydanticConfigParser + +if TYPE_CHECKING: + from .base_component import BaseComponent + + +class ApplicationContext: + + def __init__(self, *args, config: str = "", **kwargs): + parser = PydanticConfigParser(config_class=ApplicationConfig) + self.app_config: ApplicationConfig = parser.parse_args(*args, config=config, **kwargs) + self.components: dict[ComponentEnum, dict[str, BaseComponent]] = {} + + from .component_registry import R + + for component_type, component_configs in self.app_config.components.items(): + self.components[component_type] = { + name: R.get(component_type, config.get("backend"))(**config) + for name, config in component_configs.items() + } diff --git a/reme_cli/component/as_llm/__init__.py b/reme_cli/component/as_llm/__init__.py index e0b694ac..1dda9bae 100644 --- a/reme_cli/component/as_llm/__init__.py +++ b/reme_cli/component/as_llm/__init__.py @@ -2,6 +2,33 @@ from agentscope.model import OpenAIChatModel -from ..registry_factory import R +from ..base_component import BaseComponent +from ..component_registry import R +from ...enumeration import ComponentEnum -R.as_llms.register("openai")(OpenAIChatModel) + +@R.register("openai") +class AsOpenAIChatModel(BaseComponent): + """Simple wrapper for AgentScope LLM models.""" + + component_type = ComponentEnum.AS_LLM + + def __init__(self, **kwargs) -> None: + """Initialize with model configuration.""" + super().__init__(**kwargs) + self.model: OpenAIChatModel | None = None + + async def _start(self, app_context=None) -> None: + """Initialize the AgentScope model instance.""" + self.model = OpenAIChatModel(**self.kwargs) + + async def _close(self) -> None: + """Close the AgentScope model and release resources.""" + if self.model is not None: + await self.model.client.close() + self.model = None + + +__all__ = [ + "AsOpenAIChatModel", +] diff --git a/reme_cli/component/as_llm_formatter/__init__.py b/reme_cli/component/as_llm_formatter/__init__.py index 1c7eee46..42d942e4 100644 --- a/reme_cli/component/as_llm_formatter/__init__.py +++ b/reme_cli/component/as_llm_formatter/__init__.py @@ -1,9 +1,33 @@ """Module for registering AgentScope LLM formatters.""" -from agentscope.formatter import DashScopeChatFormatter +from agentscope.formatter import OpenAIChatFormatter from .reme_openai_chat_formatter import ReMeOpenAIChatFormatter -from ..registry_factory import R +from ..base_component import BaseComponent +from ..component_registry import R +from ...enumeration import ComponentEnum -R.as_llm_formatters.register("openai")(ReMeOpenAIChatFormatter) -R.as_llm_formatters.register("dashscope")(DashScopeChatFormatter) + +@R.register("openai") +class AsOpenAIChatFormatter(BaseComponent): + """Wrapper for ReMeOpenAIChatFormatter.""" + + component_type = ComponentEnum.AS_LLM_FORMATTER + + def __init__(self, **kwargs) -> None: + """Initialize with formatter configuration.""" + super().__init__(**kwargs) + self.formatter: OpenAIChatFormatter | None = None + + async def _start(self, app_context=None) -> None: + """Initialize the formatter instance.""" + self.formatter = ReMeOpenAIChatFormatter(**self.kwargs) + + async def _close(self) -> None: + """Close the formatter (no-op for formatter).""" + self.formatter = None + + +__all__ = [ + "AsOpenAIChatFormatter", +] diff --git a/reme_cli/component/base_component.py b/reme_cli/component/base_component.py index 0988320c..401bc6b2 100644 --- a/reme_cli/component/base_component.py +++ b/reme_cli/component/base_component.py @@ -1,32 +1,92 @@ """Base class for components.""" from abc import ABC, abstractmethod +from types import TracebackType +from typing import TYPE_CHECKING from ..enumeration import ComponentEnum +from ..utils.logger_utils import get_logger + +if TYPE_CHECKING: + from .application_context import ApplicationContext class BaseComponent(ABC): - """Base class supporting async start/close and async context management.""" + """Base class supporting async start/close and async context management. + + Provides lifecycle management with state tracking to prevent duplicate + start/close operations. + + Attributes: + component_type: The type identifier for this component. + _is_started: Internal flag tracking whether the component has been started. + """ component_type = ComponentEnum.BASE - @abstractmethod - async def start(self) -> None: - """Start the component asynchronously.""" + def __init__(self, **kwargs) -> None: + """Initialize the component with default state.""" + self.kwargs: dict = kwargs + self.logger = get_logger() + if hasattr(self.logger, "bind"): + self.logger = self.logger.bind(component=self.__class__.__name__) + self._is_started: bool = False @abstractmethod + async def _start(self, app_context: ApplicationContext | None = None) -> None: + """Internal method to perform the actual start logic. + + Subclasses should implement this instead of start(). + """ + + @abstractmethod + async def _close(self) -> None: + """Internal method to perform the actual close logic. + + Subclasses should implement this instead of close(). + """ + + async def start(self, app_context: ApplicationContext | None = None) -> None: + """Start the component asynchronously. + + Does nothing if already started. + """ + if self._is_started: + return + await self._start(app_context) + self._is_started = True + async def close(self) -> None: - """Close the component asynchronously.""" + """Close the component asynchronously. + + Does nothing if not started or already closed. + """ + if not self._is_started: + return + await self._close() + self._is_started = False + + async def restart(self, app_context: ApplicationContext | None = None) -> None: + """Restart the component by closing and then starting again.""" + await self.close() + await self.start(app_context) + + @property + def is_started(self) -> bool: + """Check if the component is currently started.""" + return self._is_started async def __aenter__(self) -> "BaseComponent": """Enter async context manager.""" await self.start() return self - async def __aexit__(self, exc_type, exc_val, exc_tb) -> bool: + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_val: BaseException | None, + exc_tb: TracebackType | None, + ) -> bool: """Exit async context manager.""" await self.close() - - if exc_type is not None: - return True return False diff --git a/reme_cli/component/embedding/__init__.py b/reme_cli/component/embedding/__init__.py new file mode 100644 index 00000000..4989f4a7 --- /dev/null +++ b/reme_cli/component/embedding/__init__.py @@ -0,0 +1,10 @@ +"""embedding""" + +from .base_embedding_model import BaseEmbeddingModel +from .openai_embedding_model import OpenAIEmbeddingModel + +__all__ = [ + "BaseEmbeddingModel", + "OpenAIEmbeddingModel", +] + diff --git a/reme_cli/component/embedding/base_embedding_model.py b/reme_cli/component/embedding/base_embedding_model.py new file mode 100644 index 00000000..21ed8d02 --- /dev/null +++ b/reme_cli/component/embedding/base_embedding_model.py @@ -0,0 +1,407 @@ +"""Base embedding model interface for ReMe. + +Defines the abstract base class and standard API for all embedding model implementations. +""" + +import asyncio +import hashlib +import json +import time +from abc import abstractmethod +from collections import OrderedDict +from pathlib import Path + +from ..base_component import BaseComponent +from ...schema import BaseNode + + +class BaseEmbeddingModel(BaseComponent): + """Abstract base class for embedding model implementations. + + Provides a standard interface for text-to-vector generation with + built-in batching, retry logic, and error handling. + """ + + def __init__( + self, + api_key: str | None = None, + base_url: str | None = None, + model_name: str = "", + dimensions: int = 1024, + use_dimensions: bool = False, + max_batch_size: int = 10, + max_retries: int = 3, + raise_exception: bool = True, + max_input_length: int = 8192, + cache_dir: str | Path = ".reme", + max_cache_size: int = 2000, + enable_cache: bool = True, + encoding: str = "utf-8", + **kwargs, + ): + """Initialize model configuration and parameters. + + Args: + api_key: API key for the embedding service + base_url: Base URL for the embedding service + model_name: Name of the embedding model + dimensions: Vector dimensions of the embeddings + use_dimensions: Whether to pass dimensions parameter to API (some APIs don't support it) + max_batch_size: Maximum batch size for embedding requests + max_retries: Maximum number of retry attempts on failure + raise_exception: Whether to raise exceptions on failure + max_input_length: Maximum input text length + max_cache_size: Maximum number of embeddings to cache in memory (LRU) + enable_cache: Whether to enable embedding cache + encoding: Text encoding for cache file operations + **kwargs: Additional model-specific parameters + """ + super().__init__(**kwargs) + self.api_key: str | None = api_key + self.base_url: str | None = base_url + self.model_name = model_name + self.dimensions = dimensions + self.use_dimensions = use_dimensions + self.max_batch_size = max_batch_size + self.max_retries = max_retries + self.raise_exception = raise_exception + self.max_input_length = max_input_length + self.cache_dir = cache_dir + self.max_cache_size = max_cache_size + self.enable_cache = enable_cache + self.encoding = encoding + + self._embedding_cache: OrderedDict[str, list[float]] = OrderedDict() + self._cache_hits = 0 + self._cache_misses = 0 + + self.cache_path: Path = Path(self.cache_dir) + + def _truncate_text(self, text: str) -> str: + return text[: self.max_input_length] if len(text) > self.max_input_length else text + + def _validate_and_adjust_embedding(self, embedding: list[float]) -> list[float]: + """Validate and adjust embedding dimensions to match expected dimensions. + + Args: + embedding: The embedding vector to validate + + Returns: + Embedding vector adjusted to match self.dimensions + """ + actual_len = len(embedding) + if actual_len == self.dimensions: + return embedding + + elif actual_len < self.dimensions: + self.logger.warning( + f"[ACTUAL_EMB_LENGTH]Embedding dimensions {actual_len} is less than expected {self.dimensions}, " + f"padding with zeros", + ) + return embedding + [0.0] * (self.dimensions - actual_len) + + else: + self.logger.warning( + f"[ACTUAL_EMB_LENGTH]Embedding dimensions {actual_len} is greater than expected {self.dimensions}, " + f"truncating to {self.dimensions}", + ) + return embedding[: self.dimensions] + + def _get_cache_key(self, text: str) -> str: + """Generate a cache key by hashing text + model_name + dimensions.""" + cache_string = f"{text}|{self.model_name}|{self.dimensions}" + return hashlib.sha256(cache_string.encode(self.encoding)).hexdigest() + + def _get_cache_file_path(self) -> Path: + """Get the path to the cache file. + + Returns: + Path to the embedding cache JSONL file + """ + return self.cache_path / "embedding_cache.jsonl" + + def _load_cache(self) -> None: + """Load embedding cache from disk (JSONL format). + + Each line in the JSONL file contains a JSON object with: + - key: the cache key (SHA256 hash) + - embedding: the embedding vector (list of floats) + + Loads in reverse order (newest first) to prioritize recent embeddings + when max_cache_size is smaller than the file content. + """ + if not self.enable_cache: + return + + self.cache_path.mkdir(parents=True, exist_ok=True) + cache_file = self._get_cache_file_path() + if not cache_file.exists(): + self.logger.info(f"No cache file found at {cache_file}, starting with empty cache") + return + + try: + load_start = time.time() + # Read all lines first (to load in reverse order) + with open(cache_file, "r", encoding=self.encoding) as f: + lines = f.readlines() + + loaded_count = 0 + # Load in reverse order (newest entries first) + for _, line in enumerate(reversed(lines), 1): + line = line.strip() + if not line: + continue + try: + data = json.loads(line) + except json.JSONDecodeError as e: + self.logger.warning(f"Failed to parse line in cache file: {e}") + continue + + if not data: + continue + # Each line is {cache_key: embedding} + cache_key, embedding = next(iter(data.items())) + + if cache_key and embedding and isinstance(embedding, list): + # Skip if already loaded (keep the newest) + if cache_key in self._embedding_cache: + continue + + if len(embedding) != self.dimensions: + self.logger.warning( + f"Embedding dimensions mismatch for cache key {cache_key}, " + f"expected {self.dimensions}, got {len(embedding)}", + ) + continue + + # Respect max_cache_size during loading + if len(self._embedding_cache) >= self.max_cache_size: + self.logger.info( + f"Cache size limit reached ({self.max_cache_size}), " + f"loaded {loaded_count} newest entries", + ) + break + self._embedding_cache[cache_key] = embedding + loaded_count += 1 + + self.logger.info( + f"Loaded {loaded_count} embeddings from cache file: {cache_file} in {time.time() - load_start:.2f}s", + ) + except Exception as e: + self.logger.error(f"Failed to load cache from {cache_file}: {e}, deleting cache file") + try: + cache_file.unlink() + self.logger.info(f"Deleted corrupted cache file: {cache_file}") + except Exception as del_e: + self.logger.error(f"Failed to delete cache file {cache_file}: {del_e}") + + def _save_cache(self) -> None: + """Save embedding cache to disk (JSONL format). + + Each line contains a JSON object with the cache key and embedding vector. + Only saves if cache is non-empty. + """ + if not self.enable_cache: + return + + self.logger.info(f"Attempting to save cache, current size: {len(self._embedding_cache)}") + if not self._embedding_cache: + self.logger.info("Cache is empty, skipping save") + return + + cache_file = self._get_cache_file_path() + try: + with open(cache_file, "w", encoding=self.encoding) as f: + for cache_key, embedding in self._embedding_cache.items(): + if len(embedding) != self.dimensions: + self.logger.warning( + f"Embedding dimensions mismatch for cache key {cache_key}, " + f"expected {self.dimensions}, got {len(embedding)}", + ) + continue + cache_entry = {cache_key: embedding} + f.write(json.dumps(cache_entry, ensure_ascii=False) + "\n") + + self.logger.info(f"Saved {len(self._embedding_cache)} embeddings to cache file: {cache_file}") + except Exception as e: + self.logger.error(f"Failed to save cache to {cache_file}: {e}") + + def _get_from_cache(self, text: str) -> list[float] | None: + if not self.enable_cache: + return None + + cache_key = self._get_cache_key(text) + if cache_key not in self._embedding_cache: + self._cache_misses += 1 + return None + + embeddings = self._embedding_cache[cache_key] + if len(embeddings) != self.dimensions: + self.logger.warning( + f"Cached embedding dimensions mismatch: expected {self.dimensions}, " + f"got {len(embeddings)}. Removing invalid cache entry.", + ) + del self._embedding_cache[cache_key] + self._cache_misses += 1 + return None + + self._embedding_cache.move_to_end(cache_key) + self._cache_hits += 1 + text_preview = text[:50] + "..." if len(text) > 50 else text + self.logger.info(f"Cache hit for text: {text_preview} (hits: {self._cache_hits}, misses: {self._cache_misses})") + return embeddings + + def _put_to_cache(self, text: str, embedding: list[float]) -> None: + if not self.enable_cache or self.max_cache_size <= 0: + return + + cache_key = self._get_cache_key(text) + if len(embedding) != self.dimensions: + self.logger.warning( + f"[PUT_TO_CACHE] Embedding dimensions mismatch for cache key {cache_key}, " + f"expected {self.dimensions}, got real length {len(embedding)}", + ) + return + + if len(self._embedding_cache) >= self.max_cache_size and cache_key not in self._embedding_cache: + self._embedding_cache.popitem(last=False) + + self._embedding_cache[cache_key] = embedding + self._embedding_cache.move_to_end(cache_key) + + def get_cache_stats(self) -> dict[str, int]: + """Get cache statistics. + + Returns: + Dictionary with cache size, hits, misses, and hit rate + """ + total_requests = self._cache_hits + self._cache_misses + hit_rate = self._cache_hits / total_requests if total_requests > 0 else 0.0 + return { + "cache_size": len(self._embedding_cache), + "max_cache_size": self.max_cache_size, + "cache_hits": self._cache_hits, + "cache_misses": self._cache_misses, + "hit_rate": hit_rate, + } + + def clear_cache(self) -> None: + """Clear the embedding cache and reset statistics.""" + self._embedding_cache.clear() + self._cache_hits = 0 + self._cache_misses = 0 + + @abstractmethod + async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]: + """Internal async implementation for calling the embedding API with batch input.""" + + async def get_embedding(self, input_text: str, **kwargs) -> list[float]: + truncated_text = self._truncate_text(input_text) + cached_embedding = self._get_from_cache(truncated_text) + if cached_embedding is not None: + return cached_embedding + + for retry in range(self.max_retries): + try: + result = await self._get_embeddings([truncated_text], **kwargs) + if result and len(result) == 1: + embedding = self._validate_and_adjust_embedding(result[0]) + self._put_to_cache(truncated_text, embedding) + return embedding + # Empty or mismatched result, treat as failure for retry + self.logger.warning( + f"Model {self.model_name} returned {len(result) if result else 0} results, expected 1" + ) + if retry == self.max_retries - 1: + if self.raise_exception: + raise RuntimeError("Embedding API returned empty result") + return [] + await asyncio.sleep(retry + 1) + except Exception as e: + self.logger.error(f"Model {self.model_name} failed: {e}") + if retry == self.max_retries - 1: + if self.raise_exception: + raise + return [] + await asyncio.sleep(retry + 1) + return [] + + async def get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]: + truncated_texts = [self._truncate_text(t) for t in input_text] + results: list[list[float] | None] = [None] * len(truncated_texts) + texts_to_compute: list[tuple[int, str]] = [] + + for idx, text in enumerate(truncated_texts): + cached = self._get_from_cache(text) + if cached is not None: + results[idx] = cached + else: + texts_to_compute.append((idx, text)) + + if texts_to_compute: + uncached_texts = [text for _, text in texts_to_compute] + for i in range(0, len(uncached_texts), self.max_batch_size): + batch_texts = uncached_texts[i: i + self.max_batch_size] + batch_indices = [idx for idx, _ in texts_to_compute[i: i + self.max_batch_size]] + + for retry in range(self.max_retries): + try: + batch_embeddings = await self._get_embeddings(batch_texts, **kwargs) + if batch_embeddings and len(batch_embeddings) == len(batch_texts): + for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings): + adjusted_embedding = self._validate_and_adjust_embedding(embedding) + results[orig_idx] = adjusted_embedding + self._put_to_cache(text, adjusted_embedding) + break # Success, exit retry loop + else: + self.logger.warning( + f"Batch embedding returned {len(batch_embeddings) if batch_embeddings else 0} results " + f"for {len(batch_texts)} inputs" + ) + if retry == self.max_retries - 1: + if self.raise_exception: + raise RuntimeError( + f"Batch embedding returned {len(batch_embeddings) if batch_embeddings else 0} " + f"results for {len(batch_texts)} inputs after {self.max_retries} retries" + ) + # Fill failed positions with empty lists + for orig_idx in batch_indices: + if results[orig_idx] is None: + results[orig_idx] = [] + else: + await asyncio.sleep(retry + 1) + except Exception as e: + self.logger.error(f"Model {self.model_name} batch failed: {e}") + if retry == self.max_retries - 1: + if self.raise_exception: + raise + # Fill failed positions with empty lists + for orig_idx in batch_indices: + if results[orig_idx] is None: + results[orig_idx] = [] + else: + await asyncio.sleep(retry + 1) + + return [r if r is not None else [] for r in results] + + async def get_node_embeddings(self, nodes: list[BaseNode], **kwargs) -> list[BaseNode]: + texts = [node.text for node in nodes] + embeddings = await self.get_embeddings(texts, **kwargs) + + if len(embeddings) == len(nodes): + for node, vec in zip(nodes, embeddings): + node.embedding = vec + else: + self.logger.warning( + f"Mismatch: got {len(embeddings)} vectors for {len(nodes)} nodes, " + f"skipping embedding assignment" + ) + return nodes + + async def _start(self, app_context=None) -> None: + """Initialize resources and load cache.""" + self._load_cache() + + async def _close(self) -> None: + """Release resources and save cache.""" + self._save_cache() diff --git a/reme_cli/component/embedding/openai_embedding_model.py b/reme_cli/component/embedding/openai_embedding_model.py new file mode 100644 index 00000000..cd2df068 --- /dev/null +++ b/reme_cli/component/embedding/openai_embedding_model.py @@ -0,0 +1,56 @@ +"""Asynchronous OpenAI-compatible embedding model implementation for ReMe.""" + +from openai import AsyncOpenAI + +from .base_embedding_model import BaseEmbeddingModel +from ..component_registry import R + + +@R.register("openai") +class OpenAIEmbeddingModel(BaseEmbeddingModel): + """Asynchronous embedding model implementation compatible with OpenAI-style APIs.""" + + def __init__(self, **kwargs): + """Initialize the OpenAI async embedding model with API credentials and configuration.""" + super().__init__(**kwargs) + self._client: AsyncOpenAI | None = None + + async def _start(self, app_context=None) -> None: + """Initialize the AsyncOpenAI client.""" + self._client = AsyncOpenAI(api_key=self.api_key, base_url=self.base_url, **self.kwargs) + await super()._start(app_context) + + async def _close(self) -> None: + """Close the AsyncOpenAI client and release resources.""" + if self._client is not None: + await self._client.close() + self._client = None + await super()._close() + + async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]: + """Fetch embeddings from the API for a batch of strings.""" + if self._client is None: + raise RuntimeError("Client not initialized. Call _start() first.") + + create_kwargs: dict = { + "model": self.model_name, + "input": input_text, + **kwargs, + } + if self.use_dimensions: + create_kwargs["dimensions"] = self.dimensions + + completion = await self._client.embeddings.create(**create_kwargs) + + result_emb: list[list[float] | None] = [None] * len(input_text) + for emb in completion.data: + vec = getattr(emb, "embedding", None) or getattr(emb, "dense_embedding", None) + if 0 <= emb.index < len(input_text): + if vec is not None: + result_emb[emb.index] = list(vec) + else: + self.logger.warning(f"Empty embedding returned for index {emb.index}") + else: + self.logger.warning(f"Invalid index {emb.index} for input length {len(input_text)}") + + return [r if r is not None else [] for r in result_emb] diff --git a/reme_cli/component/file_store/__init__.py b/reme_cli/component/file_store/__init__.py deleted file mode 100644 index 1358df52..00000000 --- a/reme_cli/component/file_store/__init__.py +++ /dev/null @@ -1,23 +0,0 @@ -"""File store module for persistent memory management. - -This module provides storage backends for memory chunks and file metadata, -including SQLite-based, ChromaDB-based, and pure-Python local implementations -with vector and full-text search. -""" - -from .base_file_store import BaseFileStore -from .chroma_file_store import ChromaFileStore -from .local_file_store import LocalFileStore -from .sqlite_file_store import SqliteFileStore -from ..registry_factory import R - -__all__ = [ - "BaseFileStore", - "ChromaFileStore", - "LocalFileStore", - "SqliteFileStore", -] - -R.file_stores.register("sqlite")(SqliteFileStore) -R.file_stores.register("chroma")(ChromaFileStore) -R.file_stores.register("local")(LocalFileStore) diff --git a/reme_cli/component/file_store/base_file_store.py b/reme_cli/component/file_store/base_file_store.py deleted file mode 100644 index 6a6b0901..00000000 --- a/reme_cli/component/file_store/base_file_store.py +++ /dev/null @@ -1,227 +0,0 @@ -"""Base storage interface for file store.""" - -import re -from abc import ABC, abstractmethod -from pathlib import Path - -from ..embedding import BaseEmbeddingModel -from ..enumeration import MemorySource -from ..schema import FileMetadata, MemoryChunk, MemorySearchResult -from ..utils import get_logger - -logger = get_logger() - - -class BaseFileStore(ABC): - """Abstract base class for file storage backends.""" - - def __init__( - self, - store_name: str, - db_path: str | Path, - embedding_model: BaseEmbeddingModel | None = None, - vector_enabled: bool = False, - fts_enabled: bool = True, - **kwargs, - ): - """Initialize""" - # Validate store_name to prevent SQL injection - # 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.") - - # Ensure embedding_model is provided when vector search is enabled - if vector_enabled and embedding_model is None: - raise ValueError("embedding_model is required when vector_enabled is True.") - - self.store_name: str = store_name - self.db_path: Path = Path(db_path) - self.db_path.mkdir(parents=True, exist_ok=True) - self.embedding_model: BaseEmbeddingModel | None = embedding_model - self.vector_enabled: bool = vector_enabled - self.fts_enabled: bool = fts_enabled - self.kwargs: dict = kwargs - - @property - def embedding_dim(self) -> int: - """Get the embedding model's dimensionality.""" - if self.embedding_model is None: - return 1024 - return self.embedding_model.dimensions - - def _get_mock_embedding(self) -> list[float]: - """Generate a zero vector based on embedding model dimensions.""" - return [0.0] * self.embedding_dim - - def _disable_vector_search(self, reason: str = "embedding API error") -> None: - """Disable vector search and log a warning.""" - if self.vector_enabled: - logger.warning( - f"[{self.store_name}] Disabling vector search due to {reason}. " - "Falling back to full-text search only.", - ) - self.vector_enabled = False - - async def get_embedding(self, query: str, **kwargs) -> list[float]: - """Get embedding for a single query string.""" - if not self.vector_enabled: - return self._get_mock_embedding() - try: - return await self.embedding_model.get_embedding(query, **kwargs) - except Exception as e: - self._disable_vector_search(str(e)) - return self._get_mock_embedding() - - async def get_embeddings(self, queries: list[str], **kwargs) -> list[list[float]]: - """Get embeddings for a batch of query strings.""" - if not self.vector_enabled: - return [self._get_mock_embedding() for _ in queries] - try: - return await self.embedding_model.get_embeddings(queries, **kwargs) - except Exception as e: - self._disable_vector_search(str(e)) - return [self._get_mock_embedding() for _ in queries] - - async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk: - """Generate and populate embedding field for a single MemoryChunk object.""" - if not self.vector_enabled: - chunk.embedding = self._get_mock_embedding() - return chunk - try: - return await self.embedding_model.get_chunk_embedding(chunk, **kwargs) - except Exception as e: - self._disable_vector_search(str(e)) - chunk.embedding = self._get_mock_embedding() - return chunk - - async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]: - """Generate and populate embedding fields for a batch of MemoryChunk objects.""" - if not self.vector_enabled: - mock_embedding = self._get_mock_embedding() - for chunk in chunks: - chunk.embedding = mock_embedding.copy() - return chunks - try: - return await self.embedding_model.get_chunk_embeddings(chunks, **kwargs) - except Exception as e: - self._disable_vector_search(str(e)) - mock_embedding = self._get_mock_embedding() - for chunk in chunks: - chunk.embedding = mock_embedding.copy() - return chunks - - @abstractmethod - async def start(self): - """Initialize the storage backend.""" - - @abstractmethod - async def upsert_file(self, file_meta: FileMetadata, source: MemorySource, chunks: list[MemoryChunk]): - """Insert or update a file and its chunks.""" - - @abstractmethod - async def delete_file(self, path: str, source: MemorySource): - """Delete a file and all its chunks.""" - - @abstractmethod - async def delete_file_chunks(self, path: str, chunk_ids: list[str]): - """Delete chunks for a file.""" - - @abstractmethod - async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource): - """Insert or update specific chunks without affecting other chunks.""" - - @abstractmethod - async def list_files(self, source: MemorySource) -> list[str]: - """List all indexed file paths for a source.""" - - @abstractmethod - async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None: - """Get full file metadata with statistics.""" - - @abstractmethod - async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None: - """Update file metadata without affecting chunks. - - This is useful for incremental updates where only metadata needs to be updated - (e.g., after adding/removing chunks in delta file watcher). - - Args: - file_meta: Updated file metadata (hash, mtime_ms, size, chunk_count) - source: Memory source - """ - - @abstractmethod - async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]: - """Get all chunks for a file.""" - - @abstractmethod - async def vector_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform vector similarity search. - - Args: - query: Query embedding vector - limit: Maximum number of results - sources: Optional list of sources to filter - - Returns: - List of search results sorted by similarity - """ - - @abstractmethod - async def keyword_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform keyword/full-text search. - - Args: - query: Search query text - limit: Maximum number of results - sources: Optional list of sources to filter - - Returns: - List of search results sorted by relevance - """ - - @abstractmethod - async def hybrid_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - vector_weight: float = 0.7, - candidate_multiplier: float = 3.0, - ) -> list[MemorySearchResult]: - """Perform hybrid search combining vector and keyword search. - - Args: - query: Search query text - limit: Maximum number of results - sources: Optional list of sources to filter - vector_weight: Weight for vector search results (0.0-1.0). - Keyword weight = 1.0 - vector_weight. - candidate_multiplier: Multiplier for candidate pool size. - candidates = limit * candidate_multiplier - - Returns: - List of search results sorted by combined relevance score - """ - - @abstractmethod - async def clear_all(self): - """Clear all indexed data.""" - - @abstractmethod - async def close(self): - """Close storage and release resources.""" diff --git a/reme_cli/component/file_store/chroma_file_store.py b/reme_cli/component/file_store/chroma_file_store.py deleted file mode 100644 index 6dda41b8..00000000 --- a/reme_cli/component/file_store/chroma_file_store.py +++ /dev/null @@ -1,633 +0,0 @@ -"""ChromaDB storage backend for file store.""" - -import json -import random -import time -from pathlib import Path - -from .base_file_store import BaseFileStore -from ..enumeration import MemorySource -from ..schema import FileMetadata, MemoryChunk, MemorySearchResult -from ..utils import get_logger - -logger = get_logger() - -try: - import chromadb - from chromadb.config import Settings - - _CHROMADB_IMPORT_ERROR: Exception | None = None -except Exception as e: - _CHROMADB_IMPORT_ERROR = e - chromadb = None - Settings = None - - -class ChromaFileStore(BaseFileStore): - """ChromaDB file storage with vector and full-text search. - - Inherits embedding methods from BaseFileStore: - - get_chunk_embedding / get_chunk_embeddings (async) - - get_embedding / get_embeddings (async) - - Provides ChromaDB-backed persistent storage with: - - Vector similarity search (native ChromaDB) - - Full-text search (via ChromaDB where_document filter) - - Efficient chunk and file metadata management - """ - - def __init__( - self, - **kwargs, - ): - if _CHROMADB_IMPORT_ERROR is not None: - raise _CHROMADB_IMPORT_ERROR - - super().__init__(**kwargs) - self.client: "chromadb.ClientAPI | None" = None - self.chunks_collection: "chromadb.Collection | None" = None - # Initialize metadata file path (db_path and store_name are set by base class) - self._metadata_file: Path = self.db_path.parent / f"{self.store_name}_file_metadata.json" - self._metadata_cache: dict[str, dict[str, FileMetadata]] = {} - - @property - def collection_name(self) -> str: - """Get the name of the ChromaDB collection for this store.""" - return f"chunks_{self.store_name}" - - async def _load_metadata(self) -> dict[str, dict[str, FileMetadata]]: - """Load file metadata from disk. - - Returns: - Dictionary mapping source -> path -> FileMetadata - """ - if not self._metadata_file.exists(): - return {} - - try: - data = self._metadata_file.read_text(encoding="utf-8") - metadata_dict = json.loads(data) - - # Convert dict to FileMetadata objects - result = {} - for source, files in metadata_dict.items(): - result[source] = {} - for path, meta in files.items(): - result[source][path] = FileMetadata(**meta) - - logger.debug(f"Loaded file metadata from {self._metadata_file}") - return result - except Exception as e: - logger.warning(f"Failed to load file metadata from {self._metadata_file}: {e}") - return {} - - async def _save_metadata(self, metadata: dict[str, dict[str, FileMetadata]]) -> None: - """Save file metadata to disk. - - Args: - metadata: Dictionary mapping source -> path -> FileMetadata - """ - try: - # Convert FileMetadata objects to dict for JSON serialization - metadata_dict = {} - for source, files in metadata.items(): - metadata_dict[source] = {} - for path, meta in files.items(): - metadata_dict[source][path] = { - "path": meta.path, - "hash": meta.hash, - "mtime_ms": meta.mtime_ms, - "size": meta.size, - "chunk_count": meta.chunk_count, - } - - data = json.dumps(metadata_dict, indent=2, ensure_ascii=False) - self._metadata_file.write_text(data, encoding="utf-8") - logger.debug(f"Saved file metadata to {self._metadata_file}") - except Exception as e: - logger.error(f"Failed to save file metadata to {self._metadata_file}: {e}") - - async def start(self) -> None: - """Initialize ChromaDB client and collection.""" - if self.client is not None: - return - - # Initialize persistent ChromaDB client - self.client = chromadb.PersistentClient( - path=str(self.db_path), - settings=Settings( - anonymized_telemetry=False, - allow_reset=True, - ), - ) - - # Get or create the chunks collection - # ChromaDB uses cosine distance by default for similarity - self.chunks_collection = self.client.get_or_create_collection( - name=self.collection_name, - metadata={"hnsw:space": "cosine"}, - ) - - # Load metadata into cache - self._metadata_cache = await self._load_metadata() - - logger.info(f"ChromaDB initialized with collection: {self.collection_name}") - logger.info(f"File metadata will be persisted to: {self._metadata_file}") - - async def upsert_file( - self, - file_meta: FileMetadata, - source: MemorySource, - chunks: list[MemoryChunk], - ) -> None: - """Insert or update file and its chunks.""" - if not chunks: - return - - # Delete existing chunks for this file first - await self.delete_file(file_meta.path, source) - - # Batch generate embeddings for all chunks - # (base class returns mock embeddings when vector_enabled=False) - chunks = await self.get_chunk_embeddings(chunks) - - # Prepare data for ChromaDB batch upsert - ids = [] - documents = [] - embeddings = [] - metadatas = [] - - now = int(time.time() * 1000) - for chunk in chunks: - ids.append(chunk.id) - documents.append(chunk.text) - embeddings.append(chunk.embedding) - metadatas.append( - { - "path": file_meta.path, - "source": source.value, - "start_line": chunk.start_line, - "end_line": chunk.end_line, - "hash": chunk.hash, - "updated_at": now, - }, - ) - - # Batch upsert to ChromaDB (always pass embeddings to prevent default embedding function) - self.chunks_collection.upsert( - ids=ids, - documents=documents, - embeddings=embeddings, - metadatas=metadatas, - ) - - # Update file metadata in cache - if source.value not in self._metadata_cache: - self._metadata_cache[source.value] = {} - self._metadata_cache[source.value][file_meta.path] = FileMetadata( - hash=file_meta.hash, - mtime_ms=file_meta.mtime_ms, - size=file_meta.size, - path=file_meta.path, - chunk_count=len(chunks), - ) - - async def delete_file(self, path: str, source: MemorySource) -> None: - """Delete file and all its chunks.""" - # Query for all chunks with this path and source - results = self.chunks_collection.get( - where={"$and": [{"path": path}, {"source": source.value}]}, - include=[], - ) - - if results["ids"]: - self.chunks_collection.delete( - ids=results["ids"], - ) - - # Remove from file metadata cache - if source.value in self._metadata_cache: - self._metadata_cache[source.value].pop(path, None) - - async def delete_file_chunks(self, path: str, chunk_ids: list[str]) -> None: - """Delete specific chunks for a file.""" - if not chunk_ids: - return - - self.chunks_collection.delete( - ids=chunk_ids, - ) - - # Update chunk count in file metadata cache - for source_meta in self._metadata_cache.values(): - if path in source_meta: - # Recalculate chunk count - results = self.chunks_collection.get( - where={"path": path}, - include=[], - ) - source_meta[path].chunk_count = len(results["ids"]) - break - - async def upsert_chunks( - self, - chunks: list[MemoryChunk], - source: MemorySource, - ) -> None: - """Insert or update specific chunks without affecting other chunks.""" - if not chunks: - return - - # Batch generate embeddings for all chunks - # (base class returns mock embeddings when vector_enabled=False) - chunks = await self.get_chunk_embeddings(chunks) - - ids = [] - documents = [] - embeddings = [] - metadatas = [] - - now = int(time.time() * 1000) - for chunk in chunks: - ids.append(chunk.id) - documents.append(chunk.text) - embeddings.append(chunk.embedding) - metadatas.append( - { - "path": chunk.path, - "source": source.value, - "start_line": chunk.start_line, - "end_line": chunk.end_line, - "hash": chunk.hash, - "updated_at": now, - }, - ) - - # Always pass embeddings to prevent default embedding function - self.chunks_collection.upsert( - ids=ids, - documents=documents, - embeddings=embeddings, - metadatas=metadatas, - ) - - async def list_files(self, source: MemorySource) -> list[str]: - """List all indexed files for a source.""" - if source.value not in self._metadata_cache: - return [] - return list(self._metadata_cache[source.value].keys()) - - async def get_file_metadata( - self, - path: str, - source: MemorySource, - ) -> FileMetadata | None: - """Get file metadata with chunk count.""" - if source.value not in self._metadata_cache: - return None - return self._metadata_cache[source.value].get(path) - - async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None: - """Update file metadata without affecting chunks.""" - if source.value not in self._metadata_cache: - self._metadata_cache[source.value] = {} - - self._metadata_cache[source.value][file_meta.path] = FileMetadata( - hash=file_meta.hash, - mtime_ms=file_meta.mtime_ms, - size=file_meta.size, - path=file_meta.path, - chunk_count=file_meta.chunk_count, - ) - - async def get_file_chunks( - self, - path: str, - source: MemorySource, - ) -> list[MemoryChunk]: - """Get all chunks for a file.""" - results = self.chunks_collection.get( - where={"$and": [{"path": path}, {"source": source.value}]}, - include=["documents", "embeddings", "metadatas"], - ) - - chunks = [] - for i, chunk_id in enumerate(results["ids"]): - metadata = results["metadatas"][i] - chunks.append( - MemoryChunk( - id=chunk_id, - path=metadata["path"], - source=MemorySource(metadata["source"]), - start_line=metadata["start_line"], - end_line=metadata["end_line"], - text=results["documents"][i], - hash=metadata["hash"], - embedding=results["embeddings"][i] if results["embeddings"] is not None else None, - ), - ) - - # Sort by start_line - chunks.sort(key=lambda c: c.start_line) - return chunks - - async def vector_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform vector similarity search.""" - if not self.vector_enabled or not query: - return [] - - # Get query embedding - query_embedding = await self.get_embedding(query) - if not query_embedding: - return [] - - # Build where filter for sources - where_filter = None - if sources: - if len(sources) == 1: - where_filter = {"source": sources[0].value} - else: - where_filter = {"source": {"$in": [s.value for s in sources]}} - - # Perform vector search - try: - results = self.chunks_collection.query( - query_embeddings=[query_embedding], - n_results=limit, - where=where_filter, - include=["documents", "metadatas", "distances"], - ) - except Exception as e: - logger.error(f"Vector search failed: {e}, falling back to random results") - # Fallback: get some documents without vector search and assign random scores - try: - fallback_results = self.chunks_collection.get( - where=where_filter, - limit=limit, - include=["documents", "metadatas"], - ) - search_results = [] - if fallback_results["ids"]: - for i, _ in enumerate(fallback_results["ids"]): - metadata = fallback_results["metadatas"][i] - search_results.append( - MemorySearchResult( - path=metadata["path"], - start_line=metadata["start_line"], - end_line=metadata["end_line"], - score=random.uniform(0.3, 0.7), # Random score in middle range - snippet=fallback_results["documents"][i], - source=MemorySource(metadata["source"]), - raw_metric=None, - ), - ) - return search_results - except Exception as fallback_e: - logger.error(f"Fallback search also failed: {fallback_e}") - return [] - - search_results = [] - if results["ids"] and results["ids"][0]: - for i, _ in enumerate(results["ids"][0]): - metadata = results["metadatas"][0][i] - distance = results["distances"][0][i] - - # Convert cosine distance to similarity score - # Cosine distance range is [0, 2], convert to [1, 0] score - score = max(0.0, 1.0 - distance / 2.0) - - search_results.append( - MemorySearchResult( - path=metadata["path"], - start_line=metadata["start_line"], - end_line=metadata["end_line"], - score=score, - snippet=results["documents"][0][i], - source=MemorySource(metadata["source"]), - raw_metric=distance, - ), - ) - - # Sort by score descending - search_results.sort(key=lambda r: r.score, reverse=True) - return search_results - - async def keyword_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform keyword/full-text search. - - ChromaDB supports where_document filter for text matching. - Note: ChromaDB's $contains is case-sensitive, so we generate multiple - case variants (original, lowercase, capitalized) for each word to - improve recall while maintaining case-insensitive scoring. - """ - if not self.fts_enabled or not query: - return [] - - # Normalize whitespace and split into words - words = query.split() - if not words: - return [] - - # Generate case variants for each word to handle case-sensitive $contains - # Include: original, lowercase, and capitalized forms - word_variants = set() - for word in words: - word_variants.add(word) # original - word_variants.add(word.lower()) # lowercase - word_variants.add(word.capitalize()) # Capitalized - word_variants.add(word.upper()) # UPPERCASE - word_variants_list = list(word_variants) - - # Build where filter for sources - where_filter = None - if sources: - if len(sources) == 1: - where_filter = {"source": sources[0].value} - else: - where_filter = {"source": {"$in": [s.value for s in sources]}} - - # ChromaDB where_document uses $contains for substring matching (case-sensitive) - # Use multiple case variants to improve recall - if len(word_variants_list) == 1: - where_document: dict = {"$contains": word_variants_list[0]} - else: - where_document = {"$or": [{"$contains": w} for w in word_variants_list]} - - # Get all matching documents - results = self.chunks_collection.get( - where=where_filter, - where_document=where_document, - include=["documents", "metadatas"], - ) - - search_results = [] - query_lower = query.lower() - words_lower = [w.lower() for w in words] # lowercase words for scoring - n_words = len(words) - - for i, _ in enumerate(results["ids"]): - metadata = results["metadatas"][i] - text = results["documents"][i] - text_lower = text.lower() - - # Calculate relevance score based on word matches - match_count = sum(1 for w in words_lower if w in text_lower) - base_score = match_count / n_words - - # Bonus for full phrase match (only applies to multi-word queries) - phrase_bonus = 0.2 if n_words > 1 and query_lower in text_lower else 0.0 - # Scale base_score and add phrase bonus, max score is 1.0 - score = min(1.0, base_score + phrase_bonus) - - search_results.append( - MemorySearchResult( - path=metadata["path"], - start_line=metadata["start_line"], - end_line=metadata["end_line"], - score=score, - snippet=text, - source=MemorySource(metadata["source"]), - ), - ) - - # Sort by score descending and limit results - search_results.sort(key=lambda r: r.score, reverse=True) - return search_results[:limit] - - async def hybrid_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - vector_weight: float = 0.7, - candidate_multiplier: float = 3.0, - ) -> list[MemorySearchResult]: - """Perform hybrid search combining vector and keyword search. - - Args: - query: Search query text - limit: Maximum number of results - sources: Optional list of sources to filter - vector_weight: Weight for vector search results (0.0-1.0). - Keyword weight = 1.0 - vector_weight. - candidate_multiplier: Multiplier for candidate pool size. - - Returns: - List of search results sorted by combined relevance score - """ - assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be between 0 and 1, got {vector_weight}" - - candidates = min(200, max(1, int(limit * candidate_multiplier))) - text_weight = 1.0 - vector_weight - - # Perform search based on enabled backends - if self.vector_enabled and self.fts_enabled: - keyword_results = await self.keyword_search(query, candidates, sources) - vector_results = await self.vector_search(query, candidates, sources) - - # Log original vector results - logger.info("\n=== Vector Search Results ===") - for i, r in enumerate(vector_results[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - # Log original keyword results - logger.info("\n=== Keyword Search Results ===") - for i, r in enumerate(keyword_results[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - if not keyword_results: - return vector_results[:limit] - elif not vector_results: - return keyword_results[:limit] - else: - merged = self._merge_hybrid_results( - vector=vector_results, - keyword=keyword_results, - vector_weight=vector_weight, - text_weight=text_weight, - ) - - # Log merged results - logger.info("\n=== Merged Hybrid Results ===") - for i, r in enumerate(merged[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - return merged[:limit] - elif self.vector_enabled: - vector_results = await self.vector_search(query, limit, sources) - return vector_results - elif self.fts_enabled: - keyword_results = await self.keyword_search(query, limit, sources) - return keyword_results - else: - return [] - - @staticmethod - def _merge_hybrid_results( - vector: list[MemorySearchResult], - keyword: list[MemorySearchResult], - vector_weight: float, - text_weight: float, - ) -> list[MemorySearchResult]: - """Merge vector and keyword search results with weighted scoring.""" - merged: dict[str, MemorySearchResult] = {} - - # Process vector results - for result in vector: - result.score = result.score * vector_weight - merged[result.merge_key] = result - - # Process keyword results - for result in keyword: - key = result.merge_key - if key in merged: - merged[key].score += result.score * text_weight - else: - result.score = result.score * text_weight - merged[key] = result - - # Sort by score and return - results = list(merged.values()) - results.sort(key=lambda r: r.score, reverse=True) - return results - - async def clear_all(self) -> None: - """Clear all indexed data.""" - # Delete and recreate the collection - self.client.delete_collection( - name=self.collection_name, - ) - self.chunks_collection = self.client.get_or_create_collection( - name=self.collection_name, - metadata={"hnsw:space": "cosine"}, - ) - - # Clear file metadata cache and disk - self._metadata_cache = {} - await self._save_metadata({}) - - logger.info(f"Cleared all data from ChromaDB collection: {self.collection_name}") - - async def close(self) -> None: - """Close ChromaDB client and release resources.""" - # Persist metadata cache to disk before closing - if self._metadata_cache: - await self._save_metadata(self._metadata_cache) - - # ChromaDB PersistentClient handles persistence automatically - self.client = None - self.chunks_collection = None - await super().close() diff --git a/reme_cli/component/file_store/local_file_store.py b/reme_cli/component/file_store/local_file_store.py deleted file mode 100644 index 38757d7a..00000000 --- a/reme_cli/component/file_store/local_file_store.py +++ /dev/null @@ -1,461 +0,0 @@ -"""Pure-Python in-memory storage backend for file store, with JSON file persistence.""" - -import json -from pathlib import Path - -import numpy as np -from loguru import logger - -from .base_file_store import BaseFileStore -from ..enumeration import MemorySource -from ..schema import FileMetadata, MemoryChunk, MemorySearchResult -from ..utils.common_utils import batch_cosine_similarity - - -class LocalFileStore(BaseFileStore): - """Pure-Python in-memory file storage with JSONL file persistence. - - No external dependencies required. All data lives in Python dicts; - writes are persisted to JSONL files on disk so state survives restarts. - - Inherits embedding methods from BaseFileStore: - - get_chunk_embedding / get_chunk_embeddings (async) - - get_embedding / get_embeddings (async) - - Provides: - - Vector similarity search (cosine similarity, pure Python) - - Full-text / keyword search (Python substring matching) - - Efficient chunk and file metadata management - """ - - def __init__(self, **kwargs): - super().__init__(**kwargs) - self._started: bool = False - # In-memory indexes - self._chunks: dict[str, MemoryChunk] = {} - self._files: dict[str, dict[str, FileMetadata]] = {} # source -> path -> meta - # Persistence paths (mirror ChromaFileStore convention) - self._chunks_file: Path = self.db_path / f"{self.store_name}_chunks.jsonl" - self._metadata_file: Path = self.db_path / f"{self.store_name}_file_metadata.json" - - # ------------------------------------------------------------------ - # Persistence helpers - # ------------------------------------------------------------------ - - async def _load_chunks(self) -> None: - """Load chunks from JSONL file into memory.""" - if not self._chunks_file.exists(): - return - try: - data = self._chunks_file.read_text(encoding="utf-8") - self._chunks = {} - for line in data.strip().split("\n"): - if not line: - continue - rec = json.loads(line) - chunk = MemoryChunk.model_validate(rec) - self._chunks[chunk.id] = chunk - logger.debug(f"Loaded {len(self._chunks)} chunks from {self._chunks_file}") - except Exception as e: - logger.warning(f"Failed to load chunks from {self._chunks_file}: {e}") - - async def _save_chunks(self) -> None: - """Persist chunks to JSONL file.""" - try: - lines = [] - for chunk in self._chunks.values(): - chunk_dict = chunk.model_dump(mode="json") - lines.append(json.dumps(chunk_dict, ensure_ascii=False)) - data = "\n".join(lines) - self._chunks_file.write_text(data, encoding="utf-8") - logger.debug(f"Saved {len(self._chunks)} chunks to {self._chunks_file}") - except Exception as e: - logger.error(f"Failed to save chunks to {self._chunks_file}: {e}") - - async def _load_metadata(self) -> None: - """Load file metadata from JSON file into memory.""" - if not self._metadata_file.exists(): - return - try: - data = self._metadata_file.read_text(encoding="utf-8") - raw: dict = json.loads(data) - self._files = { - source: {path: FileMetadata(**meta) for path, meta in files.items()} for source, files in raw.items() - } - logger.debug(f"Loaded file metadata from {self._metadata_file}") - except Exception as e: - logger.warning(f"Failed to load file metadata from {self._metadata_file}: {e}") - - async def _save_metadata(self) -> None: - """Persist file metadata to JSON file.""" - try: - raw: dict = {} - for source, files in self._files.items(): - raw[source] = { - path: { - "path": meta.path, - "hash": meta.hash, - "mtime_ms": meta.mtime_ms, - "size": meta.size, - "chunk_count": meta.chunk_count, - } - for path, meta in files.items() - } - data = json.dumps(raw, indent=2, ensure_ascii=False) - self._metadata_file.write_text(data, encoding="utf-8") - logger.debug(f"Saved file metadata to {self._metadata_file}") - except Exception as e: - logger.error(f"Failed to save file metadata to {self._metadata_file}: {e}") - - # ------------------------------------------------------------------ - # Lifecycle - # ------------------------------------------------------------------ - - async def start(self) -> None: - """Load persisted data into memory.""" - if self._started: - return - self._started = True - await self._load_metadata() - await self._load_chunks() - logger.info( - f"LocalFileStore '{self.store_name}' ready: " - f"{len(self._chunks)} chunks, metadata at {self._metadata_file}", - ) - - async def close(self) -> None: - """Flush state to disk and release memory.""" - await self._save_metadata() - await self._save_chunks() - self._chunks.clear() - self._files.clear() - self._started = False - - # ------------------------------------------------------------------ - # Write operations - # ------------------------------------------------------------------ - - async def upsert_file( - self, - file_meta: FileMetadata, - source: MemorySource, - chunks: list[MemoryChunk], - ) -> None: - """Insert or update file and its chunks.""" - if not chunks: - return - - # Remove existing chunks for this file/source first - await self.delete_file(file_meta.path, source) - - # Batch generate embeddings (base class returns mock embeddings when vector_enabled=False) - chunks = await self.get_chunk_embeddings(chunks) - - for chunk in chunks: - self._chunks[chunk.id] = chunk - - if source.value not in self._files: - self._files[source.value] = {} - self._files[source.value][file_meta.path] = FileMetadata( - hash=file_meta.hash, - mtime_ms=file_meta.mtime_ms, - size=file_meta.size, - path=file_meta.path, - chunk_count=len(chunks), - ) - - async def delete_file(self, path: str, source: MemorySource) -> None: - """Delete file and all its chunks.""" - to_delete = [cid for cid, chunk in self._chunks.items() if chunk.path == path and chunk.source == source] - for cid in to_delete: - del self._chunks[cid] - - if source.value in self._files: - self._files[source.value].pop(path, None) - - async def delete_file_chunks(self, path: str, chunk_ids: list[str]) -> None: - """Delete specific chunks for a file.""" - if not chunk_ids: - return - - for cid in chunk_ids: - self._chunks.pop(cid, None) - - # Recalculate chunk_count in file metadata (per source) - for source_key, source_meta in self._files.items(): - if path in source_meta: - source_meta[path].chunk_count = sum( - 1 for chunk in self._chunks.values() if chunk.path == path and chunk.source.value == source_key - ) - - async def upsert_chunks( - self, - chunks: list[MemoryChunk], - source: MemorySource, - ) -> None: - """Insert or update specific chunks without affecting other chunks.""" - if not chunks: - return - - chunks = await self.get_chunk_embeddings(chunks) - - for chunk in chunks: - self._chunks[chunk.id] = chunk - - # ------------------------------------------------------------------ - # Read operations - # ------------------------------------------------------------------ - - async def list_files(self, source: MemorySource) -> list[str]: - """List all indexed files for a source.""" - return list(self._files.get(source.value, {}).keys()) - - async def get_file_metadata( - self, - path: str, - source: MemorySource, - ) -> FileMetadata | None: - """Get file metadata.""" - return self._files.get(source.value, {}).get(path) - - async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None: - """Update file metadata without affecting chunks.""" - if source.value not in self._files: - self._files[source.value] = {} - - self._files[source.value][file_meta.path] = FileMetadata( - hash=file_meta.hash, - mtime_ms=file_meta.mtime_ms, - size=file_meta.size, - path=file_meta.path, - chunk_count=file_meta.chunk_count, - ) - - async def get_file_chunks( - self, - path: str, - source: MemorySource, - ) -> list[MemoryChunk]: - """Get all chunks for a file, sorted by start_line.""" - chunks = [chunk for chunk in self._chunks.values() if chunk.path == path and chunk.source == source] - chunks.sort(key=lambda c: c.start_line) - return chunks - - # ------------------------------------------------------------------ - # Search - # ------------------------------------------------------------------ - - async def vector_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform cosine-similarity vector search over in-memory embeddings.""" - if not self.vector_enabled or not query: - return [] - - query_embedding = await self.get_embedding(query) - if not query_embedding: - return [] - - expected_dim = self.embedding_dim - - # Collect candidate chunks with embeddings - candidates = [ - chunk for chunk in self._chunks.values() if (not sources or chunk.source in sources) and chunk.embedding - ] - - if not candidates: - return [] - - # Validate and fix chunk embedding dimensions - valid_embeddings = [] - for chunk in candidates: - emb = chunk.embedding - emb_len = len(emb) - if emb_len != expected_dim: - if emb_len < expected_dim: - emb = emb + [0.0] * (expected_dim - emb_len) - logger.warning( - f"Chunk embedding dimension {emb_len} < expected {expected_dim}, " - f"padded with zeros (chunk_id={chunk.id})", - ) - else: - emb = emb[:expected_dim] - logger.warning( - f"Chunk embedding dimension {emb_len} > expected {expected_dim}, " - f"truncated to {expected_dim} (chunk_id={chunk.id})", - ) - valid_embeddings.append(emb) - - # Build embedding matrix and compute similarities in batch - query_array = np.array([query_embedding]) # Shape: (1, emb_size) - chunk_embeddings = np.array(valid_embeddings) # Shape: (n, emb_size) - similarities = batch_cosine_similarity(query_array, chunk_embeddings)[0] # Shape: (n,) - - # Build results - results = [ - MemorySearchResult( - path=chunk.path, - start_line=chunk.start_line, - end_line=chunk.end_line, - score=float(similarity), - snippet=chunk.text, - source=chunk.source, - raw_metric=1.0 - float(similarity), - ) - for chunk, similarity in zip(candidates, similarities) - ] - - results.sort(key=lambda r: r.score, reverse=True) - return results[:limit] - - async def keyword_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform keyword/full-text search via Python substring matching.""" - if not self.fts_enabled or not query: - return [] - - words = query.split() - if not words: - return [] - - query_lower = query.lower() - words_lower = [w.lower() for w in words] - n_words = len(words) - - results = [] - for chunk in self._chunks.values(): - if sources and chunk.source not in sources: - continue - - text_lower = chunk.text.lower() - match_count = sum(1 for w in words_lower if w in text_lower) - if match_count == 0: - continue - - base_score = match_count / n_words - # Bonus for full phrase match (multi-word queries only) - phrase_bonus = 0.2 if n_words > 1 and query_lower in text_lower else 0.0 - score = min(1.0, base_score + phrase_bonus) - - results.append( - MemorySearchResult( - path=chunk.path, - start_line=chunk.start_line, - end_line=chunk.end_line, - score=score, - snippet=chunk.text, - source=chunk.source, - ), - ) - - results.sort(key=lambda r: r.score, reverse=True) - return results[:limit] - - async def hybrid_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - vector_weight: float = 0.7, - candidate_multiplier: float = 3.0, - ) -> list[MemorySearchResult]: - """Perform hybrid search combining vector and keyword search. - - Args: - query: Search query text - limit: Maximum number of results - sources: Optional list of sources to filter - vector_weight: Weight for vector search results (0.0-1.0). - Keyword weight = 1.0 - vector_weight. - candidate_multiplier: Multiplier for candidate pool size. - - Returns: - List of search results sorted by combined relevance score - """ - assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be between 0 and 1, got {vector_weight}" - - candidates = min(200, max(1, int(limit * candidate_multiplier))) - text_weight = 1.0 - vector_weight - - if self.vector_enabled and self.fts_enabled: - keyword_results = await self.keyword_search(query, candidates, sources) - vector_results = await self.vector_search(query, candidates, sources) - - logger.info("\n=== Vector Search Results ===") - for i, r in enumerate(vector_results[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - logger.info("\n=== Keyword Search Results ===") - for i, r in enumerate(keyword_results[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - if not keyword_results: - return vector_results[:limit] - elif not vector_results: - return keyword_results[:limit] - else: - merged = self._merge_hybrid_results( - vector=vector_results, - keyword=keyword_results, - vector_weight=vector_weight, - text_weight=text_weight, - ) - - logger.info("\n=== Merged Hybrid Results ===") - for i, r in enumerate(merged[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - return merged[:limit] - elif self.vector_enabled: - return await self.vector_search(query, limit, sources) - elif self.fts_enabled: - return await self.keyword_search(query, limit, sources) - else: - return [] - - @staticmethod - def _merge_hybrid_results( - vector: list[MemorySearchResult], - keyword: list[MemorySearchResult], - vector_weight: float, - text_weight: float, - ) -> list[MemorySearchResult]: - """Merge vector and keyword search results with weighted scoring.""" - merged: dict[str, MemorySearchResult] = {} - - for result in vector: - result.metadata["_weighted_score"] = result.score * vector_weight - merged[result.merge_key] = result - - for result in keyword: - key = result.merge_key - if key in merged: - merged[key].metadata["_weighted_score"] += result.score * text_weight - else: - result.metadata["_weighted_score"] = result.score * text_weight - merged[key] = result - - results = list(merged.values()) - for r in results: - r.score = r.metadata.pop("_weighted_score") - - results.sort(key=lambda r: r.score, reverse=True) - return results - - async def clear_all(self) -> None: - """Clear all indexed data from memory and disk.""" - self._chunks.clear() - self._files.clear() - await self._save_chunks() - await self._save_metadata() - logger.info(f"Cleared all data from LocalFileStore '{self.store_name}'") diff --git a/reme_cli/component/file_store/sqlite_file_store.py b/reme_cli/component/file_store/sqlite_file_store.py deleted file mode 100644 index 0a494d2c..00000000 --- a/reme_cli/component/file_store/sqlite_file_store.py +++ /dev/null @@ -1,978 +0,0 @@ -"""SQLite storage backend for file store.""" - -import json - -import struct -import time - -from loguru import logger - -from .base_file_store import BaseFileStore -from ..enumeration import MemorySource -from ..schema import FileMetadata, MemoryChunk, MemorySearchResult - - -class SqliteFileStore(BaseFileStore): - """SQLite file storage with vector and full-text search. - - Inherits embedding methods from BaseFileStore: - - get_chunk_embedding / get_chunk_embeddings (async) - - get_chunk_embedding_sync / get_chunk_embeddings_sync (sync) - - get_embedding / get_embeddings (async) - - Provides SQLite-backed persistent storage with: - - Vector similarity search (via sqlite-vec extension) - - Full-text search (via FTS5) - - Efficient chunk and file metadata management - """ - - def __init__(self, vec_ext_path: str = "", **kwargs): - super().__init__(**kwargs) - self.vec_ext_path = vec_ext_path - import sqlite3 - - self.conn: sqlite3.Connection | None = None - - @property - def vector_table_name(self) -> str: - """Get the name of the vector table for this store.""" - return f"chunks_vec_{self.store_name}" - - @property - def fts_table_name(self) -> str: - """Get the name of the FTS table for this store.""" - return f"chunks_fts_{self.store_name}" - - @property - def chunks_table_name(self) -> str: - """Get the name of the chunks table for this store.""" - return f"chunks_{self.store_name}" - - @property - def files_table_name(self) -> str: - """Get the name of the files table for this store.""" - return f"files_{self.store_name}" - - @staticmethod - def vector_to_blob(embedding: list[float]) -> bytes: - """Convert vector to binary blob for sqlite-vec.""" - return struct.pack(f"{len(embedding)}f", *embedding) - - async def start(self) -> None: - """Initialize database and load extensions.""" - if self.conn is not None: - return - import sqlite3 - - self.conn = sqlite3.connect(self.db_path / "reme.db", check_same_thread=False) - - # Only load sqlite-vec extension if vector search is enabled - if self.vector_enabled: - logger.warning( - "On macOS systems with version 14 or earlier, " - "loading the sqlite-vec vector extension carries a risk of crashes or hangs.", - ) - - 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: - logger.info("Vector search disabled, skipping sqlite-vec extension loading") - - await self._create_tables() - - async def _create_tables(self) -> None: - """Create database schema.""" - cursor = self.conn.cursor() - try: - # 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_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}] - ) - """, - ) - 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' - ) - """, - ) - 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.""" - cursor = self.conn.cursor() - - try: - cursor.execute("BEGIN") - - # Insert file - cursor.execute( - f""" - INSERT OR REPLACE INTO {self.files_table_name} (path, source, hash, mtime, size) - VALUES (?, ?, ?, ?, ?) - """, - (file_meta.path, source.value, file_meta.hash, file_meta.mtime_ms, file_meta.size), - ) - - # Insert chunks - now = int(time.time() * 1000) - for chunk in chunks: - cursor.execute( - f""" - INSERT OR REPLACE INTO {self.chunks_table_name} ( - id, path, source, start_line, end_line, - hash, text, embedding, updated_at - ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) - """, - ( - chunk.id, - file_meta.path, - source.value, - chunk.start_line, - chunk.end_line, - chunk.hash, - chunk.text, - json.dumps(chunk.embedding) if chunk.embedding else None, - now, - ), - ) - - # Insert vector (vec0 doesn't support OR REPLACE, use DELETE + INSERT) - 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_enabled: - cursor.execute( - f""" - INSERT OR REPLACE INTO {self.fts_table_name} ( - text, id, path, source, start_line, end_line - ) VALUES (?, ?, ?, ?, ?, ?) - """, - ( - chunk.text, - chunk.id, - file_meta.path, - source.value, - chunk.start_line, - chunk.end_line, - ), - ) - - cursor.execute("COMMIT") - except Exception as e: - cursor.execute("ROLLBACK") - logger.error(f"Failed to upsert file {file_meta.path}: {e}") - raise - finally: - cursor.close() - - async def delete_file(self, path: str, source: MemorySource): - """Delete file and all its chunks.""" - cursor = self.conn.cursor() - try: - cursor.execute("BEGIN") - - # Get chunk IDs for vector deletion - cursor.execute( - f"SELECT id FROM {self.chunks_table_name} WHERE path = ? AND source = ?", - (path, source.value), - ) - chunk_ids = [row[0] for row in cursor.fetchall()] - - # Delete vectors - if self.vector_enabled and chunk_ids: - for chunk_id in chunk_ids: - cursor.execute( - f"DELETE FROM {self.vector_table_name} WHERE id = ?", - (chunk_id,), - ) - - # Delete FTS entries - 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( - f"DELETE FROM {self.chunks_table_name} WHERE path = ? AND source = ?", - (path, source.value), - ) - cursor.execute( - f"DELETE FROM {self.files_table_name} WHERE path = ? AND source = ?", - (path, source.value), - ) - - cursor.execute("COMMIT") - except Exception as e: - cursor.execute("ROLLBACK") - logger.error(f"Failed to delete file {path}: {e}") - raise - finally: - cursor.close() - - async def delete_file_chunks(self, path: str, chunk_ids: list[str]): - """Delete specific chunks for a file.""" - if not chunk_ids: - return - - cursor = self.conn.cursor() - try: - cursor.execute("BEGIN") - - # Delete vectors - if self.vector_enabled: - for chunk_id in chunk_ids: - cursor.execute( - f"DELETE FROM {self.vector_table_name} WHERE id = ?", - (chunk_id,), - ) - - # Delete FTS entries - if self.fts_enabled: - placeholders = ",".join("?" * len(chunk_ids)) - cursor.execute( - f"DELETE FROM {self.fts_table_name} WHERE id IN ({placeholders})", - chunk_ids, - ) - - # Delete chunks - placeholders = ",".join("?" * len(chunk_ids)) - cursor.execute( - f"DELETE FROM {self.chunks_table_name} WHERE id IN ({placeholders})", - chunk_ids, - ) - - cursor.execute("COMMIT") - except Exception as e: - cursor.execute("ROLLBACK") - logger.error(f"Failed to delete chunks for {path}: {e}") - raise - finally: - cursor.close() - - async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource): - """Insert or update specific chunks without affecting other chunks.""" - if not chunks: - return - - cursor = self.conn.cursor() - try: - cursor.execute("BEGIN") - - now = int(time.time() * 1000) - for chunk in chunks: - # Insert/update chunk - cursor.execute( - f""" - INSERT OR REPLACE INTO {self.chunks_table_name} ( - id, path, source, start_line, end_line, - hash, text, embedding, updated_at - ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) - """, - ( - chunk.id, - chunk.path, - source.value, - chunk.start_line, - chunk.end_line, - chunk.hash, - chunk.text, - json.dumps(chunk.embedding) if chunk.embedding else None, - now, - ), - ) - - # Insert/update vector (vec0 doesn't support OR REPLACE, use DELETE + INSERT) - 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_enabled: - cursor.execute( - f""" - INSERT OR REPLACE INTO {self.fts_table_name} ( - text, id, path, source, start_line, end_line - ) VALUES (?, ?, ?, ?, ?, ?) - """, - ( - chunk.text, - chunk.id, - chunk.path, - source.value, - chunk.start_line, - chunk.end_line, - ), - ) - - cursor.execute("COMMIT") - except Exception as e: - cursor.execute("ROLLBACK") - logger.error(f"Failed to upsert chunks: {e}") - raise - finally: - cursor.close() - - async def list_files(self, source: MemorySource) -> list[str]: - """List all indexed files.""" - cursor = self.conn.cursor() - 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() - 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() - - async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None: - """Update file metadata without affecting chunks.""" - cursor = self.conn.cursor() - try: - cursor.execute( - f""" - INSERT OR REPLACE INTO {self.files_table_name} (path, source, hash, mtime, size) - VALUES (?, ?, ?, ?, ?) - """, - (file_meta.path, source.value, file_meta.hash, file_meta.mtime_ms, file_meta.size), - ) - self.conn.commit() - except Exception as e: - logger.error(f"Failed to update file metadata for {file_meta.path}: {e}") - raise - finally: - cursor.close() - - async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]: - """Get all chunks for a file.""" - cursor = self.conn.cursor() - 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), - ) - - 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, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform vector similarity search.""" - if not self.vector_enabled or not query: - return [] - - # Get query embedding - query_embedding = await self.get_embedding(query) - if not query_embedding: - return [] - - cursor = self.conn.cursor() - source_filter = "" - params: list = [] - if sources: - placeholders = ",".join("?" * len(sources)) - source_filter = f" AND c.source IN ({placeholders})" - params = [s.value for s in sources] - - try: - query_blob = self.vector_to_blob(query_embedding) - - # Correct SQLite-vec syntax for vector search with limit - # vec0 requires 'k = ?' constraint for knn queries - query_sql = f""" - SELECT c.id, c.path, c.start_line, c.end_line, c.source, c.text, v.distance - FROM {self.vector_table_name} v - JOIN {self.chunks_table_name} c ON v.id = c.id - WHERE v.embedding MATCH ? - AND k = ? - """ - query_params: list = [query_blob, limit] - - # Add source filter if specified - if source_filter: - query_sql += source_filter - query_params.extend(params) - - # Order by distance (k constraint already limits results) - query_sql += " ORDER BY v.distance" - - cursor.execute(query_sql, query_params) - - results = [] - for _, path, start, end, src, text, dist in cursor.fetchall(): - # Convert L2 distance to similarity score - # For normalized vectors, L2 distance range is [0, 2] - # Map to [1, 0] score range (higher score = more similar) - score = max(0.0, 1.0 - dist / 2.0) - snippet = text - results.append( - MemorySearchResult( - path=path, - start_line=start, - end_line=end, - score=score, - snippet=snippet, - source=MemorySource(src), - raw_metric=dist, - ), - ) - - results.sort(key=lambda r: r.score, reverse=True) - return results - except Exception as e: - logger.error(f"Vector search failed: {e}") - return [] - finally: - cursor.close() - - @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: - - * (prefix match) - - ? (not used in FTS5, but can cause issues) - - " (phrase search, needs escaping) - - : (column filter) - - ^ (start of line anchor, not standard FTS5) - - ' (single quote, causes syntax errors) - - ` (backtick, can cause issues) - - | (pipe, OR operator) - - + (plus, can be used for required terms) - - - (minus, NOT operator) - - = (equals, can cause issues) - - < > (angle brackets, comparison operators) - - ! (exclamation, NOT operator variant) - - @ # $ % & (other special chars) - - "\" - - / (slash, can interfere) - - ; (semicolon, statement separator) - - , (comma, can interfere with phrase parsing) - - Args: - query: Raw query string - - Returns: - Sanitized query string safe for FTS5 - """ - if not query: - return "" - - # Remove FTS5 special characters that we don't want users to use - # Keep only alphanumeric, spaces, periods, and underscores - special_chars = [ - "*", - "?", - ":", - "^", - "(", - ")", - "[", - "]", - "{", - "}", - "'", - '"', - "`", - "|", - "+", - "-", - "=", - "<", - ">", - "!", - "@", - "#", - "$", - "%", - "&", - "\\", - "/", - ";", - ",", - ] - cleaned = query - for char in special_chars: - cleaned = cleaned.replace(char, " ") - - # Normalize whitespace - cleaned = " ".join(cleaned.split()) - - return cleaned - - async def keyword_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """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 [] - - cleaned = self._sanitize_fts_query(query) - if not cleaned: - return [] - - words = cleaned.split() - if not words: - return [] - - # 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() - source_filter = "" - params: list = [fts_query] - if sources: - placeholders = ",".join("?" * len(sources)) - source_filter = f" AND fts.source IN ({placeholders})" - params.extend([s.value for s in sources]) - params.append(limit) - - try: - cursor.execute( - f""" - SELECT fts.id, fts.path, fts.start_line, fts.end_line, - fts.source, fts.text, rank - FROM {self.fts_table_name} fts - WHERE fts.text MATCH ?{source_filter} - ORDER BY rank - LIMIT ? - """, - params, - ) - - results = [] - for _, path, start, end, src, text, rank in cursor.fetchall(): - score = max(0.0, 1.0 / (1.0 + abs(rank))) - results.append( - MemorySearchResult( - path=path, - start_line=start, - end_line=end, - score=score, - 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"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() - - async def hybrid_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - vector_weight: float = 0.7, - candidate_multiplier: float = 3.0, - ) -> list[MemorySearchResult]: - """Perform hybrid search combining vector and keyword search. - - Args: - query: Search query text - limit: Maximum number of results - sources: Optional list of sources to filter - vector_weight: Weight for vector search results (0.0-1.0). - Keyword weight = 1.0 - vector_weight. - candidate_multiplier: Multiplier for candidate pool size. - - Returns: - List of search results sorted by combined relevance score - """ - assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be between 0 and 1, got {vector_weight}" - - candidates = min(200, max(1, int(limit * candidate_multiplier))) - text_weight = 1.0 - vector_weight - - # Perform search based on enabled backends - if self.vector_enabled and self.fts_enabled: - keyword_results = await self.keyword_search(query, candidates, sources) - vector_results = await self.vector_search(query, candidates, sources) - - # Log original vector results - logger.info("\n=== Vector Search Results ===") - for i, r in enumerate(vector_results[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - # Log original keyword results - logger.info("\n=== Keyword Search Results ===") - for i, r in enumerate(keyword_results[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - if not keyword_results: - return vector_results[:limit] - elif not vector_results: - return keyword_results[:limit] - else: - merged = self._merge_hybrid_results( - vector=vector_results, - keyword=keyword_results, - vector_weight=vector_weight, - text_weight=text_weight, - ) - - # Log merged results - logger.info("\n=== Merged Hybrid Results ===") - for i, r in enumerate(merged[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - return merged[:limit] - elif self.vector_enabled: - vector_results = await self.vector_search(query, limit, sources) - return vector_results - elif self.fts_enabled: - keyword_results = await self.keyword_search(query, limit, sources) - return keyword_results - else: - return [] - - @staticmethod - def _merge_hybrid_results( - vector: list[MemorySearchResult], - keyword: list[MemorySearchResult], - vector_weight: float, - text_weight: float, - ) -> list[MemorySearchResult]: - """Merge vector and keyword search results with weighted scoring.""" - merged: dict[str, MemorySearchResult] = {} - - # Process vector results - for result in vector: - result.score = result.score * vector_weight - merged[result.merge_key] = result - - # Process keyword results - for result in keyword: - key = result.merge_key - if key in merged: - merged[key].score += result.score * text_weight - else: - result.score = result.score * text_weight - merged[key] = result - - # Sort by score and return - results = list(merged.values()) - results.sort(key=lambda r: r.score, reverse=True) - return results - - async def clear_all(self): - """Clear all indexed data.""" - cursor = self.conn.cursor() - 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_enabled: - cursor.execute(f"DELETE FROM {self.vector_table_name}") - - if self.fts_enabled: - cursor.execute(f"DELETE FROM {self.fts_table_name}") - - cursor.execute("COMMIT") - except Exception as e: - cursor.execute("ROLLBACK") - logger.error(f"Failed to clear all data: {e}") - raise - finally: - cursor.close() - - async def close(self): - """Close database connection.""" - if self.conn: - self.conn.close() - self.conn = None - await super().close() diff --git a/reme_cli/component/file_watcher/__init__.py b/reme_cli/component/file_watcher/__init__.py deleted file mode 100644 index 020d0cfa..00000000 --- a/reme_cli/component/file_watcher/__init__.py +++ /dev/null @@ -1,19 +0,0 @@ -"""File watcher module for monitoring file system changes. - -This module provides file watcher implementations for monitoring file changes -and updating memory stores accordingly. -""" - -from .base_file_watcher import BaseFileWatcher -from .delta_file_watcher import DeltaFileWatcher -from .full_file_watcher import FullFileWatcher -from ..registry_factory import R - -__all__ = [ - "BaseFileWatcher", - "DeltaFileWatcher", - "FullFileWatcher", -] - -R.file_watchers.register("full")(FullFileWatcher) -R.file_watchers.register("delta")(DeltaFileWatcher) diff --git a/reme_cli/component/file_watcher/base_file_watcher.py b/reme_cli/component/file_watcher/base_file_watcher.py deleted file mode 100644 index b4644dd7..00000000 --- a/reme_cli/component/file_watcher/base_file_watcher.py +++ /dev/null @@ -1,240 +0,0 @@ -"""Base file watcher implementation. - -This module provides the base class for file watcher implementations -that monitor file system changes and trigger callbacks. -""" - -import asyncio -from collections.abc import Coroutine -from pathlib import Path -from typing import Any, Callable - -from loguru import logger -from watchfiles import awatch, Change - -from ..enumeration import MemorySource -from ..file_store import BaseFileStore - - -class BaseFileWatcher: - """ - Minimal file watcher base class - - This base class provides basic file monitoring functionality that can be extended - to implement specific file monitoring requirements. - """ - - def __init__( - self, - watch_paths: list[str] | str, - suffix_filters: list[str] | None = None, - recursive: bool = False, - debounce: int = 2000, - chunk_tokens: int = 400, - chunk_overlap: int = 80, - file_store: BaseFileStore | None = None, - callback: Callable[[set[tuple[Change, str]]], None | Coroutine[Any, Any, None]] | None = None, - rebuild_index_on_start: bool = True, - poll_delay_ms: int = 2000, - **kwargs, - ): - """ - Initialize the file watcher - - Args: - watch_paths: Paths to watch for changes - suffix_filters: File suffix filters (e.g., ['.py', '.txt']) - recursive: Whether to watch directories recursively - debounce: Debounce time in milliseconds - chunk_tokens: Token size for chunking - chunk_overlap: Overlap size for chunks - file_store: File store instance - callback: Callback function for changes - rebuild_index_on_start: If True, clear all indexed data on start and rescan existing files. - If False, only monitor new changes without initialization. - poll_delay_ms: Polling delay in milliseconds. If > 300ms, force_polling will be enabled automatically. - **kwargs: Additional keyword arguments - """ - self.watch_paths: list[str] = [watch_paths] if isinstance(watch_paths, str) else watch_paths - self.suffix_filters: list[str] = suffix_filters or [] - self.recursive: bool = recursive - self.debounce: int = debounce - self.chunk_tokens: int = chunk_tokens - self.chunk_overlap: int = chunk_overlap - self.file_store: BaseFileStore = file_store - self.callback = callback - self.rebuild_index_on_start: bool = rebuild_index_on_start - self.poll_delay_ms: int = poll_delay_ms - self.kwargs: dict = kwargs - - self._stop_event = asyncio.Event() - self._watch_task: asyncio.Task | None = None - self._running = False - - async def start(self): - """Start the file watcher""" - if self._running: - return - - self._running = True - - async def _initialize_and_watch(): - if self.rebuild_index_on_start: - await self.file_store.clear_all() - logger.info("Cleared all indexed data on start") - await self._scan_existing_files() - await self._watch_loop() - - self._watch_task = asyncio.create_task(_initialize_and_watch()) - logger.info(f"Started watching: {self.watch_paths}") - - async def close(self): - """Stop the file watcher""" - if not self._running: - return - - self._stop_event.set() - if self._watch_task: - await self._watch_task - self._running = False - logger.info("Stopped watching") - - def watch_filter(self, _change: Change, path: str) -> bool: - """Filter function for file watching.""" - # If no suffix filters are specified, watch all files - if not self.suffix_filters: - return True - - # Check if the file has one of the allowed suffixes - for suffix in self.suffix_filters: - if path.endswith("." + suffix.strip(".")): - return True - - return False - - async def _scan_existing_files(self): - """Scan existing files matching watch criteria and trigger on_changes with Change.added""" - existing_files: set[tuple[Change, str]] = set() - - for watch_path_str in self.watch_paths: - watch_path = Path(watch_path_str) - - if not watch_path.exists(): - logger.warning(f"Watch path does not exist: {watch_path}") - continue - - if watch_path.is_file(): - # Single file - if self.watch_filter(Change.added, str(watch_path)): - existing_files.add((Change.added, str(watch_path))) - elif watch_path.is_dir(): - # Directory - if self.recursive: - # Recursive scan - for file_path in watch_path.rglob("*"): - if file_path.is_file() and self.watch_filter(Change.added, str(file_path)): - existing_files.add((Change.added, str(file_path))) - else: - # Non-recursive scan (only immediate children) - for file_path in watch_path.iterdir(): - if file_path.is_file() and self.watch_filter(Change.added, str(file_path)): - existing_files.add((Change.added, str(file_path))) - - if existing_files: - logger.info(f"[SCAN_ON_START] Found {len(existing_files)} existing files matching watch criteria") - await self.on_changes(existing_files) - logger.info(f"[SCAN_ON_START] Added {len(existing_files)} files to memory store") - else: - logger.info("[SCAN_ON_START] No existing files found matching watch criteria") - - if self.file_store is not None: - files: list[str] = await self.file_store.list_files(MemorySource.MEMORY) - for file_path in files: - chunks = await self.file_store.get_file_chunks(file_path, MemorySource.MEMORY) - logger.info(f"Found existing file: {file_path}, {len(chunks)} chunks") - - async def _interruptible_sleep(self, seconds: float): - """Sleep that can be interrupted by stop_event.""" - try: - await asyncio.wait_for(self._stop_event.wait(), timeout=seconds) - except asyncio.TimeoutError: - pass # Normal timeout, continue - - async def _watch_loop(self): - """Core monitoring loop with auto-restart on failure""" - if not self.watch_paths: - logger.warning("No watch paths specified") - return - - while not self._stop_event.is_set(): - # Filter out non-existent paths before each watch attempt - valid_paths = [p for p in self.watch_paths if Path(p).exists()] - - if not valid_paths: - logger.warning("No valid watch paths exist, waiting 10 seconds before retry...") - await self._interruptible_sleep(10) - continue - - invalid_paths = set(self.watch_paths) - set(valid_paths) - if invalid_paths: - logger.warning(f"Skipping non-existent paths: {invalid_paths}") - - try: - logger.info(f"Starting watch on valid paths: {valid_paths}") - async for changes in awatch( - *valid_paths, - watch_filter=self.watch_filter, - recursive=self.recursive, - debounce=self.debounce, - poll_delay_ms=self.poll_delay_ms, - stop_event=self._stop_event, - ): - if self._stop_event.is_set(): - break - - await self.on_changes(changes) - - except FileNotFoundError as e: - # Watch path was deleted during monitoring - logger.error(f"Watch path no longer exists: {e}, restarting in 10 seconds...") - if not self._stop_event.is_set(): - await self._interruptible_sleep(10) - - except Exception as e: - # Log other exceptions and restart - logger.error(f"Error in watch loop: {e}, restarting in 10 seconds...", exc_info=True) - if not self._stop_event.is_set(): - await self._interruptible_sleep(10) - - async def _on_changes(self, changes: set[tuple[Change, str]]): - """Callback method to handle file changes""" - - async def on_changes(self, changes: set[tuple[Change, str]]): - """Hook method to handle file changes""" - if self.callback: - result = self.callback(changes) - if asyncio.iscoroutine(result): - await result - else: - await self._on_changes(changes) - logger.info(f"[{self.__class__.__name__}] on_changes: {changes}") - - def is_running(self) -> bool: - """Check if the watcher is running""" - return self._running - - async def add_path(self, path: str): - """Dynamically add a path to monitor""" - if path not in self.watch_paths: - self.watch_paths.append(path) - if self._running: - await self.close() - await self.start() - - async def remove_path(self, path: str): - """Remove a monitored path""" - if path in self.watch_paths: - self.watch_paths.remove(path) - if self._running: - await self.close() - await self.start() diff --git a/reme_cli/component/file_watcher/delta_file_watcher.py b/reme_cli/component/file_watcher/delta_file_watcher.py deleted file mode 100644 index 6148bd07..00000000 --- a/reme_cli/component/file_watcher/delta_file_watcher.py +++ /dev/null @@ -1,280 +0,0 @@ -"""Delta file watcher for incremental file synchronization. - -This module provides a file watcher that detects append-only changes -and only processes newly added content, avoiding redundant operations. -""" - -import asyncio -import os - -from loguru import logger -from watchfiles import Change - -from .base_file_watcher import BaseFileWatcher -from ..enumeration import MemorySource -from ..schema import FileMetadata, MemoryChunk -from ..utils import chunk_markdown, hash_text - - -class DeltaFileWatcher(BaseFileWatcher): - """Delta file watcher implementation for incremental synchronization. - - This watcher detects append-only changes (e.g., log files) and only processes - the newly added content, avoiding redundant embedding requests for unchanged content. - - Strategy: - - Detect if file is append-only (new lines added at end) - - Find the safe cutoff point (considering chunk overlap) - - Only re-chunk and embed content from cutoff to end - - Delete affected old chunks and insert new chunks - """ - - def __init__(self, overlap_lines: int = 2, **kwargs): - """ - Initialize delta file watcher. - - Args: - chunk_tokens: Maximum tokens per chunk - chunk_overlap: Overlap tokens between chunks - """ - super().__init__(**kwargs) - self.overlap_lines = overlap_lines - self.dirty = False - - @staticmethod - async def _build_file_metadata(path: str) -> FileMetadata: - """Build file metadata from filesystem.""" - - def _read_file_sync(): - stat_t = os.stat(path) - with open(path, "r", encoding="utf-8") as f: - content_t = f.read() - return stat_t, content_t - - stat, content = await asyncio.to_thread(_read_file_sync) - return FileMetadata( - hash=hash_text(content), - mtime_ms=stat.st_mtime * 1000, - size=stat.st_size, - path=path, - content=content, - ) - - def _find_cutoff_line( - self, - old_chunks: list[MemoryChunk], - old_file_meta: FileMetadata, - new_file_meta: FileMetadata, - ) -> int | None: - """Find the safe cutoff line for incremental update. - - Uses a heuristic approach: if file size increased and hash changed, - we verify by comparing content. For true append-only files (like logs), - the old content should be a prefix of new content. - - Args: - old_chunks: Existing chunks sorted by start_line - old_file_meta: Previous file metadata - new_file_meta: Current file metadata (with content) - - Returns: - Cutoff line number (1-indexed), or None if not append-only - """ - if not old_chunks: - return None - - # File shrunk - definitely not append-only - if new_file_meta.size < old_file_meta.size: - logger.debug("File shrunk, not append-only") - return None - - # File didn't grow much - might be a modification - size_growth = new_file_meta.size - old_file_meta.size - if size_growth < 10: # Less than 10 bytes growth - logger.debug("Minimal size growth, treating as modification") - return None - - # Verify append-only by checking if old content is prefix - # We need to read old file content from chunks - old_chunks_sorted = sorted(old_chunks, key=lambda c: c.start_line) - - # Simple heuristic: check if first few chunks' content matches - # This avoids reconstructing full old content - new_lines = new_file_meta.content.split("\n") - - # Sample check: verify first chunk still matches - first_chunk = old_chunks_sorted[0] - first_chunk_lines = first_chunk.text.split("\n") - new_first_lines = new_lines[first_chunk.start_line - 1 : first_chunk.end_line] - - # Compare (allowing for minor whitespace differences at boundaries) - if len(first_chunk_lines) > 0 and len(new_first_lines) > 0: - # Check if most of the lines match - matches = sum(1 for old, new in zip(first_chunk_lines, new_first_lines) if old == new) - if matches < len(first_chunk_lines) * 0.8: # Less than 80% match - logger.debug("First chunk content changed, not append-only") - return None - - # File appears to be append-only - # Find the last chunk and set cutoff considering overlap - last_chunk = max(old_chunks_sorted, key=lambda c: c.end_line) - cutoff_line = max(1, last_chunk.end_line - self.overlap_lines) - - logger.debug( - f"Append-only detected: size {old_file_meta.size} -> {new_file_meta.size}, " - f"cutoff at line {cutoff_line}", - ) - - return cutoff_line - - @staticmethod - def _extract_content_from_line(content: str, start_line: int) -> str: - """Extract content starting from a specific line number.""" - lines = content.split("\n") - if start_line <= 1: - return content - if start_line > len(lines): - return "" - # start_line is 1-indexed, array is 0-indexed - return "\n".join(lines[start_line - 1 :]) - - async def _on_changes(self, changes: set[tuple[Change, str]]): - """Handle file changes with incremental synchronization.""" - self.dirty = True - - for change_type, path in changes: - if change_type == Change.added: - # New file: process everything - file_meta = await self._build_file_metadata(path) - chunks = ( - chunk_markdown( - file_meta.content, - file_meta.path, - MemorySource.MEMORY, - self.chunk_tokens, - self.chunk_overlap, - ) - or [] - ) - - if chunks: - chunks = await self.file_store.get_chunk_embeddings(chunks) - file_meta.chunk_count = len(chunks) - await self.file_store.upsert_file(file_meta, MemorySource.MEMORY, chunks) - logger.info(f"File added: {path} ({len(chunks)} chunks)") - else: - logger.warning(f"No chunks generated for new file {path}") - - elif change_type == Change.modified: - # Get existing data - old_chunks = await self.file_store.get_file_chunks(path, MemorySource.MEMORY) - old_file_meta = await self.file_store.get_file_metadata(path, MemorySource.MEMORY) - - # Read new file - file_meta = await self._build_file_metadata(path) - - # If no old chunks, fallback to full update - if not old_chunks or not old_file_meta: - logger.debug(f"No existing chunks for {path}, doing full update") - chunks = ( - chunk_markdown( - file_meta.content, - file_meta.path, - MemorySource.MEMORY, - self.chunk_tokens, - self.chunk_overlap, - ) - or [] - ) - if chunks: - chunks = await self.file_store.get_chunk_embeddings(chunks) - file_meta.chunk_count = len(chunks) - await self.file_store.delete_file(path, MemorySource.MEMORY) - await self.file_store.upsert_file( - file_meta, - MemorySource.MEMORY, - chunks, - ) - logger.info(f"File modified (full): {path} ({len(chunks)} chunks)") - continue - - # Check if append-only and find cutoff line - old_chunks_sorted = sorted(old_chunks, key=lambda c: c.start_line) - cutoff_line = self._find_cutoff_line(old_chunks_sorted, old_file_meta, file_meta) - - if cutoff_line is None: - # Not append-only, do full update - logger.debug(f"File {path} has modifications, doing full update") - chunks = ( - chunk_markdown( - file_meta.content, - file_meta.path, - MemorySource.MEMORY, - self.chunk_tokens, - self.chunk_overlap, - ) - or [] - ) - if chunks: - chunks = await self.file_store.get_chunk_embeddings(chunks) - file_meta.chunk_count = len(chunks) - await self.file_store.delete_file(path, MemorySource.MEMORY) - await self.file_store.upsert_file(file_meta, MemorySource.MEMORY, chunks) - logger.info(f"File modified (full): {path} ({len(chunks)} chunks)") - else: - # Append-only: incremental update - new_content_part = self._extract_content_from_line(file_meta.content, cutoff_line) - - new_chunks = ( - chunk_markdown( - new_content_part, - file_meta.path, - MemorySource.MEMORY, - self.chunk_tokens, - self.chunk_overlap, - ) - or [] - ) - - if not new_chunks: - logger.debug(f"No new chunks for {path}, skipping") - continue - - for idx, chunk in enumerate(new_chunks): - chunk.start_line += cutoff_line - 1 - chunk.end_line += cutoff_line - 1 - chunk.id = hash_text( - f"{chunk.source}:{chunk.path}:{chunk.start_line}:" f"{chunk.end_line}:{chunk.hash}:{idx}", - ) - - new_chunks = await self.file_store.get_chunk_embeddings(new_chunks) - - chunks_to_delete = [c.id for c in old_chunks_sorted if c.start_line >= cutoff_line] - - # Apply incremental updates - if chunks_to_delete: - await self.file_store.delete_file_chunks(path, chunks_to_delete) - - if new_chunks: - await self.file_store.upsert_chunks(new_chunks, MemorySource.MEMORY) - - # Update file metadata to reflect the changes - # Calculate new chunk count: old chunks - deleted + new chunks - new_chunk_count = len(old_chunks) - len(chunks_to_delete) + len(new_chunks) - file_meta.chunk_count = new_chunk_count - await self.file_store.update_file_metadata(file_meta, MemorySource.MEMORY) - - logger.info( - f"File modified (incremental): {path} " - f"(cutoff: line {cutoff_line}, " - f"+{len(new_chunks)} chunks, -{len(chunks_to_delete)} chunks)", - ) - - elif change_type == Change.deleted: - await self.file_store.delete_file(path, MemorySource.MEMORY) - logger.info(f"File deleted: {path}") - - else: - logger.warning(f"Unknown change type: {change_type}") - - self.dirty = False diff --git a/reme_cli/component/file_watcher/full_file_watcher.py b/reme_cli/component/file_watcher/full_file_watcher.py deleted file mode 100644 index 5b1852b9..00000000 --- a/reme_cli/component/file_watcher/full_file_watcher.py +++ /dev/null @@ -1,79 +0,0 @@ -"""Full file watcher for complete file synchronization. - -This module provides a file watcher that processes entire files -on any change, ensuring complete synchronization. -""" - -import asyncio -from pathlib import Path - -from loguru import logger -from watchfiles import Change - -from .base_file_watcher import BaseFileWatcher -from ..enumeration import MemorySource -from ..schema import FileMetadata -from ..utils import chunk_markdown, hash_text - - -class FullFileWatcher(BaseFileWatcher): - """Full file watcher implementation for full synchronization""" - - def __init__(self, **kwargs): - """ - Initialize full file watcher""" - super().__init__(**kwargs) - self.dirty = False - - @staticmethod - async def _build_file_metadata(path: str) -> FileMetadata: - file_path = Path(path) - - def _read_file_sync(): - return file_path.stat(), file_path.read_text(encoding="utf-8") - - stat, content = await asyncio.to_thread(_read_file_sync) - return FileMetadata( - hash=hash_text(content), - mtime_ms=stat.st_mtime * 1000, - size=stat.st_size, - path=str(file_path.absolute()), - content=content, - ) - - async def _on_changes(self, changes: set[tuple[Change, str]]): - """Handle file changes with full synchronization""" - self.dirty = True - - for change_type, path in changes: - if change_type in [Change.added, Change.modified]: - file_meta = await self._build_file_metadata(path) - chunks = ( - chunk_markdown( - file_meta.content, - file_meta.path, - MemorySource.MEMORY, - self.chunk_tokens, - self.chunk_overlap, - ) - or [] - ) - if chunks: - chunks = await self.file_store.get_chunk_embeddings(chunks) - file_meta.chunk_count = len(chunks) - - await self.file_store.delete_file(file_meta.path, MemorySource.MEMORY) - logger.info(f"delete_file {file_meta.path}") - - await self.file_store.upsert_file(file_meta, MemorySource.MEMORY, chunks) - logger.info(f"Upserted {file_meta.chunk_count} chunks for {file_meta.path}") - - elif change_type == Change.deleted: - await self.file_store.delete_file(path, MemorySource.MEMORY) - logger.info(f"Deleted {path}") - - else: - logger.warning(f"Unknown change type: {change_type}") - - logger.info(f"File {change_type} changed: {path}") - self.dirty = False diff --git a/reme_cli/component/flow/__init__.py b/reme_cli/component/flow/__init__.py new file mode 100644 index 00000000..3b7e2753 --- /dev/null +++ b/reme_cli/component/flow/__init__.py @@ -0,0 +1,14 @@ +"""flow""" + +from .base_flow import BaseFlow +from .cmd_flow import CmdFlow +from .expression_flow import ExpressionFlow +from ..registry_factory import R + +__all__ = [ + "BaseFlow", + "CmdFlow", + "ExpressionFlow", +] + +R.flows.register(ExpressionFlow) diff --git a/reme_cli/component/flow/base_flow.py b/reme_cli/component/flow/base_flow.py new file mode 100644 index 00000000..68e66fd5 --- /dev/null +++ b/reme_cli/component/flow/base_flow.py @@ -0,0 +1,208 @@ +"""Base flow module providing abstract flow execution with caching and operation orchestration.""" + +import asyncio +import hashlib +import json +from abc import ABC, abstractmethod + +from loguru import logger + +from ..enumeration import ChunkEnum +from ..op import BaseOp, SequentialOp, ParallelOp +from ..registry_factory import R +from ..runtime_context import RuntimeContext +from ..schema import Response, ToolCall +from ..service_context import ServiceContext +from ..utils import camel_to_snake, CacheHandler + + +class BaseFlow(ABC): + """Abstract base class for flow execution with caching, streaming, and operation tree management.""" + + def __init__( + self, + name: str = "", + stream: bool = False, + raise_exception: bool = True, + enable_cache: bool = False, + cache_path: str = "cache/flow", + cache_expire_hours: float = 0.1, + service_context: ServiceContext | None = None, + **kwargs, + ): + """Initialize flow configuration and execution state.""" + super().__init__() + + self.name: str = name or camel_to_snake(self.__class__.__name__) + self.stream: bool = stream + self.raise_exception: bool = raise_exception + self.enable_cache: bool = enable_cache + self.cache_path: str = cache_path + self.cache_expire_hours: float = cache_expire_hours + self.service_context: ServiceContext | None = service_context + self.flow_params: dict = kwargs + + self._cache: CacheHandler | None = None + self._flow_printed: bool = False + self._flow_op: BaseOp | None = None + self._tool_call: ToolCall | None = None + + def _build_tool_call(self) -> ToolCall | None: + """Generate the tool call schema definition for this flow.""" + + @abstractmethod + def _build_flow(self) -> BaseOp: + """Construct the root operation tree for flow execution.""" + + def _compute_cache_key(self, params: dict) -> str | None: + """Generate a SHA256 hash from input parameters for caching.""" + try: + payload = json.dumps(params, sort_keys=True, ensure_ascii=False, default=str) + return hashlib.sha256(payload.encode("utf-8")).hexdigest() + except Exception as e: + logger.exception(f"[{self.__class__.__name__}] {self.name} cache key serialization failed: {e}") + return None + + def _maybe_load_cached(self, params: dict) -> Response | None: + """Retrieve a cached response if caching is enabled and available.""" + if not self.enable_cache or self.stream: + return None + + if key := self._compute_cache_key(params): + if cached := self.cache.load(key): + logger.info(f"[{self.__class__.__name__}] Loaded {self.name} response from cache.") + return Response(**cached) + return None + + def _maybe_save_cache(self, params: dict, response: Response): + """Persist the execution response to the cache.""" + if not self.enable_cache or self.stream: + return + + if key := self._compute_cache_key(params): + self.cache.save(key, response.model_dump(exclude_none=True), expire_hours=self.cache_expire_hours) + + def _print_operation_tree(self, name: str, op: BaseOp, indent: int): + """Recursively log the hierarchy of the flow's operation tree.""" + prefix = " " * indent + op_type = "sequential" if isinstance(op, SequentialOp) else "parallel" if isinstance(op, ParallelOp) else name + logger.info(f"[{self.__class__.__name__}] {prefix}{op_type} execution") + + for sub_op in op.sub_ops or []: + self._print_operation_tree(sub_op.name, sub_op, indent + 2) + + @property + def tool_call(self) -> ToolCall | None: + """Lazily construct the ToolCall schema describing this flow.""" + if hasattr(self.flow_op, "tool_call"): + return self.flow_op.tool_call + + if self._tool_call is None: + self._tool_call = self._build_tool_call() + if self._tool_call: + self._tool_call.name = self._tool_call.name or self.name + return self._tool_call + + @property + def cache(self) -> CacheHandler: + """Provide access to the internal CacheHandler instance.""" + assert self.enable_cache, "Cache usage requested while disabled." + if self._cache is None: + self._cache = CacheHandler(f"{self.cache_path}/{self.name}") + return self._cache + + @property + def flow_op(self) -> BaseOp: + """Lazily build and retrieve the root operation of the flow.""" + if self._flow_op is None: + self._flow_op = self._build_flow() + return self._flow_op + + @property + def async_mode(self) -> bool: + """Check if the current flow operation tree is asynchronous.""" + return self.flow_op.async_mode + + @staticmethod + def parse_expression(expression: str) -> BaseOp: + """Parse a string expression into an executable BaseOp instance.""" + lines = [x.strip() for x in expression.strip().splitlines() if x.strip()] + if not lines: + raise ValueError("Expression is empty") + + if len(lines) > 1: + exec("\n".join(lines[:-1]), {"__builtins__": {}}, R.ops) + + result = eval(lines[-1], {"__builtins__": {}}, R.ops) + if not isinstance(result, BaseOp): + raise TypeError(f"Expression evaluated to {type(result)}, expected BaseOp") + return result + + def print_flow(self): + """Log the visual structure of the flow once.""" + if not self._flow_printed: + logger.info(f"[{self.__class__.__name__}] ---------- [Flow Structure] {self.name} [Start] ----------") + self._print_operation_tree(self.name, self.flow_op, 0) + logger.info(f"[{self.__class__.__name__}] ---------- [Flow Structure] {self.name} [End] ----------") + self._flow_printed = True + + async def call(self, **kwargs) -> Response | asyncio.Queue: + """Execute the flow asynchronously with parameter caching.""" + kwargs["stream"] = self.stream + logger.info(f"[{self.__class__.__name__}] {self.name} incoming params: {kwargs}") + if cached := self._maybe_load_cached(kwargs): + return cached + + context = RuntimeContext(service_context=self.service_context, **kwargs) + try: + self.print_flow() + flow_op: BaseOp = self._build_flow() + assert self.flow_op.async_mode, "Async call requires an async flow operation." + await flow_op.call(context=context) + + if self.stream: + await context.add_stream_done() + return context.stream_queue + + else: + self._maybe_save_cache(kwargs, context.response) + return context.response + + except Exception as e: + logger.exception(f"[{self.__class__.__name__}] {self.name} async call failed: {e}") + if self.raise_exception: + raise e + + if self.stream: + await context.add_stream_chunk_and_type(str(e), ChunkEnum.ERROR) + await context.add_stream_done() + return context.stream_queue + + else: + context.add_response_error(e) + return context.response + + def call_sync(self, **kwargs) -> Response: + """Execute the flow synchronously with parameter caching.""" + logger.info(f"[{self.__class__.__name__}] {self.name} incoming sync params: {kwargs}") + assert not self.stream, "Synchronous call cannot be used in stream mode." + if cached := self._maybe_load_cached(kwargs): + return cached + + context = RuntimeContext(service_context=self.service_context, **kwargs) + try: + self.print_flow() + flow_op: BaseOp = self._build_flow() + assert not self.flow_op.async_mode, "Sync call requires a sync flow operation." + flow_op.call_sync(context=context) + + self._maybe_save_cache(kwargs, context.response) + return context.response + + except Exception as e: + logger.exception(f"[{self.__class__.__name__}] {self.name} sync call failed: {e}") + if self.raise_exception: + raise e + + context.add_response_error(e) + return context.response diff --git a/reme_cli/component/flow/cmd_flow.py b/reme_cli/component/flow/cmd_flow.py new file mode 100644 index 00000000..695ab48e --- /dev/null +++ b/reme_cli/component/flow/cmd_flow.py @@ -0,0 +1,18 @@ +"""Command-based flow implementation for parsing and executing operation sequences.""" + +from .base_flow import BaseFlow +from ..op import BaseOp + + +class CmdFlow(BaseFlow): + """A flow class that builds an operation chain from a string expression.""" + + def __init__(self, flow: str = "", **kwargs): + """Initialize the command flow with a string-based operation definition.""" + super().__init__(**kwargs) + self.flow = flow + assert flow, "add `cmd.flow=` in cmd!" + + def _build_flow(self) -> BaseOp: + """Parse the stored flow expression into a functional operation object.""" + return self.parse_expression(self.flow) diff --git a/reme_cli/component/flow/expression_flow.py b/reme_cli/component/flow/expression_flow.py new file mode 100644 index 00000000..a844c719 --- /dev/null +++ b/reme_cli/component/flow/expression_flow.py @@ -0,0 +1,37 @@ +"""Expression-based flow implementation driven by configuration objects.""" + +from .base_flow import BaseFlow +from ..op import BaseOp +from ..schema import FlowConfig, ToolCall +from ..service_context import ServiceContext + + +class ExpressionFlow(BaseFlow): + """A flow implementation that constructs operations from a FlowConfig definition.""" + + def __init__(self, flow_config: FlowConfig, service_context: ServiceContext): + """Initialize the flow using settings and metadata from a FlowConfig instance.""" + self.flow_config: FlowConfig = flow_config + super().__init__( + name=flow_config.name, + stream=self.flow_config.stream, + raise_exception=self.flow_config.raise_exception, + enable_cache=self.flow_config.enable_cache, + cache_path=self.flow_config.cache_path, + cache_expire_hours=self.flow_config.cache_expire_hours, + service_context=service_context, + **flow_config.model_extra, + ) + + def _build_flow(self) -> BaseOp: + """Generate the operation chain by parsing the flow content string.""" + return self.parse_expression(self.flow_config.flow_content) + + def _build_tool_call(self) -> ToolCall: + """Construct a tool call representation based on configuration parameters.""" + return ToolCall( + **{ + "description": self.flow_config.description, + "parameters": self.flow_config.parameters, + }, + ) diff --git a/reme_cli/op/__init__.py b/reme_cli/component/op/__init__.py similarity index 100% rename from reme_cli/op/__init__.py rename to reme_cli/component/op/__init__.py diff --git a/reme_cli/op/base_op.py b/reme_cli/component/op/base_op.py similarity index 100% rename from reme_cli/op/base_op.py rename to reme_cli/component/op/base_op.py diff --git a/reme_cli/component/registry_factory.py b/reme_cli/component/registry_factory.py deleted file mode 100644 index 44cf3364..00000000 --- a/reme_cli/component/registry_factory.py +++ /dev/null @@ -1,50 +0,0 @@ -"""Module providing a registry class for managing class-to-name mappings via decorators.""" - -import inspect -from typing import Callable, TypeVar - -from .base_dict import BaseDict -from ..utils import singleton - -T = TypeVar("T") - - -class Registry(BaseDict): - """A registry container that uses decorators to map and store class references.""" - - def register(self, name: str | type = "") -> Callable[[type[T]], type[T]] | type[T]: - """Return a decorator that registers a class under a specific name in the registry.""" - if inspect.isclass(name): - self[name.__name__] = name - return name - - else: - - def decorator(cls): - key: str = name if isinstance(name, str) and name else cls.__name__ - self[key] = cls - return cls - - return decorator - - -@singleton -class RegistryFactory: - """A factory class for creating registries.""" - - def __init__(self): - self.llms = Registry() - self.as_llms = Registry() - self.as_llm_formatters = Registry() - self.as_token_counters = Registry() - self.embedding_models = Registry() - self.vector_stores = Registry() - self.file_stores = Registry() - self.ops = Registry() - self.flows = Registry() - self.services = Registry() - self.token_counters = Registry() - self.file_watchers = Registry() - - -R = RegistryFactory() diff --git a/reme_cli/enumeration/__init__.py b/reme_cli/enumeration/__init__.py index f2734a36..a394e324 100644 --- a/reme_cli/enumeration/__init__.py +++ b/reme_cli/enumeration/__init__.py @@ -1,7 +1,9 @@ """enumeration""" from .component_enum import ComponentEnum +from .json_schema_enum import JsonSchemaEnum __all__ = [ "ComponentEnum", + "JsonSchemaEnum", ] diff --git a/reme_cli/enumeration/json_schema_enum.py b/reme_cli/enumeration/json_schema_enum.py new file mode 100644 index 00000000..b52a8173 --- /dev/null +++ b/reme_cli/enumeration/json_schema_enum.py @@ -0,0 +1,41 @@ +"""Defines the standard data types supported by JSON Schema. + +This enum maps common JSON Schema primitive types to their corresponding +Python runtime types, and provides a convenient string representation +compatible with JSON Schema (`"string"`, `"number"`, etc.). +""" + +from enum import Enum + + +class JsonSchemaEnum(Enum): + """Enumeration of valid JSON Schema data types. + + The enum value is the corresponding Python type, while the string + representation (`str(...)`) is the canonical JSON Schema type name. + """ + + # Textual data + STRING = str + + # Numeric values, including integers and floats + NUMBER = float + + # Integer-only numeric values + INTEGER = int + + # JSON objects (key-value mappings) + OBJECT = dict + + # Ordered JSON lists/arrays + ARRAY = list + + # Boolean values: true / false + BOOLEAN = bool + + # Null / None values + NULL = type(None) + + def __str__(self) -> str: + """Return the lowercase JSON Schema type name for this enum member.""" + return self.name.lower() diff --git a/reme_cli/schema/__init__.py b/reme_cli/schema/__init__.py index e69de29b..ed4b2c12 100644 --- a/reme_cli/schema/__init__.py +++ b/reme_cli/schema/__init__.py @@ -0,0 +1,7 @@ +from .application_config import ApplicationConfig +from .base_node import BaseNode + +__all__ = [ + "ApplicationConfig", + "BaseNode", +] diff --git a/reme_cli/schema/application_config.py b/reme_cli/schema/application_config.py new file mode 100644 index 00000000..0d439f1b --- /dev/null +++ b/reme_cli/schema/application_config.py @@ -0,0 +1,20 @@ +"""Configuration schemas for service components using Pydantic models.""" + +import os + +from pydantic import Field, BaseModel + +from ..enumeration import ComponentEnum + + +class ApplicationConfig(BaseModel): + app_name: str = Field(default=os.getenv("APP_NAME", "ReMe")) + working_dir: str = Field(default=".reme") + enable_logo: bool = Field(default=False) + language: str = Field(default="") + log_to_console: bool = Field(default=True) + mcp_servers: dict[str, dict] = Field(default_factory=dict) + service: dict = Field(default_factory=dict) + ops: dict[str, dict] = Field(default_factory=dict) + flows: dict[str, dict] = Field(default_factory=dict) + components: dict[ComponentEnum, dict[str, dict]] = Field(default_factory=dict) diff --git a/reme_cli/schema/base_node.py b/reme_cli/schema/base_node.py new file mode 100644 index 00000000..2e82d916 --- /dev/null +++ b/reme_cli/schema/base_node.py @@ -0,0 +1,10 @@ +from uuid import uuid4 + +from pydantic import BaseModel, Field + + +class BaseNode(BaseModel): + id: str = Field(default_factory=lambda: uuid4().hex) + text: str = Field(default="") + embedding: list[float] | None = Field(default=None) + metadata: dict = Field(default_factory=dict) diff --git a/reme_cli/schema/service_config.py b/reme_cli/schema/service_config.py deleted file mode 100644 index 48b97e0a..00000000 --- a/reme_cli/schema/service_config.py +++ /dev/null @@ -1,144 +0,0 @@ -"""Configuration schemas for service components using Pydantic models.""" - -import os - -from pydantic import BaseModel, Field, ConfigDict - -from .tool_call import ToolCall - - -class MCPConfig(BaseModel): - """Configuration for Model Context Protocol transport and network settings.""" - - model_config = ConfigDict(extra="allow") - - transport: str = Field(default="stdio") - host: str = Field(default="0.0.0.0") - port: int = Field(default=8001) - - -class HttpConfig(BaseModel): - """Configuration for the HTTP server interface and connection lifecycle.""" - - model_config = ConfigDict(extra="allow") - - host: str = Field(default="0.0.0.0") - port: int = Field(default=8001) - timeout_keep_alive: int = Field(default=3600) - limit_concurrency: int = Field(default=1000) - - -class CmdConfig(BaseModel): - """Configuration for command-line flow execution parameters.""" - - model_config = ConfigDict(extra="allow") - - flow: str = Field(default="") - - -class OpConfig(BaseModel): - """Configuration for op settings and parameters.""" - - model_config = ConfigDict(extra="allow") - - prompt_dict: dict[str, str] = Field(default_factory=dict) - params: dict = Field(default_factory=dict) - - -class FlowConfig(ToolCall): - """Configuration for workflow execution, caching, and error handling.""" - - model_config = ConfigDict(extra="allow") - - flow_content: str = Field(default="") - stream: bool = Field(default=False) - raise_exception: bool = Field(default=True) - enable_cache: bool = Field(default=False) - cache_path: str = Field(default="cache/flow") - cache_expire_hours: float = Field(default=0.1) - - -class BasicConfig(BaseModel): - """Configuration for basic service settings and parameters.""" - - model_config = ConfigDict(extra="allow") - - backend: str = Field(default="") - - -class ModelConfig(BasicConfig): - """Configuration for model-based services with backend and model name.""" - - model_name: str = Field(default="") - - -class LLMConfig(ModelConfig): - """Configuration for Large Language Model backend and model identification.""" - - -class EmbeddingModelConfig(ModelConfig): - """Configuration for embedding model backends and identity.""" - - -class TokenCounterConfig(ModelConfig): - """Configuration for token counting services and model mapping.""" - - -class StoreConfig(BasicConfig): - """Configuration for storage services with embedding model support.""" - - embedding_model: str = Field(default="default") - - -class VectorStoreConfig(StoreConfig): - """Configuration for vector database storage and associated embeddings.""" - - collection_name: str = Field(default="reme") - - -class FileStoreConfig(StoreConfig): - """Configuration for file store database storage and associated embeddings.""" - - store_name: str = Field(default="reme") - - -class FileWatcherConfig(BasicConfig): - """Configuration for file watcher service.""" - - file_store: str = Field(default="") - watch_paths: list[str] = Field(default_factory=list) - - -class ServiceConfig(BasicConfig): - """Root configuration schema aggregating all service-level settings and components.""" - - app_name: str = Field(default=os.getenv("APP_NAME", "ReMe")) - working_dir: str = Field(default=".reme") - enable_logo: bool = Field(default=True) - language: str = Field(default="") - thread_pool_max_workers: int = Field( - default=16, - description="Number of thread pool workers. Set to -1 to disable thread pool.", - ) - ray_max_workers: int = Field(default=-1) - log_to_console: bool = Field(default=True) - disabled_flows: list[str] = Field(default_factory=list) - enabled_flows: list[str] = Field(default_factory=list) - - mcp_servers: dict[str, dict] = Field(default_factory=dict) - mcp: MCPConfig = Field(default_factory=MCPConfig) - http: HttpConfig = Field(default_factory=HttpConfig) - cmd: CmdConfig = Field(default_factory=CmdConfig) - ops: dict[str, OpConfig] = Field(default_factory=dict) - flows: dict[str, FlowConfig] = Field(default_factory=dict) - as_llms: dict[str, BasicConfig] = Field(default_factory=dict) - as_llm_formatters: dict[str, BasicConfig] = Field(default_factory=dict) - as_token_counters: dict[str, BasicConfig] = Field(default_factory=dict) - llms: dict[str, LLMConfig] = Field(default_factory=dict) - embedding_models: dict[str, EmbeddingModelConfig] = Field(default_factory=dict) - vector_stores: dict[str, VectorStoreConfig] = Field(default_factory=dict) - file_stores: dict[str, FileStoreConfig] = Field(default_factory=dict) - token_counters: dict[str, TokenCounterConfig] = Field(default_factory=dict) - file_watchers: dict[str, FileWatcherConfig] = Field(default_factory=dict) - - metadata: dict = Field(default_factory=dict) diff --git a/reme_cli/schema/tool_call.py b/reme_cli/schema/tool_call.py new file mode 100644 index 00000000..8b0f4eb9 --- /dev/null +++ b/reme_cli/schema/tool_call.py @@ -0,0 +1,196 @@ +"""MCP Tool Schema definitions for recursive JSON Schema representation.""" + +import json +from typing import Optional + +from mcp.types import Tool +from pydantic import BaseModel, ConfigDict, Field, model_validator, field_validator + +from ..enumeration import JsonSchemaEnum + + +class ToolAttr(BaseModel): + """Recursive model representing JSON Schema attributes for tool parameters.""" + + model_config = ConfigDict(extra="allow") + + type: str = Field(default=str(JsonSchemaEnum.STRING), description="The data type of the attribute") + description: Optional[str] = Field(default=None, description="Description of the attribute") + required: Optional[list[str]] = Field(default=None, description="Required property names for object types") + properties: Optional[dict[str, "ToolAttr"]] = Field(default=None, description="Child properties for objects") + items: Optional["ToolAttr"] = Field(default=None, description="Schema for array items") + enum: Optional[list[str]] = Field(default=None, description="Allowed values for the attribute") + + @field_validator("type") + @classmethod + def validate_type_is_valid_enum(cls, v: str) -> str: + """Validates that the provided type string exists within JsonSchemaEnum values.""" + valid_types = [str(e) for e in JsonSchemaEnum] + + if v not in valid_types: + raise ValueError(f"Invalid type: '{v}'. Must be one of {valid_types}") + return v + + def simple_input_dump(self) -> dict: + """Serializes the attribute into a standard JSON Schema dictionary.""" + res: dict = {} + + # Lay down extra fields first so explicit fields can override them + if self.model_extra: + res.update(self.model_extra) + + res["type"] = self.type + if self.description: + res["description"] = self.description + if self.enum: + res["enum"] = self.enum + + if self.type == "object" and self.properties is not None: + res["properties"] = {k: v.simple_input_dump() for k, v in self.properties.items()} + if self.required is not None: + res["required"] = self.required + + if self.type == "array" and self.items is not None: + res["items"] = self.items.simple_input_dump() + + return res + + +# Enable recursive type resolution +ToolAttr.model_rebuild() + + +class ToolCall(BaseModel): + """ + Model representing a tool definition and its call structure. + Supports parsing from standard JSON Schema formats and converting to MCP Tool objects. + input: + { + "type": "function", + "function": { + "name": "get_current_weather", + "description": "It is very useful when you want to check the weather of a specified city.", + "parameters": { + "type": "object", + "properties": { + "location": { + "type": "string", + "description": "Cities or counties, such as Beijing, Hangzhou, Yuhang District, etc.", + } + }, + "required": ["location"] + } + } + } + output: + { + "index": 0, + "id": "call_6596dafa2a6a46f7a217da", + "function": { + "arguments": "{\"location\": \"Beijing\"}", + "name": "get_current_weather" + }, + "type": "function", + } + """ + + index: int = 0 + id: str = "" + type: str = "function" + name: str = "" + description: str = "" + + arguments: str = Field(default="", description="JSON string of tool execution arguments") + + parameters: ToolAttr = Field( + default_factory=lambda: ToolAttr(type="object", properties={}, required=[]), + description="Specification for input parameters", + ) + output: Optional[ToolAttr] = Field(default=None, description="Output schema") + + @model_validator(mode="before") + @classmethod + def init_tool_call(cls, data: dict) -> dict: + """Initializes the model by parsing tool-specific body data.""" + data = data.copy() + t_type = data.get("type", "function") + body = data.get(t_type, {}) + + # Extract basic metadata + data["name"] = body.get("name", data.get("name", "")) + data["arguments"] = body.get("arguments", data.get("arguments", "")) + data["description"] = body.get("description", data.get("description", "")) + + # Handle parameters mapping + if "parameters" in body: + params = body["parameters"] + # If parameters is already a dict, ensure it matches ToolAttr structure + if isinstance(params, dict): + data["parameters"] = ToolAttr(**params) + + # Handle output mapping (if provided in source) + if "output" in body and isinstance(body["output"], dict): + data["output"] = ToolAttr(**body["output"]) + + return data + + def simple_input_dump(self, as_dict: bool = True) -> dict | str: + """Returns a standardized tool definition dictionary or JSON string. + + Args: + as_dict: If True, returns dict; if False, returns JSON string. + """ + result = { + "type": self.type, + self.type: { + "name": self.name, + "description": self.description, + "parameters": self.parameters.simple_input_dump(), + }, + } + return result if as_dict else json.dumps(result, ensure_ascii=False) + + def simple_output_dump(self, as_dict: bool = True, enable_argument_dict: bool = False) -> dict | str: + """Convert ToolCall to output format dictionary or JSON string for API responses.""" + result = { + "index": self.index, + "id": self.id, + self.type: { + "arguments": self.argument_dict if enable_argument_dict else self.arguments, + "name": self.name, + }, + "type": self.type, + } + return result if as_dict else json.dumps(result, ensure_ascii=False) + + @property + def argument_dict(self) -> dict: + """Parse and return arguments as a dictionary.""" + if not self.arguments or not self.arguments.strip(): + return {} + return json.loads(self.arguments) + + def check_argument(self) -> bool: + """Check if arguments can be parsed as valid JSON.""" + try: + _ = self.argument_dict + return True + except Exception: + return False + + @classmethod + def from_mcp_tool(cls, tool: Tool) -> "ToolCall": + """Creates a ToolCall instance from an MCP Tool object.""" + return cls( + name=tool.name, + description=tool.description or "", + parameters=ToolAttr(**tool.inputSchema), + ) + + def to_mcp_tool(self) -> Tool: + """Converts the instance back into an MCP Tool object.""" + return Tool( + name=self.name, + description=self.description, + inputSchema=self.parameters.simple_input_dump(), + ) diff --git a/reme_cli/utils/__init__.py b/reme_cli/utils/__init__.py index 77c562a1..f2fdaa87 100644 --- a/reme_cli/utils/__init__.py +++ b/reme_cli/utils/__init__.py @@ -1,5 +1,9 @@ +from .logger_utils import get_logger +from .pydantic_config_parser import PydanticConfigParser from .singleton import singleton __all__ = [ + "get_logger", + "PydanticConfigParser", "singleton", ] diff --git a/reme_cli/utils/logger_utils.py b/reme_cli/utils/logger_utils.py index f6928a4a..e7636f64 100644 --- a/reme_cli/utils/logger_utils.py +++ b/reme_cli/utils/logger_utils.py @@ -4,21 +4,40 @@ import os import sys from datetime import datetime +from loguru import logger -def init_logger(log_dir: str = "logs", level: str = "INFO", log_to_console: bool = True) -> None: - """Initialize the logger with both file and console handlers. +_initialized = False + + +def get_logger( + log_dir: str = "logs", + level: str = "INFO", + log_to_console: bool = True, + force_init: bool = False, +): + """Get a configured logger instance. + + Automatically initializes on first call. Subsequent calls return + the same logger without re-initializing unless force_init=True. Args: - log_dir: Directory path for log files - level: Logging level (DEBUG, INFO, WARNING, ERROR, CRITICAL) - log_to_console: Whether to print logs to console/screen + log_dir: Directory path for log files. + level: Logging level (DEBUG, INFO, WARNING, ERROR, CRITICAL). + log_to_console: Whether to print logs to console/screen. + force_init: Force re-initialization even if already initialized. + + Returns: + The configured logger instance. """ - from loguru import logger + global _initialized + + if _initialized and not force_init: + return logger # Remove default handler to avoid duplicate logs logger.remove() - # Configure colorized standard output logging if enabled + # Configure colorized console logging if enabled if log_to_console: logger.add( sink=sys.stdout, @@ -27,18 +46,13 @@ def init_logger(log_dir: str = "logs", level: str = "INFO", log_to_console: bool colorize=True, ) - # Try to configure file-based logging (skip if permission denied) + # Configure file-based logging (skip if permission denied) try: - # Ensure the logging directory exists os.makedirs(log_dir, exist_ok=True) - # Generate filename based on the current timestamp - # Use dashes instead of colons for Windows compatibility current_ts = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") - log_filename = f"{current_ts}.log" - log_filepath = os.path.join(log_dir, log_filename) + log_filepath = os.path.join(log_dir, f"{current_ts}.log") - # Configure file-based logging with rotation and compression logger.add( log_filepath, level=level, @@ -51,9 +65,5 @@ def init_logger(log_dir: str = "logs", level: str = "INFO", log_to_console: bool except Exception as e: logger.error(f"Error configuring file logging: {e}") - -def get_logger(): - """Get a configured logger instance using loguru.""" - from loguru import logger - - return logger + _initialized = True + return logger \ No newline at end of file diff --git a/reme_cli/utils/pydantic_config_parser.py b/reme_cli/utils/pydantic_config_parser.py new file mode 100644 index 00000000..f75dd07c --- /dev/null +++ b/reme_cli/utils/pydantic_config_parser.py @@ -0,0 +1,194 @@ +"""Parser for Pydantic config models with YAML and CLI argument support.""" + +import inspect +import json +from pathlib import Path +from typing import Any, TypeVar + +import yaml +from pydantic import BaseModel +from .logger_utils import get_logger + +T = TypeVar("T", bound=BaseModel) + + +class PydanticConfigParser: + """Parser that loads and merges Pydantic configs from YAML files and CLI args.""" + + def __init__(self, config_class: type[T]): + """Initialize parser with a Pydantic config class. + + Args: + config_class: Pydantic BaseModel class to validate configs against. + """ + self.config_class = config_class + self.config_dict: dict = {} + self.logger = get_logger() + + def _deep_merge(self, base_dict: dict, update_dict: dict) -> dict: + """Recursively merge two dictionaries.""" + result = base_dict.copy() + for key, value in update_dict.items(): + if key in result and isinstance(result[key], dict) and isinstance(value, dict): + result[key] = self._deep_merge(result[key], value) + else: + result[key] = value + return result + + @staticmethod + def _convert_value(value_str: str) -> Any: + """Convert string value to appropriate Python type.""" + value_str = value_str.strip() + lower_str = value_str.lower() + + # Boolean and None conversion + if lower_str in ("true", "false"): + return lower_str == "true" + if lower_str in ("none", "null"): + return None + + # Numeric conversion + if "e" in lower_str or "." in value_str: + try: + return float(value_str) + except ValueError: + pass + else: + try: + return int(value_str) + except ValueError: + pass + + # JSON conversion for complex types + try: + return json.loads(value_str) + except (json.JSONDecodeError, ValueError): + return value_str + + @staticmethod + def load_from_yaml(yaml_path: str | Path) -> dict: + """Load configuration from YAML file. + + Args: + yaml_path: Path to YAML configuration file. + + Returns: + Dictionary containing configuration data. + + Raises: + FileNotFoundError: If YAML file does not exist. + """ + if isinstance(yaml_path, str): + yaml_path = Path(yaml_path) + + if not yaml_path.exists(): + raise FileNotFoundError(f"Configuration file does not exist: {yaml_path}") + + with yaml_path.open(encoding="utf-8") as f: + return yaml.safe_load(f) or {} + + def merge_configs(self, *config_dicts: dict) -> dict: + """Merge multiple config dictionaries in order. + + Args: + *config_dicts: Variable number of config dictionaries to merge. + + Returns: + Merged configuration dictionary. + """ + result = {} + for config_dict in config_dicts: + result = self._deep_merge(result, config_dict) + return result + + def parse_dot_notation(self, dot_list: list[str]) -> dict: + """Parse dot notation strings into nested dictionary. + + Args: + dot_list: List of strings in format "key.subkey=value". + + Returns: + Nested dictionary representation of dot notation. + """ + config_dict = {} + for item in dot_list: + if "=" not in item: + continue + + key_path, value_str = item.split("=", 1) + keys = key_path.split(".") + + # Build nested dictionary + current = config_dict + for key in keys[:-1]: + current = current.setdefault(key, {}) + current[keys[-1]] = self._convert_value(value_str) + + return config_dict + + def _find_config_path(self, config_name: str) -> Path: + """Find config file path, trying parser directory first then current directory.""" + if not config_name.endswith(".yaml"): + config_name += ".yaml" + + # Try parser class directory first + config_path = Path(inspect.getfile(self.__class__)).parent / config_name + if config_path.exists(): + self.logger.info(f"load config={config_path}") + return config_path + + # Try current directory + self.logger.warning(f"config={config_path} not found, try {config_name}") + config_path = Path(config_name) + if not config_path.exists(): + raise FileNotFoundError(f"config={config_path} not found") + return config_path + + def parse_args(self, *args: str, **kwargs) -> T: + """Parse CLI arguments and load configs from YAML files.""" + configs_to_merge = [self.config_class().model_dump()] + + # Separate config file path from other arguments + config = "" + filter_args = [] + for arg in args: + if "=" not in arg: + continue + arg = arg.lstrip("-") + if arg.startswith(("c=", "config=")): + config = arg.split("=", 1)[1] + else: + filter_args.append(arg) + + # Load each config file + for single_config in (c.strip() for c in config.split(",") if c.strip()): + config_path = self._find_config_path(single_config) + configs_to_merge.append(self.load_from_yaml(config_path)) + + # Apply CLI overrides + if filter_args: + configs_to_merge.append(self.parse_dot_notation(filter_args)) + + if kwargs: + configs_to_merge.append(kwargs) + + # Merge all configs and validate + self.config_dict = self.merge_configs(*configs_to_merge) + return self.config_class.model_validate(self.config_dict, extra="allow") + + def update_config(self, **kwargs) -> T: + """Update current config with new values using kwargs. + + Args: + **kwargs: Key-value pairs where __ in keys represents nested levels. + + Returns: + Updated and validated Pydantic config instance. + """ + # Convert kwargs to dot notation and parse + dot_list = [f"{key.replace('__', '.')}={value}" for key, value in kwargs.items()] + override_config = self.parse_dot_notation(dot_list) + + # Merge with existing config + final_config = self.merge_configs(self.config_dict, override_config) + return self.config_class.model_validate(final_config, extra="allow")