Skip to main content

trx_rs/ops/
merge.rs

1use std::collections::{BTreeSet, HashMap};
2use std::fs::{self, OpenOptions};
3use std::path::Path;
4
5use half::f16;
6use memmap2::MmapOptions;
7
8use crate::any_trx_file::{AnyTrxFile, PositionsRef};
9use crate::dtype::{DType, TrxScalar};
10use crate::error::{Result, TrxError};
11use crate::mmap_backing::MmapBacking;
12use crate::trx_file::{DataArray, DataArrayInfo, FromF32, TrxFile, TrxParts};
13
14/// Options that control how TRX files are concatenated.
15#[derive(Clone, Debug, Default)]
16pub struct ConcatenateOptions {
17    /// If true, drop all DPV arrays from the output.
18    pub delete_dpv: bool,
19    /// If true, drop all DPS arrays from the output.
20    pub delete_dps: bool,
21    /// If true, drop all group arrays from the output.
22    pub delete_groups: bool,
23    /// Override the positions dtype in the output (default: f32).
24    pub positions_dtype: Option<DType>,
25    /// Per-input group name prefix/suffix rules.
26    ///
27    /// Each element corresponds to one input in order. A `None` entry leaves
28    /// group names unchanged; `Some(name)` prefixes them with that name.
29    pub input_group_names: Vec<Option<String>>,
30}
31
32#[derive(Clone, Copy, Debug)]
33struct ArraySpec {
34    ncols: usize,
35    dtype: DType,
36}
37
38#[derive(Clone, Copy, Debug)]
39struct GroupSpec {
40    dtype: DType,
41    len: usize,
42}
43
44#[derive(Clone, Debug)]
45struct InputGroupRename {
46    aggregate_name: Option<String>,
47    renamed_groups: Vec<(String, String)>,
48}
49
50/// Concatenate multiple TRX files with Python-like semantics.
51///
52/// Header spatial metadata are copied from the first input without validation.
53/// `dpg` is intentionally dropped to match `trx-python` concatenate behavior.
54pub fn concatenate_any_trx(
55    inputs: &[&AnyTrxFile],
56    options: &ConcatenateOptions,
57) -> Result<AnyTrxFile> {
58    if inputs.is_empty() {
59        return Err(TrxError::Argument("no shards to merge".into()));
60    }
61
62    let first = inputs[0];
63    let target_dtype = options.positions_dtype.unwrap_or_else(|| first.dtype());
64    let total_streamlines = inputs.iter().map(|input| input.nb_streamlines()).sum();
65    let total_vertices = inputs.iter().map(|input| input.nb_vertices()).sum();
66    let group_renames = normalized_group_names(inputs.len(), &options.input_group_names)?
67        .into_iter()
68        .zip(inputs.iter())
69        .map(|(group_name, input)| {
70            effective_group_rename(group_name, group_map(input).keys().cloned())
71        })
72        .collect::<Vec<_>>();
73
74    let dps_specs = retained_array_specs(inputs, ArrayKind::Dps, options.delete_dps)?;
75    let dpv_specs = retained_array_specs(inputs, ArrayKind::Dpv, options.delete_dpv)?;
76    let group_specs = retained_group_specs_with_renames(inputs, &group_renames)?;
77
78    match target_dtype {
79        DType::Float16 => concatenate_into::<f16>(
80            inputs,
81            total_streamlines,
82            total_vertices,
83            dps_specs,
84            dpv_specs,
85            group_specs,
86            &group_renames,
87            options,
88        )
89        .map(AnyTrxFile::F16),
90        DType::Float32 => concatenate_into::<f32>(
91            inputs,
92            total_streamlines,
93            total_vertices,
94            dps_specs,
95            dpv_specs,
96            group_specs,
97            &group_renames,
98            options,
99        )
100        .map(AnyTrxFile::F32),
101        DType::Float64 => concatenate_into::<f64>(
102            inputs,
103            total_streamlines,
104            total_vertices,
105            dps_specs,
106            dpv_specs,
107            group_specs,
108            &group_renames,
109            options,
110        )
111        .map(AnyTrxFile::F64),
112        other => Err(TrxError::DType(format!(
113            "TRX positions must be float16, float32, or float64, got {other}"
114        ))),
115    }
116}
117
118/// Merge multiple typed TRX files into one using file-backed output.
119pub fn merge_trx_shards<P: TrxScalar>(shards: &[&TrxFile<P>]) -> Result<TrxFile<P>> {
120    if shards.is_empty() {
121        return Err(TrxError::Argument("no shards to merge".into()));
122    }
123
124    let total_streamlines = shards.iter().map(|input| input.nb_streamlines()).sum();
125    let total_vertices = shards.iter().map(|input| input.nb_vertices()).sum();
126    let dps_specs = retained_array_specs_typed(shards, ArrayKind::Dps, false)?;
127    let dpv_specs = retained_array_specs_typed(shards, ArrayKind::Dpv, false)?;
128    let group_specs = retained_group_specs_typed(shards)?;
129
130    concatenate_typed_from_trx(
131        shards,
132        total_streamlines,
133        total_vertices,
134        dps_specs,
135        dpv_specs,
136        group_specs,
137        &ConcatenateOptions {
138            positions_dtype: Some(P::DTYPE),
139            ..Default::default()
140        },
141    )
142}
143
144#[allow(clippy::too_many_arguments)]
145fn concatenate_into<P>(
146    inputs: &[&AnyTrxFile],
147    total_streamlines: usize,
148    total_vertices: usize,
149    dps_specs: HashMap<String, ArraySpec>,
150    dpv_specs: HashMap<String, ArraySpec>,
151    group_specs: HashMap<String, GroupSpec>,
152    group_renames: &[InputGroupRename],
153    options: &ConcatenateOptions,
154) -> Result<TrxFile<P>>
155where
156    P: TrxScalar + FromF32,
157{
158    let mut header = inputs[0].header().clone();
159    header.nb_streamlines = total_streamlines as u64;
160    header.nb_vertices = total_vertices as u64;
161
162    let tempdir = tempfile::TempDir::new()?;
163    let tempdir_path = tempdir.path().to_path_buf();
164
165    let mut positions_backing = create_mmap_backing(
166        &tempdir_path.join(format!("positions.3.{}", P::DTYPE.name())),
167        total_vertices * 3 * std::mem::size_of::<P>(),
168    )?;
169    let mut offsets_backing = create_mmap_backing(
170        &tempdir_path.join("offsets.uint32"),
171        (total_streamlines + 1) * 4,
172    )?;
173    let mut dps = create_data_map(&tempdir_path.join("dps"), &dps_specs, total_streamlines)?;
174    let mut dpv = create_data_map(&tempdir_path.join("dpv"), &dpv_specs, total_vertices)?;
175    let mut groups = if options.delete_groups {
176        HashMap::new()
177    } else {
178        create_group_map(&tempdir_path.join("groups"), &group_specs)?
179    };
180
181    fill_positions::<P>(inputs, &mut positions_backing)?;
182    fill_offsets(inputs, &mut offsets_backing)?;
183    copy_retained_arrays(inputs, &dps_specs, &mut dps, ArrayKind::Dps)?;
184    copy_retained_arrays(inputs, &dpv_specs, &mut dpv, ArrayKind::Dpv)?;
185    if !options.delete_groups {
186        copy_groups(inputs, group_renames, &group_specs, &mut groups)?;
187    }
188
189    Ok(TrxFile::from_parts(TrxParts {
190        header,
191        positions_backing,
192        offsets_backing,
193        dps,
194        dpv,
195        groups,
196        dpg: HashMap::new(),
197        tempdir: Some(tempdir),
198    }))
199}
200
201fn concatenate_typed_from_trx<P>(
202    inputs: &[&TrxFile<P>],
203    total_streamlines: usize,
204    total_vertices: usize,
205    dps_specs: HashMap<String, ArraySpec>,
206    dpv_specs: HashMap<String, ArraySpec>,
207    group_specs: HashMap<String, GroupSpec>,
208    options: &ConcatenateOptions,
209) -> Result<TrxFile<P>>
210where
211    P: TrxScalar,
212{
213    let mut header = inputs[0].header().clone();
214    header.nb_streamlines = total_streamlines as u64;
215    header.nb_vertices = total_vertices as u64;
216
217    let tempdir = tempfile::TempDir::new()?;
218    let tempdir_path = tempdir.path().to_path_buf();
219
220    let mut positions_backing = create_mmap_backing(
221        &tempdir_path.join(format!("positions.3.{}", P::DTYPE.name())),
222        total_vertices * 3 * std::mem::size_of::<P>(),
223    )?;
224    let mut offsets_backing = create_mmap_backing(
225        &tempdir_path.join("offsets.uint32"),
226        (total_streamlines + 1) * 4,
227    )?;
228    let mut dps = create_data_map(&tempdir_path.join("dps"), &dps_specs, total_streamlines)?;
229    let mut dpv = create_data_map(&tempdir_path.join("dpv"), &dpv_specs, total_vertices)?;
230    let mut groups = if options.delete_groups {
231        HashMap::new()
232    } else {
233        create_group_map(&tempdir_path.join("groups"), &group_specs)?
234    };
235
236    fill_positions_typed(inputs, &mut positions_backing)?;
237    fill_offsets_typed(inputs, &mut offsets_backing)?;
238    copy_retained_arrays_typed(inputs, &dps_specs, &mut dps, ArrayKind::Dps)?;
239    copy_retained_arrays_typed(inputs, &dpv_specs, &mut dpv, ArrayKind::Dpv)?;
240    if !options.delete_groups {
241        copy_groups_typed(inputs, &group_specs, &mut groups)?;
242    }
243
244    Ok(TrxFile::from_parts(TrxParts {
245        header,
246        positions_backing,
247        offsets_backing,
248        dps,
249        dpv,
250        groups,
251        dpg: HashMap::new(),
252        tempdir: Some(tempdir),
253    }))
254}
255
256#[derive(Clone, Copy)]
257enum ArrayKind {
258    Dps,
259    Dpv,
260}
261
262fn retained_array_specs(
263    inputs: &[&AnyTrxFile],
264    kind: ArrayKind,
265    delete_all: bool,
266) -> Result<HashMap<String, ArraySpec>> {
267    let union_keys = array_keys_union(inputs, kind);
268    let reference = array_infos(inputs[0], kind);
269
270    for input in inputs.iter().skip(1) {
271        let current = array_infos(input, kind);
272        for key in &union_keys {
273            match (reference.get(key), current.get(key)) {
274                (Some(lhs), Some(rhs)) => ensure_matching_array_info(kind, key, *lhs, *rhs)?,
275                _ if delete_all => {}
276                _ => {
277                    return Err(TrxError::Argument(format!(
278                        "{} key '{key}' does not exist in all TrxFile inputs",
279                        array_kind_name(kind)
280                    )))
281                }
282            }
283        }
284    }
285
286    if delete_all {
287        return Ok(HashMap::new());
288    }
289
290    Ok(reference
291        .into_iter()
292        .map(|(name, info)| {
293            (
294                name,
295                ArraySpec {
296                    ncols: info.ncols,
297                    dtype: info.dtype,
298                },
299            )
300        })
301        .collect())
302}
303
304fn retained_group_specs_with_renames(
305    inputs: &[&AnyTrxFile],
306    group_renames: &[InputGroupRename],
307) -> Result<HashMap<String, GroupSpec>> {
308    let mut specs: HashMap<String, GroupSpec> = HashMap::new();
309    let mut lengths: HashMap<String, usize> = HashMap::new();
310
311    for (input, rename) in inputs.iter().zip(group_renames.iter()) {
312        let groups = group_infos(input);
313        if groups.is_empty() {
314            if let Some(name) = &rename.aggregate_name {
315                match specs.get(name) {
316                    Some(existing) if existing.dtype != DType::UInt32 => {
317                        return Err(TrxError::Argument(format!(
318                            "group key '{name}' has different dtypes across inputs"
319                        )))
320                    }
321                    Some(_) => {}
322                    None => {
323                        specs.insert(
324                            name.clone(),
325                            GroupSpec {
326                                dtype: DType::UInt32,
327                                len: 0,
328                            },
329                        );
330                    }
331                }
332                *lengths.entry(name.clone()).or_insert(0usize) += input.nb_streamlines();
333            }
334            continue;
335        }
336
337        for (source_name, info) in groups {
338            let name = rename
339                .renamed_groups
340                .iter()
341                .find_map(|(from, to)| (from == &source_name).then_some(to))
342                .cloned()
343                .unwrap_or(source_name);
344            match specs.get(&name) {
345                Some(existing) if existing.dtype != info.dtype => {
346                    return Err(TrxError::Argument(format!(
347                        "group key '{name}' has different dtypes across inputs"
348                    )))
349                }
350                Some(_) => {}
351                None => {
352                    specs.insert(
353                        name.clone(),
354                        GroupSpec {
355                            dtype: info.dtype,
356                            len: 0,
357                        },
358                    );
359                }
360            }
361            *lengths.entry(name).or_insert(0usize) += info.nrows;
362        }
363    }
364
365    for (name, len) in lengths {
366        specs
367            .get_mut(&name)
368            .expect("lengths and specs must be in sync")
369            .len = len;
370    }
371
372    Ok(specs)
373}
374
375fn retained_array_specs_typed<P: TrxScalar>(
376    inputs: &[&TrxFile<P>],
377    kind: ArrayKind,
378    delete_all: bool,
379) -> Result<HashMap<String, ArraySpec>> {
380    let union_keys: BTreeSet<String> = inputs
381        .iter()
382        .flat_map(|input| match kind {
383            ArrayKind::Dps => input.dps_arrays().keys(),
384            ArrayKind::Dpv => input.dpv_arrays().keys(),
385        })
386        .cloned()
387        .collect();
388    let reference: HashMap<String, DataArrayInfo> = match kind {
389        ArrayKind::Dps => inputs[0]
390            .dps_arrays()
391            .iter()
392            .map(|(name, arr)| (name.clone(), arr.info()))
393            .collect(),
394        ArrayKind::Dpv => inputs[0]
395            .dpv_arrays()
396            .iter()
397            .map(|(name, arr)| (name.clone(), arr.info()))
398            .collect(),
399    };
400
401    for input in inputs.iter().skip(1) {
402        let current: HashMap<String, DataArrayInfo> = match kind {
403            ArrayKind::Dps => input
404                .dps_arrays()
405                .iter()
406                .map(|(name, arr)| (name.clone(), arr.info()))
407                .collect(),
408            ArrayKind::Dpv => input
409                .dpv_arrays()
410                .iter()
411                .map(|(name, arr)| (name.clone(), arr.info()))
412                .collect(),
413        };
414        for key in &union_keys {
415            match (reference.get(key), current.get(key)) {
416                (Some(lhs), Some(rhs)) => ensure_matching_array_info(kind, key, *lhs, *rhs)?,
417                _ if delete_all => {}
418                _ => {
419                    return Err(TrxError::Argument(format!(
420                        "{} key '{key}' does not exist in all TrxFile inputs",
421                        array_kind_name(kind)
422                    )))
423                }
424            }
425        }
426    }
427
428    if delete_all {
429        return Ok(HashMap::new());
430    }
431
432    Ok(reference
433        .into_iter()
434        .map(|(name, info)| {
435            (
436                name,
437                ArraySpec {
438                    ncols: info.ncols,
439                    dtype: info.dtype,
440                },
441            )
442        })
443        .collect())
444}
445
446fn retained_group_specs_typed<P: TrxScalar>(
447    inputs: &[&TrxFile<P>],
448) -> Result<HashMap<String, GroupSpec>> {
449    let mut specs: HashMap<String, GroupSpec> = HashMap::new();
450    let mut lengths: HashMap<String, usize> = HashMap::new();
451
452    for input in inputs {
453        for (name, arr) in input.group_arrays() {
454            match specs.get(name) {
455                Some(existing) if existing.dtype != arr.dtype() => {
456                    return Err(TrxError::Argument(format!(
457                        "group key '{name}' has different dtypes across inputs"
458                    )))
459                }
460                Some(_) => {}
461                None => {
462                    specs.insert(
463                        name.clone(),
464                        GroupSpec {
465                            dtype: arr.dtype(),
466                            len: 0,
467                        },
468                    );
469                }
470            }
471            *lengths.entry(name.clone()).or_insert(0usize) += arr.nrows();
472        }
473    }
474
475    for (name, len) in lengths {
476        specs
477            .get_mut(&name)
478            .expect("lengths and specs must be in sync")
479            .len = len;
480    }
481    Ok(specs)
482}
483
484fn array_keys_union(inputs: &[&AnyTrxFile], kind: ArrayKind) -> BTreeSet<String> {
485    let mut keys = BTreeSet::new();
486    for input in inputs {
487        for key in array_infos(input, kind).into_keys() {
488            keys.insert(key);
489        }
490    }
491    keys
492}
493
494fn array_infos(file: &AnyTrxFile, kind: ArrayKind) -> HashMap<String, DataArrayInfo> {
495    match kind {
496        ArrayKind::Dps => file.dps_entries().into_iter().collect(),
497        ArrayKind::Dpv => file.dpv_entries().into_iter().collect(),
498    }
499}
500
501fn group_infos(file: &AnyTrxFile) -> HashMap<String, DataArrayInfo> {
502    file.with_typed(
503        |trx| {
504            trx.group_arrays()
505                .iter()
506                .map(|(name, arr)| (name.clone(), arr.info()))
507                .collect()
508        },
509        |trx| {
510            trx.group_arrays()
511                .iter()
512                .map(|(name, arr)| (name.clone(), arr.info()))
513                .collect()
514        },
515        |trx| {
516            trx.group_arrays()
517                .iter()
518                .map(|(name, arr)| (name.clone(), arr.info()))
519                .collect()
520        },
521    )
522}
523
524fn ensure_matching_array_info(
525    kind: ArrayKind,
526    key: &str,
527    lhs: DataArrayInfo,
528    rhs: DataArrayInfo,
529) -> Result<()> {
530    if lhs.dtype != rhs.dtype {
531        return Err(TrxError::Argument(format!(
532            "{} key '{key}' has different dtypes across inputs",
533            array_kind_name(kind)
534        )));
535    }
536    if lhs.ncols != rhs.ncols {
537        return Err(TrxError::Argument(format!(
538            "{} key '{key}' has different column counts across inputs",
539            array_kind_name(kind)
540        )));
541    }
542    Ok(())
543}
544
545fn array_kind_name(kind: ArrayKind) -> &'static str {
546    match kind {
547        ArrayKind::Dps => "dps",
548        ArrayKind::Dpv => "dpv",
549    }
550}
551
552fn create_mmap_backing(path: &Path, len: usize) -> Result<MmapBacking> {
553    if len == 0 {
554        return Ok(MmapBacking::Owned(Vec::new()));
555    }
556    if let Some(parent) = path.parent() {
557        fs::create_dir_all(parent)?;
558    }
559    let file = OpenOptions::new()
560        .read(true)
561        .write(true)
562        .create(true)
563        .truncate(true)
564        .open(path)?;
565    file.set_len(len as u64)?;
566    let mmap = unsafe { MmapOptions::new().len(len).map_mut(&file)? };
567    Ok(MmapBacking::ReadWrite(mmap))
568}
569
570fn create_data_map(
571    dir: &Path,
572    specs: &HashMap<String, ArraySpec>,
573    rows: usize,
574) -> Result<HashMap<String, DataArray>> {
575    let mut out = HashMap::new();
576    if specs.is_empty() {
577        return Ok(out);
578    }
579    fs::create_dir_all(dir)?;
580    for (name, spec) in specs {
581        let filename = crate::io::filename::TrxFilename {
582            name: name.clone(),
583            ncols: spec.ncols,
584            dtype: spec.dtype,
585        }
586        .to_filename();
587        let len = rows
588            .checked_mul(spec.ncols)
589            .and_then(|v| v.checked_mul(spec.dtype.size_of()))
590            .ok_or_else(|| TrxError::Argument(format!("array '{name}' is too large")))?;
591        out.insert(
592            name.clone(),
593            DataArray::from_backing(
594                create_mmap_backing(&dir.join(filename), len)?,
595                spec.ncols,
596                spec.dtype,
597            ),
598        );
599    }
600    Ok(out)
601}
602
603fn create_group_map(
604    dir: &Path,
605    specs: &HashMap<String, GroupSpec>,
606) -> Result<HashMap<String, DataArray>> {
607    let mut out = HashMap::new();
608    if specs.is_empty() {
609        return Ok(out);
610    }
611    fs::create_dir_all(dir)?;
612    for (name, spec) in specs {
613        let filename = crate::io::filename::TrxFilename {
614            name: name.clone(),
615            ncols: 1,
616            dtype: spec.dtype,
617        }
618        .to_filename();
619        let len = spec
620            .len
621            .checked_mul(spec.dtype.size_of())
622            .ok_or_else(|| TrxError::Argument(format!("group '{name}' is too large")))?;
623        out.insert(
624            name.clone(),
625            DataArray::from_backing(
626                create_mmap_backing(&dir.join(filename), len)?,
627                1,
628                spec.dtype,
629            ),
630        );
631    }
632    Ok(out)
633}
634
635fn fill_positions<P>(inputs: &[&AnyTrxFile], backing: &mut MmapBacking) -> Result<()>
636where
637    P: TrxScalar + FromF32,
638{
639    let dst: &mut [[P; 3]] = backing.cast_slice_mut()?;
640    let mut cursor = 0usize;
641    for input in inputs {
642        let count = input.nb_vertices();
643        let target = &mut dst[cursor..cursor + count];
644        match input.positions_ref() {
645            PositionsRef::F16(src) => copy_positions(src, target),
646            PositionsRef::F32(src) => copy_positions(src, target),
647            PositionsRef::F64(src) => copy_positions(src, target),
648        }
649        cursor += count;
650    }
651    Ok(())
652}
653
654fn copy_positions<Src, Dst>(src: &[[Src; 3]], dst: &mut [[Dst; 3]])
655where
656    Src: TrxScalar,
657    Dst: TrxScalar + FromF32,
658{
659    for (src_row, dst_row) in src.iter().zip(dst.iter_mut()) {
660        dst_row[0] = Dst::from_f32(src_row[0].to_f32());
661        dst_row[1] = Dst::from_f32(src_row[1].to_f32());
662        dst_row[2] = Dst::from_f32(src_row[2].to_f32());
663    }
664}
665
666fn fill_offsets(inputs: &[&AnyTrxFile], backing: &mut MmapBacking) -> Result<()> {
667    let dst: &mut [u32] = backing.cast_slice_mut()?;
668    let mut cursor = 0usize;
669    let mut vertex_base = 0u32;
670    for input in inputs {
671        for offset in input.offsets_vec().into_iter().take(input.nb_streamlines()) {
672            dst[cursor] = offset
673                .checked_add(vertex_base)
674                .ok_or_else(|| TrxError::Argument("offset overflow during concatenate".into()))?;
675            cursor += 1;
676        }
677        vertex_base = vertex_base
678            .checked_add(input.nb_vertices() as u32)
679            .ok_or_else(|| TrxError::Argument("vertex count overflow during concatenate".into()))?;
680    }
681    dst[cursor] = vertex_base;
682    Ok(())
683}
684
685fn fill_positions_typed<P>(inputs: &[&TrxFile<P>], backing: &mut MmapBacking) -> Result<()>
686where
687    P: TrxScalar,
688{
689    let dst: &mut [[P; 3]] = backing.cast_slice_mut()?;
690    let mut cursor = 0usize;
691    for input in inputs {
692        let count = input.nb_vertices();
693        dst[cursor..cursor + count].copy_from_slice(input.positions());
694        cursor += count;
695    }
696    Ok(())
697}
698
699fn fill_offsets_typed<P: TrxScalar>(
700    inputs: &[&TrxFile<P>],
701    backing: &mut MmapBacking,
702) -> Result<()> {
703    let dst: &mut [u32] = backing.cast_slice_mut()?;
704    let mut cursor = 0usize;
705    let mut vertex_base = 0u32;
706    for input in inputs {
707        for &offset in input.offsets().iter().take(input.nb_streamlines()) {
708            dst[cursor] = offset
709                .checked_add(vertex_base)
710                .ok_or_else(|| TrxError::Argument("offset overflow during concatenate".into()))?;
711            cursor += 1;
712        }
713        vertex_base = vertex_base
714            .checked_add(input.nb_vertices() as u32)
715            .ok_or_else(|| TrxError::Argument("vertex count overflow during concatenate".into()))?;
716    }
717    dst[cursor] = vertex_base;
718    Ok(())
719}
720
721fn copy_retained_arrays(
722    inputs: &[&AnyTrxFile],
723    specs: &HashMap<String, ArraySpec>,
724    outputs: &mut HashMap<String, DataArray>,
725    kind: ArrayKind,
726) -> Result<()> {
727    for (name, spec) in specs {
728        let row_bytes = spec.ncols * spec.dtype.size_of();
729        let dst = outputs
730            .get_mut(name)
731            .expect("output map should contain every retained array")
732            .as_bytes_mut()?;
733        let mut byte_cursor = 0usize;
734        for input in inputs {
735            let src = array_map(input, kind)
736                .get(name)
737                .expect("retained array must exist in every input");
738            let bytes = src.as_bytes();
739            let len = src
740                .nrows()
741                .checked_mul(row_bytes)
742                .ok_or_else(|| TrxError::Argument(format!("array '{name}' is too large")))?;
743            dst[byte_cursor..byte_cursor + len].copy_from_slice(bytes);
744            byte_cursor += len;
745        }
746    }
747    Ok(())
748}
749
750fn copy_retained_arrays_typed<P: TrxScalar>(
751    inputs: &[&TrxFile<P>],
752    specs: &HashMap<String, ArraySpec>,
753    outputs: &mut HashMap<String, DataArray>,
754    kind: ArrayKind,
755) -> Result<()> {
756    for (name, spec) in specs {
757        let row_bytes = spec.ncols * spec.dtype.size_of();
758        let dst = outputs
759            .get_mut(name)
760            .expect("output map should contain every retained array")
761            .as_bytes_mut()?;
762        let mut byte_cursor = 0usize;
763        for input in inputs {
764            let src = match kind {
765                ArrayKind::Dps => input.dps_arrays().get(name),
766                ArrayKind::Dpv => input.dpv_arrays().get(name),
767            }
768            .expect("retained array must exist in every input");
769            let bytes = src.as_bytes();
770            let len = src
771                .nrows()
772                .checked_mul(row_bytes)
773                .ok_or_else(|| TrxError::Argument(format!("array '{name}' is too large")))?;
774            dst[byte_cursor..byte_cursor + len].copy_from_slice(bytes);
775            byte_cursor += len;
776        }
777    }
778    Ok(())
779}
780
781fn copy_groups(
782    inputs: &[&AnyTrxFile],
783    group_renames: &[InputGroupRename],
784    specs: &HashMap<String, GroupSpec>,
785    outputs: &mut HashMap<String, DataArray>,
786) -> Result<()> {
787    let mut positions: HashMap<String, usize> =
788        specs.keys().map(|name| (name.clone(), 0)).collect();
789    let mut streamline_base = 0u32;
790    for (input, rename) in inputs.iter().zip(group_renames.iter()) {
791        let src_groups = group_map(input);
792        if src_groups.is_empty() {
793            if let Some(name) = &rename.aggregate_name {
794                let cursor = positions
795                    .get_mut(name)
796                    .expect("output positions must exist for every group");
797                let dst = outputs
798                    .get_mut(name)
799                    .expect("output arrays must exist for every group");
800                copy_full_streamline_range(*cursor, input.nb_streamlines(), streamline_base, dst)?;
801                *cursor += input.nb_streamlines();
802            }
803        }
804        for (name, arr) in src_groups {
805            let output_name = rename
806                .renamed_groups
807                .iter()
808                .find_map(|(from, to)| (from == name).then_some(to.as_str()))
809                .unwrap_or(name);
810            let cursor = positions
811                .get_mut(output_name)
812                .expect("output positions must exist for every group");
813            let dst = outputs
814                .get_mut(output_name)
815                .expect("output arrays must exist for every group");
816            let count = arr.nrows();
817            copy_group_with_offset(arr, *cursor, streamline_base, dst)?;
818            *cursor += count;
819        }
820        streamline_base = streamline_base
821            .checked_add(input.nb_streamlines() as u32)
822            .ok_or_else(|| {
823                TrxError::Argument("streamline count overflow during concatenate".into())
824            })?;
825    }
826    Ok(())
827}
828
829fn copy_groups_typed<P: TrxScalar>(
830    inputs: &[&TrxFile<P>],
831    specs: &HashMap<String, GroupSpec>,
832    outputs: &mut HashMap<String, DataArray>,
833) -> Result<()> {
834    let mut positions: HashMap<String, usize> =
835        specs.keys().map(|name| (name.clone(), 0)).collect();
836    let mut streamline_base = 0u32;
837    for input in inputs {
838        for (name, arr) in input.group_arrays() {
839            let cursor = positions
840                .get_mut(name)
841                .expect("output positions must exist for every group");
842            let dst = outputs
843                .get_mut(name)
844                .expect("output arrays must exist for every group");
845            let count = arr.nrows();
846            copy_group_with_offset(arr, *cursor, streamline_base, dst)?;
847            *cursor += count;
848        }
849        streamline_base = streamline_base
850            .checked_add(input.nb_streamlines() as u32)
851            .ok_or_else(|| {
852                TrxError::Argument("streamline count overflow during concatenate".into())
853            })?;
854    }
855    Ok(())
856}
857
858fn copy_full_streamline_range(
859    dst_start: usize,
860    streamline_count: usize,
861    streamline_base: u32,
862    dst: &mut DataArray,
863) -> Result<()> {
864    match dst.dtype() {
865        DType::UInt32 => {
866            let values: &mut [u32] = dst.cast_slice_mut()?;
867            for idx in 0..streamline_count {
868                values[dst_start + idx] =
869                    streamline_base.checked_add(idx as u32).ok_or_else(|| {
870                        TrxError::Argument("group index overflow during concatenate".into())
871                    })?;
872            }
873            Ok(())
874        }
875        other => Err(TrxError::DType(format!(
876            "generated group arrays must use uint32 dtype, got {other}"
877        ))),
878    }
879}
880
881fn copy_group_with_offset(
882    src: &DataArray,
883    dst_start: usize,
884    streamline_base: u32,
885    dst: &mut DataArray,
886) -> Result<()> {
887    match src.dtype() {
888        DType::Int8 => copy_group_typed::<i8>(src.cast_slice(), dst_start, streamline_base, dst),
889        DType::Int16 => copy_group_typed::<i16>(src.cast_slice(), dst_start, streamline_base, dst),
890        DType::Int32 => copy_group_typed::<i32>(src.cast_slice(), dst_start, streamline_base, dst),
891        DType::Int64 => copy_group_typed::<i64>(src.cast_slice(), dst_start, streamline_base, dst),
892        DType::UInt8 => copy_group_typed::<u8>(src.cast_slice(), dst_start, streamline_base, dst),
893        DType::UInt16 => copy_group_typed::<u16>(src.cast_slice(), dst_start, streamline_base, dst),
894        DType::UInt32 => copy_group_typed::<u32>(src.cast_slice(), dst_start, streamline_base, dst),
895        DType::UInt64 => copy_group_typed::<u64>(src.cast_slice(), dst_start, streamline_base, dst),
896        other => Err(TrxError::DType(format!(
897            "group arrays must use integer dtype, got {other}"
898        ))),
899    }
900}
901
902fn copy_group_typed<T>(
903    src: &[T],
904    dst_start: usize,
905    streamline_base: u32,
906    dst: &mut DataArray,
907) -> Result<()>
908where
909    T: Copy + TryFrom<u64> + Into<i128> + bytemuck::Pod,
910    <T as TryFrom<u64>>::Error: std::fmt::Display,
911{
912    let dst_values: &mut [T] = dst.cast_slice_mut()?;
913    let addend = i128::from(streamline_base);
914    for (index, &value) in src.iter().enumerate() {
915        let current = value.into();
916        let shifted = current + addend;
917        let shifted_u64 = u64::try_from(shifted)
918            .map_err(|_| TrxError::Argument("group index underflow during concatenate".into()))?;
919        dst_values[dst_start + index] = T::try_from(shifted_u64).map_err(|err| {
920            TrxError::Argument(format!("group index overflow during concatenate: {err}"))
921        })?;
922    }
923    Ok(())
924}
925
926fn array_map(file: &AnyTrxFile, kind: ArrayKind) -> &HashMap<String, DataArray> {
927    match file {
928        AnyTrxFile::F16(trx) => match kind {
929            ArrayKind::Dps => trx.dps_arrays(),
930            ArrayKind::Dpv => trx.dpv_arrays(),
931        },
932        AnyTrxFile::F32(trx) => match kind {
933            ArrayKind::Dps => trx.dps_arrays(),
934            ArrayKind::Dpv => trx.dpv_arrays(),
935        },
936        AnyTrxFile::F64(trx) => match kind {
937            ArrayKind::Dps => trx.dps_arrays(),
938            ArrayKind::Dpv => trx.dpv_arrays(),
939        },
940    }
941}
942
943fn group_map(file: &AnyTrxFile) -> &HashMap<String, DataArray> {
944    match file {
945        AnyTrxFile::F16(trx) => trx.group_arrays(),
946        AnyTrxFile::F32(trx) => trx.group_arrays(),
947        AnyTrxFile::F64(trx) => trx.group_arrays(),
948    }
949}
950
951fn normalized_group_names(
952    input_count: usize,
953    group_names: &[Option<String>],
954) -> Result<Vec<Option<String>>> {
955    if group_names.is_empty() {
956        return Ok(vec![None; input_count]);
957    }
958    if group_names.len() != input_count {
959        return Err(TrxError::Argument(format!(
960            "input_group_names length {} must match input count {input_count}",
961            group_names.len()
962        )));
963    }
964    Ok(group_names
965        .iter()
966        .map(|value| {
967            value.as_ref().and_then(|name| {
968                let trimmed = name.trim();
969                (!trimmed.is_empty()).then(|| trimmed.to_string())
970            })
971        })
972        .collect())
973}
974
975fn effective_group_rename<I>(group_name: Option<String>, source_groups: I) -> InputGroupRename
976where
977    I: IntoIterator<Item = String>,
978{
979    let source_groups = source_groups.into_iter().collect::<Vec<_>>();
980    if source_groups.is_empty() {
981        return InputGroupRename {
982            aggregate_name: group_name,
983            renamed_groups: Vec::new(),
984        };
985    }
986
987    let renamed_groups = if let Some(group_name) = group_name {
988        source_groups
989            .into_iter()
990            .map(|existing| {
991                let renamed = format!("{group_name}{existing}");
992                (existing, renamed)
993            })
994            .collect()
995    } else {
996        Vec::new()
997    };
998
999    InputGroupRename {
1000        aggregate_name: None,
1001        renamed_groups,
1002    }
1003}
1004
1005#[cfg(test)]
1006mod tests {
1007    use super::*;
1008    use crate::header::Header;
1009    use crate::mmap_backing::vec_to_bytes;
1010
1011    fn sample_header() -> Header {
1012        Header {
1013            voxel_to_rasmm: Header::identity_affine(),
1014            dimensions: [10, 20, 30],
1015            nb_streamlines: 0,
1016            nb_vertices: 0,
1017            extra: Default::default(),
1018        }
1019    }
1020
1021    fn build_trx(
1022        positions: Vec<[f32; 3]>,
1023        offsets: Vec<u32>,
1024        dps: HashMap<String, DataArray>,
1025        dpv: HashMap<String, DataArray>,
1026        groups: HashMap<String, DataArray>,
1027        dpg: HashMap<String, HashMap<String, DataArray>>,
1028    ) -> TrxFile<f32> {
1029        let mut header = sample_header();
1030        header.nb_streamlines = offsets.len().saturating_sub(1) as u64;
1031        header.nb_vertices = positions.len() as u64;
1032        TrxFile::from_parts(TrxParts {
1033            header,
1034            positions_backing: MmapBacking::Owned(vec_to_bytes(positions)),
1035            offsets_backing: MmapBacking::Owned(vec_to_bytes(offsets)),
1036            dps,
1037            dpv,
1038            groups,
1039            dpg,
1040            tempdir: None,
1041        })
1042    }
1043
1044    fn scalar_u32(value: u32) -> DataArray {
1045        DataArray::owned_bytes(vec_to_bytes(vec![value]), 1, DType::UInt32)
1046    }
1047
1048    fn scalar_f32(value: f32) -> DataArray {
1049        DataArray::owned_bytes(vec_to_bytes(vec![value]), 1, DType::Float32)
1050    }
1051
1052    fn vertex_f32(values: Vec<f32>) -> DataArray {
1053        DataArray::owned_bytes(vec_to_bytes(values), 1, DType::Float32)
1054    }
1055
1056    #[test]
1057    fn concatenate_preserves_counts_and_group_union_and_drops_dpg() {
1058        let mut groups_a = HashMap::new();
1059        groups_a.insert(
1060            "left".into(),
1061            DataArray::owned_bytes(vec_to_bytes(vec![0u32]), 1, DType::UInt32),
1062        );
1063        let mut dpg_a = HashMap::new();
1064        dpg_a.insert(
1065            "left".into(),
1066            HashMap::from([(
1067                "color".into(),
1068                DataArray::owned_bytes(vec![1, 2, 3], 3, DType::UInt8),
1069            )]),
1070        );
1071
1072        let a = build_trx(
1073            vec![[0.0, 0.0, 0.0], [1.0, 1.0, 1.0]],
1074            vec![0, 2],
1075            HashMap::from([("weight".into(), scalar_f32(1.0))]),
1076            HashMap::from([("fa".into(), vertex_f32(vec![0.1, 0.2]))]),
1077            groups_a,
1078            dpg_a,
1079        );
1080
1081        let mut groups_b = HashMap::new();
1082        groups_b.insert(
1083            "right".into(),
1084            DataArray::owned_bytes(vec_to_bytes(vec![0u32]), 1, DType::UInt32),
1085        );
1086        let b = build_trx(
1087            vec![[2.0, 2.0, 2.0]],
1088            vec![0, 1],
1089            HashMap::from([("weight".into(), scalar_f32(2.0))]),
1090            HashMap::from([("fa".into(), vertex_f32(vec![0.3]))]),
1091            groups_b,
1092            HashMap::new(),
1093        );
1094
1095        let any_a = AnyTrxFile::F32(a);
1096        let any_b = AnyTrxFile::F32(b);
1097        let merged =
1098            concatenate_any_trx(&[&any_a, &any_b], &ConcatenateOptions::default()).unwrap();
1099
1100        match merged {
1101            AnyTrxFile::F32(trx) => {
1102                assert_eq!(trx.nb_streamlines(), 2);
1103                assert_eq!(trx.nb_vertices(), 3);
1104                assert_eq!(trx.header().dimensions, [10, 20, 30]);
1105                assert!(trx.is_file_backed());
1106                assert_eq!(trx.group("left").unwrap(), &[0]);
1107                assert_eq!(trx.group("right").unwrap(), &[1]);
1108                assert!(trx.dpg_group_names().is_empty());
1109            }
1110            _ => panic!("expected float32 output"),
1111        }
1112    }
1113
1114    #[test]
1115    fn concatenate_errors_on_missing_dps_without_delete_flag() {
1116        let a = build_trx(
1117            vec![[0.0, 0.0, 0.0]],
1118            vec![0, 1],
1119            HashMap::from([("weight".into(), scalar_f32(1.0))]),
1120            HashMap::new(),
1121            HashMap::new(),
1122            HashMap::new(),
1123        );
1124        let b = build_trx(
1125            vec![[1.0, 1.0, 1.0]],
1126            vec![0, 1],
1127            HashMap::new(),
1128            HashMap::new(),
1129            HashMap::new(),
1130            HashMap::new(),
1131        );
1132        let any_a = AnyTrxFile::F32(a);
1133        let any_b = AnyTrxFile::F32(b);
1134        let err =
1135            concatenate_any_trx(&[&any_a, &any_b], &ConcatenateOptions::default()).unwrap_err();
1136        assert!(err.to_string().contains("dps key 'weight'"));
1137    }
1138
1139    #[test]
1140    fn concatenate_delete_dps_drops_category() {
1141        let a = build_trx(
1142            vec![[0.0, 0.0, 0.0]],
1143            vec![0, 1],
1144            HashMap::from([("weight".into(), scalar_f32(1.0))]),
1145            HashMap::new(),
1146            HashMap::new(),
1147            HashMap::new(),
1148        );
1149        let b = build_trx(
1150            vec![[1.0, 1.0, 1.0]],
1151            vec![0, 1],
1152            HashMap::new(),
1153            HashMap::new(),
1154            HashMap::new(),
1155            HashMap::new(),
1156        );
1157        let any_a = AnyTrxFile::F32(a);
1158        let any_b = AnyTrxFile::F32(b);
1159        let merged = concatenate_any_trx(
1160            &[&any_a, &any_b],
1161            &ConcatenateOptions {
1162                delete_dps: true,
1163                ..Default::default()
1164            },
1165        )
1166        .unwrap();
1167        match merged {
1168            AnyTrxFile::F32(trx) => assert!(trx.dps_names().is_empty()),
1169            _ => panic!("expected float32 output"),
1170        }
1171    }
1172
1173    #[test]
1174    fn concatenate_errors_on_dpv_column_mismatch() {
1175        let a = build_trx(
1176            vec![[0.0, 0.0, 0.0]],
1177            vec![0, 1],
1178            HashMap::new(),
1179            HashMap::from([(
1180                "fa".into(),
1181                DataArray::owned_bytes(vec_to_bytes(vec![0.1f32]), 1, DType::Float32),
1182            )]),
1183            HashMap::new(),
1184            HashMap::new(),
1185        );
1186        let b = build_trx(
1187            vec![[1.0, 1.0, 1.0]],
1188            vec![0, 1],
1189            HashMap::new(),
1190            HashMap::from([(
1191                "fa".into(),
1192                DataArray::owned_bytes(vec_to_bytes(vec![0.2f32, 0.3f32]), 2, DType::Float32),
1193            )]),
1194            HashMap::new(),
1195            HashMap::new(),
1196        );
1197        let any_a = AnyTrxFile::F32(a);
1198        let any_b = AnyTrxFile::F32(b);
1199        let err =
1200            concatenate_any_trx(&[&any_a, &any_b], &ConcatenateOptions::default()).unwrap_err();
1201        assert!(err.to_string().contains("column counts"));
1202    }
1203
1204    #[test]
1205    fn concatenate_errors_on_group_dtype_mismatch() {
1206        let a = build_trx(
1207            vec![[0.0, 0.0, 0.0]],
1208            vec![0, 1],
1209            HashMap::new(),
1210            HashMap::new(),
1211            HashMap::from([("bundle".into(), scalar_u32(0))]),
1212            HashMap::new(),
1213        );
1214        let b = build_trx(
1215            vec![[1.0, 1.0, 1.0]],
1216            vec![0, 1],
1217            HashMap::new(),
1218            HashMap::new(),
1219            HashMap::from([(
1220                "bundle".into(),
1221                DataArray::owned_bytes(vec_to_bytes(vec![0u16]), 1, DType::UInt16),
1222            )]),
1223            HashMap::new(),
1224        );
1225        let any_a = AnyTrxFile::F32(a);
1226        let any_b = AnyTrxFile::F32(b);
1227        let err =
1228            concatenate_any_trx(&[&any_a, &any_b], &ConcatenateOptions::default()).unwrap_err();
1229        assert!(err.to_string().contains("group key 'bundle'"));
1230    }
1231
1232    #[test]
1233    fn concatenate_respects_positions_dtype_override() {
1234        let a = build_trx(
1235            vec![[0.0, 0.0, 0.0]],
1236            vec![0, 1],
1237            HashMap::new(),
1238            HashMap::new(),
1239            HashMap::new(),
1240            HashMap::new(),
1241        );
1242        let b = build_trx(
1243            vec![[1.0, 1.0, 1.0]],
1244            vec![0, 1],
1245            HashMap::new(),
1246            HashMap::new(),
1247            HashMap::new(),
1248            HashMap::new(),
1249        );
1250        let any_a = AnyTrxFile::F32(a);
1251        let any_b = AnyTrxFile::F32(b);
1252        let merged = concatenate_any_trx(
1253            &[&any_a, &any_b],
1254            &ConcatenateOptions {
1255                positions_dtype: Some(DType::Float16),
1256                ..Default::default()
1257            },
1258        )
1259        .unwrap();
1260        assert!(matches!(merged, AnyTrxFile::F16(_)));
1261    }
1262
1263    #[test]
1264    fn concatenate_errors_on_group_name_length_mismatch() {
1265        let a = build_trx(
1266            vec![[0.0, 0.0, 0.0]],
1267            vec![0, 1],
1268            HashMap::new(),
1269            HashMap::new(),
1270            HashMap::new(),
1271            HashMap::new(),
1272        );
1273        let b = build_trx(
1274            vec![[1.0, 1.0, 1.0]],
1275            vec![0, 1],
1276            HashMap::new(),
1277            HashMap::new(),
1278            HashMap::new(),
1279            HashMap::new(),
1280        );
1281        let any_a = AnyTrxFile::F32(a);
1282        let any_b = AnyTrxFile::F32(b);
1283        let err = concatenate_any_trx(
1284            &[&any_a, &any_b],
1285            &ConcatenateOptions {
1286                input_group_names: vec![Some("bundle".into())],
1287                ..Default::default()
1288            },
1289        )
1290        .unwrap_err();
1291        assert!(err.to_string().contains("input_group_names length"));
1292    }
1293
1294    #[test]
1295    fn concatenate_group_name_creates_aggregate_group_when_input_has_no_groups() {
1296        let a = build_trx(
1297            vec![[0.0, 0.0, 0.0], [1.0, 0.0, 0.0]],
1298            vec![0, 1, 2],
1299            HashMap::new(),
1300            HashMap::new(),
1301            HashMap::new(),
1302            HashMap::new(),
1303        );
1304        let b = build_trx(
1305            vec![[2.0, 0.0, 0.0]],
1306            vec![0, 1],
1307            HashMap::new(),
1308            HashMap::new(),
1309            HashMap::new(),
1310            HashMap::new(),
1311        );
1312        let any_a = AnyTrxFile::F32(a);
1313        let any_b = AnyTrxFile::F32(b);
1314        let merged = concatenate_any_trx(
1315            &[&any_a, &any_b],
1316            &ConcatenateOptions {
1317                input_group_names: vec![Some("first".into()), None],
1318                ..Default::default()
1319            },
1320        )
1321        .unwrap();
1322
1323        match merged {
1324            AnyTrxFile::F32(trx) => {
1325                assert_eq!(trx.group("first").unwrap(), &[0, 1]);
1326            }
1327            _ => panic!("expected float32 output"),
1328        }
1329    }
1330
1331    #[test]
1332    fn concatenate_group_name_prefixes_existing_groups_without_aggregate() {
1333        let a = build_trx(
1334            vec![[0.0, 0.0, 0.0], [1.0, 1.0, 1.0]],
1335            vec![0, 1, 2],
1336            HashMap::new(),
1337            HashMap::new(),
1338            HashMap::from([(
1339                "bundle".into(),
1340                DataArray::owned_bytes(vec_to_bytes(vec![0u32, 1u32]), 1, DType::UInt32),
1341            )]),
1342            HashMap::new(),
1343        );
1344        let b = build_trx(
1345            vec![[2.0, 2.0, 2.0]],
1346            vec![0, 1],
1347            HashMap::new(),
1348            HashMap::new(),
1349            HashMap::new(),
1350            HashMap::new(),
1351        );
1352        let any_a = AnyTrxFile::F32(a);
1353        let any_b = AnyTrxFile::F32(b);
1354        let merged = concatenate_any_trx(
1355            &[&any_a, &any_b],
1356            &ConcatenateOptions {
1357                input_group_names: vec![Some("Prefix".into()), None],
1358                ..Default::default()
1359            },
1360        )
1361        .unwrap();
1362
1363        match merged {
1364            AnyTrxFile::F32(trx) => {
1365                assert_eq!(trx.group("Prefixbundle").unwrap(), &[0, 1]);
1366                assert!(!trx.group_names().contains(&"Prefix"));
1367            }
1368            _ => panic!("expected float32 output"),
1369        }
1370    }
1371
1372    #[test]
1373    fn concatenate_group_name_blank_entry_is_ignored() {
1374        let a = build_trx(
1375            vec![[0.0, 0.0, 0.0]],
1376            vec![0, 1],
1377            HashMap::new(),
1378            HashMap::new(),
1379            HashMap::new(),
1380            HashMap::new(),
1381        );
1382        let b = build_trx(
1383            vec![[1.0, 1.0, 1.0]],
1384            vec![0, 1],
1385            HashMap::new(),
1386            HashMap::new(),
1387            HashMap::new(),
1388            HashMap::new(),
1389        );
1390        let any_a = AnyTrxFile::F32(a);
1391        let any_b = AnyTrxFile::F32(b);
1392        let merged = concatenate_any_trx(
1393            &[&any_a, &any_b],
1394            &ConcatenateOptions {
1395                input_group_names: vec![Some("   ".into()), None],
1396                ..Default::default()
1397            },
1398        )
1399        .unwrap();
1400
1401        match merged {
1402            AnyTrxFile::F32(trx) => assert!(trx.group_names().is_empty()),
1403            _ => panic!("expected float32 output"),
1404        }
1405    }
1406
1407    #[test]
1408    fn concatenate_overlapping_prefixed_group_names_union_members() {
1409        let a = build_trx(
1410            vec![[0.0, 0.0, 0.0]],
1411            vec![0, 1],
1412            HashMap::new(),
1413            HashMap::new(),
1414            HashMap::from([("One".into(), scalar_u32(0))]),
1415            HashMap::new(),
1416        );
1417        let b = build_trx(
1418            vec![[1.0, 1.0, 1.0]],
1419            vec![0, 1],
1420            HashMap::new(),
1421            HashMap::new(),
1422            HashMap::from([("One".into(), scalar_u32(0))]),
1423            HashMap::new(),
1424        );
1425        let any_a = AnyTrxFile::F32(a);
1426        let any_b = AnyTrxFile::F32(b);
1427        let merged = concatenate_any_trx(
1428            &[&any_a, &any_b],
1429            &ConcatenateOptions {
1430                input_group_names: vec![Some("Same".into()), Some("Same".into())],
1431                ..Default::default()
1432            },
1433        )
1434        .unwrap();
1435
1436        match merged {
1437            AnyTrxFile::F32(trx) => assert_eq!(trx.group("SameOne").unwrap(), &[0, 1]),
1438            _ => panic!("expected float32 output"),
1439        }
1440    }
1441}