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        })
403    }
404
405    fn scalar_f32(values: Vec<f32>) -> DataArray {
406        DataArray::owned_bytes(vec_to_bytes(values), 1, DType::Float32)
407    }
408
409    fn group_u32(indices: Vec<u32>) -> DataArray {
410        DataArray::owned_bytes(vec_to_bytes(indices), 1, DType::UInt32)
411    }
412
413    #[test]
414    fn copies_all_metadata_when_no_filters() {
415        let mut donor = build(vec![[0.0; 3], [1.0; 3], [2.0; 3]], vec![0, 2, 3]);
416        donor
417            .dps_arrays_mut()
418            .insert("weight".into(), scalar_f32(vec![0.5, 1.5]));
419        donor
420            .dpv_arrays_mut()
421            .insert("fa".into(), scalar_f32(vec![0.1, 0.2, 0.3]));
422        donor
423            .group_arrays_mut()
424            .insert("bundle".into(), group_u32(vec![0, 1]));
425
426        let target = build(vec![[0.0; 3], [1.0; 3], [2.0; 3]], vec![0, 2, 3]);
427
428        let merged = copy_metadata_any_trx(
429            AnyTrxFile::F32(target),
430            &AnyTrxFile::F32(donor),
431            &CopyMetadataOptions::default(),
432        )
433        .unwrap();
434
435        let AnyTrxFile::F32(out) = merged else {
436            panic!("dtype changed");
437        };
438        assert_eq!(
439            out.dps::<f32>("weight").unwrap().as_flat_slice(),
440            &[0.5, 1.5]
441        );
442        assert_eq!(
443            out.dpv::<f32>("fa").unwrap().as_flat_slice(),
444            &[0.1, 0.2, 0.3]
445        );
446        assert_eq!(out.group("bundle").unwrap(), &[0, 1]);
447    }
448
449    #[test]
450    fn streamline_count_mismatch_errors_when_dps_requested() {
451        let mut donor = build(vec![[0.0; 3]], vec![0, 1]);
452        donor
453            .dps_arrays_mut()
454            .insert("w".into(), scalar_f32(vec![1.0]));
455        let target = build(vec![[0.0; 3], [1.0; 3]], vec![0, 1, 2]);
456
457        let err = copy_metadata_any_trx(
458            AnyTrxFile::F32(target),
459            &AnyTrxFile::F32(donor),
460            &CopyMetadataOptions::default(),
461        )
462        .unwrap_err();
463        assert!(err.to_string().contains("streamline counts differ"));
464    }
465
466    #[test]
467    fn name_collision_errors_without_overwrite() {
468        let mut donor = build(vec![[0.0; 3]], vec![0, 1]);
469        donor
470            .dps_arrays_mut()
471            .insert("w".into(), scalar_f32(vec![2.0]));
472        let mut target = build(vec![[0.0; 3]], vec![0, 1]);
473        target
474            .dps_arrays_mut()
475            .insert("w".into(), scalar_f32(vec![1.0]));
476
477        let err = copy_metadata_any_trx(
478            AnyTrxFile::F32(target),
479            &AnyTrxFile::F32(donor),
480            &CopyMetadataOptions::default(),
481        )
482        .unwrap_err();
483        assert!(err.to_string().contains("DPS 'w'"));
484    }
485
486    #[test]
487    fn overwrite_replaces_existing_entry() {
488        let mut donor = build(vec![[0.0; 3]], vec![0, 1]);
489        donor
490            .dps_arrays_mut()
491            .insert("w".into(), scalar_f32(vec![2.0]));
492        let mut target = build(vec![[0.0; 3]], vec![0, 1]);
493        target
494            .dps_arrays_mut()
495            .insert("w".into(), scalar_f32(vec![1.0]));
496
497        let merged = copy_metadata_any_trx(
498            AnyTrxFile::F32(target),
499            &AnyTrxFile::F32(donor),
500            &CopyMetadataOptions {
501                overwrite_conflicting_metadata: true,
502                ..Default::default()
503            },
504        )
505        .unwrap();
506        let AnyTrxFile::F32(out) = merged else {
507            panic!("dtype changed");
508        };
509        assert_eq!(out.dps::<f32>("w").unwrap().as_flat_slice(), &[2.0]);
510    }
511
512    #[test]
513    fn selective_copy_only_named_dps() {
514        let mut donor = build(vec![[0.0; 3]], vec![0, 1]);
515        donor
516            .dps_arrays_mut()
517            .insert("a".into(), scalar_f32(vec![1.0]));
518        donor
519            .dps_arrays_mut()
520            .insert("b".into(), scalar_f32(vec![2.0]));
521        donor
522            .dpv_arrays_mut()
523            .insert("fa".into(), scalar_f32(vec![0.5]));
524        let target = build(vec![[0.0; 3]], vec![0, 1]);
525
526        let merged = copy_metadata_any_trx(
527            AnyTrxFile::F32(target),
528            &AnyTrxFile::F32(donor),
529            &CopyMetadataOptions {
530                dps: Some(vec!["a".into()]),
531                ..Default::default()
532            },
533        )
534        .unwrap();
535        let AnyTrxFile::F32(out) = merged else {
536            panic!("dtype changed");
537        };
538        assert_eq!(out.dps_names(), vec!["a"]);
539        assert!(
540            out.dpv_names().is_empty(),
541            "dpv should not have been copied"
542        );
543    }
544
545    #[test]
546    fn group_index_out_of_bounds_errors() {
547        let mut donor = build(vec![[0.0; 3]], vec![0, 1]);
548        donor
549            .group_arrays_mut()
550            .insert("bad".into(), group_u32(vec![5]));
551        let target = build(vec![[0.0; 3]], vec![0, 1]);
552
553        let err = copy_metadata_any_trx(
554            AnyTrxFile::F32(target),
555            &AnyTrxFile::F32(donor),
556            &CopyMetadataOptions::default(),
557        )
558        .unwrap_err();
559        assert!(err.to_string().contains("references streamline index 5"));
560    }
561}