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#[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#[derive(Clone, Debug)]
43pub struct RootReadOptions {
44 pub infer_schema: bool,
46 pub validate_all_files: bool,
48 pub sort_glob: bool,
50 pub tree: RootTreeSelection,
52 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#[derive(Clone, Debug, Default)]
70pub enum RootTreeSelection {
71 #[default]
73 First,
74 Named(Name),
76}
77
78#[derive(Clone, Debug)]
80pub struct RootFragmentKey {
81 pub file: Arc<PathBuf>,
83 pub tree_name: Name,
85}
86
87#[derive(Clone, Debug)]
89pub struct RootColumnInfo {
90 pub name: Name,
92 pub item_type_name: String,
94 pub interpretation: String,
96 pub entries: i64,
98}
99
100impl RootSource {
101 pub fn open(pattern: impl AsRef<str>) -> LadduDataResult<Self> {
108 Self::builder(pattern).build()
109 }
110
111 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 pub fn files(&self) -> &[Arc<PathBuf>] {
122 &self.files
123 }
124
125 pub fn tree_name(&self) -> &str {
127 self.tree_name.as_ref()
128 }
129
130 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 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
192pub struct RootSourceBuilder {
194 pattern: String,
195 schema: Option<Arc<Schema>>,
196 options: RootReadOptions,
197}
198
199impl RootSourceBuilder {
200 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 pub fn infer_schema(mut self, value: bool) -> Self {
209 self.options.infer_schema = value;
210 self
211 }
212
213 pub fn tree(mut self, name: impl Into<Name>) -> Self {
215 self.options.tree = RootTreeSelection::Named(name.into());
216 self
217 }
218
219 pub fn first_tree(mut self) -> Self {
221 self.options.tree = RootTreeSelection::First;
222 self
223 }
224
225 pub fn require_weight(mut self, value: bool) -> Self {
227 self.options.schema_inference.require_weight = value;
228 self
229 }
230
231 pub fn validate_all_files(mut self, value: bool) -> Self {
233 self.options.validate_all_files = value;
234 self
235 }
236
237 pub fn sort_glob(mut self, value: bool) -> Self {
239 self.options.sort_glob = value;
240 self
241 }
242
243 pub fn schema_inference(mut self, options: SchemaInferenceOptions) -> Self {
245 self.options.schema_inference = options;
246 self
247 }
248
249 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 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
852pub 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#[derive(Clone, Debug)]
865pub struct RootWriteOptions {
866 pub tree_name: Name,
868 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 pub fn create(path: impl Into<PathBuf>) -> Self {
884 Self::builder(path).build()
885 }
886
887 pub fn builder(path: impl Into<PathBuf>) -> RootSinkBuilder {
889 RootSinkBuilder {
890 output: OutputPath::new(path),
891 options: RootWriteOptions::default(),
892 }
893 }
894
895 pub fn resolved_path(&self) -> Option<&Path> {
897 self.resolved_path.as_deref()
898 }
899}
900
901pub struct RootSinkBuilder {
903 output: OutputPath,
904 options: RootWriteOptions,
905}
906
907impl RootSinkBuilder {
908 pub fn output_mode(mut self, mode: OutputMode) -> Self {
910 self.output = self.output.with_mode(mode);
911 self
912 }
913
914 pub fn single_file(self) -> Self {
916 self.output_mode(OutputMode::SingleFile)
917 }
918
919 pub fn per_rank_files(self) -> Self {
921 self.output_mode(OutputMode::PerRankFiles)
922 }
923
924 pub fn auto_output(self) -> Self {
926 self.output_mode(OutputMode::Auto)
927 }
928
929 pub fn tree(mut self, name: impl Into<Name>) -> Self {
931 self.options.tree_name = name.into();
932 self
933 }
934
935 pub fn schema_write(mut self, options: SchemaWriteOptions) -> Self {
937 self.options.schema_write = options;
938 self
939 }
940
941 pub fn column_names(mut self, column_names: SchemaColumnNames) -> Self {
943 self.options.schema_write.column_names = column_names;
944 self
945 }
946
947 pub fn precision(mut self, precision: Precision) -> Self {
949 self.options.schema_write.precision = precision;
950 self
951 }
952
953 pub fn write_weight_column(mut self, value: WriteWeightColumn) -> Self {
955 self.options.schema_write.write_weight_column = value;
956 self
957 }
958
959 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 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}