"""Resumable stage runner with durable atomic checkpoints and reports.""" from __future__ import annotations import json import hashlib import os import uuid import stat 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 JobContext: run_id: str job_type: str dry_run: bool workspace_root: Path run_dir: Path Stage = Callable[[JobContext], Any] 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) -> 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") expected = _compatibility_fingerprint(spec, requested_names) if source.compatibility_fingerprint != expected: 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) resumed_stages.append( previous if previous is not None and previous.status == "succeeded" else 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]) -> 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" 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 spec.resume_run_id is None: run = _new_run(spec, run_id, stages) 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) run = _resume_run(spec, run_id, stages, source) _persist(checkpoint_path, run) context = JobContext(run_id, spec.job_type, spec.dry_run, spec.workspace_root, run_dir) for index, stage_callable in enumerate(stages): 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_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 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 _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)