Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
111 changes: 73 additions & 38 deletions src/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,13 +20,22 @@
"""

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,
)
Comment thread
coderabbitai[bot] marked this conversation as resolved.

DATA_DIR = Path(__file__).resolve().parents[1] / "data"

# Single reliable single-file sources. wikitext is handled via HF `datasets`.
Expand Down Expand Up @@ -76,7 +85,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(
Expand Down Expand Up @@ -170,65 +179,91 @@ 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}")
if hasattr(tarfile, "data_filter"):
tf.extractall(self.data_dir, filter="data")
else:
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)
ix = torch.randint(0, n, (batch_size,)) # global RNG -> replay-exact
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

155 changes: 155 additions & 0 deletions tests/test_dataset.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,155 @@
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 or 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)):
CIFARDataset(data_dir=tmp_path)

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)

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()