254 lines
11 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 的离线测试。
用按条件 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 / ".agent" / "skills"
TEST_DIR = PROJECT_ROOT / "tests" / "skills" / "write-next-chapter"
CHECK_TEST_DIR = PROJECT_ROOT / "tests" / "skills" / "check-content-consistency"
SCRIPT_DIR = PROJECT_ROOT / "muse" / "content" / "work" / "skills" / "generate" / "write-next-chapter" / "scripts"
DETECT_DIR = PROJECT_ROOT / "muse" / "lifecycle" / "quality" / "skills" / "semantic" / "check-content-consistency" / "scripts"
READ_CONTEXT_DIR = PROJECT_ROOT / "muse" / "lifecycle" / "context" / "skills" / "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_hit_stops_at_authorization(self) -> None:
"""新合同:补证检索命中后停在授权终态,CAS 收敛 REJECTED,不开新轮。"""
context, _ = _valid_pair()
rows: dict = {}
def detector(current: dict, candidate: dict, _mechanical: dict) -> dict:
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
with self.assertRaises(PipelineError) as caught:
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(caught.exception.code, "AUTHORIZATION_REQUIRED")
self.assertEqual(caught.exception.details["nextAttempt"], 2)
chain = rows[context["runId"]]
# create→CHECKING→REJECTED:授权终态不开同运行新轮,续跑是新运行的事。
self.assertEqual((chain["state"], chain["revision"], chain["attempt"],
chain["candidate_version"]), ("REJECTED", 3, 1, 1))
if __name__ == "__main__":
unittest.main()