Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 8 additions & 9 deletions niworkflows/utils/bids.py
Original file line number Diff line number Diff line change
Expand Up @@ -189,8 +189,9 @@ def collect_data(
The BIDS directory
participant_label : :obj:`str`
The participant identifier
session_id : :obj:`str`, None, or :obj:`bids.layout.Query`
The session identifier. By default, all sessions will be used.
session_id : :obj:`str`, :obj:`list`, None, or :obj:`bids.layout.Query`
The session identifier(s). By default, all sessions will be used.
If specified, BIDS filters may only select sessions among these.
task : :obj:`str` or None
The task identifier (for BOLD queries)
echo : :obj:`int` or None
Expand Down Expand Up @@ -273,29 +274,27 @@ def collect_data(
for acq, entities in bids_filters.items():
if acq not in queries: # filter with no matching query
continue
# BIDS filters will not be able to override subject / session entities
# BIDS filters may narrow down, but not override, subject / session entities
for entity, param in reserved_entities:
if param == Query.OPTIONAL:
continue
if entity in entities and listify(param) != listify(entities[entity]):
if entity in entities and not set(listify(entities[entity])) <= set(listify(param)):
raise ValueError(
f'Conflicting entities for "{entity}" found: {entities[entity]} // {param}'
)

queries[acq].update(entities)
for entity in list(layout_get_kwargs.keys()):
if entity in entities:
# avoid clobbering layout.get
del layout_get_kwargs[entity]

if task:
queries['bold']['task'] = queries['pet']['task'] = task

if echo:
queries['bold']['echo'] = echo

# Query entities take precedence over the defaults, for that query only
subj_data = {
dtype: sorted(layout.get(**layout_get_kwargs, **query)) for dtype, query in queries.items()
dtype: sorted(layout.get(**{**layout_get_kwargs, **query}))
for dtype, query in queries.items()
}

# Special case: multi-echo BOLD, grouping echos
Expand Down
59 changes: 59 additions & 0 deletions niworkflows/utils/tests/test_bids.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,9 @@
"""Tests for :mod:`niworkflows.utils.bids`."""

from pathlib import Path

import pytest

from ..bids import collect_data
from ..testing import generate_bids_skeleton

Expand All @@ -8,6 +12,17 @@
'01': [{'anat': [{'suffix': 'T1w', 'metadata': {'EchoTime': 1}}]}],
}

SESSIONS_SKELETON = {
'01': [
{
'session': session,
'anat': [{'suffix': 'T1w'}],
'func': [{'task': 'rest', 'suffix': 'bold', 'metadata': {'RepetitionTime': 2.0}}],
}
for session in ('anat', 'func1', 'func2')
],
}


def test_collect_data_ignores_filter_without_query(tmp_path):
"""A filter naming a suffix absent from ``queries`` is dropped, not an error."""
Expand All @@ -22,3 +37,47 @@ def test_collect_data_ignores_filter_without_query(tmp_path):
bids_filters={'pet': {'session': '15'}},
)
assert len(subj_data['t1w']) == 1


@pytest.mark.parametrize(
('session_id', 'bold_session', 't1w', 'bold'),
[
(None, 'func2', ['anat', 'func1', 'func2'], ['func2']),
(['anat', 'func1'], 'func1', ['anat', 'func1'], ['func1']),
('func1', 'func1', ['func1'], ['func1']),
],
ids=['no_session_id', 'narrow_sessions', 'clobber_session'],
)
def test_collect_data_session_filters(tmp_path, session_id, bold_session, t1w, bold):
"""BIDS filters narrow down the requested sessions for their query only."""
root = tmp_path / 'bids'
generate_bids_skeleton(root, SESSIONS_SKELETON)

subj_data, _ = collect_data(
root,
'01',
session_id=session_id,
bids_validate=False,
bids_filters={'bold': {'session': bold_session}},
)

def _sessions(files):
return [Path(f).parts[-3].removeprefix('ses-') for f in files]

assert _sessions(subj_data['t1w']) == t1w
assert _sessions(subj_data['bold']) == bold


def test_collect_data_session_filter_conflict(tmp_path):
"""BIDS filters cannot select sessions outside of the requested ones."""
root = tmp_path / 'bids'
generate_bids_skeleton(root, SESSIONS_SKELETON)

with pytest.raises(ValueError, match='Conflicting entities for "session"'):
collect_data(
root,
'01',
session_id=['anat', 'func1'],
bids_validate=False,
bids_filters={'bold': {'session': 'func2'}},
)
Loading