[ENH] Move junifer.api.registry to junifer.pipeline #50

Merged
synchon merged 8 commits from refactor/api-registry into main 2022-10-18 09:03:45 +00:00
8 changed files with 26 additions and 21 deletions

View file

@ -7,4 +7,3 @@
from .cli import cli from .cli import cli
from .functions import run, collect from .functions import run, collect
from . import decorators from . import decorators
from . import registry

View file

@ -5,7 +5,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from .registry import register from ..pipeline.registry import register
def register_datagrabber(klass: type) -> type: def register_datagrabber(klass: type) -> type:

View file

@ -16,10 +16,10 @@ import yaml
from ..datagrabber.base import BaseDataGrabber from ..datagrabber.base import BaseDataGrabber
from ..markers.base import BaseMarker from ..markers.base import BaseMarker
from ..markers.collection import MarkerCollection from ..markers.collection import MarkerCollection
from ..pipeline.registry import build
from ..storage.base import BaseFeatureStorage from ..storage.base import BaseFeatureStorage
from ..utils import logger, raise_error from ..utils import logger, raise_error
from ..utils.fs import make_executable from ..utils.fs import make_executable
from .registry import build
def _get_datagrabber(datagrabber_config: Dict) -> BaseDataGrabber: def _get_datagrabber(datagrabber_config: Dict) -> BaseDataGrabber:

View file

@ -11,8 +11,8 @@ import pytest
import junifer.testing.registry # noqa: F401 import junifer.testing.registry # noqa: F401
from junifer.api.functions import collect, run from junifer.api.functions import collect, run
from junifer.api.registry import build
from junifer.datagrabber.base import BaseDataGrabber from junifer.datagrabber.base import BaseDataGrabber
from junifer.pipeline.registry import build
# Define datagrabber # Define datagrabber

View file

@ -11,12 +11,13 @@ from ..utils.logging import logger, raise_error
if TYPE_CHECKING: if TYPE_CHECKING:
from ..datagrabber.base import BaseDataGrabber from ..datagrabber import BaseDataGrabber
from ..pipeline import PipelineStepMixin from ..storage import BaseFeatureStorage
from ..storage.base import BaseFeatureStorage from .pipeline_step_mixin import PipelineStepMixin
# Define valid steps for operation # Define valid steps for operation
_valid_steps = [ _VALID_STEPS: List[str] = [
"datagrabber", "datagrabber",
"datareader", "datareader",
"preprocessing", "preprocessing",
@ -25,7 +26,7 @@ _valid_steps = [
] ]
# Define registry for 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: def register(step: str, name: str, klass: type) -> None:
@ -42,14 +43,14 @@ def register(step: str, name: str, klass: type) -> None:
""" """
# Verify step # Verify step
if step not in _valid_steps: if step not in _VALID_STEPS:
raise_error(msg=f"Invalid step: {step}", klass=ValueError) raise_error(msg=f"Invalid step: {step}", klass=ValueError)
logger.info(f"Registering {name} in {step}") 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. """Get the names of the registered functions for a given step.
Parameters Parameters
@ -64,10 +65,10 @@ def get_step_names(step: str) -> List:
""" """
# Verify step # Verify step
if step not in _valid_steps: if step not in _VALID_STEPS:
raise_error(msg=f"Invalid step: {step}", klass=ValueError) 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: def get_class(step: str, name: str) -> type:
@ -87,13 +88,13 @@ def get_class(step: str, name: str) -> type:
""" """
# Verify step # Verify step
if step not in _valid_steps: if step not in _VALID_STEPS:
raise_error(msg=f"Invalid step: {step}", klass=ValueError) raise_error(msg=f"Invalid step: {step}", klass=ValueError)
# Verify step name # Verify step name
if name not in _registry[step]: if name not in _REGISTRY[step]:
raise_error(msg=f"Invalid name: {name}", klass=ValueError) raise_error(msg=f"Invalid name: {name}", klass=ValueError)
return _registry[step][name] return _REGISTRY[step][name]
def build( def build(

View file

@ -10,8 +10,13 @@ from typing import Type
import pytest import pytest
from junifer.api.registry import build, get_class, get_step_names, register
from junifer.datagrabber import PatternDataGrabber from junifer.datagrabber import PatternDataGrabber
from junifer.pipeline.registry import (
build,
get_class,
get_step_names,
register,
)
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage

View file

@ -4,7 +4,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from ..api.registry import register from ..pipeline.registry import register
from .datagrabbers import ( from .datagrabbers import (
OasisVBMTestingDatagrabber, OasisVBMTestingDatagrabber,
SPMAuditoryTestingDatagrabber, SPMAuditoryTestingDatagrabber,

View file

@ -2,14 +2,14 @@
import importlib import importlib
from junifer.api.registry import get_step_names from junifer.pipeline.registry import get_step_names
def test_testing_registry() -> None: def test_testing_registry() -> None:
"""Test testing registry.""" """Test testing registry."""
import junifer import junifer
importlib.reload(junifer.api.registry) importlib.reload(junifer.pipeline.registry)
importlib.reload(junifer) importlib.reload(junifer)
assert "OasisVBMTestingDatagrabber" not in get_step_names("datagrabber") assert "OasisVBMTestingDatagrabber" not in get_step_names("datagrabber")