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

428 lines
18 KiB
Python

"""Resumable stage runner with durable atomic checkpoints and reports."""
from __future__ import annotations
import hashlib
import json
import os
import shutil
import stat
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 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: # noqa: BLE001 - stage failures are persisted as terminal reports
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)