From 4b82fb5cb8a22fb4fde9909c5ea0aebf5494365c Mon Sep 17 00:00:00 2001 From: frheault Date: Wed, 11 Oct 2023 17:33:18 -0400 Subject: [PATCH 01/11] lowercase --- trx/trx_file_memmap.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/trx/trx_file_memmap.py b/trx/trx_file_memmap.py index ffecf94b..24d6f70e 100644 --- a/trx/trx_file_memmap.py +++ b/trx/trx_file_memmap.py @@ -541,13 +541,13 @@ def save( the ZipFile library """ _, ext = os.path.splitext(filename) - if ext not in [".zip", ".trx", ""]: + if ext.lower() not in [".zip", ".trx", ""]: raise ValueError("Unsupported extension.") copy_trx = trx.deepcopy() copy_trx.resize() tmp_dir_name = copy_trx._uncompressed_folder_handle.name - if ext in [".zip", ".trx"]: + if ext.lower() in [".zip", ".trx"]: zip_from_folder(tmp_dir_name, filename, compression_standard) else: if os.path.isdir(filename): From ded3b3525436548888b79956e9fda3c188c05982 Mon Sep 17 00:00:00 2001 From: frheault Date: Mon, 23 Feb 2026 12:05:33 -0500 Subject: [PATCH 02/11] Added mock for utils --- trx/tests/test_utils.py | 526 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 526 insertions(+) create mode 100644 trx/tests/test_utils.py diff --git a/trx/tests/test_utils.py b/trx/tests/test_utils.py new file mode 100644 index 00000000..bc8f0d42 --- /dev/null +++ b/trx/tests/test_utils.py @@ -0,0 +1,526 @@ +# -*- coding: utf-8 -*- +"""Tests for utility functions in trx.utils.""" + +import logging +import os +import tempfile +from unittest.mock import MagicMock, patch + +import nibabel as nib +from nibabel.streamlines.array_sequence import ArraySequence +from nibabel.streamlines.tractogram import TractogramItem, Tractogram +import numpy as np +import pytest + +from trx.utils import ( + close_or_delete_mmap, + split_name_with_gz, + load_matrix_in_any_format, + get_reference_info_wrapper, + is_header_compatible, + get_axis_shift_vector, + get_axis_flip_vector, + get_shift_vector, + flip_sft, + get_reverse_enum, + convert_data_dict_to_tractogram, + append_generator_to_dict, + verify_trx_dtype, +) + +# Optional dipy import +try: + import dipy + from dipy.io.stateful_tractogram import StatefulTractogram, Space, Origin + dipy_available = True +except ImportError: + dipy_available = False + + +def test_close_or_delete_mmap_np_memmap(): + """Test close_or_delete_mmap with a numpy.memmap.""" + with tempfile.NamedTemporaryFile() as tmp: + mmap_arr = np.memmap(tmp.name, dtype='float32', mode='w+', shape=(10,)) + close_or_delete_mmap(mmap_arr) + + +def test_close_or_delete_mmap_array_sequence(): + """Test close_or_delete_mmap with an ArraySequence.""" + with (tempfile.NamedTemporaryFile() as tmp1, + tempfile.NamedTemporaryFile() as tmp2): + + data = np.memmap(tmp1.name, dtype='float32', mode='w+', shape=(10, 3)) + offsets = np.memmap(tmp2.name, dtype='uint32', mode='w+', shape=(5,)) + + seq = ArraySequence() + seq._data = data + seq._offsets = offsets + seq._lengths = np.array([2, 2, 2, 2, 2], dtype='uint32') + + close_or_delete_mmap(seq) + + +def test_close_or_delete_mmap_with_mmap_attr(): + """Test close_or_delete_mmap with an object having _mmap attribute.""" + mock_obj = MagicMock() + mock_mmap = MagicMock() + mock_obj._mmap = mock_mmap + + close_or_delete_mmap(mock_obj) + mock_mmap.close.assert_called_once() + + +def test_close_or_delete_mmap_other_type(caplog): + """Test close_or_delete_mmap with an unsupported type.""" + with caplog.at_level(logging.DEBUG): + close_or_delete_mmap("not a memmap") + assert "Object to be close or deleted must be np.memmap" in caplog.text + + +@pytest.mark.parametrize("filename,expected_base,expected_ext", [ + ("test.nii.gz", "test", ".nii.gz"), + ("test.trk.gz", "test", ".trk.gz"), + ("test.nii", "test", ".nii"), + ("test.trk", "test", ".trk"), + ("test.txt", "test", ".txt"), + ("my.file.with.dots.nii.gz", "my.file.with.dots", ".nii.gz"), + ("no_ext", "no_ext", ""), +]) +def test_split_name_with_gz(filename, expected_base, expected_ext): + """Test split_name_with_gz with various extensions.""" + base, ext = split_name_with_gz(filename) + assert base == expected_base + assert ext == expected_ext + + +def test_load_matrix_in_any_format_txt(): + """Test loading a matrix from a .txt file.""" + with tempfile.NamedTemporaryFile(suffix=".txt", mode="w", delete=False) as tmp: + tmp.write("1 2 3\n4 5 6") + tmp_name = tmp.name + + try: + matrix = load_matrix_in_any_format(tmp_name) + np.testing.assert_allclose(matrix, [[1, 2, 3], [4, 5, 6]]) + finally: + os.remove(tmp_name) + + +def test_load_matrix_in_any_format_npy(): + """Test loading a matrix from a .npy file.""" + with tempfile.NamedTemporaryFile(suffix=".npy", delete=False) as tmp: + data = np.array([[1, 2], [3, 4]]) + np.save(tmp.name, data) + tmp_name = tmp.name + + try: + matrix = load_matrix_in_any_format(tmp_name) + np.testing.assert_array_equal(matrix, data) + finally: + os.remove(tmp_name) + + +def test_load_matrix_in_any_format_error(): + """Test load_matrix_in_any_format with unsupported extension.""" + with pytest.raises(ValueError, match="Extension .invalid is not supported"): + load_matrix_in_any_format("test.invalid") + + +# --- Spatial Reference Tests --- + +@pytest.fixture +def nifti_ref(): + """Create a synthetic Nifti1Image for testing.""" + data = np.zeros((10, 20, 30), dtype=np.float32) + affine = np.diag([1.0, 2.0, 3.0, 1.0]) + affine[0:3, 3] = [1.1, 2.2, 3.3] + img = nib.Nifti1Image(data, affine) + return img + + +@pytest.fixture +def trk_header(): + """Create a synthetic TRK header for testing.""" + return { + 'voxel_to_rasmm': np.diag([1.0, 2.0, 3.0, 1.0]), + 'dimensions': np.array([10, 20, 30], dtype=np.int16), + 'voxel_sizes': np.array([1.0, 2.0, 3.0], dtype=np.float32), + 'voxel_order': 'RAS', + 'magic_number': 'TRACK' + } + + +def test_get_reference_info_wrapper_nifti_obj(nifti_ref): + """Test get_reference_info_wrapper with a Nifti1Image object.""" + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(nifti_ref) + assert np.allclose(affine, nifti_ref.affine) + assert np.array_equal(dimensions, [10, 20, 30]) + assert np.allclose(voxel_sizes, [1.0, 2.0, 3.0]) + assert voxel_order == 'RAS' + + +def test_get_reference_info_wrapper_nifti_header(nifti_ref): + """Test get_reference_info_wrapper with a Nifti1Header object.""" + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(nifti_ref.header) + assert np.allclose(affine, nifti_ref.affine) + assert np.array_equal(dimensions, [10, 20, 30]) + + +def test_get_reference_info_wrapper_nifti_file(nifti_ref): + """Test get_reference_info_wrapper with a Nifti filename.""" + with tempfile.TemporaryDirectory() as tmp_dir: + path = os.path.join(tmp_dir, "test.nii.gz") + nib.save(nifti_ref, path) + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(path) + assert np.allclose(affine, nifti_ref.affine) + assert np.array_equal(dimensions, [10, 20, 30]) + + +@patch('nibabel.streamlines.load') +def test_get_reference_info_wrapper_trk_file(mock_load): + """Test get_reference_info_wrapper with a TRK filename.""" + mock_trk = MagicMock() + mock_trk.header = { + 'voxel_to_rasmm': np.diag([1.0, 2.0, 3.0, 1.0]), + 'dimensions': np.array([10, 20, 30], dtype=np.int16), + 'voxel_sizes': np.array([1.0, 2.0, 3.0], dtype=np.float32), + 'voxel_order': 'RAS', + 'magic_number': 'TRACK' + } + mock_load.return_value = mock_trk + + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper("test.trk") + assert np.allclose(affine, mock_trk.header['voxel_to_rasmm']) + + +def test_get_reference_info_wrapper_trk_obj(): + """Test get_reference_info_wrapper with a TrkFile object.""" + mock_trk = MagicMock(spec=nib.streamlines.trk.TrkFile) + mock_trk.header = { + 'voxel_to_rasmm': np.diag([1.0, 2.0, 3.0, 1.0]), + 'dimensions': np.array([10, 20, 30], dtype=np.int16), + 'voxel_sizes': np.array([1.0, 2.0, 3.0], dtype=np.float32), + 'voxel_order': 'RAS', + 'magic_number': 'TRACK' + } + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(mock_trk) + assert np.allclose(affine, mock_trk.header['voxel_to_rasmm']) + + +def test_get_reference_info_wrapper_trk_dict(trk_header): + """Test get_reference_info_wrapper with a TRK header dict.""" + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(trk_header) + assert np.allclose(affine, trk_header['voxel_to_rasmm']) + assert np.array_equal(dimensions, trk_header['dimensions']) + + +@pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") +def test_get_reference_info_wrapper_sft(nifti_ref): + """Test get_reference_info_wrapper with a StatefulTractogram.""" + streamlines = [np.array([[0, 0, 0], [1, 1, 1]], dtype=np.float32)] + sft = StatefulTractogram(streamlines, nifti_ref, Space.RASMM) + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(sft) + assert np.allclose(affine, nifti_ref.affine) + assert np.array_equal(dimensions, [10, 20, 30]) + + +def test_get_reference_info_wrapper_trx_obj(): + """Test get_reference_info_wrapper with a TrxFile object mock.""" + from trx.trx_file_memmap import TrxFile + mock_trx = MagicMock(spec=TrxFile) + mock_trx.header = { + 'VOXEL_TO_RASMM': np.diag([1.0, 1.0, 1.0, 1.0]), + 'DIMENSIONS': np.array([10, 10, 10], dtype=np.uint16) + } + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(mock_trx) + assert np.allclose(affine, mock_trx.header['VOXEL_TO_RASMM']) + assert np.array_equal(dimensions, mock_trx.header['DIMENSIONS']) + + +@patch('trx.trx_file_memmap.load') +def test_get_reference_info_wrapper_trx_file(mock_load): + """Test get_reference_info_wrapper with a TRX filename.""" + mock_trx = MagicMock() + mock_trx.header = { + 'VOXEL_TO_RASMM': np.diag([1.0, 1.0, 1.0, 1.0]), + 'DIMENSIONS': np.array([10, 10, 10], dtype=np.uint16) + } + mock_load.return_value = mock_trx + + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper("test.trx") + assert np.allclose(affine, mock_trx.header['VOXEL_TO_RASMM']) + + +def test_get_reference_info_wrapper_trx_dict(): + """Test get_reference_info_wrapper with a TRX header dict.""" + header = { + 'VOXEL_TO_RASMM': np.diag([1.0, 1.0, 1.0, 1.0]), + 'DIMENSIONS': np.array([10, 10, 10], dtype=np.uint16), + 'NB_VERTICES': 0 + } + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(header) + assert np.allclose(affine, header['VOXEL_TO_RASMM']) + assert np.array_equal(dimensions, header['DIMENSIONS']) + + +def test_get_reference_info_wrapper_zero_affine(): + """Test get_reference_info_wrapper with an all-zero affine.""" + mock_img = MagicMock(spec=nib.Nifti1Image) + mock_header = MagicMock(spec=nib.Nifti1Header) + mock_img.header = mock_header + mock_header.get_best_affine.return_value = np.zeros((4, 4)) + mock_header.__getitem__.side_effect = lambda x: [10, 10, 10] if x == 'dim' else [1, 1, 1] + + with pytest.raises(ValueError, match="Invalid affine, contains only zeros"): + get_reference_info_wrapper(mock_img) + + +def test_get_reference_info_wrapper_binary_order(): + """Test get_reference_info_wrapper with binary voxel order.""" + header = { + 'voxel_to_rasmm': np.diag([1.0, 1.0, 1.0, 1.0]), + 'dimensions': [10, 10, 10], + 'voxel_sizes': [1, 1, 1], + 'voxel_order': np.bytes_(b'RAS'), # numpy bytes + 'magic_number': 'TRACK' + } + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(header) + assert voxel_order == 'RAS' + + +def test_get_reference_info_wrapper_error(): + """Test get_reference_info_wrapper with unsupported type.""" + with pytest.raises(TypeError, match="Input reference is not one of the supported format"): + get_reference_info_wrapper(123) + + +def test_is_header_compatible_identical(nifti_ref): + """Test is_header_compatible with identical headers.""" + assert is_header_compatible(nifti_ref, nifti_ref) + + +def test_is_header_compatible_different(nifti_ref): + """Test is_header_compatible with different headers.""" + data2 = np.zeros((10, 20, 31), dtype=np.float32) + img2 = nib.Nifti1Image(data2, nifti_ref.affine) + assert not is_header_compatible(nifti_ref, img2) + + +def test_is_header_compatible_affine_diff(nifti_ref, caplog): + """Test is_header_compatible with different affines.""" + affine2 = nifti_ref.affine.copy() + affine2[0, 0] = 5.0 + img2 = nib.Nifti1Image(nifti_ref.get_fdata(), affine2) + with caplog.at_level(logging.ERROR): + assert not is_header_compatible(nifti_ref, img2) + assert "Affine not equal" in caplog.text or "Voxel_size not equal" in caplog.text + + +def test_is_header_compatible_order_diff(caplog): + """Test is_header_compatible with different voxel orders.""" + affine1 = np.diag([1.0, 1.0, 1.0, 1.0]) + affine2 = np.diag([-1.0, 1.0, 1.0, 1.0]) # LAS instead of RAS + + header1 = { + 'voxel_to_rasmm': affine1, + 'dimensions': [10, 10, 10], + 'voxel_sizes': [1, 1, 1], + 'voxel_order': 'RAS', + 'magic_number': 'TRACK' + } + header2 = header1.copy() + header2['voxel_to_rasmm'] = affine2 + header2['voxel_order'] = 'LAS' + + with caplog.at_level(logging.ERROR): + assert not is_header_compatible(header1, header2) + assert "Voxel_order not equal" in caplog.text + + +# --- Transformation & Vector Tests --- + +def test_get_axis_shift_vector(): + """Test get_axis_shift_vector.""" + assert np.array_equal(get_axis_shift_vector(['x', 'y']), [-1.0, -1.0, 0.0]) + assert np.array_equal(get_axis_shift_vector(['z']), [0.0, 0.0, -1.0]) + + +def test_get_axis_flip_vector(): + """Test get_axis_flip_vector.""" + assert np.array_equal(get_axis_flip_vector(['x', 'z']), [-1.0, 1.0, -1.0]) + assert np.array_equal(get_axis_flip_vector([]), [1.0, 1.0, 1.0]) + + +@pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") +def test_get_shift_vector(nifti_ref): + """Test get_shift_vector.""" + sft = StatefulTractogram([], nifti_ref, Space.RASMM) + shift = get_shift_vector(sft) + assert np.array_equal(shift, [-5.0, -10.0, -15.0]) + + +@pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") +def test_flip_sft(nifti_ref): + """Test flip_sft.""" + streamlines = [np.array([[0, 0, 0], [1, 1, 1]], dtype=np.float32)] + sft = StatefulTractogram(streamlines, nifti_ref, Space.VOX) + + # Flipping X axis. Center of X is 5.0 (dim[0]=10). + # 0 -> (0 - 5) * -1 - (-5) = -5 * -1 + 5 = 10 + # 1 -> (1 - 5) * -1 - (-5) = -4 * -1 + 5 = 9 + + flipped_sft = flip_sft(sft, ['x']) + assert np.allclose(flipped_sft.streamlines[0][0, 0], 10.0) + assert np.allclose(flipped_sft.streamlines[0][1, 0], 9.0) + + +@patch('trx.utils.dipy_available', False) +def test_flip_sft_no_dipy(caplog): + """Test flip_sft when dipy is missing.""" + with caplog.at_level(logging.ERROR): + result = flip_sft(None, ['x']) + assert result is None + assert "Dipy library is missing" in caplog.text + + +@pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") +@pytest.mark.parametrize("space_str,origin_str,expected_space,expected_origin", [ + ("rasmm", "nifti", Space.RASMM, Origin.NIFTI), + ("voxmm", "trackvis", Space.VOXMM, Origin.TRACKVIS), + ("vox", "nifti", Space.VOX, Origin.NIFTI), +]) +def test_get_reverse_enum(space_str, origin_str, expected_space, expected_origin): + """Test get_reverse_enum.""" + space, origin = get_reverse_enum(space_str, origin_str) + assert space == expected_space + assert origin == expected_origin + + +@patch('trx.utils.dipy_available', False) +def test_get_reverse_enum_no_dipy(caplog): + """Test get_reverse_enum when dipy is missing.""" + with caplog.at_level(logging.ERROR): + result = get_reverse_enum("rasmm", "nifti") + assert result is None + assert "Dipy library is missing" in caplog.text + + +# --- Data Conversion & Dtype Verification Tests --- + +def test_convert_data_dict_to_tractogram(): + """Test convert_data_dict_to_tractogram.""" + data = { + 'strs': [np.array([[0, 0, 0], [1, 1, 1]]), np.array([[2, 2, 2]])], + 'dps': {'test_dps': [1, 2]}, + 'dpv': {'test_dpv': [0.1, 0.2, 0.3]} + } + obj = convert_data_dict_to_tractogram(data) + assert isinstance(obj, nib.streamlines.tractogram.Tractogram) + assert len(obj.streamlines) == 2 + assert np.array_equal(obj.data_per_streamline['test_dps'], [[1], [2]]) + # Data per vertex is returned as ArraySequence + assert np.allclose(obj.data_per_point['test_dpv'][0], [[0.1], [0.2]]) + + +def test_append_generator_to_dict_array(): + """Test append_generator_to_dict with numpy array.""" + data = {'strs': [], 'dpv': {}, 'dps': {}} + append_generator_to_dict(np.array([[0, 0, 0]]), data) + assert len(data['strs']) == 1 + + +def test_append_generator_to_dict_item(): + """Test append_generator_to_dict with TractogramItem.""" + data = {'strs': [], 'dpv': {}, 'dps': {}} + # TractogramItem(streamline, data_for_streamline=None, data_for_points=None) + item = TractogramItem( + np.array([[0, 0, 0]]), + {'s': 1}, + {'v': [0.1]} + ) + append_generator_to_dict(item, data) + assert len(data['strs']) == 1 + assert 'v' in data['dpv'] + assert 's' in data['dps'] + + +def test_verify_trx_dtype(): + """Test verify_trx_dtype.""" + # Create a mock TRX object + mock_trx = MagicMock(spec=Tractogram) + mock_trx.streamlines._data.dtype = np.float32 + mock_trx.streamlines._offsets.dtype = np.uint32 + + mock_dpv = MagicMock() + mock_dpv._data.dtype = np.uint16 + mock_trx.data_per_vertex = {'v1': mock_dpv} + + mock_trx.data_per_streamline = {'s1': np.array([1], dtype='int16')} + + # Define expected dtype dict + dtype_dict = { + 'positions': np.float32, + 'offsets': np.uint32, + 'dpv': {'v1': np.uint16}, + 'dps': {'s1': np.int16} + } + + assert verify_trx_dtype(mock_trx, dtype_dict) + + # Test mismatches for warnings + with patch('logging.warning') as mock_log: + dtype_dict['positions'] = np.float64 + assert not verify_trx_dtype(mock_trx, dtype_dict) + mock_log.assert_any_call("Positions dtype is different") + + dtype_dict['positions'] = np.float32 + dtype_dict['offsets'] = np.uint64 + assert not verify_trx_dtype(mock_trx, dtype_dict) + mock_log.assert_any_call("Offsets dtype is different") + + dtype_dict['offsets'] = np.uint32 + dtype_dict['dpv']['v1'] = np.uint32 + assert not verify_trx_dtype(mock_trx, dtype_dict) + mock_log.assert_any_call("Data per vertex (v1) dtype is different") + + dtype_dict['dpv']['v1'] = np.uint16 + dtype_dict['dps']['s1'] = np.int32 + assert not verify_trx_dtype(mock_trx, dtype_dict) + mock_log.assert_any_call("Data per streamline (s1) dtype is different") + + +def test_verify_trx_dtype_groups(): + """Test verify_trx_dtype with groups and dpg.""" + mock_trx = MagicMock(spec=Tractogram) + mock_trx.streamlines._data.dtype = np.float32 + mock_trx.streamlines._offsets.dtype = np.uint32 + + mock_g1 = MagicMock() + mock_g1._data.dtype = np.int32 + + mock_dpg_val = MagicMock() + mock_dpg_val.dtype = np.float32 + + # verify_trx_dtype expects trx.data_per_point to contain groups and dpg + mock_trx.data_per_point = { + 'g1': mock_g1, + 'g2': {'d1': mock_dpg_val} + } + + dtype_dict = { + 'groups': {'g1': np.int32}, + 'dpg': {'g2': {'d1': np.float32}} + } + + assert verify_trx_dtype(mock_trx, dtype_dict) + + # Test mismatches for dpg and groups + with patch('logging.warning') as mock_log: + dtype_dict['groups']['g1'] = np.int16 + assert not verify_trx_dtype(mock_trx, dtype_dict) + mock_log.assert_any_call("Data per group (g1) dtype is different") + + dtype_dict['groups']['g1'] = np.int32 + dtype_dict['dpg']['g2']['d1'] = np.float64 + assert not verify_trx_dtype(mock_trx, dtype_dict) + mock_log.assert_any_call("Data per group (d1) dtype is different") From 6ec4ddaae203a5c22aae8e0ef879db1b80b70cba Mon Sep 17 00:00:00 2001 From: frheault Date: Tue, 11 Aug 2026 09:24:38 -0400 Subject: [PATCH 03/11] Update mock and pre-commit check --- trx/io.py | 16 +- trx/tests/test_io.py | 27 +++- trx/tests/test_utils.py | 346 +++++++++++++++++++++------------------- 3 files changed, 219 insertions(+), 170 deletions(-) diff --git a/trx/io.py b/trx/io.py index 3d633bba..6b77d755 100644 --- a/trx/io.py +++ b/trx/io.py @@ -46,7 +46,7 @@ def get_trx_tmp_dir(): return tempfile.TemporaryDirectory(dir=trx_tmp_dir, prefix="trx_") -def load_sft_with_reference(filepath, reference=None, bbox_check=True): +def load_sft_with_reference(filepath, reference=None, bbox_check=True, **kwargs): """Load a tractogram as a StatefulTractogram with an explicit reference. Parameters @@ -59,6 +59,8 @@ def load_sft_with_reference(filepath, reference=None, bbox_check=True): bbox_check : bool, optional If True, validate that streamlines lie within the reference bounding box. Defaults to True. + **kwargs + Additional keyword arguments passed to dipy's load_tractogram. Returns ------- @@ -83,12 +85,14 @@ def load_sft_with_reference(filepath, reference=None, bbox_check=True): if ext == ".trk": if reference is not None and reference != "same": logging.warning(f"Reference is discarded for this file format {filepath}.") - sft = load_tractogram(filepath, "same", bbox_valid_check=bbox_check) + sft = load_tractogram(filepath, "same", bbox_valid_check=bbox_check, **kwargs) elif ext in [".tck", ".fib", ".vtk", ".dpy"]: if reference is None or reference == "same": raise IOError(f"--reference is required for this file format {filepath}.") else: - sft = load_tractogram(filepath, reference, bbox_valid_check=bbox_check) + sft = load_tractogram( + filepath, reference, bbox_valid_check=bbox_check, **kwargs + ) else: raise IOError(f"{filepath} is an unsupported file format") @@ -96,7 +100,7 @@ def load_sft_with_reference(filepath, reference=None, bbox_check=True): return sft -def load(tractogram_filename, reference): +def load(tractogram_filename, reference, **kwargs): """Load a tractogram from disk and return a TRX or StatefulTractogram. Parameters @@ -105,6 +109,8 @@ def load(tractogram_filename, reference): Path to the input tractogram. TRX directories are supported. reference : str or nibabel.Nifti1Image Reference image used for formats without embedded affine information. + **kwargs + Additional keyword arguments passed to dipy's load_tractogram. Returns ------- @@ -116,7 +122,7 @@ def load(tractogram_filename, reference): in_ext = split_name_with_gz(tractogram_filename)[1] if in_ext != ".trx" and not os.path.isdir(tractogram_filename): tractogram_obj = load_sft_with_reference( - tractogram_filename, reference, bbox_check=False + tractogram_filename, reference, bbox_check=False, **kwargs ) else: tractogram_obj = tmm.load(tractogram_filename) diff --git a/trx/tests/test_io.py b/trx/tests/test_io.py index 3db89a83..7d7c301a 100644 --- a/trx/tests/test_io.py +++ b/trx/tests/test_io.py @@ -17,6 +17,13 @@ except ImportError: dipy_available = False +try: + import fury # noqa: F401 + + fury_available = True +except ImportError: + fury_available = False + from trx.fetcher import fetch_data, get_home, get_testing_files_dict from trx.io import load, save import trx.trx_file_memmap as tmm @@ -28,6 +35,8 @@ @pytest.mark.parametrize("path", [("gs.trk"), ("gs.tck"), ("gs.vtk")]) @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_seq_ops_sft(path): + if path.endswith(".vtk") and not fury_available: + pytest.skip("fury is not installed") with TemporaryDirectory() as tmp_dir: gs_dir = os.path.join(get_home(), "gold_standard") path = os.path.join(tmp_dir, path) @@ -56,10 +65,17 @@ def test_seq_ops_trx(): @pytest.mark.parametrize("path", [("gs.trx"), ("gs.trk"), ("gs.tck"), ("gs.vtk")]) @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_load_vox(path): + if path.endswith(".vtk") and not fury_available: + pytest.skip("fury is not installed") gs_dir = os.path.join(get_home(), "gold_standard") path = os.path.join(gs_dir, path) coord = np.loadtxt(os.path.join(get_home(), "gold_standard", "gs_vox_space.txt")) - obj = load(path, os.path.join(gs_dir, "gs.nii")) + if path.endswith(".vtk"): + from dipy.io.stateful_tractogram import Space + + obj = load(path, os.path.join(gs_dir, "gs.nii"), from_space=Space.LPSMM) + else: + obj = load(path, os.path.join(gs_dir, "gs.nii")) sft = obj.to_sft() if isinstance(obj, TrxFile) else obj sft.to_vox() @@ -72,10 +88,17 @@ def test_load_vox(path): @pytest.mark.parametrize("path", [("gs.trx"), ("gs.trk"), ("gs.tck"), ("gs.vtk")]) @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_load_voxmm(path): + if path.endswith(".vtk") and not fury_available: + pytest.skip("fury is not installed") gs_dir = os.path.join(get_home(), "gold_standard") path = os.path.join(gs_dir, path) coord = np.loadtxt(os.path.join(get_home(), "gold_standard", "gs_voxmm_space.txt")) - obj = load(path, os.path.join(gs_dir, "gs.nii")) + if path.endswith(".vtk"): + from dipy.io.stateful_tractogram import Space + + obj = load(path, os.path.join(gs_dir, "gs.nii"), from_space=Space.LPSMM) + else: + obj = load(path, os.path.join(gs_dir, "gs.nii")) sft = obj.to_sft() if isinstance(obj, TrxFile) else obj sft.to_voxmm() diff --git a/trx/tests/test_utils.py b/trx/tests/test_utils.py index bc8f0d42..d6727efe 100644 --- a/trx/tests/test_utils.py +++ b/trx/tests/test_utils.py @@ -8,55 +8,63 @@ import nibabel as nib from nibabel.streamlines.array_sequence import ArraySequence -from nibabel.streamlines.tractogram import TractogramItem, Tractogram +from nibabel.streamlines.tractogram import Tractogram, TractogramItem import numpy as np import pytest from trx.utils import ( + append_generator_to_dict, close_or_delete_mmap, - split_name_with_gz, - load_matrix_in_any_format, - get_reference_info_wrapper, - is_header_compatible, - get_axis_shift_vector, - get_axis_flip_vector, - get_shift_vector, + convert_data_dict_to_tractogram, flip_sft, + get_axis_flip_vector, + get_axis_shift_vector, + get_reference_info_wrapper, get_reverse_enum, - convert_data_dict_to_tractogram, - append_generator_to_dict, + get_shift_vector, + is_header_compatible, + load_matrix_in_any_format, + split_name_with_gz, verify_trx_dtype, ) # Optional dipy import try: - import dipy - from dipy.io.stateful_tractogram import StatefulTractogram, Space, Origin + import dipy # noqa: F401 + from dipy.io.stateful_tractogram import Origin, Space, StatefulTractogram + dipy_available = True except ImportError: dipy_available = False + class Space: + RASMM = None + VOXMM = None + VOX = None + + class Origin: + NIFTI = None + TRACKVIS = None + def test_close_or_delete_mmap_np_memmap(): """Test close_or_delete_mmap with a numpy.memmap.""" with tempfile.NamedTemporaryFile() as tmp: - mmap_arr = np.memmap(tmp.name, dtype='float32', mode='w+', shape=(10,)) + mmap_arr = np.memmap(tmp.name, dtype="float32", mode="w+", shape=(10,)) close_or_delete_mmap(mmap_arr) def test_close_or_delete_mmap_array_sequence(): """Test close_or_delete_mmap with an ArraySequence.""" - with (tempfile.NamedTemporaryFile() as tmp1, - tempfile.NamedTemporaryFile() as tmp2): - - data = np.memmap(tmp1.name, dtype='float32', mode='w+', shape=(10, 3)) - offsets = np.memmap(tmp2.name, dtype='uint32', mode='w+', shape=(5,)) - + with tempfile.NamedTemporaryFile() as tmp1, tempfile.NamedTemporaryFile() as tmp2: + data = np.memmap(tmp1.name, dtype="float32", mode="w+", shape=(10, 3)) + offsets = np.memmap(tmp2.name, dtype="uint32", mode="w+", shape=(5,)) + seq = ArraySequence() seq._data = data seq._offsets = offsets - seq._lengths = np.array([2, 2, 2, 2, 2], dtype='uint32') - + seq._lengths = np.array([2, 2, 2, 2, 2], dtype="uint32") + close_or_delete_mmap(seq) @@ -65,7 +73,7 @@ def test_close_or_delete_mmap_with_mmap_attr(): mock_obj = MagicMock() mock_mmap = MagicMock() mock_obj._mmap = mock_mmap - + close_or_delete_mmap(mock_obj) mock_mmap.close.assert_called_once() @@ -77,15 +85,18 @@ def test_close_or_delete_mmap_other_type(caplog): assert "Object to be close or deleted must be np.memmap" in caplog.text -@pytest.mark.parametrize("filename,expected_base,expected_ext", [ - ("test.nii.gz", "test", ".nii.gz"), - ("test.trk.gz", "test", ".trk.gz"), - ("test.nii", "test", ".nii"), - ("test.trk", "test", ".trk"), - ("test.txt", "test", ".txt"), - ("my.file.with.dots.nii.gz", "my.file.with.dots", ".nii.gz"), - ("no_ext", "no_ext", ""), -]) +@pytest.mark.parametrize( + "filename,expected_base,expected_ext", + [ + ("test.nii.gz", "test", ".nii.gz"), + ("test.trk.gz", "test", ".trk.gz"), + ("test.nii", "test", ".nii"), + ("test.trk", "test", ".trk"), + ("test.txt", "test", ".txt"), + ("my.file.with.dots.nii.gz", "my.file.with.dots", ".nii.gz"), + ("no_ext", "no_ext", ""), + ], +) def test_split_name_with_gz(filename, expected_base, expected_ext): """Test split_name_with_gz with various extensions.""" base, ext = split_name_with_gz(filename) @@ -98,7 +109,7 @@ def test_load_matrix_in_any_format_txt(): with tempfile.NamedTemporaryFile(suffix=".txt", mode="w", delete=False) as tmp: tmp.write("1 2 3\n4 5 6") tmp_name = tmp.name - + try: matrix = load_matrix_in_any_format(tmp_name) np.testing.assert_allclose(matrix, [[1, 2, 3], [4, 5, 6]]) @@ -112,7 +123,7 @@ def test_load_matrix_in_any_format_npy(): data = np.array([[1, 2], [3, 4]]) np.save(tmp.name, data) tmp_name = tmp.name - + try: matrix = load_matrix_in_any_format(tmp_name) np.testing.assert_array_equal(matrix, data) @@ -128,6 +139,7 @@ def test_load_matrix_in_any_format_error(): # --- Spatial Reference Tests --- + @pytest.fixture def nifti_ref(): """Create a synthetic Nifti1Image for testing.""" @@ -142,11 +154,11 @@ def nifti_ref(): def trk_header(): """Create a synthetic TRK header for testing.""" return { - 'voxel_to_rasmm': np.diag([1.0, 2.0, 3.0, 1.0]), - 'dimensions': np.array([10, 20, 30], dtype=np.int16), - 'voxel_sizes': np.array([1.0, 2.0, 3.0], dtype=np.float32), - 'voxel_order': 'RAS', - 'magic_number': 'TRACK' + "voxel_to_rasmm": np.diag([1.0, 2.0, 3.0, 1.0]), + "dimensions": np.array([10, 20, 30], dtype=np.int16), + "voxel_sizes": np.array([1.0, 2.0, 3.0], dtype=np.float32), + "voxel_order": "RAS", + "magic_number": "TRACK", } @@ -156,12 +168,14 @@ def test_get_reference_info_wrapper_nifti_obj(nifti_ref): assert np.allclose(affine, nifti_ref.affine) assert np.array_equal(dimensions, [10, 20, 30]) assert np.allclose(voxel_sizes, [1.0, 2.0, 3.0]) - assert voxel_order == 'RAS' + assert voxel_order == "RAS" def test_get_reference_info_wrapper_nifti_header(nifti_ref): """Test get_reference_info_wrapper with a Nifti1Header object.""" - affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(nifti_ref.header) + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper( + nifti_ref.header + ) assert np.allclose(affine, nifti_ref.affine) assert np.array_equal(dimensions, [10, 20, 30]) @@ -176,42 +190,46 @@ def test_get_reference_info_wrapper_nifti_file(nifti_ref): assert np.array_equal(dimensions, [10, 20, 30]) -@patch('nibabel.streamlines.load') +@patch("nibabel.streamlines.load") def test_get_reference_info_wrapper_trk_file(mock_load): """Test get_reference_info_wrapper with a TRK filename.""" mock_trk = MagicMock() mock_trk.header = { - 'voxel_to_rasmm': np.diag([1.0, 2.0, 3.0, 1.0]), - 'dimensions': np.array([10, 20, 30], dtype=np.int16), - 'voxel_sizes': np.array([1.0, 2.0, 3.0], dtype=np.float32), - 'voxel_order': 'RAS', - 'magic_number': 'TRACK' + "voxel_to_rasmm": np.diag([1.0, 2.0, 3.0, 1.0]), + "dimensions": np.array([10, 20, 30], dtype=np.int16), + "voxel_sizes": np.array([1.0, 2.0, 3.0], dtype=np.float32), + "voxel_order": "RAS", + "magic_number": "TRACK", } mock_load.return_value = mock_trk - - affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper("test.trk") - assert np.allclose(affine, mock_trk.header['voxel_to_rasmm']) + + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper( + "test.trk" + ) + assert np.allclose(affine, mock_trk.header["voxel_to_rasmm"]) def test_get_reference_info_wrapper_trk_obj(): """Test get_reference_info_wrapper with a TrkFile object.""" mock_trk = MagicMock(spec=nib.streamlines.trk.TrkFile) mock_trk.header = { - 'voxel_to_rasmm': np.diag([1.0, 2.0, 3.0, 1.0]), - 'dimensions': np.array([10, 20, 30], dtype=np.int16), - 'voxel_sizes': np.array([1.0, 2.0, 3.0], dtype=np.float32), - 'voxel_order': 'RAS', - 'magic_number': 'TRACK' + "voxel_to_rasmm": np.diag([1.0, 2.0, 3.0, 1.0]), + "dimensions": np.array([10, 20, 30], dtype=np.int16), + "voxel_sizes": np.array([1.0, 2.0, 3.0], dtype=np.float32), + "voxel_order": "RAS", + "magic_number": "TRACK", } affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(mock_trk) - assert np.allclose(affine, mock_trk.header['voxel_to_rasmm']) + assert np.allclose(affine, mock_trk.header["voxel_to_rasmm"]) def test_get_reference_info_wrapper_trk_dict(trk_header): """Test get_reference_info_wrapper with a TRK header dict.""" - affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(trk_header) - assert np.allclose(affine, trk_header['voxel_to_rasmm']) - assert np.array_equal(dimensions, trk_header['dimensions']) + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper( + trk_header + ) + assert np.allclose(affine, trk_header["voxel_to_rasmm"]) + assert np.array_equal(dimensions, trk_header["dimensions"]) @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") @@ -227,40 +245,43 @@ def test_get_reference_info_wrapper_sft(nifti_ref): def test_get_reference_info_wrapper_trx_obj(): """Test get_reference_info_wrapper with a TrxFile object mock.""" from trx.trx_file_memmap import TrxFile + mock_trx = MagicMock(spec=TrxFile) mock_trx.header = { - 'VOXEL_TO_RASMM': np.diag([1.0, 1.0, 1.0, 1.0]), - 'DIMENSIONS': np.array([10, 10, 10], dtype=np.uint16) + "VOXEL_TO_RASMM": np.diag([1.0, 1.0, 1.0, 1.0]), + "DIMENSIONS": np.array([10, 10, 10], dtype=np.uint16), } affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(mock_trx) - assert np.allclose(affine, mock_trx.header['VOXEL_TO_RASMM']) - assert np.array_equal(dimensions, mock_trx.header['DIMENSIONS']) + assert np.allclose(affine, mock_trx.header["VOXEL_TO_RASMM"]) + assert np.array_equal(dimensions, mock_trx.header["DIMENSIONS"]) -@patch('trx.trx_file_memmap.load') +@patch("trx.trx_file_memmap.load") def test_get_reference_info_wrapper_trx_file(mock_load): """Test get_reference_info_wrapper with a TRX filename.""" mock_trx = MagicMock() mock_trx.header = { - 'VOXEL_TO_RASMM': np.diag([1.0, 1.0, 1.0, 1.0]), - 'DIMENSIONS': np.array([10, 10, 10], dtype=np.uint16) + "VOXEL_TO_RASMM": np.diag([1.0, 1.0, 1.0, 1.0]), + "DIMENSIONS": np.array([10, 10, 10], dtype=np.uint16), } mock_load.return_value = mock_trx - - affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper("test.trx") - assert np.allclose(affine, mock_trx.header['VOXEL_TO_RASMM']) + + affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper( + "test.trx" + ) + assert np.allclose(affine, mock_trx.header["VOXEL_TO_RASMM"]) def test_get_reference_info_wrapper_trx_dict(): """Test get_reference_info_wrapper with a TRX header dict.""" header = { - 'VOXEL_TO_RASMM': np.diag([1.0, 1.0, 1.0, 1.0]), - 'DIMENSIONS': np.array([10, 10, 10], dtype=np.uint16), - 'NB_VERTICES': 0 + "VOXEL_TO_RASMM": np.diag([1.0, 1.0, 1.0, 1.0]), + "DIMENSIONS": np.array([10, 10, 10], dtype=np.uint16), + "NB_VERTICES": 0, } affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(header) - assert np.allclose(affine, header['VOXEL_TO_RASMM']) - assert np.array_equal(dimensions, header['DIMENSIONS']) + assert np.allclose(affine, header["VOXEL_TO_RASMM"]) + assert np.array_equal(dimensions, header["DIMENSIONS"]) def test_get_reference_info_wrapper_zero_affine(): @@ -269,8 +290,10 @@ def test_get_reference_info_wrapper_zero_affine(): mock_header = MagicMock(spec=nib.Nifti1Header) mock_img.header = mock_header mock_header.get_best_affine.return_value = np.zeros((4, 4)) - mock_header.__getitem__.side_effect = lambda x: [10, 10, 10] if x == 'dim' else [1, 1, 1] - + mock_header.__getitem__.side_effect = ( + lambda x: [10, 10, 10] if x == "dim" else [1, 1, 1] + ) + with pytest.raises(ValueError, match="Invalid affine, contains only zeros"): get_reference_info_wrapper(mock_img) @@ -278,19 +301,21 @@ def test_get_reference_info_wrapper_zero_affine(): def test_get_reference_info_wrapper_binary_order(): """Test get_reference_info_wrapper with binary voxel order.""" header = { - 'voxel_to_rasmm': np.diag([1.0, 1.0, 1.0, 1.0]), - 'dimensions': [10, 10, 10], - 'voxel_sizes': [1, 1, 1], - 'voxel_order': np.bytes_(b'RAS'), # numpy bytes - 'magic_number': 'TRACK' + "voxel_to_rasmm": np.diag([1.0, 1.0, 1.0, 1.0]), + "dimensions": [10, 10, 10], + "voxel_sizes": [1, 1, 1], + "voxel_order": np.bytes_(b"RAS"), # numpy bytes + "magic_number": "TRACK", } affine, dimensions, voxel_sizes, voxel_order = get_reference_info_wrapper(header) - assert voxel_order == 'RAS' + assert voxel_order == "RAS" def test_get_reference_info_wrapper_error(): """Test get_reference_info_wrapper with unsupported type.""" - with pytest.raises(TypeError, match="Input reference is not one of the supported format"): + with pytest.raises( + TypeError, match="Input reference is not one of the supported format" + ): get_reference_info_wrapper(123) @@ -319,19 +344,19 @@ def test_is_header_compatible_affine_diff(nifti_ref, caplog): def test_is_header_compatible_order_diff(caplog): """Test is_header_compatible with different voxel orders.""" affine1 = np.diag([1.0, 1.0, 1.0, 1.0]) - affine2 = np.diag([-1.0, 1.0, 1.0, 1.0]) # LAS instead of RAS - + affine2 = np.diag([-1.0, 1.0, 1.0, 1.0]) # LAS instead of RAS + header1 = { - 'voxel_to_rasmm': affine1, - 'dimensions': [10, 10, 10], - 'voxel_sizes': [1, 1, 1], - 'voxel_order': 'RAS', - 'magic_number': 'TRACK' + "voxel_to_rasmm": affine1, + "dimensions": [10, 10, 10], + "voxel_sizes": [1, 1, 1], + "voxel_order": "RAS", + "magic_number": "TRACK", } header2 = header1.copy() - header2['voxel_to_rasmm'] = affine2 - header2['voxel_order'] = 'LAS' - + header2["voxel_to_rasmm"] = affine2 + header2["voxel_order"] = "LAS" + with caplog.at_level(logging.ERROR): assert not is_header_compatible(header1, header2) assert "Voxel_order not equal" in caplog.text @@ -339,15 +364,16 @@ def test_is_header_compatible_order_diff(caplog): # --- Transformation & Vector Tests --- + def test_get_axis_shift_vector(): """Test get_axis_shift_vector.""" - assert np.array_equal(get_axis_shift_vector(['x', 'y']), [-1.0, -1.0, 0.0]) - assert np.array_equal(get_axis_shift_vector(['z']), [0.0, 0.0, -1.0]) + assert np.array_equal(get_axis_shift_vector(["x", "y"]), [-1.0, -1.0, 0.0]) + assert np.array_equal(get_axis_shift_vector(["z"]), [0.0, 0.0, -1.0]) def test_get_axis_flip_vector(): """Test get_axis_flip_vector.""" - assert np.array_equal(get_axis_flip_vector(['x', 'z']), [-1.0, 1.0, -1.0]) + assert np.array_equal(get_axis_flip_vector(["x", "z"]), [-1.0, 1.0, -1.0]) assert np.array_equal(get_axis_flip_vector([]), [1.0, 1.0, 1.0]) @@ -364,31 +390,34 @@ def test_flip_sft(nifti_ref): """Test flip_sft.""" streamlines = [np.array([[0, 0, 0], [1, 1, 1]], dtype=np.float32)] sft = StatefulTractogram(streamlines, nifti_ref, Space.VOX) - + # Flipping X axis. Center of X is 5.0 (dim[0]=10). # 0 -> (0 - 5) * -1 - (-5) = -5 * -1 + 5 = 10 # 1 -> (1 - 5) * -1 - (-5) = -4 * -1 + 5 = 9 - - flipped_sft = flip_sft(sft, ['x']) + + flipped_sft = flip_sft(sft, ["x"]) assert np.allclose(flipped_sft.streamlines[0][0, 0], 10.0) assert np.allclose(flipped_sft.streamlines[0][1, 0], 9.0) -@patch('trx.utils.dipy_available', False) +@patch("trx.utils.dipy_available", False) def test_flip_sft_no_dipy(caplog): """Test flip_sft when dipy is missing.""" with caplog.at_level(logging.ERROR): - result = flip_sft(None, ['x']) + result = flip_sft(None, ["x"]) assert result is None assert "Dipy library is missing" in caplog.text @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") -@pytest.mark.parametrize("space_str,origin_str,expected_space,expected_origin", [ - ("rasmm", "nifti", Space.RASMM, Origin.NIFTI), - ("voxmm", "trackvis", Space.VOXMM, Origin.TRACKVIS), - ("vox", "nifti", Space.VOX, Origin.NIFTI), -]) +@pytest.mark.parametrize( + "space_str,origin_str,expected_space,expected_origin", + [ + ("rasmm", "nifti", Space.RASMM, Origin.NIFTI), + ("voxmm", "trackvis", Space.VOXMM, Origin.TRACKVIS), + ("vox", "nifti", Space.VOX, Origin.NIFTI), + ], +) def test_get_reverse_enum(space_str, origin_str, expected_space, expected_origin): """Test get_reverse_enum.""" space, origin = get_reverse_enum(space_str, origin_str) @@ -396,7 +425,7 @@ def test_get_reverse_enum(space_str, origin_str, expected_space, expected_origin assert origin == expected_origin -@patch('trx.utils.dipy_available', False) +@patch("trx.utils.dipy_available", False) def test_get_reverse_enum_no_dipy(caplog): """Test get_reverse_enum when dipy is missing.""" with caplog.at_level(logging.ERROR): @@ -407,41 +436,38 @@ def test_get_reverse_enum_no_dipy(caplog): # --- Data Conversion & Dtype Verification Tests --- + def test_convert_data_dict_to_tractogram(): """Test convert_data_dict_to_tractogram.""" data = { - 'strs': [np.array([[0, 0, 0], [1, 1, 1]]), np.array([[2, 2, 2]])], - 'dps': {'test_dps': [1, 2]}, - 'dpv': {'test_dpv': [0.1, 0.2, 0.3]} + "strs": [np.array([[0, 0, 0], [1, 1, 1]]), np.array([[2, 2, 2]])], + "dps": {"test_dps": [1, 2]}, + "dpv": {"test_dpv": [0.1, 0.2, 0.3]}, } obj = convert_data_dict_to_tractogram(data) assert isinstance(obj, nib.streamlines.tractogram.Tractogram) assert len(obj.streamlines) == 2 - assert np.array_equal(obj.data_per_streamline['test_dps'], [[1], [2]]) + assert np.array_equal(obj.data_per_streamline["test_dps"], [[1], [2]]) # Data per vertex is returned as ArraySequence - assert np.allclose(obj.data_per_point['test_dpv'][0], [[0.1], [0.2]]) + assert np.allclose(obj.data_per_point["test_dpv"][0], [[0.1], [0.2]]) def test_append_generator_to_dict_array(): """Test append_generator_to_dict with numpy array.""" - data = {'strs': [], 'dpv': {}, 'dps': {}} + data = {"strs": [], "dpv": {}, "dps": {}} append_generator_to_dict(np.array([[0, 0, 0]]), data) - assert len(data['strs']) == 1 + assert len(data["strs"]) == 1 def test_append_generator_to_dict_item(): """Test append_generator_to_dict with TractogramItem.""" - data = {'strs': [], 'dpv': {}, 'dps': {}} + data = {"strs": [], "dpv": {}, "dps": {}} # TractogramItem(streamline, data_for_streamline=None, data_for_points=None) - item = TractogramItem( - np.array([[0, 0, 0]]), - {'s': 1}, - {'v': [0.1]} - ) + item = TractogramItem(np.array([[0, 0, 0]]), {"s": 1}, {"v": [0.1]}) append_generator_to_dict(item, data) - assert len(data['strs']) == 1 - assert 'v' in data['dpv'] - assert 's' in data['dps'] + assert len(data["strs"]) == 1 + assert "v" in data["dpv"] + assert "s" in data["dps"] def test_verify_trx_dtype(): @@ -450,41 +476,41 @@ def test_verify_trx_dtype(): mock_trx = MagicMock(spec=Tractogram) mock_trx.streamlines._data.dtype = np.float32 mock_trx.streamlines._offsets.dtype = np.uint32 - + mock_dpv = MagicMock() mock_dpv._data.dtype = np.uint16 - mock_trx.data_per_vertex = {'v1': mock_dpv} - - mock_trx.data_per_streamline = {'s1': np.array([1], dtype='int16')} - + mock_trx.data_per_vertex = {"v1": mock_dpv} + + mock_trx.data_per_streamline = {"s1": np.array([1], dtype="int16")} + # Define expected dtype dict dtype_dict = { - 'positions': np.float32, - 'offsets': np.uint32, - 'dpv': {'v1': np.uint16}, - 'dps': {'s1': np.int16} + "positions": np.float32, + "offsets": np.uint32, + "dpv": {"v1": np.uint16}, + "dps": {"s1": np.int16}, } - + assert verify_trx_dtype(mock_trx, dtype_dict) - + # Test mismatches for warnings - with patch('logging.warning') as mock_log: - dtype_dict['positions'] = np.float64 + with patch("logging.warning") as mock_log: + dtype_dict["positions"] = np.float64 assert not verify_trx_dtype(mock_trx, dtype_dict) mock_log.assert_any_call("Positions dtype is different") - - dtype_dict['positions'] = np.float32 - dtype_dict['offsets'] = np.uint64 + + dtype_dict["positions"] = np.float32 + dtype_dict["offsets"] = np.uint64 assert not verify_trx_dtype(mock_trx, dtype_dict) mock_log.assert_any_call("Offsets dtype is different") - - dtype_dict['offsets'] = np.uint32 - dtype_dict['dpv']['v1'] = np.uint32 + + dtype_dict["offsets"] = np.uint32 + dtype_dict["dpv"]["v1"] = np.uint32 assert not verify_trx_dtype(mock_trx, dtype_dict) mock_log.assert_any_call("Data per vertex (v1) dtype is different") - - dtype_dict['dpv']['v1'] = np.uint16 - dtype_dict['dps']['s1'] = np.int32 + + dtype_dict["dpv"]["v1"] = np.uint16 + dtype_dict["dps"]["s1"] = np.int32 assert not verify_trx_dtype(mock_trx, dtype_dict) mock_log.assert_any_call("Data per streamline (s1) dtype is different") @@ -494,33 +520,27 @@ def test_verify_trx_dtype_groups(): mock_trx = MagicMock(spec=Tractogram) mock_trx.streamlines._data.dtype = np.float32 mock_trx.streamlines._offsets.dtype = np.uint32 - + mock_g1 = MagicMock() mock_g1._data.dtype = np.int32 - + mock_dpg_val = MagicMock() mock_dpg_val.dtype = np.float32 - + # verify_trx_dtype expects trx.data_per_point to contain groups and dpg - mock_trx.data_per_point = { - 'g1': mock_g1, - 'g2': {'d1': mock_dpg_val} - } - - dtype_dict = { - 'groups': {'g1': np.int32}, - 'dpg': {'g2': {'d1': np.float32}} - } - + mock_trx.data_per_point = {"g1": mock_g1, "g2": {"d1": mock_dpg_val}} + + dtype_dict = {"groups": {"g1": np.int32}, "dpg": {"g2": {"d1": np.float32}}} + assert verify_trx_dtype(mock_trx, dtype_dict) - + # Test mismatches for dpg and groups - with patch('logging.warning') as mock_log: - dtype_dict['groups']['g1'] = np.int16 + with patch("logging.warning") as mock_log: + dtype_dict["groups"]["g1"] = np.int16 assert not verify_trx_dtype(mock_trx, dtype_dict) mock_log.assert_any_call("Data per group (g1) dtype is different") - - dtype_dict['groups']['g1'] = np.int32 - dtype_dict['dpg']['g2']['d1'] = np.float64 + + dtype_dict["groups"]["g1"] = np.int32 + dtype_dict["dpg"]["g2"]["d1"] = np.float64 assert not verify_trx_dtype(mock_trx, dtype_dict) mock_log.assert_any_call("Data per group (d1) dtype is different") From 1a9efd8821d82765206f5185f599c9cb8fa2f129 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Tue, 11 Aug 2026 13:29:57 +0000 Subject: [PATCH 04/11] Initial plan From 6621ac120ac44799068afc5a3ef57d567ac03646 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Tue, 11 Aug 2026 13:34:22 +0000 Subject: [PATCH 05/11] fix: use reopenable temp memmaps in dtype workflow Co-authored-by: frheault <10820351+frheault@users.noreply.github.com> --- trx/tests/test_cli.py | 67 +++++++++++++++++++++++- trx/workflows.py | 115 ++++++++++++++++++++---------------------- 2 files changed, 122 insertions(+), 60 deletions(-) diff --git a/trx/tests/test_cli.py b/trx/tests/test_cli.py index 71a70a54..7bfa0104 100644 --- a/trx/tests/test_cli.py +++ b/trx/tests/test_cli.py @@ -3,6 +3,8 @@ import os import tempfile +from types import SimpleNamespace +from unittest.mock import patch from deepdiff import DeepDiff import numpy as np @@ -18,6 +20,7 @@ from trx.fetcher import fetch_data, get_home, get_testing_files_dict import trx.trx_file_memmap as tmm +from trx.workflows import _create_temp_memmap from trx.workflows import ( convert_dsi_studio, convert_tractogram, @@ -33,6 +36,62 @@ ) +def test_create_temp_memmap_uses_reopenable_path(tmp_path): + with patch("trx.workflows.np.memmap") as mock_memmap: + _create_temp_memmap(tmp_path, np.dtype("float32"), (10,)) + + filename = mock_memmap.call_args.args[0] + assert isinstance(filename, (str, os.PathLike)) + assert os.path.dirname(os.fspath(filename)) == os.fspath(tmp_path) + + +def test_manipulate_trx_datatype_uses_reopenable_memmaps(): + trx = SimpleNamespace( + streamlines=SimpleNamespace( + _data=np.arange(6, dtype=np.float16).reshape((2, 3)), + _offsets=np.array([0, 3], dtype=np.uint64), + ), + data_per_vertex={ + "dpv": SimpleNamespace(_data=np.arange(6, dtype=np.uint8).reshape((2, 3))) + }, + data_per_streamline={"dps": np.array([1, 2], dtype=np.uint8)}, + data_per_group={"group": {"dpg": np.array([1.0, 2.0], dtype=np.float32)}}, + groups={"group": np.array([0, 1], dtype=np.int32)}, + ) + trx.close = lambda: None + + with ( + patch("trx.workflows.tmm.load", return_value=trx), + patch("trx.workflows.tmm.save") as mock_save, + patch( + "trx.workflows.tempfile.NamedTemporaryFile", + side_effect=AssertionError( + "NamedTemporaryFile should not be used for writable memmaps" + ), + ), + ): + manipulate_trx_datatype( + "in.trx", + "out.trx", + { + "positions": np.dtype("float32"), + "offsets": np.dtype("uint32"), + "dpv": {"dpv": np.dtype("uint16")}, + "dps": {"dps": np.dtype("float32")}, + "dpg": {"group": {"dpg": np.dtype("float64")}}, + "groups": {"group": np.dtype("uint16")}, + }, + ) + + assert trx.streamlines._data.dtype == np.dtype("float32") + assert trx.streamlines._offsets.dtype == np.dtype("uint32") + assert trx.data_per_vertex["dpv"]._data.dtype == np.dtype("uint16") + assert trx.data_per_streamline["dps"].dtype == np.dtype("float32") + assert trx.data_per_group["group"]["dpg"].dtype == np.dtype("float64") + assert trx.groups["group"].dtype == np.dtype("uint16") + mock_save.assert_called_once_with(trx, "out.trx") + + def _normalize_dtype_dict(dtype_dict): """Normalize dtype dict to use explicit little-endian byte order. @@ -469,7 +528,13 @@ def test_execution_manipulate_trx_datatype(self): } out_gen_path = os.path.join(tmp_dir, "generated.trx") - manipulate_trx_datatype(expected_trx, out_gen_path, generated_dtype) + with patch( + "trx.workflows.tempfile.NamedTemporaryFile", + side_effect=AssertionError( + "NamedTemporaryFile should not be used for writable memmaps" + ), + ): + manipulate_trx_datatype(expected_trx, out_gen_path, generated_dtype) trx = tmm.load(out_gen_path) assert ( DeepDiff( diff --git a/trx/workflows.py b/trx/workflows.py index b30986a3..4df215b0 100644 --- a/trx/workflows.py +++ b/trx/workflows.py @@ -33,6 +33,12 @@ ) +def _create_temp_memmap(tmp_dir_name, dtype, shape): + fd, filename = tempfile.mkstemp(dir=tmp_dir_name, suffix=".mmap") + os.close(fd) + return np.memmap(filename, dtype=dtype, mode="w+", shape=shape) + + def convert_dsi_studio( in_dsi_tractogram, in_dsi_fa, @@ -701,66 +707,57 @@ def manipulate_trx_datatype(in_filename, out_filename, dict_dtype): # noqa: C90 # For each key in dict_dtype, we create a new memmap with the new dtype # and we copy the data from the old memmap to the new one. - for key in dict_dtype: - if key == "positions": - tmp_mm = np.memmap( - tempfile.NamedTemporaryFile(), - dtype=dict_dtype[key], - mode="w+", - shape=trx.streamlines._data.shape, - ) - tmp_mm[:] = trx.streamlines._data[:] - trx.streamlines._data = tmp_mm - elif key == "offsets": - tmp_mm = np.memmap( - tempfile.NamedTemporaryFile(), - dtype=dict_dtype[key], - mode="w+", - shape=trx.streamlines._offsets.shape, - ) - tmp_mm[:] = trx.streamlines._offsets[:] - trx.streamlines._offsets = tmp_mm - elif key == "dpv": - for key_dpv in dict_dtype[key]: - tmp_mm = np.memmap( - tempfile.NamedTemporaryFile(), - dtype=dict_dtype[key][key_dpv], - mode="w+", - shape=trx.data_per_vertex[key_dpv]._data.shape, + with get_trx_tmp_dir() as tmp_dir_name: + for key in dict_dtype: + if key == "positions": + tmp_mm = _create_temp_memmap( + tmp_dir_name, dict_dtype[key], trx.streamlines._data.shape ) - tmp_mm[:] = trx.data_per_vertex[key_dpv]._data[:] - trx.data_per_vertex[key_dpv]._data = tmp_mm - elif key == "dps": - for key_dps in dict_dtype[key]: - tmp_mm = np.memmap( - tempfile.NamedTemporaryFile(), - dtype=dict_dtype[key][key_dps], - mode="w+", - shape=trx.data_per_streamline[key_dps].shape, + tmp_mm[:] = trx.streamlines._data[:] + trx.streamlines._data = tmp_mm + elif key == "offsets": + tmp_mm = _create_temp_memmap( + tmp_dir_name, dict_dtype[key], trx.streamlines._offsets.shape ) - tmp_mm[:] = trx.data_per_streamline[key_dps][:] - trx.data_per_streamline[key_dps] = tmp_mm - elif key == "dpg": - for key_group in dict_dtype[key]: - for key_dpg in dict_dtype[key][key_group]: - tmp_mm = np.memmap( - tempfile.NamedTemporaryFile(), - dtype=dict_dtype[key][key_group][key_dpg], - mode="w+", - shape=trx.data_per_group[key_group][key_dpg].shape, + tmp_mm[:] = trx.streamlines._offsets[:] + trx.streamlines._offsets = tmp_mm + elif key == "dpv": + for key_dpv in dict_dtype[key]: + tmp_mm = _create_temp_memmap( + tmp_dir_name, + dict_dtype[key][key_dpv], + trx.data_per_vertex[key_dpv]._data.shape, ) - tmp_mm[:] = trx.data_per_group[key_group][key_dpg][:] - trx.data_per_group[key_group][key_dpg] = tmp_mm - elif key == "groups": - for key_group in dict_dtype[key]: - tmp_mm = np.memmap( - tempfile.NamedTemporaryFile(), - dtype=dict_dtype[key][key_group], - mode="w+", - shape=trx.groups[key_group].shape, - ) - tmp_mm[:] = trx.groups[key_group][:] - trx.groups[key_group] = tmp_mm + tmp_mm[:] = trx.data_per_vertex[key_dpv]._data[:] + trx.data_per_vertex[key_dpv]._data = tmp_mm + elif key == "dps": + for key_dps in dict_dtype[key]: + tmp_mm = _create_temp_memmap( + tmp_dir_name, + dict_dtype[key][key_dps], + trx.data_per_streamline[key_dps].shape, + ) + tmp_mm[:] = trx.data_per_streamline[key_dps][:] + trx.data_per_streamline[key_dps] = tmp_mm + elif key == "dpg": + for key_group in dict_dtype[key]: + for key_dpg in dict_dtype[key][key_group]: + tmp_mm = _create_temp_memmap( + tmp_dir_name, + dict_dtype[key][key_group][key_dpg], + trx.data_per_group[key_group][key_dpg].shape, + ) + tmp_mm[:] = trx.data_per_group[key_group][key_dpg][:] + trx.data_per_group[key_group][key_dpg] = tmp_mm + elif key == "groups": + for key_group in dict_dtype[key]: + tmp_mm = _create_temp_memmap( + tmp_dir_name, + dict_dtype[key][key_group], + trx.groups[key_group].shape, + ) + tmp_mm[:] = trx.groups[key_group][:] + trx.groups[key_group] = tmp_mm - tmm.save(trx, out_filename) - trx.close() + tmm.save(trx, out_filename) + trx.close() From 9b1477ce57bcca54320b0ef9d929a6d6c5a25092 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Tue, 11 Aug 2026 13:36:01 +0000 Subject: [PATCH 06/11] test: clarify mocked keys in memmap regression Co-authored-by: frheault <10820351+frheault@users.noreply.github.com> --- trx/tests/test_cli.py | 24 ++++++++++++------------ 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/trx/tests/test_cli.py b/trx/tests/test_cli.py index 7bfa0104..087c3341 100644 --- a/trx/tests/test_cli.py +++ b/trx/tests/test_cli.py @@ -52,11 +52,11 @@ def test_manipulate_trx_datatype_uses_reopenable_memmaps(): _offsets=np.array([0, 3], dtype=np.uint64), ), data_per_vertex={ - "dpv": SimpleNamespace(_data=np.arange(6, dtype=np.uint8).reshape((2, 3))) + "mock_dpv": SimpleNamespace(_data=np.arange(6, dtype=np.uint8).reshape((2, 3))) }, - data_per_streamline={"dps": np.array([1, 2], dtype=np.uint8)}, - data_per_group={"group": {"dpg": np.array([1.0, 2.0], dtype=np.float32)}}, - groups={"group": np.array([0, 1], dtype=np.int32)}, + data_per_streamline={"mock_dps": np.array([1, 2], dtype=np.uint8)}, + data_per_group={"mock_group": {"mock_dpg": np.array([1.0, 2.0], dtype=np.float32)}}, + groups={"mock_group": np.array([0, 1], dtype=np.int32)}, ) trx.close = lambda: None @@ -76,19 +76,19 @@ def test_manipulate_trx_datatype_uses_reopenable_memmaps(): { "positions": np.dtype("float32"), "offsets": np.dtype("uint32"), - "dpv": {"dpv": np.dtype("uint16")}, - "dps": {"dps": np.dtype("float32")}, - "dpg": {"group": {"dpg": np.dtype("float64")}}, - "groups": {"group": np.dtype("uint16")}, + "dpv": {"mock_dpv": np.dtype("uint16")}, + "dps": {"mock_dps": np.dtype("float32")}, + "dpg": {"mock_group": {"mock_dpg": np.dtype("float64")}}, + "groups": {"mock_group": np.dtype("uint16")}, }, ) assert trx.streamlines._data.dtype == np.dtype("float32") assert trx.streamlines._offsets.dtype == np.dtype("uint32") - assert trx.data_per_vertex["dpv"]._data.dtype == np.dtype("uint16") - assert trx.data_per_streamline["dps"].dtype == np.dtype("float32") - assert trx.data_per_group["group"]["dpg"].dtype == np.dtype("float64") - assert trx.groups["group"].dtype == np.dtype("uint16") + assert trx.data_per_vertex["mock_dpv"]._data.dtype == np.dtype("uint16") + assert trx.data_per_streamline["mock_dps"].dtype == np.dtype("float32") + assert trx.data_per_group["mock_group"]["mock_dpg"].dtype == np.dtype("float64") + assert trx.groups["mock_group"].dtype == np.dtype("uint16") mock_save.assert_called_once_with(trx, "out.trx") From d431dcd8be8d8db2521e0cdc237e72c127fea402 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Tue, 11 Aug 2026 13:36:54 +0000 Subject: [PATCH 07/11] test: isolate temp dir in memmap regression Co-authored-by: frheault <10820351+frheault@users.noreply.github.com> --- trx/tests/test_cli.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/trx/tests/test_cli.py b/trx/tests/test_cli.py index 087c3341..fba1c117 100644 --- a/trx/tests/test_cli.py +++ b/trx/tests/test_cli.py @@ -1,6 +1,7 @@ # -*- coding: utf-8 -*- """Tests for CLI commands and workflow functions.""" +from contextlib import nullcontext import os import tempfile from types import SimpleNamespace @@ -45,7 +46,7 @@ def test_create_temp_memmap_uses_reopenable_path(tmp_path): assert os.path.dirname(os.fspath(filename)) == os.fspath(tmp_path) -def test_manipulate_trx_datatype_uses_reopenable_memmaps(): +def test_manipulate_trx_datatype_uses_reopenable_memmaps(tmp_path): trx = SimpleNamespace( streamlines=SimpleNamespace( _data=np.arange(6, dtype=np.float16).reshape((2, 3)), @@ -61,6 +62,7 @@ def test_manipulate_trx_datatype_uses_reopenable_memmaps(): trx.close = lambda: None with ( + patch("trx.workflows.get_trx_tmp_dir", return_value=nullcontext(os.fspath(tmp_path))), patch("trx.workflows.tmm.load", return_value=trx), patch("trx.workflows.tmm.save") as mock_save, patch( From b3b7be0789f48d1b6e591f0c1aa9c54a32f88b37 Mon Sep 17 00:00:00 2001 From: frheault Date: Tue, 11 Aug 2026 09:37:13 -0400 Subject: [PATCH 08/11] Fix coverage and PR 107 warnings --- trx/io.py | 4 ++-- trx/tests/test_io.py | 4 ++-- trx/tests/test_utils.py | 2 +- 3 files changed, 5 insertions(+), 5 deletions(-) diff --git a/trx/io.py b/trx/io.py index 6b77d755..ebf657f2 100644 --- a/trx/io.py +++ b/trx/io.py @@ -72,7 +72,7 @@ def load_sft_with_reference(filepath, reference=None, bbox_check=True, **kwargs) IOError If the file format is unsupported or a required reference is missing. """ - if not dipy_available: + if not dipy_available: # pragma: no cover logging.error( "Dipy library is missing, cannot use functions related " "to the StatefulTractogram." @@ -151,7 +151,7 @@ def save(tractogram_obj, tractogram_filename, bbox_valid_check=False): The function writes to disk and returns ``None``. Returns ``None`` immediately when ``dipy`` is unavailable. """ - if not dipy_available: + if not dipy_available: # pragma: no cover logging.error( "Dipy library is missing, cannot use functions related " "to the StatefulTractogram." diff --git a/trx/tests/test_io.py b/trx/tests/test_io.py index 7d7c301a..11a24fae 100644 --- a/trx/tests/test_io.py +++ b/trx/tests/test_io.py @@ -14,14 +14,14 @@ from dipy.io.streamline import load_tractogram, save_tractogram dipy_available = True -except ImportError: +except ImportError: # pragma: no cover dipy_available = False try: import fury # noqa: F401 fury_available = True -except ImportError: +except ImportError: # pragma: no cover fury_available = False from trx.fetcher import fetch_data, get_home, get_testing_files_dict diff --git a/trx/tests/test_utils.py b/trx/tests/test_utils.py index d6727efe..7f835aa1 100644 --- a/trx/tests/test_utils.py +++ b/trx/tests/test_utils.py @@ -34,7 +34,7 @@ from dipy.io.stateful_tractogram import Origin, Space, StatefulTractogram dipy_available = True -except ImportError: +except ImportError: # pragma: no cover dipy_available = False class Space: From c66273618804eaf1326ae2295db4248b366a13b7 Mon Sep 17 00:00:00 2001 From: frheault Date: Tue, 11 Aug 2026 09:40:23 -0400 Subject: [PATCH 09/11] Fix errors on windows tests --- trx/tests/test_cli.py | 15 +++++++++++---- trx/tests/test_utils.py | 4 ++-- trx/workflows.py | 16 ++++++++++++++++ 3 files changed, 29 insertions(+), 6 deletions(-) diff --git a/trx/tests/test_cli.py b/trx/tests/test_cli.py index fba1c117..54a5233e 100644 --- a/trx/tests/test_cli.py +++ b/trx/tests/test_cli.py @@ -21,8 +21,8 @@ from trx.fetcher import fetch_data, get_home, get_testing_files_dict import trx.trx_file_memmap as tmm -from trx.workflows import _create_temp_memmap from trx.workflows import ( + _create_temp_memmap, convert_dsi_studio, convert_tractogram, generate_trx_from_scratch, @@ -53,16 +53,23 @@ def test_manipulate_trx_datatype_uses_reopenable_memmaps(tmp_path): _offsets=np.array([0, 3], dtype=np.uint64), ), data_per_vertex={ - "mock_dpv": SimpleNamespace(_data=np.arange(6, dtype=np.uint8).reshape((2, 3))) + "mock_dpv": SimpleNamespace( + _data=np.arange(6, dtype=np.uint8).reshape((2, 3)) + ) }, data_per_streamline={"mock_dps": np.array([1, 2], dtype=np.uint8)}, - data_per_group={"mock_group": {"mock_dpg": np.array([1.0, 2.0], dtype=np.float32)}}, + data_per_group={ + "mock_group": {"mock_dpg": np.array([1.0, 2.0], dtype=np.float32)} + }, groups={"mock_group": np.array([0, 1], dtype=np.int32)}, ) trx.close = lambda: None with ( - patch("trx.workflows.get_trx_tmp_dir", return_value=nullcontext(os.fspath(tmp_path))), + patch( + "trx.workflows.get_trx_tmp_dir", + return_value=nullcontext(os.fspath(tmp_path)), + ), patch("trx.workflows.tmm.load", return_value=trx), patch("trx.workflows.tmm.save") as mock_save, patch( diff --git a/trx/tests/test_utils.py b/trx/tests/test_utils.py index 7f835aa1..c388f1f1 100644 --- a/trx/tests/test_utils.py +++ b/trx/tests/test_utils.py @@ -290,8 +290,8 @@ def test_get_reference_info_wrapper_zero_affine(): mock_header = MagicMock(spec=nib.Nifti1Header) mock_img.header = mock_header mock_header.get_best_affine.return_value = np.zeros((4, 4)) - mock_header.__getitem__.side_effect = ( - lambda x: [10, 10, 10] if x == "dim" else [1, 1, 1] + mock_header.__getitem__.side_effect = lambda x: ( + [10, 10, 10] if x == "dim" else [1, 1, 1] ) with pytest.raises(ValueError, match="Invalid affine, contains only zeros"): diff --git a/trx/workflows.py b/trx/workflows.py index 4df215b0..f726c2e7 100644 --- a/trx/workflows.py +++ b/trx/workflows.py @@ -34,6 +34,22 @@ def _create_temp_memmap(tmp_dir_name, dtype, shape): + """Create a temporary numpy memmap array. + + Parameters + ---------- + tmp_dir_name : str + Directory to create the temporary file in. + dtype : np.dtype + Data type of the memmap array. + shape : tuple + Shape of the memmap array. + + Returns + ------- + np.memmap + The memory-mapped array. + """ fd, filename = tempfile.mkstemp(dir=tmp_dir_name, suffix=".mmap") os.close(fd) return np.memmap(filename, dtype=dtype, mode="w+", shape=shape) From 9af5bbf43a79c89bb1f0c32f8cb1207959f09fa8 Mon Sep 17 00:00:00 2001 From: frheault Date: Tue, 11 Aug 2026 09:47:47 -0400 Subject: [PATCH 10/11] Fix Windows PermissionError with NamedTemporaryFile in tests --- trx/tests/test_utils.py | 13 ++++++++----- 1 file changed, 8 insertions(+), 5 deletions(-) diff --git a/trx/tests/test_utils.py b/trx/tests/test_utils.py index c388f1f1..0e1dbd38 100644 --- a/trx/tests/test_utils.py +++ b/trx/tests/test_utils.py @@ -49,16 +49,19 @@ class Origin: def test_close_or_delete_mmap_np_memmap(): """Test close_or_delete_mmap with a numpy.memmap.""" - with tempfile.NamedTemporaryFile() as tmp: - mmap_arr = np.memmap(tmp.name, dtype="float32", mode="w+", shape=(10,)) + with tempfile.TemporaryDirectory() as tmpdir: + tmp_name = os.path.join(tmpdir, "test.mmap") + mmap_arr = np.memmap(tmp_name, dtype="float32", mode="w+", shape=(10,)) close_or_delete_mmap(mmap_arr) def test_close_or_delete_mmap_array_sequence(): """Test close_or_delete_mmap with an ArraySequence.""" - with tempfile.NamedTemporaryFile() as tmp1, tempfile.NamedTemporaryFile() as tmp2: - data = np.memmap(tmp1.name, dtype="float32", mode="w+", shape=(10, 3)) - offsets = np.memmap(tmp2.name, dtype="uint32", mode="w+", shape=(5,)) + with tempfile.TemporaryDirectory() as tmpdir: + tmp1_name = os.path.join(tmpdir, "test1.mmap") + tmp2_name = os.path.join(tmpdir, "test2.mmap") + data = np.memmap(tmp1_name, dtype="float32", mode="w+", shape=(10, 3)) + offsets = np.memmap(tmp2_name, dtype="uint32", mode="w+", shape=(5,)) seq = ArraySequence() seq._data = data From 9e31bcb71d1a2b5640adfbe943d075dc9ad6108b Mon Sep 17 00:00:00 2001 From: frheault Date: Tue, 11 Aug 2026 13:37:21 -0400 Subject: [PATCH 11/11] add assertions to validate mmap closure and temp file deletion --- trx/tests/test_utils.py | 7 +++++++ trx/utils.py | 2 -- 2 files changed, 7 insertions(+), 2 deletions(-) diff --git a/trx/tests/test_utils.py b/trx/tests/test_utils.py index 0e1dbd38..2b0753ff 100644 --- a/trx/tests/test_utils.py +++ b/trx/tests/test_utils.py @@ -53,6 +53,8 @@ def test_close_or_delete_mmap_np_memmap(): tmp_name = os.path.join(tmpdir, "test.mmap") mmap_arr = np.memmap(tmp_name, dtype="float32", mode="w+", shape=(10,)) close_or_delete_mmap(mmap_arr) + assert mmap_arr._mmap.closed + assert not os.path.exists(tmp_name) def test_close_or_delete_mmap_array_sequence(): @@ -69,6 +71,11 @@ def test_close_or_delete_mmap_array_sequence(): seq._lengths = np.array([2, 2, 2, 2, 2], dtype="uint32") close_or_delete_mmap(seq) + assert seq._data._mmap.closed + assert seq._offsets._mmap.closed + + assert not os.path.exists(tmp1_name) + assert not os.path.exists(tmp2_name) def test_close_or_delete_mmap_with_mmap_attr(): diff --git a/trx/utils.py b/trx/utils.py index 92c0f3c5..43608b41 100644 --- a/trx/utils.py +++ b/trx/utils.py @@ -31,8 +31,6 @@ def close_or_delete_mmap(obj): close_or_delete_mmap(obj._data) close_or_delete_mmap(obj._offsets) close_or_delete_mmap(obj._lengths) - elif isinstance(obj, np.memmap): - del obj else: logging.debug("Object to be close or deleted must be np.memmap")