mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
init
This commit is contained in:
parent
2665f31d86
commit
9983a9854d
35 changed files with 1414 additions and 3180 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,8 @@
|
|||
from .base_component import BaseComponent
|
||||
from .component_registry import ComponentRegistry, R
|
||||
|
||||
__all__ = [
|
||||
"BaseComponent",
|
||||
]
|
||||
"ComponentRegistry",
|
||||
"R",
|
||||
]
|
||||
|
|
|
|||
24
reme_cli/component/application_context.py
Normal file
24
reme_cli/component/application_context.py
Normal file
|
|
@ -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()
|
||||
}
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
10
reme_cli/component/embedding/__init__.py
Normal file
10
reme_cli/component/embedding/__init__.py
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
"""embedding"""
|
||||
|
||||
from .base_embedding_model import BaseEmbeddingModel
|
||||
from .openai_embedding_model import OpenAIEmbeddingModel
|
||||
|
||||
__all__ = [
|
||||
"BaseEmbeddingModel",
|
||||
"OpenAIEmbeddingModel",
|
||||
]
|
||||
|
||||
407
reme_cli/component/embedding/base_embedding_model.py
Normal file
407
reme_cli/component/embedding/base_embedding_model.py
Normal file
|
|
@ -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()
|
||||
56
reme_cli/component/embedding/openai_embedding_model.py
Normal file
56
reme_cli/component/embedding/openai_embedding_model.py
Normal file
|
|
@ -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]
|
||||
|
|
@ -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)
|
||||
|
|
@ -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."""
|
||||
|
|
@ -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()
|
||||
|
|
@ -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}'")
|
||||
|
|
@ -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()
|
||||
|
|
@ -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)
|
||||
|
|
@ -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()
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
14
reme_cli/component/flow/__init__.py
Normal file
14
reme_cli/component/flow/__init__.py
Normal file
|
|
@ -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)
|
||||
208
reme_cli/component/flow/base_flow.py
Normal file
208
reme_cli/component/flow/base_flow.py
Normal file
|
|
@ -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
|
||||
18
reme_cli/component/flow/cmd_flow.py
Normal file
18
reme_cli/component/flow/cmd_flow.py
Normal file
|
|
@ -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=<op_flow>` in cmd!"
|
||||
|
||||
def _build_flow(self) -> BaseOp:
|
||||
"""Parse the stored flow expression into a functional operation object."""
|
||||
return self.parse_expression(self.flow)
|
||||
37
reme_cli/component/flow/expression_flow.py
Normal file
37
reme_cli/component/flow/expression_flow.py
Normal file
|
|
@ -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,
|
||||
},
|
||||
)
|
||||
|
|
@ -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()
|
||||
|
|
@ -1,7 +1,9 @@
|
|||
"""enumeration"""
|
||||
|
||||
from .component_enum import ComponentEnum
|
||||
from .json_schema_enum import JsonSchemaEnum
|
||||
|
||||
__all__ = [
|
||||
"ComponentEnum",
|
||||
"JsonSchemaEnum",
|
||||
]
|
||||
|
|
|
|||
41
reme_cli/enumeration/json_schema_enum.py
Normal file
41
reme_cli/enumeration/json_schema_enum.py
Normal file
|
|
@ -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()
|
||||
|
|
@ -0,0 +1,7 @@
|
|||
from .application_config import ApplicationConfig
|
||||
from .base_node import BaseNode
|
||||
|
||||
__all__ = [
|
||||
"ApplicationConfig",
|
||||
"BaseNode",
|
||||
]
|
||||
20
reme_cli/schema/application_config.py
Normal file
20
reme_cli/schema/application_config.py
Normal file
|
|
@ -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)
|
||||
10
reme_cli/schema/base_node.py
Normal file
10
reme_cli/schema/base_node.py
Normal file
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
196
reme_cli/schema/tool_call.py
Normal file
196
reme_cli/schema/tool_call.py
Normal file
|
|
@ -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(),
|
||||
)
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
194
reme_cli/utils/pydantic_config_parser.py
Normal file
194
reme_cli/utils/pydantic_config_parser.py
Normal file
|
|
@ -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")
|
||||
Loading…
Add table
Reference in a new issue