Skip to main content

laddu_data/io/root/
mod.rs

1use std::{
2    path::{Path, PathBuf},
3    sync::{
4        Arc,
5        mpsc::{self, Receiver, Sender, SyncSender},
6    },
7    thread::{self, JoinHandle},
8};
9
10use laddu_physics::vectors::RealVec4;
11use oxyroot::{Branch, ReaderTree, RootFile, WriterTree};
12
13use crate::{
14    LadduDataError, LadduDataResult, Name,
15    columns::{ColumnBuffer, ColumnDType, ColumnValue},
16    data::EventBatch,
17    io::{
18        DataFragment, EventSink, EventSource, FragmentedSource, OutputMode, OutputPath, ReadPlan,
19        SinkState, SourceBuild, SourceBuildOptions, SourceCapabilities, WritePlan, build_source,
20        fragmented_batches, sink_error, source_error,
21    },
22    schema::{
23        ColumnInfo, ColumnType, PhysicalColumnRole, PhysicalSchemaPlan, Precision, Schema,
24        SchemaColumnNames, SchemaInferenceOptions, SchemaWriteOptions, WriteWeightColumn,
25    },
26};
27
28mod encode;
29
30use encode::root_output_columns;
31
32/// Event source backed by one or more ROOT TTrees.
33#[derive(Clone, Debug)]
34pub struct RootSource {
35    files: Arc<[Arc<PathBuf>]>,
36    tree_name: Name,
37    schema: Arc<Schema>,
38    options: RootReadOptions,
39}
40
41/// Schema, validation, glob, and tree-selection options for ROOT reads.
42#[derive(Clone, Debug)]
43pub struct RootReadOptions {
44    /// Infer a logical schema when none is supplied.
45    pub infer_schema: bool,
46    /// Validate required columns in every matched file.
47    pub validate_all_files: bool,
48    /// Sort glob results for deterministic global row order.
49    pub sort_glob: bool,
50    /// TTree selection policy.
51    pub tree: RootTreeSelection,
52    /// Logical schema inference options.
53    pub schema_inference: SchemaInferenceOptions,
54}
55
56impl Default for RootReadOptions {
57    fn default() -> Self {
58        Self {
59            infer_schema: true,
60            validate_all_files: true,
61            sort_glob: true,
62            tree: RootTreeSelection::First,
63            schema_inference: SchemaInferenceOptions::default(),
64        }
65    }
66}
67
68/// Policy for selecting a TTree from each ROOT file.
69#[derive(Clone, Debug, Default)]
70pub enum RootTreeSelection {
71    /// Select the first TTree.
72    #[default]
73    First,
74    /// Select a TTree by name.
75    Named(Name),
76}
77
78/// Key identifying one TTree within a ROOT file.
79#[derive(Clone, Debug)]
80pub struct RootFragmentKey {
81    /// Input file path.
82    pub file: Arc<PathBuf>,
83    /// TTree name.
84    pub tree_name: Name,
85}
86
87/// Introspection metadata for one ROOT branch.
88#[derive(Clone, Debug)]
89pub struct RootColumnInfo {
90    /// Branch name.
91    pub name: Name,
92    /// Rust item type reported by the reader.
93    pub item_type_name: String,
94    /// ROOT interpretation string.
95    pub interpretation: String,
96    /// Number of branch entries.
97    pub entries: i64,
98}
99
100impl RootSource {
101    /// Opens files matching a glob with default options.
102    ///
103    /// # Errors
104    ///
105    /// Returns [`LadduDataError`] when the glob is invalid or empty, a ROOT
106    /// file or tree cannot be read, or schemas are incompatible.
107    pub fn open(pattern: impl AsRef<str>) -> LadduDataResult<Self> {
108        Self::builder(pattern).build()
109    }
110
111    /// Creates a configurable source builder for a file glob.
112    pub fn builder(pattern: impl AsRef<str>) -> RootSourceBuilder {
113        RootSourceBuilder {
114            pattern: pattern.as_ref().to_owned(),
115            schema: None,
116            options: RootReadOptions::default(),
117        }
118    }
119
120    /// Returns matched files in global row order.
121    pub fn files(&self) -> &[Arc<PathBuf>] {
122        &self.files
123    }
124
125    /// Returns the selected TTree name.
126    pub fn tree_name(&self) -> &str {
127        self.tree_name.as_ref()
128    }
129
130    /// Lists TTrees in one ROOT file.
131    ///
132    /// # Errors
133    ///
134    /// Returns [`LadduDataError`] when the ROOT file cannot be opened or read.
135    pub fn tree_names(path: impl AsRef<Path>) -> LadduDataResult<Vec<Name>> {
136        let path = path.as_ref();
137        let mut file = RootFile::open(path)
138            .map_err(|error| source_error("open ROOT file", path.display(), error))?;
139        let key_names: Vec<String> = file.keys_name().map(str::to_owned).collect();
140
141        let mut out = Vec::new();
142
143        for name in key_names {
144            if file.get_tree(&name).is_ok() {
145                out.push(Name::from(name));
146            }
147        }
148
149        Ok(out)
150    }
151
152    /// Lists branch metadata for a selected or first TTree.
153    ///
154    /// # Errors
155    ///
156    /// Returns [`LadduDataError`] when the file or tree cannot be read, no tree
157    /// exists, or branch metadata is invalid.
158    pub fn columns(
159        path: impl AsRef<Path>,
160        tree: Option<&str>,
161    ) -> LadduDataResult<Vec<RootColumnInfo>> {
162        let path = path.as_ref();
163        let mut file = RootFile::open(path)
164            .map_err(|error| source_error("open ROOT file", path.display(), error))?;
165
166        let tree_name = match tree {
167            Some(name) => Name::from(name),
168            None => first_tree_name(&mut file, path)?,
169        };
170
171        let tree = file.get_tree(tree_name.as_ref()).map_err(|error| {
172            source_error(
173                "read ROOT tree",
174                format!("{}::{tree_name}", path.display()),
175                error,
176            )
177        })?;
178
179        Ok(tree
180            .branches_r()
181            .into_iter()
182            .map(|branch| RootColumnInfo {
183                name: Name::from(branch.name()),
184                item_type_name: branch.item_type_name(),
185                interpretation: branch.interpretation(),
186                entries: branch.entries(),
187            })
188            .collect())
189    }
190}
191
192/// Builder for a [`RootSource`].
193pub struct RootSourceBuilder {
194    pattern: String,
195    schema: Option<Arc<Schema>>,
196    options: RootReadOptions,
197}
198
199impl RootSourceBuilder {
200    /// Supplies an explicit logical schema and disables inference.
201    pub fn schema(mut self, schema: Arc<Schema>) -> Self {
202        self.schema = Some(schema);
203        self.options.infer_schema = false;
204        self
205    }
206
207    /// Enables or disables logical schema inference.
208    pub fn infer_schema(mut self, value: bool) -> Self {
209        self.options.infer_schema = value;
210        self
211    }
212
213    /// Selects a TTree by name.
214    pub fn tree(mut self, name: impl Into<Name>) -> Self {
215        self.options.tree = RootTreeSelection::Named(name.into());
216        self
217    }
218
219    /// Selects the first TTree.
220    pub fn first_tree(mut self) -> Self {
221        self.options.tree = RootTreeSelection::First;
222        self
223    }
224
225    /// Requires a physical weight column during inference.
226    pub fn require_weight(mut self, value: bool) -> Self {
227        self.options.schema_inference.require_weight = value;
228        self
229    }
230
231    /// Chooses whether every matched file is schema-validated eagerly.
232    pub fn validate_all_files(mut self, value: bool) -> Self {
233        self.options.validate_all_files = value;
234        self
235    }
236
237    /// Chooses whether matched paths are sorted.
238    pub fn sort_glob(mut self, value: bool) -> Self {
239        self.options.sort_glob = value;
240        self
241    }
242
243    /// Replaces logical schema inference options.
244    pub fn schema_inference(mut self, options: SchemaInferenceOptions) -> Self {
245        self.options.schema_inference = options;
246        self
247    }
248
249    /// Resolves files and tree, validates schema, and builds the source.
250    ///
251    /// # Errors
252    ///
253    /// Returns [`LadduDataError`] when the glob is invalid or empty, files or
254    /// trees cannot be read, schema inference fails, or files disagree.
255    pub fn build(self) -> LadduDataResult<RootSource> {
256        let RootSourceBuilder {
257            pattern,
258            schema: explicit_schema,
259            options,
260        } = self;
261        let tree_selection = options.tree.clone();
262        let infer_options = options.schema_inference.clone();
263        let validate_options = options.schema_inference.clone();
264        let SourceBuild {
265            files,
266            context: tree_name,
267            schema,
268        } = build_source(
269            SourceBuildOptions {
270                pattern: &pattern,
271                sort: options.sort_glob,
272                format: "ROOT",
273                explicit_schema,
274                infer_schema: options.infer_schema,
275                validate_all_files: options.validate_all_files,
276            },
277            move |path| resolve_tree_name(path, &tree_selection),
278            move |path, tree_name| {
279                let columns = root_columns(path, tree_name.as_ref())?;
280                Schema::infer_from_columns(
281                    columns.iter().map(OwnedColumnInfo::as_column_info),
282                    &infer_options,
283                )
284            },
285            move |path, schema, tree_name| {
286                validate_root_file(path, tree_name.as_ref(), schema, &validate_options)
287            },
288        )?;
289
290        Ok(RootSource {
291            files,
292            tree_name,
293            schema,
294            options,
295        })
296    }
297}
298
299impl EventSource for RootSource {
300    fn schema(&self) -> LadduDataResult<Arc<Schema>> {
301        Ok(Arc::clone(&self.schema))
302    }
303
304    fn capabilities(&self) -> SourceCapabilities {
305        SourceCapabilities {
306            exact_len: true,
307            exact_weighted_total: false,
308            random_access: false,
309            deterministic_partitioning: true,
310            predicate_pushdown: false,
311            projection_pushdown: true,
312            streaming: true,
313        }
314    }
315
316    fn num_events(&self) -> LadduDataResult<Option<u64>> {
317        Ok(Some(self.fragments()?.iter().map(|f| f.rows).sum()))
318    }
319
320    fn batches(
321        &self,
322        plan: ReadPlan,
323    ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
324        fragmented_batches(Arc::new(self.clone()), plan)
325    }
326}
327
328impl FragmentedSource for RootSource {
329    type Key = RootFragmentKey;
330
331    fn fragments(&self) -> LadduDataResult<Vec<DataFragment<Self::Key>>> {
332        let mut fragments = Vec::new();
333        let mut global_start = 0_u64;
334
335        for path in self.files.iter() {
336            let resource = path.as_ref().display().to_string();
337            let mut file = RootFile::open(path.as_ref())
338                .map_err(|error| source_error("open ROOT file", &resource, error))?;
339            let tree = file.get_tree(self.tree_name.as_ref()).map_err(|error| {
340                source_error(
341                    "read ROOT tree",
342                    format!("{resource}::{}", self.tree_name),
343                    error,
344                )
345            })?;
346
347            let tree_resource = format!("{resource}::{}", self.tree_name);
348            let rows = usize_from_i64(tree.entries(), "negative TTree entry count", &tree_resource)?
349                as u64;
350
351            fragments.push(DataFragment {
352                key: RootFragmentKey {
353                    file: Arc::clone(path),
354                    tree_name: self.tree_name.clone(),
355                },
356                global_start,
357                rows,
358            });
359
360            global_start += rows;
361        }
362
363        Ok(fragments)
364    }
365
366    fn read_fragment_range(
367        &self,
368        key: &Self::Key,
369        local_start: usize,
370        local_len: usize,
371        chunk_size: Option<usize>,
372    ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
373        if matches!(chunk_size, Some(0)) {
374            return Err(LadduDataError::InvalidArgument(
375                "chunk_size must be nonzero",
376            ));
377        }
378
379        Ok(Box::new(RootBatchIter::spawn(
380            Arc::clone(&self.schema),
381            self.options.clone(),
382            key.clone(),
383            local_start,
384            local_len,
385            chunk_size,
386        )))
387    }
388}
389
390struct RootBatchIter {
391    rx: Receiver<LadduDataResult<EventBatch>>,
392    state: RootBatchState,
393}
394
395enum RootBatchState {
396    Receiving(JoinHandle<()>),
397    Done,
398}
399
400impl RootBatchIter {
401    fn spawn(
402        schema: Arc<Schema>,
403        options: RootReadOptions,
404        key: RootFragmentKey,
405        local_start: usize,
406        local_len: usize,
407        chunk_size: Option<usize>,
408    ) -> Self {
409        // Keep at most one decoded batch ahead of the consumer so ROOT I/O
410        // cannot silently exceed the dataset's memory-derived chunk budget.
411        let (tx, rx) = mpsc::sync_channel(1);
412
413        let handle = thread::spawn(move || {
414            if let Err(err) = read_root_range_and_send_batches(
415                schema,
416                options,
417                key,
418                local_start,
419                local_len,
420                chunk_size,
421                tx.clone(),
422            ) {
423                let _ = tx.send(Err(err));
424            }
425        });
426
427        Self {
428            rx,
429            state: RootBatchState::Receiving(handle),
430        }
431    }
432
433    fn join_if_needed(&mut self) -> Option<LadduDataResult<EventBatch>> {
434        let state = std::mem::replace(&mut self.state, RootBatchState::Done);
435        match state {
436            RootBatchState::Done => None,
437            RootBatchState::Receiving(handle) => {
438                if handle.join().is_err() {
439                    Some(Err(LadduDataError::Source(
440                        "ROOT reader thread panicked".into(),
441                    )))
442                } else {
443                    None
444                }
445            }
446        }
447    }
448}
449
450impl Iterator for RootBatchIter {
451    type Item = LadduDataResult<EventBatch>;
452
453    fn next(&mut self) -> Option<Self::Item> {
454        if matches!(self.state, RootBatchState::Done) {
455            return None;
456        }
457
458        match self.rx.recv() {
459            Ok(item) => Some(item),
460            Err(_) => self.join_if_needed(),
461        }
462    }
463}
464
465fn read_root_range_and_send_batches(
466    schema: Arc<Schema>,
467    options: RootReadOptions,
468    key: RootFragmentKey,
469    local_start: usize,
470    local_len: usize,
471    chunk_size: Option<usize>,
472    tx: SyncSender<LadduDataResult<EventBatch>>,
473) -> LadduDataResult<()> {
474    let resource = format!("{}::{}", key.file.as_ref().display(), key.tree_name);
475    let mut file = RootFile::open(key.file.as_ref())
476        .map_err(|error| source_error("open ROOT file", &resource, error))?;
477    let tree = file
478        .get_tree(key.tree_name.as_ref())
479        .map_err(|error| source_error("read ROOT tree", &resource, error))?;
480
481    let mut readers =
482        RootColumnReaders::new(&tree, &schema, &options.schema_inference.column_names)?;
483
484    for _ in 0..local_start {
485        readers.skip_one()?;
486    }
487
488    let mut remaining = local_len;
489    let batch_size = chunk_size.unwrap_or(local_len.max(1));
490
491    while remaining > 0 {
492        let take = remaining.min(batch_size);
493        let batch = readers.read_batch(Arc::clone(&schema), take)?;
494
495        tx.send(Ok(batch))
496            .map_err(|error| source_error("send ROOT batch", &resource, error))?;
497
498        remaining -= take;
499    }
500
501    Ok(())
502}
503
504struct RootColumnReaders<'a> {
505    p4s: Vec<[RootFloatIter<'a>; 4]>,
506    scalars: Vec<RootFloatIter<'a>>,
507    columns: Vec<RootIntegerIter<'a>>,
508    weights: Option<RootFloatIter<'a>>,
509}
510
511impl<'a> RootColumnReaders<'a> {
512    fn new(
513        tree: &'a ReaderTree,
514        schema: &Schema,
515        column_names: &SchemaColumnNames,
516    ) -> LadduDataResult<Self> {
517        let plan = PhysicalSchemaPlan::for_read(schema, column_names);
518        let mut p4s: Vec<[Option<RootFloatIter<'a>>; 4]> = (0..schema.n_p4s())
519            .map(|_| std::array::from_fn(|_| None))
520            .collect();
521        let mut scalars: Vec<Option<RootFloatIter<'a>>> =
522            (0..schema.n_scalars()).map(|_| None).collect();
523        let mut weights = None;
524
525        for column in plan.columns() {
526            if matches!(column.role(), PhysicalColumnRole::Column { .. }) {
527                continue;
528            }
529            let reader = open_float_reader(tree, column.name().as_ref())?;
530            match column.role() {
531                PhysicalColumnRole::Column { .. } => continue,
532                PhysicalColumnRole::P4 { index, component } => {
533                    p4s[index][component] = Some(reader);
534                }
535                PhysicalColumnRole::Scalar { index } => {
536                    scalars[index] = Some(reader);
537                }
538                PhysicalColumnRole::Weight => {
539                    weights = Some(reader);
540                }
541            }
542        }
543
544        let p4s = p4s
545            .into_iter()
546            .map(|parts| {
547                let [e, px, py, pz] = parts;
548                Ok([
549                    e.ok_or_else(|| {
550                        LadduDataError::Source(
551                            "physical schema plan did not bind ROOT E branch".into(),
552                        )
553                    })?,
554                    px.ok_or_else(|| {
555                        LadduDataError::Source(
556                            "physical schema plan did not bind ROOT px branch".into(),
557                        )
558                    })?,
559                    py.ok_or_else(|| {
560                        LadduDataError::Source(
561                            "physical schema plan did not bind ROOT py branch".into(),
562                        )
563                    })?,
564                    pz.ok_or_else(|| {
565                        LadduDataError::Source(
566                            "physical schema plan did not bind ROOT pz branch".into(),
567                        )
568                    })?,
569                ])
570            })
571            .collect::<LadduDataResult<Vec<_>>>()?;
572        let scalars = scalars
573            .into_iter()
574            .map(|reader| {
575                reader.ok_or_else(|| {
576                    LadduDataError::Source(
577                        "physical schema plan did not bind ROOT scalar branch".into(),
578                    )
579                })
580            })
581            .collect::<LadduDataResult<Vec<_>>>()?;
582
583        let columns = schema
584            .columns()
585            .iter()
586            .map(|(name, dtype)| open_integer_reader(tree, name, *dtype))
587            .collect::<LadduDataResult<Vec<_>>>()?;
588        Ok(Self {
589            p4s,
590            scalars,
591            columns,
592            weights,
593        })
594    }
595
596    fn skip_one(&mut self) -> LadduDataResult<()> {
597        for [e, px, py, pz] in self.p4s.iter_mut() {
598            e.next_f64()?;
599            px.next_f64()?;
600            py.next_f64()?;
601            pz.next_f64()?;
602        }
603
604        for scalar in self.scalars.iter_mut() {
605            scalar.next_f64()?;
606        }
607
608        for column in &mut self.columns {
609            column.next_value()?;
610        }
611        if let Some(weights) = self.weights.as_mut() {
612            weights.next_f64()?;
613        }
614
615        Ok(())
616    }
617
618    fn read_batch(&mut self, schema: Arc<Schema>, len: usize) -> LadduDataResult<EventBatch> {
619        let mut p4s = Vec::with_capacity(schema.n_p4s());
620        let mut scalars = Vec::with_capacity(schema.n_scalars());
621
622        for [e, px, py, pz] in self.p4s.iter_mut() {
623            let mut col = Vec::with_capacity(len);
624            for _ in 0..len {
625                col.push(RealVec4 {
626                    e: e.next_f64()?,
627                    px: px.next_f64()?,
628                    py: py.next_f64()?,
629                    pz: pz.next_f64()?,
630                });
631            }
632            p4s.push(col);
633        }
634
635        for reader in self.scalars.iter_mut() {
636            let mut col = Vec::with_capacity(len);
637            for _ in 0..len {
638                col.push(reader.next_f64()?);
639            }
640            scalars.push(col);
641        }
642
643        let weights = if let Some(reader) = self.weights.as_mut() {
644            let mut col = Vec::with_capacity(len);
645            for _ in 0..len {
646                col.push(reader.next_f64()?);
647            }
648            Some(col)
649        } else {
650            None
651        };
652
653        let mut columns = Vec::with_capacity(self.columns.len());
654        for (reader, (_, dtype)) in self.columns.iter_mut().zip(schema.columns()) {
655            let mut buffer = ColumnBuffer::new(*dtype, len);
656            for _ in 0..len {
657                buffer.push(reader.next_value()?)?;
658            }
659            columns.push(buffer.finish());
660        }
661        EventBatch::new_with_columns_and_len(
662            schema,
663            p4s.into_iter().map(Arc::from).collect(),
664            scalars.into_iter().map(Arc::from).collect(),
665            columns,
666            weights.map(Arc::from),
667            len,
668        )
669    }
670}
671
672enum RootFloatIter<'a> {
673    F64 {
674        name: Name,
675        iter: Box<dyn Iterator<Item = f64> + 'a>,
676    },
677    F32 {
678        name: Name,
679        iter: Box<dyn Iterator<Item = f32> + 'a>,
680    },
681}
682
683impl<'a> RootFloatIter<'a> {
684    fn next_f64(&mut self) -> LadduDataResult<f64> {
685        match self {
686            Self::F64 { name, iter } => iter
687                .next()
688                .ok_or_else(|| source_error("read ROOT branch", name, "ROOT branch ended early")),
689            Self::F32 { name, iter } => iter
690                .next()
691                .map(f64::from)
692                .ok_or_else(|| source_error("read ROOT branch", name, "ROOT branch ended early")),
693        }
694    }
695}
696
697fn open_float_reader<'a>(tree: &'a ReaderTree, name: &str) -> LadduDataResult<RootFloatIter<'a>> {
698    let branch =
699        find_branch(tree, name).ok_or_else(|| LadduDataError::MissingColumn(Name::from(name)))?;
700
701    match root_column_type(branch) {
702        ColumnType::F64 => Ok(RootFloatIter::F64 {
703            name: Name::from(name),
704            iter: Box::new(
705                branch
706                    .as_iter::<f64>()
707                    .map_err(|error| source_error("open ROOT branch", name, error))?,
708            ),
709        }),
710        ColumnType::F32 => Ok(RootFloatIter::F32 {
711            name: Name::from(name),
712            iter: Box::new(
713                branch
714                    .as_iter::<f32>()
715                    .map_err(|error| source_error("open ROOT branch", name, error))?,
716            ),
717        }),
718        ColumnType::Other | ColumnType::Integer(_) => Err(source_error(
719            "decode ROOT branch",
720            name,
721            format!(
722                "column {name} has unsupported ROOT type {} interpreted as {}",
723                branch.item_type_name(),
724                branch.interpretation()
725            ),
726        )),
727    }
728}
729
730fn find_branch<'a>(tree: &'a ReaderTree, name: &str) -> Option<&'a Branch> {
731    tree.branch(name).or_else(|| {
732        tree.branches_r()
733            .into_iter()
734            .find(|branch| branch.name() == name)
735    })
736}
737
738#[derive(Clone, Debug)]
739struct OwnedColumnInfo {
740    name: Name,
741    dtype: ColumnType,
742}
743
744impl OwnedColumnInfo {
745    fn as_column_info(&self) -> ColumnInfo<'_> {
746        ColumnInfo {
747            name: self.name.as_ref(),
748            dtype: self.dtype,
749        }
750    }
751}
752
753fn root_columns(path: &Path, tree_name: &str) -> LadduDataResult<Vec<OwnedColumnInfo>> {
754    let resource = format!("{}::{tree_name}", path.display());
755    let mut file = RootFile::open(path)
756        .map_err(|error| source_error("open ROOT file", path.display(), error))?;
757    let tree = file
758        .get_tree(tree_name)
759        .map_err(|error| source_error("read ROOT tree", &resource, error))?;
760
761    Ok(tree
762        .branches_r()
763        .into_iter()
764        .map(|branch| OwnedColumnInfo {
765            name: Name::from(branch.name()),
766            dtype: root_column_type(branch),
767        })
768        .collect())
769}
770
771fn validate_root_file(
772    path: &Path,
773    tree_name: &str,
774    schema: &Schema,
775    options: &SchemaInferenceOptions,
776) -> LadduDataResult<()> {
777    let columns = root_columns(path, tree_name)?;
778
779    schema.validate_required_columns(columns.iter().map(OwnedColumnInfo::as_column_info), options)
780}
781
782fn root_column_type(branch: &Branch) -> ColumnType {
783    match branch.interpretation().as_str() {
784        "f64" => ColumnType::F64,
785        "f32" => ColumnType::F32,
786        "i8" => ColumnType::Integer(ColumnDType::I8),
787        "u8" => ColumnType::Integer(ColumnDType::U8),
788        "i16" => ColumnType::Integer(ColumnDType::I16),
789        "u16" => ColumnType::Integer(ColumnDType::U16),
790        "i32" => ColumnType::Integer(ColumnDType::I32),
791        "u32" => ColumnType::Integer(ColumnDType::U32),
792        "i64" => ColumnType::Integer(ColumnDType::I64),
793        "u64" => ColumnType::Integer(ColumnDType::U64),
794
795        _ => match branch.item_type_name().as_str() {
796            "double" | "Double_t" | "ROOT::Double_t" => ColumnType::F64,
797            "float" | "Float_t" | "ROOT::Float_t" => ColumnType::F32,
798            _ => ColumnType::Other,
799        },
800    }
801}
802
803fn resolve_tree_name(path: &Path, selection: &RootTreeSelection) -> LadduDataResult<Name> {
804    let mut file = RootFile::open(path)
805        .map_err(|error| source_error("open ROOT file", path.display(), error))?;
806
807    match selection {
808        RootTreeSelection::Named(name) => {
809            file.get_tree(name.as_ref()).map_err(|error| {
810                source_error(
811                    "read ROOT tree",
812                    format!("{}::{name}", path.display()),
813                    error,
814                )
815            })?;
816            Ok(name.clone())
817        }
818        RootTreeSelection::First => first_tree_name(&mut file, path),
819    }
820}
821
822fn first_tree_name(file: &mut RootFile, path: &Path) -> LadduDataResult<Name> {
823    let key_names: Vec<String> = file.keys_name().map(str::to_owned).collect();
824
825    for name in key_names {
826        if file.get_tree(&name).is_ok() {
827            return Ok(Name::from(name));
828        }
829    }
830
831    Err(source_error(
832        "resolve first ROOT tree",
833        path.display(),
834        "no TTree found in ROOT file",
835    ))
836}
837
838fn usize_from_i64(value: i64, message: &'static str, resource: &str) -> LadduDataResult<usize> {
839    if value < 0 {
840        return Err(source_error("read ROOT entry count", resource, message));
841    }
842
843    usize::try_from(value).map_err(|_| {
844        source_error(
845            "read ROOT entry count",
846            resource,
847            "entry count overflows usize",
848        )
849    })
850}
851
852/// Event sink that writes a ROOT TTree on a background thread.
853pub struct RootSink {
854    output: OutputPath,
855    options: RootWriteOptions,
856    resolved_path: Option<PathBuf>,
857    event_schema: Option<Arc<Schema>>,
858    senders: Option<Vec<RootColumnSender>>,
859    writer_thread: Option<JoinHandle<LadduDataResult<()>>>,
860    state: SinkState,
861}
862
863/// ROOT tree and physical schema write options.
864#[derive(Clone, Debug)]
865pub struct RootWriteOptions {
866    /// Output TTree name.
867    pub tree_name: Name,
868    /// Physical schema write options.
869    pub schema_write: SchemaWriteOptions,
870}
871
872impl Default for RootWriteOptions {
873    fn default() -> Self {
874        Self {
875            tree_name: Name::from("tree"),
876            schema_write: SchemaWriteOptions::default(),
877        }
878    }
879}
880
881impl RootSink {
882    /// Creates a sink with default options.
883    pub fn create(path: impl Into<PathBuf>) -> Self {
884        Self::builder(path).build()
885    }
886
887    /// Creates a configurable sink builder.
888    pub fn builder(path: impl Into<PathBuf>) -> RootSinkBuilder {
889        RootSinkBuilder {
890            output: OutputPath::new(path),
891            options: RootWriteOptions::default(),
892        }
893    }
894
895    /// Returns the concrete path after writing has begun.
896    pub fn resolved_path(&self) -> Option<&Path> {
897        self.resolved_path.as_deref()
898    }
899}
900
901/// Builder for a [`RootSink`].
902pub struct RootSinkBuilder {
903    output: OutputPath,
904    options: RootWriteOptions,
905}
906
907impl RootSinkBuilder {
908    /// Sets the output path mode.
909    pub fn output_mode(mut self, mode: OutputMode) -> Self {
910        self.output = self.output.with_mode(mode);
911        self
912    }
913
914    /// Selects single-file output.
915    pub fn single_file(self) -> Self {
916        self.output_mode(OutputMode::SingleFile)
917    }
918
919    /// Selects one output file per rank.
920    pub fn per_rank_files(self) -> Self {
921        self.output_mode(OutputMode::PerRankFiles)
922    }
923
924    /// Selects output mode from the write plan.
925    pub fn auto_output(self) -> Self {
926        self.output_mode(OutputMode::Auto)
927    }
928
929    /// Sets the output TTree name.
930    pub fn tree(mut self, name: impl Into<Name>) -> Self {
931        self.options.tree_name = name.into();
932        self
933    }
934
935    /// Replaces physical schema write options.
936    pub fn schema_write(mut self, options: SchemaWriteOptions) -> Self {
937        self.options.schema_write = options;
938        self
939    }
940
941    /// Sets physical column naming conventions.
942    pub fn column_names(mut self, column_names: SchemaColumnNames) -> Self {
943        self.options.schema_write.column_names = column_names;
944        self
945    }
946
947    /// Sets floating-point output precision.
948    pub fn precision(mut self, precision: Precision) -> Self {
949        self.options.schema_write.precision = precision;
950        self
951    }
952
953    /// Sets the weight-column emission policy.
954    pub fn write_weight_column(mut self, value: WriteWeightColumn) -> Self {
955        self.options.schema_write.write_weight_column = value;
956        self
957    }
958
959    /// Builds the sink.
960    pub fn build(self) -> RootSink {
961        RootSink {
962            output: self.output,
963            options: self.options,
964            resolved_path: None,
965            event_schema: None,
966            senders: None,
967            writer_thread: None,
968            state: SinkState::Idle,
969        }
970    }
971}
972
973impl EventSink for RootSink {
974    fn begin(&mut self, schema: Arc<Schema>, plan: WritePlan) -> LadduDataResult<()> {
975        schema.validate_column_names(&self.options.schema_write.column_names)?;
976        match self.state {
977            SinkState::Idle => {}
978            SinkState::Writing => {
979                return Err(LadduDataError::Sink("ROOT sink already initialized".into()));
980            }
981            SinkState::Failed => {
982                return Err(LadduDataError::Sink(
983                    "ROOT sink requires abort after failure".into(),
984                ));
985            }
986        }
987
988        let path = self.output.resolve(plan, "root")?;
989        OutputPath::create_parent_dirs(&path)?;
990
991        let columns = root_output_columns(
992            &schema,
993            self.options.schema_write.write_weight_column,
994            &self.options.schema_write,
995        );
996
997        let (senders, receivers) =
998            root_channels(&columns, &schema, self.options.schema_write.precision);
999
1000        let writer_path = path.clone();
1001        let tree_name = self.options.tree_name.clone();
1002
1003        let handle = thread::spawn(move || write_root_tree(writer_path, tree_name, receivers));
1004
1005        self.resolved_path = Some(path);
1006        self.event_schema = Some(schema);
1007        self.senders = Some(senders);
1008        self.writer_thread = Some(handle);
1009        self.state = SinkState::Writing;
1010
1011        Ok(())
1012    }
1013
1014    fn write_batch(&mut self, batch: &EventBatch) -> LadduDataResult<()> {
1015        if !matches!(self.state, SinkState::Writing) {
1016            return Err(LadduDataError::Sink(
1017                match self.state {
1018                    SinkState::Idle => "ROOT sink not initialized",
1019                    SinkState::Failed => "ROOT sink requires abort after failure",
1020                    SinkState::Writing => unreachable!(),
1021                }
1022                .into(),
1023            ));
1024        }
1025
1026        let event_schema = self
1027            .event_schema
1028            .as_ref()
1029            .ok_or_else(|| LadduDataError::Sink("ROOT sink not initialized".into()))?;
1030
1031        if event_schema.as_ref() != batch.schema().as_ref() {
1032            return Err(LadduDataError::Sink(
1033                "batch schema does not match ROOT sink schema".into(),
1034            ));
1035        }
1036
1037        let senders = self
1038            .senders
1039            .as_ref()
1040            .ok_or_else(|| LadduDataError::Sink("ROOT sink not initialized".into()))?;
1041
1042        let plan = PhysicalSchemaPlan::for_write(
1043            batch.schema(),
1044            &self.options.schema_write,
1045            self.options.schema_write.write_weight_column,
1046        );
1047
1048        for row in 0..batch.len() {
1049            for (index, column) in plan.columns().iter().enumerate() {
1050                if let PhysicalColumnRole::Column { index: logical, .. } = column.role() {
1051                    if let Err(error) = senders[index].send_integer(batch.column(logical).at(row)) {
1052                        self.state = SinkState::Failed;
1053                        return Err(sink_error("send ROOT integer column", column.name(), error));
1054                    }
1055                    continue;
1056                }
1057                let value = match column.role() {
1058                    PhysicalColumnRole::Column { .. } => unreachable!("handled above"),
1059                    PhysicalColumnRole::P4 { index, component } => {
1060                        batch.p4_at(index, row).components()[component]
1061                    }
1062                    PhysicalColumnRole::Scalar { index } => batch.scalar_at(index, row),
1063                    PhysicalColumnRole::Weight => batch.weights_at(row),
1064                };
1065                if let Err(error) = senders[index].send_float(value) {
1066                    self.state = SinkState::Failed;
1067                    return Err(sink_error("send ROOT column", column.name(), error));
1068                }
1069            }
1070        }
1071
1072        Ok(())
1073    }
1074
1075    fn finish(&mut self) -> LadduDataResult<()> {
1076        if matches!(self.state, SinkState::Idle) {
1077            return Ok(());
1078        }
1079        if matches!(self.state, SinkState::Failed) {
1080            return Err(LadduDataError::Sink(
1081                "ROOT sink requires abort after failure".into(),
1082            ));
1083        }
1084
1085        self.senders.take();
1086
1087        if let Some(handle) = self.writer_thread.take() {
1088            match handle.join() {
1089                Ok(Ok(())) => {}
1090                Ok(Err(error)) => {
1091                    self.state = SinkState::Failed;
1092                    return Err(error);
1093                }
1094                Err(_) => {
1095                    self.state = SinkState::Failed;
1096                    return Err(LadduDataError::Sink("ROOT writer thread panicked".into()));
1097                }
1098            }
1099        }
1100
1101        self.event_schema = None;
1102        self.state = SinkState::Idle;
1103        Ok(())
1104    }
1105
1106    fn abort(&mut self) -> LadduDataResult<()> {
1107        // Disconnect all channels before joining so the ROOT writer can finish
1108        // its iterator and close the file. The file is deliberately retained
1109        // and may contain an incomplete tree.
1110        self.senders.take();
1111        let result = if let Some(handle) = self.writer_thread.take() {
1112            match handle.join() {
1113                Ok(result) => result,
1114                Err(_) => Err(LadduDataError::Sink("ROOT writer thread panicked".into())),
1115            }
1116        } else {
1117            Ok(())
1118        };
1119
1120        self.event_schema = None;
1121        self.state = SinkState::Idle;
1122        result
1123    }
1124}
1125
1126impl Drop for RootSink {
1127    fn drop(&mut self) {
1128        let _ = self.abort();
1129    }
1130}
1131
1132enum RootColumnSender {
1133    F64(Sender<f64>),
1134    F32(Sender<f32>),
1135    I8(Sender<i8>),
1136    U8(Sender<u8>),
1137    I16(Sender<i16>),
1138    U16(Sender<u16>),
1139    I32(Sender<i32>),
1140    U32(Sender<u32>),
1141    I64(Sender<i64>),
1142    U64(Sender<u64>),
1143}
1144
1145impl RootColumnSender {
1146    fn send_float(&self, value: f64) -> LadduDataResult<()> {
1147        match self {
1148            Self::F64(tx) => tx
1149                .send(value)
1150                .map_err(|e| LadduDataError::Sink(e.to_string())),
1151            Self::F32(tx) => tx
1152                .send(value as f32)
1153                .map_err(|e| LadduDataError::Sink(e.to_string())),
1154            _ => Err(LadduDataError::Sink("float sent to integer column".into())),
1155        }
1156    }
1157    fn send_integer(&self, value: ColumnValue) -> LadduDataResult<()> {
1158        match (self, value) {
1159            (Self::I8(tx), ColumnValue::I8(value)) => tx
1160                .send(value)
1161                .map_err(|e| LadduDataError::Sink(e.to_string())),
1162            (Self::U8(tx), ColumnValue::U8(value)) => tx
1163                .send(value)
1164                .map_err(|e| LadduDataError::Sink(e.to_string())),
1165            (Self::I16(tx), ColumnValue::I16(value)) => tx
1166                .send(value)
1167                .map_err(|e| LadduDataError::Sink(e.to_string())),
1168            (Self::U16(tx), ColumnValue::U16(value)) => tx
1169                .send(value)
1170                .map_err(|e| LadduDataError::Sink(e.to_string())),
1171            (Self::I32(tx), ColumnValue::I32(value)) => tx
1172                .send(value)
1173                .map_err(|e| LadduDataError::Sink(e.to_string())),
1174            (Self::U32(tx), ColumnValue::U32(value)) => tx
1175                .send(value)
1176                .map_err(|e| LadduDataError::Sink(e.to_string())),
1177            (Self::I64(tx), ColumnValue::I64(value)) => tx
1178                .send(value)
1179                .map_err(|e| LadduDataError::Sink(e.to_string())),
1180            (Self::U64(tx), ColumnValue::U64(value)) => tx
1181                .send(value)
1182                .map_err(|e| LadduDataError::Sink(e.to_string())),
1183            _ => Err(LadduDataError::Sink("integer column dtype mismatch".into())),
1184        }
1185    }
1186}
1187
1188enum RootColumnReceiver {
1189    F64(Name, Receiver<f64>),
1190    F32(Name, Receiver<f32>),
1191    I8(Name, Receiver<i8>),
1192    U8(Name, Receiver<u8>),
1193    I16(Name, Receiver<i16>),
1194    U16(Name, Receiver<u16>),
1195    I32(Name, Receiver<i32>),
1196    U32(Name, Receiver<u32>),
1197    I64(Name, Receiver<i64>),
1198    U64(Name, Receiver<u64>),
1199}
1200
1201fn root_channels(
1202    columns: &[Name],
1203    schema: &Schema,
1204    precision: Precision,
1205) -> (Vec<RootColumnSender>, Vec<RootColumnReceiver>) {
1206    let mut senders = Vec::with_capacity(columns.len());
1207    let mut receivers = Vec::with_capacity(columns.len());
1208    for name in columns {
1209        match schema.column_index(name).map(|i| schema.columns()[i].1) {
1210            Some(ColumnDType::I8) => {
1211                let (tx, rx) = mpsc::channel();
1212                senders.push(RootColumnSender::I8(tx));
1213                receivers.push(RootColumnReceiver::I8(name.clone(), rx));
1214            }
1215            Some(ColumnDType::U8) => {
1216                let (tx, rx) = mpsc::channel();
1217                senders.push(RootColumnSender::U8(tx));
1218                receivers.push(RootColumnReceiver::U8(name.clone(), rx));
1219            }
1220            Some(ColumnDType::I16) => {
1221                let (tx, rx) = mpsc::channel();
1222                senders.push(RootColumnSender::I16(tx));
1223                receivers.push(RootColumnReceiver::I16(name.clone(), rx));
1224            }
1225            Some(ColumnDType::U16) => {
1226                let (tx, rx) = mpsc::channel();
1227                senders.push(RootColumnSender::U16(tx));
1228                receivers.push(RootColumnReceiver::U16(name.clone(), rx));
1229            }
1230            Some(ColumnDType::I32) => {
1231                let (tx, rx) = mpsc::channel();
1232                senders.push(RootColumnSender::I32(tx));
1233                receivers.push(RootColumnReceiver::I32(name.clone(), rx));
1234            }
1235            Some(ColumnDType::U32) => {
1236                let (tx, rx) = mpsc::channel();
1237                senders.push(RootColumnSender::U32(tx));
1238                receivers.push(RootColumnReceiver::U32(name.clone(), rx));
1239            }
1240            Some(ColumnDType::I64) => {
1241                let (tx, rx) = mpsc::channel();
1242                senders.push(RootColumnSender::I64(tx));
1243                receivers.push(RootColumnReceiver::I64(name.clone(), rx));
1244            }
1245            Some(ColumnDType::U64) => {
1246                let (tx, rx) = mpsc::channel();
1247                senders.push(RootColumnSender::U64(tx));
1248                receivers.push(RootColumnReceiver::U64(name.clone(), rx));
1249            }
1250            None => match precision {
1251                Precision::F64 => {
1252                    let (tx, rx) = mpsc::channel();
1253                    senders.push(RootColumnSender::F64(tx));
1254                    receivers.push(RootColumnReceiver::F64(name.clone(), rx));
1255                }
1256                Precision::F32 => {
1257                    let (tx, rx) = mpsc::channel();
1258                    senders.push(RootColumnSender::F32(tx));
1259                    receivers.push(RootColumnReceiver::F32(name.clone(), rx));
1260                }
1261            },
1262        }
1263    }
1264    (senders, receivers)
1265}
1266
1267fn write_root_tree(
1268    path: PathBuf,
1269    tree_name: Name,
1270    receivers: Vec<RootColumnReceiver>,
1271) -> LadduDataResult<()> {
1272    let resource = format!("{}::{tree_name}", path.display());
1273    let mut file = RootFile::create(&path)
1274        .map_err(|error| sink_error("create ROOT file", path.display(), error))?;
1275    let mut tree = WriterTree::new(tree_name.as_ref());
1276
1277    for receiver in receivers {
1278        match receiver {
1279            RootColumnReceiver::F64(name, rx) => tree.new_branch(name.as_ref(), rx.into_iter()),
1280            RootColumnReceiver::F32(name, rx) => tree.new_branch(name.as_ref(), rx.into_iter()),
1281            RootColumnReceiver::I8(name, rx) => tree.new_branch(name.as_ref(), rx.into_iter()),
1282            RootColumnReceiver::U8(name, rx) => tree.new_branch(name.as_ref(), rx.into_iter()),
1283            RootColumnReceiver::I16(name, rx) => tree.new_branch(name.as_ref(), rx.into_iter()),
1284            RootColumnReceiver::U16(name, rx) => tree.new_branch(name.as_ref(), rx.into_iter()),
1285            RootColumnReceiver::I32(name, rx) => tree.new_branch(name.as_ref(), rx.into_iter()),
1286            RootColumnReceiver::U32(name, rx) => tree.new_branch(name.as_ref(), rx.into_iter()),
1287            RootColumnReceiver::I64(name, rx) => tree.new_branch(name.as_ref(), rx.into_iter()),
1288            RootColumnReceiver::U64(name, rx) => tree.new_branch(name.as_ref(), rx.into_iter()),
1289        }
1290    }
1291
1292    tree.write(&mut file)
1293        .map_err(|error| sink_error("write ROOT tree", &resource, error))?;
1294    file.close()
1295        .map_err(|error| sink_error("close ROOT file", &resource, error))?;
1296
1297    Ok(())
1298}
1299
1300struct RootIntegerIter<'a> {
1301    name: Name,
1302    iter: Box<dyn Iterator<Item = ColumnValue> + 'a>,
1303}
1304impl RootIntegerIter<'_> {
1305    fn next_value(&mut self) -> LadduDataResult<ColumnValue> {
1306        self.iter.next().ok_or_else(|| {
1307            source_error(
1308                "read ROOT integer branch",
1309                &self.name,
1310                "ROOT branch ended early",
1311            )
1312        })
1313    }
1314}
1315fn open_integer_reader<'a>(
1316    tree: &'a ReaderTree,
1317    name: &str,
1318    dtype: ColumnDType,
1319) -> LadduDataResult<RootIntegerIter<'a>> {
1320    let branch =
1321        find_branch(tree, name).ok_or_else(|| LadduDataError::MissingColumn(Name::from(name)))?;
1322    if root_column_type(branch) != ColumnType::Integer(dtype) {
1323        return Err(source_error(
1324            "open ROOT integer branch",
1325            name,
1326            "dtype mismatch",
1327        ));
1328    }
1329    let iter: Box<dyn Iterator<Item = ColumnValue> + 'a> = match dtype {
1330        ColumnDType::I8 => Box::new(
1331            branch
1332                .as_iter::<i8>()
1333                .map_err(|e| source_error("open ROOT integer branch", name, e))?
1334                .map(ColumnValue::I8),
1335        ),
1336        ColumnDType::U8 => Box::new(
1337            branch
1338                .as_iter::<u8>()
1339                .map_err(|e| source_error("open ROOT integer branch", name, e))?
1340                .map(ColumnValue::U8),
1341        ),
1342        ColumnDType::I16 => Box::new(
1343            branch
1344                .as_iter::<i16>()
1345                .map_err(|e| source_error("open ROOT integer branch", name, e))?
1346                .map(ColumnValue::I16),
1347        ),
1348        ColumnDType::U16 => Box::new(
1349            branch
1350                .as_iter::<u16>()
1351                .map_err(|e| source_error("open ROOT integer branch", name, e))?
1352                .map(ColumnValue::U16),
1353        ),
1354        ColumnDType::I32 => Box::new(
1355            branch
1356                .as_iter::<i32>()
1357                .map_err(|e| source_error("open ROOT integer branch", name, e))?
1358                .map(ColumnValue::I32),
1359        ),
1360        ColumnDType::U32 => Box::new(
1361            branch
1362                .as_iter::<u32>()
1363                .map_err(|e| source_error("open ROOT integer branch", name, e))?
1364                .map(ColumnValue::U32),
1365        ),
1366        ColumnDType::I64 => Box::new(
1367            branch
1368                .as_iter::<i64>()
1369                .map_err(|e| source_error("open ROOT integer branch", name, e))?
1370                .map(ColumnValue::I64),
1371        ),
1372        ColumnDType::U64 => Box::new(
1373            branch
1374                .as_iter::<u64>()
1375                .map_err(|e| source_error("open ROOT integer branch", name, e))?
1376                .map(ColumnValue::U64),
1377        ),
1378    };
1379    Ok(RootIntegerIter {
1380        name: Name::from(name),
1381        iter,
1382    })
1383}
1384
1385#[cfg(test)]
1386mod tests {
1387    use std::sync::atomic::{AtomicU64, Ordering};
1388
1389    use super::*;
1390    use crate::data::{Dataset, EventBatchBuilder};
1391
1392    fn temp_path(ext: &str) -> PathBuf {
1393        static NEXT_TEMP_FILE_ID: AtomicU64 = AtomicU64::new(0);
1394
1395        let nanos = std::time::SystemTime::now()
1396            .duration_since(std::time::UNIX_EPOCH)
1397            .unwrap()
1398            .as_nanos();
1399        let id = NEXT_TEMP_FILE_ID.fetch_add(1, Ordering::Relaxed);
1400
1401        std::env::temp_dir().join(format!(
1402            "laddu-root-test-{}-{nanos}-{id}.{ext}",
1403            std::process::id()
1404        ))
1405    }
1406
1407    fn v(x: f64) -> RealVec4 {
1408        RealVec4 {
1409            e: x + 0.3,
1410            px: x,
1411            py: x + 0.1,
1412            pz: x + 0.2,
1413        }
1414    }
1415
1416    fn schema() -> Arc<Schema> {
1417        Arc::new(Schema::new(["p"], ["mass"], true).unwrap())
1418    }
1419
1420    fn batch() -> EventBatch {
1421        let schema = schema();
1422        let mut builder = EventBatchBuilder::new(schema);
1423
1424        for i in 0..4 {
1425            builder
1426                .push_weighted([v(i as f64)], [100.0 + i as f64], 10.0 + i as f64)
1427                .unwrap();
1428        }
1429
1430        builder.finish().unwrap()
1431    }
1432
1433    #[test]
1434    fn root_sink_and_source_roundtrip_named_tree_with_f32_precision() {
1435        let path = temp_path("root");
1436        let batch = batch();
1437
1438        let mut sink = RootSink::builder(path.clone())
1439            .tree("events")
1440            .precision(Precision::F32)
1441            .build();
1442
1443        sink.begin(Arc::clone(batch.schema()), WritePlan::default())
1444            .unwrap();
1445        sink.write_batch(&batch).unwrap();
1446        sink.finish().unwrap();
1447
1448        let tree_names = RootSource::tree_names(&path).unwrap();
1449        assert!(tree_names.iter().any(|name| name.as_ref() == "events"));
1450
1451        let columns = RootSource::columns(&path, Some("events")).unwrap();
1452        let names = columns
1453            .iter()
1454            .map(|col| col.name.to_string())
1455            .collect::<Vec<_>>();
1456
1457        for expected in ["p_e", "p_px", "p_py", "p_pz", "mass", "weight"] {
1458            assert!(
1459                names.iter().any(|name| name == expected),
1460                "missing {expected}"
1461            );
1462        }
1463
1464        let source = RootSource::builder(path.to_str().unwrap())
1465            .tree("events")
1466            .build()
1467            .unwrap();
1468
1469        assert_eq!(source.tree_name(), "events");
1470        assert_eq!(source.num_events().unwrap(), Some(4));
1471
1472        let read_batches: Vec<EventBatch> = source
1473            .batches(ReadPlan {
1474                chunk_size: Some(2),
1475                #[cfg(feature = "mpi")]
1476                distribution: Default::default(),
1477            })
1478            .unwrap()
1479            .map(Result::unwrap)
1480            .collect();
1481
1482        assert_eq!(
1483            read_batches.iter().map(EventBatch::len).collect::<Vec<_>>(),
1484            vec![2, 2]
1485        );
1486
1487        let read = EventBatch::concat(&read_batches).unwrap();
1488
1489        assert_eq!(read.scalar_column(0), &[100.0, 101.0, 102.0, 103.0]);
1490        assert_eq!(read.weights_column().unwrap(), &[10.0, 11.0, 12.0, 13.0]);
1491        assert!((read.p4_at(0, 2).e - 2.3).abs() < 1.0e-6);
1492
1493        let _ = std::fs::remove_file(path);
1494    }
1495
1496    #[test]
1497    fn root_source_infers_first_tree_when_no_tree_is_named() {
1498        let path = temp_path("root");
1499        let batch = batch();
1500
1501        let mut sink = RootSink::builder(path.clone()).tree("first_tree").build();
1502
1503        sink.begin(Arc::clone(batch.schema()), WritePlan::default())
1504            .unwrap();
1505        sink.write_batch(&batch).unwrap();
1506        sink.finish().unwrap();
1507
1508        let source = RootSource::builder(path.to_str().unwrap())
1509            .first_tree()
1510            .build()
1511            .unwrap();
1512
1513        assert_eq!(source.tree_name(), "first_tree");
1514
1515        let read = EventBatch::concat(
1516            &source
1517                .batches(ReadPlan::default())
1518                .unwrap()
1519                .map(Result::unwrap)
1520                .collect::<Vec<_>>(),
1521        )
1522        .unwrap();
1523
1524        assert_eq!(read.scalar_column(0), &[100.0, 101.0, 102.0, 103.0]);
1525
1526        let _ = std::fs::remove_file(path);
1527    }
1528
1529    #[test]
1530    fn negative_root_entry_count_error_includes_operation_and_resource() {
1531        let error =
1532            usize_from_i64(-1, "negative TTree entry count", "events.root::events").unwrap_err();
1533        let message = error.to_string();
1534        assert!(message.contains("read ROOT entry count `events.root::events`"));
1535        assert!(message.contains("negative TTree entry count"));
1536    }
1537
1538    #[test]
1539    fn root_source_named_missing_tree_fails() {
1540        let path = temp_path("root");
1541        let batch = batch();
1542
1543        let mut sink = RootSink::builder(path.clone()).tree("events").build();
1544
1545        sink.begin(Arc::clone(batch.schema()), WritePlan::default())
1546            .unwrap();
1547        sink.write_batch(&batch).unwrap();
1548        sink.finish().unwrap();
1549
1550        let err = RootSource::builder(path.to_str().unwrap())
1551            .tree("missing")
1552            .build()
1553            .unwrap_err();
1554
1555        assert!(matches!(err, LadduDataError::Source(_)));
1556
1557        let _ = std::fs::remove_file(path);
1558    }
1559
1560    #[test]
1561    fn root_sink_rejects_batches_with_different_schema() {
1562        let path = temp_path("root");
1563        let batch = batch();
1564
1565        let mut sink = RootSink::builder(path.clone()).tree("events").build();
1566
1567        sink.begin(Arc::clone(batch.schema()), WritePlan::default())
1568            .unwrap();
1569
1570        let other_schema = Arc::new(Schema::new(["q"], ["mass"], true).unwrap());
1571        let mut builder = EventBatchBuilder::new(other_schema);
1572        builder.push_weighted([v(1.0)], [1.0], 1.0).unwrap();
1573        let other = builder.finish().unwrap();
1574
1575        let err = sink.write_batch(&other).unwrap_err();
1576        assert!(matches!(err, LadduDataError::Sink(msg) if msg.contains("schema")));
1577
1578        sink.finish().unwrap();
1579        let _ = std::fs::remove_file(path);
1580    }
1581
1582    #[test]
1583    fn dataset_write_to_root_applies_dataset_transformations_before_writing() {
1584        let path = temp_path("root");
1585
1586        let dataset = Dataset::from_batch(batch()).filter(|ev| ev.scalar(0) >= 102.0);
1587
1588        let mut sink = RootSink::builder(path.clone()).tree("events").build();
1589
1590        dataset.write_to(&mut sink).unwrap();
1591
1592        let source = RootSource::builder(path.to_str().unwrap())
1593            .tree("events")
1594            .build()
1595            .unwrap();
1596
1597        let read = EventBatch::concat(
1598            &source
1599                .batches(ReadPlan::default())
1600                .unwrap()
1601                .map(Result::unwrap)
1602                .collect::<Vec<_>>(),
1603        )
1604        .unwrap();
1605
1606        assert_eq!(read.scalar_column(0), &[102.0, 103.0]);
1607        assert_eq!(read.weights_column().unwrap(), &[12.0, 13.0]);
1608
1609        let _ = std::fs::remove_file(path);
1610    }
1611
1612    #[test]
1613    fn root_batch_iter_disconnect_is_terminal_after_joining_reader() {
1614        let (tx, rx) = mpsc::sync_channel(1);
1615        drop(tx);
1616        let handle = thread::spawn(|| {});
1617        let mut iter = RootBatchIter {
1618            rx,
1619            state: RootBatchState::Receiving(handle),
1620        };
1621
1622        assert!(iter.next().is_none());
1623        assert!(iter.next().is_none());
1624    }
1625}