diff --git a/docs/changes/newsfragments/319.enh b/docs/changes/newsfragments/319.enh new file mode 100644 index 000000000..4c10fa6e6 --- /dev/null +++ b/docs/changes/newsfragments/319.enh @@ -0,0 +1 @@ +Raise error when partial or complete element selectors are invalid when running the pipeline by `Synchon Mandal`_ diff --git a/junifer/api/functions.py b/junifer/api/functions.py index 6ffa974cc..122e8f8c5 100644 --- a/junifer/api/functions.py +++ b/junifer/api/functions.py @@ -20,7 +20,13 @@ from ..pipeline import ( ) from ..preprocess import BasePreprocessor from ..storage import BaseFeatureStorage -from ..typing import DataGrabberLike, MarkerLike, PreprocessorLike, StorageLike +from ..typing import ( + DataGrabberLike, + Elements, + MarkerLike, + PreprocessorLike, + StorageLike, +) from ..utils import logger, raise_error, warn_with_log, yaml @@ -121,7 +127,7 @@ def run( markers: list[dict], storage: dict, preprocessors: Optional[list[dict]] = None, - elements: Optional[list[tuple[str, ...]]] = None, + elements: Optional[Elements] = None, ) -> None: """Run the pipeline on the selected element. @@ -147,7 +153,7 @@ def run( List of preprocessors to use. Each preprocessor is a dict with at least a key ``kind`` specifying the preprocessor to use. All other keys are passed to the preprocessor constructor (default None). - elements : list of tuple or None, optional + elements : list or None, optional Element(s) to process. Will be used to index the DataGrabber (default None). @@ -155,6 +161,8 @@ def run( ------ ValueError If ``workdir.cleanup=False`` when ``len(elements) > 1``. + RuntimeError + If invalid element selectors are found. """ # Conditional to handle workdir config @@ -208,10 +216,22 @@ def run( # Fit elements with datagrabber_object: if elements is not None: - for t_element in datagrabber_object.filter( - elements # type: ignore - ): + # Keep track of valid selectors + valid_elements = [] + for t_element in datagrabber_object.filter(elements): + valid_elements.append(t_element) mc.fit(datagrabber_object[t_element]) + # Compute invalid selectors + invalid_elements = set(elements) - set(valid_elements) + # Report if invalid selectors are found + if invalid_elements: + raise_error( + msg=( + "The following element selectors are invalid:\n" + f"{invalid_elements}" + ), + klass=RuntimeError, + ) else: for t_element in datagrabber_object: mc.fit(datagrabber_object[t_element]) @@ -243,7 +263,7 @@ def queue( kind: str, jobname: str = "junifer_job", overwrite: bool = False, - elements: Optional[list[tuple[str, ...]]] = None, + elements: Optional[Elements] = None, **kwargs: Union[str, int, bool, dict, tuple, list], ) -> None: """Queue a job to be executed later. @@ -258,7 +278,7 @@ def queue( The name of the job (default "junifer_job"). overwrite : bool, optional Whether to overwrite if job directory already exists (default False). - elements : list of tuple or None, optional + elements : list or None, optional Element(s) to process. Will be used to index the DataGrabber (default None). **kwargs : dict @@ -341,7 +361,7 @@ def queue( elements = dg.get_elements() # Listify elements if not isinstance(elements, list): - elements: list[Union[str, tuple]] = [elements] + elements: Elements = [elements] # Check job queueing system adapter = None @@ -406,7 +426,7 @@ def reset(config: dict) -> None: def list_elements( datagrabber: dict, - elements: Optional[list[tuple[str, ...]]] = None, + elements: Optional[Elements] = None, ) -> str: """List elements of the datagrabber filtered using `elements`. @@ -416,7 +436,7 @@ def list_elements( DataGrabber to index. Must have a key ``kind`` with the kind of DataGrabber to use. All other keys are passed to the DataGrabber constructor. - elements : list of tuple or None, optional + elements : list or None, optional Element(s) to filter using. Will be used to index the DataGrabber (default None). diff --git a/junifer/api/queue_context/gnu_parallel_local_adapter.py b/junifer/api/queue_context/gnu_parallel_local_adapter.py index a14b29057..792c2b91b 100644 --- a/junifer/api/queue_context/gnu_parallel_local_adapter.py +++ b/junifer/api/queue_context/gnu_parallel_local_adapter.py @@ -6,8 +6,9 @@ import shutil import textwrap from pathlib import Path -from typing import Optional, Union +from typing import Optional +from ...typing import Elements from ...utils import logger, make_executable, raise_error, run_ext_cmd from .queue_context_adapter import QueueContextAdapter @@ -63,7 +64,7 @@ class GnuParallelLocalAdapter(QueueContextAdapter): job_name: str, job_dir: Path, yaml_config_path: Path, - elements: list[Union[str, tuple]], + elements: Elements, pre_run: Optional[str] = None, pre_collect: Optional[str] = None, env: Optional[dict[str, str]] = None, diff --git a/junifer/api/queue_context/htcondor_adapter.py b/junifer/api/queue_context/htcondor_adapter.py index d3a004d04..d6fd44630 100644 --- a/junifer/api/queue_context/htcondor_adapter.py +++ b/junifer/api/queue_context/htcondor_adapter.py @@ -6,8 +6,9 @@ import shutil import textwrap from pathlib import Path -from typing import Optional, Union +from typing import Optional +from ...typing import Elements from ...utils import logger, make_executable, raise_error, run_ext_cmd from .queue_context_adapter import QueueContextAdapter @@ -81,7 +82,7 @@ class HTCondorAdapter(QueueContextAdapter): job_name: str, job_dir: Path, yaml_config_path: Path, - elements: list[Union[str, tuple]], + elements: Elements, pre_run: Optional[str] = None, pre_collect: Optional[str] = None, env: Optional[dict[str, str]] = None, diff --git a/junifer/api/tests/test_functions.py b/junifer/api/tests/test_functions.py index 8a42bbb05..5dc0e2c59 100644 --- a/junifer/api/tests/test_functions.py +++ b/junifer/api/tests/test_functions.py @@ -6,16 +6,19 @@ # License: AGPL import logging +from contextlib import AbstractContextManager, nullcontext from pathlib import Path -from typing import Optional, Union +from typing import Any, Optional, Union import pytest +from nibabel.filebasedimages import ImageFileError from ruamel.yaml import YAML import junifer.testing.registry # noqa: F401 from junifer.api import collect, list_elements, queue, reset, run from junifer.datagrabber.base import BaseDataGrabber from junifer.pipeline import PipelineComponentRegistry +from junifer.typing import Elements # Configure YAML class @@ -25,12 +28,37 @@ yaml.allow_unicode = True yaml.indent(mapping=2, sequence=4, offset=2) +# Kept for parametrizing +_datagrabber = { + "kind": "PartlyCloudyTestingDataGrabber", +} +_bids_ses_datagrabber = { + "kind": "PatternDataladDataGrabber", + "uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses", + "types": ["T1w", "BOLD"], + "patterns": { + "T1w": { + "pattern": ( + "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz" + ), + "space": "MNI152NLin6Asym", + }, + "BOLD": { + "pattern": ( + "{subject}/{session}/func/{subject}_{session}_task-rest_bold.nii.gz" + ), + "space": "MNI152NLin6Asym", + }, + }, + "replacements": ["subject", "session"], + "rootdir": "example_bids_ses", +} + + @pytest.fixture def datagrabber() -> dict[str, str]: """Return a datagrabber as a dictionary.""" - return { - "kind": "PartlyCloudyTestingDataGrabber", - } + return _datagrabber.copy() @pytest.fixture @@ -60,11 +88,48 @@ def storage() -> dict[str, str]: } +@pytest.mark.parametrize( + "datagrabber, element, expect", + [ + ( + _datagrabber, + [("sub-01",)], + pytest.raises(RuntimeError, match="element selectors are invalid"), + ), + ( + _datagrabber, + ["sub-01"], + nullcontext(), + ), + ( + _bids_ses_datagrabber, + ["sub-01"], + pytest.raises(ImageFileError, match="is not a gzip file"), + ), + ( + _bids_ses_datagrabber, + [("sub-01", "ses-01")], + pytest.raises(ImageFileError, match="is not a gzip file"), + ), + ( + _bids_ses_datagrabber, + [("sub-01", "ses-100")], + pytest.raises(RuntimeError, match="element selectors are invalid"), + ), + ( + _bids_ses_datagrabber, + [("sub-100", "ses-01")], + pytest.raises(RuntimeError, match="element selectors are invalid"), + ), + ], +) def test_run_single_element( tmp_path: Path, - datagrabber: dict[str, str], + datagrabber: dict[str, Any], markers: list[dict[str, str]], storage: dict[str, str], + element: Elements, + expect: AbstractContextManager, ) -> None: """Test run function with single element. @@ -78,21 +143,26 @@ def test_run_single_element( Testing markers as list of dictionary. storage : dict Testing storage as dictionary. + element : list of str or tuple + The parametrized element. + expect : typing.ContextManager + The parametrized ContextManager object. """ # Set storage storage["uri"] = str((tmp_path / "out.sqlite").resolve()) # Run operations - run( - workdir=tmp_path, - datagrabber=datagrabber, - markers=markers, - storage=storage, - elements=[("sub-01",)], - ) - # Check files - files = list(tmp_path.glob("*.sqlite")) - assert len(files) == 1 + with expect: + run( + workdir=tmp_path, + datagrabber=datagrabber, + markers=markers, + storage=storage, + elements=element, + ) + # Check files + files = list(tmp_path.glob("*.sqlite")) + assert len(files) == 1 def test_run_single_element_with_preprocessing( @@ -128,18 +198,30 @@ def test_run_single_element_with_preprocessing( "kind": "fMRIPrepConfoundRemover", } ], - elements=[("sub-01",)], + elements=["sub-01"], ) # Check files files = list(tmp_path.glob("*.sqlite")) assert len(files) == 1 +@pytest.mark.parametrize( + "element, expect", + [ + ( + [("sub-01",), ("sub-03",)], + pytest.raises(RuntimeError, match="element selectors are invalid"), + ), + (["sub-01", "sub-03"], nullcontext()), + ], +) def test_run_multi_element_multi_output( tmp_path: Path, datagrabber: dict[str, str], markers: list[dict[str, str]], storage: dict[str, str], + element: Elements, + expect: AbstractContextManager, ) -> None: """Test run function with multi element and multi output. @@ -153,22 +235,27 @@ def test_run_multi_element_multi_output( Testing markers as list of dictionary. storage : dict Testing storage as dictionary. + element : list of str or tuple + The parametrized element. + expect : typing.ContextManager + The parametrized ContextManager object. """ # Set storage storage["uri"] = str((tmp_path / "out.sqlite").resolve()) storage["single_output"] = False # type: ignore # Run operations - run( - workdir=tmp_path, - datagrabber=datagrabber, - markers=markers, - storage=storage, - elements=[("sub-01",), ("sub-03",)], - ) - # Check files - files = list(tmp_path.glob("*.sqlite")) - assert len(files) == 2 + with expect: + run( + workdir=tmp_path, + datagrabber=datagrabber, + markers=markers, + storage=storage, + elements=element, + ) + # Check files + files = list(tmp_path.glob("*.sqlite")) + assert len(files) == 2 def test_run_multi_element_single_output( @@ -200,7 +287,7 @@ def test_run_multi_element_single_output( datagrabber=datagrabber, markers=markers, storage=storage, - elements=[("sub-01",), ("sub-03",)], + elements=["sub-01", "sub-03"], ) # Check files files = list(tmp_path.glob("*.sqlite")) @@ -569,7 +656,7 @@ def test_reset_run( datagrabber=datagrabber, markers=markers, storage=storage, - elements=[("sub-01",)], + elements=["sub-01"], ) # Reset operation reset(config={"storage": storage}) diff --git a/junifer/cli/parser.py b/junifer/cli/parser.py index da64c96a7..31e745a52 100644 --- a/junifer/cli/parser.py +++ b/junifer/cli/parser.py @@ -12,6 +12,7 @@ from typing import Union import pandas as pd +from ..typing import Elements from ..utils import logger, raise_error, warn_with_log, yaml @@ -142,7 +143,7 @@ def parse_yaml(filepath: Union[str, Path]) -> dict: # noqa: C901 def parse_elements( element: tuple[str, ...], config: dict -) -> Union[list[tuple[str, ...]], None]: +) -> Union[Elements, None]: """Parse elements from cli. Parameters @@ -203,7 +204,7 @@ def parse_elements( return elements -def _parse_elements_file(filepath: Path) -> list[tuple[str, ...]]: +def _parse_elements_file(filepath: Path) -> Elements: """Parse elements from file. Parameters @@ -213,7 +214,7 @@ def _parse_elements_file(filepath: Path) -> list[tuple[str, ...]]: Returns ------- - list of tuple of str + list The element(s) as list. """ @@ -227,5 +228,8 @@ def _parse_elements_file(filepath: Path) -> list[tuple[str, ...]]: ) # Remove trailing whitespace in cell entries csv_df_trimmed = csv_df.apply(lambda x: x.str.strip()) - # Convert to list of tuple of str - return list(map(tuple, csv_df_trimmed.to_numpy())) + # Convert to list of tuple of str if more than one column else flatten + if len(csv_df_trimmed.columns) == 1: + return csv_df_trimmed.to_numpy().flatten().tolist() + else: + return list(map(tuple, csv_df_trimmed.to_numpy())) diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index 99b1f7bca..f2707ec80 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -11,6 +11,7 @@ from pathlib import Path from typing import Union from ..pipeline import UpdateMetaMixin +from ..typing import Element, Elements from ..utils import logger, raise_error @@ -67,9 +68,7 @@ class BaseDataGrabber(ABC, UpdateMetaMixin): """ yield from self.get_elements() - def __getitem__( - self, element: Union[str, tuple[str, ...]] - ) -> dict[str, dict]: + def __getitem__(self, element: Element) -> dict[str, dict]: """Enable indexing support. Parameters @@ -137,13 +136,13 @@ class BaseDataGrabber(ABC, UpdateMetaMixin): """ return self._datadir - def filter(self, selection: list[Union[str, tuple[str]]]) -> Iterator: + def filter(self, selection: Elements) -> Iterator: """Filter elements to be grabbed. Parameters ---------- - selection : list of str or tuple - The list of partial element key values to filter using. + selection : list + The list of partial or complete element selectors to filter using. Yields ------ @@ -152,7 +151,7 @@ class BaseDataGrabber(ABC, UpdateMetaMixin): """ - def filter_func(element: Union[str, tuple[str]]) -> bool: + def filter_func(element: Element) -> bool: """Filter element based on selection. Parameters @@ -201,15 +200,14 @@ class BaseDataGrabber(ABC, UpdateMetaMixin): ) # pragma: no cover @abstractmethod - def get_elements(self) -> list[Union[str, tuple[str]]]: + def get_elements(self) -> Elements: """Get elements. Returns ------- list - List of elements that can be grabbed. The elements can be strings, - tuples or any object that will be then used as a key to index the - DataGrabber. + List of elements that can be grabbed. The elements can be strings + or tuples of strings to index the DataGrabber. """ raise_error( diff --git a/junifer/datagrabber/datalad_base.py b/junifer/datagrabber/datalad_base.py index f7eb88d95..f7eb996dd 100644 --- a/junifer/datagrabber/datalad_base.py +++ b/junifer/datagrabber/datalad_base.py @@ -17,6 +17,7 @@ from datalad.support.exceptions import IncompleteResultsError from datalad.support.gitrepo import GitRepo from ..pipeline import WorkDirManager +from ..typing import Element from ..utils import config, logger, raise_error, warn_with_log from .base import BaseDataGrabber @@ -312,7 +313,7 @@ class DataladDataGrabber(BaseDataGrabber): logger.debug(f"Dropping {f}") self._dataset.drop(f, result_renderer="disabled") - def __getitem__(self, element: Union[str, tuple]) -> dict: + def __getitem__(self, element: Element) -> dict: """Implement single element indexing in the Datalad database. It will first obtain the paths from the parent class and then @@ -320,11 +321,11 @@ class DataladDataGrabber(BaseDataGrabber): Parameters ---------- - element : str or tuple + element : str or tuple of str The element to be indexed. If one string is provided, it is assumed to be a tuple with only one item. If a tuple is provided, each item in the tuple is the value for the replacement string - specified in "replacements". + specified in ``"replacements"``. Returns ------- diff --git a/junifer/datagrabber/pattern.py b/junifer/datagrabber/pattern.py index dafbe4893..7b8173cbf 100644 --- a/junifer/datagrabber/pattern.py +++ b/junifer/datagrabber/pattern.py @@ -13,7 +13,7 @@ from typing import Optional, Union import numpy as np from ..api.decorators import register_datagrabber -from ..typing import DataGrabberPatterns +from ..typing import DataGrabberPatterns, Elements from ..utils import logger, raise_error from .base import BaseDataGrabber from .pattern_validation_mixin import PatternValidationMixin @@ -367,7 +367,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin): """ return self.replacements - def get_item(self, **element: str) -> dict[str, dict]: + def get_item(self, **element: dict) -> dict[str, dict]: """Implement single element indexing for the datagrabber. This method constructs a real path to the requested item's data, by @@ -450,7 +450,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin): return out - def get_elements(self) -> list: + def get_elements(self) -> Elements: """Implement fetching list of elements in the dataset. It will use regex to search for "replacements" in the "patterns" and diff --git a/junifer/typing/__init__.pyi b/junifer/typing/__init__.pyi index bfe026fda..ff489f5e1 100644 --- a/junifer/typing/__init__.pyi +++ b/junifer/typing/__init__.pyi @@ -10,6 +10,8 @@ __all__ = [ "MarkerInOutMappings", "DataGrabberPatterns", "ConfigVal", + "Element", + "Elements", ] from ._typing import ( @@ -24,4 +26,6 @@ from ._typing import ( MarkerInOutMappings, DataGrabberPatterns, ConfigVal, + Element, + Elements, ) diff --git a/junifer/typing/_typing.py b/junifer/typing/_typing.py index d7275b597..b0887241a 100644 --- a/junifer/typing/_typing.py +++ b/junifer/typing/_typing.py @@ -24,6 +24,8 @@ __all__ = [ "DataGrabberLike", "DataGrabberPatterns", "Dependencies", + "Element", + "Elements", "ExternalDependencies", "MarkerInOutMappings", "MarkerLike", @@ -62,3 +64,5 @@ DataGrabberPatterns = dict[ str, Union[dict[str, str], Sequence[dict[str, str]]] ] ConfigVal = Union[bool, int, float] +Element = Union[str, tuple[str, ...]] +Elements = Sequence[Element]