diff --git a/docs/changes/newsfragments/450.doc b/docs/changes/newsfragments/450.doc new file mode 100644 index 000000000..a339ee6e9 --- /dev/null +++ b/docs/changes/newsfragments/450.doc @@ -0,0 +1 @@ +Add documentation on extending data registries by `Synchon Mandal`_ diff --git a/docs/changes/newsfragments/450.feature b/docs/changes/newsfragments/450.feature new file mode 100644 index 000000000..c75be19d4 --- /dev/null +++ b/docs/changes/newsfragments/450.feature @@ -0,0 +1 @@ +Enable data registries to be extended by exposing :class:`.BasePipelineDataRegistry` and introducing :func:`.register_data_registry` decorator by `Synchon Mandal`_ diff --git a/docs/extending/data_registries.rst b/docs/extending/data_registries.rst new file mode 100644 index 000000000..d2bf7ceae --- /dev/null +++ b/docs/extending/data_registries.rst @@ -0,0 +1,70 @@ +.. include:: ../links.inc + +.. _data_registries: + +Creating a Data Registry +======================== + +What is a data registry +----------------------- + +Data Registry is an object which manages pipeline data (like parcellations, +coordinates and masks) based on the scope. ``junifer`` comes with the +following in-built data registries: + +* :class:`.ParcellationRegistry` for parcellations +* :class:`.CoordinatesRegistry` for coordinates +* :class:`.MaskRegistry` for masks + +You interact with them indirectly via :func:`.get_data`, :func:`.load_data`, +:func:`.list_data`, :func:`.register_data` and :func:`.deregister_data`. The +``kind`` parameter in them directs which data registry to interact with. These +in turn supply preprocessors and markers with their necessary data for computation. + +How to make a data registry +--------------------------- + +Ideally you would not need to create a custom data registry if you work with fMRI data. +In case you work with other modalities like EEG, you might want to create a +data registry for montages. Here is how you would go about it: + +#. Check :ref:`extending junifer ` on how to create a + *junifer extension* if you have not done so. +#. Create the data registry in the *extension script* like so: + + .. code-block:: python + + from junifer.api.decorators import register_data_registry + from junifer.data import BasePipelineDataRegistry + + + @register_data_registry("montage") + class MontageDataRegistry(BasePipelineDataRegistry): + def __init__(self): + super().__init__() + + def register(self): + pass + + def deregister(self): + pass + + def load(self): + pass + + def get(self): + pass + + + * :func:`.register_data_registry` registers a class with the name passed in + the argument, ``"montage"`` in this case. + * Inheriting from :class:`.BasePipelineDataRegistry` takes care of the class + acting as a data registry. + * :meth:`.BasePipelineDataRegistry.register`, + :meth:`.BasePipelineDataRegistry.deregister`, + :meth:`.BasePipelineDataRegistry.load` and + :meth:`.BasePipelineDataRegistry.get` need to be implemented + (check other registries for reference implementations). + +#. Pass ``kind="montage"`` in ``*_data`` functions to + verify if your data registry is set up properly. diff --git a/docs/extending/extension.rst b/docs/extending/extension.rst index 2e2b52d3f..e4d7942c9 100644 --- a/docs/extending/extension.rst +++ b/docs/extending/extension.rst @@ -5,8 +5,8 @@ Creating a ``junifer`` extension ================================ -``junifer`` is designed to be easily extensible. Through the use of a registry -and decorators, you can easily add new functionality to ``junifer`` during +``junifer`` is designed to be easily extensible. Through the use of data registries, +a component registry and decorators, you can easily add new functionality to ``junifer`` during runtime. This is done by creating a new Python module and importing it before running ``junifer``. diff --git a/docs/extending/index.rst b/docs/extending/index.rst index 5a95060fb..0fba1a95d 100644 --- a/docs/extending/index.rst +++ b/docs/extending/index.rst @@ -31,3 +31,4 @@ DataGrabbers, Preprocessors, Markers, etc., following the *junifer* way. coordinates masks plugins + data_registries diff --git a/junifer/api/decorators.py b/junifer/api/decorators.py index e69e3be3c..aa0f377d5 100644 --- a/junifer/api/decorators.py +++ b/junifer/api/decorators.py @@ -5,11 +5,19 @@ # Synchon Mandal # License: AGPL +from ..data import DataDispatcher from ..pipeline import PipelineComponentRegistry -from ..typing import DataGrabberLike, MarkerLike, PreprocessorLike, StorageLike +from ..typing import ( + DataGrabberLike, + DataRegistryLike, + MarkerLike, + PreprocessorLike, + StorageLike, +) __all__ = [ + "register_data_registry", "register_datagrabber", "register_datareader", "register_marker", @@ -25,12 +33,12 @@ def register_datagrabber(klass: DataGrabberLike) -> DataGrabberLike: Parameters ---------- - klass: class + klass : class The class of the DataGrabber to register. Returns ------- - klass: class + class The unmodified input class. Notes @@ -52,12 +60,12 @@ def register_datareader(klass: type) -> type: Parameters ---------- - klass: class + klass : class The class of the DataReader to register. Returns ------- - klass: class + class The unmodified input class. Notes @@ -79,12 +87,12 @@ def register_preprocessor(klass: PreprocessorLike) -> PreprocessorLike: Parameters ---------- - klass: class + klass : class The class of the preprocessor to register. Returns ------- - klass: class + class The unmodified input class. """ @@ -102,12 +110,12 @@ def register_marker(klass: MarkerLike) -> MarkerLike: Parameters ---------- - klass: class + klass : class The class of the marker to register. Returns ------- - klass: class + class The unmodified input class. """ @@ -125,12 +133,12 @@ def register_storage(klass: StorageLike) -> StorageLike: Parameters ---------- - klass: class + klass : class The class of the storage to register. Returns ------- - klass: class + class The unmodified input class. """ @@ -139,3 +147,40 @@ def register_storage(klass: StorageLike) -> StorageLike: klass=klass, ) return klass + + +def register_data_registry(name: str) -> DataRegistryLike: + """Registry registration decorator. + + Registers the data registry as ``name``. + + Parameters + ---------- + name : str + The name of the data registry. + + Returns + ------- + class + The unmodified input class. + + """ + + def decorator(klass: DataRegistryLike) -> DataRegistryLike: + """Actual decorator. + + Parameters + ---------- + klass : class + The class of the data registry to register. + + Returns + ------- + class + The unmodified input class. + + """ + DataDispatcher()[name] = klass + return klass + + return decorator diff --git a/junifer/api/tests/test_decorators.py b/junifer/api/tests/test_decorators.py new file mode 100644 index 000000000..8dec54483 --- /dev/null +++ b/junifer/api/tests/test_decorators.py @@ -0,0 +1,32 @@ +"""Provide tests for public decorators.""" + +# Authors: Synchon Mandal +# License: AGPL + +from junifer.api.decorators import register_data_registry +from junifer.data import BasePipelineDataRegistry, DataDispatcher + + +def test_register_data_registry() -> None: + """Test data registry registration.""" + + @register_data_registry("dumb") + class DumDum(BasePipelineDataRegistry): + def __init__(self): + super().__init__() + + def register(self): + pass + + def deregister(self): + pass + + def load(self): + pass + + def get(self): + pass + + assert "dumb" in DataDispatcher() + _ = DataDispatcher().pop("dumb") + assert "dumb" not in DataDispatcher() diff --git a/junifer/data/__init__.pyi b/junifer/data/__init__.pyi index a48d55fd1..d4e8ac1cc 100644 --- a/junifer/data/__init__.pyi +++ b/junifer/data/__init__.pyi @@ -1,5 +1,7 @@ __all__ = [ + "BasePipelineDataRegistry", "CoordinatesRegistry", + "DataDispatcher", "ParcellationRegistry", "MaskRegistry", "get_data", @@ -12,11 +14,13 @@ __all__ = [ "utils", ] +from .pipeline_data_registry_base import BasePipelineDataRegistry from .coordinates import CoordinatesRegistry from .parcellations import ParcellationRegistry from .masks import MaskRegistry from ._dispatch import ( + DataDispatcher, get_data, list_data, load_data, diff --git a/junifer/data/_dispatch.py b/junifer/data/_dispatch.py index 4d94f46b4..81c1d5d4e 100644 --- a/junifer/data/_dispatch.py +++ b/junifer/data/_dispatch.py @@ -3,6 +3,7 @@ # Authors: Synchon Mandal # License: AGPL +from collections.abc import Iterator, MutableMapping from pathlib import Path from typing import ( TYPE_CHECKING, @@ -18,13 +19,14 @@ from ..utils import raise_error from .coordinates import CoordinatesRegistry from .masks import MaskRegistry from .parcellations import ParcellationRegistry +from .pipeline_data_registry_base import BasePipelineDataRegistry if TYPE_CHECKING: from nibabel.nifti1 import Nifti1Image - __all__ = [ + "DataDispatcher", "deregister_data", "get_data", "list_data", @@ -33,6 +35,77 @@ __all__ = [ ] +class DataDispatcher(MutableMapping): + """Class for helping dynamic data dispatch.""" + + _instance = None + + def __new__(cls): + # Make class singleton + if cls._instance is None: + cls._instance = super().__new__(cls) + # Set registries + cls._registries: dict[str, type[BasePipelineDataRegistry]] = {} + cls._builtin: dict[str, type[BasePipelineDataRegistry]] = {} + cls._external: dict[str, type[BasePipelineDataRegistry]] = {} + cls._builtin.update( + { + "coordinates": CoordinatesRegistry, + "parcellation": ParcellationRegistry, + "mask": MaskRegistry, + } + ) + cls._registries.update(cls._builtin) + return cls._instance + + def __getitem__(self, key: str) -> type[BasePipelineDataRegistry]: + return self._registries[key] + + def __iter__(self) -> Iterator[str]: + return iter(self._registries) + + def __len__(self) -> int: + return len(self._registries) + + def __delitem__(self, key: str) -> None: + # Internal check + if key in self._builtin: + raise_error(f"Cannot delete in-built key: {key}") + # Non-existing key + if key not in self._external: + raise_error(klass=KeyError, msg=key) + # Update external + _ = self._external.pop(key) + # Update global + _ = self._registries.pop(key) + + def __setitem__( + self, key: str, value: type[BasePipelineDataRegistry] + ) -> None: + # Internal check + if key in self._builtin: + raise_error(f"Cannot set value for in-built key: {key}") + # Value type check + if not issubclass(value, BasePipelineDataRegistry): + raise_error(f"Invalid value type: {type(value)}") + # Update external + self._external[key] = value + # Update global + self._registries[key] = value + + def popitem(): + """Not implemented.""" + pass + + def clear(self): + """Not implemented.""" + pass + + def setdefault(self, key: str, value=None): + """Not implemented.""" + pass + + def get_data( kind: str, names: Union[ @@ -76,27 +149,16 @@ def get_data( If ``kind`` is invalid value. """ - - if kind == "coordinates": - return CoordinatesRegistry().get( - coords=names, - target_data=target_data, - extra_input=extra_input, - ) - elif kind == "parcellation": - return ParcellationRegistry().get( - parcellations=names, - target_data=target_data, - extra_input=extra_input, - ) - elif kind == "mask": - return MaskRegistry().get( - masks=names, - target_data=target_data, - extra_input=extra_input, - ) - else: # pragma: no cover + try: + registry = DataDispatcher()[kind] + except KeyError: raise_error(f"Unknown data kind: {kind}") + else: + return registry().get( + names, + target_data=target_data, + extra_input=extra_input, + ) def list_data(kind: str) -> list[str]: @@ -119,14 +181,12 @@ def list_data(kind: str) -> list[str]: """ - if kind == "coordinates": - return CoordinatesRegistry().list - elif kind == "parcellation": - return ParcellationRegistry().list - elif kind == "mask": - return MaskRegistry().list - else: # pragma: no cover + try: + registry = DataDispatcher()[kind] + except KeyError: raise_error(f"Unknown data kind: {kind}") + else: + return registry().list def load_data( @@ -165,15 +225,12 @@ def load_data( If ``kind`` is invalid value. """ - - if kind == "coordinates": - return CoordinatesRegistry().load(name=name) - elif kind == "parcellation": - return ParcellationRegistry().load(name=name, **kwargs) - elif kind == "mask": - return MaskRegistry().load(name=name, **kwargs) - else: # pragma: no cover + try: + registry = DataDispatcher()[kind] + except KeyError: raise_error(f"Unknown data kind: {kind}") + else: + return registry().load(name, **kwargs) def register_data( @@ -205,20 +262,14 @@ def register_data( """ - if kind == "coordinates": - return CoordinatesRegistry().register( - name=name, space=space, overwrite=overwrite, **kwargs - ) - elif kind == "parcellation": - return ParcellationRegistry().register( - name=name, space=space, overwrite=overwrite, **kwargs - ) - elif kind == "mask": - return MaskRegistry().register( - name=name, space=space, overwrite=overwrite, **kwargs - ) - else: # pragma: no cover + try: + registry = DataDispatcher()[kind] + except KeyError: raise_error(f"Unknown data kind: {kind}") + else: + return registry().register( + name=name, space=space, overwrite=overwrite, **kwargs + ) def deregister_data(kind: str, name: str) -> None: @@ -238,11 +289,9 @@ def deregister_data(kind: str, name: str) -> None: """ - if kind == "coordinates": - return CoordinatesRegistry().deregister(name=name) - elif kind == "parcellation": - return ParcellationRegistry().deregister(name=name) - elif kind == "mask": - return MaskRegistry().deregister(name=name) - else: # pragma: no cover + try: + registry = DataDispatcher()[kind] + except KeyError: raise_error(f"Unknown data kind: {kind}") + else: + return registry().deregister(name=name) diff --git a/junifer/data/coordinates/_coordinates.py b/junifer/data/coordinates/_coordinates.py index bc88e0eb2..fc3821c72 100644 --- a/junifer/data/coordinates/_coordinates.py +++ b/junifer/data/coordinates/_coordinates.py @@ -13,7 +13,6 @@ from junifer_data import get from numpy.typing import ArrayLike from ...utils import logger, raise_error -from ...utils.singleton import Singleton from ..pipeline_data_registry_base import BasePipelineDataRegistry from ..utils import JUNIFER_DATA_PARAMS, get_dataset_path, get_native_warper from ._ants_coordinates_warper import ANTsCoordinatesWarper @@ -23,7 +22,7 @@ from ._fsl_coordinates_warper import FSLCoordinatesWarper __all__ = ["CoordinatesRegistry"] -class CoordinatesRegistry(BasePipelineDataRegistry, metaclass=Singleton): +class CoordinatesRegistry(BasePipelineDataRegistry): """Class for coordinates data registry. This class is a singleton and is used for managing available coordinates diff --git a/junifer/data/masks/_masks.py b/junifer/data/masks/_masks.py index 42c0775cb..8a7526010 100644 --- a/junifer/data/masks/_masks.py +++ b/junifer/data/masks/_masks.py @@ -24,7 +24,6 @@ from nilearn.masking import ( ) from ...utils import logger, raise_error -from ...utils.singleton import Singleton from ..pipeline_data_registry_base import BasePipelineDataRegistry from ..template_spaces import get_template from ..utils import ( @@ -216,7 +215,7 @@ def compute_brain_mask( return nimg.new_img_like(target_data["data"], mask) # type: ignore -class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): +class MaskRegistry(BasePipelineDataRegistry): """Class for mask data registry. This class is a singleton and is used for managing available mask diff --git a/junifer/data/parcellations/_parcellations.py b/junifer/data/parcellations/_parcellations.py index 71ab71b12..b9947c796 100644 --- a/junifer/data/parcellations/_parcellations.py +++ b/junifer/data/parcellations/_parcellations.py @@ -16,7 +16,6 @@ import pandas as pd from junifer_data import get from ...utils import logger, raise_error, warn_with_log -from ...utils.singleton import Singleton from ..pipeline_data_registry_base import BasePipelineDataRegistry from ..utils import ( JUNIFER_DATA_PARAMS, @@ -38,7 +37,7 @@ __all__ = [ ] -class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): +class ParcellationRegistry(BasePipelineDataRegistry): """Class for parcellation data registry. This class is a singleton and is used for managing available parcellation diff --git a/junifer/data/parcellations/tests/test_parcellations.py b/junifer/data/parcellations/tests/test_parcellations.py index be82f3643..fc47f30e8 100644 --- a/junifer/data/parcellations/tests/test_parcellations.py +++ b/junifer/data/parcellations/tests/test_parcellations.py @@ -14,6 +14,7 @@ from nilearn.image import new_img_like, resample_to_img from numpy.testing import assert_array_almost_equal, assert_array_equal from junifer.data import ( + deregister_data, get_data, list_data, load_data, @@ -1245,3 +1246,9 @@ def test_get_multi_different_space() -> None: ], target_data=element_data["VBM_GM"], ) + + +def test_deregister() -> None: + """Test parcellation deregistration.""" + deregister_data(kind="parcellation", name="testparc_3") + assert "testparc_3" not in list_data(kind="parcellation") diff --git a/junifer/data/tests/test_dispatch.py b/junifer/data/tests/test_dispatch.py new file mode 100644 index 000000000..483396a7d --- /dev/null +++ b/junifer/data/tests/test_dispatch.py @@ -0,0 +1,87 @@ +"""Provide tests for data dispatching.""" + +# Authors: Synchon Mandal +# License: AGPL + +import pytest + +from junifer.data import ( + BasePipelineDataRegistry, + deregister_data, + get_data, + list_data, + load_data, + register_data, +) +from junifer.data._dispatch import DataDispatcher + + +def test_dispatcher_addition_errors() -> None: + """Test registry addition errors.""" + with pytest.raises(ValueError, match="Cannot set"): + DataDispatcher()["mask"] = dict + + with pytest.raises(ValueError, match="Invalid"): + DataDispatcher()["masks"] = dict + + +def test_dispatcher_removal_errors() -> None: + """Test registry removal errors.""" + with pytest.raises(ValueError, match="Cannot delete"): + _ = DataDispatcher().pop("mask") + + with pytest.raises(KeyError, match="masks"): + del DataDispatcher()["masks"] + + +def test_dispatcher() -> None: + """Test registry addition and removal.""" + + class DumDum(BasePipelineDataRegistry): + def register(): + pass + + def deregister(): + pass + + def load(): + pass + + def get(): + pass + + DataDispatcher().update({"masks": DumDum}) + assert "masks" in DataDispatcher() + + _ = DataDispatcher().pop("masks") + assert "masks" not in DataDispatcher() + + +def test_get_data_error() -> None: + """Test error for get_data().""" + with pytest.raises(ValueError, match="Unknown data kind"): + get_data(kind="planet", names="neptune", target_data={}) + + +def test_list_data_error() -> None: + """Test error for list_data().""" + with pytest.raises(ValueError, match="Unknown data kind"): + list_data(kind="planet") + + +def test_load_data_error() -> None: + """Test error for load_data().""" + with pytest.raises(ValueError, match="Unknown data kind"): + load_data(kind="planet", name="neptune") + + +def test_register_data_error() -> None: + """Test error for register_data().""" + with pytest.raises(ValueError, match="Unknown data kind"): + register_data(kind="planet", name="neptune", space="milkyway") + + +def test_deregister_data_error() -> None: + """Test error for deregister_data().""" + with pytest.raises(ValueError, match="Unknown data kind"): + deregister_data(kind="planet", name="neptune") diff --git a/junifer/typing/__init__.pyi b/junifer/typing/__init__.pyi index ff489f5e1..7db1445e2 100644 --- a/junifer/typing/__init__.pyi +++ b/junifer/typing/__init__.pyi @@ -1,5 +1,6 @@ __all__ = [ "DataGrabberLike", + "DataRegistryLike", "PreprocessorLike", "MarkerLike", "StorageLike", @@ -16,6 +17,7 @@ __all__ = [ from ._typing import ( DataGrabberLike, + DataRegistryLike, PreprocessorLike, MarkerLike, StorageLike, diff --git a/junifer/typing/_typing.py b/junifer/typing/_typing.py index cd2311320..d95417ac3 100644 --- a/junifer/typing/_typing.py +++ b/junifer/typing/_typing.py @@ -11,6 +11,7 @@ from typing import ( if TYPE_CHECKING: + from ..data import BasePipelineDataRegistry from ..datagrabber import BaseDataGrabber from ..datareader import DefaultDataReader from ..markers import BaseMarker @@ -23,6 +24,7 @@ __all__ = [ "ConfigVal", "DataGrabberLike", "DataGrabberPatterns", + "DataRegistryLike", "Dependencies", "Element", "Elements", @@ -35,6 +37,7 @@ __all__ = [ ] +DataRegistryLike = type["BasePipelineDataRegistry"] DataGrabberLike = type["BaseDataGrabber"] PreprocessorLike = type["BasePreprocessor"] MarkerLike = type["BaseMarker"]