186 lines
7.1 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.

#!/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()