#!/usr/bin/env python3 """PostgresCasStateStore 的离线测试。 用按条件 WHERE 语义模拟 example_candidate_cas 行为的假连接,验证: - CAS 竞争(旧 token / 重复建链)失败关闭; - transition/start_next 的守卫与 token 递增; - 该存储可直接替换 InMemoryCasStateStore 驱动 run_writer_pipeline 全链。 跑法(仓库根目录):.venv/bin/python tests/skills/write-next-chapter/test_candidate_cas.py """ from __future__ import annotations import copy import pathlib import sys import unittest PROJECT_ROOT = pathlib.Path(__file__).resolve().parents[3] SKILLS_DIR = PROJECT_ROOT / ".claude" / "skills" TEST_DIR = PROJECT_ROOT / "tests" / "skills" / "write-next-chapter" CHECK_TEST_DIR = PROJECT_ROOT / "tests" / "skills" / "check-content-consistency" SCRIPT_DIR = SKILLS_DIR / "write-next-chapter" / "scripts" DETECT_DIR = SKILLS_DIR / "check-content-consistency" / "scripts" READ_CONTEXT_DIR = SKILLS_DIR / "assemble-context" / "scripts" for path in (READ_CONTEXT_DIR, DETECT_DIR, SCRIPT_DIR, TEST_DIR, CHECK_TEST_DIR): if str(path) not in sys.path: sys.path.insert(0, str(path)) from candidate_cas import PostgresCasStateStore # noqa: E402 from run_writer_pipeline import CasToken, PipelineError, run_writer_pipeline # noqa: E402 from run_writer_semantic_detector import canonical_sha256 # noqa: E402 from test_check_writer_candidate import _candidate_body, _requirements, _valid_pair # noqa: E402 from test_run_writer_pipeline import _semantic_pass, _writer_output # noqa: E402 from writer_contract import build_candidate_envelope, retrieval_identity # noqa: E402 class FakeCursor: """携带 rowcount 的最小游标。""" def __init__(self, rowcount: int = 0, row: tuple | None = None) -> None: self.rowcount = rowcount self._row = row def fetchone(self): return self._row class FakeCasConnection: """按条件 WHERE 语义模拟 example_candidate_cas 的短事务连接。 只模拟条件匹配的「行计数」事实(匹配则更新并计数 1,不匹配计数 0), 不重复实现 DB 触发器的迁移方向约束——那是 DB 集成测试的职责。 """ def __init__(self, rows: dict) -> None: self._rows = rows self._snapshot = copy.deepcopy(rows) def __enter__(self): return self def __exit__(self, *_): return False def commit(self) -> None: self._snapshot = copy.deepcopy(self._rows) def rollback(self) -> None: self._rows = copy.deepcopy(self._snapshot) def execute(self, sql: str, params=()): text = " ".join(sql.split()) if text.startswith("INSERT INTO example_candidate_cas"): run_id, work_id, chapter, attempt, version, creator, updater = params if run_id in self._rows: return FakeCursor(rowcount=0) self._rows[run_id] = { "work_id": work_id, "target_chapter": chapter, "attempt": attempt, "candidate_version": version, "state": "DRAFT", "revision": 1, "deleted": False, } return FakeCursor(rowcount=1) if text.startswith("UPDATE example_candidate_cas SET state=%s"): state, updater, run_id, revision, old_state, attempt, version = params row = self._rows.get(run_id) if not self._match(row, run_id, revision, old_state, attempt, version): return FakeCursor(rowcount=0) row["state"] = state row["revision"] += 1 return FakeCursor(rowcount=1) if text.startswith("UPDATE example_candidate_cas SET state='DRAFT'"): attempt, version, updater, run_id, revision, old_attempt, old_version = params row = self._rows.get(run_id) if not self._match(row, run_id, revision, "REJECTED", old_attempt, old_version): return FakeCursor(rowcount=0) row["state"] = "DRAFT" row["attempt"] = attempt row["candidate_version"] = version row["revision"] += 1 return FakeCursor(rowcount=1) if text.startswith("SELECT run_id, attempt, candidate_version, state, revision"): (run_id,) = params row = self._rows.get(run_id) if row is None or row["deleted"]: return FakeCursor(row=None) return FakeCursor(row=(run_id, row["attempt"], row["candidate_version"], row["state"], row["revision"])) raise AssertionError(f"未预期的 SQL: {text}") @staticmethod def _match(row, run_id, revision, state, attempt, version) -> bool: return ( row is not None and not row["deleted"] and row["revision"] == revision and row["state"] == state and row["attempt"] == attempt and row["candidate_version"] == version ) def _store(rows: dict, **kwargs) -> PostgresCasStateStore: """注入共享行表与假连接的被测存储。""" return PostgresCasStateStore(connect=lambda **_kw: FakeCasConnection(rows), **kwargs) class PostgresCasStateStoreTest(unittest.TestCase): def test_create_then_latest_returns_draft_token(self) -> None: rows: dict = {} store = _store(rows) token = store.create("run-1", attempt=1, candidate_version=1) self.assertEqual(token, CasToken("run-1", 1, 1, "DRAFT", 1)) self.assertEqual(store.latest("run-1"), token) def test_create_twice_same_run_fails_closed(self) -> None: rows: dict = {} store = _store(rows) store.create("run-1", attempt=1, candidate_version=1) with self.assertRaises(PipelineError) as caught: store.create("run-1", attempt=1, candidate_version=1) self.assertEqual(caught.exception.code, "CAS_CONFLICT") def test_transition_requires_exact_expected_token(self) -> None: rows: dict = {} store = _store(rows) draft = store.create("run-1", attempt=1, candidate_version=1) checking = store.transition(draft, "CHECKING") self.assertEqual(checking, CasToken("run-1", 1, 1, "CHECKING", 2)) # 旧 token 重放:revision 已前进,条件不命中,返回 None self.assertIsNone(store.transition(draft, "CHECKING")) # 非法方向同样不命中(fake 只模拟条件匹配) self.assertIsNone(store.transition(checking, "DRAFT")) def test_start_next_guards_monotonic_identity(self) -> None: rows: dict = {} store = _store(rows) draft = store.create("run-1", attempt=1, candidate_version=1) checking = store.transition(draft, "CHECKING") rejected = store.transition(checking, "REJECTED") # 非递增的 attempt / candidate_version 一律拒绝 self.assertIsNone(store.start_next(rejected, attempt=1, candidate_version=2)) self.assertIsNone(store.start_next(rejected, attempt=2, candidate_version=1)) self.assertIsNone(store.start_next(checking, attempt=2, candidate_version=2)) next_draft = store.start_next(rejected, attempt=2, candidate_version=2) self.assertEqual(next_draft, CasToken("run-1", 2, 2, "DRAFT", 4)) # 旧 REJECTED token 重放失败 self.assertIsNone(store.start_next(rejected, attempt=3, candidate_version=3)) def test_latest_missing_run_returns_none(self) -> None: self.assertIsNone(_store({}).latest("no-such-run")) def _advance_context(context: dict, attempt: int) -> dict: advanced = copy.deepcopy(context) advanced["attempt"] = attempt advanced["factEvidence"].append({ "evidenceId": f"fact-gap-{attempt}", "fact": "补充证据", "sourceType": "canonical_state", "sourceRef": {"sourceId": "state:gap", "sourceVersion": "v1"}, "contentSha256": "sha256:" + "4" * 64, "riskLevel": "low", }) advanced["contextSnapshot"]["contextSha256"] = retrieval_identity(advanced) return advanced class PipelineOnPostgresStoreTest(unittest.TestCase): """用持久存储驱动完整 pipeline,证明协议兼容与补证环行为不变。""" def test_full_pipeline_passes_with_postgres_store(self) -> None: context, _ = _valid_pair() rows: dict = {} result = run_writer_pipeline( context=context, requirements=_requirements(), writer=lambda current, version: _writer_output(current, version), evidence_provider=lambda *_args: self.fail("无 gap 不得补证"), semantic_detector=_semantic_pass, state_store=_store(rows), ) self.assertEqual(result["status"], "PASSED") self.assertEqual(result["candidateArtifact"]["schemaVersion"], "candidate-envelope-v2") chain = rows[context["runId"]] self.assertEqual(chain["state"], "PASSED") self.assertEqual(chain["revision"], 3) # create→CHECKING→PASSED def test_pipeline_evidence_loop_uses_start_next(self) -> None: context, _ = _valid_pair() rows: dict = {} def detector(current: dict, candidate: dict, _mechanical: dict) -> dict: if candidate["candidateVersion"] == 1: report = { "schemaVersion": "semantic-detection-v3", "runId": current["runId"], "sampleId": "cas-sample", "opaqueArmId": "cas-candidate", "inputSha256": "sha256:" + "1" * 64, "candidateVersion": candidate["candidateVersion"], "candidateSha256": candidate["candidateSha256"], "contextSnapshotSha256": current["contextSnapshot"]["contextSha256"], "modelReceiptSha256": "sha256:" + "2" * 64, "status": "needs_evidence", "claims": [], "findings": [], "assertionVerdicts": [], "hardConstraintVerdicts": [], "newSettingCandidates": [], "evidenceGaps": [{ "gapId": "gap-1", "query": "补证", "reason": "缺口", "priority": "high", "candidateSha256": candidate["candidateSha256"], "candidateQuote": "旧徽章", "startCodePoint": 11, "endCodePoint": 14, }], } report["reportSha256"] = canonical_sha256(report) return report return _semantic_pass(current, candidate, _mechanical) result = run_writer_pipeline( context=context, requirements=_requirements(), writer=lambda current, version: _writer_output(current, version), evidence_provider=lambda current, _gaps, attempt: _advance_context(current, attempt), semantic_detector=detector, state_store=_store(rows), ) self.assertEqual(result["status"], "PASSED") chain = rows[context["runId"]] # create→CHECKING→REJECTED→新DRAFT→CHECKING→PASSED self.assertEqual((chain["state"], chain["revision"], chain["attempt"], chain["candidate_version"]), ("PASSED", 6, 2, 2)) if __name__ == "__main__": unittest.main()