Skip to main content

trx_rs/
legacy_io.rs

1use std::fs::File;
2use std::io::{Read, Write};
3use std::path::Path;
4
5use crate::Tractogram;
6
7pub fn load_trk(path: &Path) -> Result<Tractogram, Box<dyn std::error::Error>> {
8    let mut f = File::open(path)?;
9    let mut buffer = Vec::new();
10    f.read_to_end(&mut buffer)?;
11
12    if buffer.len() < 1000 {
13        return Err("File too small".into());
14    }
15
16    let n_scalars = i16::from_le_bytes(buffer[36..38].try_into().unwrap());
17    let n_properties = i16::from_le_bytes(buffer[238..240].try_into().unwrap());
18
19    let mut voxel_sizes = [
20        f32::from_le_bytes(buffer[12..16].try_into().unwrap()),
21        f32::from_le_bytes(buffer[16..20].try_into().unwrap()),
22        f32::from_le_bytes(buffer[20..24].try_into().unwrap()),
23    ];
24
25    // Protect against division by zero for corrupted headers
26    if voxel_sizes[0] == 0.0 {
27        voxel_sizes[0] = 1.0;
28    }
29    if voxel_sizes[1] == 0.0 {
30        voxel_sizes[1] = 1.0;
31    }
32    if voxel_sizes[2] == 0.0 {
33        voxel_sizes[2] = 1.0;
34    }
35
36    let mut vox_to_ras = nalgebra::Matrix4::zeros();
37    let mut mat_offset = 440;
38    for r in 0..4 {
39        for c in 0..4 {
40            vox_to_ras[(r, c)] =
41                f32::from_le_bytes(buffer[mat_offset..mat_offset + 4].try_into().unwrap());
42            mat_offset += 4;
43        }
44    }
45
46    let mut tr = Tractogram::new();
47
48    let dims = [
49        i16::from_le_bytes(buffer[6..8].try_into().unwrap()) as u64,
50        i16::from_le_bytes(buffer[8..10].try_into().unwrap()) as u64,
51        i16::from_le_bytes(buffer[10..12].try_into().unwrap()) as u64,
52    ];
53
54    let mut vox_to_ras_f64 = [[0.0; 4]; 4];
55    for r in 0..4 {
56        for c in 0..4 {
57            vox_to_ras_f64[r][c] = vox_to_ras[(r, c)] as f64;
58        }
59    }
60    tr.set_spatial_metadata(vox_to_ras_f64, dims);
61
62    let mut offset = 1000;
63
64    while offset + 4 <= buffer.len() {
65        let n_points = i32::from_le_bytes(buffer[offset..offset + 4].try_into().unwrap());
66        offset += 4;
67
68        if n_points < 0 {
69            return Err("Negative number of points in streamline".into());
70        }
71
72        let required_bytes = (n_points as usize) * (3 + n_scalars as usize) * 4;
73        if offset + required_bytes > buffer.len() {
74            return Err("Unexpected EOF reading streamline points".into());
75        }
76
77        let mut streamline = Vec::with_capacity(n_points as usize);
78        for _ in 0..n_points {
79            let raw_x = f32::from_le_bytes(buffer[offset..offset + 4].try_into().unwrap());
80            let raw_y = f32::from_le_bytes(buffer[offset + 4..offset + 8].try_into().unwrap());
81            let raw_z = f32::from_le_bytes(buffer[offset + 8..offset + 12].try_into().unwrap());
82
83            let cx = (raw_x / voxel_sizes[0]) - 0.5;
84            let cy = (raw_y / voxel_sizes[1]) - 0.5;
85            let cz = (raw_z / voxel_sizes[2]) - 0.5;
86
87            let p_vox = nalgebra::Point3::new(cx, cy, cz);
88            let p_ras = vox_to_ras.transform_point(&p_vox);
89
90            streamline.push([p_ras.x, p_ras.y, p_ras.z]);
91            offset += (3 + n_scalars as usize) * 4;
92        }
93        tr.push_streamline(&streamline)?;
94        offset += (n_properties as usize) * 4;
95    }
96
97    Ok(tr)
98}
99
100pub fn load_vtk(path: &Path) -> Result<Tractogram, Box<dyn std::error::Error>> {
101    let mut f = File::open(path)?;
102    let mut buffer = Vec::new();
103    f.read_to_end(&mut buffer)?;
104
105    let header_str = String::from_utf8_lossy(&buffer[0..std::cmp::min(1024, buffer.len())]);
106    let points_idx = header_str.find("POINTS ").ok_or("No POINTS")?;
107
108    let points_str = header_str[points_idx..]
109        .split_whitespace()
110        .nth(1)
111        .ok_or("No POINTS count")?;
112    let num_points: usize = points_str.parse()?;
113
114    let mut is_double = false;
115    if let Some(type_str) = header_str[points_idx..].split_whitespace().nth(2) {
116        if type_str == "double" {
117            is_double = true;
118        }
119    }
120
121    let header_end = header_str[points_idx..]
122        .find('\n')
123        .ok_or("No newline after POINTS")?
124        + points_idx
125        + 1;
126    let mut pts = Vec::with_capacity(num_points * 3);
127
128    let mut offset = header_end;
129    for _ in 0..num_points * 3 {
130        if is_double {
131            let chunk = buffer
132                .get(offset..offset + 8)
133                .ok_or("Unexpected EOF reading points")?;
134            let val = f64::from_be_bytes(chunk.try_into().unwrap());
135            pts.push(val as f32);
136            offset += 8;
137        } else {
138            let chunk = buffer
139                .get(offset..offset + 4)
140                .ok_or("Unexpected EOF reading points")?;
141            let val = f32::from_be_bytes(chunk.try_into().unwrap());
142            pts.push(val);
143            offset += 4;
144        }
145    }
146
147    let search_window = std::cmp::min(offset + 1024, buffer.len());
148    let lines_str_chunk = String::from_utf8_lossy(&buffer[offset..search_window]);
149
150    let lines_idx_in_chunk = lines_str_chunk.find("LINES ").ok_or("No LINES")?;
151    let lines_idx = offset + lines_idx_in_chunk;
152
153    let lines_str = lines_str_chunk[lines_idx_in_chunk..]
154        .split_whitespace()
155        .nth(1)
156        .ok_or("No LINES count")?;
157    let num_lines: usize = lines_str.parse()?;
158
159    let lines_header_end = lines_str_chunk[lines_idx_in_chunk..]
160        .find('\n')
161        .ok_or("No newline after LINES")?
162        + lines_idx
163        + 1;
164    offset = lines_header_end;
165
166    let mut tr = Tractogram::new();
167
168    if buffer
169        .get(offset..)
170        .is_some_and(|b| b.starts_with(b"OFFSETS"))
171    {
172        let offsets_header_end = buffer[offset..]
173            .iter()
174            .position(|&c| c == b'\n')
175            .ok_or("No newline after OFFSETS")?
176            + offset
177            + 1;
178
179        let header_str = String::from_utf8_lossy(&buffer[offset..offsets_header_end]);
180        let tokens: Vec<&str> = header_str.split_whitespace().collect();
181
182        let num_offsets = if tokens.len() >= 3 {
183            tokens[2].parse().unwrap_or(num_lines)
184        } else {
185            num_lines
186        };
187
188        let is_int64 = header_str.contains("int64");
189        offset = offsets_header_end;
190
191        let mut offsets_vec = Vec::with_capacity(num_offsets);
192        for _ in 0..num_offsets {
193            if is_int64 {
194                let chunk = buffer
195                    .get(offset..offset + 8)
196                    .ok_or("Unexpected EOF reading offsets")?;
197                let val = u64::from_be_bytes(chunk.try_into().unwrap());
198                offsets_vec.push(val as usize);
199                offset += 8;
200            } else {
201                let chunk = buffer
202                    .get(offset..offset + 4)
203                    .ok_or("Unexpected EOF reading offsets")?;
204                let val = u32::from_be_bytes(chunk.try_into().unwrap());
205                offsets_vec.push(val as usize);
206                offset += 4;
207            }
208        }
209
210        let actual_num_lines = num_offsets.saturating_sub(1);
211        for i in 0..actual_num_lines {
212            let start = offsets_vec[i];
213            let end = offsets_vec[i + 1];
214
215            if end > pts.len() / 3 {
216                return Err("Offset points out of bounds".into());
217            }
218            let mut streamline = Vec::with_capacity(end.saturating_sub(start));
219            for pt_idx in start..end {
220                streamline.push([pts[pt_idx * 3], pts[pt_idx * 3 + 1], pts[pt_idx * 3 + 2]]);
221            }
222            tr.push_streamline(&streamline)?;
223        }
224        return Ok(tr);
225    }
226
227    let mut pt_idx = 0;
228    for _ in 0..num_lines {
229        if offset + 4 > buffer.len() {
230            break;
231        }
232        let n_pts = i32::from_be_bytes(buffer[offset..offset + 4].try_into().unwrap());
233        offset += 4;
234
235        if n_pts <= 0 {
236            continue;
237        }
238        if pt_idx + (n_pts as usize) > num_points {
239            break;
240        }
241
242        let mut streamline = Vec::with_capacity(n_pts as usize);
243        for _ in 0..n_pts {
244            offset += 4;
245            streamline.push([pts[pt_idx * 3], pts[pt_idx * 3 + 1], pts[pt_idx * 3 + 2]]);
246            pt_idx += 1;
247        }
248        tr.push_streamline(&streamline)?;
249    }
250
251    Ok(tr)
252}
253
254pub fn load_nifti_header(path: &Path) -> Result<crate::header::Header, Box<dyn std::error::Error>> {
255    let mut f = File::open(path)?;
256    let mut buffer = Vec::new();
257    f.read_to_end(&mut buffer)?;
258
259    if buffer.len() < 348 {
260        return Err("NIfTI file too small".into());
261    }
262
263    let mut sizeof_hdr_bytes = [0u8; 4];
264    sizeof_hdr_bytes.copy_from_slice(&buffer[0..4]);
265    let sizeof_hdr = i32::from_le_bytes(sizeof_hdr_bytes);
266
267    let (is_nifti2, swap_endian) = if sizeof_hdr == 348 {
268        (false, false)
269    } else if sizeof_hdr == 348i32.swap_bytes() {
270        (false, true)
271    } else if sizeof_hdr == 540 {
272        (true, false)
273    } else if sizeof_hdr == 540i32.swap_bytes() {
274        (true, true)
275    } else {
276        return Err(format!("Unsupported NIfTI sizeof_hdr: {}", sizeof_hdr).into());
277    };
278
279    if is_nifti2 && buffer.len() < 540 {
280        return Err("NIfTI-2 file too small".into());
281    }
282
283    let read_i16 = |offset: usize| -> Result<i16, Box<dyn std::error::Error>> {
284        let bytes = buffer.get(offset..offset + 2).ok_or("Buffer too small")?;
285        let mut arr = [0u8; 2];
286        arr.copy_from_slice(bytes);
287        let val = if swap_endian {
288            i16::from_be_bytes(arr)
289        } else {
290            i16::from_le_bytes(arr)
291        };
292        Ok(val)
293    };
294
295    let read_i32 = |offset: usize| -> Result<i32, Box<dyn std::error::Error>> {
296        let bytes = buffer.get(offset..offset + 4).ok_or("Buffer too small")?;
297        let mut arr = [0u8; 4];
298        arr.copy_from_slice(bytes);
299        let val = if swap_endian {
300            i32::from_be_bytes(arr)
301        } else {
302            i32::from_le_bytes(arr)
303        };
304        Ok(val)
305    };
306
307    let read_i64 = |offset: usize| -> Result<i64, Box<dyn std::error::Error>> {
308        let bytes = buffer.get(offset..offset + 8).ok_or("Buffer too small")?;
309        let mut arr = [0u8; 8];
310        arr.copy_from_slice(bytes);
311        let val = if swap_endian {
312            i64::from_be_bytes(arr)
313        } else {
314            i64::from_le_bytes(arr)
315        };
316        Ok(val)
317    };
318
319    let read_f32 = |offset: usize| -> Result<f32, Box<dyn std::error::Error>> {
320        let bytes = buffer.get(offset..offset + 4).ok_or("Buffer too small")?;
321        let mut arr = [0u8; 4];
322        arr.copy_from_slice(bytes);
323        let val = if swap_endian {
324            f32::from_be_bytes(arr)
325        } else {
326            f32::from_le_bytes(arr)
327        };
328        Ok(val)
329    };
330
331    let read_f64 = |offset: usize| -> Result<f64, Box<dyn std::error::Error>> {
332        let bytes = buffer.get(offset..offset + 8).ok_or("Buffer too small")?;
333        let mut arr = [0u8; 8];
334        arr.copy_from_slice(bytes);
335        let val = if swap_endian {
336            f64::from_be_bytes(arr)
337        } else {
338            f64::from_le_bytes(arr)
339        };
340        Ok(val)
341    };
342
343    let mut dimensions = [1, 1, 1];
344    let qform_code;
345    let sform_code;
346
347    let mut pixdim = [1.0; 8];
348    let mut srow_x = [0.0; 4];
349    let mut srow_y = [0.0; 4];
350    let mut srow_z = [0.0; 4];
351    let quatern_b;
352    let quatern_c;
353    let quatern_d;
354    let qoffset_x;
355    let qoffset_y;
356    let qoffset_z;
357
358    if is_nifti2 {
359        for i in 1..=3 {
360            dimensions[i - 1] = read_i64(16 + i * 8)? as u64;
361        }
362        for (i, px) in pixdim.iter_mut().enumerate() {
363            *px = read_f64(80 + i * 8)?;
364        }
365        qform_code = read_i32(344)?;
366        sform_code = read_i32(348)?;
367        quatern_b = read_f64(352)?;
368        quatern_c = read_f64(360)?;
369        quatern_d = read_f64(368)?;
370        qoffset_x = read_f64(376)?;
371        qoffset_y = read_f64(384)?;
372        qoffset_z = read_f64(392)?;
373        for i in 0..4 {
374            srow_x[i] = read_f64(400 + i * 8)?;
375            srow_y[i] = read_f64(432 + i * 8)?;
376            srow_z[i] = read_f64(464 + i * 8)?;
377        }
378    } else {
379        for i in 1..=3 {
380            dimensions[i - 1] = read_i16(40 + i * 2)? as u64;
381        }
382        for (i, px) in pixdim.iter_mut().enumerate() {
383            *px = read_f32(76 + i * 4)? as f64;
384        }
385        qform_code = read_i16(252)? as i32;
386        sform_code = read_i16(254)? as i32;
387        quatern_b = read_f32(256)? as f64;
388        quatern_c = read_f32(260)? as f64;
389        quatern_d = read_f32(264)? as f64;
390        qoffset_x = read_f32(268)? as f64;
391        qoffset_y = read_f32(272)? as f64;
392        qoffset_z = read_f32(276)? as f64;
393        for i in 0..4 {
394            srow_x[i] = read_f32(280 + i * 4)? as f64;
395            srow_y[i] = read_f32(296 + i * 4)? as f64;
396            srow_z[i] = read_f32(312 + i * 4)? as f64;
397        }
398    }
399
400    let mut voxel_to_rasmm = crate::header::Header::identity_affine();
401
402    if sform_code > 0 {
403        voxel_to_rasmm[0] = srow_x;
404        voxel_to_rasmm[1] = srow_y;
405        voxel_to_rasmm[2] = srow_z;
406        voxel_to_rasmm[3] = [0.0, 0.0, 0.0, 1.0];
407    } else if qform_code > 0 {
408        let b = quatern_b;
409        let c = quatern_c;
410        let d = quatern_d;
411        let a = (1.0 - b * b - c * c - d * d).max(0.0).sqrt();
412        let qfac = if pixdim[0] == 0.0 { 1.0 } else { pixdim[0] };
413        let dx = pixdim[1];
414        let dy = pixdim[2];
415        let dz = pixdim[3];
416
417        let r00 = a * a + b * b - c * c - d * d;
418        let r01 = 2.0 * (b * c - a * d);
419        let r02 = 2.0 * (b * d + a * c);
420
421        let r10 = 2.0 * (b * c + a * d);
422        let r11 = a * a + c * c - b * b - d * d;
423        let r12 = 2.0 * (c * d - a * b);
424
425        let r20 = 2.0 * (b * d - a * c);
426        let r21 = 2.0 * (c * d + a * b);
427        let r22 = a * a + d * d - c * c - b * b;
428
429        voxel_to_rasmm[0] = [r00 * dx, r01 * dy, r02 * qfac * dz, qoffset_x];
430        voxel_to_rasmm[1] = [r10 * dx, r11 * dy, r12 * qfac * dz, qoffset_y];
431        voxel_to_rasmm[2] = [r20 * dx, r21 * dy, r22 * qfac * dz, qoffset_z];
432        voxel_to_rasmm[3] = [0.0, 0.0, 0.0, 1.0];
433    } else {
434        return Err("NIfTI file has no valid spatial transform".into());
435    }
436
437    let header = crate::header::Header {
438        voxel_to_rasmm,
439        dimensions,
440        nb_streamlines: 0,
441        nb_vertices: 0,
442        extra: Default::default(),
443    };
444    Ok(header)
445}
446
447pub fn write_trx(
448    path: &Path,
449    tractogram: &Tractogram,
450    ref_nifti: Option<&Path>,
451) -> Result<(), Box<dyn std::error::Error>> {
452    let mut tractogram = tractogram.clone();
453    let header_empty = tractogram.header().voxel_to_rasmm
454        == crate::header::Header::identity_affine()
455        && tractogram.header().dimensions == [1, 1, 1];
456
457    if header_empty {
458        if let Some(p) = ref_nifti {
459            let hdr = load_nifti_header(p)?;
460            tractogram.set_header(hdr);
461        } else {
462            return Err("TCK -> TRX requires a reference NIfTI file".into());
463        }
464    }
465
466    let any_trx = tractogram.to_trx(crate::dtype::DType::Float32)?;
467    any_trx.save(path)?;
468    Ok(())
469}
470
471/// Derive the 3-byte voxel_order field from a 4×4 affine matrix,
472/// replicating nibabel's `io_orientation` polar-decomposition approach:
473///   1. Normalize columns of the 3×3 block by their L2 norm (removes zoom/scale).
474///   2. SVD of the normalized matrix → R = U * V^T (closest pure rotation).
475///   3. For each input axis (column of R), pick the dominant output axis
476///      (argmax of abs values) with axis-exclusion to handle oblique cases.
477fn axcodes_from_affine(aff: &[[f64; 4]; 4]) -> [u8; 3] {
478    use nalgebra::{Matrix3, SVD};
479    const POS: [u8; 3] = *b"RAS";
480    const NEG: [u8; 3] = *b"LPI";
481
482    // Step 1: build column-normalized 3×3 matrix
483    let mut rs = Matrix3::<f64>::zeros();
484    for col in 0..3 {
485        let norm = (0..3).map(|r| aff[r][col].powi(2)).sum::<f64>().sqrt();
486        let norm = if norm == 0.0 { 1.0 } else { norm };
487        for row in 0..3 {
488            rs[(row, col)] = aff[row][col] / norm;
489        }
490    }
491
492    // Step 2: SVD → R = U * V^T (polar factor, closest orthonormal matrix)
493    let svd = SVD::new(rs, true, true);
494    let u = svd.u.expect("SVD U not computed");
495    let v_t = svd.v_t.expect("SVD V^T not computed");
496    let r = u * v_t;
497
498    // Step 3: per-column argmax with axis exclusion (mirrors nibabel exactly)
499    let mut used = [false; 3];
500    let mut codes = [b'?'; 3];
501    for col in 0..3 {
502        let mut best_row = 0usize;
503        let mut best_val = -1.0f64;
504        for row in 0..3 {
505            if !used[row] && r[(row, col)].abs() > best_val {
506                best_val = r[(row, col)].abs();
507                best_row = row;
508            }
509        }
510        used[best_row] = true;
511        codes[col] = if r[(best_row, col)] >= 0.0 {
512            POS[best_row]
513        } else {
514            NEG[best_row]
515        };
516    }
517    codes
518}
519
520pub fn write_trk(
521    path: &Path,
522    tractogram: &Tractogram,
523    ref_nifti: Option<&Path>,
524) -> Result<(), Box<dyn std::error::Error>> {
525    let mut file = File::create(path)?;
526    let mut header_bytes = vec![0u8; 1000];
527
528    header_bytes[0..5].copy_from_slice(b"TRACK");
529
530    let mut header = tractogram.header().clone();
531    let header_empty = header.voxel_to_rasmm == crate::header::Header::identity_affine()
532        && header.dimensions == [1, 1, 1];
533    if header_empty {
534        if let Some(p) = ref_nifti {
535            header = load_nifti_header(p)?;
536        } else {
537            return Err("TCK -> TRK requires a reference NIfTI file".into());
538        }
539    }
540    let dims = [
541        header.dimensions[0] as i16,
542        header.dimensions[1] as i16,
543        header.dimensions[2] as i16,
544    ];
545    header_bytes[6..8].copy_from_slice(&dims[0].to_le_bytes());
546    header_bytes[8..10].copy_from_slice(&dims[1].to_le_bytes());
547    header_bytes[10..12].copy_from_slice(&dims[2].to_le_bytes());
548
549    let vox_to_ras = header.voxel_to_rasmm;
550    let voxel_sizes = [
551        ((vox_to_ras[0][0].powi(2) + vox_to_ras[1][0].powi(2) + vox_to_ras[2][0].powi(2)).sqrt())
552            as f32,
553        ((vox_to_ras[0][1].powi(2) + vox_to_ras[1][1].powi(2) + vox_to_ras[2][1].powi(2)).sqrt())
554            as f32,
555        ((vox_to_ras[0][2].powi(2) + vox_to_ras[1][2].powi(2) + vox_to_ras[2][2].powi(2)).sqrt())
556            as f32,
557    ];
558    header_bytes[12..16].copy_from_slice(&voxel_sizes[0].to_le_bytes());
559    header_bytes[16..20].copy_from_slice(&voxel_sizes[1].to_le_bytes());
560    header_bytes[20..24].copy_from_slice(&voxel_sizes[2].to_le_bytes());
561
562    let mut offset = 440;
563    for row in &vox_to_ras {
564        for &elem in row {
565            let val = elem as f32;
566            header_bytes[offset..offset + 4].copy_from_slice(&val.to_le_bytes());
567            offset += 4;
568        }
569    }
570
571    let axcodes = axcodes_from_affine(&vox_to_ras);
572    header_bytes[948..951].copy_from_slice(&axcodes);
573    header_bytes[951] = 0;
574
575    let nb_streamlines = tractogram.nb_streamlines() as i32;
576    header_bytes[988..992].copy_from_slice(&nb_streamlines.to_le_bytes());
577
578    header_bytes[992..996].copy_from_slice(&2i32.to_le_bytes());
579
580    header_bytes[996..1000].copy_from_slice(&1000i32.to_le_bytes());
581
582    file.write_all(&header_bytes)?;
583
584    let mut mat = nalgebra::Matrix4::zeros();
585    for r in 0..4 {
586        for c in 0..4 {
587            mat[(r, c)] = vox_to_ras[r][c] as f32;
588        }
589    }
590    let inv_mat = mat.try_inverse().unwrap_or(nalgebra::Matrix4::identity());
591
592    let offsets = tractogram.offsets();
593    let positions = tractogram.positions();
594    let mut chunk = Vec::with_capacity(4 * 1024 * 1024);
595
596    for i in 0..tractogram.nb_streamlines() {
597        let start = offsets[i] as usize;
598        let end = offsets[i + 1] as usize;
599        let n_points = (end - start) as i32;
600
601        chunk.extend_from_slice(&n_points.to_le_bytes());
602        for &pt in &positions[start..end] {
603            let p_ras = nalgebra::Point3::new(pt[0], pt[1], pt[2]);
604            let p_center = inv_mat.transform_point(&p_ras);
605
606            let vox_x = (p_center.x + 0.5) * voxel_sizes[0];
607            let vox_y = (p_center.y + 0.5) * voxel_sizes[1];
608            let vox_z = (p_center.z + 0.5) * voxel_sizes[2];
609
610            chunk.extend_from_slice(&vox_x.to_le_bytes());
611            chunk.extend_from_slice(&vox_y.to_le_bytes());
612            chunk.extend_from_slice(&vox_z.to_le_bytes());
613        }
614
615        if chunk.len() >= 4_000_000 {
616            file.write_all(&chunk)?;
617            chunk.clear();
618        }
619    }
620
621    if !chunk.is_empty() {
622        file.write_all(&chunk)?;
623    }
624
625    Ok(())
626}