muse-agent-example/tests/protocol/test_muse_store.py

112 lines
3.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""本地 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()