diff --git a/application/scripts/reembed.py b/application/scripts/reembed.py index bd325771..d614c846 100644 --- a/application/scripts/reembed.py +++ b/application/scripts/reembed.py @@ -323,6 +323,33 @@ def reembed_faiss(source_id: str, batch_size: int, dry_run: bool) -> Tuple[int, return len(chunks), len(docs) +def record_source_model(source_id: str) -> None: + """Stamp ``sources.model`` with the model these vectors were just built by. + + ``sources`` lives in the user-data database while the vectors may not, so + this takes its own session. Left unwritten, the column keeps naming the old + model and every consumer that trusts it -- the boot mismatch check above + all -- reports a source as stale immediately after it was migrated. + """ + from sqlalchemy import text + + from application.storage.db.session import db_session + + try: + with db_session() as conn: + conn.execute( + text("UPDATE sources SET model = :model WHERE id = :id"), + {"model": settings.EMBEDDINGS_NAME, "id": source_id}, + ) + except Exception as exc: # noqa: BLE001 — the vectors are already rewritten + logger.warning( + " %s: re-embedded, but could not update sources.model (%s). Retrieval " + "is correct; the mismatch warning may persist until it is.", + source_id, + exc, + ) + + def run( store_type: str, source_ids: Optional[Sequence[str]], @@ -355,6 +382,8 @@ def run( for position, source_id in enumerate(ids, start=1): try: seen, written = handler(source_id, batch_size, dry_run) + if written and not dry_run: + record_source_model(source_id) total_seen += seen total_written += written logger.info( diff --git a/tests/scripts/test_reembed.py b/tests/scripts/test_reembed.py index 0790ab14..2b92378c 100644 --- a/tests/scripts/test_reembed.py +++ b/tests/scripts/test_reembed.py @@ -390,3 +390,41 @@ class TestEmbedsInProcess: ) or 0): assert reembed.main(["--dry-run"]) == 0 assert seen["delegating"] is False + + +class TestRecordsTheModel: + """``sources.model`` is what the boot mismatch check reads.""" + + def test_a_re_embedded_source_is_stamped(self, monkeypatch): + from application.core.settings import settings + + monkeypatch.setattr(settings, "EMBEDDINGS_NAME", "new/model", raising=False) + conn = MagicMock() + session = MagicMock() + session.__enter__ = MagicMock(return_value=conn) + session.__exit__ = MagicMock(return_value=False) + with patch("application.storage.db.session.db_session", return_value=session): + reembed.record_source_model("src-1") + params = conn.execute.call_args.args[1] + assert params == {"model": "new/model", "id": "src-1"} + + def test_a_dry_run_stamps_nothing(self, monkeypatch): + with patch.object(reembed, "record_source_model") as record, patch.object( + reembed, "list_source_ids", return_value=["a"] + ), patch.object(reembed, "reembed_pgvector", return_value=(3, 0)): + reembed.run("pgvector", None, 64, True) + record.assert_not_called() + + def test_a_real_run_stamps_each_source(self, monkeypatch): + with patch.object(reembed, "record_source_model") as record, patch.object( + reembed, "list_source_ids", return_value=["a", "b"] + ), patch.object(reembed, "reembed_pgvector", return_value=(3, 3)): + reembed.run("pgvector", None, 64, False) + assert [c.args[0] for c in record.call_args_list] == ["a", "b"] + + def test_a_failed_source_is_not_stamped(self): + with patch.object(reembed, "record_source_model") as record, patch.object( + reembed, "list_source_ids", return_value=["a"] + ), patch.object(reembed, "reembed_pgvector", side_effect=RuntimeError("boom")): + assert reembed.run("pgvector", None, 64, False) == 1 + record.assert_not_called()