From cab312add403c7f1222828b7a22de349624bf166 Mon Sep 17 00:00:00 2001 From: dylan Date: Mon, 21 Sep 2026 18:12:35 -0500 Subject: [PATCH 1/7] Avoid reading media payloads during checkpoint startup --- README.md | 11 ++++-- src/gasbench/benchmarks/common.py | 57 ++++++++++++++++++++++++++--- tests/unit/test_benchmark_resume.py | 45 +++++++++++++++++++++++ 3 files changed, 104 insertions(+), 9 deletions(-) diff --git a/README.md b/README.md index 0a88d73..2348a08 100644 --- a/README.md +++ b/README.md @@ -85,10 +85,13 @@ before decoding, and scores and parquet output are rebuilt through the normal recorder path. A new run ID starts a new evaluation. The checkpoint freezes sample selection (including the robustness pass), source -content hashes, model/evaluator identity, seed (42 when omitted), and scoring -settings. Changed inputs fail closed. Preparing a new run reads the selected files -to fingerprint them; resumed runs validate only pending samples. Model setup and -uncommitted work may repeat after interruption. +paths and file-generation metadata (size, nanosecond modification/change times), +model/evaluator content hashes, seed (42 when omitted), and scoring settings. +Startup does not read media or augmentation payloads. Pending inputs are checked +against their frozen metadata before processing; changed inputs fail closed. +Input storage must be trusted and preserve these timestamps across mounts, as +Modal Volumes do. Metadata checks do not protect against a storage owner who can +forge file timestamps. Model setup and uncommitted work may repeat after interruption. Checkpoint storage must survive the process or container being replaced. The Python API accepts `checkpoint_persist(directory)` for filesystems requiring an diff --git a/src/gasbench/benchmarks/common.py b/src/gasbench/benchmarks/common.py index e25f355..c33db48 100644 --- a/src/gasbench/benchmarks/common.py +++ b/src/gasbench/benchmarks/common.py @@ -3,6 +3,7 @@ import hashlib import uuid import platform +import stat from importlib.metadata import version, PackageNotFoundError from pathlib import Path from typing import Callable @@ -315,9 +316,36 @@ def sample_digest(sample): return fingerprint_files(paths, paths[0].parent) +def file_metadata_digest(paths): + """Identify cached file generations without reading media payloads. + + Inputs must be on a trusted filesystem that preserves modification/change + timestamps across mounts. This detects replacement and in-place edits, + including same-size edits with restored mtime (ctime still changes). + It is not a content hash or protection against a malicious storage owner. + Sandbox inputs are read-only; checkpoint batch checksums remain unchanged. + """ + paths = sorted(paths) + if not paths: + raise CheckpointError("A selected sample is missing from the cache") + entries = [] + for path in paths: + try: + info = path.stat() + except OSError as exc: + raise CheckpointError("Cannot stat a selected input") from exc + if not stat.S_ISREG(info.st_mode): + raise CheckpointError("A selected input is not a regular file") + entries.append((str(path), info.st_size, info.st_mtime_ns, info.st_ctime_ns)) + return hashlib.sha256(json.dumps(entries, separators=(",", ":")).encode()).hexdigest() + + def verify_sample(sample): """Reject changed pending input before preprocessing; completed rows need no I/O.""" - if "content_sha256" in sample: + if "file_metadata_sha256" in sample: + if file_metadata_digest(sample_files(sample)) != sample["file_metadata_sha256"]: + raise CheckpointError("A selected sample changed since this run started") + elif "content_sha256" in sample: try: if sample_digest(sample) != sample["content_sha256"]: raise CheckpointError( @@ -330,6 +358,16 @@ def verify_sample(sample): def use_augmentation_cache(sample, path): """Use only the derived artifact selected for this run, or regenerate it.""" path = Path(path) + if "augmentation_cache_metadata_sha256" in sample: + expected = sample["augmentation_cache_metadata_sha256"] + if expected is None: + return False + try: + if file_metadata_digest([path]) != expected: + raise CheckpointError("Selected augmentation cache changed or disappeared") + except CheckpointError as exc: + raise CheckpointError("Selected augmentation cache changed or disappeared") from exc + return True if "augmentation_cache_sha256" not in sample: return path.is_file() expected = sample["augmentation_cache_sha256"] @@ -399,6 +437,7 @@ def create_tracker( ) } context = { + "input_identity": "file-metadata-v1", "settings": settings, "seed": seed, "runtime": runtime_versions(), @@ -430,7 +469,10 @@ def create_tracker( raise CheckpointError("Checkpoint has no frozen sample selection") plan.samples = previous["samples"] else: - for dataset in plan.available_datasets: + logger = get_logger(__name__) + started = time.monotonic() + logger.info("Freezing sample selection using file metadata (no media payload scan)") + for index, dataset in enumerate(plan.available_datasets, 1): selections = {} for pass_name, cap in ( ("base", plan.sampling_plan[dataset.name]), @@ -455,7 +497,7 @@ def create_tracker( ) selections[pass_name] = list(iterator) for sample in selections[pass_name]: - sample["content_sha256"] = sample_digest(sample) + sample["file_metadata_sha256"] = file_metadata_digest(sample_files(sample)) if pass_name == "aug" and config.aug_cache_dir: from .aug_cache import img_aug_cache_path, vid_aug_cache_path @@ -473,12 +515,17 @@ def create_tracker( ) # Newly generated cache files are outputs of this attempt; # only artifacts present in the frozen plan may be inputs. - sample["augmentation_cache_sha256"] = ( - fingerprint_files([path], path.parent) + sample["augmentation_cache_metadata_sha256"] = ( + file_metadata_digest([path]) if path.is_file() else None ) plan.samples[dataset.name] = selections + if index == 1 or index % 10 == 0 or index == len(plan.available_datasets): + logger.info( + "Frozen dataset %s/%s (%s) in %.1fs", + index, len(plan.available_datasets), dataset.name, time.monotonic() - started, + ) context["samples"] = plan.samples tracker = BenchmarkRunRecorder( run_id=config.run_id, diff --git a/tests/unit/test_benchmark_resume.py b/tests/unit/test_benchmark_resume.py index e61f0b7..56e3150 100644 --- a/tests/unit/test_benchmark_resume.py +++ b/tests/unit/test_benchmark_resume.py @@ -188,6 +188,51 @@ def test_changed_pending_sample_aborts_instead_of_returning_partial_score(benchm b.run(Session(b.model_dir)) +def test_checkpoint_startup_never_reads_media_or_augmentation_payloads(benchmark, tmp_path, monkeypatch): + """The old planner read every full video and augmentation before inference.""" + b = benchmark + aug_dir = tmp_path / "augmentations" + aug_dir.mkdir() + original = common.create_tracker + + def create_tracker(*args, **kwargs): + original_open = Path.open + + def guarded_open(path, *open_args, **open_kwargs): + if b.samples_dir in path.parents or aug_dir in path.parents: + pytest.fail("Checkpoint startup opened a media payload") + return original_open(path, *open_args, **open_kwargs) + + with monkeypatch.context() as patch: + patch.setattr(Path, "open", guarded_open) + tracker = original(*args, **kwargs) + assert tracker.count == 0 + raise Interrupted() + + monkeypatch.setattr(b.module, "create_tracker", create_tracker) + extra = {} if b.modality == "audio" else { + "n_aug_per_dataset": 3, "aug_cache_dir": str(aug_dir), + } + with pytest.raises(Interrupted): + b.run(Session(b.model_dir), **extra) + manifest = RecorderCheckpoint.read_manifest(b.checkpoint) + samples = manifest["benchmark"]["samples"]["tiny"]["base"] + assert samples and all("file_metadata_sha256" in sample for sample in samples) + + +def test_metadata_identity_detects_same_size_edit_with_restored_mtime(tmp_path): + import os + + path = tmp_path / "media.bin" + path.write_bytes(b"before") + before = path.stat() + sample = {"video_path": str(path), "file_metadata_sha256": common.file_metadata_digest([path])} + path.write_bytes(b"after!") + os.utime(path, ns=(before.st_atime_ns, before.st_mtime_ns)) + with pytest.raises(CheckpointError, match="changed"): + common.verify_sample(sample) + + def test_changed_model_rejected_before_inference(benchmark): b = benchmark with pytest.raises(Interrupted): From 6b1bf10b029c5512f63833c6600f3b2c79a0da19 Mon Sep 17 00:00:00 2001 From: dylan Date: Wed, 23 Sep 2026 15:58:45 -0500 Subject: [PATCH 2/7] Support mixed source formats in dataset downloads --- src/gasbench/constants.py | 4 +- src/gasbench/dataset/config.py | 17 +++++- src/gasbench/dataset/download/core.py | 20 ++++--- tests/unit/test_mixed_source_formats.py | 78 +++++++++++++++++++++++++ 4 files changed, 106 insertions(+), 13 deletions(-) create mode 100644 tests/unit/test_mixed_source_formats.py diff --git a/src/gasbench/constants.py b/src/gasbench/constants.py index dcb810f..f245c8b 100644 --- a/src/gasbench/constants.py +++ b/src/gasbench/constants.py @@ -4,7 +4,7 @@ # # Image is 3-class (no rendered): 0=real, 1=synthetic, 2=semisynthetic. # Video is 3-class: 0=real, 1=synthetic, 2=semisynthetic. -# Classical rendered video is real; rendering provenance lives in dataset metadata. +# Non-AI rendered images and video are real; rendering provenance lives in metadata. # Audio stays binary: 0=real, 1=synthetic (semisynthetic collapsed). # # Semisynthetic media retains materially captured visual content alongside @@ -12,7 +12,7 @@ # output remains synthetic, even when captured media conditions its generation; # modifying exclusively synthetic or rendered media does not make it semisynthetic. # -# Image has no rendered class: CGI/game-engine stills are excluded from image configs. +# Image has no rendered class: non-AI CGI, charts, and rendered text map to real. IMAGE_MEDIA_TYPE_TO_LABEL = { "real": 0, "synthetic": 1, diff --git a/src/gasbench/dataset/config.py b/src/gasbench/dataset/config.py index e50fcf4..10325dc 100644 --- a/src/gasbench/dataset/config.py +++ b/src/gasbench/dataset/config.py @@ -1,5 +1,5 @@ from dataclasses import dataclass, replace -from typing import Dict, List, Optional, Tuple +from typing import Dict, List, Optional, Tuple, Union from pathlib import Path import os import yaml @@ -52,7 +52,7 @@ class BenchmarkDatasetConfig: # Download parameters media_per_archive: int = 100 archives_per_dataset: int = 5 - source_format: str = "" # Auto-detected if empty + source_format: Union[str, List[str]] = "" # Auto-detected if empty source: str = "huggingface" # "huggingface", "modelscope", or "s3" hf_revision: Optional[str] = None hf_subfolders: Optional[List[str]] = None @@ -353,6 +353,13 @@ def validate_dataset_config( "media_per_archive", "archives_per_dataset", ] + source_format = config_dict.get("source_format", "") + if not isinstance(source_format, str) and not ( + isinstance(source_format, list) + and source_format + and all(isinstance(fmt, str) and fmt for fmt in source_format) + ): + errors.append(f"Dataset '{dataset_name}': source_format must be a string or a nonempty list of strings") for field in numeric_fields: if field in config_dict: value = config_dict[field] @@ -482,7 +489,11 @@ def _obfuscate_holdout_names( d.path or "", d.modality or "", d.media_type or "", - (d.source_format or ""), + ( + ",".join(sorted(set(d.source_format))) + if isinstance(d.source_format, list) + else (d.source_format or "") + ), include_paths, exclude_paths, (d.source or ""), diff --git a/src/gasbench/dataset/download/core.py b/src/gasbench/dataset/download/core.py index 5c70de8..71d0406 100644 --- a/src/gasbench/dataset/download/core.py +++ b/src/gasbench/dataset/download/core.py @@ -6,7 +6,7 @@ import shutil import tempfile from pathlib import Path -from typing import Any, Dict, Generator, List, Optional +from typing import Any, Dict, Generator, List, Optional, Union import pyarrow @@ -40,22 +40,23 @@ def _calculate_files_to_download( dataset, - source_format: str, + source_format: Union[str, List[str]], media_per_archive: int, archives_per_dataset: int, ) -> int: """Calculate # files to download based on dataset modality and source format. Returns -1 to indicate "download all files", or a positive integer for the count. """ - src_fmt = source_format.lower().lstrip(".") + formats = source_format if isinstance(source_format, list) else [source_format] + src_formats = {fmt.lower().lstrip(".") for fmt in formats} # Direct media files (jpg, png, mp4, etc.) - download equivalent to archive extraction if dataset.modality == "image": - is_direct_media = src_fmt in {ext.lstrip(".") for ext in IMAGE_FILE_EXTENSIONS} + is_direct_media = src_formats <= {ext.lstrip(".") for ext in IMAGE_FILE_EXTENSIONS} elif dataset.modality == "audio": - is_direct_media = src_fmt in {ext.lstrip(".") for ext in AUDIO_FILE_EXTENSIONS} + is_direct_media = src_formats <= {ext.lstrip(".") for ext in AUDIO_FILE_EXTENSIONS} else: - is_direct_media = src_fmt in {ext.lstrip(".") for ext in VIDEO_FILE_EXTENSIONS} + is_direct_media = src_formats <= {ext.lstrip(".") for ext in VIDEO_FILE_EXTENSIONS} if is_direct_media: if media_per_archive == -1 or archives_per_dataset == -1: @@ -187,7 +188,10 @@ def download_and_extract( fallback_formats = [".parquet", ".zip", ".tar", ".tar.gz"] requested_format = dataset.source_format - listing_formats = list(dict.fromkeys([requested_format, *fallback_formats])) + requested_formats = ( + requested_format if isinstance(requested_format, list) else [requested_format] + ) + listing_formats = list(dict.fromkeys([*requested_formats, *fallback_formats])) try: listed_filenames = _list_remote_dataset_files( dataset.path, @@ -222,7 +226,7 @@ def matches_format(filename, source_format): filenames = [ name for name in listed_filenames - if matches_format(name, requested_format) + if any(matches_format(name, fmt) for fmt in requested_formats) ] selected_format = requested_format diff --git a/tests/unit/test_mixed_source_formats.py b/tests/unit/test_mixed_source_formats.py new file mode 100644 index 0000000..24d13a6 --- /dev/null +++ b/tests/unit/test_mixed_source_formats.py @@ -0,0 +1,78 @@ +from dataclasses import replace +from types import SimpleNamespace + +import pytest + +from src.gasbench.dataset.config import ( + BenchmarkDatasetConfig, + _obfuscate_holdout_names, + validate_dataset_config, +) +from src.gasbench.dataset.download import core +from src.gasbench.dataset.download.listing import _list_remote_dataset_files +from src.gasbench.dataset.utils import s3_utils + + +def test_s3_parent_discovers_nested_shards_and_filters_sibling_datasets(monkeypatch): + keys = ["release/corpus-shard-1/a.jpg", "release/corpus-shard-2/deep/b.png", "release/other/c.jpg"] + + def paginate(**kwargs): + assert kwargs == {"Bucket": "bucket", "Prefix": "release/"} + for key in keys: + yield {"Contents": [{"Key": key}]} + + monkeypatch.setattr(s3_utils, "_get_s3_client", lambda: SimpleNamespace( + get_paginator=lambda name: SimpleNamespace(paginate=paginate), + )) + assert _list_remote_dataset_files( + "bucket/release", ["jpg", "png"], source="s3", + include_paths=["release/corpus-"], + ) == keys[:2] + + +def test_mixed_formats_download_all_matching_shards(monkeypatch, tmp_path): + dataset = BenchmarkDatasetConfig( + name="mixed", path="bucket/corpus", modality="image", media_type="real", + source="s3", source_format=["jpg", "png"], + ) + monkeypatch.setattr(core, "_list_remote_dataset_files", lambda *a, **k: [ + "corpus/shard-a/sample.jpg", "corpus/shard-b/sample.png", "corpus/backup.zip", + ]) + monkeypatch.setattr(core, "_get_download_urls", lambda path, files, *a: files) + + def download(paths, root, **kwargs): + for key in paths: + path = root / key + path.parent.mkdir(parents=True, exist_ok=True) + path.touch() + yield path + + monkeypatch.setattr(core, "_stream_downloads", download) + monkeypatch.setattr(core, "yield_media_from_source", lambda path, *a: iter([path.suffix])) + samples = list(core.download_and_extract( + dataset, media_per_archive=-1, archives_per_dataset=-1, + temp_dir=str(tmp_path), force_download=True, + )) + assert set(samples) == {".jpg", ".png"} + assert core._calculate_files_to_download(dataset, ["jpg", "png"], 3, 2) == 6 + + +def test_format_order_does_not_change_holdout_identity(): + dataset = BenchmarkDatasetConfig( + name="mixed", path="bucket/corpus", modality="image", media_type="real", + source_format=["jpg", "png"], + ) + def name(fmt): + return _obfuscate_holdout_names([replace(dataset, source_format=fmt)])[0][0].name + + assert name(["jpg", "png"]) == name(["png", "jpg", "png"]) + assert name("jpg") == name(["jpg"]) + assert name("jpg") != name(["jpg", "png"]) + + +@pytest.mark.parametrize("formats", [[], [None], [""], 42]) +def test_invalid_format_lists_are_rejected(formats): + assert any("source_format" in error for error in validate_dataset_config({ + "name": "bad", "path": "bucket/corpus", "modality": "image", + "media_type": "real", "source_format": formats, + })) From e9456a46804e9c96609e9fb7cd38283601067e09 Mon Sep 17 00:00:00 2001 From: dylan Date: Thu, 24 Sep 2026 11:12:23 -0500 Subject: [PATCH 3/7] Prefetch checkpoint batches with bounded ordered reads --- src/gasbench/benchmarks/_checkpoint.py | 95 +++++++++++++++++++------ tests/unit/test_checkpoint.py | 97 ++++++++++++++++++++++++++ 2 files changed, 170 insertions(+), 22 deletions(-) diff --git a/src/gasbench/benchmarks/_checkpoint.py b/src/gasbench/benchmarks/_checkpoint.py index c6657e1..01ebe01 100644 --- a/src/gasbench/benchmarks/_checkpoint.py +++ b/src/gasbench/benchmarks/_checkpoint.py @@ -14,8 +14,13 @@ import hashlib import json +import logging import os import tempfile +import time +from collections import deque +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager from copy import deepcopy from pathlib import Path from typing import Callable, Iterable, Mapping, Optional @@ -42,6 +47,44 @@ def _read(path: Path): raise CheckpointError(f"Cannot read checkpoint {path.name}") from exc +# Limit both active I/O and queued decoded batches on remote filesystems. +_READ_WORKERS = 16 +_READ_AHEAD = 32 + + +@contextmanager +def _prefetch_batches(paths): + """Read ahead in parallel, but expose batches in commit order. + + Keep only a bounded window of futures alive. The context owns the executor + so validation failure also cancels queued reads and joins active readers. + """ + if not paths: + yield iter(()) + return + executor = ThreadPoolExecutor(max_workers=min(_READ_WORKERS, len(paths))) + pending = deque() + remaining = iter(paths) + + def submit_next(): + path = next(remaining, None) + if path is not None: + pending.append((path, executor.submit(_read, path))) + + def ordered(): + while pending: + path, future = pending.popleft() + yield path, future.result() + submit_next() + + try: + for _ in range(min(_READ_AHEAD, len(paths))): + submit_next() + yield ordered() + finally: + executor.shutdown(wait=True, cancel_futures=True) + + def prediction_key(row: Mapping) -> tuple: """Identity of one planned prediction, using the existing recorder schema.""" fields = ("run_id", "dataset_name", "sample_id") @@ -149,35 +192,43 @@ def __init__( raise CheckpointError("Invalid or incomplete checkpoint commit head") # Files newer than the head were never acknowledged. They can be safely # replaced when the interrupted batch is replayed. - for path in batches[: head["last_batch"] + 1]: - if path.name != self._batch_name(self._next_batch): + committed = batches[: head["last_batch"] + 1] + for index, path in enumerate(committed): + if path.name != self._batch_name(index): raise CheckpointError("Checkpoint batch sequence is incomplete") - envelope = _read(path) - if not isinstance(envelope, dict): - raise CheckpointError(f"Invalid checkpoint batch {path.name}") - payload = envelope.get("payload") - try: - valid = ( - isinstance(payload, dict) - and envelope.get("sha256") == _digest(payload) - and payload.get("manifest_sha256") == self._manifest_digest - and payload.get("batch_index") == self._next_batch - ) - except (TypeError, ValueError): - valid = False - if not valid: - raise CheckpointError(f"Invalid checkpoint batch {path.name}") - records = self._validate_records(payload.get("records")) - if any(prediction_key(row) in self._records for row in records): - raise CheckpointError("Duplicate work in committed checkpoint batches") - self._records.update((prediction_key(row), row) for row in records) - self._next_batch += 1 + started = time.monotonic() + with _prefetch_batches(committed) as loaded: + for path, envelope in loaded: + if not isinstance(envelope, dict): + raise CheckpointError(f"Invalid checkpoint batch {path.name}") + payload = envelope.get("payload") + try: + valid = ( + isinstance(payload, dict) + and envelope.get("sha256") == _digest(payload) + and payload.get("manifest_sha256") == self._manifest_digest + and payload.get("batch_index") == self._next_batch + ) + except (TypeError, ValueError): + valid = False + if not valid: + raise CheckpointError(f"Invalid checkpoint batch {path.name}") + records = self._validate_records(payload.get("records")) + if any(prediction_key(row) in self._records for row in records): + raise CheckpointError("Duplicate work in committed checkpoint batches") + self._records.update((prediction_key(row), row) for row in records) + self._next_batch += 1 if ( batches and head["last_batch"] >= 0 and envelope["sha256"] != head.get("sha256") ): raise CheckpointError("Checkpoint tail differs from commit head") + if committed: + logging.getLogger(__name__).info( + "Validated %s checkpoint batches (%s rows) in %.2fs", + len(committed), len(self._records), time.monotonic() - started, + ) @staticmethod def _batch_name(index: int) -> str: diff --git a/tests/unit/test_checkpoint.py b/tests/unit/test_checkpoint.py index ff4f1ec..ad9f99f 100644 --- a/tests/unit/test_checkpoint.py +++ b/tests/unit/test_checkpoint.py @@ -236,3 +236,100 @@ def test_saved_manifest_is_checked_before_any_predictions_exist(tmp_path, manife path.write_text(json.dumps(saved)) with pytest.raises(CheckpointError, match="manifest checksum"): RecorderCheckpoint.read_manifest(tmp_path) + + +def test_parallel_reads_restore_commit_order_and_ignore_uncommitted_tail(tmp_path, manifest, monkeypatch): + from threading import Event + from gasbench.benchmarks import _checkpoint as checkpoint + + store = RecorderCheckpoint(tmp_path, manifest) + expected = [record(str(index)) for index in range(4)] + for row in expected: + store.commit_batch([row]) + tail = tmp_path / store._batch_name(len(expected)) + tail.write_text("uncommitted partial JSON") + second_finished = Event() + original = checkpoint._read + completed = [] + + def read(path): + assert path != tail + if path.name == store._batch_name(0): + assert second_finished.wait(5), "Reads were serialized behind the first batch" + result = original(path) + if path.name.startswith("batch-"): + completed.append(path.name) + if path.name == store._batch_name(1): + second_finished.set() + return result + + monkeypatch.setattr(checkpoint, "_read", read) + monkeypatch.setattr(checkpoint, "_READ_WORKERS", 2) + assert RecorderCheckpoint(tmp_path, manifest).records == expected + assert completed.index(store._batch_name(1)) < completed.index(store._batch_name(0)) + + +@pytest.mark.parametrize("stop_early", [False, True]) +def test_prefetch_is_bounded_and_releases_executor_on_consumer_failure(monkeypatch, stop_early): + from pathlib import Path + from gasbench.benchmarks import _checkpoint as checkpoint + + window = 3 + submitted = consumed = 0 + shutdown = [] + + class Future: + def __init__(self, fn, path): + self.fn, self.path = fn, path + + def result(self): + nonlocal consumed + consumed += 1 + return self.fn(self.path) + + class Executor: + def __init__(self, **kwargs): + pass + + def submit(self, fn, path): + nonlocal submitted + submitted += 1 + assert submitted - consumed <= window + return Future(fn, path) + + def shutdown(self, **kwargs): + shutdown.append(kwargs) + + monkeypatch.setattr(checkpoint, "ThreadPoolExecutor", Executor) + monkeypatch.setattr(checkpoint, "_READ_AHEAD", window) + monkeypatch.setattr(checkpoint, "_read", lambda path: path.name) + paths = [Path(str(i)) for i in range(11)] + try: + with checkpoint._prefetch_batches(paths) as loaded: + for index, (path, value) in enumerate(loaded): + assert path == paths[index] and value == path.name + if stop_early: + raise CheckpointError("validation rejected batch") + except CheckpointError: + assert stop_early + assert shutdown == [{"wait": True, "cancel_futures": True}] + if not stop_early: + assert submitted == consumed == len(paths) + + +def test_prefetched_read_failure_fails_closed(tmp_path, manifest, monkeypatch): + from gasbench.benchmarks import _checkpoint as checkpoint + + store = RecorderCheckpoint(tmp_path, manifest) + for index in range(5): + store.commit_batch([record(str(index))]) + original = checkpoint._read + + def read(path): + if path.name == store._batch_name(2): + raise CheckpointError("read failed") + return original(path) + + monkeypatch.setattr(checkpoint, "_read", read) + with pytest.raises(CheckpointError, match="read failed"): + RecorderCheckpoint(tmp_path, manifest) From 320cfaa81d7392911972d708db0024dc574193aa Mon Sep 17 00:00:00 2001 From: dylan Date: Thu, 24 Sep 2026 11:45:09 -0500 Subject: [PATCH 4/7] Bump version to 0.10.1 --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 386f3cd..d9e9bf2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "gasbench" -version = "0.10.0" +version = "0.10.1" description = "GASBench - ML model benchmark evaluation package" readme = "README.md" requires-python = ">=3.10" From 52199567d2e29cd977307db7cebf2568f1b2ea6d Mon Sep 17 00:00:00 2001 From: Dylan Uys Date: Thu, 24 Sep 2026 13:21:30 -0500 Subject: [PATCH 5/7] Bound S3 listings with explicit dataset prefixes (#167) --- src/gasbench/dataset/config.py | 10 +++ src/gasbench/dataset/download/core.py | 2 + src/gasbench/dataset/download/listing.py | 17 ++-- src/gasbench/dataset/utils/s3_utils.py | 106 +++++++++++------------ tests/unit/test_s3_listing_limits.py | 70 +++++++++++++++ 5 files changed, 140 insertions(+), 65 deletions(-) create mode 100644 tests/unit/test_s3_listing_limits.py diff --git a/src/gasbench/dataset/config.py b/src/gasbench/dataset/config.py index 10325dc..5c687bb 100644 --- a/src/gasbench/dataset/config.py +++ b/src/gasbench/dataset/config.py @@ -57,6 +57,7 @@ class BenchmarkDatasetConfig: hf_revision: Optional[str] = None hf_subfolders: Optional[List[str]] = None + s3_prefixes: Optional[List[str]] = None # Bucket-relative literal key prefixes include_paths: Optional[List[str]] = None exclude_paths: Optional[List[str]] = None # Filter media entries inside ZIP/TAR archives. Unlike include_paths and @@ -349,6 +350,12 @@ def validate_dataset_config( f"Must be one of {valid_sources}" ) + if config_dict.get("s3_prefixes") is not None: + prefixes = config_dict["s3_prefixes"] + if (config_dict.get("source") != "s3" or not isinstance(prefixes, list) + or not prefixes or any(not isinstance(p, str) or not p for p in prefixes)): + errors.append(f"Dataset '{dataset_name}': s3_prefixes requires S3 and nonempty key prefixes") + numeric_fields = [ "media_per_archive", "archives_per_dataset", @@ -387,6 +394,7 @@ def _dataset_dict_to_config(d: dict, **overrides) -> BenchmarkDatasetConfig: "source": d.get("source", "huggingface"), "hf_revision": d.get("hf_revision"), "hf_subfolders": d.get("hf_subfolders"), + "s3_prefixes": d.get("s3_prefixes"), "media_per_archive": d.get("media_per_archive", 100), "archives_per_dataset": d.get("archives_per_dataset", 5), "include_paths": d.get("include_paths"), @@ -501,6 +509,8 @@ def _obfuscate_holdout_names( hf_subfolders, ] ) + if d.s3_prefixes: + fingerprint += "|s3_prefixes=" + ",".join(sorted(set(d.s3_prefixes))) short_hash = hashlib.sha1(fingerprint.encode("utf-8")).hexdigest()[:8] new_name = f"{d.media_type}-{d.modality}-holdout-{short_hash}" obfuscated.append(replace(d, name=new_name, original_name=orig_name)) diff --git a/src/gasbench/dataset/download/core.py b/src/gasbench/dataset/download/core.py index 71d0406..1f68991 100644 --- a/src/gasbench/dataset/download/core.py +++ b/src/gasbench/dataset/download/core.py @@ -160,6 +160,7 @@ def download_and_extract( max_files=None, # need full list to shuffle and pick from hf_revision=hf_revision, hf_subfolders=hf_subfolders, + s3_prefixes=getattr(dataset, "s3_prefixes", None), ) except DatasetAccessError as e: logger.warning(f"Skipping {dataset.name}: {e}") @@ -206,6 +207,7 @@ def download_and_extract( max_files=max_files_to_list, hf_revision=hf_revision, hf_subfolders=hf_subfolders, + s3_prefixes=getattr(dataset, "s3_prefixes", None), ) except DatasetAccessError as e: logger.warning(f"Skipping dataset {dataset.name}: {e}") diff --git a/src/gasbench/dataset/download/listing.py b/src/gasbench/dataset/download/listing.py index 46af2ae..bfde8fa 100644 --- a/src/gasbench/dataset/download/listing.py +++ b/src/gasbench/dataset/download/listing.py @@ -66,6 +66,7 @@ def _list_remote_dataset_files( max_files: Optional[int] = None, hf_revision: Optional[str] = None, hf_subfolders: Optional[List[str]] = None, + s3_prefixes: Optional[List[str]] = None, ) -> List[str]: """List available files in a dataset, filtered by source_format and path patterns. @@ -123,17 +124,11 @@ def _list_remote_dataset_files( if not needs_week_filter: files = files[:max_files] elif source == "s3": - files = list_s3_files(path=dataset_path, extension=source_format) - if include_paths: - files = [ - f for f in files if any(path_seg in f for path_seg in include_paths) - ] - if exclude_paths: - files = [ - f for f in files if not any(path_seg in f for path_seg in exclude_paths) - ] - if not needs_week_filter: - files = files[:max_files] + files = list_s3_files( + path=dataset_path, extension=source_format, prefixes=s3_prefixes, + include_paths=include_paths, exclude_paths=exclude_paths, + max_files=listing_max, + ) else: # hf - supports early termination natively files = list_hf_files( repo_id=dataset_path, diff --git a/src/gasbench/dataset/utils/s3_utils.py b/src/gasbench/dataset/utils/s3_utils.py index 177e86a..d05dc8f 100644 --- a/src/gasbench/dataset/utils/s3_utils.py +++ b/src/gasbench/dataset/utils/s3_utils.py @@ -125,68 +125,66 @@ def _parse_s3_path(path: str) -> Tuple[str, str]: return bucket, prefix -def list_s3_files(path: str, extension=None) -> List[str]: - """List files from an S3 bucket. - - Args: - path: Path in format 'bucket-name/prefix/path' - extension: Filter files by extension(s) (e.g., '.parquet' or ['.tar', '.tar.gz']) - Special value 'frames' to detect frame directories - - Returns: - List of file keys (paths within the bucket) or directory paths if extension='frames' +def list_s3_files(path: str, extension=None, *, prefixes=None, + include_paths=None, exclude_paths=None, max_files=None) -> List[str]: + """List matching keys, using optional bucket-relative literal prefixes. + + Prefixes may end mid-directory name (for sharded corpora). Substring filters + retain their existing semantics. Apply the cap only after all filters. """ + bucket, base = _parse_s3_path(path) + base = base.rstrip("/") + "/" if base else "" + if prefixes is not None and ( + not isinstance(prefixes, (list, tuple)) or not prefixes + or any(not isinstance(p, str) or not p or not p.startswith(base) for p in prefixes) + ): + raise ValueError("S3 prefixes must be nonempty keys within the dataset path") + # Remove overlaps so each object is enumerated once, in global key order. + selected = [] + for prefix in sorted(set(prefixes or [base])): + if not any(prefix.startswith(parent) for parent in selected): + selected.append(prefix) + if max_files is not None and max_files < 0: + raise ValueError("max_files must be nonnegative") + if max_files == 0: + return [] + frames = extension == "frames" + extensions = IMAGE_FILE_EXTENSIONS if frames else extension + if isinstance(extensions, str): + extensions = [extensions] + extensions = tuple(e.lower() for e in extensions) if extensions else () + files, seen = [], set() try: - bucket, prefix = _parse_s3_path(path) - client = _get_s3_client() - - files = [] - paginator = client.get_paginator("list_objects_v2") - - if prefix and not prefix.endswith("/"): - prefix = prefix + "/" - - for page in paginator.paginate(Bucket=bucket, Prefix=prefix): - if "Contents" not in page: - continue - - for obj in page["Contents"]: - key = obj["Key"] - if key.endswith("/"): - continue - files.append(key) - - if extension == "frames": - frame_dirs = set() - for f in files: - if any(f.lower().endswith(ext) for ext in IMAGE_FILE_EXTENSIONS): - parent = os.path.dirname(f) - if parent: - frame_dirs.add(parent) - - logger.info(f"Found {len(frame_dirs)} frame directories in S3 bucket {bucket}/{prefix}") - return sorted(list(frame_dirs)) - - if extension and files: - if isinstance(extension, (list, tuple, set)): - exts = tuple(e.lower() for e in extension) - files = [f for f in files if f.lower().endswith(exts)] - else: - ext_lower = extension.lower() - files = [f for f in files if f.lower().endswith(ext_lower)] - - logger.info(f"Found {len(files)} files in S3 bucket {bucket}/{prefix}") - return files - + paginator = _get_s3_client().get_paginator("list_objects_v2") + for prefix in selected: + for page in paginator.paginate(Bucket=bucket, Prefix=prefix): + for obj in page.get("Contents", []): + key = obj["Key"] + if key.endswith("/") or (extensions and not key.lower().endswith(extensions)): + continue + candidate = os.path.dirname(key) if frames else key + if not candidate or candidate in seen: + continue + if include_paths and not any(p in candidate for p in include_paths): + continue + if exclude_paths and any(p in candidate for p in exclude_paths): + continue + seen.add(candidate) + files.append(candidate) + # Frame parents may interleave in key order; keep their historical + # sorted selection instead of truncating before the sort. + if not frames and max_files is not None and len(files) >= max_files: + return files + if frames: + files.sort() + return files[:max_files] if max_files is not None else files except NoCredentialsError: logger.error("S3 credentials not found or invalid") return [] except ClientError as e: logger.error(f"Failed to list S3 files from {path}: {e}") return [] - except Exception as e: - logger.error(f"Error listing S3 files from {path}: {e}") - return [] + def _get_s3_urls(path: str, filenames: List[str]) -> List[str]: diff --git a/tests/unit/test_s3_listing_limits.py b/tests/unit/test_s3_listing_limits.py new file mode 100644 index 0000000..61e4251 --- /dev/null +++ b/tests/unit/test_s3_listing_limits.py @@ -0,0 +1,70 @@ +from types import SimpleNamespace + +import pytest + +from src.gasbench.dataset.download.listing import _list_remote_dataset_files +from src.gasbench.dataset.utils import s3_utils + + +def mock_pages(monkeypatch, pages): + calls = [] + def paginate(**kwargs): + prefix = kwargs["Prefix"] + calls.append(prefix) + yield from pages(prefix) + monkeypatch.setattr(s3_utils, "_get_s3_client", lambda: SimpleNamespace( + get_paginator=lambda _: SimpleNamespace(paginate=paginate))) + return calls + + +def test_prefixes_keep_partial_shard_names_deduplicate_and_stop(monkeypatch): + def pages(prefix): + yield {"Contents": [{"Key": prefix + "001/no.txt"}, + {"Key": prefix + "001/excluded.jpg"}, + {"Key": prefix + "002/keep.PNG"}]} + raise AssertionError("Fetched a page beyond the matching limit") + calls = mock_pages(monkeypatch, pages) + assert _list_remote_dataset_files( + "bucket/root", ["jpg", "png"], source="s3", + s3_prefixes=["root/corpus-", "root/corpus-001/", "root/later-"], + exclude_paths=["excluded"], max_files=1, + ) == ["root/corpus-002/keep.PNG"] + assert calls == ["root/corpus-"] + + +def test_substring_filters_still_match_nested_keys(monkeypatch): + mock_pages(monkeypatch, lambda _: iter([ + {"Contents": [{"Key": "root/other/no.jpg"}]}, + {"Contents": [{"Key": "root/deep/corpus/a.jpg"}]}, + ])) + assert _list_remote_dataset_files( + "bucket/root", "jpg", source="s3", include_paths=["corpus"], max_files=1 + ) == ["root/deep/corpus/a.jpg"] + + +def test_week_filter_runs_before_limit(monkeypatch): + mock_pages(monkeypatch, lambda _: iter([ + {"Contents": [{"Key": "gasstation/old/a.jpg"}]}, + {"Contents": [{"Key": "gasstation/desired/b.jpg"}]}, + ])) + assert _list_remote_dataset_files( + "bucket/gasstation", "jpg", source="s3", target_week="desired", max_files=1, + ) == ["gasstation/desired/b.jpg"] + + +@pytest.mark.parametrize("prefixes", [[], ["outside/"], ["root-other/"], [""]]) +def test_prefix_cannot_escape_dataset_path(prefixes): + with pytest.raises(ValueError, match="prefixes"): + s3_utils.list_s3_files("bucket/root", prefixes=prefixes) + + +def test_prefix_config_is_loaded_and_changes_cache_identity(): + from src.gasbench.dataset.config import _dataset_dict_to_config, _obfuscate_holdout_names + row = {"name": "corpus", "path": "bucket/root", "source": "s3", + "modality": "image", "media_type": "real", "source_format": "jpg"} + original = _dataset_dict_to_config(row) + narrowed = _dataset_dict_to_config({**row, "s3_prefixes": ["root/corpus-"]}) + assert narrowed.s3_prefixes == ["root/corpus-"] + before, _ = _obfuscate_holdout_names([original]) + after, _ = _obfuscate_holdout_names([narrowed]) + assert before[0].name != after[0].name From ef4e60560c8c1eb16aa385371223f3cc16f6f784 Mon Sep 17 00:00:00 2001 From: dylan Date: Thu, 24 Sep 2026 11:12:23 -0500 Subject: [PATCH 6/7] Prefetch checkpoint batches with bounded ordered reads --- src/gasbench/benchmarks/_checkpoint.py | 95 +++++++++++++++++++------ tests/unit/test_checkpoint.py | 97 ++++++++++++++++++++++++++ 2 files changed, 170 insertions(+), 22 deletions(-) diff --git a/src/gasbench/benchmarks/_checkpoint.py b/src/gasbench/benchmarks/_checkpoint.py index c6657e1..01ebe01 100644 --- a/src/gasbench/benchmarks/_checkpoint.py +++ b/src/gasbench/benchmarks/_checkpoint.py @@ -14,8 +14,13 @@ import hashlib import json +import logging import os import tempfile +import time +from collections import deque +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager from copy import deepcopy from pathlib import Path from typing import Callable, Iterable, Mapping, Optional @@ -42,6 +47,44 @@ def _read(path: Path): raise CheckpointError(f"Cannot read checkpoint {path.name}") from exc +# Limit both active I/O and queued decoded batches on remote filesystems. +_READ_WORKERS = 16 +_READ_AHEAD = 32 + + +@contextmanager +def _prefetch_batches(paths): + """Read ahead in parallel, but expose batches in commit order. + + Keep only a bounded window of futures alive. The context owns the executor + so validation failure also cancels queued reads and joins active readers. + """ + if not paths: + yield iter(()) + return + executor = ThreadPoolExecutor(max_workers=min(_READ_WORKERS, len(paths))) + pending = deque() + remaining = iter(paths) + + def submit_next(): + path = next(remaining, None) + if path is not None: + pending.append((path, executor.submit(_read, path))) + + def ordered(): + while pending: + path, future = pending.popleft() + yield path, future.result() + submit_next() + + try: + for _ in range(min(_READ_AHEAD, len(paths))): + submit_next() + yield ordered() + finally: + executor.shutdown(wait=True, cancel_futures=True) + + def prediction_key(row: Mapping) -> tuple: """Identity of one planned prediction, using the existing recorder schema.""" fields = ("run_id", "dataset_name", "sample_id") @@ -149,35 +192,43 @@ def __init__( raise CheckpointError("Invalid or incomplete checkpoint commit head") # Files newer than the head were never acknowledged. They can be safely # replaced when the interrupted batch is replayed. - for path in batches[: head["last_batch"] + 1]: - if path.name != self._batch_name(self._next_batch): + committed = batches[: head["last_batch"] + 1] + for index, path in enumerate(committed): + if path.name != self._batch_name(index): raise CheckpointError("Checkpoint batch sequence is incomplete") - envelope = _read(path) - if not isinstance(envelope, dict): - raise CheckpointError(f"Invalid checkpoint batch {path.name}") - payload = envelope.get("payload") - try: - valid = ( - isinstance(payload, dict) - and envelope.get("sha256") == _digest(payload) - and payload.get("manifest_sha256") == self._manifest_digest - and payload.get("batch_index") == self._next_batch - ) - except (TypeError, ValueError): - valid = False - if not valid: - raise CheckpointError(f"Invalid checkpoint batch {path.name}") - records = self._validate_records(payload.get("records")) - if any(prediction_key(row) in self._records for row in records): - raise CheckpointError("Duplicate work in committed checkpoint batches") - self._records.update((prediction_key(row), row) for row in records) - self._next_batch += 1 + started = time.monotonic() + with _prefetch_batches(committed) as loaded: + for path, envelope in loaded: + if not isinstance(envelope, dict): + raise CheckpointError(f"Invalid checkpoint batch {path.name}") + payload = envelope.get("payload") + try: + valid = ( + isinstance(payload, dict) + and envelope.get("sha256") == _digest(payload) + and payload.get("manifest_sha256") == self._manifest_digest + and payload.get("batch_index") == self._next_batch + ) + except (TypeError, ValueError): + valid = False + if not valid: + raise CheckpointError(f"Invalid checkpoint batch {path.name}") + records = self._validate_records(payload.get("records")) + if any(prediction_key(row) in self._records for row in records): + raise CheckpointError("Duplicate work in committed checkpoint batches") + self._records.update((prediction_key(row), row) for row in records) + self._next_batch += 1 if ( batches and head["last_batch"] >= 0 and envelope["sha256"] != head.get("sha256") ): raise CheckpointError("Checkpoint tail differs from commit head") + if committed: + logging.getLogger(__name__).info( + "Validated %s checkpoint batches (%s rows) in %.2fs", + len(committed), len(self._records), time.monotonic() - started, + ) @staticmethod def _batch_name(index: int) -> str: diff --git a/tests/unit/test_checkpoint.py b/tests/unit/test_checkpoint.py index ff4f1ec..ad9f99f 100644 --- a/tests/unit/test_checkpoint.py +++ b/tests/unit/test_checkpoint.py @@ -236,3 +236,100 @@ def test_saved_manifest_is_checked_before_any_predictions_exist(tmp_path, manife path.write_text(json.dumps(saved)) with pytest.raises(CheckpointError, match="manifest checksum"): RecorderCheckpoint.read_manifest(tmp_path) + + +def test_parallel_reads_restore_commit_order_and_ignore_uncommitted_tail(tmp_path, manifest, monkeypatch): + from threading import Event + from gasbench.benchmarks import _checkpoint as checkpoint + + store = RecorderCheckpoint(tmp_path, manifest) + expected = [record(str(index)) for index in range(4)] + for row in expected: + store.commit_batch([row]) + tail = tmp_path / store._batch_name(len(expected)) + tail.write_text("uncommitted partial JSON") + second_finished = Event() + original = checkpoint._read + completed = [] + + def read(path): + assert path != tail + if path.name == store._batch_name(0): + assert second_finished.wait(5), "Reads were serialized behind the first batch" + result = original(path) + if path.name.startswith("batch-"): + completed.append(path.name) + if path.name == store._batch_name(1): + second_finished.set() + return result + + monkeypatch.setattr(checkpoint, "_read", read) + monkeypatch.setattr(checkpoint, "_READ_WORKERS", 2) + assert RecorderCheckpoint(tmp_path, manifest).records == expected + assert completed.index(store._batch_name(1)) < completed.index(store._batch_name(0)) + + +@pytest.mark.parametrize("stop_early", [False, True]) +def test_prefetch_is_bounded_and_releases_executor_on_consumer_failure(monkeypatch, stop_early): + from pathlib import Path + from gasbench.benchmarks import _checkpoint as checkpoint + + window = 3 + submitted = consumed = 0 + shutdown = [] + + class Future: + def __init__(self, fn, path): + self.fn, self.path = fn, path + + def result(self): + nonlocal consumed + consumed += 1 + return self.fn(self.path) + + class Executor: + def __init__(self, **kwargs): + pass + + def submit(self, fn, path): + nonlocal submitted + submitted += 1 + assert submitted - consumed <= window + return Future(fn, path) + + def shutdown(self, **kwargs): + shutdown.append(kwargs) + + monkeypatch.setattr(checkpoint, "ThreadPoolExecutor", Executor) + monkeypatch.setattr(checkpoint, "_READ_AHEAD", window) + monkeypatch.setattr(checkpoint, "_read", lambda path: path.name) + paths = [Path(str(i)) for i in range(11)] + try: + with checkpoint._prefetch_batches(paths) as loaded: + for index, (path, value) in enumerate(loaded): + assert path == paths[index] and value == path.name + if stop_early: + raise CheckpointError("validation rejected batch") + except CheckpointError: + assert stop_early + assert shutdown == [{"wait": True, "cancel_futures": True}] + if not stop_early: + assert submitted == consumed == len(paths) + + +def test_prefetched_read_failure_fails_closed(tmp_path, manifest, monkeypatch): + from gasbench.benchmarks import _checkpoint as checkpoint + + store = RecorderCheckpoint(tmp_path, manifest) + for index in range(5): + store.commit_batch([record(str(index))]) + original = checkpoint._read + + def read(path): + if path.name == store._batch_name(2): + raise CheckpointError("read failed") + return original(path) + + monkeypatch.setattr(checkpoint, "_read", read) + with pytest.raises(CheckpointError, match="read failed"): + RecorderCheckpoint(tmp_path, manifest) From c7bab881b09fd44463e9e91fe3a23bb2d76fd060 Mon Sep 17 00:00:00 2001 From: dylan Date: Fri, 25 Sep 2026 16:32:15 -0500 Subject: [PATCH 7/7] Allow audited reader-only upgrade for existing checkpoints --- pyproject.toml | 1 + src/gasbench/benchmarks/common.py | 20 ++++++++++++++ src/gasbench/checkpoint_compatibility.json | 5 ++++ tests/unit/test_benchmark_resume.py | 31 ++++++++++++++++++++++ 4 files changed, 57 insertions(+) create mode 100644 src/gasbench/checkpoint_compatibility.json diff --git a/pyproject.toml b/pyproject.toml index 386f3cd..0d54f89 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -63,6 +63,7 @@ where = ["src"] "" = "src" [tool.setuptools.package-data] +"gasbench" = ["checkpoint_compatibility.json"] "gasbench.dataset" = ["configs/*.yaml"] [tool.black] diff --git a/src/gasbench/benchmarks/common.py b/src/gasbench/benchmarks/common.py index c33db48..b6279cb 100644 --- a/src/gasbench/benchmarks/common.py +++ b/src/gasbench/benchmarks/common.py @@ -275,6 +275,20 @@ def fingerprint_files(paths, root): return digest.hexdigest() +def compatible_evaluator_identity(evaluator_dir, current, previous): + """Allow only explicitly audited, exact-build checkpoint reader upgrades.""" + if current == previous: + return True + registry = Path(evaluator_dir) / "checkpoint_compatibility.json" + if not registry.is_file(): + return False + try: + pairs = json.loads(registry.read_text()) + except (OSError, ValueError): + return False + return isinstance(pairs, dict) and isinstance(pairs.get(current), list) and previous in pairs[current] + + def runtime_versions(): versions = { "python": platform.python_version(), @@ -461,6 +475,12 @@ def create_tracker( context = json.loads(json.dumps(context)) if saved is not None: previous = saved.get("benchmark", {}) + if compatible_evaluator_identity( + evaluator_dir, context["evaluator_sha256"], previous.get("evaluator_sha256") + ): + # Preserve the original manifest identity for both restored and new batches. + # Every other context field is still compared below without exceptions. + context["evaluator_sha256"] = previous["evaluator_sha256"] if {k: v for k, v in previous.items() if k != "samples"} != context: raise CheckpointError( "Checkpoint configuration or model differs from this run" diff --git a/src/gasbench/checkpoint_compatibility.json b/src/gasbench/checkpoint_compatibility.json new file mode 100644 index 0000000..91e981d --- /dev/null +++ b/src/gasbench/checkpoint_compatibility.json @@ -0,0 +1,5 @@ +{ + "cefcbe4740db068e9e34d3f4184cbde9d65b441482bc2ecdbcfeb98584634288": [ + "ec5643b53236c144c9ed296f653a1dd39b852ee37ce4feaefa2179650b1ad2cb" + ] +} diff --git a/tests/unit/test_benchmark_resume.py b/tests/unit/test_benchmark_resume.py index 56e3150..c27a62f 100644 --- a/tests/unit/test_benchmark_resume.py +++ b/tests/unit/test_benchmark_resume.py @@ -366,3 +366,34 @@ def test_augmentation_cache_cannot_change_across_attempts(benchmark, tmp_path): path.write_bytes(b"changed") with pytest.raises(CheckpointError, match="augmentation cache"): b.run(Session(b.model_dir), n_aug_per_dataset=3, aug_cache_dir=str(aug_dir)) + + +@pytest.mark.parametrize("change", ["reader", "unapproved-evaluator", "model"]) +def test_audited_reader_upgrade_preserves_resume_guards(benchmark, monkeypatch, tmp_path, change): + b = benchmark + original_fingerprint = common.fingerprint_files + evaluator_root = Path(common.__file__).resolve().parents[1] + identity = ["old"] + + def fingerprint(paths, root): + return identity[0] if root == evaluator_root else original_fingerprint(paths, root) + + monkeypatch.setattr(common, "fingerprint_files", fingerprint) + (tmp_path / "checkpoint_compatibility.json").write_text(json.dumps({"new": ["old"]})) + compatible = common.compatible_evaluator_identity + monkeypatch.setattr(common, "compatible_evaluator_identity", lambda root, current, previous: compatible(tmp_path, current, previous)) + with pytest.raises(Interrupted): + b.run(Session(b.model_dir, stop_after=2)) + identity[0] = "unapproved" if change == "unapproved-evaluator" else "new" + if change == "model": + (b.model_dir / "model.py").write_text("changed inference implementation") + resumed = Session(b.model_dir) + if change != "reader": + with pytest.raises(CheckpointError, match="differs"): + b.run(resumed) + assert not resumed.calls + else: + b.run(resumed) + assert len(resumed.calls) == 3 + assert RecorderCheckpoint.read_manifest(b.checkpoint)["benchmark"]["evaluator_sha256"] == "old" + assert not compatible(tmp_path, "new", "unrelated-old")