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()
|
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):
|
def test_reader_lease_blocks_retain_one_publisher_until_file_reads_finish(tmp_path):
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
|
|||||||
@@ -101,38 +101,69 @@ def _read_owned(path: Path, *, readonly: bool) -> bytes:
|
|||||||
os.close(fd)
|
os.close(fd)
|
||||||
|
|
||||||
|
|
||||||
def validate_generation(target: Path) -> dict:
|
def _read_owned_at(directory_fd: int, name: str, *, readonly: bool) -> bytes:
|
||||||
manifest_path = target / GENERATION_MANIFEST
|
fd = os.open(name, os.O_RDONLY | os.O_NOFOLLOW, dir_fd=directory_fd)
|
||||||
try:
|
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 (
|
if (
|
||||||
not stat.S_ISDIR(directory_info.st_mode)
|
not stat.S_ISDIR(directory_info.st_mode)
|
||||||
or directory_info.st_uid != os.getuid()
|
or directory_info.st_uid != os.getuid()
|
||||||
or stat.S_IMODE(directory_info.st_mode) != 0o700
|
or stat.S_IMODE(directory_info.st_mode) != 0o700
|
||||||
):
|
):
|
||||||
raise ValueError
|
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"]
|
files = manifest["files"]
|
||||||
if (
|
if (
|
||||||
manifest["generation"] != target.name
|
manifest["generation"] != generation
|
||||||
or not isinstance(files, dict)
|
or not isinstance(files, dict)
|
||||||
or not re.fullmatch(r"sha256:[0-9a-f]{64}", manifest["job_spec_fingerprint"])
|
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"])
|
or not re.fullmatch(r"[0-9a-f]{64}", manifest["artifact_manifest_sha256"])
|
||||||
):
|
):
|
||||||
raise ValueError
|
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
|
raise ValueError
|
||||||
for name, expected in files.items():
|
for name, expected in files.items():
|
||||||
if name != "physical.yaml" and not _SAFE_FILE.fullmatch(name):
|
if name != "physical.yaml" and not _SAFE_FILE.fullmatch(name):
|
||||||
raise ValueError
|
raise ValueError
|
||||||
path = target / name
|
payload = _read_owned_at(directory_fd, name, readonly=True)
|
||||||
if hashlib.sha256(_read_owned(path, readonly=True)).hexdigest() != expected:
|
if hashlib.sha256(payload).hexdigest() != expected:
|
||||||
raise ValueError
|
raise ValueError
|
||||||
return manifest
|
return manifest
|
||||||
except (OSError, KeyError, TypeError, ValueError, UnicodeDecodeError, json.JSONDecodeError) as error:
|
except (OSError, KeyError, TypeError, ValueError, UnicodeDecodeError, json.JSONDecodeError) as error:
|
||||||
raise CorruptCheckpointError("published DWH generation is invalid") from 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:
|
def resolve_dwh_snapshot(cfg) -> DwhArtifactSnapshot:
|
||||||
target = active_generation_dir(cfg.paths.artifacts.parent)
|
target = active_generation_dir(cfg.paths.artifacts.parent)
|
||||||
if target is None:
|
if target is None:
|
||||||
@@ -443,25 +474,31 @@ class DwhPreprocessPipeline:
|
|||||||
except (OSError, KeyError, ValueError):
|
except (OSError, KeyError, ValueError):
|
||||||
continue
|
continue
|
||||||
generations = []
|
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)
|
root_fd = os.open(root, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
|
||||||
try:
|
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:
|
for _, name in generations:
|
||||||
if name in protected | keep_recent:
|
if name in protected | keep_recent:
|
||||||
continue
|
continue
|
||||||
|
|||||||
Reference in New Issue
Block a user