This commit is contained in:
jinli.yl 2026-05-26 16:03:55 +08:00
parent 6ad02f0a8f
commit 9f90c448f1
4 changed files with 146 additions and 128 deletions

View file

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

View file

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

View file

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

View file

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