590 lines
24 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
"""离线评测文件型 CAS 与不可变 revision journal。
状态推进全程持有 state.lock 的排他 flock。结果、安全摘要和清理标记先以不可变
artifact 落盘,再提交不可变 states/<revision>.json,最后发布内容完全相同的
state.json 副本;恢复只信任 schema、hash 链和关联 artifact 都完整的 revision。
"""
from __future__ import annotations
import fcntl
import hashlib
import json
import os
import pathlib
import secrets
import stat
from contextlib import contextmanager
from typing import Any, Iterator, Mapping
TERMINAL_STATES = frozenset({"PASSED", "REJECTED", "FAILED", "COMPLETED"})
REQUIRED_STATE_FIELDS = frozenset(
{
"schemaVersion",
"runId",
"sampleId",
"arm",
"attempt",
"candidateVersion",
"state",
"revision",
"previousStateSha256",
"resultSha256",
"safeSummarySha256",
"cleanupState",
"cleanupStateSha256",
"resultArtifact",
"safeSummaryArtifact",
"cleanupArtifact",
}
)
class CasConflictError(RuntimeError):
"""旧 revision、旧 attempt、迟到候选或重复终态的统一冲突。"""
def __init__(self, message: str = "文件 CAS 比较失败") -> None:
super().__init__(message)
self.code = "CAS_CONFLICT"
self.acceptance_eligible = False
class CasRecoveryError(RuntimeError):
"""路径替换、断链或崩溃残留无法安全恢复时的失败关闭错误。"""
def __init__(self, code: str, message: str) -> None:
super().__init__(message)
self.code = code
self.acceptance_eligible = False
def _canonical_json(value: Any) -> str:
"""生成状态与 artifact 共用的规范 JSON。"""
return json.dumps(
value,
ensure_ascii=False,
sort_keys=True,
separators=(",", ":"),
allow_nan=False,
)
def _encoded_json(value: Any) -> bytes:
"""生成带单个末尾换行的持久化 JSON 字节。"""
return (_canonical_json(value) + "\n").encode("utf-8")
def _sha256_bytes(value: bytes) -> str:
"""返回带算法前缀的字节 SHA-256。"""
return "sha256:" + hashlib.sha256(value).hexdigest()
def _validate_identity(value: Any, field: str) -> str:
"""拒绝空身份和可用于路径逃逸的控制字符。"""
if not isinstance(value, str) or not value or any(char in value for char in ("/", "\\", "\x00")):
raise CasRecoveryError("CAS_STATE_INVALID", f"{field} 非法")
return value
class FileCasStore:
"""以单目录、不可变 journal 实现跨进程文件 CAS。"""
def __init__(self, root: str | pathlib.Path) -> None:
"""创建 0700 状态目录和 0600 固定锁,并拒绝软链接根。"""
self.root = pathlib.Path(root).absolute()
if self.root.is_symlink():
raise CasRecoveryError("CAS_PATH_INVALID", "CAS 根目录不能是软链接")
try:
self.root.mkdir(parents=True, exist_ok=True, mode=0o700)
if self.root.is_symlink() or not self.root.is_dir():
raise OSError("CAS 根目录非法")
os.chmod(self.root, 0o700)
root_fd = os.open(self.root, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
except OSError as exc:
raise CasRecoveryError("CAS_PATH_INVALID", "CAS 根目录不可安全打开") from exc
try:
for name in ("states", "artifacts"):
try:
os.mkdir(name, 0o700, dir_fd=root_fd)
except FileExistsError:
pass
directory_fd = os.open(
name,
os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW,
dir_fd=root_fd,
)
try:
os.fchmod(directory_fd, 0o700)
finally:
os.close(directory_fd)
try:
lock_fd = os.open(
"state.lock",
os.O_RDWR | os.O_CREAT | os.O_EXCL | os.O_NOFOLLOW,
0o600,
dir_fd=root_fd,
)
except FileExistsError:
lock_fd = os.open("state.lock", os.O_RDWR | os.O_NOFOLLOW, dir_fd=root_fd)
try:
metadata = os.fstat(lock_fd)
if not stat.S_ISREG(metadata.st_mode):
raise OSError("锁文件不是普通文件")
os.fchmod(lock_fd, 0o600)
os.fsync(lock_fd)
finally:
os.close(lock_fd)
os.fsync(root_fd)
except OSError as exc:
raise CasRecoveryError("CAS_PATH_INVALID", "CAS 子目录或锁文件不安全") from exc
finally:
os.close(root_fd)
@staticmethod
def state_sha256(state: Mapping[str, Any]) -> str:
"""计算 revision 规范内容的稳定 SHA-256。"""
return _sha256_bytes(_encoded_json(state))
@contextmanager
def _locked_root(self) -> Iterator[int]:
"""通过 no-follow 根 fd 和固定锁包围完整 CAS 事务。"""
try:
root_fd = os.open(self.root, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
lock_fd = os.open("state.lock", os.O_RDWR | os.O_NOFOLLOW, dir_fd=root_fd)
metadata = os.fstat(lock_fd)
if not stat.S_ISREG(metadata.st_mode) or stat.S_IMODE(metadata.st_mode) != 0o600:
raise OSError("锁文件不是普通 0600 文件")
except OSError as exc:
raise CasRecoveryError("CAS_PATH_INVALID", "CAS 根或锁已被替换") from exc
try:
fcntl.flock(lock_fd, fcntl.LOCK_EX)
yield root_fd
finally:
try:
fcntl.flock(lock_fd, fcntl.LOCK_UN)
finally:
os.close(lock_fd)
os.close(root_fd)
def _open_directory(self, root_fd: int, name: str) -> int:
"""从已验证根 fd 打开 states/artifacts 子目录。"""
try:
return os.open(
name,
os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW,
dir_fd=root_fd,
)
except OSError as exc:
raise CasRecoveryError("CAS_PATH_INVALID", "CAS 子目录已被替换") from exc
def _read_json_at(self, directory_fd: int, name: str) -> dict[str, Any]:
"""以 O_NOFOLLOW 读取 0600 普通 JSON 文件。"""
try:
descriptor = os.open(name, os.O_RDONLY | os.O_NOFOLLOW, dir_fd=directory_fd)
metadata = os.fstat(descriptor)
if not stat.S_ISREG(metadata.st_mode):
raise OSError("目标不是普通文件")
with os.fdopen(descriptor, "r", encoding="utf-8") as handle:
value = json.load(handle)
except (OSError, json.JSONDecodeError) as exc:
raise CasRecoveryError("CAS_RECOVERY_FAILED", "CAS JSON 不可读取") from exc
if not isinstance(value, dict):
raise CasRecoveryError("CAS_RECOVERY_FAILED", "CAS JSON 必须是对象")
return value
def _atomic_publish(
self,
directory_fd: int,
final_name: str,
content: bytes,
*,
immutable: bool,
) -> None:
"""以 0600 排他临时文件、fsync、rename、目录 fsync 发布内容。"""
if "/" in final_name or final_name in {"", ".", ".."}:
raise CasRecoveryError("CAS_PATH_INVALID", "CAS 文件名非法")
temporary_name = f".{final_name}.{secrets.token_hex(8)}.tmp"
descriptor = os.open(
temporary_name,
os.O_WRONLY | os.O_CREAT | os.O_EXCL | os.O_NOFOLLOW,
0o600,
dir_fd=directory_fd,
)
try:
with os.fdopen(descriptor, "wb") as handle:
handle.write(content)
handle.flush()
os.fsync(handle.fileno())
if immutable:
try:
os.stat(final_name, dir_fd=directory_fd, follow_symlinks=False)
except FileNotFoundError:
pass
else:
raise CasRecoveryError("CAS_RECOVERY_FAILED", "不可变 CAS 文件已存在")
os.rename(
temporary_name,
final_name,
src_dir_fd=directory_fd,
dst_dir_fd=directory_fd,
)
os.chmod(final_name, 0o600, dir_fd=directory_fd, follow_symlinks=False)
os.fsync(directory_fd)
except BaseException:
try:
os.unlink(temporary_name, dir_fd=directory_fd)
except FileNotFoundError:
pass
raise
def _artifact_name(self, revision: int, kind: str) -> str:
"""生成与 revision 唯一绑定的不可变 artifact 文件名。"""
return f"{revision:020d}-{kind}.json"
def _revision_name(self, revision: int) -> str:
"""生成按字典序即数值序排列的 revision 文件名。"""
return f"{revision:020d}.json"
def _publish_state_copy(self, root_fd: int, state: Mapping[str, Any]) -> None:
"""把完整 revision 内容原子覆盖为 state.json 副本。"""
self._atomic_publish(root_fd, "state.json", _encoded_json(state), immutable=False)
def _commit_revision(
self,
root_fd: int,
*,
identity: Mapping[str, Any],
revision: int,
previous_state_sha256: str | None,
state: str,
result: Mapping[str, Any],
safe_summary: Mapping[str, Any],
cleanup_state: str,
) -> dict[str, Any]:
"""先提交三个 artifact,再提交 revision,最后发布 state 内容副本。"""
if not isinstance(result, Mapping) or not isinstance(safe_summary, Mapping):
raise CasRecoveryError("CAS_STATE_INVALID", "result 和 safe_summary 必须是对象")
_validate_identity(cleanup_state, "cleanupState")
artifacts_fd = self._open_directory(root_fd, "artifacts")
states_fd = self._open_directory(root_fd, "states")
try:
result_bytes = _encoded_json(result)
summary_bytes = _encoded_json(safe_summary)
cleanup_value = {"cleanupState": cleanup_state}
cleanup_bytes = _encoded_json(cleanup_value)
result_name = self._artifact_name(revision, "result")
summary_name = self._artifact_name(revision, "safe-summary")
cleanup_name = self._artifact_name(revision, "cleanup")
self._atomic_publish(artifacts_fd, result_name, result_bytes, immutable=True)
self._atomic_publish(artifacts_fd, summary_name, summary_bytes, immutable=True)
self._atomic_publish(artifacts_fd, cleanup_name, cleanup_bytes, immutable=True)
revision_state = {
"schemaVersion": "file-cas-state-v1",
"runId": identity["runId"],
"sampleId": identity["sampleId"],
"arm": identity["arm"],
"attempt": identity["attempt"],
"candidateVersion": identity["candidateVersion"],
"state": state,
"revision": revision,
"previousStateSha256": previous_state_sha256,
"resultSha256": _sha256_bytes(result_bytes),
"safeSummarySha256": _sha256_bytes(summary_bytes),
"cleanupState": cleanup_state,
"cleanupStateSha256": _sha256_bytes(cleanup_bytes),
"resultArtifact": result_name,
"safeSummaryArtifact": summary_name,
"cleanupArtifact": cleanup_name,
}
self._validate_state_shape(revision_state)
self._atomic_publish(
states_fd,
self._revision_name(revision),
_encoded_json(revision_state),
immutable=True,
)
finally:
os.close(states_fd)
os.close(artifacts_fd)
self._publish_state_copy(root_fd, revision_state)
os.fsync(root_fd)
return revision_state
def _validate_state_shape(self, state: Mapping[str, Any]) -> None:
"""校验 revision 最小 schema、身份和单调字段类型。"""
if REQUIRED_STATE_FIELDS - set(state) or state.get("schemaVersion") != "file-cas-state-v1":
raise CasRecoveryError("CAS_RECOVERY_FAILED", "CAS revision schema 非法")
for field in ("runId", "sampleId", "arm", "state", "cleanupState"):
_validate_identity(state.get(field), field)
for field in ("attempt", "candidateVersion", "revision"):
value = state.get(field)
if isinstance(value, bool) or not isinstance(value, int) or value < 1:
raise CasRecoveryError("CAS_RECOVERY_FAILED", f"{field} 非法")
previous = state.get("previousStateSha256")
hashes = ("resultSha256", "safeSummarySha256", "cleanupStateSha256")
if previous is not None and (
not isinstance(previous, str) or not __import__("re").fullmatch(r"sha256:[0-9a-f]{64}", previous)
):
raise CasRecoveryError("CAS_RECOVERY_FAILED", "previousStateSha256 非法")
if any(
not isinstance(state.get(field), str)
or not __import__("re").fullmatch(r"sha256:[0-9a-f]{64}", state[field])
for field in hashes
):
raise CasRecoveryError("CAS_RECOVERY_FAILED", "artifact hash 非法")
for field in ("resultArtifact", "safeSummaryArtifact", "cleanupArtifact"):
_validate_identity(state.get(field), field)
def _revision_files(self, states_fd: int) -> list[str]:
"""列出严格命名的不可变 revision;未知正式文件视为异常。"""
names: list[str] = []
for name in os.listdir(states_fd):
if name.startswith(".") and name.endswith(".tmp"):
continue
if not __import__("re").fullmatch(r"[0-9]{20}\.json", name):
raise CasRecoveryError("CAS_RECOVERY_FAILED", "states 目录包含未知文件")
names.append(name)
return sorted(names)
def _verify_artifact(self, artifacts_fd: int, name: str, expected_hash: str) -> None:
"""核对 revision 引用的 artifact 是普通文件且字节 hash 完整。"""
try:
descriptor = os.open(name, os.O_RDONLY | os.O_NOFOLLOW, dir_fd=artifacts_fd)
metadata = os.fstat(descriptor)
if not stat.S_ISREG(metadata.st_mode):
raise OSError("artifact 不是普通文件")
with os.fdopen(descriptor, "rb") as handle:
actual_hash = _sha256_bytes(handle.read())
except OSError as exc:
raise CasRecoveryError("CAS_RECOVERY_FAILED", "revision 引用的 artifact 缺失") from exc
if actual_hash != expected_hash:
raise CasRecoveryError("CAS_RECOVERY_FAILED", "revision 引用的 artifact hash 不一致")
def _load_valid_chain(
self, root_fd: int, *, revision_limit: int | None = None
) -> list[dict[str, Any]]:
"""加载并验证连续 revision、previous hash 链和全部关联 artifact。"""
states_fd = self._open_directory(root_fd, "states")
artifacts_fd = self._open_directory(root_fd, "artifacts")
try:
names = self._revision_files(states_fd)
if revision_limit is not None:
if revision_limit < 1 or len(names) < revision_limit:
raise CasRecoveryError("CAS_RECOVERY_FAILED", "恢复锚点 revision 不存在")
names = names[:revision_limit]
chain: list[dict[str, Any]] = []
for expected_revision, name in enumerate(names, start=1):
state = self._read_json_at(states_fd, name)
self._validate_state_shape(state)
if state["revision"] != expected_revision or name != self._revision_name(expected_revision):
raise CasRecoveryError("CAS_RECOVERY_FAILED", "CAS revision 不连续")
expected_previous = self.state_sha256(chain[-1]) if chain else None
if state["previousStateSha256"] != expected_previous:
raise CasRecoveryError("CAS_RECOVERY_FAILED", "CAS previous hash 断链")
if chain and any(
state[field] != chain[0][field] for field in ("runId", "sampleId", "arm")
):
raise CasRecoveryError("CAS_RECOVERY_FAILED", "CAS 运行身份发生漂移")
self._verify_artifact(artifacts_fd, state["resultArtifact"], state["resultSha256"])
self._verify_artifact(
artifacts_fd,
state["safeSummaryArtifact"],
state["safeSummarySha256"],
)
self._verify_artifact(
artifacts_fd,
state["cleanupArtifact"],
state["cleanupStateSha256"],
)
chain.append(state)
return chain
finally:
os.close(artifacts_fd)
os.close(states_fd)
def _read_state_copy(self, root_fd: int) -> dict[str, Any] | None:
"""读取 state.json 内容副本;不存在表示尚未初始化。"""
try:
return self._read_json_at(root_fd, "state.json")
except CasRecoveryError as exc:
try:
os.stat("state.json", dir_fd=root_fd, follow_symlinks=False)
except FileNotFoundError:
return None
raise exc
def _latest_consistent(self, root_fd: int) -> dict[str, Any] | None:
"""要求 state 副本与 journal 最新 revision 完全一致。"""
chain = self._load_valid_chain(root_fd)
state_copy = self._read_state_copy(root_fd)
if not chain:
if state_copy is not None:
raise CasRecoveryError("CAS_RECOVERY_FAILED", "state.json 没有对应 revision")
return None
if state_copy != chain[-1]:
raise CasRecoveryError("CAS_RECOVERY_REQUIRED", "state.json 与最新 revision 不一致")
return chain[-1]
def initialize(
self,
*,
run_id: str,
sample_id: str,
arm: str,
attempt: int,
candidate_version: int,
state: str,
result: Mapping[str, Any],
safe_summary: Mapping[str, Any],
cleanup_state: str,
) -> dict[str, Any]:
"""仅在空 journal 中创建 revision 1。"""
identity = {
"runId": _validate_identity(run_id, "runId"),
"sampleId": _validate_identity(sample_id, "sampleId"),
"arm": _validate_identity(arm, "arm"),
"attempt": attempt,
"candidateVersion": candidate_version,
}
if any(isinstance(value, bool) or not isinstance(value, int) or value < 1 for value in (attempt, candidate_version)):
raise CasRecoveryError("CAS_STATE_INVALID", "attempt/candidateVersion 必须是正整数")
_validate_identity(state, "state")
with self._locked_root() as root_fd:
if self._latest_consistent(root_fd) is not None:
raise CasConflictError("CAS 已初始化")
return self._commit_revision(
root_fd,
identity=identity,
revision=1,
previous_state_sha256=None,
state=state,
result=result,
safe_summary=safe_summary,
cleanup_state=cleanup_state,
)
def transition(
self,
*,
expected_revision: int,
expected_attempt: int,
expected_candidate_version: int,
state: str,
result: Mapping[str, Any],
safe_summary: Mapping[str, Any],
cleanup_state: str,
) -> dict[str, Any]:
"""仅当磁盘最新 token 完全匹配时提交下一条不可变 revision。"""
_validate_identity(state, "state")
with self._locked_root() as root_fd:
current = self._latest_consistent(root_fd)
if current is None:
raise CasConflictError("CAS 尚未初始化")
expected = (expected_revision, expected_attempt, expected_candidate_version)
actual = (current["revision"], current["attempt"], current["candidateVersion"])
if expected != actual or current["state"] in TERMINAL_STATES:
raise CasConflictError()
identity = {
"runId": current["runId"],
"sampleId": current["sampleId"],
"arm": current["arm"],
"attempt": current["attempt"],
"candidateVersion": current["candidateVersion"],
}
return self._commit_revision(
root_fd,
identity=identity,
revision=current["revision"] + 1,
previous_state_sha256=self.state_sha256(current),
state=state,
result=result,
safe_summary=safe_summary,
cleanup_state=cleanup_state,
)
def latest(self) -> dict[str, Any] | None:
"""返回与不可变 journal 一致的最新 state 内容副本。"""
with self._locked_root() as root_fd:
current = self._latest_consistent(root_fd)
return dict(current) if current is not None else None
def _clean_temporary_files(self, root_fd: int) -> None:
"""恢复前删除未被 rename 提交的临时文件。"""
directory_fds = [root_fd, self._open_directory(root_fd, "states"), self._open_directory(root_fd, "artifacts")]
try:
for directory_fd in directory_fds:
for name in os.listdir(directory_fd):
if name.startswith(".") and name.endswith(".tmp"):
os.unlink(name, dir_fd=directory_fd)
os.fsync(directory_fd)
finally:
for directory_fd in directory_fds[1:]:
os.close(directory_fd)
def _clean_orphan_artifacts(self, root_fd: int, chain: list[dict[str, Any]]) -> None:
"""删除没有任何合法 revision 引用的崩溃中间 artifact。"""
referenced = {
state[field]
for state in chain
for field in ("resultArtifact", "safeSummaryArtifact", "cleanupArtifact")
}
artifacts_fd = self._open_directory(root_fd, "artifacts")
try:
for name in os.listdir(artifacts_fd):
if name not in referenced:
metadata = os.stat(name, dir_fd=artifacts_fd, follow_symlinks=False)
if stat.S_ISDIR(metadata.st_mode):
raise CasRecoveryError("CAS_RECOVERY_FAILED", "artifacts 包含异常目录")
os.unlink(name, dir_fd=artifacts_fd)
os.fsync(artifacts_fd)
finally:
os.close(artifacts_fd)
def recover(self) -> dict[str, Any]:
"""恢复最后一条完整 revision,并清理临时文件与孤儿 artifact。"""
with self._locked_root() as root_fd:
self._clean_temporary_files(root_fd)
# 断链直接失败关闭;可变 state.json 不是不可变链的可信锚点,不能据此裁剪提交文件。
chain = self._load_valid_chain(root_fd)
if not chain:
raise CasRecoveryError("CAS_RECOVERY_FAILED", "没有可恢复的 CAS revision")
self._clean_orphan_artifacts(root_fd, chain)
latest = chain[-1]
state_copy = self._read_state_copy(root_fd)
if state_copy != latest:
self._publish_state_copy(root_fd, latest)
os.fsync(root_fd)
return dict(latest)
__all__ = ["CasConflictError", "CasRecoveryError", "FileCasStore"]