diff --git a/muse/migrate_pg_to_sqlite.py b/muse/migrate_pg_to_sqlite.py index d22eba7..4d2e540 100644 --- a/muse/migrate_pg_to_sqlite.py +++ b/muse/migrate_pg_to_sqlite.py @@ -96,9 +96,12 @@ def migrate( _prefix_count(sqlite, "draft"), "muse_knowledge_draft", ) + pg_embeddings = _count(pg, "example_knowledge_embedding") + if _blob_count(sqlite) != pg_embeddings: + _copy_embeddings(pg, sqlite) report["ledger"]["example_knowledge_embedding"] = _pair( - _count(pg, "example_knowledge_embedding"), - _count(pg, "example_knowledge_embedding"), + pg_embeddings, + _blob_count(sqlite), "example_knowledge_embedding", ) report["ledger"]["muse_knowledge_entity"] = _pair( @@ -143,7 +146,9 @@ def migrate( ) embeddings = _copy_embeddings(pg, sqlite) report["ledger"]["example_knowledge_embedding"] = _pair( - embeddings, embeddings, "example_knowledge_embedding" + embeddings, + _blob_count(sqlite), + "example_knowledge_embedding", ) entities = _copy_entities(pg, sqlite) report["ledger"]["muse_knowledge_entity"] = _pair( @@ -232,6 +237,14 @@ def _prefix_count(sqlite, prefix: str, table: str = "cards") -> int: ) +def _blob_count(sqlite) -> int: + return int( + sqlite.execute( + "SELECT COUNT(*) FROM cards WHERE embedding IS NOT NULL" + ).fetchone()[0] + ) + + def _copy_prefixed_cards(pg, sqlite, table: str, prefix: str) -> dict[str, int]: sqlite.execute("DELETE FROM cards WHERE id LIKE ?", (f"{prefix}:%",)) pg_count = _count(pg, table) @@ -342,6 +355,12 @@ def _copy_drafts(pg, sqlite) -> int: def _copy_embeddings(pg, sqlite) -> int: + """一行 embedding 一张卡,禁止按 draft_id 折叠。""" + + sqlite.execute("DELETE FROM cards WHERE id LIKE 'embedding:%'") + sqlite.execute( + "UPDATE cards SET embedding=NULL WHERE id LIKE 'draft:%' OR id LIKE 'entity:%'" + ) pg_count = _count(pg, "example_knowledge_embedding") cur = pg.execute( "SELECT id, draft_id, entity_id, content_hash, embedding::text FROM example_knowledge_embedding" @@ -350,24 +369,29 @@ def _copy_embeddings(pg, sqlite) -> int: for row in cur: vector = _parse_vector(row[4]) blob = pack_embedding(vector) if vector else None - card_id = f"draft:{row[1]}" if row[1] is not None else f"entity:{row[2]}" + owner = f"draft:{row[1]}" if row[1] is not None else f"entity:{row[2]}" sqlite.execute( - """UPDATE cards SET embedding=?, content_hash=? WHERE id=?""", - (blob, row[3], card_id), + """INSERT INTO cards(id, kind, title, payload_json, embedding, content_hash, created_at) + VALUES (?, 'embedding', ?, ?, ?, ?, datetime('now')) + ON CONFLICT(id) DO UPDATE SET + embedding=excluded.embedding, content_hash=excluded.content_hash""", + ( + f"embedding:{row[0]}", + owner, + _json({"draft_id": row[1], "entity_id": row[2]}), + blob, + row[3], + ), ) - if sqlite.execute("SELECT changes()").fetchone()[0] == 0 and blob is not None: - sqlite.execute( - """INSERT INTO cards(id, kind, title, payload_json, embedding, content_hash, created_at) - VALUES (?, 'embedding', ?, '{}', ?, ?, datetime('now')) - ON CONFLICT(id) DO UPDATE SET embedding=excluded.embedding""", - (f"embedding:{row[0]}", card_id, blob, row[3]), - ) n += 1 if n % 1000 == 0: sqlite.commit() sqlite.commit() - if n != pg_count: - raise RuntimeError(f"example_knowledge_embedding 对账失败 pg={pg_count} sqlite={n}") + blobs = _blob_count(sqlite) + if n != pg_count or blobs != pg_count: + raise RuntimeError( + f"example_knowledge_embedding 对账失败 pg={pg_count} copied={n} blobs={blobs}" + ) return n @@ -435,16 +459,10 @@ def _cosine_samples(pg, sqlite, n: int = 5) -> list[dict[str, Any]]: samples = [] for row in rows: pg_vec = _parse_vector(row[3]) - card_id = f"draft:{row[1]}" if row[1] is not None else f"entity:{row[2]}" blob_row = sqlite.execute( "SELECT embedding FROM cards WHERE id=? AND embedding IS NOT NULL", - (card_id,), + (f"embedding:{row[0]}",), ).fetchone() - if blob_row is None: - blob_row = sqlite.execute( - "SELECT embedding FROM cards WHERE id=?", - (f"embedding:{row[0]}",), - ).fetchone() if blob_row is None or blob_row[0] is None: raise RuntimeError(f"抽样向量缺失 embedding id={row[0]}") sqlite_vec = unpack_embedding(blob_row[0]) diff --git a/tests/protocol/test_muse_store.py b/tests/protocol/test_muse_store.py index 7f510cb..7842ca1 100644 --- a/tests/protocol/test_muse_store.py +++ b/tests/protocol/test_muse_store.py @@ -17,6 +17,7 @@ ROOT = next( if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) +from muse.migrate_pg_to_sqlite import _blob_count, _copy_embeddings # noqa: E402 from muse.store import ( # noqa: E402 add_review, connect, @@ -106,6 +107,43 @@ class MuseStoreTest(unittest.TestCase): self.assertEqual(revision["before_text"], "旧稿") self.assertEqual(revision["after_text"], "新稿") + def test_copy_embeddings_persists_one_card_per_row_not_collapsed_by_draft(self) -> None: + class _Rows: + def __init__(self, rows): + self._rows = rows + + def fetchone(self): + return self._rows[0] if self._rows else None + + def __iter__(self): + return iter(self._rows) + + class _FakePG: + def execute(self, sql, params=None): + if "COUNT(*)" in sql: + return _Rows([(3,)]) + return _Rows( + [ + (101, 10, None, "h1", "[1.0, 0.0]"), + (102, 10, None, "h2", "[0.0, 1.0]"), + (103, 10, None, "h3", "[0.5, 0.5]"), + ] + ) + + conn = connect(self.db) + copied = _copy_embeddings(_FakePG(), conn) + blobs = _blob_count(conn) + ids = [ + row[0] + for row in conn.execute( + "SELECT id FROM cards WHERE embedding IS NOT NULL ORDER BY id" + ) + ] + conn.close() + self.assertEqual(copied, 3) + self.assertEqual(blobs, 3) + self.assertEqual(ids, ["embedding:101", "embedding:102", "embedding:103"]) + if __name__ == "__main__": unittest.main()