Files
ThothII/harness/tht/jobs/dwh_pipeline.py
T

587 lines
23 KiB
Python

"""Crash-safe, resumable DWH catalog and LSH preprocessing stages."""
from __future__ import annotations
import hashlib
import json
import fcntl
import os
import re
import shutil
import stat
import uuid
from collections.abc import Callable
from dataclasses import dataclass
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}$")
_SAFE_FILE = re.compile(r"^[A-Za-z0-9_-]+\.(?:pkl|json)$")
GENERATION_MANIFEST = "generation-manifest.json"
@dataclass(frozen=True)
class DwhArtifactSnapshot:
generation: str | None
physical: Path
lsh_dir: Path
class DwhSnapshotLease:
def __init__(self, cfg) -> None:
self.cfg = cfg
self._fd: int | None = None
self.snapshot: DwhArtifactSnapshot | None = None
def __enter__(self) -> DwhArtifactSnapshot:
self._fd = _acquire_generation_lock(self.cfg.paths.artifacts.parent, exclusive=False)
try:
self.snapshot = resolve_dwh_snapshot(self.cfg)
return self.snapshot
except BaseException:
self.__exit__(None, None, None)
raise
def __exit__(self, *_args) -> None:
if self._fd is not None:
fd, self._fd = self._fd, None
fcntl.flock(fd, fcntl.LOCK_UN)
os.close(fd)
def lease_dwh_snapshot(cfg) -> DwhSnapshotLease:
return DwhSnapshotLease(cfg)
def _acquire_generation_lock(workspace_root: Path, *, exclusive: bool) -> int:
root = workspace_root / ".tht-dwh"
DwhPreprocessPipeline._ensure_owned_dir(root)
fd = os.open(root / "generation.lock", os.O_RDWR | os.O_CREAT | os.O_NOFOLLOW, 0o600)
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 DWH generation lock")
fcntl.flock(fd, fcntl.LOCK_EX if exclusive else fcntl.LOCK_SH)
return fd
except BaseException:
os.close(fd)
raise
def _digest(path: Path) -> str:
return hashlib.sha256(path.read_bytes()).hexdigest()
def _read_owned(path: Path, *, readonly: bool) -> bytes:
fd = os.open(path, os.O_RDONLY | os.O_NOFOLLOW)
try:
info = os.fstat(fd)
if (
not stat.S_ISREG(info.st_mode)
or info.st_uid != os.getuid()
or info.st_nlink != 1
or (readonly and bool(info.st_mode & 0o222))
):
raise OSError("unsafe DWH generation file")
chunks = []
while chunk := os.read(fd, 1024 * 1024):
chunks.append(chunk)
return b"".join(chunks)
finally:
os.close(fd)
def _read_owned_at(directory_fd: int, name: str, *, readonly: bool) -> bytes:
fd = os.open(name, os.O_RDONLY | os.O_NOFOLLOW, dir_fd=directory_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
or (readonly and bool(info.st_mode & 0o222))
):
raise OSError("unsafe DWH generation file")
chunks = []
while chunk := os.read(fd, 1024 * 1024):
chunks.append(chunk)
return b"".join(chunks)
finally:
os.close(fd)
def validate_generation_fd(directory_fd: int, generation: str) -> dict:
try:
directory_info = os.fstat(directory_fd)
if (
not stat.S_ISDIR(directory_info.st_mode)
or directory_info.st_uid != os.getuid()
or stat.S_IMODE(directory_info.st_mode) != 0o700
):
raise ValueError
manifest = json.loads(
_read_owned_at(directory_fd, GENERATION_MANIFEST, readonly=True).decode("utf-8")
)
files = manifest["files"]
if (
manifest["generation"] != generation
or not isinstance(files, dict)
or not re.fullmatch(r"sha256:[0-9a-f]{64}", manifest["job_spec_fingerprint"])
or not re.fullmatch(r"[0-9a-f]{64}", manifest["artifact_manifest_sha256"])
):
raise ValueError
if set(os.listdir(directory_fd)) != set(files) | {GENERATION_MANIFEST}:
raise ValueError
for name, expected in files.items():
if name != "physical.yaml" and not _SAFE_FILE.fullmatch(name):
raise ValueError
payload = _read_owned_at(directory_fd, name, readonly=True)
if hashlib.sha256(payload).hexdigest() != expected:
raise ValueError
return manifest
except (OSError, KeyError, TypeError, ValueError, UnicodeDecodeError, json.JSONDecodeError) as error:
raise CorruptCheckpointError("published DWH generation is invalid") from error
def validate_generation(target: Path) -> dict:
try:
fd = os.open(target, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
except OSError as error:
raise CorruptCheckpointError("published DWH generation is invalid") from error
try:
return validate_generation_fd(fd, target.name)
finally:
os.close(fd)
def resolve_dwh_snapshot(cfg) -> DwhArtifactSnapshot:
target = active_generation_dir(cfg.paths.artifacts.parent)
if target is None:
return DwhArtifactSnapshot(
None, cfg.paths.artifacts / "mschema" / "physical.yaml", cfg.paths.indexes / "lsh"
)
validate_generation(target)
return DwhArtifactSnapshot(target.name, target / "physical.yaml", target)
def active_generation_dir(workspace_root: Path) -> Path | None:
pointer = workspace_root / ".tht-dwh" / "ACTIVE"
try:
generation = _read_owned(pointer, readonly=False).decode("utf-8").strip()
except FileNotFoundError:
return None
except (OSError, UnicodeDecodeError) as error:
raise CorruptCheckpointError("DWH ACTIVE pointer is invalid") from error
if not _RUN_ID.fullmatch(generation):
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,
retain_generations: int = 3,
) -> 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
if isinstance(retain_generations, bool) or retain_generations < 1:
raise ValueError("retain_generations must be positive")
self.retain_generations = retain_generations
if len(set(self.lsh_filenames)) != 3 or any(
not _SAFE_FILE.fullmatch(name) for name in self.lsh_filenames
):
raise ValueError("LSH filenames must be unique flat safe names")
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),
reconcile_effects=self._reconcile_effects,
)
def _reconcile_effects(self, source, run_dir: Path) -> set[str]:
running = next(
(stage for stage in source.stages if stage.status == "running" and stage.effect_state == "intent"),
None,
)
if running is None:
return set()
target = self.workspace_root / ".tht-dwh" / "generations" / source.run_id
validate_generation(target)
active = active_generation_dir(self.workspace_root)
if active != target:
raise CorruptCheckpointError("sealed DWH publication is not ACTIVE")
self._validate_published(target, run_dir / "artifacts", running.artifact_files)
return {running.name}
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)
lease_fd = _acquire_generation_lock(self.workspace_root, exclusive=True)
try:
self._publish(context.run_id, artifacts, required)
if self.after_publish is not None:
self.after_publish(context.run_id)
self._cleanup_generations()
finally:
fcntl.flock(lease_fd, fcntl.LOCK_UN)
os.close(lease_fd)
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())
destination.chmod(0o400)
manifest = {
"schema_version": 1,
"generation": generation,
"files": {name: _digest(artifacts / name) for name in required},
"job_spec_fingerprint": json.loads(
(artifacts / "artifact-manifest.json").read_text(encoding="utf-8")
)["spec_fingerprint"],
"artifact_manifest_sha256": _digest(
artifacts / "artifact-manifest.json"
),
}
manifest_path = temporary / GENERATION_MANIFEST
manifest_path.write_text(
json.dumps(manifest, sort_keys=True, separators=(",", ":")) + "\n",
encoding="utf-8",
)
manifest_path.chmod(0o400)
with manifest_path.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 = root / "ACTIVE"
try:
previous = pointer.read_text(encoding="utf-8")
except FileNotFoundError:
previous = None
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, pointer)
try:
self._fsync(root)
except BaseException:
self._restore_pointer(root, pointer, previous)
raise
except BaseException:
pointer_tmp.unlink(missing_ok=True)
raise
@staticmethod
def _validate_published(target: Path, artifacts: Path, required: tuple[str, ...]) -> None:
manifest = validate_generation(target)
if set(manifest["files"]) != set(required):
raise CorruptCheckpointError("published DWH generation is invalid")
if manifest["artifact_manifest_sha256"] != _digest(
artifacts / "artifact-manifest.json"
):
raise CorruptCheckpointError("published DWH job manifest digest mismatch")
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")
def _restore_pointer(self, root: Path, pointer: Path, previous: str | None) -> None:
if previous is None:
pointer.unlink(missing_ok=True)
else:
restore = root / f".ACTIVE.restore.{uuid.uuid4().hex}.tmp"
restore.write_text(previous, encoding="utf-8")
with restore.open("rb") as stream:
os.fsync(stream.fileno())
os.replace(restore, pointer)
self._fsync(root)
def _cleanup_generations(self) -> None:
root = self.workspace_root / ".tht-dwh" / "generations"
if not root.exists():
return
active = active_generation_dir(self.workspace_root)
active_name = active.name if active else None
protected = {active_name} if active_name else set()
runs = self.workspace_root / ".tht-jobs" / "dwh" / "runs"
for checkpoint in runs.glob("*/checkpoint.json") if runs.exists() else ():
try:
status = json.loads(checkpoint.read_text(encoding="utf-8"))["status"]
if status in {"running", "failed"}:
protected.add(checkpoint.parent.name)
except (OSError, KeyError, ValueError):
continue
generations = []
root_fd = os.open(root, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
try:
for name in os.listdir(root_fd):
if not _RUN_ID.fullmatch(name):
continue
try:
candidate_fd = os.open(
name, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW, dir_fd=root_fd
)
except OSError:
continue
try:
info = os.fstat(candidate_fd)
validate_generation_fd(candidate_fd, name)
except CorruptCheckpointError:
continue
finally:
os.close(candidate_fd)
generations.append((info.st_mtime_ns, name))
generations.sort(key=lambda value: (value[0], value[1]))
rollback = [value for value in generations if value[1] != active_name]
rollback_count = self.retain_generations - 1
keep_recent = {
name for _, name in (rollback[-rollback_count:] if rollback_count else ())
}
for _, name in generations:
if name in protected | keep_recent:
continue
self._safe_delete_generation(root_fd, name)
os.fsync(root_fd)
finally:
os.close(root_fd)
@staticmethod
def _safe_delete_generation(root_fd: int, name: str) -> None:
try:
generation_fd = os.open(
name, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW, dir_fd=root_fd
)
except OSError:
return
try:
info = os.fstat(generation_fd)
if not stat.S_ISDIR(info.st_mode) or info.st_uid != os.getuid():
return
entries = os.listdir(generation_fd)
opened = []
try:
for child in entries:
try:
fd = os.open(
child, os.O_RDONLY | os.O_NOFOLLOW, dir_fd=generation_fd
)
except OSError:
return
child_info = os.fstat(fd)
if (
not stat.S_ISREG(child_info.st_mode)
or child_info.st_uid != os.getuid()
or child_info.st_nlink != 1
):
os.close(fd)
return
opened.append((child, fd))
for child, fd in opened:
os.fchmod(fd, 0o600)
os.close(fd)
os.unlink(child, dir_fd=generation_fd)
opened.clear()
finally:
for _, fd in opened:
os.close(fd)
finally:
os.close(generation_fd)
try:
os.rmdir(name, dir_fd=root_fd)
except OSError:
return
@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()