"""本地 muse.db 合同:WAL、五表、向量暴力扫、revise 必带 diff。""" from __future__ import annotations import math import pathlib import sys import tempfile import unittest ROOT = next( parent for parent in (pathlib.Path(__file__).resolve().parent, *pathlib.Path(__file__).resolve().parents) if (parent / "AGENTS.md").is_file() and (parent / ".git").exists() ) 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, cosine, pack_embedding, search_card_vectors, unpack_embedding, upsert_card, ) class MuseStoreTest(unittest.TestCase): def setUp(self) -> None: self.tmp = pathlib.Path(tempfile.mkdtemp()) self.db = self.tmp / "muse.db" def test_wal_and_required_tables(self) -> None: with connect(self.db) as conn: mode = conn.execute("PRAGMA journal_mode").fetchone()[0] self.assertEqual(str(mode).lower(), "wal") names = { row[0] for row in conn.execute( "SELECT name FROM sqlite_master WHERE type='table'" ) } for table in ("runs", "events", "reviews", "revisions", "cards"): self.assertIn(table, names) def test_brute_force_vector_search_ranks_by_cosine(self) -> None: query = [1.0, 0.0, 0.0] upsert_card( card_id="near", kind="craft", title="近", payload={"型": "craft", "名称": "近", "一句话摘要": "同向"}, embedding=[0.9, 0.1, 0.0], path=self.db, ) upsert_card( card_id="far", kind="craft", title="远", payload={"型": "craft", "名称": "远", "一句话摘要": "正交"}, embedding=[0.0, 1.0, 0.0], path=self.db, ) hits = search_card_vectors(query, kind="craft", top=2, path=self.db) self.assertEqual([item["cardId"] for item in hits], ["near", "far"]) packed = pack_embedding(query) self.assertEqual(unpack_embedding(packed), query) self.assertTrue(hits[0]["score"] > hits[1]["score"]) self.assertTrue(math.isclose(hits[0]["score"], cosine(query, [0.9, 0.1, 0.0]), rel_tol=1e-6)) def test_revise_requires_diff_and_writes_revision(self) -> None: with connect(self.db) as conn: conn.execute( """INSERT INTO runs(id, created_at, kind, input_json, output_text, meta_json, skill_set_hash) VALUES ('r1', '2026-08-28T00:00:00+00:00', 'agent.writer', '{}', '旧稿', '{}', 'x')""" ) conn.commit() with self.assertRaises(ValueError): add_review( run_id="r1", target="candidate", action="revise", reviewer="qingse", path=self.db, ) review_id = add_review( run_id="r1", target="candidate", action="revise", reviewer="qingse", reason="收紧钩子", before_text="旧稿", after_text="新稿", path=self.db, ) with connect(self.db) as conn: review = conn.execute("SELECT * FROM reviews WHERE id=?", (review_id,)).fetchone() revision = conn.execute( "SELECT * FROM revisions WHERE review_id=?", (review_id,) ).fetchone() self.assertEqual(review["action"], "revise") self.assertEqual(review["reviewer"], "qingse") 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()