fix(jobs): bind completed effects to checkpoints

This commit is contained in:
2026-07-12 04:54:00 +02:00
parent c964920f16
commit b6a52995ae
5 changed files with 174 additions and 7 deletions
+15 -1
View File
@@ -15,6 +15,7 @@ _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"]
EffectState = Literal["intent", "completed"]
def utc_now() -> datetime:
@@ -110,13 +111,22 @@ class StageRun(_FrozenModel):
started_at: datetime | None = None
finished_at: datetime | None = None
error: StageError | None = None
effect_state: EffectState | None = None
artifact_manifest_digest: str | None = None
artifact_files: tuple[str, ...] = ()
_name_key = field_validator("name")(_validate_job_key)
_artifact_digest = field_validator("artifact_manifest_digest")(
lambda value: value if value is None or _FINGERPRINT.fullmatch(value) else _invalid_fingerprint()
)
@model_validator(mode="after")
def state_shape(self) -> "StageRun":
if self.status == "pending" and any(
value is not None for value in (self.started_at, self.finished_at, self.error)
value is not None for value in (
self.started_at, self.finished_at, self.error, self.effect_state,
self.artifact_manifest_digest,
)
):
raise ValueError("pending stage cannot contain timestamps or error")
if self.status == "running" and (
@@ -131,6 +141,10 @@ class StageRun(_FrozenModel):
self.started_at is None or self.finished_at is None or self.error is None
):
raise ValueError("failed stage requires timestamps and safe error")
if (self.effect_state is None) != (self.artifact_manifest_digest is None):
raise ValueError("effect state and artifact manifest digest must be persisted together")
if self.artifact_files and self.effect_state is None:
raise ValueError("artifact files require a persisted effect state")
return self
+62 -6
View File
@@ -35,6 +35,14 @@ class JobContext:
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]
@@ -45,7 +53,7 @@ def _artifact_digest(path: Path) -> dict[str, Any]:
return {"sha256": hashlib.sha256(payload).hexdigest(), "size": len(payload)}
def _seal_artifacts(context: JobContext, stage: str, result: Any, spec: JobSpec) -> None:
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)
@@ -68,14 +76,16 @@ def _seal_artifacts(context: JobContext, stage: str, result: Any, spec: JobSpec)
prior["required"].remove(relative)
prior["files"].pop(relative, None)
manifest["stages"][stage] = {"required": list(required), "files": files}
_atomic_write(manifest_path, json.dumps(manifest, sort_keys=True, separators=(",", ":")) + "\n")
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."""
_seal_artifacts(context, stage, StageArtifacts(required), spec)
context.record_artifacts(stage, required, "intent")
def _validate_artifacts(run_dir: Path, spec: JobSpec, source: JobRun) -> set[str]:
@@ -86,6 +96,7 @@ def _validate_artifacts(run_dir: Path, spec: JobSpec, source: JobRun) -> set[str
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"])
@@ -100,6 +111,16 @@ def _validate_artifacts(run_dir: Path, spec: JobSpec, source: JobRun) -> set[str
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
@@ -113,7 +134,10 @@ def _validate_artifacts(run_dir: Path, spec: JobSpec, source: JobRun) -> set[str
allowed.add("compensated.json")
if actual != allowed:
raise ValueError
return sealed
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
@@ -199,6 +223,9 @@ def _resume_run(
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))
@@ -254,9 +281,36 @@ def run_job(
if source_artifacts.exists():
shutil.copytree(source_artifacts, run_dir / "artifacts")
_persist(checkpoint_path, run)
context = JobContext(run_id, spec.job_type, spec.dry_run, spec.workspace_root, run_dir)
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(
@@ -285,7 +339,9 @@ def run_job(
)
_persist(checkpoint_path, run)
break
_seal_artifacts(context, stage.name, stage_result, spec)
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()})