428 lines
17 KiB
Python
428 lines
17 KiB
Python
"""Resumable stage runner with durable atomic checkpoints and reports."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import hashlib
|
|
import os
|
|
import uuid
|
|
import stat
|
|
import shutil
|
|
from collections.abc import Callable, Sequence
|
|
from dataclasses import dataclass
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
from pydantic import ValidationError
|
|
|
|
from tht.jobs.locking import WorkspaceJobLock
|
|
from tht.jobs.models import JobReport, JobRun, JobSpec, StageError, StageRun, utc_now
|
|
|
|
|
|
class CorruptCheckpointError(RuntimeError):
|
|
pass
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class StageArtifacts:
|
|
required: tuple[str, ...] = ()
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class JobContext:
|
|
run_id: str
|
|
job_type: str
|
|
dry_run: bool
|
|
workspace_root: Path
|
|
run_dir: Path
|
|
_record_artifacts: Callable[[str, tuple[str, ...], str], None] | None = None
|
|
|
|
def record_artifacts(
|
|
self, stage: str, required: tuple[str, ...], effect_state: str,
|
|
) -> None:
|
|
if self._record_artifacts is None:
|
|
raise RuntimeError("artifact recorder is unavailable")
|
|
self._record_artifacts(stage, required, effect_state)
|
|
|
|
|
|
Stage = Callable[[JobContext], Any]
|
|
|
|
|
|
def _artifact_digest(path: Path) -> dict[str, Any]:
|
|
payload = path.read_bytes()
|
|
return {"sha256": hashlib.sha256(payload).hexdigest(), "size": len(payload)}
|
|
|
|
|
|
def _seal_artifacts(context: JobContext, stage: str, result: Any, spec: JobSpec) -> str:
|
|
required = result.required if isinstance(result, StageArtifacts) else ()
|
|
root = context.run_dir / "artifacts"
|
|
root.mkdir(exist_ok=True)
|
|
manifest_path = root / "artifact-manifest.json"
|
|
manifest = json.loads(manifest_path.read_text()) if manifest_path.exists() else {
|
|
"schema_version": 1,
|
|
"spec_fingerprint": _compatibility_fingerprint(spec, list(spec.stage_ids)),
|
|
"stages": {},
|
|
}
|
|
files = {}
|
|
for relative in required:
|
|
candidate = root / relative
|
|
if Path(relative).is_absolute() or ".." in Path(relative).parts or candidate.is_symlink():
|
|
raise CorruptCheckpointError("artifact path is unsafe")
|
|
if not candidate.is_file():
|
|
raise CorruptCheckpointError("required stage artifact is missing")
|
|
files[relative] = _artifact_digest(candidate)
|
|
for prior in manifest["stages"].values():
|
|
if relative in prior.get("required", []):
|
|
prior["required"].remove(relative)
|
|
prior["files"].pop(relative, None)
|
|
manifest["stages"][stage] = {"required": list(required), "files": files}
|
|
canonical = json.dumps(manifest, sort_keys=True, separators=(",", ":")) + "\n"
|
|
_atomic_write(manifest_path, canonical)
|
|
return _value_fingerprint(canonical)
|
|
|
|
|
|
def seal_stage_artifacts(
|
|
context: JobContext, stage: str, required: tuple[str, ...], spec: JobSpec,
|
|
) -> None:
|
|
"""Durably record external-effect intent before a stage performs that effect."""
|
|
context.record_artifacts(stage, required, "intent")
|
|
|
|
|
|
def _validate_artifacts(run_dir: Path, spec: JobSpec, source: JobRun) -> set[str]:
|
|
root = run_dir / "artifacts"
|
|
manifest_path = root / "artifact-manifest.json"
|
|
successful = {stage.name for stage in source.stages if stage.status == "succeeded"}
|
|
if not successful and not manifest_path.exists():
|
|
return set()
|
|
try:
|
|
manifest = json.loads(manifest_path.read_text())
|
|
root_digest = _value_fingerprint(manifest_path.read_text())
|
|
if manifest["spec_fingerprint"] != _compatibility_fingerprint(spec, list(spec.stage_ids)):
|
|
raise ValueError
|
|
sealed = set(manifest["stages"])
|
|
allowed = {"artifact-manifest.json"}
|
|
for stage, record in manifest["stages"].items():
|
|
for relative in record["required"]:
|
|
candidate = root / relative
|
|
if Path(relative).is_absolute() or ".." in Path(relative).parts or candidate.is_symlink():
|
|
raise ValueError
|
|
if not candidate.is_file() or _artifact_digest(candidate) != record["files"][relative]:
|
|
raise ValueError
|
|
allowed.add(relative)
|
|
if not successful.issubset(sealed):
|
|
raise ValueError
|
|
for stage in source.stages:
|
|
if stage.effect_state is None:
|
|
continue
|
|
record = manifest["stages"].get(stage.name)
|
|
if (
|
|
stage.artifact_manifest_digest != root_digest
|
|
or record is None
|
|
or tuple(record["required"]) != stage.artifact_files
|
|
):
|
|
raise ValueError
|
|
entries = list(root.iterdir())
|
|
if any(path.is_symlink() or not path.is_file() for path in entries):
|
|
raise ValueError
|
|
actual = {path.name for path in entries}
|
|
incomplete = next((stage for stage in source.stages if stage.status != "succeeded"), None)
|
|
marker = root / "compensated.json"
|
|
if incomplete is not None and incomplete.status in {"failed", "running"} and marker.is_file():
|
|
payload = json.loads(marker.read_text())
|
|
if not isinstance(payload.get("generation"), str):
|
|
raise ValueError
|
|
allowed.add("compensated.json")
|
|
if actual != allowed:
|
|
raise ValueError
|
|
return {
|
|
stage.name for stage in source.stages
|
|
if stage.effect_state == "completed"
|
|
}
|
|
except (OSError, KeyError, TypeError, ValueError, json.JSONDecodeError) as error:
|
|
raise CorruptCheckpointError("resume artifact manifest is invalid") from error
|
|
|
|
|
|
def _atomic_write(path: Path, payload: str) -> None:
|
|
temporary = path.with_name(f".{path.name}.{uuid.uuid4().hex}.tmp")
|
|
fd = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600)
|
|
try:
|
|
with os.fdopen(fd, "w", encoding="utf-8") as stream:
|
|
stream.write(payload)
|
|
stream.flush()
|
|
os.fsync(stream.fileno())
|
|
os.replace(temporary, path)
|
|
directory_fd = os.open(path.parent, os.O_RDONLY)
|
|
try:
|
|
os.fsync(directory_fd)
|
|
finally:
|
|
os.close(directory_fd)
|
|
except BaseException:
|
|
try:
|
|
temporary.unlink()
|
|
except FileNotFoundError:
|
|
pass
|
|
raise
|
|
|
|
|
|
def _persist(path: Path, run: JobRun) -> None:
|
|
_atomic_write(path, run.model_dump_json(indent=2) + "\n")
|
|
|
|
|
|
def _load_checkpoint(path: Path) -> JobRun:
|
|
try:
|
|
return JobRun.model_validate_json(path.read_text(encoding="utf-8"))
|
|
except (OSError, ValidationError, ValueError, json.JSONDecodeError) as error:
|
|
raise CorruptCheckpointError("checkpoint is invalid and cannot be resumed") from error
|
|
|
|
|
|
def _new_run(spec: JobSpec, run_id: str, stages: Sequence[Stage]) -> JobRun:
|
|
names = list(spec.stage_ids)
|
|
if len(names) != len(stages):
|
|
raise ValueError("stage_ids must identify every stage exactly once")
|
|
if len(names) != len(set(names)):
|
|
raise ValueError("stage names must be unique")
|
|
return JobRun(
|
|
run_id=run_id,
|
|
compatibility_fingerprint=_compatibility_fingerprint(spec, names),
|
|
workspace_fingerprint=_value_fingerprint(spec.workspace_id),
|
|
job_type=spec.job_type,
|
|
spec_version=spec.spec_version,
|
|
pipeline_version=spec.pipeline_version,
|
|
config_fingerprint=spec.config_fingerprint,
|
|
input_fingerprint=spec.input_fingerprint,
|
|
dry_run=spec.dry_run,
|
|
status="running",
|
|
started_at=utc_now(),
|
|
resumed_from=spec.resume_run_id,
|
|
stages=tuple(StageRun(name=name) for name in names),
|
|
)
|
|
|
|
|
|
def _resume_run(
|
|
spec: JobSpec, run_id: str, stages: Sequence[Stage], source: JobRun,
|
|
effect_completed: set[str] | None = None,
|
|
) -> JobRun:
|
|
requested_names = list(spec.stage_ids)
|
|
if len(requested_names) != len(stages) or len(requested_names) != len(set(requested_names)):
|
|
raise CorruptCheckpointError("resume checkpoint is incompatible with requested stages")
|
|
source_fingerprint = _source_compatibility_fingerprint(source)
|
|
if source.compatibility_fingerprint != source_fingerprint:
|
|
raise CorruptCheckpointError("resume checkpoint compatibility fingerprint is invalid")
|
|
expected = _compatibility_fingerprint(spec, requested_names)
|
|
if source_fingerprint != expected or [stage.name for stage in source.stages] != requested_names:
|
|
raise CorruptCheckpointError(
|
|
"resume checkpoint is incompatible; start an intentional new run without resume"
|
|
)
|
|
source_by_name = {stage.name: stage for stage in source.stages}
|
|
resumed_stages = []
|
|
for name in requested_names:
|
|
previous = source_by_name.get(name)
|
|
if previous is not None and previous.status == "succeeded":
|
|
resumed_stages.append(previous)
|
|
elif previous is not None and previous.status == "running" and name in (effect_completed or set()):
|
|
resumed_stages.append(StageRun(
|
|
name=name, status="succeeded", started_at=previous.started_at or utc_now(),
|
|
finished_at=utc_now(),
|
|
effect_state="completed",
|
|
artifact_manifest_digest=previous.artifact_manifest_digest,
|
|
artifact_files=previous.artifact_files,
|
|
))
|
|
else:
|
|
resumed_stages.append(StageRun(name=name))
|
|
return JobRun(
|
|
run_id=run_id,
|
|
compatibility_fingerprint=source.compatibility_fingerprint,
|
|
workspace_fingerprint=source.workspace_fingerprint,
|
|
job_type=spec.job_type,
|
|
spec_version=spec.spec_version,
|
|
pipeline_version=spec.pipeline_version,
|
|
config_fingerprint=spec.config_fingerprint,
|
|
input_fingerprint=spec.input_fingerprint,
|
|
dry_run=spec.dry_run,
|
|
status="running",
|
|
started_at=utc_now(),
|
|
resumed_from=source.run_id,
|
|
stages=tuple(resumed_stages),
|
|
)
|
|
|
|
|
|
def run_job(
|
|
spec: JobSpec, stages: Sequence[Stage], *,
|
|
after_stage_return: Callable[[JobContext, str], Any] | None = None,
|
|
reconcile_effects: Callable[[JobRun, Path], set[str]] | None = None,
|
|
) -> JobReport:
|
|
"""Run stages once, returning a terminal report instead of leaking stage exceptions."""
|
|
with WorkspaceJobLock(spec.workspace_root, spec.workspace_id, spec.job_type):
|
|
jobs_root = spec.workspace_root / ".tht-jobs" / spec.job_type / "runs"
|
|
if spec.resume_run_id is None:
|
|
source = None
|
|
else:
|
|
source_path = jobs_root / spec.resume_run_id / "checkpoint.json"
|
|
if not source_path.exists():
|
|
matches = list(
|
|
(spec.workspace_root / ".tht-jobs").glob(
|
|
f"*/runs/{spec.resume_run_id}/checkpoint.json"
|
|
)
|
|
)
|
|
if len(matches) == 1:
|
|
source_path = matches[0]
|
|
source = _load_checkpoint(source_path)
|
|
_validate_resume_source(spec, stages, source)
|
|
effect_completed = _validate_artifacts(source_path.parent, spec, source)
|
|
if reconcile_effects is not None:
|
|
effect_completed |= reconcile_effects(source, source_path.parent)
|
|
|
|
run_id = uuid.uuid4().hex
|
|
run_dir = jobs_root / run_id
|
|
_prepare_run_directory(spec.workspace_root, spec.job_type, run_id)
|
|
checkpoint_path = run_dir / "checkpoint.json"
|
|
if source is None:
|
|
run = _new_run(spec, run_id, stages)
|
|
else:
|
|
run = _resume_run(spec, run_id, stages, source, effect_completed)
|
|
source_artifacts = jobs_root / source.run_id / "artifacts"
|
|
if source_artifacts.exists():
|
|
shutil.copytree(source_artifacts, run_dir / "artifacts")
|
|
_persist(checkpoint_path, run)
|
|
current_index = -1
|
|
|
|
def record_artifacts(stage_name: str, required: tuple[str, ...], effect_state: str) -> None:
|
|
nonlocal run
|
|
if current_index < 0 or run.stages[current_index].name != stage_name:
|
|
raise CorruptCheckpointError("artifact producer does not match running stage")
|
|
digest = _seal_artifacts(context, stage_name, StageArtifacts(required), spec)
|
|
manifest = json.loads((run_dir / "artifacts" / "artifact-manifest.json").read_text())
|
|
updated = []
|
|
for position, value in enumerate(run.stages):
|
|
record = manifest["stages"].get(value.name)
|
|
if record is not None and (value.status == "succeeded" or position == current_index):
|
|
state = effect_state if position == current_index else value.effect_state
|
|
updated.append(value.model_copy(update={
|
|
"effect_state": state,
|
|
"artifact_manifest_digest": digest,
|
|
"artifact_files": tuple(record["required"]),
|
|
}))
|
|
else:
|
|
updated.append(value)
|
|
run = run.model_copy(update={"stages": tuple(updated)})
|
|
_persist(checkpoint_path, run)
|
|
|
|
context = JobContext(
|
|
run_id, spec.job_type, spec.dry_run, spec.workspace_root, run_dir,
|
|
record_artifacts,
|
|
)
|
|
|
|
for index, stage_callable in enumerate(stages):
|
|
current_index = index
|
|
if run.stages[index].status == "succeeded":
|
|
continue
|
|
stage = run.stages[index].model_copy(
|
|
update={"status": "running", "started_at": utc_now()}
|
|
)
|
|
run = run.model_copy(
|
|
update={"stages": run.stages[:index] + (stage,) + run.stages[index + 1 :]}
|
|
)
|
|
_persist(checkpoint_path, run)
|
|
try:
|
|
stage_result = stage_callable(context)
|
|
except Exception:
|
|
failed = stage.model_copy(
|
|
update={
|
|
"status": "failed",
|
|
"finished_at": utc_now(),
|
|
"error": StageError(),
|
|
}
|
|
)
|
|
run = run.model_copy(
|
|
update={
|
|
"status": "failed",
|
|
"finished_at": utc_now(),
|
|
"stages": run.stages[:index] + (failed,) + run.stages[index + 1 :],
|
|
}
|
|
)
|
|
_persist(checkpoint_path, run)
|
|
break
|
|
required = stage_result.required if isinstance(stage_result, StageArtifacts) else ()
|
|
context.record_artifacts(stage.name, required, "completed")
|
|
stage = run.stages[index]
|
|
if after_stage_return is not None:
|
|
after_stage_return(context, stage.name)
|
|
succeeded = stage.model_copy(update={"status": "succeeded", "finished_at": utc_now()})
|
|
run = run.model_copy(
|
|
update={"stages": run.stages[:index] + (succeeded,) + run.stages[index + 1 :]}
|
|
)
|
|
_persist(checkpoint_path, run)
|
|
else:
|
|
run = run.model_copy(update={"status": "succeeded", "finished_at": utc_now()})
|
|
_persist(checkpoint_path, run)
|
|
|
|
report = JobReport.model_validate(run.model_dump())
|
|
_atomic_write(run_dir / "report.json", report.model_dump_json(indent=2) + "\n")
|
|
return report
|
|
|
|
|
|
def _value_fingerprint(value: str) -> str:
|
|
return "sha256:" + hashlib.sha256(value.encode("utf-8")).hexdigest()
|
|
|
|
|
|
def _compatibility_fingerprint(spec: JobSpec, stage_ids: list[str]) -> str:
|
|
payload = {
|
|
"schema_version": 1,
|
|
"workspace": _value_fingerprint(spec.workspace_id),
|
|
"job_type": spec.job_type,
|
|
"dry_run": spec.dry_run,
|
|
"spec_version": spec.spec_version,
|
|
"pipeline_version": spec.pipeline_version,
|
|
"config_fingerprint": spec.config_fingerprint,
|
|
"input_fingerprint": spec.input_fingerprint,
|
|
"stage_ids": stage_ids,
|
|
}
|
|
canonical = json.dumps(payload, sort_keys=True, separators=(",", ":"))
|
|
return _value_fingerprint(canonical)
|
|
|
|
|
|
def _source_compatibility_fingerprint(source: JobRun) -> str:
|
|
payload = {
|
|
"schema_version": source.schema_version,
|
|
"workspace": source.workspace_fingerprint,
|
|
"job_type": source.job_type,
|
|
"dry_run": source.dry_run,
|
|
"spec_version": source.spec_version,
|
|
"pipeline_version": source.pipeline_version,
|
|
"config_fingerprint": source.config_fingerprint,
|
|
"input_fingerprint": source.input_fingerprint,
|
|
"stage_ids": [stage.name for stage in source.stages],
|
|
}
|
|
canonical = json.dumps(payload, sort_keys=True, separators=(",", ":"))
|
|
return _value_fingerprint(canonical)
|
|
|
|
|
|
def _validate_resume_source(spec: JobSpec, stages: Sequence[Stage], source: JobRun) -> None:
|
|
_resume_run(spec, "0" * 32, stages, source)
|
|
|
|
|
|
def _prepare_run_directory(workspace_root: Path, job_type: str, run_id: str) -> None:
|
|
parent_fd = os.open(workspace_root, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
|
|
try:
|
|
for component in (".tht-jobs", job_type, "runs", run_id):
|
|
try:
|
|
os.mkdir(component, 0o700, dir_fd=parent_fd)
|
|
os.fsync(parent_fd)
|
|
except FileExistsError:
|
|
pass
|
|
child_fd = os.open(
|
|
component,
|
|
os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW,
|
|
dir_fd=parent_fd,
|
|
)
|
|
info = os.fstat(child_fd)
|
|
if not stat.S_ISDIR(info.st_mode) or info.st_uid != os.getuid():
|
|
os.close(child_fd)
|
|
raise OSError("unsafe job run directory")
|
|
os.fchmod(child_fd, 0o700)
|
|
os.close(parent_fd)
|
|
parent_fd = child_fd
|
|
os.fsync(parent_fd)
|
|
finally:
|
|
os.close(parent_fd)
|