fix(preprocess): anchor DWH retention validation
This commit is contained in:
@@ -260,6 +260,71 @@ def test_corrupt_newer_directory_does_not_consume_rollback_slot(tmp_path):
|
||||
assert corrupt.is_dir()
|
||||
|
||||
|
||||
def test_retention_n_counts_active_plus_n_minus_one_rollbacks_even_if_active_is_old(tmp_path):
|
||||
import os
|
||||
|
||||
run_ids = []
|
||||
pipeline = None
|
||||
for index in range(3):
|
||||
pipeline = DwhPreprocessPipeline(
|
||||
workspace_id="demo", workspace_root=tmp_path,
|
||||
config_fingerprint=FP, input_fingerprint=FP,
|
||||
introspect=lambda output, i=index: output.write_text(str(i)),
|
||||
build_lsh=lambda physical, output, i=index: [
|
||||
(output / name).write_text(str(i))
|
||||
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
|
||||
], retain_generations=3,
|
||||
)
|
||||
run_ids.append(pipeline.run().run_id)
|
||||
generations = tmp_path / ".tht-dwh" / "generations"
|
||||
os.utime(generations / run_ids[-1], ns=(1, 1))
|
||||
|
||||
pipeline.retain_generations = 2
|
||||
pipeline._cleanup_generations()
|
||||
|
||||
remaining = {path.name for path in generations.iterdir() if path.is_dir()}
|
||||
assert remaining == {run_ids[-1], run_ids[-2]}
|
||||
|
||||
|
||||
def test_retention_candidate_swap_to_symlink_is_never_followed(monkeypatch, tmp_path):
|
||||
import tht.jobs.dwh_pipeline as module
|
||||
|
||||
pipeline = DwhPreprocessPipeline(
|
||||
workspace_id="demo", workspace_root=tmp_path,
|
||||
config_fingerprint=FP, input_fingerprint=FP,
|
||||
introspect=lambda output: output.write_text("active"),
|
||||
build_lsh=lambda physical, output: [
|
||||
(output / name).write_text("active")
|
||||
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
|
||||
], retain_generations=1,
|
||||
)
|
||||
pipeline.run()
|
||||
generations = tmp_path / ".tht-dwh" / "generations"
|
||||
candidate_name = "e" * 32
|
||||
candidate = generations / candidate_name
|
||||
candidate.mkdir(mode=0o700)
|
||||
external = tmp_path / "external-crafted"
|
||||
external.mkdir()
|
||||
sentinel = external / "sentinel"
|
||||
sentinel.write_text("must-not-read-or-mutate")
|
||||
real_open = module.os.open
|
||||
swapped = False
|
||||
|
||||
def swapping_open(path, flags, *args, **kwargs):
|
||||
nonlocal swapped
|
||||
if path == candidate_name and kwargs.get("dir_fd") is not None and not swapped:
|
||||
swapped = True
|
||||
candidate.rmdir()
|
||||
candidate.symlink_to(external, target_is_directory=True)
|
||||
return real_open(path, flags, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(module.os, "open", swapping_open)
|
||||
pipeline._cleanup_generations()
|
||||
assert swapped
|
||||
assert sentinel.read_text() == "must-not-read-or-mutate"
|
||||
assert candidate.is_symlink()
|
||||
|
||||
|
||||
def test_reader_lease_blocks_retain_one_publisher_until_file_reads_finish(tmp_path):
|
||||
import threading
|
||||
import time
|
||||
|
||||
@@ -101,38 +101,69 @@ def _read_owned(path: Path, *, readonly: bool) -> bytes:
|
||||
os.close(fd)
|
||||
|
||||
|
||||
def validate_generation(target: Path) -> dict:
|
||||
manifest_path = target / GENERATION_MANIFEST
|
||||
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:
|
||||
directory_info = target.lstat()
|
||||
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(manifest_path, readonly=True).decode("utf-8"))
|
||||
manifest = json.loads(
|
||||
_read_owned_at(directory_fd, GENERATION_MANIFEST, readonly=True).decode("utf-8")
|
||||
)
|
||||
files = manifest["files"]
|
||||
if (
|
||||
manifest["generation"] != target.name
|
||||
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(path.name for path in target.iterdir()) != set(files) | {GENERATION_MANIFEST}:
|
||||
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
|
||||
path = target / name
|
||||
if hashlib.sha256(_read_owned(path, readonly=True)).hexdigest() != expected:
|
||||
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:
|
||||
@@ -443,25 +474,31 @@ class DwhPreprocessPipeline:
|
||||
except (OSError, KeyError, ValueError):
|
||||
continue
|
||||
generations = []
|
||||
for path in root.iterdir():
|
||||
try:
|
||||
info = path.lstat()
|
||||
except OSError:
|
||||
continue
|
||||
if (
|
||||
_RUN_ID.fullmatch(path.name)
|
||||
and stat.S_ISDIR(info.st_mode)
|
||||
and info.st_uid == os.getuid()
|
||||
):
|
||||
try:
|
||||
validate_generation(path)
|
||||
except CorruptCheckpointError:
|
||||
continue
|
||||
generations.append((info.st_mtime_ns, path.name))
|
||||
generations.sort()
|
||||
keep_recent = {name for _, name in generations[-self.retain_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
|
||||
|
||||
Reference in New Issue
Block a user