diff --git a/junifer/api/__init__.py b/junifer/api/__init__.py index 422fd632a..bb976649a 100644 --- a/junifer/api/__init__.py +++ b/junifer/api/__init__.py @@ -7,4 +7,3 @@ from .cli import cli from .functions import run, collect from . import decorators -from . import registry diff --git a/junifer/api/decorators.py b/junifer/api/decorators.py index 49011871d..a5252486a 100644 --- a/junifer/api/decorators.py +++ b/junifer/api/decorators.py @@ -5,7 +5,7 @@ # Synchon Mandal # License: AGPL -from .registry import register +from ..pipeline.registry import register def register_datagrabber(klass: type) -> type: diff --git a/junifer/api/functions.py b/junifer/api/functions.py index bb9fb64e6..6e3796c57 100644 --- a/junifer/api/functions.py +++ b/junifer/api/functions.py @@ -16,10 +16,10 @@ import yaml from ..datagrabber.base import BaseDataGrabber from ..markers.base import BaseMarker from ..markers.collection import MarkerCollection +from ..pipeline.registry import build from ..storage.base import BaseFeatureStorage from ..utils import logger, raise_error from ..utils.fs import make_executable -from .registry import build def _get_datagrabber(datagrabber_config: Dict) -> BaseDataGrabber: diff --git a/junifer/api/tests/test_functions.py b/junifer/api/tests/test_functions.py index 3b175c901..284058a1c 100644 --- a/junifer/api/tests/test_functions.py +++ b/junifer/api/tests/test_functions.py @@ -11,8 +11,8 @@ import pytest import junifer.testing.registry # noqa: F401 from junifer.api.functions import collect, run -from junifer.api.registry import build from junifer.datagrabber.base import BaseDataGrabber +from junifer.pipeline.registry import build # Define datagrabber diff --git a/junifer/api/registry.py b/junifer/pipeline/registry.py similarity index 85% rename from junifer/api/registry.py rename to junifer/pipeline/registry.py index c1398d211..07cb40dc9 100644 --- a/junifer/api/registry.py +++ b/junifer/pipeline/registry.py @@ -11,12 +11,13 @@ from ..utils.logging import logger, raise_error if TYPE_CHECKING: - from ..datagrabber.base import BaseDataGrabber - from ..pipeline import PipelineStepMixin - from ..storage.base import BaseFeatureStorage + from ..datagrabber import BaseDataGrabber + from ..storage import BaseFeatureStorage + from .pipeline_step_mixin import PipelineStepMixin + # Define valid steps for operation -_valid_steps = [ +_VALID_STEPS: List[str] = [ "datagrabber", "datareader", "preprocessing", @@ -25,7 +26,7 @@ _valid_steps = [ ] # Define registry for valid steps -_registry = {x: {} for x in _valid_steps} +_REGISTRY: Dict[str, Dict[str, type]] = {x: {} for x in _VALID_STEPS} def register(step: str, name: str, klass: type) -> None: @@ -42,14 +43,14 @@ def register(step: str, name: str, klass: type) -> None: """ # Verify step - if step not in _valid_steps: + if step not in _VALID_STEPS: raise_error(msg=f"Invalid step: {step}", klass=ValueError) logger.info(f"Registering {name} in {step}") - _registry[step][name] = klass + _REGISTRY[step][name] = klass -def get_step_names(step: str) -> List: +def get_step_names(step: str) -> List[str]: """Get the names of the registered functions for a given step. Parameters @@ -64,10 +65,10 @@ def get_step_names(step: str) -> List: """ # Verify step - if step not in _valid_steps: + if step not in _VALID_STEPS: raise_error(msg=f"Invalid step: {step}", klass=ValueError) - return list(_registry[step].keys()) + return list(_REGISTRY[step].keys()) def get_class(step: str, name: str) -> type: @@ -87,13 +88,13 @@ def get_class(step: str, name: str) -> type: """ # Verify step - if step not in _valid_steps: + if step not in _VALID_STEPS: raise_error(msg=f"Invalid step: {step}", klass=ValueError) # Verify step name - if name not in _registry[step]: + if name not in _REGISTRY[step]: raise_error(msg=f"Invalid name: {name}", klass=ValueError) - return _registry[step][name] + return _REGISTRY[step][name] def build( diff --git a/junifer/api/tests/test_registry.py b/junifer/pipeline/tests/test_registry.py similarity index 97% rename from junifer/api/tests/test_registry.py rename to junifer/pipeline/tests/test_registry.py index 7a45c0d5f..a84a9b7ee 100644 --- a/junifer/api/tests/test_registry.py +++ b/junifer/pipeline/tests/test_registry.py @@ -10,8 +10,13 @@ from typing import Type import pytest -from junifer.api.registry import build, get_class, get_step_names, register from junifer.datagrabber import PatternDataGrabber +from junifer.pipeline.registry import ( + build, + get_class, + get_step_names, + register, +) from junifer.storage import SQLiteFeatureStorage diff --git a/junifer/testing/registry.py b/junifer/testing/registry.py index 7e9c84a74..014aa50d4 100644 --- a/junifer/testing/registry.py +++ b/junifer/testing/registry.py @@ -4,7 +4,7 @@ # Synchon Mandal # License: AGPL -from ..api.registry import register +from ..pipeline.registry import register from .datagrabbers import ( OasisVBMTestingDatagrabber, SPMAuditoryTestingDatagrabber, diff --git a/junifer/testing/tests/test_testing_registry.py b/junifer/testing/tests/test_testing_registry.py index 090b6709f..6d2f3efc9 100644 --- a/junifer/testing/tests/test_testing_registry.py +++ b/junifer/testing/tests/test_testing_registry.py @@ -2,14 +2,14 @@ import importlib -from junifer.api.registry import get_step_names +from junifer.pipeline.registry import get_step_names def test_testing_registry() -> None: """Test testing registry.""" import junifer - importlib.reload(junifer.api.registry) + importlib.reload(junifer.pipeline.registry) importlib.reload(junifer) assert "OasisVBMTestingDatagrabber" not in get_step_names("datagrabber")