Files
ThothII/harness/tht/jobs/runner.py
T

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