Skip to main content

trx_rs/ops/
copy_metadata.rs

1//! Copy metadata (DPS / DPV / groups, optionally DPG) from a donor TRX onto a
2//! target TRX whose streamlines and vertices already match.
3//!
4//! Unlike [`concatenate_any_trx`](crate::concatenate_any_trx) and
5//! [`subset_streamlines`](crate::subset_streamlines), this operation does not
6//! touch positions, offsets, or the header — it only grafts named arrays from
7//! the source onto the target's metadata `HashMap`s.
8
9use std::collections::HashMap;
10
11use crate::any_trx_file::AnyTrxFile;
12use crate::dtype::{DType, TrxScalar};
13use crate::error::{Result, TrxError};
14use crate::trx_file::{DataArray, DataPerGroup, TrxFile};
15
16/// Options controlling [`copy_metadata_any_trx`].
17///
18/// When `dps`, `dpv`, and `groups` are all `None`, every array in each
19/// category is copied. Once any of them is `Some`, only the kinds with an
20/// explicit filter are copied; unfiltered kinds are skipped.
21#[derive(Clone, Debug, Default)]
22pub struct CopyMetadataOptions {
23    /// Names of DPS arrays to copy, or `None` to fall back to the global
24    /// "copy everything" rule described above.
25    pub dps: Option<Vec<String>>,
26    pub dpv: Option<Vec<String>>,
27    pub groups: Option<Vec<String>>,
28    /// If `true`, also copy data-per-group entries for the selected groups.
29    pub copy_dpg: bool,
30    /// If `true`, donor entries replace existing target entries with the same
31    /// name. If `false`, name collisions are an error.
32    pub overwrite_conflicting_metadata: bool,
33    /// If `true`, donor arrays whose row count does not match the target's
34    /// streamline / vertex count are skipped with a warning rather than
35    /// aborting the operation.
36    pub skip_mismatched: bool,
37}
38
39/// Copy metadata from `source` onto `target`, returning the modified target.
40pub fn copy_metadata_any_trx(
41    target: AnyTrxFile,
42    source: &AnyTrxFile,
43    opts: &CopyMetadataOptions,
44) -> Result<AnyTrxFile> {
45    match target {
46        AnyTrxFile::F16(trx) => copy_metadata(trx, source, opts).map(AnyTrxFile::F16),
47        AnyTrxFile::F32(trx) => copy_metadata(trx, source, opts).map(AnyTrxFile::F32),
48        AnyTrxFile::F64(trx) => copy_metadata(trx, source, opts).map(AnyTrxFile::F64),
49    }
50}
51
52/// Typed variant of [`copy_metadata_any_trx`].
53pub fn copy_metadata<P: TrxScalar>(
54    mut target: TrxFile<P>,
55    source: &AnyTrxFile,
56    opts: &CopyMetadataOptions,
57) -> Result<TrxFile<P>> {
58    let plan = SelectionPlan::from_options(opts);
59
60    if (plan.copies(Kind::Dps) || plan.copies(Kind::Group))
61        && target.nb_streamlines() != source.nb_streamlines()
62    {
63        return Err(TrxError::Argument(format!(
64            "donor and target streamline counts differ ({} vs {}); cannot copy DPS or groups",
65            source.nb_streamlines(),
66            target.nb_streamlines()
67        )));
68    }
69    if plan.copies(Kind::Dpv) && target.nb_vertices() != source.nb_vertices() {
70        return Err(TrxError::Argument(format!(
71            "donor and target vertex counts differ ({} vs {}); cannot copy DPV",
72            source.nb_vertices(),
73            target.nb_vertices()
74        )));
75    }
76
77    let dps = collect_named(source, Kind::Dps, plan.filter(Kind::Dps))?;
78    let dpv = collect_named(source, Kind::Dpv, plan.filter(Kind::Dpv))?;
79    let groups = collect_named(source, Kind::Group, plan.filter(Kind::Group))?;
80
81    let dps = filter_by_rows(
82        dps,
83        target.nb_streamlines(),
84        Kind::Dps,
85        opts.skip_mismatched,
86    )?;
87    let dpv = filter_by_rows(dpv, target.nb_vertices(), Kind::Dpv, opts.skip_mismatched)?;
88    validate_group_indices(&groups, target.nb_streamlines())?;
89
90    if !opts.overwrite_conflicting_metadata {
91        check_no_conflict(target.dps_arrays(), &dps, Kind::Dps)?;
92        check_no_conflict(target.dpv_arrays(), &dpv, Kind::Dpv)?;
93        check_no_conflict(target.group_arrays(), &groups, Kind::Group)?;
94    }
95
96    extend_map(target.dps_arrays_mut(), dps);
97    extend_map(target.dpv_arrays_mut(), dpv);
98    extend_map(target.group_arrays_mut(), groups);
99
100    if opts.copy_dpg {
101        let dpg = collect_dpg(source, plan.filter(Kind::Group));
102        if !opts.overwrite_conflicting_metadata {
103            check_no_dpg_conflict(target.dpg_arrays(), &dpg)?;
104        }
105        merge_dpg(target.dpg_arrays_mut(), dpg);
106    }
107
108    Ok(target)
109}
110
111#[derive(Copy, Clone, Eq, PartialEq)]
112enum Kind {
113    Dps,
114    Dpv,
115    Group,
116}
117
118impl Kind {
119    fn label(self) -> &'static str {
120        match self {
121            Kind::Dps => "DPS",
122            Kind::Dpv => "DPV",
123            Kind::Group => "group",
124        }
125    }
126}
127
128/// What to copy, per kind: nothing, every entry, or a specific name list.
129enum Selection<'a> {
130    Skip,
131    All,
132    Named(&'a [String]),
133}
134
135struct SelectionPlan<'a> {
136    dps: Selection<'a>,
137    dpv: Selection<'a>,
138    groups: Selection<'a>,
139}
140
141impl<'a> SelectionPlan<'a> {
142    fn from_options(opts: &'a CopyMetadataOptions) -> Self {
143        let any_filter = opts.dps.is_some() || opts.dpv.is_some() || opts.groups.is_some();
144        Self {
145            dps: Self::select(opts.dps.as_deref(), any_filter),
146            dpv: Self::select(opts.dpv.as_deref(), any_filter),
147            groups: Self::select(opts.groups.as_deref(), any_filter),
148        }
149    }
150
151    fn select(filter: Option<&'a [String]>, any_filter: bool) -> Selection<'a> {
152        match (filter, any_filter) {
153            (Some(names), _) => Selection::Named(names),
154            (None, false) => Selection::All,
155            (None, true) => Selection::Skip,
156        }
157    }
158
159    fn filter(&self, kind: Kind) -> &Selection<'_> {
160        match kind {
161            Kind::Dps => &self.dps,
162            Kind::Dpv => &self.dpv,
163            Kind::Group => &self.groups,
164        }
165    }
166
167    fn copies(&self, kind: Kind) -> bool {
168        !matches!(self.filter(kind), Selection::Skip)
169    }
170}
171
172fn collect_named(
173    source: &AnyTrxFile,
174    kind: Kind,
175    selection: &Selection<'_>,
176) -> Result<Vec<(String, DataArray)>> {
177    let pick = |arrays: &HashMap<String, DataArray>| -> Result<Vec<(String, DataArray)>> {
178        match selection {
179            Selection::Skip => Ok(Vec::new()),
180            Selection::All => Ok(arrays
181                .iter()
182                .map(|(name, arr)| (name.clone(), arr.clone_owned()))
183                .collect()),
184            Selection::Named(names) => names
185                .iter()
186                .map(|name| {
187                    arrays
188                        .get(name.as_str())
189                        .map(|arr| (name.clone(), arr.clone_owned()))
190                        .ok_or_else(|| {
191                            TrxError::Argument(format!(
192                                "donor has no {} named '{name}'",
193                                kind.label()
194                            ))
195                        })
196                })
197                .collect(),
198        }
199    };
200
201    source.with_typed(
202        |s| pick(arrays_for(s, kind)),
203        |s| pick(arrays_for(s, kind)),
204        |s| pick(arrays_for(s, kind)),
205    )
206}
207
208fn arrays_for<P: TrxScalar>(trx: &TrxFile<P>, kind: Kind) -> &HashMap<String, DataArray> {
209    match kind {
210        Kind::Dps => trx.dps_arrays(),
211        Kind::Dpv => trx.dpv_arrays(),
212        Kind::Group => trx.group_arrays(),
213    }
214}
215
216fn filter_by_rows(
217    entries: Vec<(String, DataArray)>,
218    expected_rows: usize,
219    kind: Kind,
220    skip_mismatched: bool,
221) -> Result<Vec<(String, DataArray)>> {
222    entries
223        .into_iter()
224        .filter_map(|(name, arr)| {
225            if arr.nrows() == expected_rows {
226                return Some(Ok((name, arr)));
227            }
228            if skip_mismatched {
229                eprintln!(
230                    "warning: skipping {} '{name}' (rows {} != target {expected_rows})",
231                    kind.label(),
232                    arr.nrows()
233                );
234                return None;
235            }
236            Some(Err(TrxError::Argument(format!(
237                "{} '{name}' has {} rows, target expects {expected_rows}",
238                kind.label(),
239                arr.nrows()
240            ))))
241        })
242        .collect()
243}
244
245fn validate_group_indices(groups: &[(String, DataArray)], nb_streamlines: usize) -> Result<()> {
246    for (name, arr) in groups {
247        if let Some(max) = max_group_index(arr, name)? {
248            if max >= nb_streamlines as u64 {
249                return Err(TrxError::Argument(format!(
250                    "group '{name}' references streamline index {max} but target has only \
251                     {nb_streamlines} streamlines"
252                )));
253            }
254        }
255    }
256    Ok(())
257}
258
259fn max_group_index(arr: &DataArray, name: &str) -> Result<Option<u64>> {
260    match arr.dtype() {
261        DType::Int8 => Ok(arr
262            .cast_slice::<i8>()
263            .iter()
264            .map(|&v| v as i64)
265            .max()
266            .map(check_nonneg)
267            .transpose()?),
268        DType::Int16 => Ok(arr
269            .cast_slice::<i16>()
270            .iter()
271            .map(|&v| v as i64)
272            .max()
273            .map(check_nonneg)
274            .transpose()?),
275        DType::Int32 => Ok(arr
276            .cast_slice::<i32>()
277            .iter()
278            .map(|&v| v as i64)
279            .max()
280            .map(check_nonneg)
281            .transpose()?),
282        DType::Int64 => Ok(arr
283            .cast_slice::<i64>()
284            .iter()
285            .copied()
286            .max()
287            .map(check_nonneg)
288            .transpose()?),
289        DType::UInt8 => Ok(arr.cast_slice::<u8>().iter().map(|&v| v as u64).max()),
290        DType::UInt16 => Ok(arr.cast_slice::<u16>().iter().map(|&v| v as u64).max()),
291        DType::UInt32 => Ok(arr.cast_slice::<u32>().iter().map(|&v| v as u64).max()),
292        DType::UInt64 => Ok(arr.cast_slice::<u64>().iter().copied().max()),
293        other => Err(TrxError::DType(format!(
294            "group '{name}' uses non-integer dtype {other}"
295        ))),
296    }
297}
298
299fn check_nonneg(value: i64) -> Result<u64> {
300    u64::try_from(value).map_err(|_| TrxError::Argument(format!("group index {value} is negative")))
301}
302
303fn check_no_conflict(
304    target: &HashMap<String, DataArray>,
305    incoming: &[(String, DataArray)],
306    kind: Kind,
307) -> Result<()> {
308    if let Some((name, _)) = incoming.iter().find(|(name, _)| target.contains_key(name)) {
309        return Err(TrxError::Argument(format!(
310            "{} '{name}' already exists in target; pass overwrite_conflicting_metadata to replace",
311            kind.label()
312        )));
313    }
314    Ok(())
315}
316
317fn extend_map(target: &mut HashMap<String, DataArray>, incoming: Vec<(String, DataArray)>) {
318    target.extend(incoming);
319}
320
321fn collect_dpg(source: &AnyTrxFile, selection: &Selection<'_>) -> DataPerGroup {
322    let pick = |dpg: &DataPerGroup| -> DataPerGroup {
323        let want_group = |group: &str| -> bool {
324            match selection {
325                Selection::Skip => false,
326                Selection::All => true,
327                Selection::Named(names) => names.iter().any(|n| n == group),
328            }
329        };
330        dpg.iter()
331            .filter(|(group, _)| want_group(group.as_str()))
332            .map(|(group, entries)| {
333                let cloned = entries
334                    .iter()
335                    .map(|(name, arr)| (name.clone(), arr.clone_owned()))
336                    .collect();
337                (group.clone(), cloned)
338            })
339            .collect()
340    };
341
342    source.with_typed(
343        |s| pick(s.dpg_arrays()),
344        |s| pick(s.dpg_arrays()),
345        |s| pick(s.dpg_arrays()),
346    )
347}
348
349fn check_no_dpg_conflict(target: &DataPerGroup, incoming: &DataPerGroup) -> Result<()> {
350    for (group, entries) in incoming {
351        let Some(existing) = target.get(group) else {
352            continue;
353        };
354        if let Some(name) = entries
355            .keys()
356            .find(|name| existing.contains_key(name.as_str()))
357        {
358            return Err(TrxError::Argument(format!(
359                "DPG '{group}/{name}' already exists in target; pass \
360                 overwrite_conflicting_metadata to replace"
361            )));
362        }
363    }
364    Ok(())
365}
366
367fn merge_dpg(target: &mut DataPerGroup, incoming: DataPerGroup) {
368    for (group, entries) in incoming {
369        target.entry(group).or_default().extend(entries);
370    }
371}
372
373#[cfg(test)]
374mod tests {
375    use super::*;
376    use crate::header::Header;
377    use crate::mmap_backing::{vec_to_bytes, MmapBacking};
378    use crate::trx_file::TrxParts;
379
380    fn header() -> Header {
381        Header {
382            voxel_to_rasmm: Header::identity_affine(),
383            dimensions: [10, 10, 10],
384            nb_streamlines: 0,
385            nb_vertices: 0,
386            extra: Default::default(),
387        }
388    }
389
390    fn build(positions: Vec<[f32; 3]>, offsets: Vec<u32>) -> TrxFile<f32> {
391        let mut h = header();
392        h.nb_streamlines = offsets.len().saturating_sub(1) as u64;
393        h.nb_vertices = positions.len() as u64;
394        TrxFile::from_parts(TrxParts {
395            header: h,
396            positions_backing: MmapBacking::Owned(vec_to_bytes(positions)),
397            offsets_backing: MmapBacking::Owned(vec_to_bytes(offsets)),
398            dps: HashMap::new(),
399            dpv: HashMap::new(),
400            groups: HashMap::new(),
401            dpg: HashMap::new(),
402            tempdir: None,
403        })
404    }
405
406    fn scalar_f32(values: Vec<f32>) -> DataArray {
407        DataArray::owned_bytes(vec_to_bytes(values), 1, DType::Float32)
408    }
409
410    fn group_u32(indices: Vec<u32>) -> DataArray {
411        DataArray::owned_bytes(vec_to_bytes(indices), 1, DType::UInt32)
412    }
413
414    #[test]
415    fn copies_all_metadata_when_no_filters() {
416        let mut donor = build(vec![[0.0; 3], [1.0; 3], [2.0; 3]], vec![0, 2, 3]);
417        donor
418            .dps_arrays_mut()
419            .insert("weight".into(), scalar_f32(vec![0.5, 1.5]));
420        donor
421            .dpv_arrays_mut()
422            .insert("fa".into(), scalar_f32(vec![0.1, 0.2, 0.3]));
423        donor
424            .group_arrays_mut()
425            .insert("bundle".into(), group_u32(vec![0, 1]));
426
427        let target = build(vec![[0.0; 3], [1.0; 3], [2.0; 3]], vec![0, 2, 3]);
428
429        let merged = copy_metadata_any_trx(
430            AnyTrxFile::F32(target),
431            &AnyTrxFile::F32(donor),
432            &CopyMetadataOptions::default(),
433        )
434        .unwrap();
435
436        let AnyTrxFile::F32(out) = merged else {
437            panic!("dtype changed");
438        };
439        assert_eq!(
440            out.dps::<f32>("weight").unwrap().as_flat_slice(),
441            &[0.5, 1.5]
442        );
443        assert_eq!(
444            out.dpv::<f32>("fa").unwrap().as_flat_slice(),
445            &[0.1, 0.2, 0.3]
446        );
447        assert_eq!(out.group("bundle").unwrap(), &[0, 1]);
448    }
449
450    #[test]
451    fn streamline_count_mismatch_errors_when_dps_requested() {
452        let mut donor = build(vec![[0.0; 3]], vec![0, 1]);
453        donor
454            .dps_arrays_mut()
455            .insert("w".into(), scalar_f32(vec![1.0]));
456        let target = build(vec![[0.0; 3], [1.0; 3]], vec![0, 1, 2]);
457
458        let err = copy_metadata_any_trx(
459            AnyTrxFile::F32(target),
460            &AnyTrxFile::F32(donor),
461            &CopyMetadataOptions::default(),
462        )
463        .unwrap_err();
464        assert!(err.to_string().contains("streamline counts differ"));
465    }
466
467    #[test]
468    fn name_collision_errors_without_overwrite() {
469        let mut donor = build(vec![[0.0; 3]], vec![0, 1]);
470        donor
471            .dps_arrays_mut()
472            .insert("w".into(), scalar_f32(vec![2.0]));
473        let mut target = build(vec![[0.0; 3]], vec![0, 1]);
474        target
475            .dps_arrays_mut()
476            .insert("w".into(), scalar_f32(vec![1.0]));
477
478        let err = copy_metadata_any_trx(
479            AnyTrxFile::F32(target),
480            &AnyTrxFile::F32(donor),
481            &CopyMetadataOptions::default(),
482        )
483        .unwrap_err();
484        assert!(err.to_string().contains("DPS 'w'"));
485    }
486
487    #[test]
488    fn overwrite_replaces_existing_entry() {
489        let mut donor = build(vec![[0.0; 3]], vec![0, 1]);
490        donor
491            .dps_arrays_mut()
492            .insert("w".into(), scalar_f32(vec![2.0]));
493        let mut target = build(vec![[0.0; 3]], vec![0, 1]);
494        target
495            .dps_arrays_mut()
496            .insert("w".into(), scalar_f32(vec![1.0]));
497
498        let merged = copy_metadata_any_trx(
499            AnyTrxFile::F32(target),
500            &AnyTrxFile::F32(donor),
501            &CopyMetadataOptions {
502                overwrite_conflicting_metadata: true,
503                ..Default::default()
504            },
505        )
506        .unwrap();
507        let AnyTrxFile::F32(out) = merged else {
508            panic!("dtype changed");
509        };
510        assert_eq!(out.dps::<f32>("w").unwrap().as_flat_slice(), &[2.0]);
511    }
512
513    #[test]
514    fn selective_copy_only_named_dps() {
515        let mut donor = build(vec![[0.0; 3]], vec![0, 1]);
516        donor
517            .dps_arrays_mut()
518            .insert("a".into(), scalar_f32(vec![1.0]));
519        donor
520            .dps_arrays_mut()
521            .insert("b".into(), scalar_f32(vec![2.0]));
522        donor
523            .dpv_arrays_mut()
524            .insert("fa".into(), scalar_f32(vec![0.5]));
525        let target = build(vec![[0.0; 3]], vec![0, 1]);
526
527        let merged = copy_metadata_any_trx(
528            AnyTrxFile::F32(target),
529            &AnyTrxFile::F32(donor),
530            &CopyMetadataOptions {
531                dps: Some(vec!["a".into()]),
532                ..Default::default()
533            },
534        )
535        .unwrap();
536        let AnyTrxFile::F32(out) = merged else {
537            panic!("dtype changed");
538        };
539        assert_eq!(out.dps_names(), vec!["a"]);
540        assert!(
541            out.dpv_names().is_empty(),
542            "dpv should not have been copied"
543        );
544    }
545
546    #[test]
547    fn group_index_out_of_bounds_errors() {
548        let mut donor = build(vec![[0.0; 3]], vec![0, 1]);
549        donor
550            .group_arrays_mut()
551            .insert("bad".into(), group_u32(vec![5]));
552        let target = build(vec![[0.0; 3]], vec![0, 1]);
553
554        let err = copy_metadata_any_trx(
555            AnyTrxFile::F32(target),
556            &AnyTrxFile::F32(donor),
557            &CopyMetadataOptions::default(),
558        )
559        .unwrap_err();
560        assert!(err.to_string().contains("references streamline index 5"));
561    }
562}