From b607eda0d460c59d82f066b0dbb48391b063177a Mon Sep 17 00:00:00 2001 From: Dhritiman Das <14159298+dhritimandas@users.noreply.github.com> Date: Fri, 10 Apr 2026 16:52:29 +0530 Subject: [PATCH 1/7] feat(data): add DataSpec validation and inspection system Add versioned dataset contract with DataLad-aware preflight validation: - DataSpec dataclass with JSON serialization - validate() with symlink-aware file checks, spacing/orientation/label validation - inspect_entry() for NIfTI/Zarr metadata inspection - CLI commands: validate, inspect, --skip-validate on predict/convert-to-zarr - 33 unit tests covering all branches --- nobrainer/cli/main.py | 180 ++++++++- nobrainer/data/__init__.py | 1 + nobrainer/data/spec.py | 520 +++++++++++++++++++++++++ nobrainer/tests/unit/test_data_spec.py | 397 +++++++++++++++++++ 4 files changed, 1097 insertions(+), 1 deletion(-) create mode 100644 nobrainer/data/__init__.py create mode 100644 nobrainer/data/spec.py create mode 100644 nobrainer/tests/unit/test_data_spec.py diff --git a/nobrainer/cli/main.py b/nobrainer/cli/main.py index 3b6ea348..2f087989 100644 --- a/nobrainer/cli/main.py +++ b/nobrainer/cli/main.py @@ -104,6 +104,12 @@ def cli(): @click.option( "-v", "--verbose", is_flag=True, help="Print progress messages.", **_option_kwds ) +@click.option( + "--skip-validate", + is_flag=True, + help="Skip preflight input-file validation.", + **_option_kwds, +) def predict( *, infile, @@ -117,11 +123,30 @@ def predict( n_samples, device, verbose, + skip_validate, ): """Predict labels from a NIfTI volume using a trained PyTorch model. The predictions are saved to OUTFILE. """ + # Preflight validation + if not skip_validate: + from ..data.spec import FileStatus, check_file_presence + + status = check_file_presence(infile) + if status == FileStatus.NOT_FOUND: + click.echo(click.style(f"ERROR: Input file not found: {infile}", fg="red")) + raise SystemExit(1) + if status == FileStatus.ANNEX_MISSING: + click.echo( + click.style( + f"ERROR: Input is a git-annex symlink with missing content.\n" + f" Run: datalad get {infile}", + fg="red", + ) + ) + raise SystemExit(1) + if os.path.exists(outfile): raise FileExistsError(f"Output file already exists: {outfile}") @@ -271,7 +296,14 @@ def convert_tfrecords(*, input_paths, output_dir, output_format, verbose): ) @click.option("--no-conform", is_flag=True, help="Disable auto-conforming.") @click.option("-v", "--verbose", is_flag=True, help="Print progress.") -def convert_to_zarr(*, output, images, labels, chunk_shape, no_conform, verbose): +@click.option( + "--skip-validate", + is_flag=True, + help="Skip preflight input-file validation.", +) +def convert_to_zarr( + *, output, images, labels, chunk_shape, no_conform, verbose, skip_validate +): """Convert NIfTI image+label pairs to a sharded Zarr3 store.""" from ..datasets.zarr_store import create_zarr_store @@ -283,6 +315,32 @@ def convert_to_zarr(*, output, images, labels, chunk_shape, no_conform, verbose) ) sys.exit(1) + # Preflight validation + if not skip_validate: + from ..data.spec import FileStatus, check_file_presence + + missing: list[str] = [] + annex_missing: list[str] = [] + for path in (*images, *labels): + status = check_file_presence(path) + if status == FileStatus.NOT_FOUND: + missing.append(path) + elif status == FileStatus.ANNEX_MISSING: + annex_missing.append(path) + if missing: + click.echo(click.style("ERROR: Files not found:", fg="red")) + for p in missing: + click.echo(f" {p}") + sys.exit(1) + if annex_missing: + click.echo( + click.style("ERROR: git-annex symlinks with missing content:", fg="red") + ) + for p in annex_missing: + click.echo(f" {p}") + click.echo(f"\n Run: datalad get {' '.join(annex_missing)}") + sys.exit(1) + pairs = list(zip(images, labels)) chunks = tuple(int(x) for x in chunk_shape.split(",")) @@ -623,6 +681,126 @@ def info(): click.echo(s) +# --------------------------------------------------------------------------- +# validate / inspect +# --------------------------------------------------------------------------- + + +@cli.command() +@click.argument("manifest", type=click.Path(exists=True)) +@click.option("--json", "as_json", is_flag=True, help="Output results as JSON.") +@click.option("-v", "--verbose", is_flag=True, help="Show per-subject details.") +def validate(*, manifest, as_json, verbose): + """Validate a dataset manifest against its DataSpec contract. + + MANIFEST is the path to a JSON manifest file. Exits non-zero if any + errors are found. Warnings alone do not cause a non-zero exit. + """ + import json as _json + + from ..data.spec import DataSpec + from ..data.spec import Severity as _Sev + from ..data.spec import validate as _validate + + spec = DataSpec.from_json(manifest) + findings = _validate(spec) + + if as_json: + click.echo(_json.dumps([f.to_dict() for f in findings], indent=2)) + else: + errs = [f for f in findings if f.severity == _Sev.ERROR] + warns = [f for f in findings if f.severity == _Sev.WARNING] + + for f in errs: + click.echo(click.style(f"ERROR {f.field}: {f.message}", fg="red")) + for f in warns: + click.echo(click.style(f"WARN {f.field}: {f.message}", fg="yellow")) + + if not findings: + click.echo(click.style("Validation passed.", fg="green")) + else: + click.echo(f"\n{len(errs)} error(s), {len(warns)} warning(s).") + + has_errors = any(f.severity == _Sev.ERROR for f in findings) + sys.exit(1 if has_errors else 0) + + +@cli.command() +@click.argument("path", type=click.Path(exists=True)) +@click.option("--json", "as_json", is_flag=True, help="Output as JSON.") +def inspect(*, path, as_json): + """Inspect NIfTI / Zarr files or a dataset manifest. + + PATH may be a JSON manifest, a single NIfTI file, a .zarr store, + or a directory containing NIfTI / Zarr files. + """ + import json as _json + from pathlib import Path as _Path + + from ..data.spec import inspect_entry + + target = _Path(path) + entries: list[dict] = [] + + if target.suffix == ".json": + with open(target) as fh: + manifest = _json.load(fh) + base_dir = target.parent + for entry in manifest.get("entries", []): + for key in ("image", "label"): + if key in entry: + p = _Path(entry[key]) + if not p.is_absolute(): + p = base_dir / p + entries.append(inspect_entry(p)) + elif target.is_dir(): + for child in sorted(target.iterdir()): + if child.suffix in (".gz", ".nii", ".zarr") or child.suffixes == [ + ".nii", + ".gz", + ]: + entries.append(inspect_entry(child)) + else: + entries.append(inspect_entry(target)) + + if as_json: + click.echo(_json.dumps(entries, indent=2)) + else: + counts = {"present": 0, "annex_missing": 0, "not_found": 0} + for e in entries: + status = e["file_status"] + counts[status] = counts.get(status, 0) + 1 + if status == "present": + shape = tuple(e["shape"]) if e["shape"] else "?" + spacing = tuple(e["spacing"]) if e["spacing"] else "?" + orient = e["orientation"] or "?" + size_mb = ( + f"{e['file_size_bytes'] / 1_048_576:.1f}MB" + if e["file_size_bytes"] + else "?" + ) + click.echo( + f"{e['path']} PRESENT shape={shape} " + f"spacing={spacing} orient={orient} {size_mb}" + ) + elif status == "annex_missing": + click.echo( + click.style( + f"{e['path']} ANNEX_MISSING " + f"(run: datalad get {e['path']})", + fg="yellow", + ) + ) + else: + click.echo(click.style(f"{e['path']} NOT_FOUND", fg="red")) + click.echo("---") + click.echo( + f"{len(entries)} entries: {counts['present']} present, " + f"{counts['annex_missing']} annex_missing, " + f"{counts['not_found']} not_found" + ) + + # --------------------------------------------------------------------------- # zarr subcommands # --------------------------------------------------------------------------- diff --git a/nobrainer/data/__init__.py b/nobrainer/data/__init__.py new file mode 100644 index 00000000..6dc1cfde --- /dev/null +++ b/nobrainer/data/__init__.py @@ -0,0 +1 @@ +"""Dataset specifications and static data for nobrainer.""" diff --git a/nobrainer/data/spec.py b/nobrainer/data/spec.py new file mode 100644 index 00000000..e5b461a7 --- /dev/null +++ b/nobrainer/data/spec.py @@ -0,0 +1,520 @@ +"""Dataset specification, validation, and inspection. + +Self-contained module — imports only stdlib + nibabel + numpy. +Does NOT import from ``nobrainer.*`` so it can be used by external +CI scripts without installing torch/monai. +""" + +from __future__ import annotations + +import dataclasses +import enum +import json +import logging +from pathlib import Path + +import nibabel as nib +import numpy as np + +logger = logging.getLogger(__name__) + +# --------------------------------------------------------------------------- +# Enums +# --------------------------------------------------------------------------- + +VALID_AXIS_CODES = frozenset("RLAPIS") + + +class FileStatus(enum.Enum): + """Result of a symlink-aware file presence check.""" + + PRESENT = "present" + ANNEX_MISSING = "annex_missing" + NOT_FOUND = "not_found" + + +class Severity(enum.Enum): + """Severity level for a validation finding.""" + + ERROR = "error" + WARNING = "warning" + + +# --------------------------------------------------------------------------- +# Dataclasses +# --------------------------------------------------------------------------- + + +@dataclasses.dataclass(frozen=True) +class ValidationError: + """A single validation finding. + + Attributes: + field: Dotted path to the offending field, e.g. ``"entries[3].image"``. + subject_index: Entry index, or ``None`` for spec-level errors. + message: Human-readable description of the problem. + severity: ``Severity.ERROR`` or ``Severity.WARNING``. + """ + + field: str + subject_index: int | None + message: str + severity: Severity + + def to_dict(self) -> dict: + """Serialize to a JSON-friendly dict.""" + return { + "field": self.field, + "subject_index": self.subject_index, + "message": self.message, + "severity": self.severity.value, + } + + +@dataclasses.dataclass +class DataSpec: + """Dataset manifest contract. + + Attributes: + entries: List of ``{"image": path}`` or ``{"image": path, "label": path}`` + dicts. Mirrors the MONAI-style data dicts used by + ``nobrainer.dataset.get_dataset()``. + expected_classes: If set, the allowed integer label values. + spacing_range: ``(min_mm, max_mm)`` — every voxel spacing axis must + fall within this range. + orientation: Expected 3-character orientation code (e.g. ``"RAS"``). + zarr_chunk_shape: Expected Zarr inner-chunk dimensions. + zarr_levels: Expected number of Zarr pyramid levels. + """ + + entries: list[dict[str, str]] = dataclasses.field(default_factory=list) + expected_classes: set[int] | None = None + spacing_range: tuple[float, float] | None = None + orientation: str | None = None + zarr_chunk_shape: tuple[int, ...] | None = None + zarr_levels: int | None = None + + # -- Serialization ------------------------------------------------------- + + @classmethod + def from_json(cls, path: str | Path) -> DataSpec: + """Load a manifest JSON file and return a ``DataSpec``. + + Relative entry paths are resolved against the manifest's parent + directory. + """ + path = Path(path) + with open(path) as fh: + raw = json.load(fh) + + base_dir = path.parent + + entries: list[dict[str, str]] = [] + for entry in raw.get("entries", []): + resolved: dict[str, str] = {} + for key in ("image", "label"): + if key in entry: + p = Path(entry[key]) + if not p.is_absolute(): + p = base_dir / p + resolved[key] = str(p) + entries.append(resolved) + + expected_classes = raw.get("expected_classes") + if expected_classes is not None: + expected_classes = set(expected_classes) + + spacing_range = raw.get("spacing_range") + if spacing_range is not None: + spacing_range = tuple(spacing_range) + + zarr_chunk_shape = raw.get("zarr_chunk_shape") + if zarr_chunk_shape is not None: + zarr_chunk_shape = tuple(zarr_chunk_shape) + + zarr_levels = raw.get("zarr_levels") + + return cls( + entries=entries, + expected_classes=expected_classes, + spacing_range=spacing_range, + orientation=raw.get("orientation"), + zarr_chunk_shape=zarr_chunk_shape, + zarr_levels=zarr_levels, + ) + + def to_json(self, path: str | Path) -> None: + """Write this spec to a JSON manifest file.""" + data: dict = {"entries": self.entries} + if self.expected_classes is not None: + data["expected_classes"] = sorted(self.expected_classes) + if self.spacing_range is not None: + data["spacing_range"] = list(self.spacing_range) + if self.orientation is not None: + data["orientation"] = self.orientation + if self.zarr_chunk_shape is not None: + data["zarr_chunk_shape"] = list(self.zarr_chunk_shape) + if self.zarr_levels is not None: + data["zarr_levels"] = self.zarr_levels + with open(path, "w") as fh: + json.dump(data, fh, indent=2) + fh.write("\n") + + +# --------------------------------------------------------------------------- +# File-presence helper +# --------------------------------------------------------------------------- + + +def check_file_presence(path: str | Path) -> FileStatus: + """Symlink-aware file presence check. + + Returns: + ``FileStatus.PRESENT`` if the file exists and is readable. + ``FileStatus.ANNEX_MISSING`` if the path is a symlink whose target + does not exist (typical of un-fetched git-annex / DataLad content). + ``FileStatus.NOT_FOUND`` if the path does not exist at all. + """ + p = Path(path) + if p.is_symlink() and not p.exists(): + return FileStatus.ANNEX_MISSING + if p.exists(): + return FileStatus.PRESENT + return FileStatus.NOT_FOUND + + +# --------------------------------------------------------------------------- +# Validation +# --------------------------------------------------------------------------- + + +def validate(spec: DataSpec) -> list[ValidationError]: + """Check a ``DataSpec`` against its constraints. + + Returns a (possibly empty) list of ``ValidationError`` findings. + Header-only reads for spacing/orientation; label-data read only when + ``expected_classes`` is set. + """ + errors: list[ValidationError] = [] + + # -- Phase 1: spec-level checks (no I/O) -------------------------------- + + if not spec.entries: + errors.append( + ValidationError( + field="entries", + subject_index=None, + message="Entries list is empty.", + severity=Severity.ERROR, + ) + ) + return errors # nothing else to check + + for i, entry in enumerate(spec.entries): + if "image" not in entry: + errors.append( + ValidationError( + field=f"entries[{i}]", + subject_index=i, + message="Entry is missing required 'image' key.", + severity=Severity.ERROR, + ) + ) + + if spec.spacing_range is not None: + lo, hi = spec.spacing_range + if lo <= 0 or hi <= 0: + errors.append( + ValidationError( + field="spacing_range", + subject_index=None, + message=f"Spacing range values must be positive, got ({lo}, {hi}).", + severity=Severity.ERROR, + ) + ) + elif lo > hi: + errors.append( + ValidationError( + field="spacing_range", + subject_index=None, + message=f"Spacing range min ({lo}) > max ({hi}).", + severity=Severity.ERROR, + ) + ) + + if spec.orientation is not None: + if len(spec.orientation) != 3 or not all( + c in VALID_AXIS_CODES for c in spec.orientation.upper() + ): + errors.append( + ValidationError( + field="orientation", + subject_index=None, + message=( + f"Invalid orientation code '{spec.orientation}'. " + f"Expected 3 characters from {{R,L,A,P,I,S}}." + ), + severity=Severity.ERROR, + ) + ) + + if spec.zarr_chunk_shape is not None: + if any(d <= 0 for d in spec.zarr_chunk_shape): + errors.append( + ValidationError( + field="zarr_chunk_shape", + subject_index=None, + message="All chunk shape dimensions must be > 0.", + severity=Severity.ERROR, + ) + ) + + if spec.zarr_levels is not None and spec.zarr_levels < 1: + errors.append( + ValidationError( + field="zarr_levels", + subject_index=None, + message=f"zarr_levels must be >= 1, got {spec.zarr_levels}.", + severity=Severity.ERROR, + ) + ) + + # -- Phase 2 & 3: per-entry file presence + header validation ------------ + + for i, entry in enumerate(spec.entries): + if "image" not in entry: + continue # already reported in Phase 1 + + # Track which files are present so Phase 3 can read them. + image_present = _check_entry_file( + errors, entry["image"], f"entries[{i}].image", i + ) + + label_path = entry.get("label") + label_present = False + if label_path is not None: + label_present = _check_entry_file( + errors, label_path, f"entries[{i}].label", i + ) + + # Phase 3: header validation for present files + if image_present: + _validate_nifti_header( + errors, entry["image"], f"entries[{i}].image", i, spec + ) + + if label_present and spec.expected_classes is not None: + _validate_label_classes( + errors, label_path, f"entries[{i}].label", i, spec.expected_classes + ) + + return errors + + +# --------------------------------------------------------------------------- +# Validation helpers (private) +# --------------------------------------------------------------------------- + + +def _check_entry_file( + errors: list[ValidationError], + path: str, + field: str, + index: int, +) -> bool: + """Check file presence; append errors; return True if PRESENT.""" + status = check_file_presence(path) + if status == FileStatus.NOT_FOUND: + errors.append( + ValidationError( + field=field, + subject_index=index, + message=f"File does not exist: {path}", + severity=Severity.ERROR, + ) + ) + return False + if status == FileStatus.ANNEX_MISSING: + errors.append( + ValidationError( + field=field, + subject_index=index, + message=( + "File is a git-annex symlink but content is not present " + f"locally. Run: datalad get {path}" + ), + severity=Severity.ERROR, + ) + ) + return False + return True + + +def _validate_nifti_header( + errors: list[ValidationError], + path: str, + field: str, + index: int, + spec: DataSpec, +) -> None: + """Read NIfTI header (no full data load) and check spacing/orientation.""" + try: + img = nib.load(path) + except Exception as exc: + errors.append( + ValidationError( + field=field, + subject_index=index, + message=f"Failed to read NIfTI header: {exc}", + severity=Severity.ERROR, + ) + ) + return + + # Spacing check + if spec.spacing_range is not None: + lo, hi = spec.spacing_range + zooms = img.header.get_zooms()[:3] + for axis_idx, z in enumerate(zooms): + if not (lo <= z <= hi): + errors.append( + ValidationError( + field=field, + subject_index=index, + message=( + f"Voxel spacing axis {axis_idx} = {z:.4f} mm " + f"is outside range [{lo}, {hi}]." + ), + severity=Severity.WARNING, + ) + ) + + # Orientation check + if spec.orientation is not None: + axcodes = "".join(nib.aff2axcodes(img.affine)) + if axcodes != spec.orientation.upper(): + errors.append( + ValidationError( + field=field, + subject_index=index, + message=( + f"Orientation is '{axcodes}', " + f"expected '{spec.orientation}'." + ), + severity=Severity.WARNING, + ) + ) + + +def _validate_label_classes( + errors: list[ValidationError], + path: str, + field: str, + index: int, + expected_classes: set[int], +) -> None: + """Read label volume data and check for unexpected class values.""" + try: + img = nib.load(path) + data = np.asarray(img.dataobj) + found = set(np.unique(data).astype(int).tolist()) + del data # free memory immediately + except Exception as exc: + errors.append( + ValidationError( + field=field, + subject_index=index, + message=f"Failed to read label data: {exc}", + severity=Severity.ERROR, + ) + ) + return + + unexpected = found - expected_classes + if unexpected: + errors.append( + ValidationError( + field=field, + subject_index=index, + message=( + f"Label contains unexpected classes: {sorted(unexpected)}. " + f"Expected: {sorted(expected_classes)}." + ), + severity=Severity.WARNING, + ) + ) + + +# --------------------------------------------------------------------------- +# Inspection +# --------------------------------------------------------------------------- + + +def inspect_entry(path: str | Path) -> dict: + """Return a summary dict for a single NIfTI or Zarr file. + + The dict always contains ``"path"`` and ``"file_status"`` keys. + Shape, spacing, orientation, dtype, and file_size_bytes are populated + only when the file is present and readable. + """ + p = Path(path) + status = check_file_presence(p) + result: dict = { + "path": str(p), + "file_status": status.value, + "shape": None, + "spacing": None, + "orientation": None, + "dtype": None, + "file_size_bytes": None, + } + + if status != FileStatus.PRESENT: + return result + + # File size (follows symlinks) + try: + result["file_size_bytes"] = p.stat().st_size + except OSError: + pass + + # Zarr store + suffix = "".join(p.suffixes) + if suffix == ".zarr" or p.is_dir(): + return _inspect_zarr(result, p) + + # NIfTI + return _inspect_nifti(result, p) + + +def _inspect_nifti(result: dict, p: Path) -> dict: + """Populate result dict from a NIfTI header.""" + try: + img = nib.load(str(p)) + result["shape"] = list(img.shape[:3]) + result["spacing"] = [round(float(z), 4) for z in img.header.get_zooms()[:3]] + result["orientation"] = "".join(nib.aff2axcodes(img.affine)) + result["dtype"] = str(img.header.get_data_dtype()) + except Exception: + logger.debug("Could not read NIfTI header for %s", p, exc_info=True) + return result + + +def _inspect_zarr(result: dict, p: Path) -> dict: + """Populate result dict from Zarr store metadata.""" + try: + import zarr + + store = zarr.open_group(str(p), mode="r") + attrs = dict(store.attrs) + result["shape"] = attrs.get("volume_shape") + result["dtype"] = attrs.get("image_dtype") + if "n_levels" in attrs: + result["zarr_levels"] = attrs["n_levels"] + if "chunk_shape" in attrs: + result["zarr_chunk_shape"] = attrs["chunk_shape"] + if "n_subjects" in attrs: + result["n_subjects"] = attrs["n_subjects"] + except Exception: + logger.debug("Could not read Zarr metadata for %s", p, exc_info=True) + return result diff --git a/nobrainer/tests/unit/test_data_spec.py b/nobrainer/tests/unit/test_data_spec.py new file mode 100644 index 00000000..356a4cca --- /dev/null +++ b/nobrainer/tests/unit/test_data_spec.py @@ -0,0 +1,397 @@ +"""Tests for nobrainer.data.spec — DataSpec validation and inspection.""" + +from __future__ import annotations + +import json +import os +from pathlib import Path + +from click.testing import CliRunner +import nibabel as nib +import numpy as np +import pytest + +from nobrainer.cli.main import cli +from nobrainer.data.spec import ( + DataSpec, + FileStatus, + Severity, + ValidationError, + check_file_presence, + inspect_entry, + validate, +) + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_nifti( + tmp_path: Path, + name: str, + shape: tuple[int, ...] = (16, 16, 16), + spacing: tuple[float, ...] = (1.0, 1.0, 1.0), + label_classes: list[int] | None = None, +) -> str: + """Create a NIfTI file and return its path as a string.""" + affine = np.diag([*spacing, 1.0]) + if label_classes is not None: + data = np.random.choice(label_classes, size=shape).astype(np.int32) + else: + data = np.random.rand(*shape).astype(np.float32) + path = tmp_path / name + nib.save(nib.Nifti1Image(data, affine), str(path)) + return str(path) + + +def _make_manifest(tmp_path: Path, spec_dict: dict) -> str: + """Write a manifest JSON and return its path.""" + path = tmp_path / "manifest.json" + with open(path, "w") as fh: + json.dump(spec_dict, fh) + return str(path) + + +# --------------------------------------------------------------------------- +# check_file_presence +# --------------------------------------------------------------------------- + + +class TestCheckFilePresence: + def test_present_file(self, tmp_path: Path) -> None: + f = tmp_path / "real.nii.gz" + f.write_bytes(b"data") + assert check_file_presence(f) == FileStatus.PRESENT + + def test_not_found(self, tmp_path: Path) -> None: + assert check_file_presence(tmp_path / "nope.nii.gz") == FileStatus.NOT_FOUND + + def test_annex_missing(self, tmp_path: Path) -> None: + link = tmp_path / "annex_file.nii.gz" + os.symlink("/nonexistent/.git/annex/objects/xx/yy/data", str(link)) + assert check_file_presence(link) == FileStatus.ANNEX_MISSING + + def test_valid_symlink(self, tmp_path: Path) -> None: + real = tmp_path / "real.nii.gz" + real.write_bytes(b"data") + link = tmp_path / "link.nii.gz" + os.symlink(str(real), str(link)) + assert check_file_presence(link) == FileStatus.PRESENT + + +# --------------------------------------------------------------------------- +# ValidationError +# --------------------------------------------------------------------------- + + +class TestValidationError: + def test_frozen(self) -> None: + err = ValidationError( + field="entries[0].image", + subject_index=0, + message="bad", + severity=Severity.ERROR, + ) + with pytest.raises(AttributeError): + err.field = "other" # type: ignore[misc] + + def test_to_dict(self) -> None: + err = ValidationError("f", 1, "msg", Severity.WARNING) + d = err.to_dict() + assert d == { + "field": "f", + "subject_index": 1, + "message": "msg", + "severity": "warning", + } + + +# --------------------------------------------------------------------------- +# DataSpec serialization +# --------------------------------------------------------------------------- + + +class TestDataSpecSerialization: + def test_from_json_roundtrip(self, tmp_path: Path) -> None: + img = _make_nifti(tmp_path, "img.nii.gz") + lbl = _make_nifti(tmp_path, "lbl.nii.gz", label_classes=[0, 1, 2]) + spec = DataSpec( + entries=[{"image": img, "label": lbl}], + expected_classes={0, 1, 2}, + spacing_range=(0.5, 2.0), + orientation="RAS", + zarr_chunk_shape=(32, 32, 32), + zarr_levels=3, + ) + out = tmp_path / "spec.json" + spec.to_json(out) + loaded = DataSpec.from_json(out) + assert loaded.expected_classes == spec.expected_classes + assert loaded.spacing_range == spec.spacing_range + assert loaded.orientation == spec.orientation + assert loaded.zarr_chunk_shape == spec.zarr_chunk_shape + assert loaded.zarr_levels == spec.zarr_levels + assert len(loaded.entries) == 1 + + def test_relative_paths_resolved(self, tmp_path: Path) -> None: + _make_nifti(tmp_path, "sub01.nii.gz") + manifest = {"entries": [{"image": "sub01.nii.gz"}]} + mpath = _make_manifest(tmp_path, manifest) + spec = DataSpec.from_json(mpath) + assert Path(spec.entries[0]["image"]).is_absolute() + + +# --------------------------------------------------------------------------- +# validate() +# --------------------------------------------------------------------------- + + +class TestValidate: + def test_valid_spec_passes(self, tmp_path: Path) -> None: + img = _make_nifti(tmp_path, "img.nii.gz") + lbl = _make_nifti(tmp_path, "lbl.nii.gz", label_classes=[0, 1, 2]) + spec = DataSpec( + entries=[{"image": img, "label": lbl}], + expected_classes={0, 1, 2}, + spacing_range=(0.5, 1.5), + orientation="RAS", + ) + errors = validate(spec) + assert errors == [] + + def test_empty_entries_error(self) -> None: + spec = DataSpec(entries=[]) + errors = validate(spec) + assert len(errors) == 1 + assert errors[0].severity == Severity.ERROR + assert errors[0].field == "entries" + + def test_missing_image_key(self) -> None: + spec = DataSpec(entries=[{"label": "/some/path"}]) + errors = validate(spec) + assert any( + e.severity == Severity.ERROR and "'image'" in e.message for e in errors + ) + + def test_missing_file_detected(self, tmp_path: Path) -> None: + spec = DataSpec(entries=[{"image": str(tmp_path / "nope.nii.gz")}]) + errors = validate(spec) + assert any( + e.severity == Severity.ERROR and "does not exist" in e.message + for e in errors + ) + + def test_annex_missing_detected(self, tmp_path: Path) -> None: + link = tmp_path / "annex.nii.gz" + os.symlink("/nonexistent/.git/annex/objects/xx/yy/data", str(link)) + spec = DataSpec(entries=[{"image": str(link)}]) + errors = validate(spec) + assert len(errors) == 1 + assert errors[0].severity == Severity.ERROR + assert "datalad get" in errors[0].message + # Must not crash trying to read the header + + def test_annex_present_passes(self, tmp_path: Path) -> None: + real = _make_nifti(tmp_path, "real.nii.gz") + link = tmp_path / "link.nii.gz" + os.symlink(real, str(link)) + spec = DataSpec(entries=[{"image": str(link)}]) + errors = validate(spec) + assert errors == [] + + def test_spacing_mismatch_detected(self, tmp_path: Path) -> None: + img = _make_nifti(tmp_path, "wide.nii.gz", spacing=(3.0, 3.0, 3.0)) + spec = DataSpec( + entries=[{"image": img}], + spacing_range=(0.5, 2.0), + ) + errors = validate(spec) + warnings = [e for e in errors if e.severity == Severity.WARNING] + assert len(warnings) >= 1 + assert any("outside range" in w.message for w in warnings) + + def test_spacing_in_range(self, tmp_path: Path) -> None: + img = _make_nifti(tmp_path, "ok.nii.gz", spacing=(1.0, 1.0, 1.0)) + spec = DataSpec( + entries=[{"image": img}], + spacing_range=(0.5, 2.0), + ) + errors = validate(spec) + assert not any( + e.severity == Severity.WARNING and "spacing" in e.message.lower() + for e in errors + ) + + def test_orientation_mismatch(self, tmp_path: Path) -> None: + img = _make_nifti(tmp_path, "ras.nii.gz") # identity affine → RAS + spec = DataSpec( + entries=[{"image": img}], + orientation="LPI", + ) + errors = validate(spec) + warnings = [e for e in errors if e.severity == Severity.WARNING] + assert any("Orientation" in w.message for w in warnings) + + def test_orientation_match(self, tmp_path: Path) -> None: + img = _make_nifti(tmp_path, "ras.nii.gz") # identity affine → RAS + spec = DataSpec( + entries=[{"image": img}], + orientation="RAS", + ) + errors = validate(spec) + assert not any( + e.severity == Severity.WARNING and "Orientation" in e.message + for e in errors + ) + + def test_label_class_violation(self, tmp_path: Path) -> None: + img = _make_nifti(tmp_path, "img.nii.gz") + lbl = _make_nifti(tmp_path, "lbl.nii.gz", label_classes=[0, 1, 2, 99]) + spec = DataSpec( + entries=[{"image": img, "label": lbl}], + expected_classes={0, 1, 2}, + ) + errors = validate(spec) + warnings = [e for e in errors if e.severity == Severity.WARNING] + assert any("unexpected classes" in w.message for w in warnings) + + def test_label_classes_match(self, tmp_path: Path) -> None: + img = _make_nifti(tmp_path, "img.nii.gz") + lbl = _make_nifti(tmp_path, "lbl.nii.gz", label_classes=[0, 1, 2]) + spec = DataSpec( + entries=[{"image": img, "label": lbl}], + expected_classes={0, 1, 2}, + ) + errors = validate(spec) + assert not any("unexpected" in e.message for e in errors) + + def test_invalid_spacing_range(self) -> None: + spec = DataSpec( + entries=[{"image": "/dummy"}], + spacing_range=(2.0, 0.5), + ) + errors = validate(spec) + assert any( + e.severity == Severity.ERROR and "spacing_range" in e.field for e in errors + ) + + def test_invalid_chunk_shape(self) -> None: + spec = DataSpec( + entries=[{"image": "/dummy"}], + zarr_chunk_shape=(32, 0, 32), + ) + errors = validate(spec) + assert any( + e.severity == Severity.ERROR and "zarr_chunk_shape" in e.field + for e in errors + ) + + +# --------------------------------------------------------------------------- +# inspect_entry +# --------------------------------------------------------------------------- + + +class TestInspectEntry: + def test_present_nifti(self, tmp_path: Path) -> None: + img = _make_nifti(tmp_path, "vol.nii.gz", shape=(32, 32, 32)) + result = inspect_entry(img) + assert result["file_status"] == "present" + assert result["shape"] == [32, 32, 32] + assert result["spacing"] is not None + assert result["orientation"] == "RAS" + assert result["file_size_bytes"] > 0 + + def test_not_found(self, tmp_path: Path) -> None: + result = inspect_entry(tmp_path / "nope.nii.gz") + assert result["file_status"] == "not_found" + assert result["shape"] is None + + def test_annex_missing(self, tmp_path: Path) -> None: + link = tmp_path / "annex.nii.gz" + os.symlink("/nonexistent/.git/annex/objects/xx", str(link)) + result = inspect_entry(link) + assert result["file_status"] == "annex_missing" + assert result["shape"] is None + + +# --------------------------------------------------------------------------- +# CLI: validate +# --------------------------------------------------------------------------- + + +class TestCLIValidate: + def test_valid_manifest_exit_zero(self, tmp_path: Path) -> None: + _make_nifti(tmp_path, "img.nii.gz") + manifest = _make_manifest(tmp_path, {"entries": [{"image": "img.nii.gz"}]}) + runner = CliRunner() + result = runner.invoke(cli, ["validate", manifest]) + assert result.exit_code == 0 + assert "passed" in result.output.lower() + + def test_invalid_manifest_exit_nonzero(self, tmp_path: Path) -> None: + manifest = _make_manifest( + tmp_path, {"entries": [{"image": "nonexistent.nii.gz"}]} + ) + runner = CliRunner() + result = runner.invoke(cli, ["validate", manifest]) + assert result.exit_code == 1 + assert "ERROR" in result.output + + def test_json_output(self, tmp_path: Path) -> None: + _make_nifti(tmp_path, "img.nii.gz") + manifest = _make_manifest(tmp_path, {"entries": [{"image": "img.nii.gz"}]}) + runner = CliRunner() + result = runner.invoke(cli, ["validate", manifest, "--json"]) + assert result.exit_code == 0 + parsed = json.loads(result.output) + assert isinstance(parsed, list) + + def test_annex_missing_shows_datalad_get(self, tmp_path: Path) -> None: + link = tmp_path / "annex.nii.gz" + os.symlink("/nonexistent/.git/annex/objects/xx", str(link)) + manifest = _make_manifest(tmp_path, {"entries": [{"image": "annex.nii.gz"}]}) + runner = CliRunner() + result = runner.invoke(cli, ["validate", manifest]) + assert result.exit_code == 1 + assert "datalad get" in result.output + + +# --------------------------------------------------------------------------- +# CLI: inspect +# --------------------------------------------------------------------------- + + +class TestCLIInspect: + def test_single_nifti(self, tmp_path: Path) -> None: + img = _make_nifti(tmp_path, "vol.nii.gz") + runner = CliRunner() + result = runner.invoke(cli, ["inspect", img]) + assert result.exit_code == 0 + assert "PRESENT" in result.output + + def test_manifest_json(self, tmp_path: Path) -> None: + _make_nifti(tmp_path, "img.nii.gz") + manifest = _make_manifest(tmp_path, {"entries": [{"image": "img.nii.gz"}]}) + runner = CliRunner() + result = runner.invoke(cli, ["inspect", manifest]) + assert result.exit_code == 0 + assert "1 entries" in result.output + + def test_json_output_has_file_status(self, tmp_path: Path) -> None: + img = _make_nifti(tmp_path, "vol.nii.gz") + runner = CliRunner() + result = runner.invoke(cli, ["inspect", img, "--json"]) + assert result.exit_code == 0 + parsed = json.loads(result.output) + assert isinstance(parsed, list) + assert parsed[0]["file_status"] == "present" + + def test_directory_scan(self, tmp_path: Path) -> None: + _make_nifti(tmp_path, "a.nii.gz") + _make_nifti(tmp_path, "b.nii.gz") + runner = CliRunner() + result = runner.invoke(cli, ["inspect", str(tmp_path)]) + assert result.exit_code == 0 + assert "2 entries" in result.output From a2d8c3312f84c66635b039eb436354ebcd52e623 Mon Sep 17 00:00:00 2001 From: Dhritiman Das <14159298+dhritimandas@users.noreply.github.com> Date: Fri, 10 Apr 2026 16:52:57 +0530 Subject: [PATCH 2/7] fix(croissant): use official Croissant 1.0 JSON-LD context - Replace inline @context dicts with CROISSANT_CONTEXT constant - Change @type from cr:Dataset to sc:Dataset per Croissant 1.0 spec - Add conformsTo and sha256 fields required by mlcroissant validation - Update test assertions to match corrected @type --- nobrainer/processing/croissant.py | 63 ++++++++++++++++++-- nobrainer/tests/unit/test_croissant.py | 4 +- nobrainer/tests/unit/test_dataset_builder.py | 2 +- 3 files changed, 60 insertions(+), 9 deletions(-) diff --git a/nobrainer/processing/croissant.py b/nobrainer/processing/croissant.py index 8b337b74..3485ce17 100644 --- a/nobrainer/processing/croissant.py +++ b/nobrainer/processing/croissant.py @@ -8,6 +8,44 @@ from pathlib import Path from typing import Any +CROISSANT_CONTEXT = { + "@language": "en", + "@vocab": "https://schema.org/", + "citeAs": "cr:citeAs", + "column": "cr:column", + "conformsTo": "dct:conformsTo", + "cr": "http://mlcommons.org/croissant/", + "rai": "http://mlcommons.org/croissant/RAI/", + "data": {"@id": "cr:data", "@type": "@json"}, + "dataType": {"@id": "cr:dataType", "@type": "@vocab"}, + "dct": "http://purl.org/dc/terms/", + "examples": {"@id": "cr:examples", "@type": "@json"}, + "extract": "cr:extract", + "field": "cr:field", + "fileProperty": "cr:fileProperty", + "fileObject": "cr:fileObject", + "fileSet": "cr:fileSet", + "format": "cr:format", + "includes": "cr:includes", + "isLiveDataset": "cr:isLiveDataset", + "jsonPath": "cr:jsonPath", + "key": "cr:key", + "md5": "cr:md5", + "parentField": "cr:parentField", + "path": "cr:path", + "recordSet": "cr:recordSet", + "references": "cr:references", + "regex": "cr:regex", + "repeated": "cr:repeated", + "replace": "cr:replace", + "samplingRate": "cr:samplingRate", + "sc": "https://schema.org/", + "separator": "cr:separator", + "source": "cr:source", + "subField": "cr:subField", + "transform": "cr:transform", +} + def _sha256(path: str | Path) -> str: """Compute SHA-256 hex digest of a file.""" @@ -53,8 +91,9 @@ def write_model_croissant( loss_name = getattr(estimator, "_loss_name", "unknown") metadata = { - "@context": {"@vocab": "http://mlcommons.org/croissant/"}, - "@type": "cr:Dataset", + "@context": CROISSANT_CONTEXT, + "@type": "sc:Dataset", + "conformsTo": "http://mlcommons.org/croissant/1.0", "name": f"nobrainer-{getattr(estimator, 'base_model', 'model')}", "description": ( f"Trained {getattr(estimator, 'base_model', 'model')} model " @@ -66,6 +105,11 @@ def write_model_croissant( "name": "model.pth", "contentUrl": "model.pth", "encodingFormat": "application/x-pytorch", + "sha256": ( + _sha256(save_dir / "model.pth") + if (save_dir / "model.pth").exists() + else "" + ), } ], "nobrainer:provenance": { @@ -123,8 +167,9 @@ def write_checkpoint_croissant( checkpoint_dir = Path(checkpoint_dir) metadata = { - "@context": {"@vocab": "http://mlcommons.org/croissant/"}, - "@type": "cr:Dataset", + "@context": CROISSANT_CONTEXT, + "@type": "sc:Dataset", + "conformsTo": "http://mlcommons.org/croissant/1.0", "name": f"nobrainer-{type(model).__name__}", "description": f"Trained {type(model).__name__} checkpoint via nobrainer", "distribution": [ @@ -133,6 +178,11 @@ def write_checkpoint_croissant( "name": "best_model.pth", "contentUrl": "best_model.pth", "encodingFormat": "application/x-pytorch", + "sha256": ( + _sha256(checkpoint_dir / "best_model.pth") + if (checkpoint_dir / "best_model.pth").exists() + else "" + ), } ], "nobrainer:provenance": { @@ -172,8 +222,9 @@ def write_dataset_croissant( ) -> Path: """Write Croissant-ML JSON-LD for a Dataset.""" metadata = { - "@context": {"@vocab": "http://mlcommons.org/croissant/"}, - "@type": "cr:Dataset", + "@context": CROISSANT_CONTEXT, + "@type": "sc:Dataset", + "conformsTo": "http://mlcommons.org/croissant/1.0", "name": "nobrainer-dataset", "description": "Brain MRI dataset for nobrainer", "distribution": [], diff --git a/nobrainer/tests/unit/test_croissant.py b/nobrainer/tests/unit/test_croissant.py index 9dc4b7bf..a8fdc507 100644 --- a/nobrainer/tests/unit/test_croissant.py +++ b/nobrainer/tests/unit/test_croissant.py @@ -81,7 +81,7 @@ def test_creates_valid_jsonld(self, tmp_path): data = json.loads(out.read_text()) assert "@context" in data assert "@type" in data - assert data["@type"] == "cr:Dataset" + assert data["@type"] == "sc:Dataset" def test_required_provenance_fields(self, tmp_path): """Provenance must contain all required fields.""" @@ -164,7 +164,7 @@ def test_writes_dataset_metadata(self, tmp_path): data = json.loads(out.read_text()) assert "@context" in data assert "@type" in data - assert data["@type"] == "cr:Dataset" + assert data["@type"] == "sc:Dataset" def test_dataset_info_present(self, tmp_path): ds = _make_fake_dataset(tmp_path) diff --git a/nobrainer/tests/unit/test_dataset_builder.py b/nobrainer/tests/unit/test_dataset_builder.py index 1e4e2ce1..b148fc13 100644 --- a/nobrainer/tests/unit/test_dataset_builder.py +++ b/nobrainer/tests/unit/test_dataset_builder.py @@ -159,7 +159,7 @@ def test_writes_valid_jsonld(self, tmp_path): data = json.loads(out.read_text()) assert "@context" in data assert "@type" in data - assert data["@type"] == "cr:Dataset" + assert data["@type"] == "sc:Dataset" def test_has_dataset_info(self, tmp_path): pairs = _make_file_pairs(2, (16, 16, 16), tmp_path) From e8281aa8662ba6e1b5f9b70ef3f68b84ca5c55e1 Mon Sep 17 00:00:00 2001 From: Dhritiman Das <14159298+dhritimandas@users.noreply.github.com> Date: Fri, 10 Apr 2026 16:53:06 +0530 Subject: [PATCH 3/7] fix(tests): disable MPS when Conv3D unsupported --- conftest.py | 22 ++++++++++++++++++++++ 1 file changed, 22 insertions(+) diff --git a/conftest.py b/conftest.py index 1471a88e..69577dee 100644 --- a/conftest.py +++ b/conftest.py @@ -6,6 +6,28 @@ import torch +def _mps_supports_conv3d() -> bool: + """Return True if Conv3D works on the MPS backend.""" + if not (hasattr(torch.backends, "mps") and torch.backends.mps.is_available()): + return False + try: + m = torch.nn.Conv3d(1, 1, 1).to("mps") + x = torch.randn(1, 1, 2, 2, 2, device="mps") + m(x) + return True + except RuntimeError: + return False + + +# Disable MPS auto-detection when Conv3D is unsupported (PyTorch < 2.3 on +# Apple Silicon). Without this, ``nobrainer.gpu.get_device()`` returns MPS +# and every 3D convolution raises ``RuntimeError: Conv3D is not supported on +# MPS``. +if hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): + if not _mps_supports_conv3d(): + torch.backends.mps.is_available = lambda: False # type: ignore[assignment] + + def pytest_collection_modifyitems(config, items): """Skip tests marked with @pytest.mark.gpu when CUDA is not available.""" if torch.cuda.is_available(): From f5b7580567d71288c6b1b45fea20be4378860d1b Mon Sep 17 00:00:00 2001 From: Dhritiman Das <14159298+dhritimandas@users.noreply.github.com> Date: Fri, 10 Apr 2026 16:53:21 +0530 Subject: [PATCH 4/7] docs: add development guidelines to CLAUDE.md --- CLAUDE.md | 30 ++++++++++++++++++++++++++++++ 1 file changed, 30 insertions(+) diff --git a/CLAUDE.md b/CLAUDE.md index ed09251c..b616c0ba 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -128,3 +128,33 @@ Every plan MUST include a Constitution Check evaluated before research and after ``` If speckit is not installed, follow the principles above manually. + +## Dependency constraint + +Do not add external dependencies unless they are already declared in `setup.cfg` or `pyproject.toml`. Current core deps include: torch, monai, nibabel, numpy, click, zarr, fsspec, joblib, scikit-image, psutil. If you think a new dep is justified, flag it explicitly — do not silently add it. + +## MONAI-first rule + +Before implementing any data loading, transform, loss, metric, or inference utility, check whether MONAI already provides it. Prefer wrapping MONAI over reimplementing. Check `monai.transforms`, `monai.data`, `monai.inferers`, `monai.losses`, `monai.metrics` before writing new code. If MONAI's version is insufficient, document why in a code comment. + +## Read before write + +Before modifying any existing module: +1. Read the module completely. List its implicit assumptions. +2. Check what MONAI provides that overlaps. +3. Do not guess function signatures — read the source. +4. If the module has existing tests, read those too to understand expected behavior. + +## Testing philosophy + +Every test should fail if the feature is removed. If a test passes regardless of whether your code exists, it's not testing your code. + +## Review process + +A reviewer subagent runs automatically (via stop hook) when you finish a task. It executes tests, inspects the diff, and writes a verdict to `.claude/reviews/`. If the reviewer blocks: +1. Read the review file it points to. +2. Fix every CRITICAL finding. +3. Address WARNING findings where reasonable. +4. The hook will re-run on your next stop. + +Do not commit until the reviewer passes. From fd69f6bd6d90b6c65605c2859a92c5a2a12ee4ba Mon Sep 17 00:00:00 2001 From: Dhritiman Das <14159298+dhritimandas@users.noreply.github.com> Date: Fri, 10 Apr 2026 16:53:39 +0530 Subject: [PATCH 5/7] chore: update .gitignore --- .gitignore | 3 +++ 1 file changed, 3 insertions(+) diff --git a/.gitignore b/.gitignore index 2b6b30ed..dc80d69c 100644 --- a/.gitignore +++ b/.gitignore @@ -151,3 +151,6 @@ data/ # Model artifacts *.pth brain_mask_extraction_model/ + +# Claude Code review artifacts +.claude/REVIEW.MD/ From 1828fe25388c5b624d3f2f134a62ac19810ddfb2 Mon Sep 17 00:00:00 2001 From: Dhritiman Das <14159298+dhritimandas@users.noreply.github.com> Date: Sun, 23 Aug 2026 18:42:19 +0530 Subject: [PATCH 6/7] feat(provenance): add PROV-O RDF export for BrainKB ingestion No structured provenance graph existed for a training run -- croissant.json records model/dataset facts but no prov:Activity linking them. This adds `nobrainer provenance export`, reading a saved model bundle (and optionally a DataSpec manifest for image+label dataset provenance) into a PROV-O graph: the run as prov:Activity, used dataset entity, generated model entity, and associated software/caller agents. Every IRI is content-derived (sha256 for files/models, stable digests for runs/datasets/agents) -- no blank nodes, no timestamps or UUIDs in identity, so re-export of the same input is byte-identical regardless of its path on disk. Emits DOMAIN provenance only: named-graph registration and POSTing to BrainKB's ingestion API are separate steps against BrainKB's own API, enforced at runtime (not just documented) by rejecting a second Activity, any non-SoftwareAgent floating free of the run, or any brainkb-named term. rdflib is the one new dependency, gated behind a new [provenance] extra (it was already present transitively via mlcroissant, but undeclared). --- nobrainer/cli/main.py | 81 ++ nobrainer/provenance/__init__.py | 19 + nobrainer/provenance/rdf_export.py | 1004 +++++++++++++++++++++++ nobrainer/tests/unit/test_rdf_export.py | 653 +++++++++++++++ pyproject.toml | 3 +- 5 files changed, 1759 insertions(+), 1 deletion(-) create mode 100644 nobrainer/provenance/__init__.py create mode 100644 nobrainer/provenance/rdf_export.py create mode 100644 nobrainer/tests/unit/test_rdf_export.py diff --git a/nobrainer/cli/main.py b/nobrainer/cli/main.py index 2f087989..485900ef 100644 --- a/nobrainer/cli/main.py +++ b/nobrainer/cli/main.py @@ -846,6 +846,87 @@ def zarr_suggest_shards(n_volumes, volume_shape, dtype, n_input_files, levels): click.echo(json.dumps(result, indent=2)) +# --------------------------------------------------------------------------- +# provenance subcommands +# --------------------------------------------------------------------------- + + +@cli.group() +def provenance(): + """Emit training-run provenance as PROV-O RDF.""" + + +@provenance.command("export") +@click.option( + "--bundle", + required=True, + type=click.Path(exists=True, file_okay=False), + help="Segmentation.save() directory (model.pth + croissant.json).", + **_option_kwds, +) +@click.option( + "--dataspec", + default=None, + type=click.Path(exists=True, dir_okay=False), + help="Optional DataSpec manifest JSON for image+label dataset provenance.", + **_option_kwds, +) +@click.option( + "--out", + default=None, + type=click.Path(), + help="Output file path. Omit to write to stdout.", +) +@click.option( + "--format", + "fmt", + type=click.Choice(["turtle", "json-ld"]), + default="turtle", + help="Serialization format.", + **_option_kwds, +) +@click.option( + "--base-iri", + default="https://neuronets.dev/nobrainer/", + help="Base IRI for minted instance identifiers.", + **_option_kwds, +) +@click.option( + "--agent", + default=None, + help="Optional caller-supplied human/organization agent identity " + "('Name ' or an IRI).", +) +@click.option( + "--strict", + is_flag=True, + help="Fail instead of eliding unrepresentable values.", +) +def provenance_export(*, bundle, dataspec, out, fmt, base_iri, agent, strict): + """Export PROV-O RDF provenance for a saved nobrainer training run.""" + from ..provenance import ProvenanceError, export_provenance + + try: + text = export_provenance( + bundle, + dataspec_path=dataspec, + fmt=fmt, + base_iri=base_iri, + agent=agent, + strict=strict, + ) + except ProvenanceError as exc: + click.echo(click.style(f"ERROR: {exc}", fg="red")) + raise SystemExit(1) from exc + + if out: + with open(out, "w") as fh: + fh.write(text) + click.echo(f"Wrote {fmt} provenance to {out}") + else: + click.echo(text) + + # For debugging only. if __name__ == "__main__": cli() diff --git a/nobrainer/provenance/__init__.py b/nobrainer/provenance/__init__.py new file mode 100644 index 00000000..abaae909 --- /dev/null +++ b/nobrainer/provenance/__init__.py @@ -0,0 +1,19 @@ +"""PROV-O RDF provenance export for nobrainer training runs.""" + +from __future__ import annotations + +from .rdf_export import ( + ProvenanceError, + build_graph, + export_provenance, + to_jsonld, + to_turtle, +) + +__all__ = [ + "ProvenanceError", + "build_graph", + "export_provenance", + "to_jsonld", + "to_turtle", +] diff --git a/nobrainer/provenance/rdf_export.py b/nobrainer/provenance/rdf_export.py new file mode 100644 index 00000000..04106a3d --- /dev/null +++ b/nobrainer/provenance/rdf_export.py @@ -0,0 +1,1004 @@ +"""PROV-O RDF provenance export for nobrainer training runs. + +Reads a saved model bundle (a ``Segmentation.save()`` directory: ``model.pth`` ++ ``croissant.json``) and, optionally, a ``DataSpec`` dataset manifest +(:mod:`nobrainer.data.spec`), and emits a PROV-O graph describing the +training run as a ``prov:Activity`` that ``prov:used`` a dataset entity and +``prov:generated`` a model entity, both ``prov:wasAssociatedWith`` one or +more agents. + +Scope: DOMAIN provenance only +------------------------------ +This module emits facts about what nobrainer itself did -- the run, its +inputs, its outputs, its agents. It does **not** emit BrainKB +ingestion-activity triples (for example, an activity describing "this +graph was loaded into BrainKB at time T"). Named-graph registration and +POSTing this module's output to BrainKB's ingestion API are separate steps +performed by the caller against BrainKB's own API, entirely outside this +module's scope. ``_assert_domain_only`` enforces this boundary on every +call to :func:`build_graph`. + +Determinism +----------- +Every IRI minted by this module is derived from content -- a sha256 digest +of a file, or a stable hash of a canonical-JSON payload -- never from a +timestamp, a random UUID, or an absolute filesystem path. Re-exporting the +same ``--bundle``/``--dataspec`` input from a different location on disk +produces a byte-identical graph. There are no blank nodes anywhere in the +emitted graph. +""" + +from __future__ import annotations + +import hashlib +import json +from pathlib import Path +import re +from typing import Any + +from rdflib import Graph, Literal, Namespace, URIRef +from rdflib.namespace import DCTERMS, PROV, RDF, RDFS, XSD + +import nobrainer +from nobrainer.data.spec import DataSpec, FileStatus, check_file_presence + +__all__ = [ + "ProvenanceError", + "build_graph", + "export_provenance", + "to_jsonld", + "to_turtle", +] + +SCHEMA = Namespace("https://schema.org/") + +DEFAULT_BASE_IRI = "https://neuronets.dev/nobrainer/" +DEFAULT_VOCAB_IRI = "https://neuronets.dev/ns/nobrainer#" + +# NOTE (needs owner sign-off before this is used as a permanent identifier): +# neuronets.dev currently serves the nobrainer book, not a term dereferencer. +# Both constants above are overridable via --base-iri specifically so this +# is a one-line change, not a re-mint of every IRI this module has produced. + +_MEMORY_ADDRESS_RE = re.compile(r"<[^>]* at 0x[0-9a-fA-F]+>") +_GIT_SHA_RE = re.compile(r"\+.*?g([0-9a-fA-F]{7,40})") +_EMAIL_RE = re.compile(r"[^<\s]+@[^>\s]+") + +# The PROV-O predicates/types this module allows itself to emit. Adding a +# predicate here is the reviewable seam that keeps the DOMAIN-only boundary +# from drifting silently -- see _assert_domain_only. +_FORBIDDEN_PREDICATES = frozenset( + { + PROV.wasInformedBy, + PROV.wasStartedBy, + PROV.wasEndedBy, + PROV.qualifiedAssociation, + PROV.hadPlan, + } +) + + +class ProvenanceError(RuntimeError): + """Raised when a bundle or dataspec cannot be turned into a provenance graph.""" + + +# --------------------------------------------------------------------------- +# Hashing / canonicalization -- the basis of every minted IRI +# --------------------------------------------------------------------------- + + +def canonical_json_bytes(obj: Any) -> bytes: + """Serialize an object to canonical, sorted, compact UTF-8 JSON bytes. + + Parameters + ---------- + obj : Any + A JSON-serializable object (dict, list, str, int, float, bool, None). + + Returns + ------- + bytes + UTF-8 encoded JSON with sorted keys and no incidental whitespace, so + that the same logical payload always produces the same bytes + regardless of dict insertion order. + + Raises + ------ + ValueError + If ``obj`` contains a non-finite float (``allow_nan=False``): a NaN + or Infinity must never silently enter a content hash. + """ + return json.dumps( + obj, + sort_keys=True, + separators=(",", ":"), + ensure_ascii=False, + allow_nan=False, + ).encode("utf-8") + + +def digest(kind: str, payload: Any, length: int = 32) -> str: + """Domain-separated, content-derived hex digest of a JSON payload. + + Parameters + ---------- + kind : str + A short tag identifying what is being hashed (e.g. ``"run"``, + ``"dataset"``). Mixed into the hash ahead of an ASCII Unit + Separator byte (``\\x1f``, which cannot occur in the JSON output), + so a run and a dataset built from identical payloads never collide. + payload : Any + JSON-serializable payload to hash. + length : int, optional + Number of hex characters to keep from the SHA-256 digest. The + default, 32 hex characters (128 bits), is birthday-safe for any + plausible corpus size and should not be lowered. + + Returns + ------- + str + A stable, deterministic hex digest. + """ + h = hashlib.sha256() + h.update(kind.encode("utf-8")) + h.update(b"\x1f") + h.update(canonical_json_bytes(payload)) + return h.hexdigest()[:length] + + +def sha256_file(path: str | Path) -> str: + """Stream a file in 64 KiB chunks and return its hex SHA-256 digest. + + Parameters + ---------- + path : str or Path + Path to an existing, readable file. + + Returns + ------- + str + Lowercase hex SHA-256 digest. + """ + h = hashlib.sha256() + with open(path, "rb") as fh: + for chunk in iter(lambda: fh.read(1 << 16), b""): + h.update(chunk) + return h.hexdigest() + + +def normalize_float(value: float | None) -> str | None: + """Convert a float to a stable string form for hashing, rejecting non-finite. + + Parameters + ---------- + value : float or None + A value read from run metadata (e.g. a training loss). + + Returns + ------- + str or None + ``None`` if ``value`` is ``None``; otherwise ``repr(float(value))``, + a string, so the exact decimal representation is pinned independent + of any future change to Python's own float-repr algorithm. + + Raises + ------ + ProvenanceError + If ``value`` is NaN or +/-Infinity. rdflib serializes a non-finite + float literal to bare ``NaN``/``Infinity`` in JSON-LD, which is not + valid JSON (verified) -- such a value must never reach a hash or a + literal, silently or otherwise. + """ + if value is None: + return None + fv = float(value) + if fv != fv or fv in (float("inf"), float("-inf")): + raise ProvenanceError(f"non-finite float cannot be normalized: {value!r}") + return repr(fv) + + +def _elide_memory_addresses(value: Any) -> tuple[Any, bool]: + """Detect and elide a Python ``repr(obj)`` memory-address string. + + ``write_model_croissant`` serializes with ``json.dumps(..., default=str)``, + so a value without a stable ``__str__`` (a callable, an unpickleable + object) can land in ``model_args``/``optimizer.args`` as + ``""`` -- a string that differs on every + process. If such a value entered a content hash, run IRIs would stop + being deterministic across machines. + + Parameters + ---------- + value : Any + A value read from croissant.json's hyperparameter dicts. + + Returns + ------- + tuple[Any, bool] + ``(value, False)`` unchanged, or ``(None, True)`` if ``value`` (or + any string it contains) matched the memory-address pattern. + """ + if isinstance(value, str) and _MEMORY_ADDRESS_RE.search(value): + return None, True + return value, False + + +def _parse_git_sha(version: str) -> str | None: + """Extract a short git commit SHA from a hatch-vcs local version segment. + + Parameters + ---------- + version : str + A version string such as ``"2.0.0a17.dev6+gb85a1ca5b"``. + + Returns + ------- + str or None + The commit SHA (e.g. ``"b85a1ca5b"``), or ``None`` if the version + string has no ``+g`` local segment. + """ + match = _GIT_SHA_RE.search(version) + return match.group(1) if match else None + + +# --------------------------------------------------------------------------- +# IRI minting -- every subject in the emitted graph is one of these +# --------------------------------------------------------------------------- + + +def _content_iri( + base_iri: str, kind: str, sha256: str | None, nosha_payload: dict +) -> URIRef: + """Mint a content-addressed IRI, or a visibly-unaddressed fallback. + + Parameters + ---------- + base_iri : str + The export's base IRI (trailing slash expected). + kind : str + ``"model"`` or ``"file"`` -- the path segment used for both the + addressed and unaddressed forms. + sha256 : str or None + The file's SHA-256 digest, if known. + nosha_payload : dict + A JSON payload identifying the entity when no checksum is + available (e.g. its recorded path). Hashed under ``"{kind}-nosha"`` + so it cannot collide with a real checksum. + + Returns + ------- + URIRef + ``{base}{kind}/sha256/{sha256}`` if a checksum is known, else + ``{base}{kind}/x-nosha/{digest}`` -- the ``x-nosha`` segment is + deliberately visible so any consumer can tell at a glance that the + entity is not content-addressed. + """ + if sha256: + return URIRef(f"{base_iri}{kind}/sha256/{sha256}") + return URIRef(f"{base_iri}{kind}/x-nosha/{digest(f'{kind}-nosha', nosha_payload)}") + + +def _dataset_iri(base_iri: str, member_iris: list[URIRef]) -> URIRef: + """Mint a dataset-collection IRI from the sorted set of its member IRIs.""" + payload = {"members": sorted(str(m) for m in member_iris)} + return URIRef(f"{base_iri}dataset/{digest('dataset', payload)}") + + +def _run_iri(base_iri: str, payload: dict) -> URIRef: + """Mint a training-run IRI from its canonical identity payload.""" + return URIRef(f"{base_iri}run/{digest('run', payload)}") + + +def _software_agent_iri(base_iri: str, name: str, version: str) -> URIRef: + """Mint a software-agent IRI from its name and version.""" + payload = {"name": name, "version": version} + return URIRef(f"{base_iri}agent/software/{digest('agent-software', payload)}") + + +def _human_agent_iri(base_iri: str, agent_text: str) -> URIRef: + """Mint a caller-supplied agent IRI from the raw --agent text.""" + return URIRef( + f"{base_iri}agent/caller/{digest('agent-caller', {'text': agent_text})}" + ) + + +def _hparam_iri( + base_iri: str, run_iri: URIRef, scope: str, name: str, value_payload: Any +) -> URIRef: + """Mint a hyperparameter-node IRI scoped to its owning run.""" + payload = {"scope": scope, "name": name, "value": value_payload} + return URIRef(f"{run_iri}/hparam/{scope}/{digest('hparam', payload, length=16)}") + + +# --------------------------------------------------------------------------- +# Reading the bundle / dataspec inputs +# --------------------------------------------------------------------------- + + +def _read_bundle(bundle_dir: Path) -> dict[str, Any]: + """Read and validate a ``Segmentation.save()`` directory's croissant.json. + + Reads ``croissant.json`` with plain :func:`json.load` -- deliberately + **not** via :mod:`nobrainer.processing.croissant`, because importing + that module executes ``nobrainer/processing/__init__.py``, which pulls + in torch (verified). This module stays stdlib + rdflib only. + + Parameters + ---------- + bundle_dir : Path + Directory expected to contain ``croissant.json`` (and normally + ``model.pth``, though this function does not require the weights + file to exist). + + Returns + ------- + dict[str, Any] + The full parsed croissant.json document. + + Raises + ------ + ProvenanceError + If ``croissant.json`` is missing, malformed, or describes a + dataset (``nobrainer:dataset_info``) rather than a training run + (``nobrainer:provenance``). + """ + croissant_path = bundle_dir / "croissant.json" + if not croissant_path.is_file(): + raise ProvenanceError(f"No croissant.json found under {bundle_dir}") + try: + doc = json.loads(croissant_path.read_text()) + except json.JSONDecodeError as exc: + raise ProvenanceError( + f"Malformed croissant.json at {croissant_path}: {exc}" + ) from exc + + if "nobrainer:provenance" not in doc: + if "nobrainer:dataset_info" in doc: + raise ProvenanceError( + f"{croissant_path} describes a dataset (nobrainer:dataset_info), " + "not a training run. --bundle must point at a Segmentation.save() " + "directory, whose croissant.json carries nobrainer:provenance." + ) + raise ProvenanceError(f"{croissant_path} has no 'nobrainer:provenance' key.") + return doc + + +def _model_distribution(doc: dict[str, Any]) -> dict[str, Any]: + """Return the first distribution entry (the model weights file).""" + dist = doc.get("distribution") or [] + if not dist: + raise ProvenanceError("croissant.json has an empty 'distribution' list.") + return dist[0] + + +def _resolve_model_sha256( + bundle_dir: Path, distribution: dict[str, Any] +) -> tuple[str | None, str]: + """Resolve a model checksum via the recorded-value / recompute / unavailable ladder. + + Parameters + ---------- + bundle_dir : Path + The bundle directory (used to locate the weights file for + recomputation). + distribution : dict[str, Any] + The ``distribution[0]`` entry from croissant.json. + + Returns + ------- + tuple[str or None, str] + ``(sha256, status)`` where ``status`` is one of ``"present"`` + (recorded in croissant.json), ``"recomputed"`` (recorded value was + empty but the file exists on disk), or ``"unavailable"`` (neither). + """ + recorded = distribution.get("sha256") or "" + if recorded: + return recorded, "present" + content_url = distribution.get("contentUrl", "model.pth") + candidate = bundle_dir / content_url + if candidate.is_file(): + return sha256_file(candidate), "recomputed" + return None, "unavailable" + + +def _read_dataspec(path: Path) -> DataSpec: + """Load a DataSpec manifest, raising ProvenanceError on failure. + + Parameters + ---------- + path : Path + Path to a DataSpec JSON manifest (:meth:`DataSpec.to_json` format). + + Returns + ------- + DataSpec + The parsed dataset specification. + + Raises + ------ + ProvenanceError + If the file is missing or malformed. + """ + try: + return DataSpec.from_json(path) + except (OSError, json.JSONDecodeError) as exc: + raise ProvenanceError( + f"Could not read DataSpec manifest {path}: {exc}" + ) from exc + + +# --------------------------------------------------------------------------- +# Graph construction +# --------------------------------------------------------------------------- + + +def build_graph( + bundle_dir: str | Path, + dataspec_path: str | Path | None = None, + base_iri: str = DEFAULT_BASE_IRI, + agent: str | None = None, + strict: bool = False, +) -> Graph: + """Build a DOMAIN-only PROV-O graph for one saved nobrainer training run. + + Parameters + ---------- + bundle_dir : str or Path + A ``Segmentation.save()`` directory (``model.pth`` + ``croissant.json``). + dataspec_path : str, Path, or None, optional + Optional ``DataSpec`` manifest JSON. When given, dataset provenance + is built from its ``entries`` (both image and label paths, each + checksummed here) instead of croissant's ``source_datasets``, which + only ever records image paths. When omitted, falls back to + croissant's ``source_datasets``. + base_iri : str, optional + Base IRI for minted instance identifiers. Default + ``"https://neuronets.dev/nobrainer/"``. + agent : str or None, optional + An optional caller-supplied human or organization identity (a free + text ``"Name "`` string, or an IRI). If given, adds a + ``prov:Person`` (free text) or a plain ``prov:Agent`` (an IRI -- + this module cannot tell whether an opaque IRI denotes a person or + an organization) ``prov:wasAssociatedWith`` the run, **in addition + to** the always-emitted software agents. This is caller-asserted, + never inferred: nothing in the underlying data identifies a human, + so this module never invents one on its own. + strict : bool, optional + If ``True``, raise :class:`ProvenanceError` on conditions that + would otherwise be silently elided (an unavailable checksum, a + non-finite loss value, a memory-address-poisoned hyperparameter). + Default ``False``. + + Returns + ------- + rdflib.Graph + A graph containing exactly one ``prov:Activity`` (the run), a + ``prov:Entity`` for the model it ``prov:generated``, a + ``prov:Entity``/``prov:Collection`` for the dataset it + ``prov:used`` (omitted if there are zero dataset members), and one + or more ``prov:Agent`` nodes it is ``prov:wasAssociatedWith``. No + blank nodes. Verified DOMAIN-only via ``_assert_domain_only`` + before being returned. + """ + bundle_dir = Path(bundle_dir) + doc = _read_bundle(bundle_dir) + prov_meta = doc["nobrainer:provenance"] + distribution = _model_distribution(doc) + + g = Graph() + g.bind("prov", PROV) + g.bind("dcterms", DCTERMS) + g.bind("rdfs", RDFS) + g.bind("schema", SCHEMA) + g.bind("xsd", XSD) + nb = Namespace(DEFAULT_VOCAB_IRI) + g.bind("nb", nb) + + # -- Dataset members ----------------------------------------------------- + dataset_members: list[URIRef] = [] + checksum_coverage = "none" + if dataspec_path is not None: + spec = _read_dataspec(Path(dataspec_path)) + checksum_coverage = "images-and-labels" + for entry in spec.entries: + for key in ("image", "label"): + if key not in entry: + continue + file_path = entry[key] + status = check_file_presence(file_path) + sha = sha256_file(file_path) if status == FileStatus.PRESENT else None + fiu = _content_iri(base_iri, "file", sha, {"path": file_path}) + g.add((fiu, RDF.type, PROV.Entity)) + g.add((fiu, RDF.type, nb.SourceFile)) + g.add((fiu, nb.sourcePath, Literal(file_path))) + g.add((fiu, nb.sourceRole, Literal(key))) + if sha: + g.add((fiu, nb.sha256, Literal(sha, datatype=XSD.hexBinary))) + g.add((fiu, nb.checksumStatus, Literal("present"))) + else: + g.add((fiu, nb.checksumStatus, Literal("unavailable"))) + if strict: + raise ProvenanceError( + f"Cannot checksum {key} file: {file_path}" + ) + dataset_members.append(fiu) + else: + source_datasets = prov_meta.get("source_datasets") or [] + if source_datasets: + checksum_coverage = "images-only" + for item in source_datasets: + path = item.get("path", "") + sha = item.get("sha256") or None + fiu = _content_iri(base_iri, "file", sha, {"path": path}) + g.add((fiu, RDF.type, PROV.Entity)) + g.add((fiu, RDF.type, nb.SourceFile)) + g.add((fiu, nb.sourcePath, Literal(path))) + g.add((fiu, nb.sourceRole, Literal("image"))) + if sha: + g.add((fiu, nb.sha256, Literal(sha, datatype=XSD.hexBinary))) + g.add((fiu, nb.checksumStatus, Literal("present"))) + else: + g.add((fiu, nb.checksumStatus, Literal("unavailable"))) + dataset_members.append(fiu) + + dataset_iri: URIRef | None = None + if dataset_members: + dataset_iri = _dataset_iri(base_iri, dataset_members) + g.add((dataset_iri, RDF.type, PROV.Entity)) + g.add((dataset_iri, RDF.type, PROV.Collection)) + g.add((dataset_iri, RDF.type, nb.TrainingDataset)) + g.add( + ( + dataset_iri, + nb.memberCount, + Literal(len(dataset_members), datatype=XSD.nonNegativeInteger), + ) + ) + g.add((dataset_iri, nb.checksumCoverage, Literal(checksum_coverage))) + for m in dataset_members: + g.add((dataset_iri, PROV.hadMember, m)) + + # -- Model entity --------------------------------------------------------- + model_sha, checksum_status = _resolve_model_sha256(bundle_dir, distribution) + if checksum_status == "unavailable" and strict: + raise ProvenanceError( + "Model weights checksum is unavailable and --strict was set." + ) + model_iri = _content_iri( + base_iri, + "model", + model_sha, + {"contentUrl": distribution.get("contentUrl", "model.pth")}, + ) + g.add((model_iri, RDF.type, PROV.Entity)) + g.add((model_iri, RDF.type, nb.TrainedModel)) + g.add((model_iri, SCHEMA.name, Literal(distribution.get("name", "model.pth")))) + g.add( + ( + model_iri, + nb.relativePath, + Literal(distribution.get("contentUrl", "model.pth")), + ) + ) + if distribution.get("encodingFormat"): + g.add( + (model_iri, SCHEMA.encodingFormat, Literal(distribution["encodingFormat"])) + ) + if model_sha: + g.add((model_iri, nb.sha256, Literal(model_sha, datatype=XSD.hexBinary))) + g.add((model_iri, nb.checksumStatus, Literal(checksum_status))) + training_date = prov_meta.get("training_date") + if training_date: + g.add( + ( + model_iri, + PROV.generatedAtTime, + Literal(training_date, datatype=XSD.dateTime), + ) + ) + g.add( + (model_iri, DCTERMS.created, Literal(training_date, datatype=XSD.dateTime)) + ) + + # -- Architecture vocabulary discrimination ------------------------------- + architecture = prov_meta.get("model_architecture") + is_model_flavor = any( + prov_meta.get(k) + for k in ("source_datasets", "model_args", "n_classes", "block_shape") + ) + architecture_vocab = "unknown" + if architecture: + architecture_vocab = ( + "nobrainer-model-registry" if is_model_flavor else "torch-class-name" + ) + + # -- Hyperparameters ------------------------------------------------------- + hparam_iris: list[URIRef] = [] + + def _add_hparams( + run_iri_placeholder: URIRef, scope: str, values: dict[str, Any] + ) -> None: + for name, value in (values or {}).items(): + clean, elided = _elide_memory_addresses(value) + if elided: + if strict: + raise ProvenanceError( + f"Hyperparameter {scope}.{name} is memory-address-poisoned." + ) + hiu = _hparam_iri(base_iri, run_iri_placeholder, scope, name, None) + g.add((hiu, RDF.type, nb.Hyperparameter)) + g.add((hiu, SCHEMA.name, Literal(name))) + g.add((hiu, nb.hyperparameterScope, Literal(scope))) + g.add( + (hiu, nb.hyperparameterElided, Literal(True, datatype=XSD.boolean)) + ) + hparam_iris.append(hiu) + continue + hiu = _hparam_iri(base_iri, run_iri_placeholder, scope, name, clean) + g.add((hiu, RDF.type, nb.Hyperparameter)) + g.add((hiu, SCHEMA.name, Literal(name))) + g.add((hiu, nb.hyperparameterScope, Literal(scope))) + if isinstance(clean, bool): + g.add((hiu, SCHEMA.value, Literal(clean, datatype=XSD.boolean))) + elif isinstance(clean, int): + g.add((hiu, SCHEMA.value, Literal(clean, datatype=XSD.integer))) + elif isinstance(clean, float): + g.add( + ( + hiu, + SCHEMA.value, + Literal(normalize_float(clean), datatype=XSD.double), + ) + ) + elif isinstance(clean, str): + g.add((hiu, SCHEMA.value, Literal(clean))) + else: + g.add((hiu, nb.valueJson, Literal(json.dumps(clean, sort_keys=True)))) + hparam_iris.append(hiu) + + # -- Run identity payload + IRI ------------------------------------------- + optimizer = prov_meta.get("optimizer") or {} + final_loss = prov_meta.get("final_loss") + best_loss = prov_meta.get("best_loss") + loss_status = "finite" + try: + final_loss_norm = normalize_float(final_loss) + best_loss_norm = normalize_float(best_loss) + except ProvenanceError: + if strict: + raise + loss_status = "non-finite" + final_loss_norm = None + best_loss_norm = None + + run_payload = { + "schema": 1, + "model": str(model_iri), + "dataset": str(dataset_iri) if dataset_iri else None, + "training_date": training_date, + "nobrainer_version": prov_meta.get("nobrainer_version"), + "pytorch_version": prov_meta.get("pytorch_version"), + "model_architecture": architecture, + "architecture_vocab": architecture_vocab, + "loss_function": prov_meta.get("loss_function"), + "optimizer": optimizer, + "epochs_trained": prov_meta.get("epochs_trained"), + "final_loss": final_loss_norm, + "best_loss": best_loss_norm, + "n_classes": prov_meta.get("n_classes"), + "block_shape": list(prov_meta.get("block_shape") or []), + "model_args": prov_meta.get("model_args") or {}, + "gpu_count": prov_meta.get("gpu_count"), + } + run_iri = _run_iri(base_iri, run_payload) + + _add_hparams(run_iri, "optimizer", optimizer.get("args") or {}) + _add_hparams(run_iri, "model", prov_meta.get("model_args") or {}) + + g.add((run_iri, RDF.type, PROV.Activity)) + g.add((run_iri, RDF.type, nb.TrainingRun)) + if doc.get("name"): + g.add((run_iri, RDFS.label, Literal(doc["name"]))) + if doc.get("description"): + g.add((run_iri, DCTERMS.description, Literal(doc["description"]))) + if doc.get("conformsTo"): + g.add( + ( + run_iri, + nb.croissantConformsTo, + Literal(doc["conformsTo"], datatype=XSD.anyURI), + ) + ) + g.add((run_iri, nb.provenanceScope, Literal("domain"))) + g.add( + ( + run_iri, + nb.provenanceSchemaVersion, + Literal(1, datatype=XSD.nonNegativeInteger), + ) + ) + g.add((run_iri, nb.runIdentifierSource, Literal("derived"))) + if dataset_iri is not None: + g.add((run_iri, PROV.used, dataset_iri)) + g.add((run_iri, PROV.generated, model_iri)) + g.add((model_iri, PROV.wasGeneratedBy, run_iri)) + if training_date: + g.add( + ( + run_iri, + nb.metadataRecordedAt, + Literal(training_date, datatype=XSD.dateTime), + ) + ) + if prov_meta.get("nobrainer_version"): + g.add((run_iri, nb.nobrainerVersion, Literal(prov_meta["nobrainer_version"]))) + git_sha = _parse_git_sha(prov_meta["nobrainer_version"]) + if git_sha: + g.add((run_iri, nb.gitCommitId, Literal(git_sha))) + if prov_meta.get("pytorch_version"): + g.add((run_iri, nb.pytorchVersion, Literal(prov_meta["pytorch_version"]))) + if optimizer.get("class"): + g.add((run_iri, nb.optimizerClass, Literal(optimizer["class"]))) + if prov_meta.get("loss_function"): + g.add((run_iri, nb.lossFunction, Literal(prov_meta["loss_function"]))) + if prov_meta.get("epochs_trained") is not None: + g.add( + ( + run_iri, + nb.epochsTrained, + Literal(prov_meta["epochs_trained"], datatype=XSD.nonNegativeInteger), + ) + ) + if final_loss_norm is not None: + g.add((run_iri, nb.finalLoss, Literal(float(final_loss), datatype=XSD.double))) + if best_loss_norm is not None: + g.add((run_iri, nb.bestLoss, Literal(float(best_loss), datatype=XSD.double))) + if loss_status == "non-finite": + g.add((run_iri, nb.lossStatus, Literal("non-finite"))) + if architecture: + g.add((run_iri, nb.modelArchitecture, Literal(architecture))) + g.add((run_iri, nb.architectureVocabulary, Literal(architecture_vocab))) + if architecture_vocab == "nobrainer-model-registry": + g.add( + (run_iri, nb.modelArchitectureNormalized, Literal(architecture.lower())) + ) + if prov_meta.get("n_classes") is not None: + g.add( + ( + run_iri, + nb.numberOfClasses, + Literal(prov_meta["n_classes"], datatype=XSD.nonNegativeInteger), + ) + ) + block_shape = prov_meta.get("block_shape") or [] + if block_shape: + g.add((run_iri, nb.blockShape, Literal(json.dumps(list(block_shape))))) + g.add( + ( + run_iri, + nb.blockShapeRank, + Literal(len(block_shape), datatype=XSD.nonNegativeInteger), + ) + ) + if prov_meta.get("gpu_count") is not None: + g.add( + ( + run_iri, + nb.gpuCount, + Literal(prov_meta["gpu_count"], datatype=XSD.nonNegativeInteger), + ) + ) + g.add((run_iri, nb.sourceChecksumCoverage, Literal(checksum_coverage))) + for hiu in hparam_iris: + g.add((run_iri, nb.hasHyperparameter, hiu)) + + # -- Agents ----------------------------------------------------------------- + nobrainer_version = prov_meta.get("nobrainer_version") or nobrainer.__version__ + nb_agent_iri = _software_agent_iri(base_iri, "nobrainer", nobrainer_version) + g.add((nb_agent_iri, RDF.type, PROV.Agent)) + g.add((nb_agent_iri, RDF.type, PROV.SoftwareAgent)) + g.add((nb_agent_iri, SCHEMA.name, Literal("nobrainer"))) + g.add((nb_agent_iri, SCHEMA.softwareVersion, Literal(nobrainer_version))) + g.add((nb_agent_iri, RDFS.label, Literal(f"nobrainer {nobrainer_version}"))) + g.add((run_iri, PROV.wasAssociatedWith, nb_agent_iri)) + + pytorch_version = prov_meta.get("pytorch_version") + if pytorch_version: + pt_agent_iri = _software_agent_iri(base_iri, "pytorch", pytorch_version) + g.add((pt_agent_iri, RDF.type, PROV.Agent)) + g.add((pt_agent_iri, RDF.type, PROV.SoftwareAgent)) + g.add((pt_agent_iri, SCHEMA.name, Literal("pytorch"))) + g.add((pt_agent_iri, SCHEMA.softwareVersion, Literal(pytorch_version))) + g.add((pt_agent_iri, RDFS.label, Literal(f"pytorch {pytorch_version}"))) + g.add((run_iri, PROV.wasAssociatedWith, pt_agent_iri)) + + if agent: + if agent.startswith("http://") or agent.startswith("https://"): + agent_iri = URIRef(agent) + g.add((agent_iri, RDF.type, PROV.Agent)) + else: + agent_iri = _human_agent_iri(base_iri, agent) + g.add((agent_iri, RDF.type, PROV.Agent)) + g.add((agent_iri, RDF.type, PROV.Person)) + g.add((agent_iri, RDFS.label, Literal(agent))) + email_match = _EMAIL_RE.search(agent) + if email_match: + g.add((agent_iri, SCHEMA.email, Literal(email_match.group(0)))) + g.add((run_iri, PROV.wasAssociatedWith, agent_iri)) + + _assert_domain_only(g) + return g + + +def _assert_domain_only(g: Graph) -> None: + """Verify a graph contains DOMAIN provenance only -- no BrainKB ingestion triples. + + Parameters + ---------- + g : rdflib.Graph + The graph to check. + + Raises + ------ + ProvenanceError + If any of the following hold: more than one ``prov:Activity`` + exists; any ``prov:Agent`` is not a ``prov:SoftwareAgent``, + ``prov:Person``, or plain ``prov:Agent`` associated with the run; + a forbidden PROV predicate (``wasInformedBy``, ``wasStartedBy``, + ``wasEndedBy``, ``qualifiedAssociation``, ``hadPlan``) is present; + any term (subject, predicate, or object) contains ``"brainkb"`` + case-insensitively; or a blank node is present anywhere. + """ + from rdflib import BNode + + activities = set(g.subjects(RDF.type, PROV.Activity)) + if len(activities) != 1: + raise ProvenanceError( + f"Expected exactly one prov:Activity (the training run); found {len(activities)}." + ) + + for s, p, o in g: + if isinstance(s, BNode) or isinstance(o, BNode): + raise ProvenanceError( + "Graph contains a blank node; all IRIs must be content-derived." + ) + if p in _FORBIDDEN_PREDICATES: + raise ProvenanceError(f"Forbidden predicate present: {p}") + for term in (s, p, o): + if "brainkb" in str(term).lower(): + raise ProvenanceError( + f"Term contains 'brainkb': {term!r}. This module emits DOMAIN " + "provenance only; BrainKB ingestion-activity triples are added " + "by the caller against BrainKB's own API, not by this module." + ) + + agents = set(g.subjects(RDF.type, PROV.Agent)) + for agent_iri in agents: + associated_with_run = any( + (run, PROV.wasAssociatedWith, agent_iri) in g for run in activities + ) + if not associated_with_run: + raise ProvenanceError( + f"Agent {agent_iri} is not associated with the training run." + ) + + +# --------------------------------------------------------------------------- +# Serialization +# --------------------------------------------------------------------------- + + +def to_turtle(g: Graph) -> str: + """Serialize a graph to Turtle. + + Parameters + ---------- + g : rdflib.Graph + The graph to serialize. + + Returns + ------- + str + Turtle-formatted text. + """ + return g.serialize(format="turtle") + + +def to_jsonld(g: Graph) -> str: + """Serialize a graph to JSON-LD. + + Parameters + ---------- + g : rdflib.Graph + The graph to serialize. + + Returns + ------- + str + JSON-LD formatted text. + """ + return g.serialize(format="json-ld") + + +def _reparse_or_raise(text: str, fmt: str) -> Graph: + """Parse serialized RDF text with a fresh Graph and raise on failure. + + Parameters + ---------- + text : str + Serialized RDF (Turtle or JSON-LD). + fmt : str + ``"turtle"`` or ``"json-ld"``. + + Returns + ------- + rdflib.Graph + The re-parsed graph. + + Raises + ------ + ProvenanceError + If the text does not parse as valid RDF in the given format. This + is a correctness self-check performed on every export, not just in + tests: an export this module cannot re-parse must never be handed + to a caller. + """ + try: + return Graph().parse(data=text, format=fmt) + except Exception as exc: # noqa: BLE001 - any parse failure is a real bug here + raise ProvenanceError( + f"Serialized {fmt} output failed to re-parse: {exc}" + ) from exc + + +def export_provenance( + bundle_dir: str | Path, + dataspec_path: str | Path | None = None, + fmt: str = "turtle", + base_iri: str = DEFAULT_BASE_IRI, + agent: str | None = None, + strict: bool = False, +) -> str: + """Build a provenance graph and return it serialized, verified re-parseable. + + Parameters + ---------- + bundle_dir : str or Path + A ``Segmentation.save()`` directory. + dataspec_path : str, Path, or None, optional + Optional DataSpec manifest JSON for image+label dataset provenance. + fmt : {"turtle", "json-ld"}, optional + Output serialization. Default ``"turtle"``. + base_iri : str, optional + Base IRI for minted identifiers. + agent : str or None, optional + Optional caller-supplied human/organization agent identity. + strict : bool, optional + If ``True``, raise on conditions this module would otherwise elide. + + Returns + ------- + str + The serialized graph, already verified to re-parse with a fresh + ``rdflib.Graph().parse()``. + + Raises + ------ + ProvenanceError + On any input, construction, boundary, or serialization failure. + ValueError + If ``fmt`` is not ``"turtle"`` or ``"json-ld"``. + """ + if fmt not in ("turtle", "json-ld"): + raise ValueError(f"fmt must be 'turtle' or 'json-ld', got {fmt!r}") + + g = build_graph( + bundle_dir, + dataspec_path=dataspec_path, + base_iri=base_iri, + agent=agent, + strict=strict, + ) + text = to_turtle(g) if fmt == "turtle" else to_jsonld(g) + _reparse_or_raise(text, fmt) + return text diff --git a/nobrainer/tests/unit/test_rdf_export.py b/nobrainer/tests/unit/test_rdf_export.py new file mode 100644 index 00000000..11d8166a --- /dev/null +++ b/nobrainer/tests/unit/test_rdf_export.py @@ -0,0 +1,653 @@ +"""Tests for nobrainer.provenance.rdf_export. + +Placed under nobrainer/tests/unit/ (not a top-level tests/ dir, which does +not exist in this repo) to match pyproject.toml's +``testpaths = ["nobrainer/tests"]`` and the existing unit-test convention. +Run with: uv run pytest nobrainer/tests/unit/test_rdf_export.py -q +""" + +from __future__ import annotations + +import json +from pathlib import Path +import subprocess +import sys + +import pytest +from rdflib import RDF, Graph +from rdflib.namespace import PROV + +from nobrainer.provenance import ( + ProvenanceError, + build_graph, + export_provenance, + to_jsonld, + to_turtle, +) +from nobrainer.provenance.rdf_export import ( + _assert_domain_only, + canonical_json_bytes, + digest, + normalize_float, +) + + +def _write_bundle( + tmp_path: Path, + *, + name: str = "run1", + source_datasets: list[dict] | None = None, + model_sha256: str = "c4f0deadbeef", + training_date: str = "2026-02-11T18:22:04.113295+00:00", + nobrainer_version: str = "2.0.0a17.dev6+gb85a1ca5b", + pytorch_version: str = "2.9.0+cu128", + model_architecture: str = "unet", + model_args: dict | None = None, + n_classes: int | None = 2, + block_shape: list[int] | None = None, + final_loss: float | None = 0.1873, + best_loss: float | None = 0.1712, + write_weights: bool = True, +) -> Path: + """Write a minimal Segmentation.save()-shaped bundle directory.""" + bundle_dir = tmp_path / name + bundle_dir.mkdir(parents=True, exist_ok=True) + if source_datasets is None: + source_datasets = [ + {"path": "/data/sub-01_T1w.nii.gz", "sha256": "3a7bd3e2360a3d29"} + ] + if model_args is None: + model_args = {"channels": [4, 8], "strides": [2]} + if block_shape is None: + block_shape = [16, 16, 16] + + doc = { + "@context": {"@vocab": "https://schema.org/"}, + "@type": "sc:Dataset", + "conformsTo": "http://mlcommons.org/croissant/1.0", + "name": f"nobrainer-{model_architecture}", + "description": f"Trained {model_architecture} model via nobrainer", + "distribution": [ + { + "@type": "cr:FileObject", + "name": "model.pth", + "contentUrl": "model.pth", + "encodingFormat": "application/x-pytorch", + "sha256": model_sha256, + } + ], + "nobrainer:provenance": { + "source_datasets": source_datasets, + "training_date": training_date, + "nobrainer_version": nobrainer_version, + "pytorch_version": pytorch_version, + "optimizer": {"class": "Adam", "args": {"lr": "0.001"}}, + "loss_function": "CrossEntropyLoss", + "epochs_trained": 12, + "final_loss": final_loss, + "best_loss": best_loss, + "model_architecture": model_architecture, + "model_args": model_args, + "n_classes": n_classes, + "block_shape": block_shape, + "gpu_count": 1, + }, + } + (bundle_dir / "croissant.json").write_text(json.dumps(doc, indent=2)) + if write_weights: + (bundle_dir / "model.pth").write_bytes(b"dummy-weights") + return bundle_dir + + +def _write_dataspec(tmp_path: Path, image_path: Path, label_path: Path) -> Path: + spec_path = tmp_path / "manifest.json" + spec_path.write_text( + json.dumps({"entries": [{"image": str(image_path), "label": str(label_path)}]}) + ) + return spec_path + + +# --------------------------------------------------------------------------- +# Hashing primitives +# --------------------------------------------------------------------------- + + +class TestHashingPrimitives: + def test_canonical_json_is_stable_across_key_order(self) -> None: + a = canonical_json_bytes({"b": 1, "a": 2}) + b = canonical_json_bytes({"a": 2, "b": 1}) + assert a == b + + def test_canonical_json_rejects_nan(self) -> None: + with pytest.raises(ValueError): + canonical_json_bytes({"x": float("nan")}) + + def test_digest_domain_separation(self) -> None: + payload = {"a": 1} + assert digest("entity", payload) != digest("activity", payload) + + def test_digest_length_is_128_bits(self) -> None: + assert len(digest("k", {"a": 1})) == 32 + + def test_digest_is_deterministic(self) -> None: + payload = {"a": 1, "b": [1, 2, 3]} + assert digest("k", payload) == digest("k", payload) + + def test_normalize_float_passes_none(self) -> None: + assert normalize_float(None) is None + + def test_normalize_float_rejects_nan(self) -> None: + with pytest.raises(ProvenanceError): + normalize_float(float("nan")) + + def test_normalize_float_rejects_inf(self) -> None: + with pytest.raises(ProvenanceError): + normalize_float(float("inf")) + + def test_normalize_float_returns_repr_string(self) -> None: + assert normalize_float(0.5) == repr(0.5) + + +# --------------------------------------------------------------------------- +# build_graph: required shape (the /goal's four graph requirements) +# --------------------------------------------------------------------------- + + +class TestGraphShape: + def test_run_is_prov_activity(self, tmp_path: Path) -> None: + g = build_graph(_write_bundle(tmp_path)) + activities = list(g.subjects(RDF.type, PROV.Activity)) + assert len(activities) == 1 + + def test_dataset_entity_is_used_by_run(self, tmp_path: Path) -> None: + g = build_graph(_write_bundle(tmp_path)) + run = next(g.subjects(RDF.type, PROV.Activity)) + used = list(g.objects(run, PROV.used)) + assert len(used) == 1 + assert (used[0], RDF.type, PROV.Entity) in g + + def test_model_entity_was_generated_by_run(self, tmp_path: Path) -> None: + g = build_graph(_write_bundle(tmp_path)) + run = next(g.subjects(RDF.type, PROV.Activity)) + generated = list(g.objects(run, PROV.generated)) + assert len(generated) == 1 + assert (generated[0], PROV.wasGeneratedBy, run) in g + + def test_agent_is_associated_with_run(self, tmp_path: Path) -> None: + g = build_graph(_write_bundle(tmp_path)) + run = next(g.subjects(RDF.type, PROV.Activity)) + agents = list(g.objects(run, PROV.wasAssociatedWith)) + assert len(agents) >= 1 + assert (agents[0], RDF.type, PROV.Agent) in g + + def test_no_blank_nodes(self, tmp_path: Path) -> None: + from rdflib import BNode + + g = build_graph(_write_bundle(tmp_path)) + for s, p, o in g: + assert not isinstance(s, BNode) + assert not isinstance(o, BNode) + + def test_model_sha256_is_recorded(self, tmp_path: Path) -> None: + g = build_graph(_write_bundle(tmp_path, model_sha256="c4f0deadbeef")) + run = next(g.subjects(RDF.type, PROV.Activity)) + model = next(g.objects(run, PROV.generated)) + # The model IRI itself is content-addressed under sha256/. + assert "sha256/c4f0deadbeef" in str(model) + + +# --------------------------------------------------------------------------- +# --agent +# --------------------------------------------------------------------------- + + +class TestAgentOption: + def test_no_agent_still_has_software_agents(self, tmp_path: Path) -> None: + g = build_graph(_write_bundle(tmp_path), agent=None) + run = next(g.subjects(RDF.type, PROV.Activity)) + agents = list(g.objects(run, PROV.wasAssociatedWith)) + assert len(agents) >= 1 + + def test_human_agent_text_adds_person(self, tmp_path: Path) -> None: + g = build_graph(_write_bundle(tmp_path), agent="Jane Doe ") + run = next(g.subjects(RDF.type, PROV.Activity)) + agents = list(g.objects(run, PROV.wasAssociatedWith)) + persons = [a for a in agents if (a, RDF.type, PROV.Person) in g] + assert len(persons) == 1 + + def test_agent_email_is_extracted(self, tmp_path: Path) -> None: + from rdflib.namespace import Namespace + + schema = Namespace("https://schema.org/") + g = build_graph(_write_bundle(tmp_path), agent="Jane Doe ") + run = next(g.subjects(RDF.type, PROV.Activity)) + persons = [ + a + for a in g.objects(run, PROV.wasAssociatedWith) + if (a, RDF.type, PROV.Person) in g + ] + emails = list(g.objects(persons[0], schema.email)) + assert str(emails[0]) == "jane@lab.org" + + def test_iri_agent_used_directly(self, tmp_path: Path) -> None: + from rdflib import URIRef + + g = build_graph( + _write_bundle(tmp_path), agent="https://orcid.org/0000-0000-0000-0000" + ) + run = next(g.subjects(RDF.type, PROV.Activity)) + agents = list(g.objects(run, PROV.wasAssociatedWith)) + assert URIRef("https://orcid.org/0000-0000-0000-0000") in agents + + def test_agent_is_not_invented_when_absent(self, tmp_path: Path) -> None: + g = build_graph(_write_bundle(tmp_path), agent=None) + assert list(g.subjects(RDF.type, PROV.Person)) == [] + + +# --------------------------------------------------------------------------- +# --dataspec +# --------------------------------------------------------------------------- + + +class TestDataspecOption: + def test_dataspec_produces_image_and_label_members(self, tmp_path: Path) -> None: + img = tmp_path / "sub-01_T1w.nii.gz" + lbl = tmp_path / "sub-01_aseg.nii.gz" + img.write_bytes(b"image-bytes") + lbl.write_bytes(b"label-bytes") + bundle = _write_bundle(tmp_path) + spec_path = _write_dataspec(tmp_path, img, lbl) + + g = build_graph(bundle, dataspec_path=spec_path) + run = next(g.subjects(RDF.type, PROV.Activity)) + dataset = next(g.objects(run, PROV.used)) + members = list(g.objects(dataset, PROV.hadMember)) + assert len(members) == 2 + + def test_without_dataspec_falls_back_to_croissant_source_datasets( + self, tmp_path: Path + ) -> None: + g = build_graph(_write_bundle(tmp_path), dataspec_path=None) + run = next(g.subjects(RDF.type, PROV.Activity)) + dataset = next(g.objects(run, PROV.used)) + members = list(g.objects(dataset, PROV.hadMember)) + assert len(members) == 1 + + +# --------------------------------------------------------------------------- +# Edge cases +# --------------------------------------------------------------------------- + + +class TestEdgeCases: + def test_empty_block_shape_emits_no_triple(self, tmp_path: Path) -> None: + from rdflib.namespace import Namespace + + nb = Namespace("https://neuronets.dev/ns/nobrainer#") + g = build_graph(_write_bundle(tmp_path, block_shape=[])) + run = next(g.subjects(RDF.type, PROV.Activity)) + assert (run, nb.blockShape, None) not in g + + def test_none_n_classes_emits_no_triple(self, tmp_path: Path) -> None: + from rdflib.namespace import Namespace + + nb = Namespace("https://neuronets.dev/ns/nobrainer#") + g = build_graph(_write_bundle(tmp_path, n_classes=None)) + run = next(g.subjects(RDF.type, PROV.Activity)) + assert (run, nb.numberOfClasses, None) not in g + + def test_missing_sha256_falls_back_to_recompute(self, tmp_path: Path) -> None: + from rdflib.namespace import Namespace + + nb = Namespace("https://neuronets.dev/ns/nobrainer#") + bundle = _write_bundle(tmp_path, model_sha256="", write_weights=True) + g = build_graph(bundle) + run = next(g.subjects(RDF.type, PROV.Activity)) + model = next(g.objects(run, PROV.generated)) + status = list(g.objects(model, nb.checksumStatus)) + assert str(status[0]) == "recomputed" + + def test_missing_sha256_and_no_file_is_unavailable(self, tmp_path: Path) -> None: + from rdflib.namespace import Namespace + + nb = Namespace("https://neuronets.dev/ns/nobrainer#") + bundle = _write_bundle(tmp_path, model_sha256="", write_weights=False) + g = build_graph(bundle) + run = next(g.subjects(RDF.type, PROV.Activity)) + model = next(g.objects(run, PROV.generated)) + status = list(g.objects(model, nb.checksumStatus)) + assert str(status[0]) == "unavailable" + + def test_strict_raises_on_unavailable_checksum(self, tmp_path: Path) -> None: + bundle = _write_bundle(tmp_path, model_sha256="", write_weights=False) + with pytest.raises(ProvenanceError): + build_graph(bundle, strict=True) + + def test_non_finite_loss_is_omitted_not_raised(self, tmp_path: Path) -> None: + from rdflib.namespace import Namespace + + nb = Namespace("https://neuronets.dev/ns/nobrainer#") + bundle = _write_bundle(tmp_path, final_loss=float("nan")) + g = build_graph(bundle, strict=False) + run = next(g.subjects(RDF.type, PROV.Activity)) + assert (run, nb.finalLoss, None) not in g + assert (run, nb.lossStatus, None) in g + + def test_non_finite_loss_raises_under_strict(self, tmp_path: Path) -> None: + bundle = _write_bundle(tmp_path, final_loss=float("nan")) + with pytest.raises(ProvenanceError): + build_graph(bundle, strict=True) + + def test_memory_address_poisoned_hyperparameter_is_elided( + self, tmp_path: Path + ) -> None: + from rdflib.namespace import Namespace + + nb = Namespace("https://neuronets.dev/ns/nobrainer#") + bundle = _write_bundle( + tmp_path, model_args={"callback": ""} + ) + g = build_graph(bundle, strict=False) + elided = list(g.subjects(nb.hyperparameterElided, None)) + assert len(elided) == 1 + + def test_poisoned_hyperparameter_raises_under_strict(self, tmp_path: Path) -> None: + bundle = _write_bundle( + tmp_path, model_args={"callback": ""} + ) + with pytest.raises(ProvenanceError): + build_graph(bundle, strict=True) + + def test_missing_croissant_json_raises(self, tmp_path: Path) -> None: + empty_dir = tmp_path / "empty" + empty_dir.mkdir() + with pytest.raises(ProvenanceError): + build_graph(empty_dir) + + def test_dataset_flavor_croissant_is_rejected(self, tmp_path: Path) -> None: + bundle_dir = tmp_path / "dsflavor" + bundle_dir.mkdir() + (bundle_dir / "croissant.json").write_text( + json.dumps( + { + "@type": "sc:Dataset", + "nobrainer:dataset_info": {"n_volumes": 3}, + } + ) + ) + with pytest.raises(ProvenanceError): + build_graph(bundle_dir) + + def test_checkpoint_flavor_architecture_is_not_tagged_registry( + self, tmp_path: Path + ) -> None: + """write_checkpoint_croissant's subset omits source_datasets/model_args/ + n_classes/block_shape and puts a torch class name in model_architecture.""" + from rdflib.namespace import Namespace + + nb = Namespace("https://neuronets.dev/ns/nobrainer#") + bundle_dir = tmp_path / "checkpoint_flavor" + bundle_dir.mkdir() + doc = { + "name": "nobrainer-MeshNet", + "description": "Trained MeshNet checkpoint via nobrainer", + "distribution": [ + { + "name": "best_model.pth", + "contentUrl": "best_model.pth", + "sha256": "aabbcc", + } + ], + "nobrainer:provenance": { + "training_date": "2026-01-01T00:00:00+00:00", + "nobrainer_version": "2.0.0a17.dev6+gb85a1ca5b", + "pytorch_version": "2.9.0", + "optimizer": {"class": "Adam", "args": {}}, + "loss_function": "CrossEntropyLoss", + "epochs_trained": 5, + "final_loss": 0.2, + "best_loss": 0.2, + "model_architecture": "MeshNet", + "gpu_count": 0, + }, + } + (bundle_dir / "croissant.json").write_text(json.dumps(doc)) + g = build_graph(bundle_dir) + run = next(g.subjects(RDF.type, PROV.Activity)) + vocab = list(g.objects(run, nb.architectureVocabulary)) + assert str(vocab[0]) == "torch-class-name" + + +# --------------------------------------------------------------------------- +# The DOMAIN-only boundary +# --------------------------------------------------------------------------- + + +class TestDomainOnlyBoundary: + def test_default_export_passes_boundary_check(self, tmp_path: Path) -> None: + g = build_graph(_write_bundle(tmp_path)) + _assert_domain_only(g) # must not raise + + def test_positive_control_second_activity_is_rejected(self, tmp_path: Path) -> None: + """Without this test, the boundary checks above could pass vacuously.""" + from rdflib import RDF, URIRef + + g = build_graph(_write_bundle(tmp_path)) + g.add( + (URIRef("https://example.org/ingestion-activity"), RDF.type, PROV.Activity) + ) + with pytest.raises(ProvenanceError): + _assert_domain_only(g) + + def test_no_brainkb_term_anywhere(self, tmp_path: Path) -> None: + g = build_graph(_write_bundle(tmp_path)) + for s, p, o in g: + assert "brainkb" not in str(s).lower() + assert "brainkb" not in str(p).lower() + assert "brainkb" not in str(o).lower() + + def test_module_docstring_states_brainkb_is_separate_step(self) -> None: + from nobrainer.provenance import rdf_export + + doc = (rdf_export.__doc__ or "").lower() + assert "brainkb" in doc + assert "separate" in doc + + +# --------------------------------------------------------------------------- +# Serialization: turtle/json-ld equivalence (required by the plan) +# --------------------------------------------------------------------------- + + +class TestSerialization: + def test_turtle_reparses(self, tmp_path: Path) -> None: + g = build_graph(_write_bundle(tmp_path)) + ttl = to_turtle(g) + reparsed = Graph().parse(data=ttl, format="turtle") + assert len(reparsed) == len(g) + + def test_jsonld_reparses(self, tmp_path: Path) -> None: + g = build_graph(_write_bundle(tmp_path)) + jsonld = to_jsonld(g) + reparsed = Graph().parse(data=jsonld, format="json-ld") + assert len(reparsed) == len(g) + + def test_jsonld_is_strict_valid_json(self, tmp_path: Path) -> None: + g = build_graph(_write_bundle(tmp_path)) + text = to_jsonld(g) + + def _boom(x): + raise ValueError(f"non-JSON constant: {x}") + + json.loads(text, parse_constant=_boom) # must not raise + + def test_turtle_and_jsonld_have_same_triple_count(self, tmp_path: Path) -> None: + bundle = _write_bundle(tmp_path) + ttl = export_provenance(bundle, fmt="turtle") + jsonld = export_provenance(bundle, fmt="json-ld") + g_ttl = Graph().parse(data=ttl, format="turtle") + g_jsonld = Graph().parse(data=jsonld, format="json-ld") + assert len(g_ttl) == len(g_jsonld) + + def test_export_provenance_rejects_bad_format(self, tmp_path: Path) -> None: + bundle = _write_bundle(tmp_path) + with pytest.raises(ValueError): + export_provenance(bundle, fmt="xml") + + +# --------------------------------------------------------------------------- +# Determinism +# --------------------------------------------------------------------------- + + +class TestDeterminism: + def test_repeated_export_is_byte_identical(self, tmp_path: Path) -> None: + bundle = _write_bundle(tmp_path) + a = export_provenance(bundle, fmt="turtle") + b = export_provenance(bundle, fmt="turtle") + assert a == b + + def test_run_iri_unchanged_after_directory_move(self, tmp_path: Path) -> None: + import shutil + + bundle = _write_bundle(tmp_path, name="orig") + moved = tmp_path / "moved" + shutil.copytree(bundle, moved) + + g1 = build_graph(bundle) + g2 = build_graph(moved) + run1 = next(g1.subjects(RDF.type, PROV.Activity)) + run2 = next(g2.subjects(RDF.type, PROV.Activity)) + assert run1 == run2 + + def test_run_iri_changes_when_epochs_trained_changes(self, tmp_path: Path) -> None: + """Negative control: without this, the identity payload could be constant.""" + bundle_a = _write_bundle(tmp_path, name="a") + bundle_b = tmp_path / "b" + bundle_b.mkdir() + doc = json.loads((bundle_a / "croissant.json").read_text()) + doc["nobrainer:provenance"]["epochs_trained"] = 999 + (bundle_b / "croissant.json").write_text(json.dumps(doc)) + (bundle_b / "model.pth").write_bytes(b"dummy-weights") + + g_a = build_graph(bundle_a) + g_b = build_graph(bundle_b) + run_a = next(g_a.subjects(RDF.type, PROV.Activity)) + run_b = next(g_b.subjects(RDF.type, PROV.Activity)) + assert run_a != run_b + + def test_base_iri_changes_instance_iris(self, tmp_path: Path) -> None: + bundle = _write_bundle(tmp_path) + g1 = build_graph(bundle, base_iri="https://neuronets.dev/nobrainer/") + g2 = build_graph(bundle, base_iri="https://example.org/nb/") + run1 = next(g1.subjects(RDF.type, PROV.Activity)) + run2 = next(g2.subjects(RDF.type, PROV.Activity)) + assert str(run1).startswith("https://neuronets.dev/nobrainer/") + assert str(run2).startswith("https://example.org/nb/") + + +# --------------------------------------------------------------------------- +# No-torch-import constraint +# --------------------------------------------------------------------------- + + +class TestImportIsolation: + def test_provenance_module_imports_without_torch_being_required(self) -> None: + result = subprocess.run( + [ + sys.executable, + "-c", + "import sys\n" + "import nobrainer.provenance\n" + "assert 'torch' not in sys.modules, " + "'nobrainer.provenance must not import torch'\n" + "print('OK')", + ], + capture_output=True, + text=True, + ) + assert result.returncode == 0, result.stdout + result.stderr + assert "OK" in result.stdout + + +# --------------------------------------------------------------------------- +# CLI +# --------------------------------------------------------------------------- + + +class TestCli: + def test_export_help_exits_zero(self) -> None: + result = subprocess.run( + [ + sys.executable, + "-m", + "nobrainer.cli.main", + "provenance", + "export", + "--help", + ], + capture_output=True, + text=True, + ) + assert result.returncode == 0, result.stderr + + def test_export_turtle_to_stdout(self, tmp_path: Path) -> None: + bundle = _write_bundle(tmp_path) + result = subprocess.run( + [ + sys.executable, + "-m", + "nobrainer.cli.main", + "provenance", + "export", + "--bundle", + str(bundle), + "--format", + "turtle", + ], + capture_output=True, + text=True, + ) + assert result.returncode == 0, result.stderr + g = Graph().parse(data=result.stdout, format="turtle") + assert len(list(g.subjects(RDF.type, PROV.Activity))) == 1 + + def test_export_writes_to_out_file(self, tmp_path: Path) -> None: + bundle = _write_bundle(tmp_path) + out_path = tmp_path / "out.ttl" + result = subprocess.run( + [ + sys.executable, + "-m", + "nobrainer.cli.main", + "provenance", + "export", + "--bundle", + str(bundle), + "--format", + "turtle", + "--out", + str(out_path), + ], + capture_output=True, + text=True, + ) + assert result.returncode == 0, result.stderr + assert out_path.exists() + Graph().parse(str(out_path), format="turtle") # must not raise + + def test_export_missing_bundle_dir_fails_cleanly(self, tmp_path: Path) -> None: + result = subprocess.run( + [ + sys.executable, + "-m", + "nobrainer.cli.main", + "provenance", + "export", + "--bundle", + str(tmp_path / "does-not-exist"), + ], + capture_output=True, + text=True, + ) + assert result.returncode != 0 diff --git a/pyproject.toml b/pyproject.toml index ecac034d..e2d80f82 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -62,11 +62,12 @@ generative = ["pytorch-lightning >= 2.0"] lightning = ["pytorch-lightning >= 2.0"] zarr = ["zarr >= 3.0", "nifti-zarr", "ome-zarr >= 0.14.0", "dask[array]", "scipy >= 1.11"] croissant = ["mlcroissant"] +provenance = ["rdflib >= 7.0"] versioning = ["datalad >= 0.19"] tfrecord = ["tfrecord >= 1.14"] dev = ["pre-commit", "pytest", "pytest-cov", "scipy"] all = [ - "nobrainer[bayesian,generative,zarr,croissant,versioning,tfrecord,dev]", + "nobrainer[bayesian,generative,zarr,croissant,provenance,versioning,tfrecord,dev]", ] [tool.hatch.version] From 53021f94888343f711751e2456ffba26f7b039e9 Mon Sep 17 00:00:00 2001 From: Dhritiman Das <14159298+dhritimandas@users.noreply.github.com> Date: Wed, 2 Sep 2026 11:48:37 +0530 Subject: [PATCH 7/7] fix(ci): install provenance extra so rdflib is available for tests nobrainer/tests/unit/test_rdf_export.py imports rdflib, but the CI and EC2 GPU workflows only installed [bayesian,generative,zarr,dev] -- missing rdflib fails collection and breaks the whole unit-test run on all three Python versions. Co-Authored-By: Claude Sonnet 5 Claude-Session: https://claude.ai/code/session_01VQwt3eBXLsr9sgvLM46B1P --- .github/workflows/ci.yml | 2 +- .github/workflows/guide-notebooks-ec2.yml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index d1a0ffcd..18c1f1cb 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -36,7 +36,7 @@ jobs: - name: Install dependencies run: | uv pip install \ - ".[bayesian,generative,zarr,dev]" \ + ".[bayesian,generative,zarr,provenance,dev]" \ monai \ pyro-ppl diff --git a/.github/workflows/guide-notebooks-ec2.yml b/.github/workflows/guide-notebooks-ec2.yml index 6787a10e..6b74d1bc 100644 --- a/.github/workflows/guide-notebooks-ec2.yml +++ b/.github/workflows/guide-notebooks-ec2.yml @@ -124,7 +124,7 @@ jobs: # Install nobrainer from checkout on top of the base layer uv pip install \ - ".[bayesian,generative,zarr,dev]" \ + ".[bayesian,generative,zarr,provenance,dev]" \ monai \ pyro-ppl \ matplotlib