diff --git a/trx/io.py b/trx/io.py index 3d633bb..ebf657f 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 ------- @@ -70,7 +72,7 @@ def load_sft_with_reference(filepath, reference=None, bbox_check=True): 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." @@ -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) @@ -145,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_cli.py b/trx/tests/test_cli.py index 71a70a5..54a5233 100644 --- a/trx/tests/test_cli.py +++ b/trx/tests/test_cli.py @@ -1,8 +1,11 @@ # -*- coding: utf-8 -*- """Tests for CLI commands and workflow functions.""" +from contextlib import nullcontext import os import tempfile +from types import SimpleNamespace +from unittest.mock import patch from deepdiff import DeepDiff import numpy as np @@ -19,6 +22,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, convert_dsi_studio, convert_tractogram, generate_trx_from_scratch, @@ -33,6 +37,70 @@ ) +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(tmp_path): + 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={ + "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)} + }, + 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.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": {"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["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") + + def _normalize_dtype_dict(dtype_dict): """Normalize dtype dict to use explicit little-endian byte order. @@ -469,7 +537,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/tests/test_io.py b/trx/tests/test_io.py index 3db89a8..11a24fa 100644 --- a/trx/tests/test_io.py +++ b/trx/tests/test_io.py @@ -14,9 +14,16 @@ 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: # pragma: no cover + 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 new file mode 100644 index 0000000..2b0753f --- /dev/null +++ b/trx/tests/test_utils.py @@ -0,0 +1,556 @@ +# -*- 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 Tractogram, TractogramItem +import numpy as np +import pytest + +from trx.utils import ( + append_generator_to_dict, + close_or_delete_mmap, + convert_data_dict_to_tractogram, + flip_sft, + get_axis_flip_vector, + get_axis_shift_vector, + get_reference_info_wrapper, + get_reverse_enum, + get_shift_vector, + is_header_compatible, + load_matrix_in_any_format, + split_name_with_gz, + verify_trx_dtype, +) + +# Optional dipy import +try: + import dipy # noqa: F401 + from dipy.io.stateful_tractogram import Origin, Space, StatefulTractogram + + dipy_available = True +except ImportError: # pragma: no cover + 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.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) + assert mmap_arr._mmap.closed + assert not os.path.exists(tmp_name) + + +def test_close_or_delete_mmap_array_sequence(): + """Test close_or_delete_mmap with an ArraySequence.""" + 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 + seq._offsets = offsets + 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(): + """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") diff --git a/trx/trx_file_memmap.py b/trx/trx_file_memmap.py index 04a35b5..412923d 100644 --- a/trx/trx_file_memmap.py +++ b/trx/trx_file_memmap.py @@ -848,13 +848,13 @@ def save( Default is zipfile.ZIP_STORED. """ _, 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): diff --git a/trx/utils.py b/trx/utils.py index 92c0f3c..43608b4 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") diff --git a/trx/workflows.py b/trx/workflows.py index b30986a..f726c2e 100644 --- a/trx/workflows.py +++ b/trx/workflows.py @@ -33,6 +33,28 @@ ) +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) + + def convert_dsi_studio( in_dsi_tractogram, in_dsi_fa, @@ -701,66 +723,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()