149 lines
5.0 KiB
Python
149 lines
5.0 KiB
Python
"""本地 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()
|