1use 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#[derive(Clone, Debug, Default)]
22pub struct CopyMetadataOptions {
23 pub dps: Option<Vec<String>>,
26 pub dpv: Option<Vec<String>>,
27 pub groups: Option<Vec<String>>,
28 pub copy_dpg: bool,
30 pub overwrite_conflicting_metadata: bool,
33 pub skip_mismatched: bool,
37}
38
39pub 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
52pub 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
128enum 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}