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