poincare-hyper / tests /test_contracts.py
DHDRL's picture
Upload 30 files
ae73c7f verified
Raw
History Blame Contribute Delete
12.4 kB
"""
Test scenarios required before this codebase's core contracts could be
considered complete (per project guidelines: nothing is done without
test scenarios for duplicate requests, edge cases, and race conditions).
Run: python -m pytest tests/test_contracts.py -v
"""
from __future__ import annotations
import os
import sys
import shutil
import tempfile
from pathlib import Path
import numpy as np
import pytest
import torch
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from src.data_real import LocalWellHDF5, get_dataset, get_synthetic_dataset
from src.provenance import (
DataLoadError,
SchemaValidationError,
TrajectoryTooShortError,
EmptyDatasetError,
DatasetAlreadyUsedError,
DatasetInProgressError,
DatasetRegistry,
CheckpointStore,
hash_config,
hash_dataset,
hash_code,
combined_identity_hash,
validate_trajectory_lengths,
)
@pytest.fixture
def tmpdir():
d = tempfile.mkdtemp()
yield d
shutil.rmtree(d, ignore_errors=True)
# --------------------------------------------------------------------- #
# Scenario 1: missing real-data directory -> hard fail, no substitution
# --------------------------------------------------------------------- #
def test_get_dataset_hard_fails_when_no_real_data_directory_exists(tmpdir):
nonexistent = os.path.join(tmpdir, "does_not_exist")
with pytest.raises(DataLoadError) as exc_info:
get_dataset(search_roots=[nonexistent])
assert exc_info.value.outcome_code == "NO_DATA_DIRECTORY"
def test_get_dataset_hard_fails_on_empty_real_data_directory(tmpdir):
empty_dir = os.path.join(tmpdir, "real")
os.makedirs(empty_dir)
with pytest.raises(DataLoadError) as exc_info:
get_dataset(search_roots=[empty_dir])
assert exc_info.value.outcome_code == "NO_FILES_FOUND"
def test_get_dataset_never_silently_returns_synthetic_data(tmpdir):
nonexistent = os.path.join(tmpdir, "nope")
try:
ds, provenance = get_dataset(search_roots=[nonexistent])
assert False, "should have raised, not returned a dataset"
except DataLoadError:
pass # correct: hard failure, no dataset returned at all
# --------------------------------------------------------------------- #
# Scenario 2: HDF5 with schema that doesn't match declared expectations
# --------------------------------------------------------------------- #
def test_local_hdf5_rejects_wrong_channel_count(tmpdir):
h5py = pytest.importorskip("h5py")
real_dir = os.path.join(tmpdir, "real")
os.makedirs(real_dir)
fp = os.path.join(real_dir, "sample.hdf5")
with h5py.File(fp, "w") as f:
# (T=10, C=5, H=16, W=16) but caller expects 2 channels
f.create_dataset("fields", data=np.random.randn(10, 5, 16, 16).astype("float32"))
with pytest.raises(SchemaValidationError) as exc_info:
LocalWellHDF5(real_dir, expected_channels=2, channel_layout="channels_first", strict=True)
assert exc_info.value.outcome_code in ("CHANNEL_COUNT_MISMATCH", "NO_VALID_TRAJECTORIES")
def test_local_hdf5_rejects_wrong_ndim(tmpdir):
h5py = pytest.importorskip("h5py")
real_dir = os.path.join(tmpdir, "real")
os.makedirs(real_dir)
fp = os.path.join(real_dir, "sample.hdf5")
with h5py.File(fp, "w") as f:
f.create_dataset("fields", data=np.random.randn(10, 16, 16).astype("float32")) # 3D, no channel axis
with pytest.raises(SchemaValidationError):
LocalWellHDF5(real_dir, expected_channels=2, strict=True)
def test_local_hdf5_accepts_correctly_shaped_data(tmpdir):
h5py = pytest.importorskip("h5py")
real_dir = os.path.join(tmpdir, "real")
os.makedirs(real_dir)
fp = os.path.join(real_dir, "sample.hdf5")
with h5py.File(fp, "w") as f:
f.create_dataset("fields", data=np.random.randn(10, 2, 16, 16).astype("float32"))
ds = LocalWellHDF5(real_dir, expected_channels=2, channel_layout="channels_first", strict=True)
assert len(ds) == 1
item = ds[0]
assert item["fields"].shape == (10, 2, 16, 16)
# --------------------------------------------------------------------- #
# Scenario 3: trajectory-length validation fails loudly, before training
# --------------------------------------------------------------------- #
def test_validate_trajectory_lengths_raises_on_short_trajectories():
ds, _ = get_synthetic_dataset(max_samples=8, n_steps=5) # too short
with pytest.raises(TrajectoryTooShortError) as exc_info:
validate_trajectory_lengths(ds, required_length=10)
assert exc_info.value.outcome_code == "TRAJECTORY_TOO_SHORT"
def test_validate_trajectory_lengths_passes_on_sufficient_trajectories():
ds, _ = get_synthetic_dataset(max_samples=8, n_steps=14)
result = validate_trajectory_lengths(ds, required_length=8)
assert result["success"] is True
def test_validate_trajectory_lengths_raises_on_empty_dataset():
class Empty:
def __len__(self):
return 0
with pytest.raises(EmptyDatasetError):
validate_trajectory_lengths(Empty(), required_length=8)
# --------------------------------------------------------------------- #
# Scenario 4: checkpoint atomicity — no partial/corrupt file ever visible
# --------------------------------------------------------------------- #
def test_checkpoint_save_is_atomic_no_temp_file_left_behind(tmpdir):
store = CheckpointStore(checkpoints_dir=os.path.join(tmpdir, "ckpts"))
result = store.save(
model_state={"w": torch.randn(4, 4)},
config={"lr": 0.001, "hidden": 32},
dataset_hash="a" * 64,
code_hash="b" * 64,
data_provenance="SYNTHETIC",
)
assert result["success"] is True
assert result["outcome_code"] == "SAVED"
ckpt_dir = Path(tmpdir) / "ckpts"
files = list(ckpt_dir.iterdir())
# No .tmp_* files should remain after a successful save.
assert not any(f.name.startswith(".tmp_") for f in files)
# Exactly the final .pt and .meta.json should exist.
assert any(f.suffix == ".pt" for f in files)
assert any(f.name.endswith(".meta.json") for f in files)
def test_checkpoint_save_is_idempotent_for_identical_inputs(tmpdir):
store = CheckpointStore(checkpoints_dir=os.path.join(tmpdir, "ckpts"))
kwargs = dict(
model_state={"w": torch.randn(4, 4)},
config={"lr": 0.001},
dataset_hash="c" * 64,
code_hash="d" * 64,
data_provenance="SYNTHETIC",
)
r1 = store.save(**kwargs)
r2 = store.save(**kwargs)
assert r1["identity_hash"] == r2["identity_hash"]
assert r1["outcome_code"] == "SAVED"
assert r2["outcome_code"] == "DUPLICATE_EXISTS"
ckpt_dir = Path(tmpdir) / "ckpts"
pt_files = [f for f in ckpt_dir.iterdir() if f.suffix == ".pt"]
assert len(pt_files) == 1 # not duplicated
def test_checkpoint_save_rejects_missing_dataset_hash(tmpdir):
store = CheckpointStore(checkpoints_dir=os.path.join(tmpdir, "ckpts"))
with pytest.raises(Exception) as exc_info:
store.save(
model_state={"w": torch.randn(2, 2)},
config={"lr": 0.001},
dataset_hash="",
code_hash="e" * 64,
data_provenance="SYNTHETIC",
)
assert "MISSING_DATASET_HASH" in str(exc_info.value)
def test_checkpoint_load_detects_missing_meta_as_integrity_failure(tmpdir):
from src.provenance import CheckpointIntegrityError
store = CheckpointStore(checkpoints_dir=os.path.join(tmpdir, "ckpts"))
with pytest.raises(CheckpointIntegrityError):
store.load("nonexistent" * 8)
def test_checkpoint_extra_with_tensors_does_not_crash_meta_write(tmpdir):
store = CheckpointStore(checkpoints_dir=os.path.join(tmpdir, "ckpts"))
result = store.save(
model_state={"w": torch.randn(3, 3)},
config={"lr": 0.001},
dataset_hash="9" * 64,
code_hash="8" * 64,
data_provenance="SYNTHETIC",
extra={
"normalizer": {"mean": torch.tensor([0.1, 0.2]), "std": torch.tensor([1.0, 1.0])},
"ppo_returns": [1.0, 2.5, -3.2],
},
)
assert result["success"] is True
loaded = store.load(result["identity_hash"])
assert loaded["meta"]["extra"]["normalizer"]["mean"] == [pytest.approx(0.1), pytest.approx(0.2)]
def test_checkpoint_pt_without_meta_is_integrity_failure(tmpdir):
from src.provenance import CheckpointIntegrityError
ckpt_dir = Path(tmpdir) / "ckpts"
ckpt_dir.mkdir(parents=True)
fake_identity = "z" * 64
torch.save({"w": torch.randn(2, 2)}, ckpt_dir / f"{fake_identity}.pt")
# deliberately do NOT write the .meta.json sidecar
store = CheckpointStore(checkpoints_dir=str(ckpt_dir))
with pytest.raises(CheckpointIntegrityError):
store.load(fake_identity)
# --------------------------------------------------------------------- #
# Scenario 5: dataset-reuse registry blocks retraining on consumed data
# --------------------------------------------------------------------- #
def test_registry_blocks_retraining_on_already_consumed_dataset(tmpdir):
registry = DatasetRegistry(registry_dir=os.path.join(tmpdir, "registry"))
dataset_hash = "f" * 64
registry.claim(dataset_hash, experiment_id="exp1")
registry.mark_consumed(dataset_hash)
with pytest.raises(DatasetAlreadyUsedError):
registry.claim(dataset_hash, experiment_id="exp2")
def test_registry_blocks_concurrent_in_progress_claim(tmpdir):
registry = DatasetRegistry(registry_dir=os.path.join(tmpdir, "registry"))
dataset_hash = "0" * 64
registry.claim(dataset_hash, experiment_id="exp1") # first claim succeeds
with pytest.raises(DatasetInProgressError):
registry.claim(dataset_hash, experiment_id="exp2") # second is blocked
def test_registry_allows_retry_after_explicit_human_action(tmpdir):
registry = DatasetRegistry(registry_dir=os.path.join(tmpdir, "registry"))
dataset_hash = "1" * 64
registry.claim(dataset_hash, experiment_id="exp1")
registry.mark_failed(dataset_hash, error_detail="crashed")
# Second claim after a FAILED run is blocked until an explicit human
# action (allow_retry) is taken — not automatic.
with pytest.raises(DatasetInProgressError):
registry.claim(dataset_hash, experiment_id="exp2")
registry.allow_retry(dataset_hash)
result = registry.claim(dataset_hash, experiment_id="exp2") # now succeeds
assert result.success is True
def test_registry_status_none_for_unknown_dataset(tmpdir):
registry = DatasetRegistry(registry_dir=os.path.join(tmpdir, "registry"))
assert registry.status("nonexistent" * 8) is None
# --------------------------------------------------------------------- #
# Scenario 6: hashing determinism (identity must be stable and content-based)
# --------------------------------------------------------------------- #
def test_hash_config_is_order_independent():
h1 = hash_config({"lr": 0.001, "hidden": 64})
h2 = hash_config({"hidden": 64, "lr": 0.001})
assert h1 == h2
def test_hash_config_differs_for_different_values():
h1 = hash_config({"lr": 0.001})
h2 = hash_config({"lr": 0.002})
assert h1 != h2
def test_hash_dataset_is_content_based_not_path_based():
ds_a, _ = get_synthetic_dataset(max_samples=4, n_steps=8)
ds_b, _ = get_synthetic_dataset(max_samples=4, n_steps=8)
# Different random seeds inside SyntheticWellLike -> different content
# (this asserts the hash actually reflects content, not just shape).
assert hash_dataset(ds_a) != hash_dataset(ds_b) or torch.allclose(
ds_a[0]["fields"], ds_b[0]["fields"]
)
def test_hash_dataset_rejects_non_tensor_items():
class BadDataset:
def __len__(self):
return 1
def __getitem__(self, idx):
return {"fields": "not a tensor"}
with pytest.raises(SchemaValidationError):
hash_dataset(BadDataset())
def test_combined_identity_hash_changes_if_any_component_changes():
base = combined_identity_hash("cfg", "code", "data")
assert combined_identity_hash("cfg2", "code", "data") != base
assert combined_identity_hash("cfg", "code2", "data") != base
assert combined_identity_hash("cfg", "code", "data2") != base
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v"]))