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