[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 .functions import run, collect
from . import decorators
from . import registry

View file

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

View file

@ -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:

View file

@ -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

View file

@ -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(

View file

@ -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

View file

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

View file

@ -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")