"""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 ( 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: """Stage a complete artifact bundle, then publish it through one atomic pointer.""" def __init__( self, *, workspace_id: str, workspace_root: Path, config_fingerprint: str, input_fingerprint: str, 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.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") def fingerprint(value: str) -> str: return "sha256:" + hashlib.sha256(value.encode("utf-8")).hexdigest()