From 4a29086fe49e96d09dc852e8b69d4544a6379442 Mon Sep 17 00:00:00 2001 From: mptyl Date: Sun, 12 Jul 2026 06:15:27 +0200 Subject: [PATCH] fix(preprocess): anchor DWH retention validation --- harness/tests/test_dwh_preprocess_job.py | 65 ++++++++++++++++++ harness/tht/jobs/dwh_pipeline.py | 87 +++++++++++++++++------- 2 files changed, 127 insertions(+), 25 deletions(-) diff --git a/harness/tests/test_dwh_preprocess_job.py b/harness/tests/test_dwh_preprocess_job.py index 79d6812f..7c85c7aa 100644 --- a/harness/tests/test_dwh_preprocess_job.py +++ b/harness/tests/test_dwh_preprocess_job.py @@ -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 diff --git a/harness/tht/jobs/dwh_pipeline.py b/harness/tht/jobs/dwh_pipeline.py index e0815f37..af3a1488 100644 --- a/harness/tht/jobs/dwh_pipeline.py +++ b/harness/tht/jobs/dwh_pipeline.py @@ -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