diff --git a/reme4/components/base_component.py b/reme4/components/base_component.py index 9b59dcfe..ba238cb0 100644 --- a/reme4/components/base_component.py +++ b/reme4/components/base_component.py @@ -134,6 +134,11 @@ class BaseComponent(ABC): return Path.cwd() / "metadata" return self.vault_path / self.app_context.app_config.metadata_dir + @property + def component_metadata_path(self) -> Path: + """Resolved component metadata directory: vault_metadata_path / component_type.""" + return self.vault_metadata_path / self.component_type.value + def to_vault_relative(self, path: str | Path) -> str: """Return path relative to vault_path; absolute path string if outside.""" abs_path = Path(path).absolute() diff --git a/reme4/components/embedding/base_embedding_model.py b/reme4/components/embedding/base_embedding_model.py index 64add5b5..3f4356cc 100644 --- a/reme4/components/embedding/base_embedding_model.py +++ b/reme4/components/embedding/base_embedding_model.py @@ -13,9 +13,11 @@ from ..base_component import BaseComponent from ...enumeration import ComponentEnum from ...schema import EmbNode +Miss = tuple[int, str, str] # (result_index, text, cache_key) + class BaseEmbeddingModel(BaseComponent): - """Embedding model with LRU cache and disk persistence.""" + """Embedding model with LRU cache, disk persistence, and concurrent batching.""" component_type = ComponentEnum.EMBEDDING_MODEL @@ -29,6 +31,7 @@ class BaseEmbeddingModel(BaseComponent): max_batch_size: int = 10, max_input_length: int = 8192, max_cache_size: int = 10000, + max_concurrency: int = 2, enable_cache: bool = True, cache_version: str = "v1", max_retries: int = 3, @@ -43,21 +46,24 @@ class BaseEmbeddingModel(BaseComponent): self.max_batch_size = max_batch_size self.max_input_length = max_input_length self.max_cache_size = max_cache_size + self.max_concurrency = max_concurrency self.enable_cache = enable_cache self.cache_version = cache_version self.max_retries = max_retries - self._embedding_cache: OrderedDict[str, np.ndarray] = OrderedDict() + self._cache: OrderedDict[str, np.ndarray] = OrderedDict() + self._key_suffix = f"|{model_name}|{dimensions}".encode() self.is_healthy: bool = True @property def cache_path(self) -> Path: - """Disk path for the embedding cache file.""" return self.vault_metadata_path / "embedding_cache" / f"{self.name}_{self.cache_version}.npz" async def _start(self) -> None: - """Load cache from disk on startup.""" await self.load() + async def _close(self) -> None: + await self.dump() + async def health_check(self, timeout: float = 2.0) -> bool: """Probe the provider; sets and returns is_healthy.""" tag = f"[EMBEDDING HEALTH CHECK] name={self.name} model={self.model_name}" @@ -75,71 +81,21 @@ class BaseEmbeddingModel(BaseComponent): self.logger.error(f"{tag} -> FAIL {type(e).__name__}: {e}") return self.is_healthy - async def _close(self) -> None: - """Persist cache to disk on shutdown.""" - await self.dump() - # -- Public API -- async def get_embedding(self, input_text: str, **kwargs) -> np.ndarray | None: - """Get embedding for a single text.""" results = await self.get_embeddings([input_text], **kwargs) return results[0] if results else None async def get_embeddings(self, input_text: list[str], **kwargs) -> list[np.ndarray | None]: - """Get embeddings for a list of texts, with caching and batching.""" - truncated = [t[: self.max_input_length] for t in input_text] - results: list[np.ndarray | None] = [None] * len(truncated) - to_compute: list[tuple[int, str]] = [] - - # Split into cache hits and misses - for idx, text in enumerate(truncated): - cached = self._get_from_cache(text) - if cached is not None: - results[idx] = cached - else: - to_compute.append((idx, text)) - - # Batch-compute misses with retry - if to_compute: - for i in range(0, len(to_compute), self.max_batch_size): - batch = to_compute[i : i + self.max_batch_size] - indices = [idx for idx, _ in batch] - texts = [text for _, text in batch] - - embeddings = None - for attempt in range(self.max_retries): - try: - embeddings = await self._get_embeddings(texts, **kwargs) - if embeddings and len(embeddings) == len(texts): - break - except (TimeoutError, ConnectionError, OSError): - if attempt < self.max_retries - 1: - await asyncio.sleep(2**attempt) - except Exception: - self.logger.exception("Embedding request failed") - break - - if not embeddings or len(embeddings) != len(texts): - continue - - # Normalize dimensions and cache - for orig_idx, text, emb in zip(indices, texts, embeddings): - if emb is None: - continue - emb_array = np.asarray(emb, dtype=np.float16) - if len(emb_array) != self.dimensions: - if len(emb_array) < self.dimensions: - emb_array = np.pad(emb_array, (0, self.dimensions - len(emb_array))) - else: - emb_array = emb_array[: self.dimensions] - results[orig_idx] = emb_array - self._put_to_cache(text, emb_array) - + """Get embeddings for texts. Cache hits return immediately; misses run concurrently.""" + texts = [self._truncate(t) for t in input_text] + results, misses = self._partition_by_cache(texts) + if misses: + await self._fill_misses(misses, results, **kwargs) return results async def get_node_embeddings(self, nodes: list[EmbNode], **kwargs) -> list[EmbNode]: - """Compute and assign embeddings for EmbNode objects.""" embeddings = await self.get_embeddings([n.text for n in nodes], **kwargs) if len(embeddings) == len(nodes): for node, vec in zip(nodes, embeddings): @@ -151,64 +107,134 @@ class BaseEmbeddingModel(BaseComponent): async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float] | None]: """Get raw embeddings from the underlying provider.""" - # -- Cache Operations -- + # -- Batching -- - def _get_from_cache(self, text: str) -> np.ndarray | None: - """Lookup text in LRU cache, promoting on hit.""" - if not self.enable_cache: - return None - key = self._get_cache_key(text) - if key not in self._embedding_cache: - return None - self._embedding_cache.move_to_end(key) - return self._embedding_cache[key] + def _truncate(self, text: str) -> str: + return text if len(text) <= self.max_input_length else text[: self.max_input_length] - def _put_to_cache(self, text: str, embedding: np.ndarray) -> None: - """Insert into LRU cache, evicting oldest if full.""" + def _partition_by_cache(self, texts: list[str]) -> tuple[list[np.ndarray | None], list[Miss]]: + """Split texts into pre-filled results (hits) and a miss list to compute.""" + results: list[np.ndarray | None] = [None] * len(texts) + misses: list[Miss] = [] + for idx, text in enumerate(texts): + key = self._cache_key(text) + hit = self._cache_get(key) + if hit is not None: + results[idx] = hit + else: + misses.append((idx, text, key)) + return results, misses + + async def _fill_misses(self, misses: list[Miss], results: list[np.ndarray | None], **kwargs) -> None: + """Compute miss embeddings in concurrent batches and write into results + cache.""" + size = self.max_batch_size + batches = [misses[i : i + size] for i in range(0, len(misses), size)] + sem = asyncio.Semaphore(self.max_concurrency) + + async def run(batch: list[Miss]) -> list[tuple[int, str, np.ndarray]]: + async with sem: + return await self._compute_batch(batch, **kwargs) + + for done in await asyncio.gather(*(run(b) for b in batches)): + for idx, key, emb in done: + results[idx] = emb + self._cache_put(key, emb) + + async def _compute_batch(self, batch: list[Miss], **kwargs) -> list[tuple[int, str, np.ndarray]]: + """Call provider for one batch with retry; returns [(idx, key, embedding)].""" + texts = [text for _, text, _ in batch] + embeddings = await self._call_with_retry(texts, **kwargs) + if not embeddings or len(embeddings) != len(texts): + return [] + out: list[tuple[int, str, np.ndarray]] = [] + for (idx, _text, key), raw in zip(batch, embeddings): + if raw is None: + continue + emb = self._normalize_dim(np.asarray(raw, dtype=np.float16)) + out.append((idx, key, emb)) + return out + + async def _call_with_retry(self, texts: list[str], **kwargs) -> list[list[float] | None] | None: + """Call provider with exponential backoff on transient errors.""" + for attempt in range(self.max_retries): + try: + result = await self._get_embeddings(texts, **kwargs) + if result and len(result) == len(texts): + return result + except (TimeoutError, ConnectionError, OSError): + if attempt < self.max_retries - 1: + await asyncio.sleep(2**attempt) + except Exception: + self.logger.exception("Embedding request failed") + return None + return None + + def _normalize_dim(self, emb: np.ndarray) -> np.ndarray: + if len(emb) == self.dimensions: + return emb + if len(emb) < self.dimensions: + return np.pad(emb, (0, self.dimensions - len(emb))) + return emb[: self.dimensions] + + # -- Cache -- + + def _cache_key(self, text: str) -> str: + return hashlib.sha256(text.encode() + self._key_suffix).hexdigest() + + def _cache_get(self, key: str) -> np.ndarray | None: + if not self.enable_cache or key not in self._cache: + return None + self._cache.move_to_end(key) + return self._cache[key] + + def _cache_put(self, key: str, embedding: np.ndarray) -> None: if not self.enable_cache or self.max_cache_size <= 0 or len(embedding) != self.dimensions: return - key = self._get_cache_key(text) - if len(self._embedding_cache) >= self.max_cache_size and key not in self._embedding_cache: - self._embedding_cache.popitem(last=False) - self._embedding_cache[key] = embedding - self._embedding_cache.move_to_end(key) + cache = self._cache + if key in cache: + cache.move_to_end(key) + cache[key] = embedding + return + if len(cache) >= self.max_cache_size: + cache.popitem(last=False) + cache[key] = embedding - def _get_cache_key(self, text: str) -> str: - """Generate cache key from text, model name, and dimensions.""" - return hashlib.sha256(f"{text}|{self.model_name}|{self.dimensions}".encode()).hexdigest() - - # -- Cache Persistence -- + # -- Persistence -- async def load(self) -> None: - """Load cached embeddings from disk (npz format); replaces in-memory cache.""" - self._embedding_cache.clear() + """Load cached embeddings from disk (npz); replaces in-memory cache.""" + self._cache.clear() if not self.enable_cache or not self.cache_path.exists(): return + await asyncio.to_thread(self._load_sync) + def _load_sync(self) -> None: try: data = np.load(self.cache_path) except Exception: self.logger.exception("Failed to load embedding cache, removing") self.cache_path.unlink(missing_ok=True) return - for key, emb in zip(data["keys"], data["embeddings"]): if len(emb) != self.dimensions: continue - if len(self._embedding_cache) >= self.max_cache_size: + if len(self._cache) >= self.max_cache_size: break - self._embedding_cache[str(key)] = emb.astype(np.float16) - self.logger.info(f"Loaded {len(self._embedding_cache)} embeddings from {self.cache_path}") + self._cache[str(key)] = emb.astype(np.float16) + self.logger.info(f"Loaded {len(self._cache)} embeddings from {self.cache_path}") async def dump(self) -> None: - """Persist in-memory cache to disk (npz format).""" - if not self.enable_cache or not self._embedding_cache: + """Persist in-memory cache to disk (npz).""" + if not self.enable_cache or not self._cache: return + await asyncio.to_thread(self._dump_sync) + + def _dump_sync(self) -> None: self.cache_path.parent.mkdir(parents=True, exist_ok=True) - keys = list(self._embedding_cache.keys()) - embeddings = np.stack(list(self._embedding_cache.values())) + keys = np.array(list(self._cache.keys()), dtype=str) + embeddings = np.stack(list(self._cache.values())) try: - np.savez(self.cache_path, keys=np.array(keys, dtype=str), embeddings=embeddings) - self.logger.info(f"Saved {len(self._embedding_cache)} embeddings to {self.cache_path}") + np.savez(self.cache_path, keys=keys, embeddings=embeddings) + self.logger.info(f"Saved {len(self._cache)} embeddings to {self.cache_path}") except Exception: self.logger.exception("Failed to save embedding cache") diff --git a/reme4/components/file_catalog/base_file_catalog.py b/reme4/components/file_catalog/base_file_catalog.py index 390fe619..661af8d9 100644 --- a/reme4/components/file_catalog/base_file_catalog.py +++ b/reme4/components/file_catalog/base_file_catalog.py @@ -1,7 +1,4 @@ -"""Abstract base for file-catalog backends.""" - from abc import abstractmethod -from pathlib import Path from ..base_component import BaseComponent from ...enumeration import ComponentEnum @@ -9,22 +6,10 @@ from ...schema import FileNode class BaseFileCatalog(BaseComponent): - """Abstract base for file-catalog backends. - - A catalog records FileNode entries keyed by path — a lightweight - counterpart to FileStore that drops chunk/embedding/keyword/link - machinery and exposes only node upsert / delete / lookup. - """ + """File-catalog backend recording FileNode entries keyed by path.""" component_type = ComponentEnum.FILE_CATALOG - def __init__(self, **kwargs): - super().__init__(**kwargs) - self.catalog_path: Path = self.vault_metadata_path / self.component_type.value - self.catalog_path.mkdir(parents=True, exist_ok=True) - - # -- Lifecycle --------------------------------------------------------- - async def _start(self) -> None: await super()._start() await self.load() @@ -34,12 +19,10 @@ class BaseFileCatalog(BaseComponent): await super()._close() async def load(self) -> None: - """Load persisted state. No-op for backends without local files.""" + """Load persisted state. No-op without local files.""" async def dump(self) -> None: - """Persist state. No-op for backends without local files.""" - - # -- CRUD -------------------------------------------------------------- + """Persist state. No-op without local files.""" @abstractmethod async def upsert(self, nodes: list[FileNode]) -> None: @@ -51,4 +34,4 @@ class BaseFileCatalog(BaseComponent): @abstractmethod async def get_nodes(self, paths: list[str] | None = None) -> list[FileNode]: - """Return nodes by paths; None = all nodes; [] = []; missing paths are skipped.""" + """Return nodes by paths; None = all; missing paths are skipped.""" diff --git a/reme4/components/file_catalog/local_file_catalog.py b/reme4/components/file_catalog/local_file_catalog.py index e0563640..9f35364d 100644 --- a/reme4/components/file_catalog/local_file_catalog.py +++ b/reme4/components/file_catalog/local_file_catalog.py @@ -1,5 +1,3 @@ -"""In-memory file catalog with JSONL persistence.""" - import aiofiles from .base_file_catalog import BaseFileCatalog @@ -9,34 +7,27 @@ from ...schema import FileNode @R.register("local") class LocalFileCatalog(BaseFileCatalog): - """Dict-backed file catalog persisted as JSONL on close.""" + """Dict-backed catalog persisted as JSONL.""" def __init__(self, encoding: str = "utf-8", **kwargs): super().__init__(**kwargs) self.encoding = encoding self._nodes: dict[str, FileNode] = {} - self._catalog_file = self.catalog_path / f"{self.name}.jsonl" + self.component_metadata_path.mkdir(parents=True, exist_ok=True) + self._catalog_file = self.component_metadata_path / f"{self.name}.jsonl" async def load(self) -> None: if not self._catalog_file.exists(): return try: - async with aiofiles.open(self._catalog_file, encoding=self.encoding) as f: - async for line in f: - line = line.strip() - if line: - node = FileNode.model_validate_json(line) - self._nodes[node.path] = node + await self._read_jsonl() self.logger.info(f"Loaded {len(self._nodes)} nodes from {self._catalog_file}") except Exception as e: self.logger.exception(f"Failed to load {self._catalog_file}: {e}") async def dump(self) -> None: try: - tmp = self._catalog_file.with_suffix(".tmp") - async with aiofiles.open(tmp, "w", encoding=self.encoding) as f: - await f.write("\n".join(n.model_dump_json() for n in self._nodes.values())) - tmp.replace(self._catalog_file) + await self._write_jsonl() self.logger.info(f"Saved {len(self._nodes)} nodes to {self._catalog_file}") except Exception as e: self.logger.exception(f"Failed to write {self._catalog_file}: {e}") @@ -54,3 +45,16 @@ class LocalFileCatalog(BaseFileCatalog): if paths is None: return list(self._nodes.values()) return [self._nodes[p] for p in paths if p in self._nodes] + + async def _read_jsonl(self) -> None: + async with aiofiles.open(self._catalog_file, encoding=self.encoding) as f: + async for line in f: + if stripped := line.strip(): + node = FileNode.model_validate_json(stripped) + self._nodes[node.path] = node + + async def _write_jsonl(self) -> None: + tmp = self._catalog_file.with_suffix(".tmp") + async with aiofiles.open(tmp, "w", encoding=self.encoding) as f: + await f.write("\n".join(n.model_dump_json() for n in self._nodes.values())) + tmp.replace(self._catalog_file)