diff --git a/docs/changes/latest.inc b/docs/changes/latest.inc index 234590313..ab910e434 100644 --- a/docs/changes/latest.inc +++ b/docs/changes/latest.inc @@ -83,6 +83,8 @@ Enhancements - Rename "atlas" to "parcellation" (:gh:`116` by `Fede Raimondo`_). +- Refactor the :class:`junifer.datagrabber.BaseDataGrabber` class to allow for easier subclassing (:gh:`123` by `Fede Raimondo`_). + Bugs ~~~~ @@ -91,6 +93,8 @@ Bugs - Fix a bug in which AOMIC PIOP2 datagrabber did not use user input to constrain elements based on tasks (:gh:`105` by `Leonard Sasse`_) +- Fix a bug in which a datalad dataset could remove a user-cloned dataset (:gh:`53` by `Fede Raimondo`_) + API changes ~~~~~~~~~~~ diff --git a/junifer/datagrabber/aomic/id1000.py b/junifer/datagrabber/aomic/id1000.py index 250a78594..2266ff04b 100644 --- a/junifer/datagrabber/aomic/id1000.py +++ b/junifer/datagrabber/aomic/id1000.py @@ -24,15 +24,11 @@ class DataladAOMICID1000(PatternDataladDataGrabber): The directory where the datalad dataset will be cloned. If None, the datalad dataset will be cloned into a temporary directory (default None). - **kwargs - Keyword arguments passed to superclass. - """ def __init__( self, datadir: Union[str, Path, None] = None, - **kwargs, ) -> None: # The types of data types = [ diff --git a/junifer/datagrabber/aomic/piop1.py b/junifer/datagrabber/aomic/piop1.py index d633ef3b4..0e65acee8 100644 --- a/junifer/datagrabber/aomic/piop1.py +++ b/junifer/datagrabber/aomic/piop1.py @@ -8,7 +8,7 @@ from itertools import product from pathlib import Path -from typing import Dict, List, Tuple, Union +from typing import Dict, List, Union from junifer.datagrabber import PatternDataladDataGrabber @@ -30,16 +30,12 @@ class DataladAOMICPIOP1(PatternDataladDataGrabber): "gstroop", "workingmemory"} or list of the options, optional AOMIC PIOP1 task sessions. If None, all available task sessions are selected (default None). - **kwargs - Keyword arguments passed to superclass. - """ def __init__( self, datadir: Union[str, Path, None] = None, tasks: Union[str, List[str], None] = None, - **kwargs, ) -> None: # The types of data types = [ @@ -122,25 +118,24 @@ class DataladAOMICPIOP1(PatternDataladDataGrabber): replacements=replacements, ) - def __getitem__(self, element: Tuple[str, str]) -> Dict[str, Path]: + def get_item(self, subject: str, task: str) -> Dict[str, Path]: """Index one element in the dataset. Parameters ---------- - element : tuple of str - The element to be indexed. First element in the tuple is the - subject, second element is the task. + subject : str + The subject ID. + task : {"restingstate", "anticipation", "emomatching", "faces", \ + "gstroop", "workingmemory"} + The task to get. Returns ------- out : dict Dictionary of paths for each type of data required for the specified element. - """ - sub, task = element - # depending on task 'acquisition is different' task_acqs = { "anticipation": "seq", "emomatching": "seq", @@ -152,8 +147,7 @@ class DataladAOMICPIOP1(PatternDataladDataGrabber): acq = task_acqs[task] new_task = f"{task}_acq-{acq}" - out = super().__getitem__((sub, new_task)) - out["meta"]["element"] = {"subject": sub, "task": task} + out = super().get_item(subject=subject, task=new_task) return out def get_elements(self) -> List: diff --git a/junifer/datagrabber/aomic/piop2.py b/junifer/datagrabber/aomic/piop2.py index a13354227..9d39eb024 100644 --- a/junifer/datagrabber/aomic/piop2.py +++ b/junifer/datagrabber/aomic/piop2.py @@ -29,16 +29,12 @@ class DataladAOMICPIOP2(PatternDataladDataGrabber): or list of the options, optional AOMIC PIOP2 task sessions. If None, all available task sessions are selected (default None). - **kwargs - Keyword arguments passed to superclass. - """ def __init__( self, datadir: Union[str, Path, None] = None, tasks: Union[str, List[str], None] = None, - **kwargs, ) -> None: # The types of data types = [ diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index ef5cc3ec7..3fcbad91c 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -63,10 +63,7 @@ class BaseDataGrabber(ABC): Parameters ---------- element : str or tuple - 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". + The element to be indexed. Returns ------- @@ -76,8 +73,14 @@ class BaseDataGrabber(ABC): """ logger.info(f"Getting element {element}") - out = {} - out["meta"] = {"datagrabber": self.get_meta()} + if not isinstance(element, tuple): + element = (element,) + named_element = dict(zip(self.get_element_keys(), element)) + out = self.get_item(**named_element) + out["meta"] = { + "datagrabber": self.get_meta(), + "element": named_element + } return out def __enter__(self) -> "BaseDataGrabber": @@ -115,18 +118,6 @@ class BaseDataGrabber(ABC): t_meta[k] = v return t_meta - # TODO: what is the final functionality? - def get_element_keys(self) -> str: - """Get element keys. - - Returns - ------- - str - The element keys. - - """ - return "element" - @property def datadir(self) -> Path: """Get data directory path. @@ -139,6 +130,24 @@ class BaseDataGrabber(ABC): """ return self._datadir + @abstractmethod + def get_element_keys(self) -> str: + """Get element keys. + + For each item in the ``element`` tuple passed to ``__getitem__()``, + this method returns the corresponding key(s). + + Returns + ------- + str + The element keys. + + """ + raise_error( + msg="Concrete classes need to implement get_element_keys().", + klass=NotImplementedError, + ) + @abstractmethod def get_elements(self) -> List: """Get elements. @@ -155,3 +164,24 @@ class BaseDataGrabber(ABC): msg="Concrete classes need to implement get_elements().", klass=NotImplementedError, ) + + @abstractmethod + def get_item(self, **element: Dict) -> Dict[str, Dict]: + """Get the specified item from the dataset. + + Parameters + ---------- + element : dict + The element to be indexed. + + Returns + ------- + dict + Dictionary of paths for each type of data required for the + specified element. + + """ + raise_error( + msg="Concrete classes need to implement get_item().", + klass=NotImplementedError, + ) diff --git a/junifer/datagrabber/datalad_base.py b/junifer/datagrabber/datalad_base.py index 2377736ba..1e5fa924e 100644 --- a/junifer/datagrabber/datalad_base.py +++ b/junifer/datagrabber/datalad_base.py @@ -10,11 +10,12 @@ from pathlib import Path from typing import Dict, Optional, Tuple, Union import datalad.api as dl +from datalad.support.gitrepo import GitRepo from ..api.decorators import register_datagrabber from ..utils import logger from .base import BaseDataGrabber -from .utils import raise_error +from ..utils import raise_error, warn_with_log @register_datagrabber @@ -81,12 +82,37 @@ class DataladDataGrabber(BaseDataGrabber): logger.debug(f"\t_rootdir = {rootdir}") self.uri = uri self._rootdir = rootdir + # Flag to indicate if the dataset was cloned before and it might be + # dirty + self._dataset_dirty = False @property def datadir(self) -> Path: """Get data directory path.""" return super().datadir / self._rootdir + def _get_dataset_id_remote(self) -> str: + """Get the dataset id from the remote. + + Returns + ------- + str + The dataset id. + + """ + remote_id = None + with tempfile.TemporaryDirectory() as tmpdir: + logger.debug(f"Querying {self.uri} for dataset ID") + repo = GitRepo.clone( + self.uri, path=tmpdir, + clone_options=["-n", "--depth=1"]) + repo.checkout(name=".datalad/config", options=["HEAD"]) + remote_id = repo.config.get("datalad.dataset.id", None) + logger.debug(f"Got remote dataset ID = {remote_id}") + if remote_id is None: + raise_error("Could not get dataset ID from remote") + return remote_id + def _dataset_get(self, out: Dict) -> Dict: """Get the dataset found from the path in `out`. @@ -98,39 +124,90 @@ class DataladDataGrabber(BaseDataGrabber): Returns ------- dict - The modified dictionary with version appended. + The modified dictionary with meta updated. """ - for _, v in out.items(): - if "path" in v: - logger.debug(f"Getting {v['path']}") - # Note, that `self._dataset.get` without an option would get - # the content of all files in a (sub-dataset) if `v["path"]` - # was to point to a subdataset rather than a file. This may be - # a source of confusion (+ performance/storage issue) when - # implementing a grabber. - self._dataset.get(v["path"]) - logger.debug("Get done") + to_get = [v["path"] for v in out.values() if "path" in v] + + if len(to_get) > 0: + logger.debug(f"Getting {len(to_get)} files using datalad:") + for fname in to_get: + logger.debug(f"\t: {fname}") + + dl_out = self._dataset.get(to_get, result_renderer="disabled") + if not self._was_cloned: + # If the dataset was already installed, check that the + # file was actually downloaded to avoid removing a + # file that was already there. + for t_out in dl_out: + t_path = Path(t_out["path"]) + if t_out["status"] == "ok": + logger.debug(f"File {t_path} downloaded") + self._got_files.append(t_path) + elif t_out["status"] == "notneeded": + logger.debug( + f"File {t_path} was already present" + ) + else: + raise_error(f"File download failed: {t_out}") + logger.debug("Get done") - # append the version of the dataset - out["meta"]["datagrabber"][ - "dataset_commit_id" - ] = self._dataset.repo.get_hexsha( - self._dataset.repo.get_corresponding_branch() - ) return out def install(self) -> None: - """Install the datalad dataset into the datadir.""" - logger.debug(f"Installing dataset {self.uri} to {self._datadir}") - self._dataset: dl.Dataset = dl.clone(self.uri, self._datadir) - logger.debug("Dataset installed") + """Install the datalad dataset into the datadir. - def remove(self): - """Remove the datalad dataset from the datadir.""" - # This probably wants to use `reckless='kill'` or similar. - # See issue #53 - self._dataset.remove(recursive=True) + Raises + ------ + ValueError + If the dataset is already installed but with a different ID. + """ + isinstalled = dl.Dataset(self._datadir).is_installed() + if isinstalled: + logger.debug("Dataset already installed") + self._got_files = [] + self._dataset: dl.Dataset = dl.Dataset(self._datadir) + + remote_id = self._get_dataset_id_remote() + if remote_id != self._dataset.id: + raise_error( + "Dataset already installed but with a different " + f"ID: {self._dataset.id} (local) != {remote_id} (remote)" + ) + + # Check for dirty datasets: + status = self._dataset.status() + if any([x["state"] != "clean" for x in status]): + self._dataset_dirty = True + warn_with_log( + "At least one file is not clean, Junifer will " + "consider this dataset as dirty." + ) + else: + logger.debug("Dataset is clean") + + else: + logger.debug(f"Installing dataset {self.uri} to {self._datadir}") + self._dataset: dl.Dataset = dl.clone( # type: ignore + self.uri, self._datadir, result_renderer="disabled" + ) + logger.debug("Dataset installed") + self._was_cloned = not isinstalled + + self._datalad_commit_id = self._dataset.repo.get_hexsha( + self._dataset.repo.get_corresponding_branch() + ) + + def cleanup(self) -> None: + """Cleanup the datalad dataset.""" + if self._was_cloned: + logger.debug("Removing dataset with reckless='kill'") + self._dataset.remove(reckless="kill", result_renderer="disabled") + else: + logger.debug("Dropping files that were downloaded") + for f in self._got_files: + logger.debug(f"Dropping {f}") + self._dataset.drop(f, result_renderer="disabled") def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Path]: """Implement single element indexing in the Datalad database. @@ -166,6 +243,24 @@ class DataladDataGrabber(BaseDataGrabber): def __exit__(self, exc_type, exc_value, exc_traceback): """Implement context exit.""" - logger.debug("Removing dataset") - self.remove() - logger.debug("Dataset removed") + logger.debug("Cleaning up dataset") + self.cleanup() + logger.debug("Dataset state restored") + + def get_meta(self) -> Dict: + """Get metadata. + + Returns + ------- + dict + The metadata as dictionary. + + """ + t_meta = super().get_meta() + t_meta["datalad_commit_id"] = self._datalad_commit_id + + t_meta["datalad_id"] = self._dataset.id + + # Set a flag to indicate that the dataset was dirty + t_meta["datalad_dirty"] = self._dataset_dirty + return t_meta diff --git a/junifer/datagrabber/hcp.py b/junifer/datagrabber/hcp.py index b8e468ec6..bb67d32e7 100644 --- a/junifer/datagrabber/hcp.py +++ b/junifer/datagrabber/hcp.py @@ -2,7 +2,7 @@ from itertools import product from pathlib import Path -from typing import Dict, List, Tuple, Union +from typing import Dict, List, Union from junifer.datagrabber.datalad_base import DataladDataGrabber @@ -104,15 +104,20 @@ class HCP1200(PatternDataGrabber): ) self.phase_encodings = phase_encodings - def __getitem__(self, element: Tuple[str, str, str]) -> Dict[str, Path]: + def get_item( + self, subject: str, task: str, phase_encoding: str + ) -> Dict[str, Path]: """Index one element in the dataset. Parameters ---------- - element : triple of str - The element to be indexed. First element in the tuple is the - subject, second element is the task, third element is the - phase encoding direction. + subject : str + The subject ID. + task : {"REST1", "REST2", "SOCIAL", "WM", "RELATIONAL", "EMOTION", \ + "LANGUAGE", "GAMBLING", "MOTOR"} + The task. + phase_encoding : {"LR", "RL"} + The phase encoding. Returns ------- @@ -121,20 +126,15 @@ class HCP1200(PatternDataGrabber): specified element. """ - sub, task, phase_encoding = element - # Resting task if "REST" in task: new_task = f"rfMRI_{task}" else: new_task = f"tfMRI_{task}" - out = super().__getitem__((sub, new_task, phase_encoding)) - out["meta"]["element"] = { - "subject": sub, - "task": task, - "phase_encoding": phase_encoding, - } + out = super().get_item( + subject=subject, task=new_task, phase_encoding=phase_encoding + ) return out def get_elements(self) -> List: @@ -173,17 +173,13 @@ class DataladHCP1200(DataladDataGrabber, HCP1200): phase_encodings : {"LR", "RL"} or list of the options, optional HCP phase encoding directions. If None, both will be used (default None). - **kwargs - Keyword arguments passed to superclass. - """ def __init__( self, datadir: Union[str, Path, None] = None, tasks: Union[str, List[str], None] = None, - phase_encodings: Union[str, List[str], None] = None, - **kwargs, + phase_encodings: Union[str, List[str], None] = None ) -> None: uri = ( "https://github.com/datalad-datasets/" diff --git a/junifer/datagrabber/multiple.py b/junifer/datagrabber/multiple.py index ac4370102..68fad38bc 100644 --- a/junifer/datagrabber/multiple.py +++ b/junifer/datagrabber/multiple.py @@ -26,9 +26,16 @@ class MultipleDataGrabber(BaseDataGrabber): """ def __init__(self, datagrabbers: List[BaseDataGrabber], **kwargs) -> None: - # TODO: Check datagrabbers consistency - # - same element keys - # - no overlapping types + # Check datagrabbers consistency + # 1) same element keys + first_keys = datagrabbers[0].get_element_keys() + for dg in datagrabbers[1:]: + if dg.get_element_keys() != first_keys: + raise ValueError("Datagrabbers have different element keys.") + # 2) no overlapping types + types = [x for dg in datagrabbers for x in dg.get_types()] + if len(types) != len(set(types)): + raise ValueError("Datagrabbers have overlapping types.") self._datagrabbers = datagrabbers def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Path]: @@ -49,12 +56,34 @@ class MultipleDataGrabber(BaseDataGrabber): specified element. """ + out = {} for dg in self._datagrabbers: t_out = dg[element] out.update(t_out) return out + def get_item(self, **element: Dict) -> Dict[str, Dict]: + """Get item. + + Parameters + ---------- + element : dict + The element to be indexed. + + Returns + ------- + dict + Dictionary of paths for each type of data required for the + specified element. + + Notes + ----- + This function is not implemented for this class as it is useless. + """ + raise NotImplementedError( + "get_item() is not useful for this class, hence not implemented.") + def __enter__(self) -> "BaseDataGrabber": """Implement context entry.""" for dg in self._datagrabbers: @@ -82,6 +111,19 @@ class MultipleDataGrabber(BaseDataGrabber): elements.intersection_update(s) return list(elements) + def get_element_keys(self) -> List[str]: + """Get element keys. + + For each item in the ``element`` tuple passed to ``__getitem__()``, + this method returns the corresponding key(s). + + Returns + ------- + list of str + The element keys. + """ + return self._datagrabbers[0].get_element_keys() + def get_types(self) -> List[str]: """Get types. diff --git a/junifer/datagrabber/pattern.py b/junifer/datagrabber/pattern.py index e67bcb573..65bd43bfc 100644 --- a/junifer/datagrabber/pattern.py +++ b/junifer/datagrabber/pattern.py @@ -105,12 +105,12 @@ class PatternDataGrabber(BaseDataGrabber): glob_pattern = glob_pattern.replace(f"{{{t_r}}}", "*") return re_pattern, glob_pattern, t_replacements - def _replace_patterns_glob(self, element: Tuple, pattern: str) -> str: + def _replace_patterns_glob(self, element: Dict, pattern: str) -> str: """Replace patterns with the element so it can be globbed. Parameters ---------- - element : tuple + element : dict The element to be used in the replacement. pattern : str The pattern to be replaced. @@ -121,27 +121,39 @@ class PatternDataGrabber(BaseDataGrabber): The pattern with the element replaced. """ - if len(element) != len(self.replacements): + if list(element.keys()) != self.replacements: raise_error( - f"The element length must be {len(self.replacements)}, " - f"indicating {self.replacements}." + f"The element keys must be {self.replacements}, " + f"element has {list(element.keys())}." ) - to_replace = dict(zip(self.replacements, element)) - return pattern.format(**to_replace) + return pattern.format(**element) - def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Dict]: + def get_element_keys(self) -> List[str]: + """Get element keys. + + For each item in the "element" tuple, this functions returns the + corresponding key, that is, the ``replacements`` of patterns defined + in the constructor. + + Returns + ------- + list of str + The element keys. + + """ + return self.replacements + + def get_item(self, **element: Dict) -> Dict[str, Dict]: """Implement single element indexing in the database. - Each occurrence of the strings in "replacements" is replaced by the - corresponding item in the element tuple. + This method constructs a real path to the requested item's data, by + replacing the ``patterns`` with actual values passed via ``**element``. Parameters ---------- - element : str or tuple - 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". + element : dict + The element to be indexed. The keys must be the same as the + replacements. Returns ------- @@ -150,9 +162,7 @@ class PatternDataGrabber(BaseDataGrabber): specified element. """ - out = super().__getitem__(element) - if not isinstance(element, tuple): - element = (element,) + out = {} for t_type in self.types: t_pattern = self.patterns[t_type] t_replace = self._replace_patterns_glob(element, t_pattern) @@ -175,8 +185,6 @@ class PatternDataGrabber(BaseDataGrabber): f"File {t_out} does not exist" ) out[t_type] = {"path": t_out} - # Meta here is element and types - out["meta"]["element"] = dict(zip(self.replacements, element)) return out def get_elements(self) -> List: diff --git a/junifer/datagrabber/tests/test_base.py b/junifer/datagrabber/tests/test_base.py index 91d622d62..946a1f71a 100644 --- a/junifer/datagrabber/tests/test_base.py +++ b/junifer/datagrabber/tests/test_base.py @@ -22,18 +22,24 @@ def test_BaseDataGrabber() -> None: """Test BaseDataGrabber.""" # Create concrete class. class MyDataGrabber(BaseDataGrabber): - def __getitem__(self, element): - return super().__getitem__(element) + def get_item(self, subject): + return {} def get_elements(self): return super().get_elements() + def get_element_keys(self): + return ["subject"] + dg = MyDataGrabber(datadir="/tmp", types=["func"]) elem = dg["elem"] assert "meta" in elem assert "datagrabber" in elem["meta"] assert "class" in elem["meta"]["datagrabber"] assert MyDataGrabber.__name__ in elem["meta"]["datagrabber"]["class"] + assert "element" in elem["meta"] + assert "subject" in elem["meta"]["element"] + assert "elem" in elem["meta"]["element"]["subject"] with pytest.raises(NotImplementedError): dg.get_elements() @@ -41,3 +47,19 @@ def test_BaseDataGrabber() -> None: with dg: assert dg.datadir == Path("/tmp") assert dg.types == ["func"] + + class MyDataGrabber2(BaseDataGrabber): + def get_item(self, subject): + return super().get_item(subject=subject) + + def get_elements(self): + return super().get_elements() + + def get_element_keys(self): + return super().get_element_keys() + dg = MyDataGrabber2(datadir="/tmp", types=["func"]) + with pytest.raises(NotImplementedError): + dg.get_element_keys() + + with pytest.raises(NotImplementedError): + dg.get_item(subject=1) # type: ignore diff --git a/junifer/datagrabber/tests/test_datalad_base.py b/junifer/datagrabber/tests/test_datalad_base.py index 2197fc097..4eaebf638 100644 --- a/junifer/datagrabber/tests/test_datalad_base.py +++ b/junifer/datagrabber/tests/test_datalad_base.py @@ -5,18 +5,419 @@ import pytest +from pathlib import Path + +import datalad.api as dl + from junifer.datagrabber.datalad_base import DataladDataGrabber +_testing_dataset = { + "example_bids": { + "uri": "https://gin.g-node.org/juaml/datalad-example-bids", + "commit": "522dfb203afcd2cd55799bf347f9b211919a7338", + "id": "fec92475-d9c0-4409-92ba-f041b6a12c40", + }, + "example_bids_ses": { + "uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses", + "commit": "3d08d55d1faad4f12ab64ac9497544a0d924d47a", + "id": "c83500d0-532f-45be-baf1-0dab703bdc2a", + }, +} + + def test_datalad_base_abstractness() -> None: """Test datalad base is abstract.""" with pytest.raises(TypeError, match=r"abstract"): DataladDataGrabber() -# def test_datalad_base_missing_uri() -> None: -# """Test proper check of missing URI in datalad base initialization.""" -# with pytest.raises(ValueError, match=r"`uri` must be provided"): -# DataladDataGrabber( +@pytest.fixture +def concrete_datagrabber() -> DataladDataGrabber: + """Return a concrete datagrabber class. -# ) + Returns + ------- + DataladDataGrabber + A concrete datagrabber class. + + """ + + class MyDataGrabber(DataladDataGrabber): # type: ignore + def __init__(self, datadir, uri): + super().__init__( + datadir=datadir, + rootdir="example_bids", + uri=uri, + types=["T1w", "BOLD"], + ) + + def get_item(self, subject): + out = { + "T1w": { + "path": self.datadir + / f"{subject}/anat/{subject}_T1w.nii.gz" + }, + "BOLD": { + "path": self.datadir + / f"{subject}/func/{subject}_task-rest_bold.nii.gz" + }, + } + return out + + def get_elements(self): + return [f"sub-{i:02d}" for i in range(1, 10)] + + def get_element_keys(self): + return ["subject"] + + return MyDataGrabber + + +def test_datalad_install_errors( + tmp_path: Path, concrete_datagrabber: DataladDataGrabber +) -> None: + """Test datalad base install errors / warnings. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + concrete_datagrabber : DataladDataGrabber + A concrete datagrabber class to use. + """ + + # Dataset cloned outside of datagrabber + datadir = tmp_path / "cloned_uri" + uri = _testing_dataset["example_bids"]["uri"] + uri2 = _testing_dataset["example_bids_ses"]["uri"] + + # Files are not there + assert datadir.exists() is False + # Clone dataset + dl.clone(uri, datadir) # type: ignore + dg = concrete_datagrabber(datadir=datadir, uri=uri2) + with pytest.raises(ValueError, match=r"different ID"): + with dg: + pass + + elem1_t1w = datadir / "example_bids/sub-01/anat/sub-01_T1w.nii.gz" + elem1_t1w.unlink() + with open(elem1_t1w, "w") as f: + f.write("modified!") + + dg = concrete_datagrabber(datadir=datadir, uri=uri) + with pytest.warns(RuntimeWarning, match=r"one file is not clean"): + with dg: + pass + + +def test_datalad_clone_cleanup( + tmp_path: Path, concrete_datagrabber: DataladDataGrabber +) -> None: + """Test datalad base clone and remove. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + concrete_datagrabber : DataladDataGrabber + A concrete datagrabber class to use. + """ + + # Clone whole dataset + datadir = tmp_path / "newclone" + uri = _testing_dataset["example_bids"]["uri"] + elem1_bold = ( + datadir / "example_bids/sub-01/func/sub-01_task-rest_bold.nii.gz" + ) + elem1_t1w = datadir / "example_bids/sub-01/anat/sub-01_T1w.nii.gz" + + assert datadir.exists() is False + assert elem1_bold.is_file() is False + assert elem1_t1w.is_file() is False + with concrete_datagrabber(datadir=datadir, uri=uri) as dg: + assert datadir.exists() is True + assert dg._was_cloned is True + assert elem1_bold.is_file() is False + assert elem1_bold.is_symlink() is True + assert elem1_t1w.is_file() is False + assert elem1_t1w.is_symlink() is True + elem1 = dg["sub-01"] + assert "meta" in elem1 + assert "datagrabber" in elem1["meta"] + assert "datalad_dirty" in elem1["meta"]["datagrabber"] + assert elem1["meta"]["datagrabber"]["datalad_dirty"] is False + assert hasattr(dg, "_got_files") is False + assert datadir.exists() is True + assert elem1_bold.is_file() is True + assert elem1_bold.is_symlink() is True + assert elem1_t1w.is_file() is True + assert elem1_t1w.is_symlink() is True + + assert datadir.exists() is False + assert len(list(datadir.glob("*"))) == 0 + + +def test_datalad_previously_cloned( + tmp_path: Path, concrete_datagrabber: DataladDataGrabber +) -> None: + """Test datalad base on cloned dataset. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + concrete_datagrabber : DataladDataGrabber + A concrete datagrabber class to use. + """ + + # Dataset cloned outside of datagrabber + datadir = tmp_path / "cloned" + elem1_bold = ( + datadir / "example_bids/sub-01/func/sub-01_task-rest_bold.nii.gz" + ) + elem1_t1w = datadir / "example_bids/sub-01/anat/sub-01_T1w.nii.gz" + uri = _testing_dataset["example_bids"]["uri"] + commit = _testing_dataset["example_bids"]["commit"] + remote_id = _testing_dataset["example_bids"]["id"] + # Files are not there + assert datadir.exists() is False + assert elem1_bold.exists() is False + assert elem1_t1w.exists() is False + + # Clone dataset + dl.clone(uri, datadir, result_renderer="disabled") # type: ignore + + # Files are there, but are empty symbolic links + assert datadir.exists() is True + assert elem1_bold.is_symlink() is True + assert elem1_bold.is_file() is False + assert elem1_t1w.is_symlink() is True + assert elem1_t1w.is_file() is False + with concrete_datagrabber(datadir=datadir, uri=uri) as dg: + assert datadir.exists() is True + assert dg._was_cloned is False + elem1 = dg["sub-01"] + assert "meta" in elem1 + assert "datagrabber" in elem1["meta"] + assert "datalad_dirty" in elem1["meta"]["datagrabber"] + assert elem1["meta"]["datagrabber"]["datalad_dirty"] is False + assert "datalad_commit_id" in elem1["meta"]["datagrabber"] + assert elem1["meta"]["datagrabber"]["datalad_commit_id"] == commit + assert "datalad_id" in elem1["meta"]["datagrabber"] + assert elem1["meta"]["datagrabber"]["datalad_id"] == remote_id + + assert hasattr(dg, "_got_files") is True + # Files are there and symlinks are fixed + assert elem1["BOLD"]["path"].is_file() is True + assert elem1["BOLD"]["path"].is_symlink() is True + assert elem1["T1w"]["path"].is_file() is True + assert elem1["T1w"]["path"].is_symlink() is True + + # Datagrabber fetched two files + assert len(dg._got_files) == 2 + assert any(x.name == "sub-01_T1w.nii.gz" for x in dg._got_files) + assert any( + x.name == "sub-01_task-rest_bold.nii.gz" for x in dg._got_files + ) + + assert datadir.exists() is True + assert len(list(datadir.glob("*"))) > 0 + + +def test_datalad_previously_cloned_and_get( + tmp_path: Path, concrete_datagrabber: DataladDataGrabber +) -> None: + """Test datalad base on cloned dataset with files present. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + concrete_datagrabber : DataladDataGrabber + A concrete datagrabber class to use. + """ + + # Dataset cloned outside of datagrabber with some files present + datadir = tmp_path / "cloned_clean" + elem1_bold = ( + datadir / "example_bids/sub-01/func/sub-01_task-rest_bold.nii.gz" + ) + elem1_t1w = datadir / "example_bids/sub-01/anat/sub-01_T1w.nii.gz" + uri = _testing_dataset["example_bids"]["uri"] + commit = _testing_dataset["example_bids"]["commit"] + remote_id = _testing_dataset["example_bids"]["id"] + + # Files are not there + assert datadir.exists() is False + assert elem1_bold.exists() is False + assert elem1_t1w.exists() is False + + # Clone dataset + dl.clone(uri, datadir, result_renderer="disabled") # type: ignore + + # Files are there, but are empty symbolic links + assert datadir.exists() is True + assert elem1_bold.is_symlink() is True + assert elem1_bold.is_file() is False + assert elem1_t1w.is_symlink() is True + assert elem1_t1w.is_file() is False + + dl.get( # type: ignore + elem1_t1w, dataset=datadir, result_renderer="disabled" + ) + + assert elem1_bold.is_symlink() is True + assert elem1_bold.is_file() is False + assert elem1_t1w.is_symlink() is True + assert elem1_t1w.is_file() is True + + with concrete_datagrabber(datadir=datadir, uri=uri) as dg: + assert datadir.exists() is True + assert dg._was_cloned is False + elem1 = dg["sub-01"] + assert "meta" in elem1 + assert "datagrabber" in elem1["meta"] + assert "datalad_dirty" in elem1["meta"]["datagrabber"] + assert elem1["meta"]["datagrabber"]["datalad_dirty"] is False + assert "datalad_commit_id" in elem1["meta"]["datagrabber"] + assert elem1["meta"]["datagrabber"]["datalad_commit_id"] == commit + assert "datalad_id" in elem1["meta"]["datagrabber"] + assert elem1["meta"]["datagrabber"]["datalad_id"] == remote_id + + assert hasattr(dg, "_got_files") is True + # Files are there and symlinks are fixed + assert elem1["BOLD"]["path"].is_file() is True + assert elem1["BOLD"]["path"].is_symlink() is True + assert elem1["T1w"]["path"].is_file() is True + assert elem1["T1w"]["path"].is_symlink() is True + + # Datagrabber fetched two files + assert len(dg._got_files) == 1 + assert dg._got_files[0].name == "sub-01_task-rest_bold.nii.gz" + + assert datadir.exists() is True + assert len(list(datadir.glob("*"))) > 0 + + # Same state as before using the datagrabber + assert elem1_bold.is_symlink() is True + assert elem1_bold.is_file() is False + assert elem1_t1w.is_symlink() is True + assert elem1_t1w.is_file() is True + + +def test_datalad_previously_cloned_and_get_dirty( + tmp_path: Path, concrete_datagrabber: DataladDataGrabber +) -> None: + """Test datalad base on a dirty cloned dataset. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + concrete_datagrabber : DataladDataGrabber + A concrete datagrabber class to use. + """ + + # Dataset cloned outside of datagrabber with some files present and dirty + datadir = tmp_path / "cloned_dirty" + elem1_bold = ( + datadir / "example_bids/sub-01/func/sub-01_task-rest_bold.nii.gz" + ) + elem1_t1w = datadir / "example_bids/sub-01/anat/sub-01_T1w.nii.gz" + uri = _testing_dataset["example_bids"]["uri"] + commit = _testing_dataset["example_bids"]["commit"] + remote_id = _testing_dataset["example_bids"]["id"] + + # Files are not there + assert datadir.exists() is False + assert elem1_bold.exists() is False + assert elem1_t1w.exists() is False + + # Clone dataset + dl.clone(uri, datadir, result_renderer="disabled") # type: ignore + + # Files are there, but are empty symbolic links + assert datadir.exists() is True + assert elem1_bold.is_symlink() is True + assert elem1_bold.is_file() is False + assert elem1_t1w.is_symlink() is True + assert elem1_t1w.is_file() is False + + dl.get( # type: ignore + elem1_t1w, dataset=datadir, result_renderer="disabled" + ) + + assert elem1_bold.is_symlink() is True + assert elem1_bold.is_file() is False + assert elem1_t1w.is_symlink() is True + assert elem1_t1w.is_file() is True + + elem1_t1w.unlink() + with open(elem1_t1w, "w") as f: + f.write("modified!") + + with concrete_datagrabber(datadir=datadir, uri=uri) as dg: + assert datadir.exists() is True + assert dg._was_cloned is False + elem1 = dg["sub-01"] + assert "meta" in elem1 + assert "datagrabber" in elem1["meta"] + assert "datalad_dirty" in elem1["meta"]["datagrabber"] + assert elem1["meta"]["datagrabber"]["datalad_dirty"] is True + assert "datalad_commit_id" in elem1["meta"]["datagrabber"] + assert elem1["meta"]["datagrabber"]["datalad_commit_id"] == commit + assert "datalad_id" in elem1["meta"]["datagrabber"] + assert elem1["meta"]["datagrabber"]["datalad_id"] == remote_id + + assert hasattr(dg, "_got_files") is True + # Files are there and symlinks are fixed + assert elem1["BOLD"]["path"].is_file() is True + assert elem1["BOLD"]["path"].is_symlink() is True + assert elem1["T1w"]["path"].is_file() is True + assert elem1["T1w"]["path"].is_symlink() is False + + # Datagrabber fetched two files + assert len(dg._got_files) == 1 + assert dg._got_files[0].name == "sub-01_task-rest_bold.nii.gz" + + # Now get another subject that has not been modified + with concrete_datagrabber(datadir=datadir, uri=uri) as dg: + assert datadir.exists() is True + assert dg._was_cloned is False + elem2 = dg["sub-02"] + assert "meta" in elem2 + assert "datagrabber" in elem2["meta"] + assert "datalad_dirty" in elem2["meta"]["datagrabber"] + + # Dataset is still dirty due to subject sub-01 + assert elem2["meta"]["datagrabber"]["datalad_dirty"] is True + + assert "datalad_commit_id" in elem2["meta"]["datagrabber"] + assert elem2["meta"]["datagrabber"]["datalad_commit_id"] == commit + assert "datalad_id" in elem2["meta"]["datagrabber"] + assert elem2["meta"]["datagrabber"]["datalad_id"] == remote_id + + assert hasattr(dg, "_got_files") is True + # Files are there and symlinks are fixed + assert elem2["BOLD"]["path"].is_file() is True + assert elem2["BOLD"]["path"].is_symlink() is True + assert elem2["T1w"]["path"].is_file() is True + assert elem2["T1w"]["path"].is_symlink() is True + + # Datagrabber fetched two files + assert len(dg._got_files) == 2 + assert any(x.name == "sub-02_T1w.nii.gz" for x in dg._got_files) + assert any( + x.name == "sub-02_task-rest_bold.nii.gz" for x in dg._got_files + ) + + assert datadir.exists() is True + assert len(list(datadir.glob("*"))) > 0 + + # Same state as before using the datagrabber + assert elem1_bold.is_symlink() is True + assert elem1_bold.is_file() is False + assert elem1_t1w.is_symlink() is False + assert elem1_t1w.is_file() is True diff --git a/junifer/datagrabber/tests/test_multiple.py b/junifer/datagrabber/tests/test_multiple.py index 5e1298d3a..badfcea34 100644 --- a/junifer/datagrabber/tests/test_multiple.py +++ b/junifer/datagrabber/tests/test_multiple.py @@ -3,13 +3,15 @@ # Authors: Federico Raimondo # License: AGPL +import pytest + from junifer.datagrabber import MultipleDataGrabber, PatternDataladDataGrabber _testing_dataset = { "example_bids": { "uri": "https://gin.g-node.org/juaml/datalad-example-bids", - "id": "e2ce149bd723088769a86c72e57eded009258c6b", + "id": "522dfb203afcd2cd55799bf347f9b211919a7338", }, "example_bids_ses": { "uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses", @@ -109,3 +111,60 @@ def test_multiple_no_intersection() -> None: with dg: subs = [x for x in dg] assert set(subs) == set(expected_subs) + + +def test_multiple_get_item() -> None: + """Test a multiple datagrabber get_item error.""" + repo_uri1 = _testing_dataset["example_bids"]["uri"] + rootdir = "example_bids_ses" + replacements = ["subject", "session"] + pattern1 = { + "T1w": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", + } + dg1 = PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri1, + types=["T1w"], + patterns=pattern1, + replacements=replacements, + ) + + dg = MultipleDataGrabber([dg1]) + with pytest.raises(NotImplementedError): + dg.get_item(subject="sub-01") # type: ignore + + +def test_multiple_validation() -> None: + """Test a multiple datagrabber 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": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", + } + pattern2 = { + "bold": "{subject}/func/{subject}_task-rest_bold.nii.gz", + } + dg1 = PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri1, + types=["T1w"], + patterns=pattern1, + replacements=replacement1, + ) + + dg2 = PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri2, + types=["bold"], + patterns=pattern2, + replacements=replacement2, + ) + + with pytest.raises(ValueError, match="different element key"): + MultipleDataGrabber([dg1, dg2]) + + with pytest.raises(ValueError, match="overlapping types"): + MultipleDataGrabber([dg1, dg1]) diff --git a/junifer/datagrabber/tests/test_pattern.py b/junifer/datagrabber/tests/test_pattern.py index 472934d2c..551cc2571 100644 --- a/junifer/datagrabber/tests/test_pattern.py +++ b/junifer/datagrabber/tests/test_pattern.py @@ -105,7 +105,7 @@ def test_PatternDataGrabber_errors(tmp_path: Path) -> None: replacements=["subject", "session"], ) - with pytest.raises(ValueError, match="element length must be"): + with pytest.raises(ValueError, match="element keys must be"): datagrabber["sub001"] # This should not work, file does not exists diff --git a/junifer/datagrabber/tests/test_pattern_datalad.py b/junifer/datagrabber/tests/test_pattern_datalad.py index 7cacf7b60..ec8f3bb37 100644 --- a/junifer/datagrabber/tests/test_pattern_datalad.py +++ b/junifer/datagrabber/tests/test_pattern_datalad.py @@ -15,11 +15,13 @@ from junifer.datagrabber.pattern_datalad import PatternDataladDataGrabber _testing_dataset = { "example_bids": { "uri": "https://gin.g-node.org/juaml/datalad-example-bids", - "id": "e2ce149bd723088769a86c72e57eded009258c6b", + "commit": "522dfb203afcd2cd55799bf347f9b211919a7338", + "id": "fec92475-d9c0-4409-92ba-f041b6a12c40", }, "example_bids_ses": { "uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses", - "id": "3d08d55d1faad4f12ab64ac9497544a0d924d47a", + "commit": "3d08d55d1faad4f12ab64ac9497544a0d924d47a", + "id": "c83500d0-532f-45be-baf1-0dab703bdc2a", }, } @@ -56,7 +58,7 @@ def test_bids_PatternDataladDataGrabber(tmp_path: Path) -> None: repo_uri = _testing_dataset["example_bids"]["uri"] rootdir = "example_bids" - repo_commit = _testing_dataset["example_bids"]["id"] + repo_commit = _testing_dataset["example_bids"]["commit"] with PatternDataladDataGrabber( rootdir=rootdir, @@ -87,11 +89,11 @@ def test_bids_PatternDataladDataGrabber(tmp_path: Path) -> None: assert dg_meta["class"] == "PatternDataladDataGrabber" assert "uri" in dg_meta assert dg_meta["uri"] == repo_uri - assert "dataset_commit_id" in dg_meta - assert dg_meta["dataset_commit_id"] == repo_commit + assert "datalad_commit_id" in dg_meta + assert dg_meta["datalad_commit_id"] == repo_commit with open(t_sub["T1w"]["path"], "r") as f: - assert f.readlines()[0] == "placeholder" + assert f.readlines()[0].startswith("placeholder") def test_bids_PatternDataladDataGrabber_datadir(tmp_path: Path) -> None: diff --git a/junifer/testing/datagrabbers.py b/junifer/testing/datagrabbers.py index 57192f23a..7253ab7a7 100644 --- a/junifer/testing/datagrabbers.py +++ b/junifer/testing/datagrabbers.py @@ -24,13 +24,24 @@ class OasisVBMTestingDatagrabber(BaseDataGrabber): types = ["VBM_GM"] super().__init__(types=types, datadir=datadir) - def __getitem__(self, element: str) -> Dict: + def get_element_keys(self) -> List[str]: + """Get element keys. + + Returns + ------- + list of str + The element keys. + + """ + return ["subject"] + + def get_item(self, subject: str) -> Dict[str, Dict]: """Implement indexing support. Parameters ---------- - element : str - The element to retrieve. + subject : str + The subject to retrieve. Returns ------- @@ -38,11 +49,10 @@ class OasisVBMTestingDatagrabber(BaseDataGrabber): The data along with the metadata. """ - out = super().__getitem__(element) - i_sub = int(element.split("-")[1]) - 1 + out = {} + i_sub = int(subject.split("-")[1]) - 1 out["VBM_GM"] = {"path": Path(self._dataset.gray_matter_maps[i_sub])} - # Set the element accordingly - out["meta"]["element"] = {"subject": element} + return out def __enter__(self) -> "OasisVBMTestingDatagrabber": @@ -82,6 +92,17 @@ class SPMAuditoryTestingDatagrabber(BaseDataGrabber): types = ["BOLD", "T1w"] # TODO: Check that they are T1w super().__init__(types=types, datadir=datadir) + def get_element_keys(self) -> List[str]: + """Get element keys. + + Returns + ------- + list of str + The element keys. + + """ + return ["subject"] + def get_elements(self) -> List[str]: """Get elements. @@ -93,13 +114,13 @@ class SPMAuditoryTestingDatagrabber(BaseDataGrabber): """ return [f"sub{x:03d}" for x in list(range(1, 11))] - def __getitem__(self, element: str) -> Dict: + def get_item(self, subject: str) -> Dict[str, Dict]: """Implement indexing support. Parameters ---------- - element : str - The element to retrieve. + subject : str + The subject to retrieve. Returns ------- @@ -107,20 +128,17 @@ class SPMAuditoryTestingDatagrabber(BaseDataGrabber): The data along with the metadata. """ - out = super().__getitem__(element) - - nilearn_data = datasets.fetch_spm_auditory(subject_id=element) + out = {} + nilearn_data = datasets.fetch_spm_auditory(subject_id=subject) fmri_img = image.concat_imgs(nilearn_data.func) # type: ignore anat_img = image.concat_imgs(nilearn_data.anat) # type: ignore - fmri_fname = self.datadir / f"{element}_bold.nii.gz" - anat_fname = self.datadir / f"{element}_T1w.nii.gz" + fmri_fname = self.datadir / f"{subject}_bold.nii.gz" + anat_fname = self.datadir / f"{subject}_T1w.nii.gz" nib.save(fmri_img, fmri_fname) nib.save(anat_img, anat_fname) out["BOLD"] = {"path": fmri_fname} out["T1w"] = {"path": anat_fname} - # Set the element accordingly - out["meta"]["element"] = {"subject": element} return out @@ -176,6 +194,17 @@ class PartlyCloudyTestingDataGrabber(BaseDataGrabber): ) return self + def get_element_keys(self) -> List[str]: + """Get element keys. + + Returns + ------- + list of str + The element keys. + + """ + return ["subject"] + def get_elements(self) -> List[str]: """Get elements. @@ -187,13 +216,13 @@ class PartlyCloudyTestingDataGrabber(BaseDataGrabber): """ return [f"sub-{x:02d}" for x in list(range(1, 11))] - def __getitem__(self, element: str) -> Dict: + def get_item(self, subject: str) -> Dict[str, Dict]: """Implement indexing support. Parameters ---------- - element : str - The element to retrieve. + subject : str + The subject to retrieve. Returns ------- @@ -201,8 +230,8 @@ class PartlyCloudyTestingDataGrabber(BaseDataGrabber): The data along with the metadata. """ - out = super().__getitem__(element) - i_sub = int(element.split("-")[1]) - 1 + out = {} + i_sub = int(subject.split("-")[1]) - 1 out["BOLD"] = {"path": Path(self._dataset["func"][i_sub])} conf_format = "fmriprep" @@ -211,6 +240,4 @@ class PartlyCloudyTestingDataGrabber(BaseDataGrabber): "format": conf_format, } - # Set the element accordingly - out["meta"]["element"] = {"subject": element} return out diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index 0cceed96b..679f0797f 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -11,10 +11,13 @@ from pathlib import Path from subprocess import PIPE, Popen, TimeoutExpired from typing import Dict, NoReturn, Optional, Type, Union from warnings import warn - +import datalad logger = logging.getLogger("JUNIFER") +# Set up datalad logger level to warning by default +datalad.log.lgr.setLevel(logging.WARNING) + _logging_types = { "DEBUG": logging.DEBUG, "INFO": logging.INFO, @@ -253,6 +256,7 @@ def configure_logging( lh.setFormatter(formatter) # set formatter logger.setLevel(level) # set level + datalad.log.lgr.setLevel(level) # set level for datalad logger.addHandler(lh) # set handler log_versions() # log versions of installed packages diff --git a/tools/create_bids_example_dataset.py b/tools/create_bids_example_dataset.py index e985bc29c..fc665ca05 100644 --- a/tools/create_bids_example_dataset.py +++ b/tools/create_bids_example_dataset.py @@ -27,7 +27,7 @@ with TemporaryDirectory() as tmpdir_name: f'func/{t_sub}_task-rest_bold.json'] for fname in fnames: with open(sub_dir / fname, 'w') as f: - f.write('placeholder') + f.write(f'placeholder-{fname}') ds.save(recursive=True) ds.siblings('add', name='gin', url=dst) diff --git a/tools/create_bids_example_dataset_sessions.py b/tools/create_bids_example_dataset_sessions.py index 3301c6f8a..9d798b2c5 100644 --- a/tools/create_bids_example_dataset_sessions.py +++ b/tools/create_bids_example_dataset_sessions.py @@ -34,7 +34,7 @@ with TemporaryDirectory() as tmpdir_name: f'func/{t_sub}_{t_ses}_task-rest_bold.json']) for fname in fnames: with open(ses_dir / fname, 'w') as f: - f.write('placeholder') + f.write('placeholder-{fname}') ds.save(recursive=True) ds.siblings('add', name='gin', url=dst)