From 5dc237e0c3bcdee45f68122c40a42fe4d8d74ec5 Mon Sep 17 00:00:00 2001 From: Hrishikesh Yadav Date: Sun, 4 Oct 2026 13:05:37 +0530 Subject: [PATCH 1/3] fix(dataset): distinguish offline fallback from corrupted CIFAR archives and batches - Narrow CIFAR-10 fallback catch strictly to expected network/connectivity errors during remote archive download. - Ensure corrupted tar archives, unsafe path traversal, missing batch files, unpickling errors, and malformed dictionary payloads raise explicitly rather than silently substituting synthetic random tensors. - Clean up partial archive downloads if network fails during retrieval. - Add regression tests in tests/test_dataset.py covering offline fallback, corrupt archives, path traversal guards, corrupt pickle batches, malformed payloads, and dedicated RNG generator isolation. Fixes #99. --- src/dataset.py | 109 +++++++++++++++++++----------- tests/test_dataset.py | 149 ++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 220 insertions(+), 38 deletions(-) create mode 100644 tests/test_dataset.py diff --git a/src/dataset.py b/src/dataset.py index 7e41580e..975a1162 100644 --- a/src/dataset.py +++ b/src/dataset.py @@ -20,13 +20,23 @@ """ import hashlib +import http.client import sys +import urllib.error import urllib.request import zipfile from pathlib import Path import torch +_NETWORK_ERRORS = ( + urllib.error.URLError, + TimeoutError, + ConnectionError, + http.client.HTTPException, + OSError, +) + DATA_DIR = Path(__file__).resolve().parents[1] / "data" # Single reliable single-file sources. wikitext is handled via HF `datasets`. @@ -76,7 +86,7 @@ def load_corpus(name): print(" ~> downloading tinyshakespeare ...", file=sys.stderr) _download(_SOURCES["shakespeare"]["url"], local, expected_hash=_SOURCES["shakespeare"]["sha256"]) - except Exception as exc: # offline / blocked -> bundled sample + except _NETWORK_ERRORS as exc: # offline / blocked -> bundled sample sample = DATA_DIR / "shakespeare_sample.txt" if sample.exists(): print( @@ -170,52 +180,74 @@ class CIFARDataset: URL = "https://www.cs.toronto.edu/~kriz/cifar-10-python.tar.gz" - def __init__(self, image_size=32): + def __init__(self, image_size=32, data_dir=None): self.name = "cifar" self.vocab_size = 10 # num classes; named vocab_size for a uniform API self.num_classes = 10 self.image_size = image_size + self.data_dir = Path(data_dir) if data_dir is not None else DATA_DIR self._images, self._labels = self._load() self.encoded = self._labels # for manifest hashing + @staticmethod + def _synthetic_dataset(num_samples=2048): + """Fixed synthetic image dataset for offline / sandbox environments. + + Uses a dedicated generator so global torch RNG state is preserved. + """ + g = torch.Generator().manual_seed(0) # dedicated -> global RNG untouched + x = torch.randn(num_samples, 3, 32, 32, generator=g) + y = torch.randint(0, 10, (num_samples,), generator=g) + return x, y + def _load(self): - local_dir = DATA_DIR / "cifar-10-batches-py" - try: - if not local_dir.exists(): - import tarfile + local_dir = self.data_dir / "cifar-10-batches-py" + tgz = self.data_dir / "cifar-10-python.tar.gz" - tgz = DATA_DIR / "cifar-10-python.tar.gz" - if not tgz.exists(): + if not local_dir.exists(): + if not tgz.exists(): + try: print(" ~> downloading CIFAR-10 (~170 MB) ...", file=sys.stderr) _download(self.URL, tgz) - with tarfile.open(tgz) as tf: - # Validate all members to prevent path traversal attacks - for member in tf.getmembers(): - target_path = (DATA_DIR / member.name).resolve() - if not target_path.is_relative_to(DATA_DIR.resolve()): - raise ValueError(f"Unsafe path in archive: {member.name}") - tf.extractall(DATA_DIR) - import pickle - - xs, ys = [], [] - for i in range(1, 6): - with open(local_dir / f"data_batch_{i}", "rb") as f: - d = pickle.load(f, encoding="bytes") - xs.append(torch.tensor(d[b"data"], dtype=torch.float32)) - ys.extend(d[b"labels"]) - x = torch.cat(xs).view(-1, 3, 32, 32) / 255.0 - y = torch.tensor(ys, dtype=torch.long) - return x, y - except Exception as exc: - print( - f" ~> CIFAR download/parse failed ({type(exc).__name__}); using a fixed " - f"synthetic image set (conv path only).", - file=sys.stderr, - ) - g = torch.Generator().manual_seed(0) # dedicated -> global RNG untouched - x = torch.randn(2048, 3, 32, 32, generator=g) - y = torch.randint(0, 10, (2048,), generator=g) - return x, y + except _NETWORK_ERRORS as exc: + if tgz.exists(): + try: + tgz.unlink() + except OSError: + pass + print( + f" ~> CIFAR download failed ({type(exc).__name__}); using a fixed " + f"synthetic image set (conv path only).", + file=sys.stderr, + ) + return self._synthetic_dataset() + + import tarfile + + with tarfile.open(tgz) as tf: + # Validate all members to prevent path traversal attacks + for member in tf.getmembers(): + target_path = (self.data_dir / member.name).resolve() + if not target_path.is_relative_to(self.data_dir.resolve()): + raise ValueError(f"Unsafe path in archive: {member.name}") + tf.extractall(self.data_dir) + + import pickle + + xs, ys = [], [] + for i in range(1, 6): + batch_path = local_dir / f"data_batch_{i}" + with open(batch_path, "rb") as f: + d = pickle.load(f, encoding="bytes") + if not isinstance(d, dict) or b"data" not in d or b"labels" not in d: + raise ValueError( + f"Corrupted or invalid CIFAR batch format in {batch_path.name}" + ) + xs.append(torch.tensor(d[b"data"], dtype=torch.float32)) + ys.extend(d[b"labels"]) + x = torch.cat(xs).view(-1, 3, 32, 32) / 255.0 + y = torch.tensor(ys, dtype=torch.long) + return x, y def get_batch(self, batch_size=64, block_size=None, device="cpu"): n = self._images.size(0) @@ -223,12 +255,13 @@ def get_batch(self, batch_size=64, block_size=None, device="cpu"): return self._images[ix].to(device), self._labels[ix].to(device) -def get_dataset(name, block_size=128): +def get_dataset(name, block_size=128, data_dir=None): """Factory: text corpora -> CharDataset, 'cifar' -> CIFARDataset.""" if name.lower() == "cifar": - return CIFARDataset() + return CIFARDataset(data_dir=data_dir) return CharDataset(name=name, block_size=block_size) + # get_batch draws from the GLOBAL torch RNG -> replay-exact when called first in-loop diff --git a/tests/test_dataset.py b/tests/test_dataset.py new file mode 100644 index 00000000..5848de84 --- /dev/null +++ b/tests/test_dataset.py @@ -0,0 +1,149 @@ +import io +import pickle +import sys +import tarfile +import tempfile +import unittest +import urllib.error +from pathlib import Path +from unittest.mock import patch + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) + +from dataset import CIFARDataset, get_dataset # noqa: E402 + + +class TestCIFARDatasetIntegrity(unittest.TestCase): + def test_cifar_offline_fallback_on_network_error(self): + """Simulate an offline environment where CIFAR download fails. + + Should cleanly fall back to fixed synthetic dataset and clean up partial archives. + """ + with tempfile.TemporaryDirectory() as tmp_dir: + tmp_path = Path(tmp_dir) + with patch("dataset._download", side_effect=urllib.error.URLError("Network unreachable")): + dataset = CIFARDataset(data_dir=tmp_path) + self.assertEqual(dataset._images.shape, (2048, 3, 32, 32)) + self.assertEqual(dataset._labels.shape, (2048,)) + self.assertEqual(dataset.encoded.shape, (2048,)) + # Ensure no partial tar.gz artifact remained + self.assertFalse((tmp_path / "cifar-10-python.tar.gz").exists()) + + def test_cifar_raises_on_corrupt_tar_archive(self): + """Pre-existing corrupted archive must raise, never fall back to synthetic.""" + with tempfile.TemporaryDirectory() as tmp_dir: + tmp_path = Path(tmp_dir) + corrupt_tgz = tmp_path / "cifar-10-python.tar.gz" + corrupt_tgz.write_bytes(b"not a valid tar.gz file content") + + with self.assertRaises((tarfile.TarError, EOFError, OSError)): + CIFARDataset(data_dir=tmp_path) + + def test_cifar_raises_on_unsafe_tar_traversal(self): + """Tar archives containing path traversal members must raise ValueError.""" + with tempfile.TemporaryDirectory() as tmp_dir: + tmp_path = Path(tmp_dir) + tgz_path = tmp_path / "cifar-10-python.tar.gz" + + with tarfile.open(tgz_path, "w:gz") as tar: + payload = b"dummy" + info = tarfile.TarInfo(name="../escaped_file.txt") + info.size = len(payload) + tar.addfile(info, io.BytesIO(payload)) + + with self.assertRaises(ValueError) as ctx: + CIFARDataset(data_dir=tmp_path) + self.assertIn("Unsafe path in archive", str(ctx.exception)) + + def test_cifar_raises_on_corrupted_pickle_batch(self): + """Corrupted local batch files must raise UnpicklingError / EOFError.""" + with tempfile.TemporaryDirectory() as tmp_dir: + tmp_path = Path(tmp_dir) + batch_dir = tmp_path / "cifar-10-batches-py" + batch_dir.mkdir(parents=True) + + # Write corrupted garbage to data_batch_1 + (batch_dir / "data_batch_1").write_bytes(b"garbage-non-pickle-data") + + with self.assertRaises((pickle.UnpicklingError, EOFError, Exception)) as ctx: + CIFARDataset(data_dir=tmp_path) + # Must NOT be swallowed into synthetic fallback + self.assertNotIsInstance(ctx.exception, AssertionError) + + def test_cifar_raises_on_malformed_batch_structure(self): + """Valid pickle but invalid internal CIFAR dictionary structure must raise ValueError.""" + with tempfile.TemporaryDirectory() as tmp_dir: + tmp_path = Path(tmp_dir) + batch_dir = tmp_path / "cifar-10-batches-py" + batch_dir.mkdir(parents=True) + + # Valid pickle, but missing expected b"data" and b"labels" keys + malformed_dict = {b"wrong_key": [1, 2, 3]} + with open(batch_dir / "data_batch_1", "wb") as f: + pickle.dump(malformed_dict, f) + + with self.assertRaises(ValueError) as ctx: + CIFARDataset(data_dir=tmp_path) + self.assertIn("Corrupted or invalid CIFAR batch format", str(ctx.exception)) + + def test_cifar_raises_on_missing_batch_file(self): + """If batch directory exists but a required batch is missing, must raise FileNotFoundError.""" + with tempfile.TemporaryDirectory() as tmp_dir: + tmp_path = Path(tmp_dir) + batch_dir = tmp_path / "cifar-10-batches-py" + batch_dir.mkdir(parents=True) + + # Only write batch 1, batches 2..5 missing + valid_batch = {b"data": [[0] * 3072], b"labels": [1]} + with open(batch_dir / "data_batch_1", "wb") as f: + pickle.dump(valid_batch, f) + + with self.assertRaises(FileNotFoundError): + CIFARDataset(data_dir=tmp_path) + + def test_cifar_loads_valid_batches_correctly(self): + """When valid batch files 1..5 exist, load and concatenate tensors correctly.""" + with tempfile.TemporaryDirectory() as tmp_dir: + tmp_path = Path(tmp_dir) + batch_dir = tmp_path / "cifar-10-batches-py" + batch_dir.mkdir(parents=True) + + # Write valid batches 1 through 5 with 4 samples each + for i in range(1, 6): + sample_data = [[float(i * 10)] * 3072 for _ in range(4)] + sample_labels = [i] * 4 + with open(batch_dir / f"data_batch_{i}", "wb") as f: + pickle.dump({b"data": sample_data, b"labels": sample_labels}, f) + + dataset = CIFARDataset(data_dir=tmp_path) + self.assertEqual(dataset._images.shape, (20, 3, 32, 32)) + self.assertEqual(dataset._labels.shape, (20,)) + self.assertEqual(dataset.encoded.shape, (20,)) + + # Test get_batch sampling + xb, yb = dataset.get_batch(batch_size=8, device="cpu") + self.assertEqual(xb.shape, (8, 3, 32, 32)) + self.assertEqual(yb.shape, (8,)) + + def test_synthetic_generator_does_not_mutate_global_rng(self): + """Synthetic generator must use dedicated seed and leave global torch RNG untouched.""" + state_before = torch.get_rng_state() + CIFARDataset._synthetic_dataset(num_samples=16) + state_after = torch.get_rng_state() + self.assertTrue(torch.equal(state_before, state_after)) + + def test_get_dataset_factory(self): + """Factory function get_dataset handles 'cifar' with custom data_dir.""" + with tempfile.TemporaryDirectory() as tmp_dir: + tmp_path = Path(tmp_dir) + with patch("dataset._download", side_effect=urllib.error.URLError("Offline")): + ds = get_dataset("cifar", data_dir=tmp_path) + self.assertIsInstance(ds, CIFARDataset) + self.assertEqual(ds.name, "cifar") + self.assertEqual(ds.vocab_size, 10) + + +if __name__ == "__main__": + unittest.main() From fafe67391d0f6695b7be1bfd2d3b049f09f9328b Mon Sep 17 00:00:00 2001 From: Hrishikesh Yadav Date: Sun, 4 Oct 2026 13:12:37 +0530 Subject: [PATCH 2/3] test(dataset): tighten corrupted pickle assertion to explicit unpickling errors --- tests/test_dataset.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/tests/test_dataset.py b/tests/test_dataset.py index 5848de84..91fc8322 100644 --- a/tests/test_dataset.py +++ b/tests/test_dataset.py @@ -58,7 +58,7 @@ def test_cifar_raises_on_unsafe_tar_traversal(self): self.assertIn("Unsafe path in archive", str(ctx.exception)) def test_cifar_raises_on_corrupted_pickle_batch(self): - """Corrupted local batch files must raise UnpicklingError / EOFError.""" + """Corrupted local batch files must raise UnpicklingError or EOFError.""" with tempfile.TemporaryDirectory() as tmp_dir: tmp_path = Path(tmp_dir) batch_dir = tmp_path / "cifar-10-batches-py" @@ -67,10 +67,8 @@ def test_cifar_raises_on_corrupted_pickle_batch(self): # Write corrupted garbage to data_batch_1 (batch_dir / "data_batch_1").write_bytes(b"garbage-non-pickle-data") - with self.assertRaises((pickle.UnpicklingError, EOFError, Exception)) as ctx: + with self.assertRaises((pickle.UnpicklingError, EOFError)): CIFARDataset(data_dir=tmp_path) - # Must NOT be swallowed into synthetic fallback - self.assertNotIsInstance(ctx.exception, AssertionError) def test_cifar_raises_on_malformed_batch_structure(self): """Valid pickle but invalid internal CIFAR dictionary structure must raise ValueError.""" From ae3a33ec88f41693751d2817d4867e952e012cfd Mon Sep 17 00:00:00 2001 From: Hrishikesh Yadav Date: Sun, 4 Oct 2026 13:19:27 +0530 Subject: [PATCH 3/3] fix(dataset): scope network error handling and use safe tar data filter --- pyproject.toml | 2 +- src/dataset.py | 6 ++++-- tests/test_dataset.py | 8 ++++++++ 3 files changed, 13 insertions(+), 3 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index b131e105..ddfc2ce2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -7,7 +7,7 @@ name = "openverifiablellm" version = "0.1.0" description = "One-command verification for small open model artifacts." readme = "README.md" -requires-python = ">=3.10" +requires-python = ">=3.10.12" dependencies = [ "torch>=2.10.0,<3.0", "numpy>=2.4,<3.0", diff --git a/src/dataset.py b/src/dataset.py index 975a1162..18fca463 100644 --- a/src/dataset.py +++ b/src/dataset.py @@ -34,7 +34,6 @@ TimeoutError, ConnectionError, http.client.HTTPException, - OSError, ) DATA_DIR = Path(__file__).resolve().parents[1] / "data" @@ -230,7 +229,10 @@ def _load(self): target_path = (self.data_dir / member.name).resolve() if not target_path.is_relative_to(self.data_dir.resolve()): raise ValueError(f"Unsafe path in archive: {member.name}") - tf.extractall(self.data_dir) + if hasattr(tarfile, "data_filter"): + tf.extractall(self.data_dir, filter="data") + else: + tf.extractall(self.data_dir) import pickle diff --git a/tests/test_dataset.py b/tests/test_dataset.py index 91fc8322..e0b9bed8 100644 --- a/tests/test_dataset.py +++ b/tests/test_dataset.py @@ -142,6 +142,14 @@ def test_get_dataset_factory(self): self.assertEqual(ds.name, "cifar") self.assertEqual(ds.vocab_size, 10) + def test_cifar_raises_on_filesystem_permission_error(self): + """Filesystem errors during download must propagate and not trigger synthetic fallback.""" + with tempfile.TemporaryDirectory() as tmp_dir: + tmp_path = Path(tmp_dir) + with patch("dataset._download", side_effect=PermissionError("Read-only filesystem")): + with self.assertRaises(PermissionError): + CIFARDataset(data_dir=tmp_path) + if __name__ == "__main__": unittest.main()