GNN4Colliders / tests /unit /data /test_huggingface_source.py
ho22joshua's picture
PR 2 — Portable Public End-to-End Validation Fixture (#11)
c20af2a
Raw History Blame Contribute Delete
1.72 kB
from __future__ import annotations
import hashlib
from pathlib import Path
import pytest
from gnn4colliders.data.huggingface import resolve_data_files
def test_huggingface_source_resolves_files_and_labels(monkeypatch, tmp_path):
root_file = tmp_path / "sample.root"
root_file.write_bytes(b"fixture")
monkeypatch.setattr(
"gnn4colliders.data.huggingface.hf_hub_download",
lambda **kwargs: str(root_file),
)
paths, labels = resolve_data_files(
{
"source": {
"type": "huggingface",
"repo_id": "org/data",
"revision": "abc123",
"files": [{"path": "sample.root", "label": 4}],
}
}
)
assert paths == [root_file]
assert labels == [4]
def test_huggingface_source_verifies_checksum(monkeypatch, tmp_path):
root_file = tmp_path / "sample.root"
root_file.write_bytes(b"fixture")
monkeypatch.setattr(
"gnn4colliders.data.huggingface.hf_hub_download",
lambda **kwargs: str(root_file),
)
checksum = hashlib.sha256(b"wrong").hexdigest()
with pytest.raises(ValueError, match="SHA-256 mismatch"):
resolve_data_files(
{
"source": {
"type": "huggingface",
"repo_id": "org/data",
"revision": "abc123",
"files": [{"path": "sample.root", "sha256": checksum}],
}
}
)
def test_local_source_remains_unchanged(tmp_path: Path):
local = tmp_path / "events.root"
paths, label = resolve_data_files({"files": [str(local)], "label": 2})
assert paths == [local]
assert label == 2