fix(preprocess): publish DWH artifacts atomically
This commit is contained in:
@@ -1,20 +1,47 @@
|
||||
"""Resumable DWH catalog and LSH preprocessing stages."""
|
||||
"""Crash-safe, resumable DWH catalog and LSH preprocessing stages."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import stat
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
|
||||
from tht.jobs.models import JobReport, JobSpec
|
||||
from tht.jobs.runner import run_job
|
||||
from tht.jobs.runner import (
|
||||
CorruptCheckpointError,
|
||||
JobContext,
|
||||
StageArtifacts,
|
||||
run_job,
|
||||
seal_stage_artifacts,
|
||||
)
|
||||
|
||||
|
||||
DWH_STAGE_IDS = ("introspect", "lsh")
|
||||
_RUN_ID = re.compile(r"^[0-9a-f]{32}$")
|
||||
|
||||
|
||||
def active_generation_dir(workspace_root: Path) -> Path | None:
|
||||
pointer = workspace_root / ".tht-dwh" / "ACTIVE"
|
||||
try:
|
||||
generation = pointer.read_text(encoding="utf-8").strip()
|
||||
except FileNotFoundError:
|
||||
return None
|
||||
if not _RUN_ID.fullmatch(generation) or pointer.is_symlink():
|
||||
raise CorruptCheckpointError("DWH ACTIVE pointer is invalid")
|
||||
target = pointer.parent / "generations" / generation
|
||||
if not target.is_dir() or target.is_symlink():
|
||||
raise CorruptCheckpointError("active DWH generation is missing")
|
||||
return target
|
||||
|
||||
|
||||
class DwhPreprocessPipeline:
|
||||
"""Adapt existing DWH preprocessing operations to the shared job envelope."""
|
||||
"""Stage a complete artifact bundle, then publish it through one atomic pointer."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -23,42 +50,218 @@ class DwhPreprocessPipeline:
|
||||
workspace_root: Path,
|
||||
config_fingerprint: str,
|
||||
input_fingerprint: str,
|
||||
introspect: Callable[[], object],
|
||||
build_lsh: Callable[[], object],
|
||||
introspect: Callable[[Path], object],
|
||||
build_lsh: Callable[[Path, Path], object],
|
||||
lsh_filenames: tuple[str, str, str] | None = None,
|
||||
current_physical: Path | None = None,
|
||||
current_lsh_dir: Path | None = None,
|
||||
after_publish: Callable[[str], object] | None = None,
|
||||
) -> None:
|
||||
self.workspace_id = workspace_id
|
||||
self.workspace_root = workspace_root
|
||||
self.config_fingerprint = config_fingerprint
|
||||
self.input_fingerprint = input_fingerprint
|
||||
self._operations = {"introspect": introspect, "lsh": build_lsh}
|
||||
self.introspect = introspect
|
||||
self.build_lsh = build_lsh
|
||||
self.lsh_filenames = lsh_filenames or (
|
||||
f"{workspace_id}_lsh.pkl", f"{workspace_id}_minhashes.pkl",
|
||||
f"{workspace_id}_meta.json",
|
||||
)
|
||||
self.current_physical = current_physical
|
||||
self.current_lsh_dir = current_lsh_dir
|
||||
self.after_publish = after_publish
|
||||
|
||||
def run(
|
||||
self, steps: tuple[str, ...] = DWH_STAGE_IDS, *, resume_run_id: str | None = None
|
||||
) -> JobReport:
|
||||
self._validate_steps(steps)
|
||||
if resume_run_id is not None:
|
||||
self._validate_resume_publication(resume_run_id)
|
||||
spec = JobSpec(
|
||||
workspace_id=self.workspace_id,
|
||||
job_type="dwh",
|
||||
workspace_root=self.workspace_root,
|
||||
spec_version="jobs-v1",
|
||||
pipeline_version="dwh-v2",
|
||||
config_fingerprint=self.config_fingerprint,
|
||||
input_fingerprint=self.input_fingerprint,
|
||||
stage_ids=steps,
|
||||
resume_run_id=resume_run_id,
|
||||
)
|
||||
|
||||
def introspect_stage(context: JobContext):
|
||||
artifacts = self._artifacts(context)
|
||||
physical = artifacts / "physical.yaml"
|
||||
try:
|
||||
self.introspect(physical)
|
||||
except BaseException:
|
||||
physical.unlink(missing_ok=True)
|
||||
raise
|
||||
if steps[-1] == "introspect":
|
||||
self._copy_current_lsh(artifacts)
|
||||
return self._publish_stage(context, "introspect", spec)
|
||||
return StageArtifacts(("physical.yaml",))
|
||||
|
||||
def lsh_stage(context: JobContext):
|
||||
artifacts = self._artifacts(context)
|
||||
physical = artifacts / "physical.yaml"
|
||||
if not physical.exists():
|
||||
source = self.current_physical
|
||||
if source is None or not source.is_file():
|
||||
raise FileNotFoundError("physical catalog is missing")
|
||||
shutil.copyfile(source, physical)
|
||||
try:
|
||||
self.build_lsh(physical, artifacts)
|
||||
except BaseException:
|
||||
for name in self.lsh_filenames:
|
||||
(artifacts / name).unlink(missing_ok=True)
|
||||
raise
|
||||
return self._publish_stage(context, "lsh", spec)
|
||||
|
||||
implementations = {"introspect": introspect_stage, "lsh": lsh_stage}
|
||||
return run_job(spec, tuple(implementations[step] for step in steps))
|
||||
|
||||
def _validate_resume_publication(self, run_id: str) -> None:
|
||||
run_dir = self.workspace_root / ".tht-jobs" / "dwh" / "runs" / run_id
|
||||
try:
|
||||
checkpoint = json.loads((run_dir / "checkpoint.json").read_text(encoding="utf-8"))
|
||||
except (OSError, ValueError, TypeError) as error:
|
||||
raise CorruptCheckpointError("checkpoint is invalid and cannot be resumed") from error
|
||||
target = self.workspace_root / ".tht-dwh" / "generations" / run_id
|
||||
if not target.exists():
|
||||
if checkpoint.get("status") == "succeeded":
|
||||
raise CorruptCheckpointError("published DWH generation is missing")
|
||||
return
|
||||
try:
|
||||
manifest = json.loads(
|
||||
(run_dir / "artifacts" / "artifact-manifest.json").read_text(encoding="utf-8")
|
||||
)
|
||||
required = tuple(
|
||||
name
|
||||
for stage in manifest["stages"].values()
|
||||
for name in stage["required"]
|
||||
)
|
||||
except (OSError, KeyError, ValueError, TypeError) as error:
|
||||
raise CorruptCheckpointError("resume artifact manifest is invalid") from error
|
||||
self._validate_published(target, run_dir / "artifacts", required)
|
||||
|
||||
@staticmethod
|
||||
def _artifacts(context: JobContext) -> Path:
|
||||
root = context.run_dir / "artifacts"
|
||||
root.mkdir(exist_ok=True)
|
||||
return root
|
||||
|
||||
def _required(self, artifacts: Path) -> tuple[str, ...]:
|
||||
required = ("physical.yaml",) + tuple(
|
||||
name for name in self.lsh_filenames if (artifacts / name).is_file()
|
||||
)
|
||||
if not (artifacts / "physical.yaml").is_file():
|
||||
raise CorruptCheckpointError("staged physical catalog is missing")
|
||||
lsh_count = len(required) - 1
|
||||
if lsh_count not in (0, len(self.lsh_filenames)):
|
||||
raise CorruptCheckpointError("staged LSH artifact set is incomplete")
|
||||
return required
|
||||
|
||||
def _copy_current_lsh(self, artifacts: Path) -> None:
|
||||
if self.current_lsh_dir is None:
|
||||
return
|
||||
existing = [self.current_lsh_dir / name for name in self.lsh_filenames]
|
||||
if not any(path.exists() for path in existing):
|
||||
return
|
||||
if not all(path.is_file() for path in existing):
|
||||
raise CorruptCheckpointError("current LSH artifact set is incomplete")
|
||||
for source in existing:
|
||||
shutil.copyfile(source, artifacts / source.name)
|
||||
|
||||
def _publish_stage(self, context: JobContext, stage: str, spec: JobSpec):
|
||||
artifacts = self._artifacts(context)
|
||||
required = self._required(artifacts)
|
||||
seal_stage_artifacts(context, stage, required, spec)
|
||||
self._publish(context.run_id, artifacts, required)
|
||||
if self.after_publish is not None:
|
||||
self.after_publish(context.run_id)
|
||||
return StageArtifacts(required)
|
||||
|
||||
def _publish(self, generation: str, artifacts: Path, required: tuple[str, ...]) -> None:
|
||||
root = self.workspace_root / ".tht-dwh"
|
||||
generations = root / "generations"
|
||||
self._ensure_owned_dir(root)
|
||||
self._ensure_owned_dir(generations)
|
||||
target = generations / generation
|
||||
if target.exists():
|
||||
self._validate_published(target, artifacts, required)
|
||||
else:
|
||||
temporary = generations / f".{generation}.{uuid.uuid4().hex}.tmp"
|
||||
temporary.mkdir(mode=0o700)
|
||||
try:
|
||||
for name in required:
|
||||
destination = temporary / name
|
||||
shutil.copyfile(artifacts / name, destination)
|
||||
with destination.open("rb") as stream:
|
||||
os.fsync(stream.fileno())
|
||||
self._fsync(temporary)
|
||||
os.replace(temporary, target)
|
||||
self._fsync(generations)
|
||||
except BaseException:
|
||||
shutil.rmtree(temporary, ignore_errors=True)
|
||||
raise
|
||||
pointer_tmp = root / f".ACTIVE.{uuid.uuid4().hex}.tmp"
|
||||
fd = os.open(pointer_tmp, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600)
|
||||
try:
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as stream:
|
||||
stream.write(generation + "\n")
|
||||
stream.flush()
|
||||
os.fsync(stream.fileno())
|
||||
os.replace(pointer_tmp, root / "ACTIVE")
|
||||
self._fsync(root)
|
||||
except BaseException:
|
||||
pointer_tmp.unlink(missing_ok=True)
|
||||
raise
|
||||
|
||||
@staticmethod
|
||||
def _validate_published(target: Path, artifacts: Path, required: tuple[str, ...]) -> None:
|
||||
if (
|
||||
target.is_symlink()
|
||||
or not target.is_dir()
|
||||
or {path.name for path in target.iterdir()} != set(required)
|
||||
):
|
||||
raise CorruptCheckpointError("published DWH generation is invalid")
|
||||
for name in required:
|
||||
source, published = artifacts / name, target / name
|
||||
if published.is_symlink() or not published.is_file():
|
||||
raise CorruptCheckpointError("published DWH artifact is invalid")
|
||||
if hashlib.sha256(source.read_bytes()).digest() != hashlib.sha256(
|
||||
published.read_bytes()
|
||||
).digest():
|
||||
raise CorruptCheckpointError("published DWH artifact digest mismatch")
|
||||
|
||||
@staticmethod
|
||||
def _ensure_owned_dir(path: Path) -> None:
|
||||
try:
|
||||
path.mkdir(mode=0o700)
|
||||
except FileExistsError:
|
||||
pass
|
||||
info = path.lstat()
|
||||
if not stat.S_ISDIR(info.st_mode) or info.st_uid != os.getuid():
|
||||
raise OSError("unsafe DWH publication directory")
|
||||
path.chmod(0o700)
|
||||
|
||||
@staticmethod
|
||||
def _fsync(path: Path) -> None:
|
||||
fd = os.open(path, os.O_RDONLY | os.O_DIRECTORY)
|
||||
try:
|
||||
os.fsync(fd)
|
||||
finally:
|
||||
os.close(fd)
|
||||
|
||||
@staticmethod
|
||||
def _validate_steps(steps: tuple[str, ...]) -> None:
|
||||
if not steps or len(steps) != len(set(steps)) or any(
|
||||
step not in DWH_STAGE_IDS for step in steps
|
||||
):
|
||||
raise ValueError("DWH preprocessing steps must be unique introspect/lsh stages")
|
||||
if tuple(sorted(steps, key=DWH_STAGE_IDS.index)) != steps:
|
||||
raise ValueError("DWH preprocessing steps must follow introspect,lsh order")
|
||||
spec = JobSpec(
|
||||
workspace_id=self.workspace_id,
|
||||
job_type="dwh",
|
||||
workspace_root=self.workspace_root,
|
||||
spec_version="jobs-v1",
|
||||
pipeline_version="dwh-v1",
|
||||
config_fingerprint=self.config_fingerprint,
|
||||
input_fingerprint=self.input_fingerprint,
|
||||
stage_ids=steps,
|
||||
resume_run_id=resume_run_id,
|
||||
)
|
||||
def stage(operation):
|
||||
def execute(_context):
|
||||
return operation()
|
||||
|
||||
return execute
|
||||
|
||||
return run_job(spec, tuple(stage(self._operations[step]) for step in steps))
|
||||
|
||||
|
||||
def fingerprint(value: str) -> str:
|
||||
|
||||
Reference in New Issue
Block a user