259 lines
10 KiB
Python

#!/usr/bin/env python3
"""文件 CAS 的不可变 journal、并发、崩溃恢复与迟到结果测试。"""
from __future__ import annotations
import json
import pathlib
import stat
import sys
import tempfile
import threading
import unittest
from unittest import mock
PROJECT_ROOT = pathlib.Path(__file__).resolve().parents[3]
SCRIPT_DIR = PROJECT_ROOT / "muse" / "authority" / "evidence" / "skills" / "record-run-evidence" / "scripts"
sys.path.insert(0, str(SCRIPT_DIR))
from file_cas import CasConflictError, CasRecoveryError, FileCasStore # noqa: E402
class FileCasTest(unittest.TestCase):
"""验证磁盘 journal 是唯一状态事实源。"""
def setUp(self) -> None:
"""为每个测试创建独立状态目录。"""
self.temporary = tempfile.TemporaryDirectory()
self.root = pathlib.Path(self.temporary.name) / "cas"
self.store = FileCasStore(self.root)
def tearDown(self) -> None:
"""删除独立状态目录。"""
self.temporary.cleanup()
def _initialize(self) -> dict[str, object]:
"""创建第一条 DRAFT revision。"""
return self.store.initialize(
run_id="run-1",
sample_id="sample-1",
arm="A",
attempt=1,
candidate_version=1,
state="DRAFT",
result={"status": "draft"},
safe_summary={"bodySha256": "sha256:" + "a" * 64},
cleanup_state="pending",
)
def test_revision_is_immutable_and_state_is_full_content_copy(self):
"""state.json 必须复制最新 revision 内容,不能只保存路径指针。"""
first = self._initialize()
second = self.store.transition(
expected_revision=1,
expected_attempt=1,
expected_candidate_version=1,
state="CHECKING",
result={"status": "checking"},
safe_summary={"findingCount": 0},
cleanup_state="cleaned",
)
revisions = sorted((self.root / "states").glob("*.json"))
self.assertEqual(len(revisions), 2)
self.assertEqual(json.loads(revisions[0].read_text(encoding="utf-8")), first)
self.assertEqual(json.loads(revisions[1].read_text(encoding="utf-8")), second)
self.assertEqual(json.loads((self.root / "state.json").read_text(encoding="utf-8")), second)
self.assertEqual(stat.S_IMODE(revisions[0].stat().st_mode), 0o600)
self.assertEqual(second["previousStateSha256"], self.store.state_sha256(first))
def test_concurrent_compare_and_swap_allows_exactly_one_winner(self):
"""两个调用方持有同一旧 revision 时,只能有一个推进成功。"""
self._initialize()
barrier = threading.Barrier(2)
successes: list[str] = []
conflicts: list[str] = []
def advance(target: str) -> None:
"""等待同时起跑,然后尝试以同一 expected token 推进。"""
barrier.wait()
try:
self.store.transition(
expected_revision=1,
expected_attempt=1,
expected_candidate_version=1,
state=target,
result={"winner": target},
safe_summary={"target": target},
cleanup_state="pending",
)
successes.append(target)
except CasConflictError:
conflicts.append(target)
threads = [threading.Thread(target=advance, args=(state,)) for state in ("CHECKING", "REJECTED")]
for thread in threads:
thread.start()
for thread in threads:
thread.join()
self.assertEqual(len(successes), 1)
self.assertEqual(len(conflicts), 1)
self.assertEqual(self.store.latest()["revision"], 2)
def test_recovery_promotes_valid_revision_after_state_copy_crash(self):
"""revision 已提交而 state.json 未覆盖时,恢复必须选中最后一条完整链。"""
self._initialize()
with mock.patch.object(self.store, "_publish_state_copy", side_effect=OSError("crash")):
with self.assertRaises(OSError):
self.store.transition(
expected_revision=1,
expected_attempt=1,
expected_candidate_version=1,
state="CHECKING",
result={"status": "checking"},
safe_summary={"findingCount": 0},
cleanup_state="pending",
)
self.assertEqual(json.loads((self.root / "state.json").read_text(encoding="utf-8"))["revision"], 1)
recovered = self.store.recover()
self.assertEqual(recovered["revision"], 2)
self.assertEqual(json.loads((self.root / "state.json").read_text(encoding="utf-8")), recovered)
def test_recovery_preserves_corrupted_revision_tail_for_audit(self):
"""可变 state 副本不能作为裁剪锚点,断链失败关闭时保留不可变文件。"""
self._initialize()
self.store.transition(
expected_revision=1,
expected_attempt=1,
expected_candidate_version=1,
state="CHECKING",
result={"status": "checking"},
safe_summary={"findingCount": 0},
cleanup_state="pending",
)
third = self.store.transition(
expected_revision=2,
expected_attempt=1,
expected_candidate_version=1,
state="PASSED",
result={"status": "passed"},
safe_summary={"findingCount": 0},
cleanup_state="cleaned",
)
state_copy = self.root / "state.json"
second = json.loads(
(self.root / "states" / "00000000000000000002.json").read_text(encoding="utf-8")
)
state_copy.write_text(json.dumps(second, ensure_ascii=False) + "\n", encoding="utf-8")
broken = self.root / "states" / "00000000000000000003.json"
broken_payload = json.loads(broken.read_text(encoding="utf-8"))
broken_payload["previousStateSha256"] = "sha256:" + "f" * 64
broken.write_text(json.dumps(broken_payload, ensure_ascii=False) + "\n", encoding="utf-8")
immutable_files = sorted((self.root / "states").glob("*.json"))
immutable_files.extend(
self.root / "artifacts" / str(third[field])
for field in ("resultArtifact", "safeSummaryArtifact", "cleanupArtifact")
)
broken_bytes = broken.read_bytes()
with self.assertRaises(CasRecoveryError) as raised:
self.store.recover()
self.assertEqual(raised.exception.code, "CAS_RECOVERY_FAILED")
for path in immutable_files:
self.assertTrue(path.exists(), path)
self.assertEqual(broken.read_bytes(), broken_bytes)
self.assertEqual(json.loads(state_copy.read_text(encoding="utf-8")), second)
def test_late_old_revision_and_repeated_terminal_result_are_rejected(self):
"""旧 revision、迟到 candidate 和重复终态都统一返回 CAS_CONFLICT。"""
self._initialize()
terminal = self.store.transition(
expected_revision=1,
expected_attempt=1,
expected_candidate_version=1,
state="PASSED",
result={"status": "passed"},
safe_summary={"findingCount": 0},
cleanup_state="cleaned",
)
cases = (
{"expected_revision": 1, "expected_attempt": 1, "expected_candidate_version": 1},
{"expected_revision": 2, "expected_attempt": 0, "expected_candidate_version": 1},
{"expected_revision": 2, "expected_attempt": 1, "expected_candidate_version": 0},
)
for expected in cases:
with self.subTest(expected=expected), self.assertRaises(CasConflictError) as raised:
self.store.transition(
**expected,
state="PASSED",
result={"status": "late"},
safe_summary={"findingCount": 1},
cleanup_state="cleaned",
)
self.assertEqual(raised.exception.code, "CAS_CONFLICT")
self.assertEqual(self.store.latest(), terminal)
def test_recovery_rejects_broken_hash_chain_and_preserves_immutable_files(self):
"""断链 journal 必须失败关闭,临时文件清理但不可变文件必须保留。"""
self._initialize()
temporary = self.root / "states" / ".orphan.tmp"
temporary.write_text("partial", encoding="utf-8")
broken = self.root / "states" / "00000000000000000002.json"
payload = dict(self.store.latest())
payload["revision"] = 2
payload["previousStateSha256"] = "sha256:" + "f" * 64
broken.write_text(json.dumps(payload), encoding="utf-8")
broken.chmod(0o600)
broken_bytes = broken.read_bytes()
with self.assertRaises(CasRecoveryError) as raised:
self.store.recover()
self.assertEqual(raised.exception.code, "CAS_RECOVERY_FAILED")
self.assertFalse(temporary.exists())
self.assertTrue(broken.exists())
self.assertEqual(broken.read_bytes(), broken_bytes)
def test_symlink_state_root_is_rejected(self):
"""状态根路径被软链接替换时不得跟随到攻击者目录。"""
other = pathlib.Path(self.temporary.name) / "other"
other.mkdir()
link = pathlib.Path(self.temporary.name) / "linked-cas"
link.symlink_to(other, target_is_directory=True)
with self.assertRaises(CasRecoveryError) as raised:
FileCasStore(link)
self.assertEqual(raised.exception.code, "CAS_PATH_INVALID")
def test_abnormal_lock_permissions_fail_closed(self):
"""固定锁被放宽权限后,后续事务必须拒绝继续读取或写入状态。"""
self._initialize()
(self.root / "state.lock").chmod(0o644)
with self.assertRaises(CasRecoveryError) as raised:
self.store.latest()
self.assertEqual(raised.exception.code, "CAS_PATH_INVALID")
if __name__ == "__main__":
unittest.main()