172 lines
5.9 KiB
Python
172 lines
5.9 KiB
Python
"""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 == "<lambda>":
|
|
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
|