fix(jobs): harden resume locks and durability

This commit is contained in:
2026-07-12 04:05:43 +02:00
parent a4acee4c70
commit 9f069cdd5b
6 changed files with 348 additions and 24 deletions
+54 -6
View File
@@ -6,6 +6,7 @@ import fcntl
import hashlib
import os
import re
import stat
from pathlib import Path
from types import TracebackType
@@ -36,13 +37,42 @@ class WorkspaceJobLock:
def acquire(self) -> "WorkspaceJobLock":
if self._fd is not None:
raise RuntimeError("job lock is already held by this object")
self.path.parent.mkdir(parents=True, exist_ok=True)
fd = os.open(self.path, os.O_RDWR | os.O_CREAT, 0o600)
root_fd = os.open(self.path.parents[2], os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
try:
fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
except BlockingIOError as error:
os.close(fd)
raise JobAlreadyRunningError("this workspace job is already running") from error
jobs_fd = _open_owned_directory(root_fd, ".tht-jobs")
try:
locks_fd = _open_owned_directory(jobs_fd, ".locks")
try:
fd = os.open(
self.path.name,
os.O_RDWR | os.O_CREAT | os.O_NOFOLLOW | os.O_CLOEXEC,
0o600,
dir_fd=locks_fd,
)
try:
info = os.fstat(fd)
if (
not stat.S_ISREG(info.st_mode)
or info.st_uid != os.getuid()
or info.st_nlink != 1
):
raise OSError("unsafe job lock file")
os.fchmod(fd, 0o600)
try:
fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
except BlockingIOError as error:
raise JobAlreadyRunningError(
"this workspace job is already running"
) from error
except BaseException:
os.close(fd)
raise
finally:
os.close(locks_fd)
finally:
os.close(jobs_fd)
finally:
os.close(root_fd)
self._fd = fd
return self
@@ -65,3 +95,21 @@ class WorkspaceJobLock:
traceback: TracebackType | None,
) -> None:
self.release()
def _open_owned_directory(parent_fd: int, name: str) -> int:
try:
os.mkdir(name, 0o700, dir_fd=parent_fd)
os.fsync(parent_fd)
except FileExistsError:
pass
fd = os.open(name, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW, dir_fd=parent_fd)
try:
info = os.fstat(fd)
if not stat.S_ISDIR(info.st_mode) or info.st_uid != os.getuid():
raise OSError("unsafe job lock directory")
os.fchmod(fd, 0o700)
except BaseException:
os.close(fd)
raise
return fd
+40 -1
View File
@@ -12,6 +12,7 @@ from pydantic import BaseModel, ConfigDict, Field, field_serializer, field_valid
_JOB_KEY = re.compile(r"^[a-z][a-z0-9_-]{0,63}$")
_RUN_ID = re.compile(r"^[0-9a-f]{32}$")
_FINGERPRINT = re.compile(r"^sha256:[0-9a-f]{64}$")
JobStatus = Literal["pending", "running", "succeeded", "failed"]
StageStatus = Literal["pending", "running", "succeeded", "failed"]
@@ -54,18 +55,38 @@ class JobSpec(_FrozenModel):
workspace_id: str
job_type: str
workspace_root: Path = Field(exclude=True)
spec_version: str = Field(min_length=1, max_length=64)
pipeline_version: str = Field(min_length=1, max_length=64)
config_fingerprint: str
input_fingerprint: str
stage_ids: tuple[str, ...]
dry_run: bool = False
resume_run_id: str | None = None
_workspace_key = field_validator("workspace_id")(_validate_job_key)
_job_type_key = field_validator("job_type")(_validate_job_key)
_version_keys = field_validator("spec_version", "pipeline_version")(_validate_job_key)
_resume_id = field_validator("resume_run_id")(_validate_run_id)
_config_fingerprint = field_validator("config_fingerprint")(
lambda value: value if _FINGERPRINT.fullmatch(value) else _invalid_fingerprint()
)
_input_fingerprint = field_validator("input_fingerprint")(
lambda value: value if _FINGERPRINT.fullmatch(value) else _invalid_fingerprint()
)
_stage_ids = field_validator("stage_ids")(
lambda values: tuple(_validate_job_key(value) for value in values)
)
def model_copy(self, *, update=None, deep: bool = False) -> Self:
data = {
"workspace_id": self.workspace_id,
"job_type": self.job_type,
"workspace_root": self.workspace_root,
"spec_version": self.spec_version,
"pipeline_version": self.pipeline_version,
"config_fingerprint": self.config_fingerprint,
"input_fingerprint": self.input_fingerprint,
"stage_ids": self.stage_ids,
"dry_run": self.dry_run,
"resume_run_id": self.resume_run_id,
}
@@ -78,7 +99,8 @@ class JobSpec(_FrozenModel):
class StageError(_FrozenModel):
category: str = Field(pattern=r"^[A-Za-z][A-Za-z0-9_]{0,127}$")
category: Literal["internal"] = "internal"
code: Literal["stage_exception"] = "stage_exception"
message: Literal["stage execution failed"] = "stage execution failed"
@@ -97,7 +119,13 @@ class JobRun(_FrozenModel):
schema_version: Literal[1] = 1
run_id: str
compatibility_fingerprint: str
workspace_fingerprint: str
job_type: str
spec_version: str
pipeline_version: str
config_fingerprint: str
input_fingerprint: str
dry_run: bool
status: JobStatus
started_at: datetime
@@ -106,9 +134,20 @@ class JobRun(_FrozenModel):
stages: tuple[StageRun, ...] = ()
_run_id = field_validator("run_id")(_validate_run_id)
_compatibility = field_validator("compatibility_fingerprint", "workspace_fingerprint")(
lambda value: value if _FINGERPRINT.fullmatch(value) else _invalid_fingerprint()
)
_input_fingerprints = field_validator("config_fingerprint", "input_fingerprint")(
lambda value: value if _FINGERPRINT.fullmatch(value) else _invalid_fingerprint()
)
_job_type = field_validator("job_type")(_validate_job_key)
_persisted_versions = field_validator("spec_version", "pipeline_version")(_validate_job_key)
_resumed_from = field_validator("resumed_from")(_validate_run_id)
class JobReport(JobRun):
"""Public machine-readable terminal report (contains no paths or stage outputs)."""
def _invalid_fingerprint():
raise ValueError("fingerprint must be sha256 followed by 64 lowercase hexadecimal characters")
+82 -14
View File
@@ -3,8 +3,10 @@
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
@@ -33,7 +35,6 @@ 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:
@@ -66,20 +67,21 @@ def _load_checkpoint(path: Path) -> JobRun:
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]
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(),
@@ -89,9 +91,14 @@ def _new_run(spec: JobSpec, run_id: str, stages: Sequence[Stage]) -> JobRun:
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]
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:
@@ -103,7 +110,13 @@ def _resume_run(spec: JobSpec, run_id: str, stages: Sequence[Stage], source: Job
)
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(),
@@ -118,11 +131,20 @@ def run_job(spec: JobSpec, stages: Sequence[Stage]) -> JobReport:
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)
@@ -140,12 +162,12 @@ def run_job(spec: JobSpec, stages: Sequence[Stage]) -> JobReport:
_persist(checkpoint_path, run)
try:
stage_callable(context)
except Exception as error:
except Exception:
failed = stage.model_copy(
update={
"status": "failed",
"finished_at": utc_now(),
"error": StageError(category=type(error).__name__),
"error": StageError(),
}
)
run = run.model_copy(
@@ -169,3 +191,49 @@ def run_job(spec: JobSpec, stages: Sequence[Stage]) -> JobReport:
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)