1use std::{
2 collections::{BTreeMap, HashMap},
3 sync::Arc,
4};
5
6use crate::{LadduDataError, LadduDataResult, Name, columns::ColumnDType};
7
8#[derive(Clone, Debug)]
10pub struct Schema {
11 p4s: Vec<Name>,
12 scalars: Vec<Name>,
13 columns: Vec<(Name, ColumnDType)>,
14 has_weight: bool,
15
16 p4_index: Arc<HashMap<Name, usize>>,
17 scalar_index: Arc<HashMap<Name, usize>>,
18 column_index: Arc<HashMap<Name, usize>>,
19}
20
21#[derive(Clone, Debug, PartialEq)]
28pub struct ScalarBinding {
29 schema: Arc<Schema>,
30 index: usize,
31}
32
33impl ScalarBinding {
34 pub fn index(&self) -> usize {
36 self.index
37 }
38
39 pub(crate) fn matches(&self, schema: &Schema) -> bool {
40 self.schema.as_ref() == schema
41 }
42}
43
44#[derive(Clone, Debug, PartialEq)]
51pub struct P4Binding {
52 schema: Arc<Schema>,
53 index: usize,
54}
55
56impl P4Binding {
57 pub fn index(&self) -> usize {
59 self.index
60 }
61
62 pub(crate) fn matches(&self, schema: &Schema) -> bool {
63 self.schema.as_ref() == schema
64 }
65}
66
67impl PartialEq for Schema {
68 fn eq(&self, other: &Self) -> bool {
69 self.p4s == other.p4s
70 && self.scalars == other.scalars
71 && self.columns == other.columns
72 && self.has_weight == other.has_weight
73 }
74}
75
76impl Schema {
77 pub fn new(
84 p4s: impl IntoIterator<Item = impl Into<Name>>,
85 scalars: impl IntoIterator<Item = impl Into<Name>>,
86 has_weight: bool,
87 ) -> LadduDataResult<Self> {
88 let p4s: Vec<Name> = p4s.into_iter().map(Into::into).collect();
89 let scalars: Vec<Name> = scalars.into_iter().map(Into::into).collect();
90 let p4_index = Arc::new(make_index(&p4s, "p4")?);
91 let scalar_index = Arc::new(make_index(&scalars, "scalar")?);
92 Ok(Self {
93 p4s,
94 scalars,
95 columns: Vec::new(),
96 has_weight,
97 p4_index,
98 scalar_index,
99 column_index: Arc::new(HashMap::new()),
100 })
101 }
102
103 pub fn with_columns(
108 mut self,
109 columns: impl IntoIterator<Item = (impl Into<Name>, ColumnDType)>,
110 ) -> LadduDataResult<Self> {
111 self.columns = columns
112 .into_iter()
113 .map(|(name, dtype)| (name.into(), dtype))
114 .collect();
115 let names: Vec<Name> = self
116 .columns
117 .iter()
118 .map(|(name, _)| Arc::clone(name))
119 .collect();
120 for name in &names {
121 if name.is_empty() || self.p4_index(name).is_some() || self.scalar_index(name).is_some()
122 {
123 return Err(LadduDataError::Schema(format!(
124 "invalid or conflicting column name: {name}"
125 )));
126 }
127 }
128 self.column_index = Arc::new(make_index(&names, "typed")?);
129 Ok(self)
130 }
131
132 pub fn columns(&self) -> &[(Name, ColumnDType)] {
134 &self.columns
135 }
136
137 pub fn column_index(&self, name: &str) -> Option<usize> {
139 self.column_index.get(name).copied()
140 }
141
142 pub fn n_columns(&self) -> usize {
144 self.columns.len()
145 }
146
147 pub fn validate_column_names(&self, names: &SchemaColumnNames) -> LadduDataResult<()> {
152 for (name, _) in &self.columns {
153 if name == &names.weight_column
154 || self.p4s.iter().any(|p4| {
155 names
156 .p4_suffixes
157 .physical_p4_names(p4)
158 .iter()
159 .any(|physical| physical == name.as_ref())
160 })
161 {
162 return Err(LadduDataError::Schema(format!(
163 "conflicting physical column name: {name}"
164 )));
165 }
166 }
167 Ok(())
168 }
169
170 pub fn p4_index(&self, name: &str) -> Option<usize> {
172 self.p4_index.get(name).copied()
173 }
174
175 pub fn bind_p4(&self, name: &str) -> Option<P4Binding> {
177 self.p4_index(name).map(|index| P4Binding {
178 schema: Arc::new(self.clone()),
179 index,
180 })
181 }
182
183 pub fn scalar_index(&self, name: &str) -> Option<usize> {
185 self.scalar_index.get(name).copied()
186 }
187
188 pub fn bind_scalar(&self, name: &str) -> Option<ScalarBinding> {
190 self.scalar_index(name).map(|index| ScalarBinding {
191 schema: Arc::new(self.clone()),
192 index,
193 })
194 }
195
196 pub fn p4s(&self) -> &[Name] {
198 &self.p4s
199 }
200
201 pub fn scalars(&self) -> &[Name] {
203 &self.scalars
204 }
205
206 pub fn has_weight(&self) -> bool {
208 self.has_weight
209 }
210
211 pub fn n_p4s(&self) -> usize {
213 self.p4s.len()
214 }
215
216 pub fn n_scalars(&self) -> usize {
218 self.scalars.len()
219 }
220
221 pub fn require_p4(&self, name: &str) -> LadduDataResult<usize> {
228 self.p4_index(name)
229 .ok_or_else(|| LadduDataError::MissingColumn(Name::from(name)))
230 }
231
232 pub fn require_scalar(&self, name: &str) -> LadduDataResult<usize> {
239 self.scalar_index(name)
240 .ok_or_else(|| LadduDataError::MissingColumn(Name::from(name)))
241 }
242}
243
244fn make_index(names: &[Name], kind: &'static str) -> LadduDataResult<HashMap<Name, usize>> {
245 let mut out = HashMap::with_capacity(names.len());
246 for (i, name) in names.iter().cloned().enumerate() {
247 if out.insert(name.clone(), i).is_some() {
248 return Err(LadduDataError::Schema(format!(
249 "duplicate {kind} column: {name}"
250 )));
251 }
252 }
253 Ok(out)
254}
255
256#[derive(Clone, Debug)]
258pub struct SchemaColumnNames {
259 pub weight_column: Name,
261 pub p4_suffixes: P4Suffixes,
263}
264
265impl Default for SchemaColumnNames {
266 fn default() -> Self {
267 Self {
268 weight_column: Name::from("weight"),
269 p4_suffixes: P4Suffixes::default(),
270 }
271 }
272}
273
274#[derive(Clone, Debug)]
276pub struct SchemaInferenceOptions {
277 pub column_names: SchemaColumnNames,
279 pub require_weight: bool,
281 pub incomplete_p4_components_are_scalars: bool,
283}
284
285impl Default for SchemaInferenceOptions {
286 fn default() -> Self {
287 Self {
288 column_names: SchemaColumnNames::default(),
289 require_weight: false,
290 incomplete_p4_components_are_scalars: true,
291 }
292 }
293}
294
295#[derive(Clone, Debug)]
297pub struct P4Suffixes {
298 pub e: &'static str,
300 pub px: &'static str,
302 pub py: &'static str,
304 pub pz: &'static str,
306}
307
308impl Default for P4Suffixes {
309 fn default() -> Self {
310 Self {
311 e: "_e",
312 px: "_px",
313 py: "_py",
314 pz: "_pz",
315 }
316 }
317}
318
319impl P4Suffixes {
320 pub fn component<'a>(&'a self, name: &'a str) -> Option<(&'a str, usize)> {
322 if let Some(prefix) = name.strip_suffix(self.e) {
323 Some((prefix, 0))
324 } else if let Some(prefix) = name.strip_suffix(self.px) {
325 Some((prefix, 1))
326 } else if let Some(prefix) = name.strip_suffix(self.py) {
327 Some((prefix, 2))
328 } else if let Some(prefix) = name.strip_suffix(self.pz) {
329 Some((prefix, 3))
330 } else {
331 None
332 }
333 }
334
335 pub fn physical_p4_names(&self, prefix: &str) -> [String; 4] {
337 [
338 format!("{prefix}{}", self.e),
339 format!("{prefix}{}", self.px),
340 format!("{prefix}{}", self.py),
341 format!("{prefix}{}", self.pz),
342 ]
343 }
344}
345
346#[derive(Clone, Copy, Debug, PartialEq, Eq)]
348pub enum ColumnType {
349 Integer(ColumnDType),
351 F64,
353 F32,
355 Other,
357}
358
359impl ColumnType {
360 pub fn is_supported_float(self) -> bool {
362 matches!(self, Self::F64 | Self::F32)
363 }
364}
365
366#[derive(Clone, Copy, Debug)]
368pub struct ColumnInfo<'a> {
369 pub name: &'a str,
371 pub dtype: ColumnType,
373}
374
375impl Schema {
376 pub fn infer_from_columns<'a>(
383 columns: impl IntoIterator<Item = ColumnInfo<'a>>,
384 options: &SchemaInferenceOptions,
385 ) -> LadduDataResult<Self> {
386 let mut p4_candidates = BTreeMap::<String, [bool; 4]>::new();
387 let mut scalar_names = Vec::<Name>::new();
388 let mut has_weight = false;
389 let mut typed_columns = Vec::new();
390
391 for col in columns {
392 if let ColumnType::Integer(dtype) = col.dtype {
393 typed_columns.push((Name::from(col.name), dtype));
394 continue;
395 }
396 if !col.dtype.is_supported_float() {
397 continue;
398 }
399
400 if col.name == options.column_names.weight_column.as_ref() {
401 has_weight = true;
402 continue;
403 }
404
405 if let Some((prefix, component)) = options.column_names.p4_suffixes.component(col.name)
406 {
407 p4_candidates.entry(prefix.to_owned()).or_default()[component] = true;
408 } else {
409 scalar_names.push(Name::from(col.name));
410 }
411 }
412
413 let mut p4s = Vec::<Name>::new();
414
415 for (prefix, seen) in p4_candidates {
416 if seen == [true, true, true, true] {
417 p4s.push(Name::from(prefix));
418 } else if options.incomplete_p4_components_are_scalars {
419 let names = options.column_names.p4_suffixes.physical_p4_names(&prefix);
420
421 for (i, name) in names.into_iter().enumerate() {
422 if seen[i] {
423 scalar_names.push(Name::from(name));
424 }
425 }
426 }
427 }
428
429 if options.require_weight && !has_weight {
430 return Err(LadduDataError::MissingColumn(Arc::clone(
431 &options.column_names.weight_column,
432 )));
433 }
434
435 let schema = Schema::new(p4s, scalar_names, has_weight)?.with_columns(typed_columns)?;
436 schema.validate_column_names(&options.column_names)?;
437 Ok(schema)
438 }
439
440 pub fn physical_columns(&self, column_names: &SchemaColumnNames) -> Vec<Name> {
442 PhysicalSchemaPlan::for_read(self, column_names)
443 .columns()
444 .iter()
445 .map(|column| Arc::clone(column.name()))
446 .collect()
447 }
448
449 pub fn validate_required_columns<'a>(
456 &self,
457 available: impl IntoIterator<Item = ColumnInfo<'a>>,
458 options: &SchemaInferenceOptions,
459 ) -> LadduDataResult<()> {
460 self.validate_column_names(&options.column_names)?;
461 let available: HashMap<&str, ColumnType> =
462 available.into_iter().map(|c| (c.name, c.dtype)).collect();
463
464 for required in PhysicalSchemaPlan::for_read(self, &options.column_names).columns() {
465 let dtype = available.get(required.name().as_ref());
466 let matches = match required.role() {
467 PhysicalColumnRole::Column {
468 dtype: expected, ..
469 } => dtype == Some(&ColumnType::Integer(expected)),
470 _ => dtype.is_some_and(|dtype| dtype.is_supported_float()),
471 };
472 if !matches {
473 if let PhysicalColumnRole::Column {
474 dtype: expected, ..
475 } = required.role()
476 && let Some(actual) = dtype
477 {
478 return Err(LadduDataError::Schema(format!(
479 "column {} has dtype {actual:?}, expected {expected}",
480 required.name()
481 )));
482 }
483 return Err(LadduDataError::MissingColumn(Arc::clone(required.name())));
484 }
485 }
486
487 Ok(())
488 }
489}
490
491#[derive(Copy, Clone, Debug, Default)]
493pub enum Precision {
494 #[default]
496 F64,
497 F32,
499}
500
501#[derive(Clone, Copy, Debug, Default)]
503pub enum WriteWeightColumn {
504 #[default]
506 Always,
507 OnlyIfPresent,
509}
510
511#[derive(Clone, Debug, Default)]
513pub struct SchemaWriteOptions {
514 pub column_names: SchemaColumnNames,
516 pub precision: Precision,
518 pub write_weight_column: WriteWeightColumn,
520}
521
522#[derive(Clone, Copy, Debug, PartialEq, Eq)]
528pub(crate) enum PhysicalColumnRole {
529 Column { index: usize, dtype: ColumnDType },
531 P4 {
533 index: usize,
535 component: usize,
537 },
538 Scalar {
540 index: usize,
542 },
543 Weight,
545}
546
547#[derive(Clone, Debug)]
553pub(crate) struct PhysicalSchemaPlan {
554 columns: Vec<PhysicalColumn>,
555}
556
557#[derive(Clone, Debug)]
558pub(crate) struct PhysicalColumn {
559 name: Name,
560 role: PhysicalColumnRole,
561}
562
563impl PhysicalSchemaPlan {
564 pub(crate) fn for_read(schema: &Schema, column_names: &SchemaColumnNames) -> Self {
566 Self::build(schema, column_names, schema.has_weight())
567 }
568
569 pub(crate) fn for_write(
571 schema: &Schema,
572 options: &SchemaWriteOptions,
573 write_weight: WriteWeightColumn,
574 ) -> Self {
575 let should_write_weight =
576 matches!(write_weight, WriteWeightColumn::Always) || schema.has_weight();
577 Self::build(schema, &options.column_names, should_write_weight)
578 }
579
580 fn build(schema: &Schema, column_names: &SchemaColumnNames, include_weight: bool) -> Self {
581 let mut columns = Vec::with_capacity(
582 4 * schema.n_p4s() + schema.n_scalars() + usize::from(include_weight),
583 );
584
585 for (index, p4) in schema.p4s().iter().enumerate() {
586 for (component, name) in column_names
587 .p4_suffixes
588 .physical_p4_names(p4)
589 .into_iter()
590 .enumerate()
591 {
592 columns.push(PhysicalColumn {
593 name: Name::from(name),
594 role: PhysicalColumnRole::P4 { index, component },
595 });
596 }
597 }
598
599 for (index, name) in schema.scalars().iter().cloned().enumerate() {
600 columns.push(PhysicalColumn {
601 name,
602 role: PhysicalColumnRole::Scalar { index },
603 });
604 }
605
606 for (index, (name, dtype)) in schema.columns().iter().enumerate() {
607 columns.push(PhysicalColumn {
608 name: Arc::clone(name),
609 role: PhysicalColumnRole::Column {
610 index,
611 dtype: *dtype,
612 },
613 });
614 }
615
616 if include_weight {
617 columns.push(PhysicalColumn {
618 name: Arc::clone(&column_names.weight_column),
619 role: PhysicalColumnRole::Weight,
620 });
621 }
622
623 Self { columns }
624 }
625
626 pub(crate) fn columns(&self) -> &[PhysicalColumn] {
628 &self.columns
629 }
630}
631
632impl PhysicalColumn {
633 pub(crate) fn name(&self) -> &Name {
635 &self.name
636 }
637
638 pub(crate) fn role(&self) -> PhysicalColumnRole {
640 self.role
641 }
642}
643
644#[cfg(test)]
645mod tests {
646 use super::*;
647
648 fn col(name: &'static str, dtype: ColumnType) -> ColumnInfo<'static> {
649 ColumnInfo { name, dtype }
650 }
651
652 #[test]
653 fn schema_new_rejects_duplicates_and_required_lookup_reports_missing_column() {
654 let duplicate_p4 = Schema::new(["p", "p"], ["mass"], false);
655 assert!(matches!(duplicate_p4, Err(LadduDataError::Schema(_))));
656
657 let duplicate_scalar = Schema::new(["p"], ["mass", "mass"], false);
658 assert!(matches!(duplicate_scalar, Err(LadduDataError::Schema(_))));
659
660 let schema = Schema::new(["beam", "recoil"], ["mass", "costheta"], true).unwrap();
661
662 assert_eq!(schema.require_p4("recoil").unwrap(), 1);
663 assert_eq!(schema.require_scalar("costheta").unwrap(), 1);
664
665 let err = schema.require_scalar("missing").unwrap_err();
666 assert!(matches!(err, LadduDataError::MissingColumn(name) if name.as_ref() == "missing"));
667 }
668
669 #[test]
670 fn schema_bindings_resolve_typed_indices_once() {
671 let schema = Schema::new(["beam", "recoil"], ["mass", "costheta"], true).unwrap();
672
673 assert_eq!(schema.bind_p4("recoil").unwrap().index(), 1);
674 assert_eq!(schema.bind_scalar("costheta").unwrap().index(), 1);
675 assert_eq!(schema.bind_p4("missing"), None);
676 assert_eq!(schema.bind_scalar("missing"), None);
677 }
678
679 #[test]
680 fn infer_from_columns_groups_complete_p4s_keeps_incomplete_components_as_scalars_and_ignores_nonfloats()
681 {
682 let options = SchemaInferenceOptions::default();
683
684 let schema = Schema::infer_from_columns(
685 [
686 col("gamma_px", ColumnType::F64),
687 col("gamma_py", ColumnType::F64),
688 col("gamma_pz", ColumnType::F32),
689 col("gamma_e", ColumnType::F64),
690 col("partial_px", ColumnType::F64),
691 col("partial_e", ColumnType::F64),
692 col("mass", ColumnType::F32),
693 col("ignored", ColumnType::Other),
694 col("weight", ColumnType::F64),
695 ],
696 &options,
697 )
698 .unwrap();
699
700 assert_eq!(
701 schema
702 .p4s()
703 .iter()
704 .map(|n| n.to_string())
705 .collect::<Vec<_>>(),
706 vec!["gamma"]
707 );
708
709 assert_eq!(
710 schema
711 .scalars()
712 .iter()
713 .map(|n| n.to_string())
714 .collect::<Vec<_>>(),
715 vec!["mass", "partial_e", "partial_px"]
716 );
717
718 assert!(schema.has_weight());
719 }
720
721 #[test]
722 fn infer_from_columns_can_discard_incomplete_p4_components_and_require_weight() {
723 let options = SchemaInferenceOptions {
724 incomplete_p4_components_are_scalars: false,
725 ..Default::default()
726 };
727
728 let schema = Schema::infer_from_columns(
729 [
730 col("partial_px", ColumnType::F64),
731 col("partial_e", ColumnType::F64),
732 col("mass", ColumnType::F64),
733 col("weight", ColumnType::F64),
734 ],
735 &options,
736 )
737 .unwrap();
738
739 assert!(schema.p4s().is_empty());
740 assert_eq!(
741 schema
742 .scalars()
743 .iter()
744 .map(|n| n.to_string())
745 .collect::<Vec<_>>(),
746 vec!["mass"]
747 );
748
749 let require_weight = SchemaInferenceOptions {
750 require_weight: true,
751 ..Default::default()
752 };
753
754 let err = Schema::infer_from_columns([col("mass", ColumnType::F64)], &require_weight)
755 .unwrap_err();
756
757 assert!(matches!(err, LadduDataError::MissingColumn(name) if name.as_ref() == "weight"));
758 }
759
760 #[test]
761 fn physical_columns_and_validation_respect_custom_names_and_float_types_only() {
762 let schema = Schema::new(["p"], ["mass"], true).unwrap();
763
764 let names = SchemaColumnNames {
765 weight_column: Name::from("event_weight"),
766 ..Default::default()
767 };
768
769 let physical = schema
770 .physical_columns(&names)
771 .into_iter()
772 .map(|n| n.to_string())
773 .collect::<Vec<_>>();
774
775 assert_eq!(
776 physical,
777 vec!["p_e", "p_px", "p_py", "p_pz", "mass", "event_weight"]
778 );
779
780 let options = SchemaInferenceOptions {
781 column_names: names,
782 ..Default::default()
783 };
784
785 let ok = schema.validate_required_columns(
786 [
787 col("p_e", ColumnType::F32),
788 col("p_px", ColumnType::F64),
789 col("p_py", ColumnType::F64),
790 col("p_pz", ColumnType::F32),
791 col("mass", ColumnType::F64),
792 col("event_weight", ColumnType::F64),
793 ],
794 &options,
795 );
796
797 assert!(ok.is_ok());
798
799 let missing_because_not_float = schema
800 .validate_required_columns(
801 [
802 col("p_e", ColumnType::F32),
803 col("p_px", ColumnType::F64),
804 col("p_py", ColumnType::Other),
805 col("p_pz", ColumnType::F32),
806 col("mass", ColumnType::F64),
807 col("event_weight", ColumnType::F64),
808 ],
809 &options,
810 )
811 .unwrap_err();
812
813 assert!(
814 matches!(missing_because_not_float, LadduDataError::MissingColumn(name) if name.as_ref() == "p_py")
815 );
816 }
817
818 #[test]
819 fn physical_schema_plan_preserves_order_roles_and_weight_policy() {
820 let schema = Schema::new(["p"], ["mass"], false).unwrap();
821 let options = SchemaWriteOptions {
822 column_names: SchemaColumnNames {
823 weight_column: Name::from("event_weight"),
824 ..Default::default()
825 },
826 ..Default::default()
827 };
828
829 let only_if_present =
830 PhysicalSchemaPlan::for_write(&schema, &options, WriteWeightColumn::OnlyIfPresent);
831 assert_eq!(
832 only_if_present
833 .columns()
834 .iter()
835 .map(|column| column.name().to_string())
836 .collect::<Vec<_>>(),
837 ["p_e", "p_px", "p_py", "p_pz", "mass"]
838 );
839 assert_eq!(
840 only_if_present
841 .columns()
842 .iter()
843 .map(PhysicalColumn::role)
844 .collect::<Vec<_>>(),
845 [
846 PhysicalColumnRole::P4 {
847 index: 0,
848 component: 0,
849 },
850 PhysicalColumnRole::P4 {
851 index: 0,
852 component: 1,
853 },
854 PhysicalColumnRole::P4 {
855 index: 0,
856 component: 2,
857 },
858 PhysicalColumnRole::P4 {
859 index: 0,
860 component: 3,
861 },
862 PhysicalColumnRole::Scalar { index: 0 },
863 ]
864 );
865
866 let always = PhysicalSchemaPlan::for_write(&schema, &options, WriteWeightColumn::Always);
867 assert_eq!(
868 always.columns().last().map(|column| column.name().as_ref()),
869 Some("event_weight")
870 );
871 assert_eq!(
872 always.columns().last().map(PhysicalColumn::role),
873 Some(PhysicalColumnRole::Weight)
874 );
875 }
876}