forked from chiki2bum2/SniperGold_ML
116 lines
3.9 KiB
Python
116 lines
3.9 KiB
Python
"""Checkpoint / resume state (Layer 7, spec 13).
| |||
| |||
CP_V1 checkpoint with non-self-referential ``checkpoint_hash`` and atomic
| |||
tmp+fsync+rename transitions. The status machine is fail-closed: any illegal
| |||
transition raises CheckpointTransitionError; corrupt payloads raise
| |||
CheckpointCorrupt; no silent recovery is ever performed.
| |||
"""
| |||
| |||
import os
| |||
| |||
from . import versions
| |||
from .util import canonical_json_sha256, atomic_write_json
| |||
| |||
| |||
class CheckpointCorrupt(Exception):
| |||
pass
| |||
| |||
| |||
class CheckpointTransitionError(Exception):
| |||
pass
| |||
| |||
| |||
STATUS_INITIALIZED = "INITIALIZED"
| |||
STATUS_RUNNING = "RUNNING"
| |||
STATUS_CHUNK_COMMITTED = "CHUNK_COMMITTED"
| |||
STATUS_PAUSED = "PAUSED"
| |||
STATUS_FAILED = "FAILED"
| |||
STATUS_COMPLETED = "COMPLETED"
| |||
STATUS_RESUME_BLOCKED = "RESUME_BLOCKED"
| |||
| |||
ALL_STATUSES = (STATUS_INITIALIZED, STATUS_RUNNING, STATUS_CHUNK_COMMITTED,
| |||
STATUS_PAUSED, STATUS_FAILED, STATUS_COMPLETED,
| |||
STATUS_RESUME_BLOCKED)
| |||
| |||
# Allowed transition table (spec 13.3). Key: (from, to) -> allowed.
| |||
_ALLOWED = {
| |||
(STATUS_INITIALIZED, STATUS_RUNNING): True,
| |||
(STATUS_INITIALIZED, STATUS_PAUSED): True,
| |||
(STATUS_INITIALIZED, STATUS_FAILED): True,
| |||
(STATUS_INITIALIZED, STATUS_RESUME_BLOCKED): True,
| |||
(STATUS_INITIALIZED, STATUS_COMPLETED): True, # empty-source degenerate path
| |||
(STATUS_RUNNING, STATUS_CHUNK_COMMITTED): True,
| |||
(STATUS_RUNNING, STATUS_PAUSED): True,
| |||
(STATUS_RUNNING, STATUS_FAILED): True,
| |||
(STATUS_RUNNING, STATUS_RESUME_BLOCKED): True,
| |||
(STATUS_CHUNK_COMMITTED, STATUS_RUNNING): True,
| |||
(STATUS_CHUNK_COMMITTED, STATUS_PAUSED): True,
| |||
(STATUS_CHUNK_COMMITTED, STATUS_FAILED): True,
| |||
(STATUS_CHUNK_COMMITTED, STATUS_RESUME_BLOCKED): True,
| |||
(STATUS_CHUNK_COMMITTED, STATUS_COMPLETED): True,
| |||
(STATUS_PAUSED, STATUS_RUNNING): True,
| |||
(STATUS_PAUSED, STATUS_FAILED): True,
| |||
(STATUS_PAUSED, STATUS_RESUME_BLOCKED): True,
| |||
(STATUS_FAILED, STATUS_RESUME_BLOCKED): True,
| |||
(STATUS_FAILED, STATUS_RUNNING): True, # after human resolution, explicit resume
| |||
(STATUS_RESUME_BLOCKED, STATUS_RESUME_BLOCKED): True,
| |||
(STATUS_RESUME_BLOCKED, STATUS_RUNNING): True, # human gate re-init path
| |||
}
| |||
| |||
| |||
def empty_malformed_counts():
| |||
"""Cumulative malformed counter skeleton shared with parse.py classes."""
| |||
cls = {}
| |||
for c in _malformed_classes():
| |||
cls[c] = 0
| |||
return cls
| |||
| |||
| |||
def _malformed_classes():
| |||
from .parse import ALL_MALFORMED_CLASSES
| |||
return ALL_MALFORMED_CLASSES
| |||
| |||
| |||
def make_checkpoint(payload):
| |||
"""Inject checkpoint_hash over the payload minus the hash field."""
| |||
body = {k: v for k, v in payload.items() if k != "checkpoint_hash"}
| |||
payload = dict(payload)
| |||
payload["checkpoint_hash"] = canonical_json_sha256(body)
| |||
return payload
| |||
| |||
| |||
def validate_transition(current, target):
| |||
if current == target:
| |||
return
| |||
if not _ALLOWED.get((current, target)):
| |||
raise CheckpointTransitionError(
| |||
"illegal checkpoint transition %s -> %s" % (current, target))
| |||
| |||
| |||
def write_checkpoint(path, payload):
| |||
os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True)
| |||
atomic_write_json(path, make_checkpoint(payload))
| |||
| |||
| |||
def load_checkpoint(path):
| |||
"""Read + verify checkpoint_hash. Raises CheckpointCorrupt on mismatch."""
| |||
import json
| |||
if not os.path.exists(path):
| |||
raise CheckpointCorrupt("checkpoint file missing: %s" % path)
| |||
with open(path, "rb") as fh:
| |||
text = fh.read()
| |||
try:
| |||
payload = json.loads(text.decode("utf-8"))
| |||
except Exception as e:
| |||
raise CheckpointCorrupt("checkpoint unparseable: %s" % e)
| |||
body = {k: v for k, v in payload.items() if k != "checkpoint_hash"}
| |||
expected = payload.get("checkpoint_hash")
| |||
actual = canonical_json_sha256(body)
| |||
if expected != actual:
| |||
raise CheckpointCorrupt("checkpoint_hash mismatch")
| |||
return payload
| |||
| |||
| |||
def checkpoint_hash_of(payload):
| |||
body = {k: v for k, v in payload.items() if k != "checkpoint_hash"}
| |||
return canonical_json_sha256(body)
|