fix(preprocess): anchor DWH retention validation

This commit is contained in:
2026-07-12 06:15:27 +02:00
parent 27697e5db2
commit 4a29086fe4
2 changed files with 127 additions and 25 deletions
+65
View File
@@ -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
+62 -25
View File
@@ -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