Skip to main content

laddu_data/
schema.rs

1use std::{
2    collections::{BTreeMap, HashMap},
3    sync::Arc,
4};
5
6use crate::{LadduDataError, LadduDataResult, Name, columns::ColumnDType};
7
8/// Logical names and lookup tables for four-momentum, scalar, and weight columns.
9#[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/// A schema-resolved scalar column binding.
22///
23/// Resolve bindings once against the schema used by an
24/// [`EventBatch`](crate::data::EventBatch), then reuse them for every row.
25/// Bindings must not be used with a batch whose schema has a different column
26/// ordering.
27#[derive(Clone, Debug, PartialEq)]
28pub struct ScalarBinding {
29    schema: Arc<Schema>,
30    index: usize,
31}
32
33impl ScalarBinding {
34    /// Returns the physical scalar-column index.
35    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/// A schema-resolved four-momentum column binding.
45///
46/// Resolve bindings once against the schema used by an
47/// [`EventBatch`](crate::data::EventBatch), then reuse them for every row.
48/// Bindings must not be used with a batch whose schema has a different column
49/// ordering.
50#[derive(Clone, Debug, PartialEq)]
51pub struct P4Binding {
52    schema: Arc<Schema>,
53    index: usize,
54}
55
56impl P4Binding {
57    /// Returns the physical four-momentum-column index.
58    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    /// Validates and constructs a logical schema.
78    ///
79    /// # Errors
80    ///
81    /// Returns [`LadduDataError`] when four-momentum or scalar column names are
82    /// duplicated.
83    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    /// Adds exact non-expression column declarations in their supplied order.
104    ///
105    /// # Errors
106    /// Returns an error for empty/duplicate names or logical name collisions.
107    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    /// Returns exact column declarations in schema order.
133    pub fn columns(&self) -> &[(Name, ColumnDType)] {
134        &self.columns
135    }
136
137    /// Returns the exact column index for a name.
138    pub fn column_index(&self, name: &str) -> Option<usize> {
139        self.column_index.get(name).copied()
140    }
141
142    /// Returns the number of exact row-data columns.
143    pub fn n_columns(&self) -> usize {
144        self.columns.len()
145    }
146
147    /// Validates typed names against the configured physical fields.
148    ///
149    /// # Errors
150    /// Returns an error if a typed name collides with a physical weight or p4 field.
151    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    /// Returns the four-momentum column index for `name`.
171    pub fn p4_index(&self, name: &str) -> Option<usize> {
172        self.p4_index.get(name).copied()
173    }
174
175    /// Resolves a four-momentum column name once for repeated row access.
176    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    /// Returns the scalar column index for `name`.
184    pub fn scalar_index(&self, name: &str) -> Option<usize> {
185        self.scalar_index.get(name).copied()
186    }
187
188    /// Resolves a scalar column name once for repeated row access.
189    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    /// Returns four-momentum names in column order.
197    pub fn p4s(&self) -> &[Name] {
198        &self.p4s
199    }
200
201    /// Returns scalar names in column order.
202    pub fn scalars(&self) -> &[Name] {
203        &self.scalars
204    }
205
206    /// Returns whether events carry explicit weights.
207    pub fn has_weight(&self) -> bool {
208        self.has_weight
209    }
210
211    /// Returns the number of four-momentum columns.
212    pub fn n_p4s(&self) -> usize {
213        self.p4s.len()
214    }
215
216    /// Returns the number of scalar columns.
217    pub fn n_scalars(&self) -> usize {
218        self.scalars.len()
219    }
220
221    /// Requires and returns a four-momentum column index.
222    ///
223    /// # Errors
224    ///
225    /// Returns [`LadduDataError::MissingColumn`] when `name` is not a
226    /// four-momentum column.
227    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    /// Requires and returns a scalar column index.
233    ///
234    /// # Errors
235    ///
236    /// Returns [`LadduDataError::MissingColumn`] when `name` is not a scalar
237    /// column.
238    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/// Physical naming conventions used to map a logical schema to storage columns.
257#[derive(Clone, Debug)]
258pub struct SchemaColumnNames {
259    /// Physical weight-column name.
260    pub weight_column: Name,
261    /// Suffixes for four-momentum components.
262    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/// Options controlling logical schema inference from physical columns.
275#[derive(Clone, Debug)]
276pub struct SchemaInferenceOptions {
277    /// Physical naming conventions.
278    pub column_names: SchemaColumnNames,
279    /// Whether inference fails if the weight column is absent.
280    pub require_weight: bool,
281    /// Whether incomplete four-momenta become independent scalar columns.
282    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/// Physical suffixes for `(E, px, py, pz)` columns.
296#[derive(Clone, Debug)]
297pub struct P4Suffixes {
298    /// Energy suffix.
299    pub e: &'static str,
300    /// X-momentum suffix.
301    pub px: &'static str,
302    /// Y-momentum suffix.
303    pub py: &'static str,
304    /// Z-momentum suffix.
305    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    /// Splits a matching physical name into logical prefix and component index.
321    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    /// Produces the four physical names for a logical four-momentum prefix.
336    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/// Physical storage type relevant to schema inference.
347#[derive(Clone, Copy, Debug, PartialEq, Eq)]
348pub enum ColumnType {
349    /// An exact integer row-data column.
350    Integer(ColumnDType),
351    /// 64-bit floating point.
352    F64,
353    /// 32-bit floating point.
354    F32,
355    /// Any unsupported type.
356    Other,
357}
358
359impl ColumnType {
360    /// Returns whether this type can populate an event-data column.
361    pub fn is_supported_float(self) -> bool {
362        matches!(self, Self::F64 | Self::F32)
363    }
364}
365
366/// Name and physical type of one available storage column.
367#[derive(Clone, Copy, Debug)]
368pub struct ColumnInfo<'a> {
369    /// Physical column name.
370    pub name: &'a str,
371    /// Physical column type.
372    pub dtype: ColumnType,
373}
374
375impl Schema {
376    /// Infers a logical schema from available physical columns.
377    ///
378    /// # Errors
379    ///
380    /// Returns [`LadduDataError`] when required momentum components or weights
381    /// are missing, names are ambiguous, or inferred logical names conflict.
382    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    /// Returns all physical columns required to store this schema.
441    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    /// Validates that all physical columns required by this schema are available.
450    ///
451    /// # Errors
452    ///
453    /// Returns [`LadduDataError::MissingColumn`] when a required physical
454    /// column is absent or has an unsupported type.
455    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/// Floating-point precision used when writing physical columns.
492#[derive(Copy, Clone, Debug, Default)]
493pub enum Precision {
494    /// Write 64-bit floating-point values.
495    #[default]
496    F64,
497    /// Write 32-bit floating-point values.
498    F32,
499}
500
501/// Policy controlling whether sinks emit a weight column.
502#[derive(Clone, Copy, Debug, Default)]
503pub enum WriteWeightColumn {
504    /// Always write weights, using unit weights when the schema has none.
505    #[default]
506    Always,
507    /// Write weights only when the logical schema contains them.
508    OnlyIfPresent,
509}
510
511/// Physical naming, precision, and weight policy for event sinks.
512#[derive(Clone, Debug, Default)]
513pub struct SchemaWriteOptions {
514    /// Physical column naming conventions.
515    pub column_names: SchemaColumnNames,
516    /// Floating-point output precision.
517    pub precision: Precision,
518    /// Weight-column emission policy.
519    pub write_weight_column: WriteWeightColumn,
520}
521
522/// The semantic role of one physical storage column.
523///
524/// This is deliberately crate-private: callers work with logical [`Schema`]
525/// values while the file-format adapters use this plan to keep physical
526/// column ordering and binding identical across backends.
527#[derive(Clone, Copy, Debug, PartialEq, Eq)]
528pub(crate) enum PhysicalColumnRole {
529    /// An exact non-expression column and its storage dtype.
530    Column { index: usize, dtype: ColumnDType },
531    /// One component of a logical four-momentum column.
532    P4 {
533        /// Logical four-momentum column index.
534        index: usize,
535        /// Component index in `(E, px, py, pz)` order.
536        component: usize,
537    },
538    /// One logical scalar column.
539    Scalar {
540        /// Logical scalar column index.
541        index: usize,
542    },
543    /// The physical event-weight column.
544    Weight,
545}
546
547/// Ordered physical representation of a logical event schema.
548///
549/// A plan is built once at each storage boundary and then consumed by schema
550/// creation, projection, encoding, and decoding.  Keeping the role alongside
551/// the name avoids each backend reimplementing the canonical ordering.
552#[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    /// Builds the physical columns used when reading or validating a schema.
565    pub(crate) fn for_read(schema: &Schema, column_names: &SchemaColumnNames) -> Self {
566        Self::build(schema, column_names, schema.has_weight())
567    }
568
569    /// Builds the physical columns emitted by a sink.
570    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    /// Returns physical columns in canonical storage order.
627    pub(crate) fn columns(&self) -> &[PhysicalColumn] {
628        &self.columns
629    }
630}
631
632impl PhysicalColumn {
633    /// Returns the physical storage name.
634    pub(crate) fn name(&self) -> &Name {
635        &self.name
636    }
637
638    /// Returns the logical role of this physical column.
639    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}