mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-03 07:11:56 +00:00
feat: warn when an index is queried by a different embedding model
Changing EMBEDDINGS_NAME on a populated index is the one failure the width check cannot catch. Two models of the same width -- mpnet and granite are both 768 -- swap without raising anything, and every query is then embedded by a different model than the stored vectors were. Nothing fails; answers just get worse. Boot now compares what each source was built with against the active model and names the mismatched sources and the command that fixes them. The comparison goes through the registry rather than string equality, so a stored alias is not read as a different model. A source with no recorded model pre-dates the column and is therefore the legacy model, not unknown. That check is only as good as sources.model, which reembed was not maintaining: it rewrote the vectors and left the column naming the old model, so a source would be reported stale immediately after being migrated. It is now stamped after each source succeeds, from its own session -- sources lives in the user-data database while the vectors may not.
This commit is contained in:
1 parent
3242d68fd2
commit
cb62dea701
2 files changed
+67
No files matched your search
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
Reference in new issue
Block a user