diff --git a/docs/en/memory_search.md b/docs/en/memory_search.md index c4fc8437..bdc4692e 100644 --- a/docs/en/memory_search.md +++ b/docs/en/memory_search.md @@ -121,6 +121,10 @@ The embedding store accepts `health_check_timeout` for its startup probe. A temp backfill while keeping BM25 available; a later successful provider request resumes the missing-vector backfill automatically. +Embedded integrations that have already verified a provider can call `resume_embedding(verified=True)`. When changing +the embedding vector space, pass `rebuild=True`; persisted vectors are invalidated before a serial background rebuild, +and vector search remains unavailable until the rebuilt vectors are safely persisted. + ## How to Search The `search` Job is also configured in `default.yaml`: diff --git a/docs/zh/memory_search.md b/docs/zh/memory_search.md index f65900e6..67fdac8c 100644 --- a/docs/zh/memory_search.md +++ b/docs/zh/memory_search.md @@ -109,6 +109,9 @@ file_store: Embedding store 可通过 `health_check_timeout` 配置启动探测。临时失败只会跳过本次向量回填,BM25 仍可使用; 后续真实请求成功后会自动恢复缺失向量的回填。 +已经完成真实服务验证的嵌入式集成可以调用 `resume_embedding(verified=True)`。切换 Embedding 向量空间时应同时传入 +`rebuild=True`;ReMe 会先使旧向量失效,再串行后台重建,并在新向量安全持久化前暂停向量搜索。 + ## 怎么搜索 `search` Job 也是在 `default.yaml` 中配置: diff --git a/reme/components/file_store/faiss_local_file_store.py b/reme/components/file_store/faiss_local_file_store.py index 48174e4e..2e02bfe2 100644 --- a/reme/components/file_store/faiss_local_file_store.py +++ b/reme/components/file_store/faiss_local_file_store.py @@ -252,6 +252,11 @@ class FaissLocalFileStore(LocalFileStore): self._add_to_index([c.id for c in to_add], vectors) self._compact_if_needed() + async def _reset_vector_index(self) -> None: + """Discard all vectors before rebuilding a changed vector space.""" + await self._stop_reindex_worker() + self._rebuild_index() + # -- async reindex ---------------------------------------------------- def _submit_reindex(self) -> None: @@ -626,8 +631,9 @@ class FaissLocalFileStore(LocalFileStore): async def vector_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]: index_empty = self._faiss_index is None or self._faiss_index.ntotal == 0 + embedding_unavailable = self.embedding_store is None or self._embedding_rebuild_pending if ( - self.embedding_store is None + embedding_unavailable or not query or limit <= 0 or (index_empty and getattr(self.embedding_store, "is_healthy", True)) @@ -646,7 +652,7 @@ class FaissLocalFileStore(LocalFileStore): f"search: query embedding dimension {len(query_embedding)} != {self.embedding_store.dimensions}", ) return [] - self._recover_after_real_request(was_healthy) + await self._recover_after_real_request(was_healthy) # get_embedding above yielded control; a concurrent clear() drops the # index to None once embedding is disabled, and a reindex may have swapped diff --git a/reme/components/file_store/local_file_store.py b/reme/components/file_store/local_file_store.py index 6a5294c0..e87ca36e 100644 --- a/reme/components/file_store/local_file_store.py +++ b/reme/components/file_store/local_file_store.py @@ -67,6 +67,8 @@ class LocalFileStore(BaseFileStore): self.file_chunks: dict[str, FileChunk] = {} self.chunks_path = self.component_metadata_path / f"file_chunks_{self.name}_{self.store_version}.jsonl.zst" self._embedding_backfill_task: asyncio.Task | None = None + self._embedding_backfill_pending: tuple[bool, bool] | None = None + self._embedding_rebuild_pending = False self._closing = False # -- lifecycle ------------------------------------------------------------ @@ -109,24 +111,31 @@ class LocalFileStore(BaseFileStore): self.embedding_store.is_healthy = False self.logger.error(f"{self.name}: embedding unavailable, {reason}; keyword search remains active") - def _recover_after_real_request(self, was_healthy: bool) -> None: + async def _recover_after_real_request(self, was_healthy: bool) -> None: """Schedule repair when a real, non-cache provider request recovers.""" if self.embedding_store is None or was_healthy or not getattr(self.embedding_store, "is_healthy", True): return self.logger.info(f"{self.name}: embedding provider recovered; scheduling missing-vector backfill") - self._start_embedding_backfill(skip_health_check=True) + await self.resume_embedding(verified=True) - async def resume_embedding(self, *, verified: bool = False) -> bool: + async def resume_embedding(self, *, verified: bool = False, rebuild: bool = False) -> bool: """Resume a configured provider and schedule a deduplicated repair. Embedded applications may pass ``verified=True`` after they have already - made a successful real provider request, avoiding a redundant ping. + made a successful real provider request, avoiding a redundant ping. Pass + ``rebuild=True`` when the active vector space changed; existing vectors + are derived data and are discarded before a full background rebuild. """ if self.embedding_store is None or self._closing: return False if verified: self.embedding_store.is_healthy = True - self._start_embedding_backfill(skip_health_check=verified) + if rebuild: + await self._prepare_embedding_rebuild() + if not self.file_chunks: + self._embedding_rebuild_pending = False + return True + self._start_embedding_backfill(skip_health_check=verified, rebuild=rebuild) return True def _embedding_dim_matches(self, embedding: np.ndarray | None) -> bool: @@ -251,7 +260,7 @@ class LocalFileStore(BaseFileStore): return self._drop_stale_embeddings(self.file_chunks.values(), "load") - def _start_embedding_backfill(self, *, skip_health_check: bool = False) -> None: + def _start_embedding_backfill(self, *, skip_health_check: bool = False, rebuild: bool = False) -> None: """Schedule startup embedding repair without delaying component readiness.""" started_at = time.monotonic() if self._closing: @@ -270,13 +279,18 @@ class LocalFileStore(BaseFileStore): ) return if self._embedding_backfill_task is not None and not self._embedding_backfill_task.done(): + if skip_health_check: + pending_rebuild = rebuild or bool( + self._embedding_backfill_pending and self._embedding_backfill_pending[1], + ) + self._embedding_backfill_pending = (True, pending_rebuild) self.logger.info( f"{self.name}: embedding backfill scheduling skipped: reason=already_running, " f"elapsed={time.monotonic() - started_at:.3f}s", ) return self._embedding_backfill_task = asyncio.create_task( - self._backfill_missing_embeddings(skip_health_check=skip_health_check), + self._run_embedding_backfill(skip_health_check=skip_health_check, rebuild=rebuild), name=f"embedding-backfill:{self.name}", ) self.logger.info( @@ -284,10 +298,37 @@ class LocalFileStore(BaseFileStore): f"elapsed={time.monotonic() - started_at:.3f}s", ) + async def _run_embedding_backfill(self, *, skip_health_check: bool, rebuild: bool) -> None: + """Run one repair and honor a verified request queued behind it.""" + current_task = asyncio.current_task() + try: + if rebuild: + # A task that was already running when rebuild was requested + # may have written a stale provider result after the first + # invalidation. Clear once more at the queue boundary. + await self._prepare_embedding_rebuild() + await self._backfill_missing_embeddings(skip_health_check=skip_health_check) + finally: + if self._embedding_backfill_task is current_task: + self._embedding_backfill_task = None + pending = self._embedding_backfill_pending + self._embedding_backfill_pending = None + if pending is not None and not self._closing and self.embedding_store is not None: + pending_verified, pending_rebuild = pending + if pending_verified: + self.embedding_store.is_healthy = True + if pending_rebuild: + self._embedding_rebuild_pending = True + self._start_embedding_backfill( + skip_health_check=pending_verified, + rebuild=pending_rebuild, + ) + async def _cancel_embedding_backfill(self) -> None: """Cancel and collect the startup repair task during component shutdown.""" task = self._embedding_backfill_task self._embedding_backfill_task = None + self._embedding_backfill_pending = None if task is None: return if not task.done(): @@ -329,6 +370,10 @@ class LocalFileStore(BaseFileStore): f"missing={len(missing)}, elapsed={time.monotonic() - scan_started_at:.3f}s", ) if not missing: + if self._embedding_rebuild_pending: + await self._after_embedding_backfill() + await self.dump() + self._embedding_rebuild_pending = False self.logger.info( f"{self.name}: embedding backfill complete: filled=0/0, " f"elapsed={time.monotonic() - started_at:.3f}s", @@ -374,6 +419,7 @@ class LocalFileStore(BaseFileStore): f"{total}, elapsed={elapsed:.2f}s", ) raise + except Exception as e: self._mark_embedding_unhealthy(f"backfill: {type(e).__name__}: {e}") elapsed = time.monotonic() - started_at @@ -388,13 +434,25 @@ class LocalFileStore(BaseFileStore): self.logger.info( f"{self.name}: embedding backfill complete: filled={filled}/{total}, elapsed={elapsed:.2f}s", ) - if filled: + if filled or self._embedding_rebuild_pending: try: await self._after_embedding_backfill() await self.dump() + self._embedding_rebuild_pending = False except Exception: self.logger.exception(f"{self.name}: failed to persist completed embedding backfill") + async def _prepare_embedding_rebuild(self) -> None: + """Invalidate and persist vectors from the previous vector space.""" + self._embedding_rebuild_pending = True + for chunk in self.file_chunks.values(): + chunk.embedding = None + await self._reset_vector_index() + await self.dump() + + async def _reset_vector_index(self) -> None: + """Drop a derived vector index before rebuilding a changed vector space.""" + async def _after_embedding_backfill(self) -> None: """Backend hook for refreshing derived vector indexes after backfill.""" @@ -574,7 +632,7 @@ class LocalFileStore(BaseFileStore): return self._drop_stale_embeddings(chunks, "upsert") if any(chunk.embedding is not None for chunk in chunks): - self._recover_after_real_request(was_healthy) + await self._recover_after_real_request(was_healthy) async def delete(self, path: str | list[str]) -> None: assert self.file_graph is not None @@ -630,7 +688,7 @@ class LocalFileStore(BaseFileStore): # -- search --------------------------------------------------------------- async def vector_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]: - if self.embedding_store is None or not query or limit <= 0: + if self.embedding_store is None or self._embedding_rebuild_pending or not query or limit <= 0: return [] was_healthy = bool(getattr(self.embedding_store, "is_healthy", True)) @@ -646,7 +704,7 @@ class LocalFileStore(BaseFileStore): f"search: query embedding dimension {len(query_embedding)} != {self.embedding_store.dimensions}", ) return [] - self._recover_after_real_request(was_healthy) + await self._recover_after_real_request(was_healthy) top: list[tuple[float, int, FileChunk]] = [] candidates: list[FileChunk] = [] diff --git a/reme/components/file_store/zvec_local_file_store.py b/reme/components/file_store/zvec_local_file_store.py index 652795c0..4b234a3e 100644 --- a/reme/components/file_store/zvec_local_file_store.py +++ b/reme/components/file_store/zvec_local_file_store.py @@ -176,6 +176,10 @@ class ZvecLocalFileStore(LocalFileStore): ] self._upsert_docs(to_add) + async def _reset_vector_index(self) -> None: + """Discard all vectors before rebuilding a changed vector space.""" + self._collection = self._create_collection() + # -- maintenance ------------------------------------------------------ async def optimize_index(self) -> None: @@ -412,7 +416,7 @@ class ZvecLocalFileStore(LocalFileStore): # -- search ----------------------------------------------------------- async def vector_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]: - if self.embedding_store is None or not query or limit <= 0: + if self.embedding_store is None or self._embedding_rebuild_pending or not query or limit <= 0: return [] index_empty = self._collection is None or not self._indexed_ids if index_empty and getattr(self.embedding_store, "is_healthy", True): @@ -430,7 +434,7 @@ class ZvecLocalFileStore(LocalFileStore): f"search: query embedding dimension {len(query_embedding)} != {self.embedding_store.dimensions}", ) return [] - self._recover_after_real_request(was_healthy) + await self._recover_after_real_request(was_healthy) # get_embedding above yielded control; a concurrent clear() may have # swapped or dropped the collection. Re-read before dereferencing. diff --git a/tests/unit/test_file_store_consistency.py b/tests/unit/test_file_store_consistency.py index ba011912..83a21096 100644 --- a/tests/unit/test_file_store_consistency.py +++ b/tests/unit/test_file_store_consistency.py @@ -68,6 +68,7 @@ class CountingFakeEmbeddingStore(FakeEmbeddingStore): def __init__(self): self.node_embedding_calls: list[list[str]] = [] + self.is_healthy = True async def get_node_embeddings(self, nodes: list[FileChunk], **_kwargs) -> list[FileChunk]: self.node_embedding_calls.append([node.id for node in nodes]) @@ -130,6 +131,44 @@ class BlockingEmbeddingStore(FakeEmbeddingStore): return await super().get_node_embeddings(nodes, **kwargs) +class CancellationResistantHealthStore(CountingFakeEmbeddingStore): + """Startup probe that completes stale after cancellation is requested.""" + + def __init__(self): + super().__init__() + self.is_healthy = True + self.health_started = asyncio.Event() + self.release_health = asyncio.Event() + + async def health_check(self, _timeout: float = 2.0) -> bool: + self.health_started.set() + try: + await self.release_health.wait() + except asyncio.CancelledError: + await self.release_health.wait() + self.is_healthy = False + return False + + +class DelayedOldVectorStore(CountingFakeEmbeddingStore): + """First batch returns an old-space vector after rebuild was requested.""" + + def __init__(self): + super().__init__() + self.first_batch_started = asyncio.Event() + self.release_first_batch = asyncio.Event() + + async def get_node_embeddings(self, nodes: list[FileChunk], **_kwargs) -> list[FileChunk]: + self.node_embedding_calls.append([node.id for node in nodes]) + if len(self.node_embedding_calls) == 1: + self.first_batch_started.set() + await self.release_first_batch.wait() + for chunk_node in nodes: + chunk_node.embedding = np.array([0.0, 1.0], dtype=np.float16) + return nodes + return await FakeEmbeddingStore.get_node_embeddings(self, nodes) + + class WrongDimEmbeddingStore(FakeEmbeddingStore): """Fake embedding store that returns vectors with the wrong dimension.""" @@ -653,6 +692,101 @@ def test_load_skips_backfill_when_embedding_health_check_fails(): run(go()) +def test_verified_resume_supersedes_inflight_startup_health_check(): + """A stale startup probe cannot consume or overwrite verified recovery.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + store = _new_local_store("t_embedding_verified_resume_race") + await store.start() + await set_chunks_with_graph(store, {"a": chunk("a", "a.md", "alpha text")}) + fake = CancellationResistantHealthStore() + store.embedding_store = fake + store._start_embedding_backfill() + startup_task = store._embedding_backfill_task + await fake.health_started.wait() + + recovery = asyncio.create_task(store.resume_embedding(verified=True)) + await asyncio.sleep(0) + assert await recovery is True + assert store._embedding_backfill_pending == (True, False) + fake.release_health.set() + + await startup_task + assert store._embedding_backfill_task is not startup_task + if store._embedding_backfill_task is not None: + await store._embedding_backfill_task + assert fake.is_healthy is True + assert fake.node_embedding_calls == [["a"]] + assert store.file_chunks["a"].embedding.tolist() == [1.0, 0.0] + await store.close() + + run(go()) + + +@pytest.mark.parametrize("store_factory", [_new_local_store, _new_faiss_store, _new_zvec_store]) +def test_verified_rebuild_discards_same_dimension_vectors_before_backfill(store_factory): + """A changed vector space never searches compatible-shaped stale vectors.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + store = store_factory("t_embedding_verified_rebuild") + await store.start() + stale = chunk("a", "a.md", "alpha text") + stale.embedding = np.array([0.0, 1.0], dtype=np.float16) + await set_chunks_with_graph(store, {"a": stale}) + fake = CountingFakeEmbeddingStore() + fake.is_healthy = False + store.embedding_store = fake + if isinstance(store, FaissLocalFileStore): + store._rebuild_index() + elif isinstance(store, ZvecLocalFileStore): + store._rebuild_collection() + + assert await store.resume_embedding(verified=True, rebuild=True) is True + assert store._embedding_rebuild_pending is True + assert store.file_chunks["a"].embedding is None + assert await store.vector_search("alpha", 5, {}) == [] + + await store._embedding_backfill_task + assert store._embedding_rebuild_pending is False + assert fake.node_embedding_calls == [["a"]] + assert store.file_chunks["a"].embedding.tolist() == [1.0, 0.0] + assert [item.id for item in await store.vector_search("alpha", 5, {})] == ["a"] + await store.close() + + run(go()) + + +def test_verified_rebuild_discards_late_result_from_previous_vector_space(): + """A queued rebuild clears old-space vectors written by an in-flight batch.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + store = _new_local_store("t_embedding_verified_rebuild_race") + await store.start() + await set_chunks_with_graph(store, {"a": chunk("a", "a.md", "alpha text")}) + fake = DelayedOldVectorStore() + store.embedding_store = fake + store._start_embedding_backfill(skip_health_check=True) + old_task = store._embedding_backfill_task + await fake.first_batch_started.wait() + + assert await store.resume_embedding(verified=True, rebuild=True) is True + assert store._embedding_backfill_pending == (True, True) + fake.release_first_batch.set() + + await old_task + if store._embedding_backfill_task is not None: + await store._embedding_backfill_task + assert fake.node_embedding_calls == [["a"], ["a"]] + assert store.file_chunks["a"].embedding.tolist() == [1.0, 0.0] + assert store._embedding_rebuild_pending is False + await store.close() + + run(go()) + + @pytest.mark.parametrize("store_factory", [_new_local_store, _new_faiss_store, _new_zvec_store]) def test_search_recovery_schedules_backfill_without_another_health_check(store_factory): """A successful real search request repairs historical missing vectors."""