From c4dca4d4ed09854d81b1ace1f8f29f659965bf23 Mon Sep 17 00:00:00 2001 From: omar-nahhas Date: Wed, 29 Jul 2026 13:51:06 +0100 Subject: [PATCH 1/2] fix(_memory): make interrupted reindex resumable, verified by model identity Memory.initialize() used to embed every document in one pass when a reindex was needed (embedding model changed) and only persist the result at the very end. Interrupting that pass -- a crash, a restart, a slow or shared embeddings backend timing out mid-batch -- threw away every already-embedded document; the next attempt started over from zero. On a large memory store against a slow backend this can cost hours of redone work every time the process is interrupted. Reindexing now embeds in batches (Memory._REBUILD_BATCH_SIZE) and checkpoints to disk after each one (index.rebuilding.faiss/.pkl), so an interrupted rebuild resumes from where it left off instead of restarting. The checkpoint is not trusted blindly: it's tagged with the embedding model it was built under (index.rebuilding.json), and a resume is only used if that matches the model we're about to (re)index with. If the target model changed again while a rebuild was interrupted, the stale checkpoint is discarded and a fresh rebuild starts -- otherwise resuming would add new-model vectors into an old-model FAISS index, which fails a dimension assertion on the very next batch. Adds tests/test_memory_rebuild_resume.py covering the checkpoint save/load/clear helpers directly and Memory.initialize() end-to-end for both the resume-under-same-model and discard-under-changed-model cases. --- plugins/_memory/helpers/memory.py | 163 +++++++++++- tests/test_memory_rebuild_resume.py | 384 ++++++++++++++++++++++++++++ 2 files changed, 536 insertions(+), 11 deletions(-) create mode 100644 tests/test_memory_rebuild_resume.py diff --git a/plugins/_memory/helpers/memory.py b/plugins/_memory/helpers/memory.py index 364bfe76cd..d92cff0c27 100644 --- a/plugins/_memory/helpers/memory.py +++ b/plugins/_memory/helpers/memory.py @@ -60,6 +60,113 @@ class Area(Enum): index: dict[str, "MyFaiss"] = {} + # Name (without extension) of the on-disk checkpoint a reindex writes + # after every batch, so an interrupted rebuild can resume instead of + # re-embedding everything from scratch. + _REBUILD_INDEX_NAME = "index.rebuilding" + # Documents embedded per batch during a reindex. Keeping this modest + # bounds how much work a crash mid-batch can lose, and how large a single + # embeddings-API call gets (a large batch of long documents can make one + # call take a very long time on a slow/shared backend). + _REBUILD_BATCH_SIZE = 20 + + @staticmethod + def _rebuild_meta_path(db_dir: str) -> str: + return files.get_abs_path(db_dir, f"{Memory._REBUILD_INDEX_NAME}.json") + + @staticmethod + def _load_rebuild_checkpoint( + db_dir: str, + model_config: "models.ModelConfig", + embedder: Any, + ) -> tuple["MyFaiss | None", set[str]]: + """Load a partially-rebuilt index left behind by an interrupted + reindex, but only if it was built with the SAME embedding model as + the one we're about to (re)index with. + + A stale checkpoint left over from a *different* model (e.g. the + model was switched again while the previous rebuild was interrupted) + holds vectors of a possibly different dimension/semantics and must + not be reused: resuming into it would raise a FAISS dimension + assertion the moment a new document is added, rather than making + progress. Returns (db, resumed_ids); db is None (and resumed_ids + empty) if there is nothing usable to resume from. + """ + name = Memory._REBUILD_INDEX_NAME + if not ( + files.exists(db_dir, f"{name}.faiss") and files.exists(db_dir, f"{name}.pkl") + ): + return None, set() + + meta_path = Memory._rebuild_meta_path(db_dir) + if not os.path.exists(meta_path): + PrintStyle.standard( + "Found a partial memory rebuild with no model marker -- discarding it and starting fresh" + ) + return None, set() + try: + meta = json.loads(files.read_file(meta_path)) + except Exception: + PrintStyle.standard( + "Partial memory rebuild marker is unreadable -- discarding it and starting fresh" + ) + return None, set() + + if ( + meta.get("model_provider") != model_config.provider + or meta.get("model_name") != model_config.name + ): + PrintStyle.standard( + "Partial memory rebuild was for a different embedding model " + f"({meta.get('model_provider')}/{meta.get('model_name')}) -- " + "discarding it and starting fresh" + ) + return None, set() + + try: + db = MyFaiss.load_local( + folder_path=db_dir, + index_name=name, + embeddings=embedder, + allow_dangerous_deserialization=True, + distance_strategy=DistanceStrategy.COSINE, + relevance_score_fn=Memory._cosine_normalizer, + ) + except Exception as e: + PrintStyle.standard( + f"Could not load partial memory rebuild ({e}) -- starting fresh" + ) + return None, set() + + return db, set(db.get_all_docs().keys()) + + @staticmethod + def _write_rebuild_checkpoint( + db: "MyFaiss", db_dir: str, model_config: "models.ModelConfig" + ) -> None: + name = Memory._REBUILD_INDEX_NAME + db.save_local(folder_path=db_dir, index_name=name) + files.write_file( + Memory._rebuild_meta_path(db_dir), + json.dumps( + { + "model_provider": model_config.provider, + "model_name": model_config.name, + } + ), + ) + + @staticmethod + def _clear_rebuild_checkpoint(db_dir: str) -> None: + name = Memory._REBUILD_INDEX_NAME + for ext in ("faiss", "pkl", "json"): + path = files.get_abs_path(db_dir, f"{name}.{ext}") + try: + if os.path.exists(path): + os.remove(path) + except Exception: + pass + @staticmethod def _get_embedding_config(agent=None): from plugins._model_config.helpers.model_config import get_embedding_model_config_object @@ -211,24 +318,58 @@ def initialize( # DB not loaded, create one if not db: - index = faiss.IndexFlatIP(len(embedder.embed_query("example"))) + resumed_db: MyFaiss | None = None + resumed_ids: set[str] = set() + if docs: + resumed_db, resumed_ids = Memory._load_rebuild_checkpoint( + db_dir, model_config, embedder + ) - db = MyFaiss( - embedding_function=embedder, - index=index, - docstore=InMemoryDocstore(), - index_to_docstore_id={}, - distance_strategy=DistanceStrategy.COSINE, - # normalize_L2=True, - relevance_score_fn=Memory._cosine_normalizer, - ) + if resumed_ids: + db = resumed_db # type: ignore + PrintStyle.standard( + f"Resuming memory rebuild: {len(resumed_ids)}/{len(docs)} already embedded" + ) + else: + index = faiss.IndexFlatIP(len(embedder.embed_query("example"))) + + db = MyFaiss( + embedding_function=embedder, + index=index, + docstore=InMemoryDocstore(), + index_to_docstore_id={}, + distance_strategy=DistanceStrategy.COSINE, + # normalize_L2=True, + relevance_score_fn=Memory._cosine_normalizer, + ) # insert docs if reindexing if docs: PrintStyle.standard("Indexing memories...") if log_item: log_item.stream(progress="\nIndexing memories") - db.add_documents(documents=list(docs.values()), ids=list(docs.keys())) + + # Embed in small batches, checkpointing to disk after each + # one, so an interrupted rebuild (crash, restart, a slow/ + # shared embeddings backend timing out) resumes from where + # it left off instead of re-embedding everything -- which, + # for a large memory store on a slow backend, can otherwise + # cost hours of redone work every time it's interrupted. + remaining = [(k, v) for k, v in docs.items() if k not in resumed_ids] + batch_size = Memory._REBUILD_BATCH_SIZE + done = len(resumed_ids) + total = len(docs) + for i in range(0, len(remaining), batch_size): + chunk = remaining[i : i + batch_size] + db.add_documents( + documents=[d for _, d in chunk], + ids=[k for k, _ in chunk], + ) + done += len(chunk) + PrintStyle.standard(f"Indexed {done}/{total} memories...") + Memory._write_rebuild_checkpoint(db, db_dir, model_config) + + Memory._clear_rebuild_checkpoint(db_dir) # save DB Memory._save_db_file(db, memory_subdir) diff --git a/tests/test_memory_rebuild_resume.py b/tests/test_memory_rebuild_resume.py new file mode 100644 index 0000000000..8c8ab279b2 --- /dev/null +++ b/tests/test_memory_rebuild_resume.py @@ -0,0 +1,384 @@ +"""Tests for the memory-rebuild checkpoint/resume mechanism in +plugins/_memory/helpers/memory.py. + +Context: a reindex (triggered when the configured embedding model no longer +matches the one a memory store was built with) used to embed every document +in a single pass and only persist the result at the very end. Interrupting +that pass (crash, restart, a slow/shared embeddings backend) threw away all +already-embedded work. The fix batches the reindex and checkpoints to disk +after every batch (`index.rebuilding.faiss`/`.pkl`), so a resumed run can +skip documents it already embedded. + +That checkpoint must not be trusted blindly: if the target embedding model +changes *again* while a rebuild is interrupted, the checkpoint's vectors are +of a different dimension/semantic space than what the new target model +would produce. These tests cover both the happy path (resume under the same +model) and that guard (refuse a stale, wrong-model checkpoint and rebuild +fresh instead of corrupting/crashing on it). +""" + +import json +import sys +import types +from pathlib import Path + +import pytest + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + +# --- Stub out heavy/unrelated modules memory.py's import chain pulls in +# transitively via `import models` / `from agent import Agent, AgentContext`, +# so these tests exercise the real memory.py + real FAISS without needing a +# full runtime (network clients, MCP, litellm providers, etc.) to be usable. +sys.modules.setdefault("giturlparse", types.SimpleNamespace(parse=lambda *a, **k: None)) + + +from langchain_core.documents import Document # noqa: E402 +from langchain_core.embeddings import Embeddings # noqa: E402 + +from plugins._memory.helpers import memory as memory_module # noqa: E402 +from plugins._memory.helpers.memory import Memory, MyFaiss # noqa: E402 + + +class FakeModelConfig: + """Minimal stand-in for models.ModelConfig -- memory.py only reads + .provider, .name and calls .build_kwargs().""" + + def __init__(self, provider: str, name: str): + self.provider = provider + self.name = name + + def build_kwargs(self): + return {} + + +class FakeEmbeddings(Embeddings): + """Deterministic, network-free stand-in for a real embedding model. + Encodes the text length into the vector so different texts get distinct + (but reproducible) embeddings; dimension is fixed per instance so tests + can simulate two "different models" by using two different dimensions. + """ + + def __init__(self, dim: int = 4): + self.dim = dim + + def embed_documents(self, texts): + return [self.embed_query(t) for t in texts] + + def embed_query(self, text): + seed = sum(text.encode()) + return [((seed + i) % 97) / 97.0 for i in range(self.dim)] + + +@pytest.fixture(autouse=True) +def _patch_embedding_backend(monkeypatch): + """Every test gets a fresh FakeEmbeddings via models.get_embedding_model, + keyed off model name so two different "models" (by name) get two + different, mutually-incompatible embedding dimensions -- mirroring a + real embedding-model swap (e.g. Qwen 4096-dim -> Nemotron 2048-dim).""" + + dims_by_name = {"model-a": 4, "model-b": 6} + + def fake_get_embedding_model(provider, name, model_config=None, **kwargs): + return FakeEmbeddings(dim=dims_by_name.get(name, 4)) + + monkeypatch.setattr(memory_module.models, "get_embedding_model", fake_get_embedding_model) + # Memory.index is a module-level cache keyed by subdir; make sure one + # test's cached db can't leak into the next. + memory_module.Memory.index = {} + + +def _make_docs(n: int) -> dict[str, Document]: + return {f"id-{i}": Document(page_content=f"document number {i}") for i in range(n)} + + +def test_load_rebuild_checkpoint_returns_nothing_when_absent(tmp_path): + db, resumed_ids = Memory._load_rebuild_checkpoint( + str(tmp_path), FakeModelConfig("other", "model-a"), FakeEmbeddings(4) + ) + assert db is None + assert resumed_ids == set() + + +def test_write_then_load_rebuild_checkpoint_round_trips(tmp_path): + embedder = FakeEmbeddings(4) + index = memory_module.faiss.IndexFlatIP(embedder.dim) + from langchain_community.docstore.in_memory import InMemoryDocstore + from langchain_community.vectorstores.utils import DistanceStrategy + + db = MyFaiss( + embedding_function=embedder, + index=index, + docstore=InMemoryDocstore(), + index_to_docstore_id={}, + distance_strategy=DistanceStrategy.COSINE, + relevance_score_fn=Memory._cosine_normalizer, + ) + docs = _make_docs(3) + db.add_documents(documents=list(docs.values()), ids=list(docs.keys())) + + model_config = FakeModelConfig("other", "model-a") + Memory._write_rebuild_checkpoint(db, str(tmp_path), model_config) + + assert (tmp_path / "index.rebuilding.faiss").exists() + assert (tmp_path / "index.rebuilding.pkl").exists() + meta = json.loads((tmp_path / "index.rebuilding.json").read_text()) + assert meta == {"model_provider": "other", "model_name": "model-a"} + + loaded_db, resumed_ids = Memory._load_rebuild_checkpoint( + str(tmp_path), model_config, embedder + ) + assert loaded_db is not None + assert resumed_ids == set(docs.keys()) + + +def test_load_rebuild_checkpoint_rejects_different_model(tmp_path): + """The core bug fix: a checkpoint recorded for one embedding model must + never be resumed into when the currently-configured model is different + -- even though both files are present and loadable, the vectors are from + an incompatible embedding space.""" + embedder_a = FakeEmbeddings(4) + index = memory_module.faiss.IndexFlatIP(embedder_a.dim) + from langchain_community.docstore.in_memory import InMemoryDocstore + from langchain_community.vectorstores.utils import DistanceStrategy + + db = MyFaiss( + embedding_function=embedder_a, + index=index, + docstore=InMemoryDocstore(), + index_to_docstore_id={}, + distance_strategy=DistanceStrategy.COSINE, + relevance_score_fn=Memory._cosine_normalizer, + ) + docs = _make_docs(3) + db.add_documents(documents=list(docs.values()), ids=list(docs.keys())) + Memory._write_rebuild_checkpoint(db, str(tmp_path), FakeModelConfig("other", "model-a")) + + # A DIFFERENT model is now the target (different name -> different dim + # via the fixture's fake backend). + embedder_b = FakeEmbeddings(6) + loaded_db, resumed_ids = Memory._load_rebuild_checkpoint( + str(tmp_path), FakeModelConfig("other", "model-b"), embedder_b + ) + assert loaded_db is None + assert resumed_ids == set() + # The stale checkpoint is left in place (only initialize()'s caller + # decides to overwrite/clear it, once it actually rebuilds); this helper + # only refuses to *use* it. + assert (tmp_path / "index.rebuilding.faiss").exists() + + +def test_load_rebuild_checkpoint_discards_when_meta_missing(tmp_path): + """A checkpoint with no model marker (e.g. left by a version of this + code predating the marker, or a corrupted write) must not be trusted + either -- treat it the same as a model mismatch, not as a free pass.""" + embedder = FakeEmbeddings(4) + index = memory_module.faiss.IndexFlatIP(embedder.dim) + from langchain_community.docstore.in_memory import InMemoryDocstore + from langchain_community.vectorstores.utils import DistanceStrategy + + db = MyFaiss( + embedding_function=embedder, + index=index, + docstore=InMemoryDocstore(), + index_to_docstore_id={}, + distance_strategy=DistanceStrategy.COSINE, + relevance_score_fn=Memory._cosine_normalizer, + ) + docs = _make_docs(2) + db.add_documents(documents=list(docs.values()), ids=list(docs.keys())) + db.save_local(folder_path=str(tmp_path), index_name=Memory._REBUILD_INDEX_NAME) + # deliberately no .json meta file written + + loaded_db, resumed_ids = Memory._load_rebuild_checkpoint( + str(tmp_path), FakeModelConfig("other", "model-a"), embedder + ) + assert loaded_db is None + assert resumed_ids == set() + + +def test_clear_rebuild_checkpoint_removes_all_three_files(tmp_path): + embedder = FakeEmbeddings(4) + index = memory_module.faiss.IndexFlatIP(embedder.dim) + from langchain_community.docstore.in_memory import InMemoryDocstore + from langchain_community.vectorstores.utils import DistanceStrategy + + db = MyFaiss( + embedding_function=embedder, + index=index, + docstore=InMemoryDocstore(), + index_to_docstore_id={}, + distance_strategy=DistanceStrategy.COSINE, + relevance_score_fn=Memory._cosine_normalizer, + ) + Memory._write_rebuild_checkpoint(db, str(tmp_path), FakeModelConfig("other", "model-a")) + assert (tmp_path / "index.rebuilding.faiss").exists() + + Memory._clear_rebuild_checkpoint(str(tmp_path)) + + assert not (tmp_path / "index.rebuilding.faiss").exists() + assert not (tmp_path / "index.rebuilding.pkl").exists() + assert not (tmp_path / "index.rebuilding.json").exists() + + +def test_clear_rebuild_checkpoint_is_a_noop_when_absent(tmp_path): + # Must not raise just because there was nothing to resume from. + Memory._clear_rebuild_checkpoint(str(tmp_path)) + + +class _NullLogItem: + def stream(self, **kwargs): + pass + + def update(self, **kwargs): + pass + + +def _initialize(tmp_path, monkeypatch, model_config, in_memory=False): + monkeypatch.setattr(memory_module, "abs_db_dir", lambda subdir: str(tmp_path)) + return Memory.initialize(_NullLogItem(), model_config, "default", in_memory) + + +def test_initialize_creates_fresh_index_when_none_exists(tmp_path, monkeypatch): + db, created = _initialize(tmp_path, monkeypatch, FakeModelConfig("other", "model-a")) + assert created is True + assert db.get_all_docs() == {} + assert (tmp_path / "index.faiss").exists() + assert not (tmp_path / "index.rebuilding.faiss").exists() + + +def test_initialize_resumes_interrupted_rebuild_under_same_model(tmp_path, monkeypatch): + """End-to-end: an existing (stale-model) index plus a partial rebuild + checkpoint for the CURRENT model must result in only the not-yet- + embedded documents being (re)embedded, not all of them.""" + monkeypatch.setattr(memory_module, "abs_db_dir", lambda subdir: str(tmp_path)) + + old_embedder = FakeEmbeddings(4) # pretend this is "model-a" too, just stale content + old_docs = _make_docs(5) + from langchain_community.docstore.in_memory import InMemoryDocstore + from langchain_community.vectorstores.utils import DistanceStrategy + + old_index = memory_module.faiss.IndexFlatIP(old_embedder.dim) + old_db = MyFaiss( + embedding_function=old_embedder, + index=old_index, + docstore=InMemoryDocstore(), + index_to_docstore_id={}, + distance_strategy=DistanceStrategy.COSINE, + relevance_score_fn=Memory._cosine_normalizer, + ) + old_db.add_documents(documents=list(old_docs.values()), ids=list(old_docs.keys())) + # This is the STALE on-disk index: recorded under a DIFFERENT model name + # so initialize() decides a reindex is needed. + Memory._save_db_file(old_db, "default") + memory_module.files.write_file( + memory_module.files.get_abs_path(str(tmp_path), "embedding.json"), + json.dumps({"model_provider": "other", "model_name": "stale-model"}), + ) + + # A partial rebuild already embedded the first 3 of those 5 documents, + # under the model we're about to target ("model-a") -- simulating a + # crash partway through a previous reindex attempt. + partial_ids = list(old_docs.keys())[:3] + partial_index = memory_module.faiss.IndexFlatIP(4) + partial_db = MyFaiss( + embedding_function=FakeEmbeddings(4), + index=partial_index, + docstore=InMemoryDocstore(), + index_to_docstore_id={}, + distance_strategy=DistanceStrategy.COSINE, + relevance_score_fn=Memory._cosine_normalizer, + ) + partial_db.add_documents( + documents=[old_docs[i] for i in partial_ids], ids=partial_ids + ) + Memory._write_rebuild_checkpoint( + partial_db, str(tmp_path), FakeModelConfig("other", "model-a") + ) + + added_calls = [] + original_add_documents = MyFaiss.add_documents + + def spy_add_documents(self, documents, ids=None, **kwargs): + added_calls.append(list(ids or [])) + return original_add_documents(self, documents, ids=ids, **kwargs) + + monkeypatch.setattr(MyFaiss, "add_documents", spy_add_documents) + + db, created = Memory.initialize( + _NullLogItem(), FakeModelConfig("other", "model-a"), "default", False + ) + + assert created is True + # Only the 2 documents NOT already in the checkpoint were (re)embedded. + all_added_ids = {i for batch in added_calls for i in batch} + assert all_added_ids == set(old_docs.keys()) - set(partial_ids) + # The final index nonetheless contains all 5 original documents. + assert set(db.get_all_docs().keys()) == set(old_docs.keys()) + # The checkpoint is cleaned up once the rebuild completes successfully. + assert not (tmp_path / "index.rebuilding.faiss").exists() + # The final embedding.json now reflects the model we just rebuilt under. + final_meta = json.loads((tmp_path / "embedding.json").read_text()) + assert final_meta == {"model_provider": "other", "model_name": "model-a"} + + +def test_initialize_discards_stale_checkpoint_when_model_changed_again(tmp_path, monkeypatch): + """The scenario this fix exists for: a rebuild was interrupted, then the + target embedding model changed AGAIN before the next attempt. The stale + checkpoint (built for the now-abandoned intermediate model) must be + discarded, not resumed into -- resuming would otherwise mix embedding + spaces or crash on a FAISS dimension mismatch the moment a new batch is + added.""" + monkeypatch.setattr(memory_module, "abs_db_dir", lambda subdir: str(tmp_path)) + + old_docs = _make_docs(4) + from langchain_community.docstore.in_memory import InMemoryDocstore + from langchain_community.vectorstores.utils import DistanceStrategy + + old_index = memory_module.faiss.IndexFlatIP(4) + old_db = MyFaiss( + embedding_function=FakeEmbeddings(4), + index=old_index, + docstore=InMemoryDocstore(), + index_to_docstore_id={}, + distance_strategy=DistanceStrategy.COSINE, + relevance_score_fn=Memory._cosine_normalizer, + ) + old_db.add_documents(documents=list(old_docs.values()), ids=list(old_docs.keys())) + Memory._save_db_file(old_db, "default") + memory_module.files.write_file( + memory_module.files.get_abs_path(str(tmp_path), "embedding.json"), + json.dumps({"model_provider": "other", "model_name": "stale-model"}), + ) + + # An interrupted rebuild checkpoint exists, but it was for "model-a" + # (dim=4) -- NOT the model we're now targeting ("model-b", dim=6). + partial_index = memory_module.faiss.IndexFlatIP(4) + partial_db = MyFaiss( + embedding_function=FakeEmbeddings(4), + index=partial_index, + docstore=InMemoryDocstore(), + index_to_docstore_id={}, + distance_strategy=DistanceStrategy.COSINE, + relevance_score_fn=Memory._cosine_normalizer, + ) + partial_db.add_documents( + documents=[list(old_docs.values())[0]], ids=[list(old_docs.keys())[0]] + ) + Memory._write_rebuild_checkpoint( + partial_db, str(tmp_path), FakeModelConfig("other", "model-a") + ) + + # Must not raise (e.g. a FAISS "assert d == self.d" dimension error from + # blindly adding dim=6 vectors into the dim=4 stale checkpoint index). + db, created = Memory.initialize( + _NullLogItem(), FakeModelConfig("other", "model-b"), "default", False + ) + + assert created is True + assert set(db.get_all_docs().keys()) == set(old_docs.keys()) + assert db.index.d == 6 + final_meta = json.loads((tmp_path / "embedding.json").read_text()) + assert final_meta == {"model_provider": "other", "model_name": "model-b"} From 20d9b3567f501c2eec50b8e2eff9f4fbf058509a Mon Sep 17 00:00:00 2001 From: omar-nahhas Date: Wed, 29 Jul 2026 14:53:04 +0100 Subject: [PATCH 2/2] =?UTF-8?q?fix(=5Fmemory):=20address=20review=20?= =?UTF-8?q?=E2=80=94=20checkpoint=20signature=20+=20clear-after-save=20ord?= =?UTF-8?q?ering?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Two gaps flagged by review on the resumable-reindex fix: 1. The rebuild checkpoint's model-identity check only compared provider/name. A kwarg that changes the embedder's actual output (e.g. an OpenAI-style `dimensions` override, or an api_base pointed at a different backend) could change while provider/name stayed the same, letting a checkpoint built under the old kwargs get resumed into under the new ones. The signature now includes build_kwargs() (minus api_key, which is deliberately excluded since it's persisted to disk in the memory dir and a credential rotation alone isn't "a different model"). 2. The rebuild checkpoint was cleared right after the batch loop, before the final index/embedding.json were durably written. A crash in that gap left neither a valid final index nor a resumable checkpoint, reintroducing the full re-embed-from-scratch failure mode the checkpoint exists to prevent. Clearing now happens strictly after the final save succeeds. Adds 3 tests: api_key exclusion, kwarg-change rejection, and checkpoint-survives-a-crash-during-final-save. --- plugins/_memory/helpers/memory.py | 66 +++++++++----- tests/test_memory_rebuild_resume.py | 131 +++++++++++++++++++++++++++- 2 files changed, 171 insertions(+), 26 deletions(-) diff --git a/plugins/_memory/helpers/memory.py b/plugins/_memory/helpers/memory.py index d92cff0c27..0b1247ab43 100644 --- a/plugins/_memory/helpers/memory.py +++ b/plugins/_memory/helpers/memory.py @@ -74,6 +74,26 @@ class Area(Enum): def _rebuild_meta_path(db_dir: str) -> str: return files.get_abs_path(db_dir, f"{Memory._REBUILD_INDEX_NAME}.json") + @staticmethod + def _rebuild_model_signature(model_config: "models.ModelConfig") -> dict: + """Everything about model_config that can change what the embedder + actually produces, excluding secrets. provider/name alone aren't + enough: some providers accept kwargs (e.g. an OpenAI `dimensions` + override, or an `api_base` pointed at a different backend/model + version) that change the output dimension/semantics without + changing provider or name. api_key is deliberately excluded -- + this signature is persisted to disk in the memory dir, and a + credential rotation alone must not be treated as "a different + model".""" + kwargs = { + k: v for k, v in model_config.build_kwargs().items() if k != "api_key" + } + return { + "model_provider": model_config.provider, + "model_name": model_config.name, + "model_kwargs": kwargs, + } + @staticmethod def _load_rebuild_checkpoint( db_dir: str, @@ -81,16 +101,17 @@ def _load_rebuild_checkpoint( embedder: Any, ) -> tuple["MyFaiss | None", set[str]]: """Load a partially-rebuilt index left behind by an interrupted - reindex, but only if it was built with the SAME embedding model as - the one we're about to (re)index with. - - A stale checkpoint left over from a *different* model (e.g. the - model was switched again while the previous rebuild was interrupted) - holds vectors of a possibly different dimension/semantics and must - not be reused: resuming into it would raise a FAISS dimension - assertion the moment a new document is added, rather than making - progress. Returns (db, resumed_ids); db is None (and resumed_ids - empty) if there is nothing usable to resume from. + reindex, but only if it was built under the SAME effective embedding + configuration as the one we're about to (re)index with. + + A stale checkpoint left over from a *different* configuration (the + model was switched again, or a dimension-affecting kwarg changed, + while the previous rebuild was interrupted) holds vectors of a + possibly different dimension/semantics and must not be reused: + resuming into it would raise a FAISS dimension assertion the moment + a new document is added, rather than making progress. Returns (db, + resumed_ids); db is None (and resumed_ids empty) if there is + nothing usable to resume from. """ name = Memory._REBUILD_INDEX_NAME if not ( @@ -112,12 +133,9 @@ def _load_rebuild_checkpoint( ) return None, set() - if ( - meta.get("model_provider") != model_config.provider - or meta.get("model_name") != model_config.name - ): + if meta != Memory._rebuild_model_signature(model_config): PrintStyle.standard( - "Partial memory rebuild was for a different embedding model " + "Partial memory rebuild was for a different embedding configuration " f"({meta.get('model_provider')}/{meta.get('model_name')}) -- " "discarding it and starting fresh" ) @@ -148,12 +166,7 @@ def _write_rebuild_checkpoint( db.save_local(folder_path=db_dir, index_name=name) files.write_file( Memory._rebuild_meta_path(db_dir), - json.dumps( - { - "model_provider": model_config.provider, - "model_name": model_config.name, - } - ), + json.dumps(Memory._rebuild_model_signature(model_config)), ) @staticmethod @@ -369,8 +382,6 @@ def initialize( PrintStyle.standard(f"Indexed {done}/{total} memories...") Memory._write_rebuild_checkpoint(db, db_dir, model_config) - Memory._clear_rebuild_checkpoint(db_dir) - # save DB Memory._save_db_file(db, memory_subdir) # save meta file @@ -385,6 +396,15 @@ def initialize( ), ) + # Only drop the rebuild checkpoint once the final index and meta + # file are durably written. Clearing it any earlier (e.g. right + # after the batch loop) leaves a window where a crash between + # "checkpoint cleared" and "final save complete" loses BOTH the + # resumable checkpoint and the persisted index -- reintroducing + # the full re-embed-from-scratch failure mode this exists to fix. + if docs: + Memory._clear_rebuild_checkpoint(db_dir) + created = True return db, created diff --git a/tests/test_memory_rebuild_resume.py b/tests/test_memory_rebuild_resume.py index 8c8ab279b2..b91f515897 100644 --- a/tests/test_memory_rebuild_resume.py +++ b/tests/test_memory_rebuild_resume.py @@ -46,12 +46,17 @@ class FakeModelConfig: """Minimal stand-in for models.ModelConfig -- memory.py only reads .provider, .name and calls .build_kwargs().""" - def __init__(self, provider: str, name: str): + def __init__(self, provider: str, name: str, kwargs: dict | None = None, api_key: str = ""): self.provider = provider self.name = name + self.kwargs = kwargs or {} + self.api_key = api_key def build_kwargs(self): - return {} + kwargs = dict(self.kwargs) + if self.api_key: + kwargs["api_key"] = self.api_key + return kwargs class FakeEmbeddings(Embeddings): @@ -125,7 +130,7 @@ def test_write_then_load_rebuild_checkpoint_round_trips(tmp_path): assert (tmp_path / "index.rebuilding.faiss").exists() assert (tmp_path / "index.rebuilding.pkl").exists() meta = json.loads((tmp_path / "index.rebuilding.json").read_text()) - assert meta == {"model_provider": "other", "model_name": "model-a"} + assert meta == {"model_provider": "other", "model_name": "model-a", "model_kwargs": {}} loaded_db, resumed_ids = Memory._load_rebuild_checkpoint( str(tmp_path), model_config, embedder @@ -382,3 +387,123 @@ def test_initialize_discards_stale_checkpoint_when_model_changed_again(tmp_path, assert db.index.d == 6 final_meta = json.loads((tmp_path / "embedding.json").read_text()) assert final_meta == {"model_provider": "other", "model_name": "model-b"} + + +def test_checkpoint_signature_excludes_api_key(tmp_path): + """A credential rotation alone (same provider/name/other-kwargs, only + api_key differs) must not invalidate an otherwise-matching checkpoint -- + it has no bearing on the embedding output, and the signature is + persisted to disk in the memory dir where a plaintext key doesn't + belong.""" + embedder = FakeEmbeddings(4) + index = memory_module.faiss.IndexFlatIP(embedder.dim) + from langchain_community.docstore.in_memory import InMemoryDocstore + from langchain_community.vectorstores.utils import DistanceStrategy + + db = MyFaiss( + embedding_function=embedder, + index=index, + docstore=InMemoryDocstore(), + index_to_docstore_id={}, + distance_strategy=DistanceStrategy.COSINE, + relevance_score_fn=Memory._cosine_normalizer, + ) + docs = _make_docs(2) + db.add_documents(documents=list(docs.values()), ids=list(docs.keys())) + Memory._write_rebuild_checkpoint( + db, str(tmp_path), FakeModelConfig("other", "model-a", api_key="old-secret-123") + ) + + meta = json.loads((tmp_path / "index.rebuilding.json").read_text()) + assert "old-secret-123" not in json.dumps(meta) + + loaded_db, resumed_ids = Memory._load_rebuild_checkpoint( + str(tmp_path), + FakeModelConfig("other", "model-a", api_key="rotated-secret-456"), + embedder, + ) + assert loaded_db is not None + assert resumed_ids == set(docs.keys()) + + +def test_checkpoint_signature_rejects_dimension_affecting_kwarg_change(tmp_path): + """The bug the automated review caught: provider+name alone aren't + enough to key the checkpoint on. A kwarg that changes the embedder's + actual output (e.g. an OpenAI-style `dimensions` override) must be part + of the signature, or a checkpoint built under the old kwargs gets + resumed into under the new ones -- silently mixing embedding spaces.""" + embedder = FakeEmbeddings(4) + index = memory_module.faiss.IndexFlatIP(embedder.dim) + from langchain_community.docstore.in_memory import InMemoryDocstore + from langchain_community.vectorstores.utils import DistanceStrategy + + db = MyFaiss( + embedding_function=embedder, + index=index, + docstore=InMemoryDocstore(), + index_to_docstore_id={}, + distance_strategy=DistanceStrategy.COSINE, + relevance_score_fn=Memory._cosine_normalizer, + ) + docs = _make_docs(2) + db.add_documents(documents=list(docs.values()), ids=list(docs.keys())) + Memory._write_rebuild_checkpoint( + db, str(tmp_path), FakeModelConfig("other", "model-a", kwargs={"dimensions": 1536}) + ) + + loaded_db, resumed_ids = Memory._load_rebuild_checkpoint( + str(tmp_path), + FakeModelConfig("other", "model-a", kwargs={"dimensions": 512}), + embedder, + ) + assert loaded_db is None + assert resumed_ids == set() + + +def test_initialize_does_not_clear_checkpoint_before_final_save_succeeds( + tmp_path, monkeypatch +): + """The other bug the automated review caught: clearing the rebuild + checkpoint must happen strictly AFTER the final index/meta are durably + written. If `_save_db_file` (or anything after it) fails or the process + is killed first, the checkpoint must still be there to resume from -- + otherwise that crash loses both the checkpoint and the final index.""" + monkeypatch.setattr(memory_module, "abs_db_dir", lambda subdir: str(tmp_path)) + + # A genuine reindex has to be triggered for the batch loop (and its + # checkpoint writes) to run at all -- set up a stale-model existing + # index, same as the other end-to-end tests. + from langchain_community.docstore.in_memory import InMemoryDocstore + from langchain_community.vectorstores.utils import DistanceStrategy + + old_docs = _make_docs(2) + old_index = memory_module.faiss.IndexFlatIP(4) + old_db = MyFaiss( + embedding_function=FakeEmbeddings(4), + index=old_index, + docstore=InMemoryDocstore(), + index_to_docstore_id={}, + distance_strategy=DistanceStrategy.COSINE, + relevance_score_fn=Memory._cosine_normalizer, + ) + old_db.add_documents(documents=list(old_docs.values()), ids=list(old_docs.keys())) + Memory._save_db_file(old_db, "default") + memory_module.files.write_file( + memory_module.files.get_abs_path(str(tmp_path), "embedding.json"), + json.dumps({"model_provider": "other", "model_name": "stale-model"}), + ) + + def boom(*args, **kwargs): + raise RuntimeError("simulated crash during final save") + + monkeypatch.setattr(Memory, "_save_db_file", staticmethod(boom)) + + with pytest.raises(RuntimeError, match="simulated crash"): + Memory.initialize( + _NullLogItem(), FakeModelConfig("other", "model-a"), "default", False + ) + + # The rebuild checkpoint (written by the batch loop before the crash) + # must still be on disk -- it's the only resumable state left. + assert (tmp_path / "index.rebuilding.faiss").exists() + assert (tmp_path / "index.rebuilding.json").exists()