254 lines
11 KiB
Python
254 lines
11 KiB
Python
#!/usr/bin/env python3
|
||
"""PostgresCasStateStore 的离线测试。
|
||
|
||
用按条件 WHERE 语义模拟 example_candidate_cas 行为的假连接,验证:
|
||
- CAS 竞争(旧 token / 重复建链)失败关闭;
|
||
- transition/start_next 的守卫与 token 递增;
|
||
- 该存储可直接替换 InMemoryCasStateStore 驱动 run_writer_pipeline 全链。
|
||
|
||
跑法(仓库根目录):.venv/bin/python tests/skills/写下一章/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" / "写下一章"
|
||
CHECK_TEST_DIR = PROJECT_ROOT / "tests" / "skills" / "核对内容一致性"
|
||
SCRIPT_DIR = PROJECT_ROOT / "muse" / "content" / "work" / "skills" / "generate" / "写下一章" / "scripts"
|
||
DETECT_DIR = PROJECT_ROOT / "muse" / "lifecycle" / "quality" / "skills" / "semantic" / "核对内容一致性" / "scripts"
|
||
READ_CONTEXT_DIR = PROJECT_ROOT / "muse" / "lifecycle" / "context" / "skills" / "准备任务上下文" / "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 _requirements, _valid_pair # noqa: E402
|
||
from test_run_writer_pipeline import _semantic_pass, _writer_output # noqa: E402
|
||
from writer_contract import 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()
|