fix(file_store): make embedding recovery race-safe

This commit is contained in:
jinli.yl 2026-08-20 20:55:16 +08:00
parent 097f6c8270
commit 7832a469b8
6 changed files with 224 additions and 15 deletions

View file

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

View file

@ -109,6 +109,9 @@ file_store:
Embedding store 可通过 `health_check_timeout` 配置启动探测。临时失败只会跳过本次向量回填BM25 仍可使用;
后续真实请求成功后会自动恢复缺失向量的回填。
已经完成真实服务验证的嵌入式集成可以调用 `resume_embedding(verified=True)`。切换 Embedding 向量空间时应同时传入
`rebuild=True`ReMe 会先使旧向量失效,再串行后台重建,并在新向量安全持久化前暂停向量搜索。
## 怎么搜索
`search` Job 也是在 `default.yaml` 中配置:

View file

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

View file

@ -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] = []

View file

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

View file

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