112 lines
3.7 KiB
Python
112 lines
3.7 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.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"], "新稿")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
unittest.main()
|