"""Resumable stage runner with durable atomic checkpoints and reports.""" from __future__ import annotations import json import os import uuid 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: path.parent.mkdir(parents=True, exist_ok=True) 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 _stage_name(stage: Stage) -> str: name = getattr(stage, "__name__", "") if name == "": name = "stage" return name.replace("_", "-") def _new_run(spec: JobSpec, run_id: str, stages: Sequence[Stage]) -> JobRun: names = [_stage_name(stage) for stage in stages] if len(names) != len(set(names)): raise ValueError("stage names must be unique") return JobRun( run_id=run_id, job_type=spec.job_type, 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: if source.job_type != spec.job_type: raise CorruptCheckpointError("checkpoint job type does not match resume request") requested_names = [_stage_name(stage) for stage in stages] 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, job_type=spec.job_type, 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 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" 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 as error: failed = stage.model_copy( update={ "status": "failed", "finished_at": utc_now(), "error": StageError(category=type(error).__name__), } ) 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