From d6bac966c78248892d95bd916622ceba79461adf Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 10 Jan 2023 11:35:54 +0100 Subject: [PATCH 01/20] fix: register MultipleDataGrabber --- junifer/datagrabber/multiple.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/junifer/datagrabber/multiple.py b/junifer/datagrabber/multiple.py index fef8f0da1..c79655661 100644 --- a/junifer/datagrabber/multiple.py +++ b/junifer/datagrabber/multiple.py @@ -7,6 +7,7 @@ from typing import Dict, List, Tuple, Union +from ..api.decorators import register_datagrabber from ..utils import raise_error from .base import BaseDataGrabber @@ -14,6 +15,7 @@ from .base import BaseDataGrabber __all__ = ["MultipleDataGrabber"] +@register_datagrabber class MultipleDataGrabber(BaseDataGrabber): """Concrete implementation for multi sourced data fetching. -- 2.52.0 From a5a72e2bb0097f9b09bdd377e26e3dbfa802ab0d Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 11 Jul 2024 12:29:15 +0200 Subject: [PATCH 02/20] refactor: move pattern validation for datagrabbers to PatternValidationMixin --- junifer/datagrabber/pattern.py | 26 +- .../datagrabber/pattern_validation_mixin.py | 388 ++++++++++++++++++ ...ls.py => test_pattern_validation_mixin.py} | 241 ++++++----- junifer/datagrabber/utils.py | 317 -------------- 4 files changed, 539 insertions(+), 433 deletions(-) create mode 100644 junifer/datagrabber/pattern_validation_mixin.py rename junifer/datagrabber/tests/{test_datagrabber_utils.py => test_pattern_validation_mixin.py} (53%) delete mode 100644 junifer/datagrabber/utils.py diff --git a/junifer/datagrabber/pattern.py b/junifer/datagrabber/pattern.py index 3947a3953..6d8b39f1f 100644 --- a/junifer/datagrabber/pattern.py +++ b/junifer/datagrabber/pattern.py @@ -15,7 +15,7 @@ import numpy as np from ..api.decorators import register_datagrabber from ..utils import logger, raise_error from .base import BaseDataGrabber -from .utils import validate_patterns, validate_replacements +from .pattern_validation_mixin import PatternValidationMixin __all__ = ["PatternDataGrabber"] @@ -26,7 +26,7 @@ _CONFOUNDS_FORMATS = ("fmriprep", "adhoc") @register_datagrabber -class PatternDataGrabber(BaseDataGrabber): +class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin): """Concrete implementation for pattern-based data fetching. Implements a DataGrabber that understands patterns to grab data. @@ -142,6 +142,13 @@ class PatternDataGrabber(BaseDataGrabber): The directory where the data is / will be stored. confounds_format : {"fmriprep", "adhoc"} or None, optional The format of the confounds for the dataset (default None). + partial_pattern_ok : bool, optional + Whether to raise error if partial pattern for a data type is found. + This allows to bypass mandatory key check and issue a warning + instead of raising error. This allows one to have a DataGrabber + with data types without the corresponding mandatory keys and is + powerful when used with :class:`.MultipleDataGrabber` + (default True). Raises ------ @@ -157,17 +164,20 @@ class PatternDataGrabber(BaseDataGrabber): replacements: Union[List[str], str], datadir: Union[str, Path], confounds_format: Optional[str] = None, + partial_pattern_ok: bool = False, ) -> None: - # Validate patterns - validate_patterns(types=types, patterns=patterns) - self.patterns = patterns - # Convert replacements to list if not already if not isinstance(replacements, list): replacements = [replacements] - # Validate replacements - validate_replacements(replacements=replacements, patterns=patterns) + # Validate patterns + self.validate_patterns( + types=types, + replacements=replacements, + patterns=patterns, + partial_pattern_ok=partial_pattern_ok, + ) self.replacements = replacements + self.patterns = patterns # Validate confounds format if ( diff --git a/junifer/datagrabber/pattern_validation_mixin.py b/junifer/datagrabber/pattern_validation_mixin.py new file mode 100644 index 000000000..9c130b902 --- /dev/null +++ b/junifer/datagrabber/pattern_validation_mixin.py @@ -0,0 +1,388 @@ +"""Provide mixin validation class for pattern-based DataGrabber.""" + +# Authors: Synchon Mandal +# License: AGPL + +from typing import Dict, List + +from ..utils import logger, raise_error, warn_with_log + + +__all__ = ["PatternValidationMixin"] + + +# Define schema for pattern-based datagrabber's patterns +PATTERNS_SCHEMA = { + "T1w": { + "mandatory": ["pattern", "space"], + "optional": { + "mask": {"mandatory": ["pattern", "space"], "optional": []}, + }, + }, + "T2w": { + "mandatory": ["pattern", "space"], + "optional": { + "mask": {"mandatory": ["pattern", "space"], "optional": []}, + }, + }, + "BOLD": { + "mandatory": ["pattern", "space"], + "optional": { + "mask": {"mandatory": ["pattern", "space"], "optional": []}, + "confounds": { + "mandatory": ["pattern", "format"], + "optional": ["mappings"], + }, + }, + }, + "Warp": { + "mandatory": ["pattern", "src", "dst"], + "optional": {}, + }, + "VBM_GM": { + "mandatory": ["pattern", "space"], + "optional": {}, + }, + "VBM_WM": { + "mandatory": ["pattern", "space"], + "optional": {}, + }, + "VBM_CSF": { + "mandatory": ["pattern", "space"], + "optional": {}, + }, + "DWI": { + "mandatory": ["pattern"], + "optional": {}, + }, + "FreeSurfer": { + "mandatory": ["pattern"], + "optional": { + "aseg": {"mandatory": ["pattern"], "optional": []}, + "norm": {"mandatory": ["pattern"], "optional": []}, + "lh_white": {"mandatory": ["pattern"], "optional": []}, + "rh_white": {"mandatory": ["pattern"], "optional": []}, + "lh_pial": {"mandatory": ["pattern"], "optional": []}, + "rh_pial": {"mandatory": ["pattern"], "optional": []}, + }, + }, +} + + +class PatternValidationMixin: + """Mixin class for pattern validation.""" + + def _validate_types(self, types: List[str]) -> None: + """Validate the types. + + Parameters + ---------- + types : list of str + The data types to validate. + + Raises + ------ + TypeError + If ``types`` is not a list or if the values are not string. + + """ + if not isinstance(types, list): + raise_error(msg="`types` must be a list", klass=TypeError) + if any(not isinstance(x, str) for x in types): + raise_error( + msg="`types` must be a list of strings", klass=TypeError + ) + + def _validate_replacements( + self, + replacements: List[str], + patterns: Dict[str, Dict[str, str]], + partial_pattern_ok: bool, + ) -> None: + """Validate the replacements. + + Parameters + ---------- + replacements : list of str + The replacements to validate. + patterns : dict + The patterns to validate replacements against. + partial_pattern_ok : bool + Whether to raise error if partial pattern for a data type is found. + + Raises + ------ + TypeError + If ``replacements`` is not a list or if the values are not string. + ValueError + If a value in ``replacements`` is not part of a data type pattern + and ``partial_pattern_ok=False`` or + if no data type patterns contain all values in ``replacements`` and + ``partial_pattern_ok=False``. + + Warns + ----- + RuntimeWarning + If a value in ``replacements`` is not part of the data type pattern + and ``partial_pattern_ok=True``. + + """ + if not isinstance(replacements, list): + raise_error(msg="`replacements` must be a list.", klass=TypeError) + + if any(not isinstance(x, str) for x in replacements): + raise_error( + msg="`replacements` must be a list of strings.", + klass=TypeError, + ) + + for x in replacements: + if all( + x not in y + for y in [ + data_type_val.get("pattern", "") + for data_type_val in patterns.values() + ] + ): + if partial_pattern_ok: + warn_with_log( + f"Replacement: `{x}` is not part of any pattern, " + "things might not work as expected if you are unsure " + "of what you are doing" + ) + else: + raise_error( + msg=f"Replacement: {x} is not part of any pattern." + ) + + # Check that at least one pattern has all the replacements + at_least_one = False + for data_type_val in patterns.values(): + if all( + x in data_type_val.get("pattern", "") for x in replacements + ): + at_least_one = True + if not at_least_one and not partial_pattern_ok: + raise_error( + msg="At least one pattern must contain all replacements." + ) + + def _validate_mandatory_keys( + self, + keys: List[str], + schema: List[str], + data_type: str, + partial_pattern_ok: bool = False, + ) -> None: + """Validate mandatory keys. + + Parameters + ---------- + keys : list of str + The keys to validate. + schema : list of str + The schema to validate against. + data_type : str + The data type being validated. + partial_pattern_ok : bool, optional + Whether to raise error if partial pattern for a data type is found + (default True). + + Raises + ------ + KeyError + If any mandatory key is missing for a data type and + ``partial_pattern_ok=False``. + + Warns + ----- + RuntimeWarning + If any mandatory key is missing for a data type and + ``partial_pattern_ok=True``. + + """ + for key in schema: + if key not in keys: + if partial_pattern_ok: + warn_with_log( + f"Mandatory key: `{key}` not found for {data_type}, " + "things might not work as expected if you are unsure " + "of what you are doing" + ) + else: + raise_error( + msg=f"Mandatory key: `{key}` missing for {data_type}", + klass=KeyError, + ) + else: + logger.debug(f"Mandatory key: `{key}` found for {data_type}") + + def _identify_stray_keys( + self, keys: List[str], schema: List[str], data_type: str + ) -> None: + """Identify stray keys. + + Parameters + ---------- + keys : list of str + The keys to check. + schema : list of str + The schema to check against. + data_type : str + The data type being checked. + + Raises + ------ + RuntimeError + If an unknown key is found for a data type. + + """ + for key in keys: + if key not in schema: + raise_error( + msg=( + f"Key: {key} not accepted for {data_type} " + "pattern, remove it to proceed" + ), + klass=RuntimeError, + ) + + def validate_patterns( + self, + types: List[str], + replacements: List[str], + patterns: Dict[str, Dict[str, str]], + partial_pattern_ok: bool = False, + ) -> None: + """Validate the patterns. + + Parameters + ---------- + types : list of str + The data types to check patterns of. + replacements : list of str + The replacements to be replaced in the patterns. + patterns : dict + The patterns to validate. + partial_pattern_ok : bool, optional + Whether to raise error if partial pattern for a data type is found. + If False, a warning is issued instead of raising an error + (default False). + + Raises + ------ + TypeError + If ``patterns`` is not a dictionary. + ValueError + If length of ``types`` and ``patterns`` are different or + if ``patterns`` is missing entries from ``types`` or + if unknown data type is found in ``patterns`` or + if data type pattern key contains '*' as value. + + """ + # Validate types + self._validate_types(types=types) + + # Validate patterns + if not isinstance(patterns, dict): + raise_error(msg="`patterns` must be a dict", klass=TypeError) + # Unequal length of objects + if len(types) > len(patterns): + raise_error( + msg="Length of `types` more than that of `patterns`", + klass=ValueError, + ) + # Missing type in patterns + if any(x not in patterns for x in types): + raise_error( + msg="`patterns` must contain all `types`", klass=ValueError + ) + # Check against schema + for data_type_key, data_type_val in patterns.items(): + # Check if valid data type is provided + if data_type_key not in PATTERNS_SCHEMA: + raise_error( + f"Unknown data type: {data_type_key}, " + f"should be one of: {list(PATTERNS_SCHEMA.keys())}" + ) + # Check mandatory keys for data type + self._validate_mandatory_keys( + keys=list(data_type_val), + schema=PATTERNS_SCHEMA[data_type_key]["mandatory"], + data_type=data_type_key, + partial_pattern_ok=partial_pattern_ok, + ) + # Check optional keys for data type + for optional_key, optional_val in PATTERNS_SCHEMA[data_type_key][ + "optional" + ].items(): + if optional_key not in data_type_val: + logger.debug( + f"Optional key: `{optional_key}` missing for " + f"{data_type_key}" + ) + else: + logger.debug( + f"Optional key: `{optional_key}` found for " + f"{data_type_key}" + ) + # Set nested type name for easier access + nested_data_type = f"{data_type_key}.{optional_key}" + nested_mandatory_keys_schema = PATTERNS_SCHEMA[ + data_type_key + ]["optional"][optional_key]["mandatory"] + nested_optional_keys_schema = PATTERNS_SCHEMA[ + data_type_key + ]["optional"][optional_key]["optional"] + # Check mandatory keys for nested type + self._validate_mandatory_keys( + keys=list(optional_val["mandatory"]), + schema=nested_mandatory_keys_schema, + data_type=nested_data_type, + partial_pattern_ok=partial_pattern_ok, + ) + # Check optional keys for nested type + for nested_optional_key in nested_optional_keys_schema: + if nested_optional_key not in optional_val["optional"]: + logger.debug( + f"Optional key: `{nested_optional_key}` " + f"missing for {nested_data_type}" + ) + else: + logger.debug( + f"Optional key: `{nested_optional_key}` found " + f"for {nested_data_type}" + ) + # Check stray key for nested data type + self._identify_stray_keys( + keys=optional_val["mandatory"] + + optional_val["optional"], + schema=nested_mandatory_keys_schema + + nested_optional_keys_schema, + data_type=nested_data_type, + ) + # Check stray key for data type + self._identify_stray_keys( + keys=list(data_type_val.keys()), + schema=( + PATTERNS_SCHEMA[data_type_key]["mandatory"] + + list(PATTERNS_SCHEMA[data_type_key]["optional"].keys()) + ), + data_type=data_type_key, + ) + # Wildcard check in patterns + if "}*" in data_type_val.get("pattern", ""): + raise_error( + msg=( + f"`{data_type_key}.pattern` must not contain `*` " + "following a replacement" + ), + klass=ValueError, + ) + + # Validate replacements + self._validate_replacements( + replacements=replacements, + patterns=patterns, + partial_pattern_ok=partial_pattern_ok, + ) diff --git a/junifer/datagrabber/tests/test_datagrabber_utils.py b/junifer/datagrabber/tests/test_pattern_validation_mixin.py similarity index 53% rename from junifer/datagrabber/tests/test_datagrabber_utils.py rename to junifer/datagrabber/tests/test_pattern_validation_mixin.py index ec5c8e613..1b9c08376 100644 --- a/junifer/datagrabber/tests/test_datagrabber_utils.py +++ b/junifer/datagrabber/tests/test_pattern_validation_mixin.py @@ -1,6 +1,7 @@ -"""Provide tests for utils.""" +"""Provide tests for PatternValidationMixin.""" # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL from contextlib import nullcontext @@ -8,136 +9,57 @@ from typing import ContextManager, Dict, List, Union import pytest -from junifer.datagrabber.utils import ( - validate_patterns, - validate_replacements, - validate_types, -) +from junifer.datagrabber.pattern_validation_mixin import PatternValidationMixin @pytest.mark.parametrize( - "types, expect", - [ - ("wrong", pytest.raises(TypeError, match="must be a list")), - ([1], pytest.raises(TypeError, match="must be a list of strings")), - (["T1w", "BOLD"], nullcontext()), - ], -) -def test_validate_types( - types: Union[str, List[str], List[int]], - expect: ContextManager, -) -> None: - """Test validation of types. - - Parameters - ---------- - types : str, list of int or str - The parametrized data types to validate. - expect : typing.ContextManager - The parametrized ContextManager object. - - """ - with expect: - validate_types(types) # type: ignore - - -@pytest.mark.parametrize( - "replacements, patterns, expect", + "types, replacements, patterns, expect", [ ( "wrong", - "also wrong", - pytest.raises(TypeError, match="must be a list"), + [], + {}, + pytest.raises(TypeError, match="`types` must be a list"), ), ( [1], - { - "T1w": {"pattern": "{subject}/anat/{subject}_T1w.nii.gz"}, - "BOLD": { - "pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz" - }, - }, - pytest.raises(TypeError, match="must be a list of strings"), + [], + {}, + pytest.raises( + TypeError, match="`types` must be a list of strings" + ), ), ( - ["session"], - { - "T1w": {"pattern": "{subject}/anat/{subject}_T1w.nii.gz"}, - "BOLD": { - "pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz" - }, - }, - pytest.raises(ValueError, match="is not part of"), - ), - ( - ["subject", "session"], - { - "T1w": {"pattern": "{subject}/anat/_T1w.nii.gz"}, - "BOLD": {"pattern": "{session}/func/_task-rest_bold.nii.gz"}, - }, - pytest.raises(ValueError, match="At least one pattern"), - ), - ( - ["subject"], - { - "T1w": {"pattern": "{subject}/anat/{subject}_T1w.nii.gz"}, - "BOLD": { - "pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz" - }, - }, - nullcontext(), - ), - ], -) -def test_validate_replacements( - replacements: Union[str, List[str], List[int]], - patterns: Union[str, Dict[str, Dict[str, str]]], - expect: ContextManager, -) -> None: - """Test validation of replacements. - - Parameters - ---------- - replacements : str, list of str or int - The parametrized pattern replacements to validate. - patterns : str, dict - The parametrized patterns to validate against. - expect : typing.ContextManager - The parametrized ContextManager object. - - """ - with expect: - validate_replacements(replacements=replacements, patterns=patterns) # type: ignore - - -@pytest.mark.parametrize( - "types, patterns, expect", - [ - ( - ["T1w", "BOLD"], + ["BOLD"], + [], "wrong", - pytest.raises(TypeError, match="must be a dict"), + pytest.raises(TypeError, match="`patterns` must be a dict"), ), ( ["T1w", "BOLD"], + "", { "T1w": {"pattern": "{subject}/anat/{subject}_T1w.nii.gz"}, }, pytest.raises( ValueError, - match="Length of `types` more than that of `patterns`.", + match="Length of `types` more than that of `patterns`", ), ), ( ["T1w", "BOLD"], + "", { "T1w": {"pattern": "{subject}/anat/{subject}_T1w.nii.gz"}, "T2w": {"pattern": "{subject}/anat/{subject}_T2w.nii.gz"}, }, - pytest.raises(ValueError, match="contain all"), + pytest.raises( + ValueError, match="`patterns` must contain all `types`" + ), ), ( ["T3w"], + "", { "T3w": {"pattern": "{subject}/anat/{subject}_T3w.nii.gz"}, }, @@ -145,6 +67,7 @@ def test_validate_replacements( ), ( ["BOLD"], + "", { "BOLD": {"patterns": "{subject}/func/{subject}_BOLD.nii.gz"}, }, @@ -152,6 +75,7 @@ def test_validate_replacements( ), ( ["BOLD"], + "", { "BOLD": { "pattern": ( @@ -169,6 +93,7 @@ def test_validate_replacements( ), ( ["T1w"], + "", { "T1w": { "pattern": "{subject}/anat/{subject}*.nii", @@ -177,8 +102,65 @@ def test_validate_replacements( }, pytest.raises(ValueError, match="following a replacement"), ), + ( + ["T1w"], + "wrong", + { + "T1w": { + "pattern": "{subject}/anat/{subject}_T1w.nii", + "space": "native", + }, + }, + pytest.raises(TypeError, match="`replacements` must be a list"), + ), + ( + ["T1w"], + [1], + { + "T1w": { + "pattern": "{subject}/anat/{subject}_T1w.nii", + "space": "native", + }, + }, + pytest.raises( + TypeError, match="`replacements` must be a list of strings" + ), + ), + ( + ["T1w", "BOLD"], + ["subject", "session"], + { + "T1w": { + "pattern": "{subject}/anat/{subject}_T1w.nii.gz", + "space": "native", + }, + "BOLD": { + "pattern": ( + "{subject}/func/{subject}_task-rest_bold.nii.gz" + ), + "space": "MNI152NLin6Asym", + }, + }, + pytest.raises(ValueError, match="is not part of any pattern"), + ), + ( + ["BOLD"], + ["subject", "session"], + { + "T1w": { + "pattern": "{subject}/anat/_T1w.nii.gz", + "space": "native", + }, + "BOLD": { + "pattern": "{session}/func/_task-rest_bold.nii.gz", + "space": "MNI152NLin6Asym", + }, + }, + pytest.raises(ValueError, match="At least one pattern"), + ), ( ["T1w", "T2w", "BOLD"], + ["subject"], { "T1w": { "pattern": "{subject}/anat/{subject}_T1w.nii.gz", @@ -190,7 +172,7 @@ def test_validate_replacements( }, "BOLD": { "pattern": ( - "{subject}/func/{subject}_task-rest_bold.nii.gz" + "{subject}/func/{session}/{subject}_task-rest_bold.nii.gz" ), "space": "MNI152NLin6Asym", "confounds": { @@ -203,22 +185,65 @@ def test_validate_replacements( ), ], ) -def test_validate_patterns( - types: List[str], +def test_PatternValidationMixin( + types: Union[str, List[str], List[int]], + replacements: Union[str, List[str], List[int]], patterns: Union[str, Dict[str, Dict[str, str]]], expect: ContextManager, ) -> None: - """Test validation of patterns. + """Test validation. Parameters ---------- - types : list of str - The parametrized data types. + types : str, list of int or str + The parametrized data types to validate. + replacements : str, list of str or int + The parametrized pattern replacements to validate. patterns : str, dict - The patterns to validate. + The parametrized patterns to validate against. expect : typing.ContextManager The parametrized ContextManager object. """ + + class MockDataGrabber(PatternValidationMixin): + def __init__( + self, + types, + replacements, + patterns, + ) -> None: + self.types = types + self.replacements = replacements + self.patterns = patterns + + def validate(self) -> None: + self.validate_patterns( + types=self.types, + replacements=self.replacements, + patterns=self.patterns, + ) + + dg = MockDataGrabber(types, replacements, patterns) with expect: - validate_patterns(types=types, patterns=patterns) # type: ignore + dg.validate() + + +# This test is kept separate as bool doesn't support context manager protocol, +# used in the earlier test +def test_PatternValidationMixin_partial_pattern_check() -> None: + """Test validation for partial patterns.""" + with pytest.warns(RuntimeWarning, match="might not work as expected"): + PatternValidationMixin().validate_patterns( + types=["BOLD"], + replacements=["subject"], + patterns={ + "BOLD": { + "mask": { + "pattern": "{subject}/func/{subject}_BOLD.nii.gz", + "space": "MNI152NLin6Asym", + }, + }, + }, # type: ignore + partial_pattern_ok=True, + ) diff --git a/junifer/datagrabber/utils.py b/junifer/datagrabber/utils.py deleted file mode 100644 index 5caec4e85..000000000 --- a/junifer/datagrabber/utils.py +++ /dev/null @@ -1,317 +0,0 @@ -"""Provide utility functions for the datagrabber sub-package.""" - -# Authors: Federico Raimondo -# Synchon Mandal -# License: AGPL - -from typing import Dict, List - -from ..utils import logger, raise_error - - -__all__ = ["validate_types", "validate_replacements", "validate_patterns"] - - -# Define schema for pattern-based datagrabber's patterns -PATTERNS_SCHEMA = { - "T1w": { - "mandatory": ["pattern", "space"], - "optional": { - "mask": {"mandatory": ["pattern", "space"], "optional": []}, - }, - }, - "T2w": { - "mandatory": ["pattern", "space"], - "optional": { - "mask": {"mandatory": ["pattern", "space"], "optional": []}, - }, - }, - "BOLD": { - "mandatory": ["pattern", "space"], - "optional": { - "mask": {"mandatory": ["pattern", "space"], "optional": []}, - "confounds": { - "mandatory": ["pattern", "format"], - "optional": ["mappings"], - }, - }, - }, - "Warp": { - "mandatory": ["pattern", "src", "dst"], - "optional": {}, - }, - "VBM_GM": { - "mandatory": ["pattern", "space"], - "optional": {}, - }, - "VBM_WM": { - "mandatory": ["pattern", "space"], - "optional": {}, - }, - "VBM_CSF": { - "mandatory": ["pattern", "space"], - "optional": {}, - }, - "DWI": { - "mandatory": ["pattern"], - "optional": {}, - }, - "FreeSurfer": { - "mandatory": ["pattern"], - "optional": { - "aseg": {"mandatory": ["pattern"], "optional": []}, - "norm": {"mandatory": ["pattern"], "optional": []}, - "lh_white": {"mandatory": ["pattern"], "optional": []}, - "rh_white": {"mandatory": ["pattern"], "optional": []}, - "lh_pial": {"mandatory": ["pattern"], "optional": []}, - "rh_pial": {"mandatory": ["pattern"], "optional": []}, - }, - }, -} - - -def validate_types(types: List[str]) -> None: - """Validate the types. - - Parameters - ---------- - types : list of str - The object to validate. - - Raises - ------ - TypeError - If ``types`` is not a list or if the values are not string. - - """ - if not isinstance(types, list): - raise_error(msg="`types` must be a list", klass=TypeError) - if any(not isinstance(x, str) for x in types): - raise_error(msg="`types` must be a list of strings", klass=TypeError) - - -def validate_replacements( - replacements: List[str], patterns: Dict[str, Dict[str, str]] -) -> None: - """Validate the replacements. - - Parameters - ---------- - replacements : list of str - The object to validate. - patterns : dict - The patterns to validate against. - - Raises - ------ - TypeError - If ``replacements`` is not a list or if the values are not string. - ValueError - If a value in ``replacements`` is not part of a data type pattern or - if no data type patterns contain all values in ``replacements``. - - """ - if not isinstance(replacements, list): - raise_error(msg="`replacements` must be a list.", klass=TypeError) - - if any(not isinstance(x, str) for x in replacements): - raise_error( - msg="`replacements` must be a list of strings.", klass=TypeError - ) - - for x in replacements: - if all( - x not in y - for y in [ - data_type_val["pattern"] for data_type_val in patterns.values() - ] - ): - raise_error(msg=f"Replacement: {x} is not part of any pattern.") - - # Check that at least one pattern has all the replacements - at_least_one = False - for data_type_val in patterns.values(): - if all(x in data_type_val["pattern"] for x in replacements): - at_least_one = True - if at_least_one is False: - raise_error(msg="At least one pattern must contain all replacements.") - - -def _validate_mandatory_keys( - keys: List[str], schema: List[str], data_type: str -) -> None: - """Validate mandatory keys. - - Parameters - ---------- - keys : list of str - The keys to validate. - schema : list of str - The schema to validate against. - data_type : str - The data type being validated. - - Raises - ------ - KeyError - If any mandatory key is missing for a data type. - - """ - for key in schema: - if key not in keys: - raise_error( - msg=f"Mandatory key: `{key}` missing for {data_type}", - klass=KeyError, - ) - else: - logger.debug(f"Mandatory key: `{key}` found for {data_type}") - - -def _identify_stray_keys( - keys: List[str], schema: List[str], data_type: str -) -> None: - """Identify stray keys. - - Parameters - ---------- - keys : list of str - The keys to check. - schema : list of str - The schema to check against. - data_type : str - The data type being checked. - - Raises - ------ - RuntimeError - If an unknown key is found for a data type. - - """ - for key in keys: - if key not in schema: - raise_error( - msg=( - f"Key: {key} not accepted for {data_type} " - "pattern, remove it to proceed" - ), - klass=RuntimeError, - ) - - -def validate_patterns( - types: List[str], patterns: Dict[str, Dict[str, str]] -) -> None: - """Validate the patterns. - - Parameters - ---------- - types : list of str - The types list. - patterns : dict - The object to validate. - - Raises - ------ - TypeError - If ``patterns`` is not a dictionary. - ValueError - If length of ``types`` and ``patterns`` are different or - if ``patterns`` is missing entries from ``types`` or - if unknown data type is found in ``patterns`` or - if data type pattern key contains '*' as value. - - """ - # Validate the types - validate_types(types) - if not isinstance(patterns, dict): - raise_error(msg="`patterns` must be a dict.", klass=TypeError) - # Unequal length of objects - if len(types) > len(patterns): - raise_error( - msg="Length of `types` more than that of `patterns`.", - klass=ValueError, - ) - # Missing type in patterns - if any(x not in patterns for x in types): - raise_error( - msg="`patterns` must contain all `types`", klass=ValueError - ) - # Check against schema - for data_type_key, data_type_val in patterns.items(): - # Check if valid data type is provided - if data_type_key not in PATTERNS_SCHEMA: - raise_error( - f"Unknown data type: {data_type_key}, " - f"should be one of: {list(PATTERNS_SCHEMA.keys())}" - ) - # Check mandatory keys for data type - _validate_mandatory_keys( - keys=list(data_type_val), - schema=PATTERNS_SCHEMA[data_type_key]["mandatory"], - data_type=data_type_key, - ) - # Check optional keys for data type - for optional_key, optional_val in PATTERNS_SCHEMA[data_type_key][ - "optional" - ].items(): - if optional_key not in data_type_val: - logger.debug( - f"Optional key: `{optional_key}` missing for " - f"{data_type_key}" - ) - else: - logger.debug( - f"Optional key: `{optional_key}` found for " - f"{data_type_key}" - ) - # Set nested type name for easier access - nested_data_type = f"{data_type_key}.{optional_key}" - nested_mandatory_keys_schema = PATTERNS_SCHEMA[data_type_key][ - "optional" - ][optional_key]["mandatory"] - nested_optional_keys_schema = PATTERNS_SCHEMA[data_type_key][ - "optional" - ][optional_key]["optional"] - # Check mandatory keys for nested type - _validate_mandatory_keys( - keys=list(optional_val["mandatory"]), - schema=nested_mandatory_keys_schema, - data_type=nested_data_type, - ) - # Check optional keys for nested type - for nested_optional_key in nested_optional_keys_schema: - if nested_optional_key not in optional_val["optional"]: - logger.debug( - f"Optional key: `{nested_optional_key}` missing " - f"for {nested_data_type}" - ) - else: - logger.debug( - f"Optional key: `{nested_optional_key}` found for " - f"{nested_data_type}" - ) - # Check stray key for nested data type - _identify_stray_keys( - keys=optional_val["mandatory"] + optional_val["optional"], - schema=nested_mandatory_keys_schema - + nested_optional_keys_schema, - data_type=nested_data_type, - ) - # Check stray key for data type - _identify_stray_keys( - keys=list(data_type_val.keys()), - schema=( - PATTERNS_SCHEMA[data_type_key]["mandatory"] - + list(PATTERNS_SCHEMA[data_type_key]["optional"].keys()) - ), - data_type=data_type_key, - ) - # Wildcard check in patterns - if "}*" in data_type_val["pattern"]: - raise_error( - msg=( - f"`{data_type_key}.pattern` must not contain `*` " - "following a replacement" - ), - klass=ValueError, - ) -- 2.52.0 From 3bce11b1e97ea34f1d6516d4c8c1ce8ca20b5a97 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 11 Jul 2024 12:30:05 +0200 Subject: [PATCH 03/20] update: adapt MultipleDataGrabber to respect nested data types --- junifer/datagrabber/multiple.py | 41 ++++- junifer/datagrabber/tests/test_multiple.py | 194 ++++++++++++--------- 2 files changed, 148 insertions(+), 87 deletions(-) diff --git a/junifer/datagrabber/multiple.py b/junifer/datagrabber/multiple.py index c79655661..6efeeaa2b 100644 --- a/junifer/datagrabber/multiple.py +++ b/junifer/datagrabber/multiple.py @@ -29,19 +29,52 @@ class MultipleDataGrabber(BaseDataGrabber): **kwargs Keyword arguments passed to superclass. + Raises + ------ + RuntimeError + If ``datagrabbers`` have different element keys or + overlapping data types or nested data types. + """ def __init__(self, datagrabbers: List[BaseDataGrabber], **kwargs) -> None: # Check datagrabbers consistency - # 1) same element keys + # Check for same element keys first_keys = datagrabbers[0].get_element_keys() for dg in datagrabbers[1:]: if dg.get_element_keys() != first_keys: - raise_error("DataGrabbers have different element keys.") - # 2) no overlapping types + raise_error( + msg="DataGrabbers have different element keys", + klass=RuntimeError, + ) + # Check for no overlapping types (and nested data types) types = [x for dg in datagrabbers for x in dg.get_types()] if len(types) != len(set(types)): - raise_error("DataGrabbers have overlapping types.") + if all(hasattr(dg, "patterns") for dg in datagrabbers): + first_patterns = datagrabbers[0].patterns + for dg in datagrabbers[1:]: + for data_type in set(types): + patterns = dg.patterns + dtype_pattern = patterns[data_type] + # Check if first-level keys of data type are same + if ( + dtype_pattern.keys() + == first_patterns[data_type].keys() + ): + raise_error( + msg=( + "DataGrabbers have overlapping mandatory " + "and / or optional key(s) for data type: " + f"`{data_type}`" + ), + klass=RuntimeError, + ) + else: + # Can't check further + raise_error( + msg="DataGrabbers have overlapping types", + klass=RuntimeError, + ) self._datagrabbers = datagrabbers def __getitem__(self, element: Union[str, Tuple]) -> Dict: diff --git a/junifer/datagrabber/tests/test_multiple.py b/junifer/datagrabber/tests/test_multiple.py index 4ec7155e0..cd7e27326 100644 --- a/junifer/datagrabber/tests/test_multiple.py +++ b/junifer/datagrabber/tests/test_multiple.py @@ -25,28 +25,19 @@ def test_MultipleDataGrabber() -> None: repo_uri = _testing_dataset["example_bids_ses"]["uri"] rootdir = "example_bids_ses" replacements = ["subject", "session"] - pattern1 = { - "T1w": { - "pattern": ( - "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz" - ), - "space": "native", - }, - } - pattern2 = { - "BOLD": { - "pattern": ( - "{subject}/{session}/func/" - "{subject}_{session}_task-rest_bold.nii.gz" - ), - "space": "MNI152NLin6Asym", - }, - } + dg1 = PatternDataladDataGrabber( rootdir=rootdir, uri=repo_uri, types=["T1w"], - patterns=pattern1, + patterns={ + "T1w": { + "pattern": ( + "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz" + ), + "space": "native", + }, + }, replacements=replacements, ) @@ -54,7 +45,15 @@ def test_MultipleDataGrabber() -> None: rootdir=rootdir, uri=repo_uri, types=["BOLD"], - patterns=pattern2, + patterns={ + "BOLD": { + "pattern": ( + "{subject}/{session}/func/" + "{subject}_{session}_task-rest_bold.nii.gz" + ), + "space": "MNI152NLin6Asym", + }, + }, replacements=replacements, ) @@ -89,40 +88,37 @@ def test_MultipleDataGrabber() -> None: def test_MultipleDataGrabber_no_intersection() -> None: """Test MultipleDataGrabber without intersection (0 elements).""" - repo_uri1 = _testing_dataset["example_bids"]["uri"] - repo_uri2 = _testing_dataset["example_bids_ses"]["uri"] rootdir = "example_bids_ses" replacements = ["subject", "session"] - pattern1 = { - "T1w": { - "pattern": ( - "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz" - ), - "space": "native", - }, - } - pattern2 = { - "BOLD": { - "pattern": ( - "{subject}/{session}/func/" - "{subject}_{session}_task-rest_bold.nii.gz" - ), - "space": "MNI152NLin6Asym", - }, - } + dg1 = PatternDataladDataGrabber( rootdir=rootdir, - uri=repo_uri1, + uri=_testing_dataset["example_bids"]["uri"], types=["T1w"], - patterns=pattern1, + patterns={ + "T1w": { + "pattern": ( + "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz" + ), + "space": "native", + }, + }, replacements=replacements, ) dg2 = PatternDataladDataGrabber( rootdir=rootdir, - uri=repo_uri2, + uri=_testing_dataset["example_bids_ses"]["uri"], types=["BOLD"], - patterns=pattern2, + patterns={ + "BOLD": { + "pattern": ( + "{subject}/{session}/func/" + "{subject}_{session}_task-rest_bold.nii.gz" + ), + "space": "MNI152NLin6Asym", + }, + }, replacements=replacements, ) @@ -135,23 +131,19 @@ def test_MultipleDataGrabber_no_intersection() -> None: def test_MultipleDataGrabber_get_item() -> None: """Test MultipleDataGrabber get_item() error.""" - repo_uri1 = _testing_dataset["example_bids"]["uri"] - rootdir = "example_bids_ses" - replacements = ["subject", "session"] - pattern1 = { - "T1w": { - "pattern": ( - "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz" - ), - "space": "native", - }, - } dg1 = PatternDataladDataGrabber( - rootdir=rootdir, - uri=repo_uri1, + rootdir="example_bids_ses", + uri=_testing_dataset["example_bids"]["uri"], types=["T1w"], - patterns=pattern1, - replacements=replacements, + patterns={ + "T1w": { + "pattern": ( + "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz" + ), + "space": "native", + }, + }, + replacements=["subject", "session"], ) dg = MultipleDataGrabber([dg1]) @@ -161,43 +153,79 @@ def test_MultipleDataGrabber_get_item() -> None: def test_MultipleDataGrabber_validation() -> None: """Test MultipleDataGrabber init validation.""" - repo_uri1 = _testing_dataset["example_bids"]["uri"] - repo_uri2 = _testing_dataset["example_bids_ses"]["uri"] rootdir = "example_bids_ses" - replacement1 = ["subject", "session"] - replacement2 = ["subject"] - pattern1 = { - "T1w": { - "pattern": ( - "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz" - ), - "space": "native", - }, - } - pattern2 = { - "BOLD": { - "pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz", - "space": "MNI152NLin6Asym", - }, - } + dg1 = PatternDataladDataGrabber( rootdir=rootdir, - uri=repo_uri1, + uri=_testing_dataset["example_bids"]["uri"], types=["T1w"], - patterns=pattern1, - replacements=replacement1, + patterns={ + "T1w": { + "pattern": ( + "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz" + ), + "space": "native", + }, + }, + replacements=["subject", "session"], ) dg2 = PatternDataladDataGrabber( rootdir=rootdir, - uri=repo_uri2, + uri=_testing_dataset["example_bids_ses"]["uri"], types=["BOLD"], - patterns=pattern2, - replacements=replacement2, + patterns={ + "BOLD": { + "pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz", + "space": "MNI152NLin6Asym", + }, + }, + replacements=["subject"], ) - with pytest.raises(ValueError, match="different element key"): + with pytest.raises(RuntimeError, match="have different element keys"): MultipleDataGrabber([dg1, dg2]) - with pytest.raises(ValueError, match="overlapping types"): + with pytest.raises(RuntimeError, match="have overlapping mandatory"): MultipleDataGrabber([dg1, dg1]) + + +def test_MultipleDataGrabber_partial_pattern() -> None: + """Test MultipleDataGrabber partial pattern.""" + dg1 = PatternDataladDataGrabber( + rootdir=".", + uri="data://uri1", + types=["BOLD"], + patterns={ + "BOLD": { + "pattern": ( + "{subject}/{session}/func/" + "{subject}_{session}_task-rest_bold.nii.gz" + ), + "space": "MNI152NLin6Asym", + }, + }, + replacements=["subject", "session"], + ) + + dg2 = PatternDataladDataGrabber( + rootdir=".", + uri="data://uri2", + types=["BOLD"], + patterns={ + "BOLD": { + "confounds": { + "pattern": ( + "{subject}/{session}/func/" + "{subject}_{session}_confounds.tsv" + ), + "format": "fmriprep", + }, + }, + }, + replacements=["subject", "session"], + partial_pattern_ok=True, + ) + + # Test validation works + MultipleDataGrabber([dg1, dg2]) -- 2.52.0 From a1472547fe230982663a6458bbc98746d8709656 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 11 Jul 2024 12:56:47 +0200 Subject: [PATCH 04/20] update: adapt BaseDataGrabber types validation --- junifer/datagrabber/base.py | 16 ++++++++++------ 1 file changed, 10 insertions(+), 6 deletions(-) diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index e9b4dfbc8..9e6d0246f 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -11,7 +11,6 @@ from typing import Dict, Iterator, List, Tuple, Union from ..pipeline import UpdateMetaMixin from ..utils import logger, raise_error -from .utils import validate_types __all__ = ["BaseDataGrabber"] @@ -30,16 +29,21 @@ class BaseDataGrabber(ABC, UpdateMetaMixin): datadir : str or pathlib.Path The directory where the data is / will be stored. - Attributes - ---------- - datadir : pathlib.Path - The directory where the data is / will be stored. + Raises + ------ + TypeError + If ``types`` is not a list or if the values are not string. """ def __init__(self, types: List[str], datadir: Union[str, Path]) -> None: # Validate types - validate_types(types) + if not isinstance(types, list): + raise_error(msg="`types` must be a list", klass=TypeError) + if any(not isinstance(x, str) for x in types): + raise_error( + msg="`types` must be a list of strings", klass=TypeError + ) self.types = types # Convert str to Path -- 2.52.0 From 28de66e759acb1dc5c62cc95e88bf921708f2fae Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 11 Jul 2024 12:58:50 +0200 Subject: [PATCH 05/20] fix: correct raise_error() import in HCP1200 --- junifer/datagrabber/hcp1200/hcp1200.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/datagrabber/hcp1200/hcp1200.py b/junifer/datagrabber/hcp1200/hcp1200.py index bee9405ce..0fa8a8018 100644 --- a/junifer/datagrabber/hcp1200/hcp1200.py +++ b/junifer/datagrabber/hcp1200/hcp1200.py @@ -10,8 +10,8 @@ from pathlib import Path from typing import Dict, List, Union from ...api.decorators import register_datagrabber +from ...utils import raise_error from ..pattern import PatternDataGrabber -from ..utils import raise_error __all__ = ["HCP1200"] -- 2.52.0 From 17355fd020903f05707ba7c9698d528fc9fd0b9f Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 11 Jul 2024 14:02:53 +0200 Subject: [PATCH 06/20] fix: correct pattern validation for MultipleDataGrabber --- junifer/datagrabber/multiple.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/junifer/datagrabber/multiple.py b/junifer/datagrabber/multiple.py index 6efeeaa2b..da3ac54c7 100644 --- a/junifer/datagrabber/multiple.py +++ b/junifer/datagrabber/multiple.py @@ -54,8 +54,9 @@ class MultipleDataGrabber(BaseDataGrabber): first_patterns = datagrabbers[0].patterns for dg in datagrabbers[1:]: for data_type in set(types): - patterns = dg.patterns - dtype_pattern = patterns[data_type] + dtype_pattern = dg.patterns.get(data_type) + if dtype_pattern is None: + continue # Check if first-level keys of data type are same if ( dtype_pattern.keys() -- 2.52.0 From 581ce411deeb8b7861a75c14b264e30545431243 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 12 Jul 2024 08:07:22 +0200 Subject: [PATCH 07/20] chore: add missing import for PatternValidationMixin in datagrabber.__init__ --- junifer/datagrabber/__init__.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/junifer/datagrabber/__init__.py b/junifer/datagrabber/__init__.py index fb5e276ab..1ee601f56 100644 --- a/junifer/datagrabber/__init__.py +++ b/junifer/datagrabber/__init__.py @@ -17,6 +17,7 @@ from .hcp1200 import HCP1200, DataladHCP1200 from .multiple import MultipleDataGrabber from .dmcc13_benchmark import DMCC13Benchmark +from .pattern_validation_mixin import PatternValidationMixin __all__ = [ "BaseDataGrabber", @@ -30,4 +31,5 @@ __all__ = [ "DataladHCP1200", "MultipleDataGrabber", "DMCC13Benchmark", + "PatternValidationMixin", ] -- 2.52.0 From bc6d086abc3127775cb55af4f71773e45ab71fed Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 12 Jul 2024 08:08:03 +0200 Subject: [PATCH 08/20] chore: add changelogs 351.{change,enh,feature} --- docs/changes/newsfragments/351.change | 1 + docs/changes/newsfragments/351.enh | 1 + docs/changes/newsfragments/351.feature | 1 + 3 files changed, 3 insertions(+) create mode 100644 docs/changes/newsfragments/351.change create mode 100644 docs/changes/newsfragments/351.enh create mode 100644 docs/changes/newsfragments/351.feature diff --git a/docs/changes/newsfragments/351.change b/docs/changes/newsfragments/351.change new file mode 100644 index 000000000..1f1539bd9 --- /dev/null +++ b/docs/changes/newsfragments/351.change @@ -0,0 +1 @@ +Add ``partial_pattern_ok`` argument to :class:`.PatternDataGrabber` to not raise error on missing mandatory key checks for data types by `Synchon Mandal`_ diff --git a/docs/changes/newsfragments/351.enh b/docs/changes/newsfragments/351.enh new file mode 100644 index 000000000..462dca670 --- /dev/null +++ b/docs/changes/newsfragments/351.enh @@ -0,0 +1 @@ +Adapt :class:`.MultipleDataGrabber` to handle "nested types" introduced in :gh:`341` by `Synchon Mandal`_ diff --git a/docs/changes/newsfragments/351.feature b/docs/changes/newsfragments/351.feature new file mode 100644 index 000000000..1826bccbe --- /dev/null +++ b/docs/changes/newsfragments/351.feature @@ -0,0 +1 @@ +Introduce :class:`.PatternValidationMixin` to simplify validation for pattern-based DataGrabbers by `Synchon Mandal`_ -- 2.52.0 From c15090d38546ccd7554627d6bcaeca8db71f0d1b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Jul 2024 12:44:08 +0200 Subject: [PATCH 09/20] update: allow pattern retrieval for elements fetch of PatternDataGrabber with partial pattern --- junifer/datagrabber/pattern.py | 17 +++++++++++++++-- 1 file changed, 15 insertions(+), 2 deletions(-) diff --git a/junifer/datagrabber/pattern.py b/junifer/datagrabber/pattern.py index 6d8b39f1f..9790e7dba 100644 --- a/junifer/datagrabber/pattern.py +++ b/junifer/datagrabber/pattern.py @@ -178,6 +178,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin): ) self.replacements = replacements self.patterns = patterns + self.partial_pattern_ok = partial_pattern_ok # Validate confounds format if ( @@ -446,14 +447,26 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin): for t_idx in reversed(order): t_type = self.types[t_idx] types_element = set() - # Get the pattern + + # Get the pattern dict t_pattern = self.patterns[t_type] + # Conditional fetch of base pattern for getting elements + pattern = None + # Try for data type pattern + pattern = t_pattern.get("pattern") + # Try for nested data type pattern + if pattern is None and self.partial_pattern_ok: + for v in t_pattern.values(): + if isinstance(v, dict) and "pattern" in v: + pattern = v["pattern"] + break + # Replace the pattern ( re_pattern, glob_pattern, t_replacements, - ) = self._replace_patterns_regex(t_pattern["pattern"]) + ) = self._replace_patterns_regex(pattern) for fname in self.datadir.glob(glob_pattern): suffix = fname.relative_to(self.datadir).as_posix() m = re.match(re_pattern, suffix) -- 2.52.0 From f5e28ecb12625ae0db206517a77d770e04ede7bf Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Jul 2024 12:45:56 +0200 Subject: [PATCH 10/20] feat: add deep_update() helper for dict update in varying width and depth --- junifer/utils/__init__.py | 3 ++- junifer/utils/helpers.py | 32 ++++++++++++++++++++++++++++++-- 2 files changed, 32 insertions(+), 3 deletions(-) diff --git a/junifer/utils/__init__.py b/junifer/utils/__init__.py index c11541231..084379611 100644 --- a/junifer/utils/__init__.py +++ b/junifer/utils/__init__.py @@ -6,7 +6,7 @@ from .fs import make_executable from .logging import configure_logging, logger, raise_error, warn_with_log -from .helpers import run_ext_cmd +from .helpers import run_ext_cmd, deep_update __all__ = [ @@ -16,4 +16,5 @@ __all__ = [ "raise_error", "warn_with_log", "run_ext_cmd", + "deep_update", ] diff --git a/junifer/utils/helpers.py b/junifer/utils/helpers.py index 3946ee000..352d357ce 100644 --- a/junifer/utils/helpers.py +++ b/junifer/utils/helpers.py @@ -3,13 +3,14 @@ # Authors: Synchon Mandal # License: AGPL +import collections.abc import subprocess -from typing import List +from typing import Dict, List from .logging import logger, raise_error -__all__ = ["run_ext_cmd"] +__all__ = ["run_ext_cmd", "deep_update"] def run_ext_cmd(name: str, cmd: List[str]) -> None: @@ -54,3 +55,30 @@ def run_ext_cmd(name: str, cmd: List[str]) -> None: ), klass=RuntimeError, ) + + +def deep_update(d: Dict, u: Dict) -> Dict: + """Deep update `d` with `u`. + + From: "https://stackoverflow.com/questions/3232943/update-value-of-a-nested + -dictionary-of-varying-depth" + + Parameters + ---------- + d : dict + The dictionary to deep-update. + u : dict + The dictionary to deep-update `d` with. + + Returns + ------- + dict + The updated dictionary. + + """ + for k, v in u.items(): + if isinstance(v, collections.abc.Mapping): + d[k] = deep_update(d.get(k, {}), v) + else: + d[k] = v + return d -- 2.52.0 From 52c8bf71ed92a12e92a17bb2dee9a88da83e9509 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Jul 2024 12:48:32 +0200 Subject: [PATCH 11/20] chore: update 351.feature --- docs/changes/newsfragments/351.feature | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/changes/newsfragments/351.feature b/docs/changes/newsfragments/351.feature index 1826bccbe..939e87941 100644 --- a/docs/changes/newsfragments/351.feature +++ b/docs/changes/newsfragments/351.feature @@ -1 +1 @@ -Introduce :class:`.PatternValidationMixin` to simplify validation for pattern-based DataGrabbers by `Synchon Mandal`_ +Introduce :class:`.PatternValidationMixin` to simplify validation for pattern-based DataGrabbers and :func:`.deep_update` for updating dictionary with varying width and depth by `Synchon Mandal`_ -- 2.52.0 From 71c372715dd4e22e4c52a4073ecffc407629d02a Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Jul 2024 12:49:03 +0200 Subject: [PATCH 12/20] update: use deep_update() in MultipleDataGrabber --- junifer/datagrabber/multiple.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/junifer/datagrabber/multiple.py b/junifer/datagrabber/multiple.py index da3ac54c7..7370dc234 100644 --- a/junifer/datagrabber/multiple.py +++ b/junifer/datagrabber/multiple.py @@ -8,7 +8,7 @@ from typing import Dict, List, Tuple, Union from ..api.decorators import register_datagrabber -from ..utils import raise_error +from ..utils import deep_update, raise_error from .base import BaseDataGrabber @@ -101,7 +101,7 @@ class MultipleDataGrabber(BaseDataGrabber): metas = [] for dg in self._datagrabbers: t_out = dg[element] - out.update(t_out) + deep_update(out, t_out) # Now get the meta for this datagrabber t_meta = {} dg.update_meta(t_meta, "datagrabber") -- 2.52.0 From e11846b6e5b3dd7bc7bc327dada459065d3c893a Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Jul 2024 12:49:37 +0200 Subject: [PATCH 13/20] update: improve tests for MultipleDataGrabber --- junifer/datagrabber/tests/test_multiple.py | 67 +++++++++++++++++++--- 1 file changed, 58 insertions(+), 9 deletions(-) diff --git a/junifer/datagrabber/tests/test_multiple.py b/junifer/datagrabber/tests/test_multiple.py index cd7e27326..24e7eb301 100644 --- a/junifer/datagrabber/tests/test_multiple.py +++ b/junifer/datagrabber/tests/test_multiple.py @@ -36,6 +36,13 @@ def test_MultipleDataGrabber() -> None: "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz" ), "space": "native", + "mask": { + "pattern": ( + "{subject}/{session}/anat/{subject}_{session}_" + "brain_mask.nii.gz" + ), + "space": "native", + }, }, }, replacements=replacements, @@ -52,6 +59,13 @@ def test_MultipleDataGrabber() -> None: "{subject}_{session}_task-rest_bold.nii.gz" ), "space": "MNI152NLin6Asym", + "mask": { + "pattern": ( + "{subject}/{session}/func/" + "{subject}_{session}_task-rest_brain_mask.nii.gz" + ), + "space": "MNI152NLin6Asym", + }, }, }, replacements=replacements, @@ -72,14 +86,17 @@ def test_MultipleDataGrabber() -> None: with dg: subs = list(dg) assert set(subs) == set(expected_subs) - + # Check data type elem = dg[("sub-01", "ses-01")] + # Check data types assert "T1w" in elem assert "BOLD" in elem + # Check meta assert "meta" in elem["BOLD"] meta = elem["BOLD"]["meta"]["datagrabber"] assert "class" in meta assert meta["class"] == "MultipleDataGrabber" + # Check datagrabbers assert "datagrabbers" in meta assert len(meta["datagrabbers"]) == 2 assert meta["datagrabbers"][0]["class"] == "PatternDataladDataGrabber" @@ -192,9 +209,13 @@ def test_MultipleDataGrabber_validation() -> None: def test_MultipleDataGrabber_partial_pattern() -> None: """Test MultipleDataGrabber partial pattern.""" + repo_uri = _testing_dataset["example_bids_ses"]["uri"] + rootdir = "example_bids_ses" + replacements = ["subject", "session"] + dg1 = PatternDataladDataGrabber( - rootdir=".", - uri="data://uri1", + rootdir=rootdir, + uri=repo_uri, types=["BOLD"], patterns={ "BOLD": { @@ -205,19 +226,20 @@ def test_MultipleDataGrabber_partial_pattern() -> None: "space": "MNI152NLin6Asym", }, }, - replacements=["subject", "session"], + replacements=replacements, ) dg2 = PatternDataladDataGrabber( - rootdir=".", - uri="data://uri2", + rootdir=rootdir, + uri=repo_uri, types=["BOLD"], patterns={ "BOLD": { "confounds": { "pattern": ( "{subject}/{session}/func/" - "{subject}_{session}_confounds.tsv" + "{subject}_{session}_task-rest_" + "confounds_regressors.tsv" ), "format": "fmriprep", }, @@ -227,5 +249,32 @@ def test_MultipleDataGrabber_partial_pattern() -> None: partial_pattern_ok=True, ) - # Test validation works - MultipleDataGrabber([dg1, dg2]) + dg = MultipleDataGrabber([dg1, dg2]) + + types = dg.get_types() + assert "BOLD" in types + + expected_subs = [ + (f"sub-{i:02d}", f"ses-{j:02d}") + for j in range(1, 3) + for i in range(1, 10) + ] + + with dg: + subs = list(dg) + assert set(subs) == set(expected_subs) + # Fetch element + elem = dg[("sub-01", "ses-01")] + # Check data type and nested data type + assert "BOLD" in elem + assert "confounds" in elem["BOLD"] + # Check meta + assert "meta" in elem["BOLD"] + meta = elem["BOLD"]["meta"]["datagrabber"] + assert "class" in meta + assert meta["class"] == "MultipleDataGrabber" + # Check datagrabbers + assert "datagrabbers" in meta + assert len(meta["datagrabbers"]) == 2 + assert meta["datagrabbers"][0]["class"] == "PatternDataladDataGrabber" + assert meta["datagrabbers"][1]["class"] == "PatternDataladDataGrabber" -- 2.52.0 From dc5ca201eb438919379a16acb3a7e3e8ddfeeed1 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Jul 2024 12:49:59 +0200 Subject: [PATCH 14/20] chore: update scripts for example bids datasets in gin --- tools/create_bids_example_dataset.py | 30 +++++++++------ tools/create_bids_example_dataset_sessions.py | 37 ++++++++++++------- 2 files changed, 41 insertions(+), 26 deletions(-) diff --git a/tools/create_bids_example_dataset.py b/tools/create_bids_example_dataset.py index fc665ca05..7b755eccb 100644 --- a/tools/create_bids_example_dataset.py +++ b/tools/create_bids_example_dataset.py @@ -1,34 +1,40 @@ # Authors: Federico Raimondo # License: AGPL -from tempfile import TemporaryDirectory from pathlib import Path +from tempfile import TemporaryDirectory import datalad.api as dl -dst = 'git@gin.g-node.org:/juaml/datalad-example-bids.git' + +dst = "git@gin.g-node.org:/juaml/datalad-example-bids.git" with TemporaryDirectory() as tmpdir_name: tmpdir = Path(tmpdir_name) ds = dl.create(tmpdir) # type: ignore - base_dir = tmpdir / 'example_bids' + base_dir = tmpdir / "example_bids" base_dir.mkdir() for i_sub in range(1, 10): - t_sub = f'sub-{i_sub:02d}' + t_sub = f"sub-{i_sub:02d}" sub_dir = base_dir / t_sub sub_dir.mkdir() - for dname in ['anat', 'func']: + for dname in ["anat", "func"]: (sub_dir / dname).mkdir() - fnames = [f'anat/{t_sub}_T1w.nii.gz', - f'func/{t_sub}_task-rest_bold.nii.gz', - f'func/{t_sub}_task-rest_bold.json'] + fnames = [ + f"anat/{t_sub}_T1w.nii.gz", + f"anat/{t_sub}_brain_mask.nii.gz", + f"func/{t_sub}_task-rest_bold.nii.gz", + f"func/{t_sub}_task-rest_bold.json", + f"func/{t_sub}_task-rest_brain_mask.nii.gz", + f"func/{t_sub}_task-rest_confounds_regressors.tsv", + ] for fname in fnames: - with open(sub_dir / fname, 'w') as f: - f.write(f'placeholder-{fname}') + with open(sub_dir / fname, "w") as f: + f.write(f"placeholder-{fname}") ds.save(recursive=True) - ds.siblings('add', name='gin', url=dst) - ds.push(to='gin', force='all') + ds.siblings("add", name="gin", url=dst) + ds.push(to="gin", force="all") diff --git a/tools/create_bids_example_dataset_sessions.py b/tools/create_bids_example_dataset_sessions.py index 9d798b2c5..b37120bb2 100644 --- a/tools/create_bids_example_dataset_sessions.py +++ b/tools/create_bids_example_dataset_sessions.py @@ -1,41 +1,50 @@ # Authors: Federico Raimondo # License: AGPL -from tempfile import TemporaryDirectory from pathlib import Path +from tempfile import TemporaryDirectory import datalad.api as dl -dst = 'git@gin.g-node.org:/juaml/datalad-example-bids-ses.git' + +dst = "git@gin.g-node.org:/juaml/datalad-example-bids-ses.git" with TemporaryDirectory() as tmpdir_name: tmpdir = Path(tmpdir_name) ds = dl.create(tmpdir) # type: ignore - base_dir = tmpdir / 'example_bids_ses' + base_dir = tmpdir / "example_bids_ses" base_dir.mkdir() for i_sub in range(1, 10): - t_sub = f'sub-{i_sub:02d}' + t_sub = f"sub-{i_sub:02d}" sub_dir = base_dir / t_sub sub_dir.mkdir() for i_ses in range(1, 4): - t_ses = f'ses-{i_ses:02d}' + t_ses = f"ses-{i_ses:02d}" ses_dir = sub_dir / t_ses ses_dir.mkdir() - for dname in ['anat', 'func']: + for dname in ["anat", "func"]: (ses_dir / dname).mkdir() - fnames = [f'anat/{t_sub}_{t_ses}_T1w.nii.gz'] + fnames = [ + f"anat/{t_sub}_{t_ses}_T1w.nii.gz", + f"anat/{t_sub}_{t_ses}_brain_mask.nii.gz", + ] if i_ses != 3: # Session 3 does not have functional data - fnames.extend([ - f'func/{t_sub}_{t_ses}_task-rest_bold.nii.gz', - f'func/{t_sub}_{t_ses}_task-rest_bold.json']) + fnames.extend( + [ + f"func/{t_sub}_{t_ses}_task-rest_bold.nii.gz", + f"func/{t_sub}_{t_ses}_task-rest_bold.json", + f"func/{t_sub}_{t_ses}_task-rest_brain_mask.nii.gz", + f"func/{t_sub}_{t_ses}_task-rest_confounds_regressors.tsv", + ] + ) for fname in fnames: - with open(ses_dir / fname, 'w') as f: - f.write('placeholder-{fname}') + with open(ses_dir / fname, "w") as f: + f.write("placeholder-{fname}") ds.save(recursive=True) - ds.siblings('add', name='gin', url=dst) - ds.push(to='gin', force='all') + ds.siblings("add", name="gin", url=dst) + ds.push(to="gin", force="all") -- 2.52.0 From 2a68002c0a4ca35177a8a74c6226f8392f419cb4 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Jul 2024 12:50:28 +0200 Subject: [PATCH 15/20] chore: format tools/ --- tools/create_aomic1000_example_dataset.py | 3 ++- tools/create_aomicpiop1_example_dataset.py | 1 + tools/create_aomicpiop2_example_dataset.py | 1 + 3 files changed, 4 insertions(+), 1 deletion(-) diff --git a/tools/create_aomic1000_example_dataset.py b/tools/create_aomic1000_example_dataset.py index 55aee2a31..11bc9576a 100644 --- a/tools/create_aomic1000_example_dataset.py +++ b/tools/create_aomic1000_example_dataset.py @@ -4,11 +4,12 @@ # Vera Komeyer # Xuan Li # License: AGPL -from tempfile import TemporaryDirectory from pathlib import Path +from tempfile import TemporaryDirectory import datalad.api as dl + # repo has to be created on gin manually beforehand if not owner dst = "git@gin.g-node.org:/juaml/datalad-example-aomic1000.git" diff --git a/tools/create_aomicpiop1_example_dataset.py b/tools/create_aomicpiop1_example_dataset.py index 42898b75b..6b8bd420b 100644 --- a/tools/create_aomicpiop1_example_dataset.py +++ b/tools/create_aomicpiop1_example_dataset.py @@ -1,4 +1,5 @@ """Create an example/testing dataset for PIOP1 with mock data.""" + # Authors: Federico Raimondo # Vera Komeyer # Xuan Li diff --git a/tools/create_aomicpiop2_example_dataset.py b/tools/create_aomicpiop2_example_dataset.py index 0b58e8620..d4c76a12c 100644 --- a/tools/create_aomicpiop2_example_dataset.py +++ b/tools/create_aomicpiop2_example_dataset.py @@ -1,4 +1,5 @@ """Create an example/testing dataset for PIOP2 with mock data.""" + # Authors: Federico Raimondo # Vera Komeyer # Xuan Li -- 2.52.0 From 2407ff19b06a60054cc2a1f76580e68ff2bac67b Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 18 Jul 2024 11:55:25 +0200 Subject: [PATCH 16/20] Suppress warnings from the logger if they are filtered --- junifer/utils/logging.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index 2589a93c9..552ac5121 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -13,6 +13,7 @@ else: # pragma: no cover from looseversion import LooseVersion import logging +import warnings from pathlib import Path from subprocess import PIPE, Popen, TimeoutExpired from typing import Dict, NoReturn, Optional, Type, Union @@ -44,6 +45,14 @@ _logging_types = { } +def _showwarning(message, category, filename, lineno, file=None, line=None): + s = warnings.formatwarning(message, category, filename, lineno, line) + logger.warning(str(s)) + + +warnings.showwarning = _showwarning + + class WrapStdOut(logging.StreamHandler): """Dynamically wrap to sys.stdout. @@ -325,5 +334,5 @@ def warn_with_log( The warning subclass (default RuntimeWarning). """ - logger.warning(msg) + # logger.warning(msg) warn(msg, category=category, stacklevel=2) -- 2.52.0 From f189a33e8658ea5a807f665d074ad3fa488404cf Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 18 Jul 2024 12:36:24 +0200 Subject: [PATCH 17/20] chore: cleanup logging.py --- junifer/utils/logging.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index 552ac5121..3e17bc2cd 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -45,11 +45,13 @@ _logging_types = { } +# Copied over from stdlib and tweaked to our use-case. def _showwarning(message, category, filename, lineno, file=None, line=None): s = warnings.formatwarning(message, category, filename, lineno, line) logger.warning(str(s)) +# Overwrite warnings display to integrate with logging warnings.showwarning = _showwarning @@ -334,5 +336,4 @@ def warn_with_log( The warning subclass (default RuntimeWarning). """ - # logger.warning(msg) warn(msg, category=category, stacklevel=2) -- 2.52.0 From 0e830ce41a6f1fce36afbc6fbb9cd347d604af9b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 18 Jul 2024 12:41:56 +0200 Subject: [PATCH 18/20] chore: add changelog 351.misc --- docs/changes/newsfragments/351.misc | 1 + 1 file changed, 1 insertion(+) create mode 100644 docs/changes/newsfragments/351.misc diff --git a/docs/changes/newsfragments/351.misc b/docs/changes/newsfragments/351.misc new file mode 100644 index 000000000..5e7f255a3 --- /dev/null +++ b/docs/changes/newsfragments/351.misc @@ -0,0 +1 @@ +Integrate ``warnings`` with ``logging`` respecting filters by `Fede Raimondo`_ -- 2.52.0 From 19f51660076d29122d51b71a0afdfa0e1bb67d3c Mon Sep 17 00:00:00 2001 From: Fede Date: Thu, 18 Jul 2024 16:17:07 +0300 Subject: [PATCH 19/20] Workaround for pytest and logging with captureWarnings --- junifer/utils/logging.py | 9 ++++++++- junifer/utils/tests/test_logging.py | 9 ++++++++- 2 files changed, 16 insertions(+), 2 deletions(-) diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index 3e17bc2cd..894d71c49 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -52,7 +52,14 @@ def _showwarning(message, category, filename, lineno, file=None, line=None): # Overwrite warnings display to integrate with logging -warnings.showwarning = _showwarning + + +def capture_warnings(): + """Capture warnings and log them.""" + warnings.showwarning = _showwarning + + +capture_warnings() class WrapStdOut(logging.StreamHandler): diff --git a/junifer/utils/tests/test_logging.py b/junifer/utils/tests/test_logging.py index ea8d51e2a..b1515bbfc 100644 --- a/junifer/utils/tests/test_logging.py +++ b/junifer/utils/tests/test_logging.py @@ -9,7 +9,6 @@ import logging from pathlib import Path import pytest - from junifer.utils.logging import ( _close_handlers, configure_logging, @@ -145,8 +144,16 @@ def test_log_file(tmp_path: Path) -> None: assert any("Warn3 message" in line for line in lines) assert any("Error3 message" in line for line in lines) + # This should raise a warning (test that it was raised) with pytest.warns(RuntimeWarning, match=r"Warn raised"): warn_with_log("Warn raised") + + # This should log the warning (workaround for pytest messing with logging) + from junifer.utils.logging import capture_warnings + + capture_warnings() + + warn_with_log("Warn raised 2") with pytest.raises(ValueError, match=r"Error raised"): raise_error("Error raised") with open(tmp_path / "test4.log") as f: -- 2.52.0 From f2d8d08984f3a47c67e18db24c6ab6b3e82b3c77 Mon Sep 17 00:00:00 2001 From: Fede Date: Thu, 18 Jul 2024 16:17:59 +0300 Subject: [PATCH 20/20] organise imports --- junifer/utils/tests/test_logging.py | 1 + 1 file changed, 1 insertion(+) diff --git a/junifer/utils/tests/test_logging.py b/junifer/utils/tests/test_logging.py index b1515bbfc..92a1102ef 100644 --- a/junifer/utils/tests/test_logging.py +++ b/junifer/utils/tests/test_logging.py @@ -9,6 +9,7 @@ import logging from pathlib import Path import pytest + from junifer.utils.logging import ( _close_handlers, configure_logging, -- 2.52.0