186 lines
7.1 KiB
Python
186 lines
7.1 KiB
Python
#!/usr/bin/env python3
|
||
"""PostgresCasStateStore 对 muse-example 真实库的集成测试。
|
||
|
||
验证离线 fake 无法覆盖的部分:DB 触发器的迁移方向闭集、revision 恰好 +1、
|
||
PASSED 终态、身份不可变,以及条件 UPDATE 在真实并发语义下的行计数。
|
||
测试行用 unittest-cas- 前缀的 run_id,结束前物理清理(本表是可变注册表,无 append-only 约束)。
|
||
|
||
跑法(需 Tailscale 内网可达 muse-example):
|
||
.venv/bin/python tests/skills/write-next-chapter/test_candidate_cas_db.py
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import pathlib
|
||
import sys
|
||
import uuid
|
||
|
||
from psycopg.errors import RaiseException
|
||
|
||
PROJECT_ROOT = pathlib.Path(__file__).resolve().parents[3]
|
||
SKILLS_DIR = PROJECT_ROOT / ".agent" / "skills"
|
||
SCRIPT_DIR = PROJECT_ROOT / "muse" / "content" / "work" / "skills" / "generate" / "write-next-chapter" / "scripts"
|
||
for path in (SCRIPT_DIR,):
|
||
if str(path) not in sys.path:
|
||
sys.path.insert(0, str(path))
|
||
|
||
from candidate_cas import PostgresCasStateStore # noqa: E402
|
||
from muse_db import connect # noqa: E402
|
||
|
||
PREFIX = f"unittest-cas-{uuid.uuid4().hex[:8]}-"
|
||
_run_ids: list[str] = []
|
||
|
||
|
||
def _cleanup() -> None:
|
||
if not _run_ids:
|
||
return
|
||
with connect() as conn:
|
||
for run_id in _run_ids:
|
||
conn.execute("DELETE FROM example_candidate_cas WHERE run_id=%s", (run_id,))
|
||
conn.commit()
|
||
|
||
|
||
def _new_run_id(tag: str) -> str:
|
||
run_id = f"{PREFIX}{tag}"
|
||
_run_ids.append(run_id)
|
||
return run_id
|
||
|
||
|
||
def test_full_lifecycle_and_contention() -> None:
|
||
store = PostgresCasStateStore(work_id=0, target_chapter=None, creator="unittest")
|
||
run_id = _new_run_id("lifecycle")
|
||
draft = store.create(run_id, attempt=1, candidate_version=1)
|
||
assert (draft.state, draft.revision) == ("DRAFT", 1)
|
||
# 重复建链失败关闭
|
||
try:
|
||
store.create(run_id, attempt=1, candidate_version=1)
|
||
raise AssertionError("重复建链必须失败")
|
||
except Exception as exc:
|
||
assert getattr(exc, "code", None) == "CAS_CONFLICT", exc
|
||
checking = store.transition(draft, "CHECKING")
|
||
assert checking is not None and checking.revision == 2
|
||
# 旧 token 重放不命中
|
||
assert store.transition(draft, "CHECKING") is None
|
||
rejected = store.transition(checking, "REJECTED")
|
||
assert rejected is not None and rejected.state == "REJECTED"
|
||
# start_next 守卫:非递增身份拒绝
|
||
assert store.start_next(rejected, attempt=1, candidate_version=2) is None
|
||
assert store.start_next(rejected, attempt=2, candidate_version=1) is None
|
||
second = store.start_next(rejected, attempt=2, candidate_version=2)
|
||
assert second is not None and (second.state, second.revision) == ("DRAFT", 4)
|
||
checking2 = store.transition(second, "CHECKING")
|
||
passed = store.transition(checking2, "PASSED")
|
||
assert passed is not None and passed.state == "PASSED"
|
||
assert store.latest(run_id) == passed
|
||
|
||
|
||
def test_passed_is_terminal() -> None:
|
||
store = PostgresCasStateStore(creator="unittest")
|
||
run_id = _new_run_id("terminal")
|
||
draft = store.create(run_id, attempt=1, candidate_version=1)
|
||
checking = store.transition(draft, "CHECKING")
|
||
passed = store.transition(checking, "PASSED")
|
||
# 存储层对 PASSED 出发的迁移方向闭集拒绝(返回 None,不下探到库)
|
||
assert store.transition(passed, "REJECTED") is None
|
||
assert store.transition(passed, "CHECKING") is None
|
||
assert store.start_next(passed, attempt=2, candidate_version=2) is None
|
||
assert store.latest(run_id).state == "PASSED"
|
||
|
||
|
||
def test_db_trigger_rejects_malformed_direct_updates() -> None:
|
||
store = PostgresCasStateStore(creator="unittest")
|
||
run_id = _new_run_id("trigger")
|
||
store.create(run_id, attempt=1, candidate_version=1)
|
||
# 绕过存储直接写非法形状,触发器必须拒绝
|
||
bad_updates = (
|
||
# revision 跳号
|
||
"UPDATE example_candidate_cas SET state='CHECKING', revision=revision+2 WHERE run_id=%s",
|
||
# DRAFT 直接到 PASSED
|
||
"UPDATE example_candidate_cas SET state='PASSED', revision=revision+1 WHERE run_id=%s",
|
||
# 同轮迁移改身份
|
||
"UPDATE example_candidate_cas SET attempt=9, revision=revision+1 WHERE run_id=%s",
|
||
# run_id 身份漂移
|
||
"UPDATE example_candidate_cas SET run_id=run_id || 'x' WHERE run_id=%s",
|
||
)
|
||
for sql in bad_updates:
|
||
try:
|
||
with connect() as conn:
|
||
conn.execute(sql, (run_id,))
|
||
conn.commit()
|
||
raise AssertionError(f"触发器必须拒绝: {sql}")
|
||
except AssertionError:
|
||
raise
|
||
except RaiseException:
|
||
pass
|
||
latest = store.latest(run_id)
|
||
assert (latest.state, latest.revision) == ("DRAFT", 1), "非法 UPDATE 不得改变链"
|
||
|
||
|
||
def test_start_next_monotonic_at_db_level() -> None:
|
||
store = PostgresCasStateStore(creator="unittest")
|
||
run_id = _new_run_id("startnext")
|
||
draft = store.create(run_id, attempt=1, candidate_version=1)
|
||
checking = store.transition(draft, "CHECKING")
|
||
rejected = store.transition(checking, "REJECTED")
|
||
assert rejected is not None and rejected.state == "REJECTED"
|
||
# 绕过存储直接开新轮但不递增 attempt/candidate_version,触发器必须拒绝
|
||
try:
|
||
with connect() as conn:
|
||
conn.execute(
|
||
"UPDATE example_candidate_cas SET state='DRAFT', revision=revision+1 "
|
||
"WHERE run_id=%s AND state='REJECTED'", (run_id,))
|
||
conn.commit()
|
||
raise AssertionError("REJECTED 开新轮必须递增 attempt/candidate_version")
|
||
except AssertionError:
|
||
raise
|
||
except RaiseException:
|
||
pass
|
||
assert store.latest(run_id).state == "REJECTED", "非法开新轮不得改变链"
|
||
|
||
|
||
def test_insert_must_start_draft_revision_one() -> None:
|
||
run_id = _new_run_id("insert-guard")
|
||
try:
|
||
with connect() as conn:
|
||
conn.execute(
|
||
"INSERT INTO example_candidate_cas(run_id, attempt, candidate_version, state, revision, creator) "
|
||
"VALUES (%s,1,1,'CHECKING',1,'unittest')", (run_id,))
|
||
conn.commit()
|
||
raise AssertionError("初始行必须是 DRAFT")
|
||
except AssertionError:
|
||
raise
|
||
except RaiseException:
|
||
pass
|
||
try:
|
||
with connect() as conn:
|
||
conn.execute(
|
||
"INSERT INTO example_candidate_cas(run_id, attempt, candidate_version, state, revision, creator) "
|
||
"VALUES (%s,1,1,'DRAFT',5,'unittest')", (run_id,))
|
||
conn.commit()
|
||
raise AssertionError("初始 revision 必须是 1")
|
||
except AssertionError:
|
||
raise
|
||
except RaiseException:
|
||
pass
|
||
|
||
|
||
def main() -> None:
|
||
tests = (
|
||
test_full_lifecycle_and_contention,
|
||
test_passed_is_terminal,
|
||
test_db_trigger_rejects_malformed_direct_updates,
|
||
test_start_next_monotonic_at_db_level,
|
||
test_insert_must_start_draft_revision_one,
|
||
)
|
||
try:
|
||
for test in tests:
|
||
test()
|
||
print(f"PASS: {test.__name__}")
|
||
print("PASS:PostgresCasStateStore 真实库集成测试全部通过")
|
||
finally:
|
||
_cleanup()
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|