This commit is contained in:
jinli.yl 2026-04-09 10:38:04 +08:00
parent 2665f31d86
commit 9983a9854d
35 changed files with 1414 additions and 3180 deletions

View file

@ -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

View file

@ -1,5 +1,8 @@
from .base_component import BaseComponent
from .component_registry import ComponentRegistry, R
__all__ = [
"BaseComponent",
]
"ComponentRegistry",
"R",
]

View 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()
}

View file

@ -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",
]

View file

@ -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",
]

View file

@ -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

View file

@ -0,0 +1,10 @@
"""embedding"""
from .base_embedding_model import BaseEmbeddingModel
from .openai_embedding_model import OpenAIEmbeddingModel
__all__ = [
"BaseEmbeddingModel",
"OpenAIEmbeddingModel",
]

View 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()

View 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]

View file

@ -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)

View file

@ -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."""

View file

@ -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()

View file

@ -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}'")

View file

@ -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()

View file

@ -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)

View file

@ -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()

View file

@ -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

View file

@ -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

View 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)

View 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

View 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)

View 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,
},
)

View file

@ -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()

View file

@ -1,7 +1,9 @@
"""enumeration"""
from .component_enum import ComponentEnum
from .json_schema_enum import JsonSchemaEnum
__all__ = [
"ComponentEnum",
"JsonSchemaEnum",
]

View 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()

View file

@ -0,0 +1,7 @@
from .application_config import ApplicationConfig
from .base_node import BaseNode
__all__ = [
"ApplicationConfig",
"BaseNode",
]

View 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)

View 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)

View file

@ -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)

View 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(),
)

View file

@ -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",
]

View file

@ -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

View 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")