Skip to main content

laddu_data/data/
event.rs

1use std::fmt;
2use std::sync::Arc;
3
4use laddu_physics::vectors::RealVec4;
5
6use crate::{
7    BatchLayout, LadduDataError, LadduDataResult,
8    columns::{Column, ColumnBuffer, ColumnValue},
9    schema::{P4Binding, Precision, ScalarBinding, Schema},
10};
11
12#[derive(Clone, Debug)]
13struct BatchParts {
14    p4s: Arc<[Arc<[RealVec4]>]>,
15    scalars: Arc<[Arc<[f64]>]>,
16    columns: Arc<[Column]>,
17    weights: Weights,
18}
19
20#[derive(Clone, Debug)]
21enum Weights {
22    ImplicitUnit,
23    Explicit(Arc<[f64]>),
24}
25
26impl Weights {
27    fn from_option(weights: Option<Arc<[f64]>>) -> Self {
28        match weights {
29            Some(weights) => Self::Explicit(weights),
30            None => Self::ImplicitUnit,
31        }
32    }
33
34    fn as_slice(&self) -> Option<&[f64]> {
35        match self {
36            Self::ImplicitUnit => None,
37            Self::Explicit(weights) => Some(weights),
38        }
39    }
40
41    fn at(&self, row: usize) -> f64 {
42        self.as_slice().map_or(1.0, |weights| weights[row])
43    }
44
45    fn is_explicit(&self) -> bool {
46        matches!(self, Self::Explicit(_))
47    }
48
49    fn select(&self, rows: &[usize]) -> Self {
50        match self {
51            Self::ImplicitUnit => Self::ImplicitUnit,
52            Self::Explicit(weights) => {
53                let selected: Arc<[f64]> = rows.iter().map(|&row| weights[row]).collect();
54                Self::Explicit(selected)
55            }
56        }
57    }
58
59    fn slice(&self, start: usize, end: usize) -> Self {
60        match self {
61            Self::ImplicitUnit => Self::ImplicitUnit,
62            Self::Explicit(weights) => Self::Explicit(Arc::from(&weights[start..end])),
63        }
64    }
65
66    fn reweight<F>(&self, len: usize, f: F) -> Self
67    where
68        F: Fn(usize, f64) -> f64,
69    {
70        let weights: Arc<[f64]> = (0..len).map(|i| f(i, self.at(i))).collect();
71        Self::Explicit(weights)
72    }
73}
74
75impl BatchParts {
76    fn from_columns(p4s: Vec<Arc<[RealVec4]>>, scalars: Vec<Arc<[f64]>>, weights: Weights) -> Self {
77        Self {
78            p4s: p4s.into(),
79            scalars: scalars.into(),
80            columns: Arc::from([]),
81            weights,
82        }
83    }
84
85    fn validate(&self, schema: &Schema, expected_len: Option<usize>) -> LadduDataResult<usize> {
86        if self.p4s.len() != schema.n_p4s() {
87            return Err(LadduDataError::Schema(
88                "wrong number of vec4 columns".into(),
89            ));
90        }
91
92        if self.scalars.len() != schema.n_scalars() {
93            return Err(LadduDataError::Schema(
94                "wrong number of scalar columns".into(),
95            ));
96        }
97
98        if self.columns.len() != schema.n_columns() {
99            return Err(LadduDataError::Schema(
100                "wrong number of typed columns".into(),
101            ));
102        }
103        for (column, (name, dtype)) in self.columns.iter().zip(schema.columns()) {
104            if column.dtype() != *dtype {
105                return Err(LadduDataError::Schema(format!(
106                    "column {name:?} dtype does not match schema"
107                )));
108            }
109        }
110        let len = infer_len(
111            &self.p4s,
112            &self.scalars,
113            &self.columns,
114            self.weights.as_slice(),
115        )?;
116        if let Some(expected_len) = expected_len {
117            let has_columns = !self.p4s.is_empty()
118                || !self.scalars.is_empty()
119                || !self.columns.is_empty()
120                || self.weights.is_explicit();
121            if has_columns && len != expected_len {
122                return Err(LadduDataError::Schema("inconsistent batch length".into()));
123            }
124            return Ok(expected_len);
125        }
126        Ok(len)
127    }
128
129    fn select(&self, rows: &[usize]) -> Self {
130        let p4s = self
131            .p4s
132            .iter()
133            .map(|col| rows.iter().map(|&i| col[i]).collect())
134            .collect();
135        let scalars = self
136            .scalars
137            .iter()
138            .map(|col| rows.iter().map(|&i| col[i]).collect())
139            .collect();
140
141        Self {
142            p4s,
143            scalars,
144            weights: self.weights.select(rows),
145            columns: self
146                .columns
147                .iter()
148                .map(|column| column.select(rows))
149                .collect(),
150        }
151    }
152
153    fn slice(&self, start: usize, end: usize) -> Self {
154        let p4s = self
155            .p4s
156            .iter()
157            .map(|col| Arc::<[RealVec4]>::from(&col[start..end]))
158            .collect();
159        let scalars = self
160            .scalars
161            .iter()
162            .map(|col| Arc::<[f64]>::from(&col[start..end]))
163            .collect();
164
165        Self {
166            p4s,
167            scalars,
168            weights: self.weights.slice(start, end),
169            columns: self
170                .columns
171                .iter()
172                .map(|column| column.slice(start, end))
173                .collect(),
174        }
175    }
176
177    fn reweight<F>(&self, len: usize, f: F) -> Self
178    where
179        F: Fn(usize, f64) -> f64,
180    {
181        Self {
182            p4s: Arc::clone(&self.p4s),
183            scalars: Arc::clone(&self.scalars),
184            columns: Arc::clone(&self.columns),
185            weights: self.weights.reweight(len, f),
186        }
187    }
188
189    fn concat(batches: &[(&Self, usize)]) -> LadduDataResult<Self> {
190        let len: usize = batches.iter().map(|(_, len)| *len).sum();
191        let n_p4s = batches.first().map_or(0, |(batch, _)| batch.p4s.len());
192        let n_scalars = batches.first().map_or(0, |(batch, _)| batch.scalars.len());
193
194        let mut p4s = Vec::with_capacity(n_p4s);
195        for col in 0..n_p4s {
196            let mut out = Vec::with_capacity(len);
197            for (batch, _) in batches {
198                out.extend_from_slice(&batch.p4s[col]);
199            }
200            p4s.push(Arc::from(out));
201        }
202
203        let mut scalars = Vec::with_capacity(n_scalars);
204        for col in 0..n_scalars {
205            let mut out = Vec::with_capacity(len);
206            for (batch, _) in batches {
207                out.extend_from_slice(&batch.scalars[col]);
208            }
209            scalars.push(Arc::from(out));
210        }
211
212        let weights = if batches.iter().any(|(batch, _)| batch.weights.is_explicit()) {
213            let mut out = Vec::with_capacity(len);
214            for (batch, batch_len) in batches {
215                for row in 0..*batch_len {
216                    out.push(batch.weights.at(row));
217                }
218            }
219            Weights::Explicit(Arc::from(out))
220        } else {
221            Weights::ImplicitUnit
222        };
223
224        let mut parts = Self::from_columns(p4s, scalars, weights);
225        if let Some((first, _)) = batches.first() {
226            parts.columns = first
227                .columns
228                .iter()
229                .enumerate()
230                .map(|(index, column)| {
231                    let inputs = batches
232                        .iter()
233                        .map(|(batch, _)| &batch.columns[index])
234                        .collect::<Vec<_>>();
235                    Column::concat(column.dtype(), &inputs)
236                })
237                .collect::<LadduDataResult<Vec<_>>>()?
238                .into();
239        }
240        Ok(parts)
241    }
242}
243
244#[derive(Default)]
245enum WeightAssembler {
246    #[default]
247    ImplicitUnit,
248    Explicit(Vec<f64>),
249}
250
251impl WeightAssembler {
252    fn push(&mut self, weight: Option<f64>, len: usize) -> LadduDataResult<()> {
253        match self {
254            Self::Explicit(weights) => match weight {
255                Some(weight) => weights.push(weight),
256                None => {
257                    return Err(LadduDataError::InvalidArgument(
258                        "cannot mix weighted and unweighted events in one batch",
259                    ));
260                }
261            },
262            Self::ImplicitUnit => match weight {
263                Some(weight) if len == 0 => *self = Self::Explicit(vec![weight]),
264                Some(_) => {
265                    return Err(LadduDataError::InvalidArgument(
266                        "cannot mix unweighted and weighted events in one batch",
267                    ));
268                }
269                None => {}
270            },
271        }
272
273        Ok(())
274    }
275
276    fn finish(self) -> Weights {
277        match self {
278            Self::ImplicitUnit => Weights::ImplicitUnit,
279            Self::Explicit(weights) => Weights::Explicit(Arc::from(weights)),
280        }
281    }
282}
283
284/// Immutable columnar batch of events sharing one schema.
285#[derive(Clone)]
286pub struct EventBatch {
287    schema: Arc<Schema>,
288    len: usize,
289    parts: BatchParts,
290    row_ids: Option<Arc<[u64]>>,
291}
292
293impl fmt::Debug for EventBatch {
294    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
295        formatter
296            .debug_struct("EventBatch")
297            .field("schema", &self.schema)
298            .field("len", &self.len)
299            .field("p4s", &self.parts.p4s)
300            .field("scalars", &self.parts.scalars)
301            .field("columns", &self.parts.columns)
302            .field("weights", &self.parts.weights.as_slice())
303            .finish()
304    }
305}
306
307impl EventBatch {
308    /// Validates column counts and lengths and constructs a batch.
309    ///
310    /// # Errors
311    ///
312    /// Returns [`LadduDataError`] when column counts do not match `schema` or
313    /// column and weight lengths are inconsistent.
314    pub fn new(
315        schema: Arc<Schema>,
316        p4s: Vec<Arc<[RealVec4]>>,
317        scalars: Vec<Arc<[f64]>>,
318        weights: Option<Arc<[f64]>>,
319    ) -> LadduDataResult<Self> {
320        BatchAssembler::from_columns(schema, p4s, scalars, weights)
321    }
322
323    /// Constructs a batch including exact non-expression columns.
324    ///
325    /// # Errors
326    /// Returns an error for schema, dtype, or column-length mismatches.
327    pub fn new_with_columns(
328        schema: Arc<Schema>,
329        p4s: Vec<Arc<[RealVec4]>>,
330        scalars: Vec<Arc<[f64]>>,
331        columns: Vec<Column>,
332        weights: Option<Arc<[f64]>>,
333    ) -> LadduDataResult<Self> {
334        let mut parts = BatchParts::from_columns(p4s, scalars, Weights::from_option(weights));
335        parts.columns = columns.into();
336        Self::from_parts(schema, parts)
337    }
338
339    /// Constructs a typed batch with an explicit row count.
340    ///
341    /// # Errors
342    /// Returns an error if columns do not match the schema or row count.
343    pub fn new_with_columns_and_len(
344        schema: Arc<Schema>,
345        p4s: Vec<Arc<[RealVec4]>>,
346        scalars: Vec<Arc<[f64]>>,
347        columns: Vec<Column>,
348        weights: Option<Arc<[f64]>>,
349        len: usize,
350    ) -> LadduDataResult<Self> {
351        let mut parts = BatchParts::from_columns(p4s, scalars, Weights::from_option(weights));
352        parts.columns = columns.into();
353        Self::from_parts_with_len(schema, parts, len)
354    }
355
356    /// Constructs a batch with an explicit event count, including column-free events.
357    ///
358    /// # Errors
359    ///
360    /// Returns an error if the columns do not match the schema or event count.
361    pub fn new_with_len(
362        schema: Arc<Schema>,
363        p4s: Vec<Arc<[RealVec4]>>,
364        scalars: Vec<Arc<[f64]>>,
365        weights: Option<Arc<[f64]>>,
366        len: usize,
367    ) -> LadduDataResult<Self> {
368        Self::from_parts_with_len(
369            schema,
370            BatchParts::from_columns(p4s, scalars, Weights::from_option(weights)),
371            len,
372        )
373    }
374
375    fn from_parts(schema: Arc<Schema>, parts: BatchParts) -> LadduDataResult<Self> {
376        let len = parts.validate(&schema, None)?;
377        Ok(Self {
378            schema,
379            len,
380            parts,
381            row_ids: None,
382        })
383    }
384
385    fn from_parts_with_len(
386        schema: Arc<Schema>,
387        parts: BatchParts,
388        expected_len: usize,
389    ) -> LadduDataResult<Self> {
390        let len = parts.validate(&schema, Some(expected_len))?;
391        Ok(Self {
392            schema,
393            len,
394            parts,
395            row_ids: None,
396        })
397    }
398
399    /// Collects owned row events into a columnar batch.
400    ///
401    /// # Errors
402    ///
403    /// Returns [`LadduDataError`] when an event has the wrong number of values
404    /// or weighted and unweighted events are mixed.
405    pub fn from_events<I>(schema: Arc<Schema>, events: I) -> LadduDataResult<Self>
406    where
407        I: IntoIterator<Item = OwnedEvent>,
408    {
409        let mut builder = EventBatchBuilder::new(schema);
410        builder.extend(events)?;
411        builder.finish()
412    }
413
414    /// Returns the shared logical schema.
415    pub fn schema(&self) -> &Arc<Schema> {
416        &self.schema
417    }
418
419    /// Returns the number of rows.
420    pub fn len(&self) -> usize {
421        self.len
422    }
423
424    /// Returns optional stable global source-row identities for distributed transforms.
425    pub fn row_ids(&self) -> Option<&[u64]> {
426        self.row_ids.as_deref()
427    }
428
429    /// Attaches stable source-row identities without changing the logical schema.
430    ///
431    /// # Errors
432    /// Returns an error when the identity count differs from the event count.
433    pub fn with_row_ids(mut self, ids: Arc<[u64]>) -> LadduDataResult<Self> {
434        if ids.len() != self.len {
435            return Err(LadduDataError::Schema(
436                "row identity count does not match event count".into(),
437            ));
438        }
439        self.row_ids = Some(ids);
440        Ok(self)
441    }
442
443    /// Returns the logical payload bytes per event represented by this batch.
444    pub fn bytes_per_event(&self) -> usize {
445        BatchLayout::from_batch(self)
446            .bytes_per_event(Precision::F64)
447            .ok()
448            .and_then(|bytes| usize::try_from(bytes).ok())
449            .unwrap_or(usize::MAX)
450    }
451
452    /// Returns the retained column and row-identity payload size in bytes.
453    ///
454    /// Shared schema metadata, allocation headers, and other owners of shared
455    /// columns are not included.
456    pub fn resident_bytes(&self) -> usize {
457        BatchLayout::from_batch(self)
458            .footprint(Precision::F64)
459            .and_then(|footprint| footprint.checked_peak_bytes(self.len))
460            .ok()
461            .and_then(|bytes| usize::try_from(bytes).ok())
462            .and_then(|bytes| {
463                let identities = self
464                    .row_ids
465                    .as_ref()
466                    .map_or(0, |ids| std::mem::size_of_val(ids.as_ref()));
467                bytes.checked_add(identities)
468            })
469            .unwrap_or(usize::MAX)
470    }
471
472    /// Returns whether the batch contains no rows.
473    pub fn is_empty(&self) -> bool {
474        self.len == 0
475    }
476
477    /// Returns a four-momentum column by index.
478    pub fn vec4_column(&self, index: usize) -> &[RealVec4] {
479        &self.parts.p4s[index]
480    }
481
482    /// Returns a scalar column by index.
483    pub fn scalar_column(&self, index: usize) -> &[f64] {
484        &self.parts.scalars[index]
485    }
486
487    /// Returns an exact row-data column by index.
488    ///
489    /// # Panics
490    /// Panics when the column index is outside the schema.
491    pub fn column(&self, index: usize) -> &Column {
492        &self.parts.columns[index]
493    }
494
495    /// Returns an exact row-data column by name.
496    pub fn column_named(&self, name: &str) -> Option<&Column> {
497        self.schema
498            .column_index(name)
499            .map(|index| self.column(index))
500    }
501
502    /// Returns the optional explicit weight column.
503    pub fn weights_column(&self) -> Option<&[f64]> {
504        self.parts.weights.as_slice()
505    }
506
507    /// Returns a four-momentum column by logical name.
508    pub fn vec4_column_named(&self, name: &str) -> Option<&[RealVec4]> {
509        let i = self.schema.p4_index(name)?;
510        Some(self.vec4_column(i))
511    }
512
513    /// Returns a scalar column by logical name.
514    pub fn scalar_column_named(&self, name: &str) -> Option<&[f64]> {
515        let i = self.schema.scalar_index(name)?;
516        Some(self.scalar_column(i))
517    }
518
519    /// Returns one four-momentum cell.
520    pub fn p4_at(&self, col: usize, row: usize) -> RealVec4 {
521        self.parts.p4s[col][row]
522    }
523
524    /// Returns a four-momentum column through a schema-resolved binding.
525    ///
526    /// Schema compatibility is checked once, so callers can safely reuse the
527    /// returned slice for every row in a packing loop.
528    ///
529    /// # Errors
530    ///
531    /// Returns [`LadduDataError::Schema`] when the binding belongs to another
532    /// schema.
533    pub fn p4_column_bound(&self, binding: &P4Binding) -> LadduDataResult<&[RealVec4]> {
534        if !binding.matches(&self.schema) {
535            return Err(LadduDataError::Schema(
536                "column binding belongs to a different schema".into(),
537            ));
538        }
539        Ok(self.vec4_column(binding.index()))
540    }
541
542    /// Returns one scalar cell.
543    pub fn scalar_at(&self, col: usize, row: usize) -> f64 {
544        self.parts.scalars[col][row]
545    }
546
547    /// Returns a scalar column through a schema-resolved binding.
548    ///
549    /// Schema compatibility is checked once, so callers can safely reuse the
550    /// returned slice for every row in a packing loop.
551    ///
552    /// # Errors
553    ///
554    /// Returns [`LadduDataError::Schema`] when the binding belongs to another
555    /// schema.
556    pub fn scalar_column_bound(&self, binding: &ScalarBinding) -> LadduDataResult<&[f64]> {
557        if !binding.matches(&self.schema) {
558            return Err(LadduDataError::Schema(
559                "column binding belongs to a different schema".into(),
560            ));
561        }
562        Ok(self.scalar_column(binding.index()))
563    }
564
565    /// Returns the explicit row weight, or one when weights are absent.
566    pub fn weights_at(&self, row: usize) -> f64 {
567        self.parts.weights.at(row)
568    }
569
570    /// Returns a borrowed view of one row.
571    pub fn event(&self, row: usize) -> BatchEvent<'_> {
572        BatchEvent { batch: self, row }
573    }
574
575    /// Iterates over borrowed event views.
576    pub fn iter(&self) -> impl Iterator<Item = BatchEvent<'_>> {
577        (0..self.len()).map(|i| self.event(i))
578    }
579
580    /// Copies selected rows into a new batch in the requested order.
581    ///
582    /// # Panics
583    ///
584    /// Panics when a selected row is outside the batch.
585    pub fn select(&self, rows: &[usize]) -> Self {
586        let mut selected = Self::from_parts_with_len(
587            Arc::clone(&self.schema),
588            self.parts.select(rows),
589            rows.len(),
590        )
591        .expect("select preserves EventBatch invariants");
592        selected.row_ids = self
593            .row_ids
594            .as_ref()
595            .map(|ids| rows.iter().map(|&i| ids[i]).collect());
596        selected
597    }
598
599    /// Copies rows satisfying `keep` into a new batch.
600    pub fn filter<F>(&self, keep: F) -> Self
601    where
602        F: Fn(BatchEvent<'_>) -> bool,
603    {
604        let rows: Vec<usize> = (0..self.len).filter(|&i| keep(self.event(i))).collect();
605
606        self.select(&rows)
607    }
608
609    /// Returns a batch sharing value columns with newly computed weights.
610    ///
611    /// # Panics
612    ///
613    /// Panics if an internal batch invariant is violated while rebuilding the
614    /// batch.
615    pub fn reweight<F>(&self, f: F) -> Self
616    where
617        F: Fn(usize, f64) -> f64,
618    {
619        let mut reweighted = Self::from_parts_with_len(
620            Arc::clone(&self.schema),
621            self.parts.reweight(self.len, f),
622            self.len,
623        )
624        .expect("reweight preserves EventBatch invariants");
625        reweighted.row_ids = self.row_ids.clone();
626        reweighted
627    }
628
629    /// Copies the half-open row range `start..end` into a new batch.
630    ///
631    /// # Panics
632    ///
633    /// Panics when `start > end` or `end` exceeds the batch length.
634    pub fn slice(&self, start: usize, end: usize) -> Self {
635        assert!(start <= end);
636        assert!(end <= self.len);
637
638        if start == 0 && end == self.len {
639            return self.clone();
640        }
641
642        let mut sliced = Self::from_parts_with_len(
643            Arc::clone(&self.schema),
644            self.parts.slice(start, end),
645            end - start,
646        )
647        .expect("slice preserves EventBatch invariants");
648        sliced.row_ids = self.row_ids.as_ref().map(|ids| Arc::from(&ids[start..end]));
649        sliced
650    }
651
652    /// Concatenates schema-compatible batches.
653    ///
654    /// # Errors
655    ///
656    /// Returns [`LadduDataError`] when `batches` is empty or contains
657    /// incompatible schemas.
658    pub fn concat(batches: &[Self]) -> LadduDataResult<Self> {
659        if batches.is_empty() {
660            return Err(LadduDataError::InvalidArgument(
661                "cannot concatenate zero batches",
662            ));
663        }
664
665        let schema = Arc::clone(&batches[0].schema);
666
667        for batch in batches {
668            if schema != batch.schema {
669                return Err(LadduDataError::Schema(
670                    "cannot concatenate batches with different schemas".into(),
671                ));
672            }
673        }
674
675        let parts = batches
676            .iter()
677            .map(|batch| (&batch.parts, batch.len))
678            .collect::<Vec<_>>();
679        let len = batches.iter().map(|batch| batch.len).sum();
680        let mut combined = Self::from_parts_with_len(schema, BatchParts::concat(&parts)?, len)?;
681        if batches.iter().all(|b| b.row_ids.is_some()) {
682            combined.row_ids = Some(
683                batches
684                    .iter()
685                    .flat_map(|b| b.row_ids.iter().flat_map(|ids| ids.iter().copied()))
686                    .collect(),
687            );
688        }
689        Ok(combined)
690    }
691}
692
693fn infer_len(
694    vec4s: &[Arc<[RealVec4]>],
695    scalars: &[Arc<[f64]>],
696    columns: &[Column],
697    weight: Option<&[f64]>,
698) -> LadduDataResult<usize> {
699    let len = vec4s
700        .first()
701        .map(|c| c.len())
702        .or_else(|| scalars.first().map(|c| c.len()))
703        .or_else(|| columns.first().map(Column::len))
704        .or_else(|| weight.map(|w| w.len()))
705        .unwrap_or(0);
706
707    for col in vec4s {
708        if col.len() != len {
709            return Err(LadduDataError::Schema(
710                "inconsistent vec4 column length".into(),
711            ));
712        }
713    }
714
715    for col in scalars {
716        if col.len() != len {
717            return Err(LadduDataError::Schema(
718                "inconsistent scalar column length".into(),
719            ));
720        }
721    }
722
723    if columns.iter().any(|column| column.len() != len) {
724        return Err(LadduDataError::Schema(
725            "inconsistent typed column length".into(),
726        ));
727    }
728
729    if let Some(w) = weight
730        && w.len() != len
731    {
732        return Err(LadduDataError::Schema("inconsistent weight length".into()));
733    }
734
735    Ok(len)
736}
737
738/// Borrowed view of one row in an [`EventBatch`].
739#[derive(Copy, Clone, Debug)]
740pub struct BatchEvent<'a> {
741    batch: &'a EventBatch,
742    row: usize,
743}
744
745impl<'a> BatchEvent<'a> {
746    /// Returns the row index.
747    pub fn row(&self) -> usize {
748        self.row
749    }
750
751    /// Returns the backing batch.
752    pub fn batch(&self) -> &'a EventBatch {
753        self.batch
754    }
755
756    /// Returns a four-momentum value by column index.
757    pub fn p4(&self, col: usize) -> RealVec4 {
758        self.batch.p4_at(col, self.row)
759    }
760
761    /// Returns a scalar value by column index.
762    pub fn scalar(&self, col: usize) -> f64 {
763        self.batch.scalar_at(col, self.row)
764    }
765
766    /// Returns an exact row-data value by name.
767    pub fn column_named(&self, name: &str) -> Option<ColumnValue> {
768        self.batch
769            .column_named(name)
770            .map(|column| column.at(self.row))
771    }
772
773    /// Returns the row weight, defaulting to one.
774    pub fn weight(&self) -> f64 {
775        self.batch.weights_at(self.row)
776    }
777
778    /// Returns a four-momentum value by logical name.
779    pub fn p4_named(&self, name: &str) -> Option<RealVec4> {
780        let col = self.batch.schema.p4_index(name)?;
781        Some(self.p4(col))
782    }
783
784    /// Returns a scalar value by logical name.
785    pub fn scalar_named(&self, name: &str) -> Option<f64> {
786        let col = self.batch.schema.scalar_index(name)?;
787        Some(self.scalar(col))
788    }
789}
790
791/// Borrowed event view with a possibly transformed weight.
792#[derive(Copy, Clone, Debug)]
793pub struct Event<'a> {
794    pub(super) batch: &'a EventBatch,
795    pub(super) row: usize,
796    pub(super) weight: f64,
797}
798
799impl<'a> Event<'a> {
800    /// Returns the row index in the backing batch.
801    pub fn row(&self) -> usize {
802        self.row
803    }
804
805    /// Returns a four-momentum value by column index.
806    pub fn p4(&self, col: usize) -> RealVec4 {
807        self.batch.p4_at(col, self.row)
808    }
809
810    /// Returns a scalar value by column index.
811    pub fn scalar(&self, col: usize) -> f64 {
812        self.batch.scalar_at(col, self.row)
813    }
814
815    /// Returns an exact row-data value by name.
816    pub fn column_named(&self, name: &str) -> Option<ColumnValue> {
817        self.batch
818            .column_named(name)
819            .map(|column| column.at(self.row))
820    }
821
822    /// Returns this view's effective weight.
823    pub fn weight(&self) -> f64 {
824        self.weight
825    }
826
827    /// Returns a four-momentum value by logical name.
828    pub fn p4_named(&self, name: &str) -> Option<RealVec4> {
829        let col = self.batch.schema.p4_index(name)?;
830        Some(self.p4(col))
831    }
832
833    /// Returns a scalar value by logical name.
834    pub fn scalar_named(&self, name: &str) -> Option<f64> {
835        let col = self.batch.schema.scalar_index(name)?;
836        Some(self.scalar(col))
837    }
838}
839
840/// Owned row-oriented event used while constructing batches.
841#[derive(Clone, Debug)]
842pub struct OwnedEvent {
843    /// Four-momentum values in schema order.
844    pub p4s: Vec<RealVec4>,
845    /// Scalar values in schema order.
846    pub scalars: Vec<f64>,
847    /// Optional explicit event weight.
848    pub weight: Option<f64>,
849}
850
851impl OwnedEvent {
852    /// Creates an unweighted owned event.
853    pub fn new(p4s: Vec<RealVec4>, scalars: Vec<f64>) -> Self {
854        Self {
855            p4s,
856            scalars,
857            weight: None,
858        }
859    }
860
861    /// Creates an owned event with an explicit weight.
862    pub fn weighted(p4s: Vec<RealVec4>, scalars: Vec<f64>, weight: f64) -> Self {
863        Self {
864            p4s,
865            scalars,
866            weight: Some(weight),
867        }
868    }
869}
870
871/// Shared checked assembly for row-oriented batch producers.
872pub(crate) struct BatchAssembler {
873    schema: Arc<Schema>,
874    p4s: Vec<Vec<RealVec4>>,
875    scalars: Vec<Vec<f64>>,
876    columns: Vec<ColumnBuffer>,
877    weights: WeightAssembler,
878    len: usize,
879}
880
881impl BatchAssembler {
882    pub(crate) fn new(schema: Arc<Schema>, capacity: usize) -> Self {
883        let p4s = (0..schema.n_p4s())
884            .map(|_| Vec::with_capacity(capacity))
885            .collect();
886        let scalars = (0..schema.n_scalars())
887            .map(|_| Vec::with_capacity(capacity))
888            .collect();
889
890        Self {
891            columns: schema
892                .columns()
893                .iter()
894                .map(|(_, dtype)| ColumnBuffer::new(*dtype, capacity))
895                .collect(),
896            schema,
897            p4s,
898            scalars,
899            weights: WeightAssembler::default(),
900            len: 0,
901        }
902    }
903
904    pub(crate) fn with_weight_mode(
905        schema: Arc<Schema>,
906        capacity: usize,
907        explicit_weights: bool,
908    ) -> Self {
909        let mut assembler = Self::new(schema, capacity);
910        if explicit_weights {
911            assembler.weights = WeightAssembler::Explicit(Vec::with_capacity(capacity));
912        }
913        assembler
914    }
915
916    pub(crate) fn from_columns(
917        schema: Arc<Schema>,
918        p4s: Vec<Arc<[RealVec4]>>,
919        scalars: Vec<Arc<[f64]>>,
920        weights: Option<Arc<[f64]>>,
921    ) -> LadduDataResult<EventBatch> {
922        EventBatch::from_parts(
923            schema,
924            BatchParts::from_columns(p4s, scalars, Weights::from_option(weights)),
925        )
926    }
927
928    fn push_owned(&mut self, event: OwnedEvent) -> LadduDataResult<()> {
929        if self.schema.n_columns() != 0 {
930            return Err(LadduDataError::Unsupported(
931                "owned float events cannot populate typed columns; construct typed columnar batches",
932            ));
933        }
934        if event.p4s.len() != self.schema.n_p4s() {
935            return Err(LadduDataError::Schema(
936                "wrong number of event vec4 values".into(),
937            ));
938        }
939
940        if event.scalars.len() != self.schema.n_scalars() {
941            return Err(LadduDataError::Schema(
942                "wrong number of event scalar values".into(),
943            ));
944        }
945
946        self.weights.push(event.weight, self.len)?;
947
948        for (col, value) in event.p4s.into_iter().enumerate() {
949            self.p4s[col].push(value);
950        }
951        for (col, value) in event.scalars.into_iter().enumerate() {
952            self.scalars[col].push(value);
953        }
954
955        self.len += 1;
956        Ok(())
957    }
958
959    pub(crate) fn push_borrowed(
960        &mut self,
961        event: Event<'_>,
962        explicit_weight: bool,
963    ) -> LadduDataResult<()> {
964        if self.schema.columns() != event.batch.schema().columns() {
965            return Err(LadduDataError::Schema(
966                "typed event columns do not match schema".into(),
967            ));
968        }
969        if event.batch.schema().n_p4s() != self.schema.n_p4s() {
970            return Err(LadduDataError::Schema(
971                "wrong number of event vec4 values".into(),
972            ));
973        }
974
975        if event.batch.schema().n_scalars() != self.schema.n_scalars() {
976            return Err(LadduDataError::Schema(
977                "wrong number of event scalar values".into(),
978            ));
979        }
980
981        self.weights
982            .push(explicit_weight.then_some(event.weight()), self.len)?;
983
984        for col in 0..self.schema.n_p4s() {
985            self.p4s[col].push(event.p4(col));
986        }
987        for col in 0..self.schema.n_scalars() {
988            self.scalars[col].push(event.scalar(col));
989        }
990
991        for (index, column) in self.columns.iter_mut().enumerate() {
992            column.push(event.batch.column(index).at(event.row))?;
993        }
994
995        self.len += 1;
996        Ok(())
997    }
998
999    pub(crate) fn finish(self) -> LadduDataResult<EventBatch> {
1000        let mut parts = BatchParts::from_columns(
1001            self.p4s.into_iter().map(Arc::from).collect(),
1002            self.scalars.into_iter().map(Arc::from).collect(),
1003            self.weights.finish(),
1004        );
1005        parts.columns = self.columns.into_iter().map(ColumnBuffer::finish).collect();
1006        EventBatch::from_parts_with_len(self.schema, parts, self.len)
1007    }
1008}
1009
1010/// Incremental builder for a columnar [`EventBatch`].
1011pub struct EventBatchBuilder {
1012    assembler: BatchAssembler,
1013}
1014
1015impl EventBatchBuilder {
1016    /// Creates an empty builder.
1017    pub fn new(schema: Arc<Schema>) -> Self {
1018        Self::with_capacity(schema, 0)
1019    }
1020
1021    /// Creates an empty builder with per-column capacity.
1022    pub fn with_capacity(schema: Arc<Schema>, capacity: usize) -> Self {
1023        Self {
1024            assembler: BatchAssembler::new(schema, capacity),
1025        }
1026    }
1027
1028    /// Appends an unweighted event from ordered values.
1029    ///
1030    /// # Errors
1031    ///
1032    /// Returns [`LadduDataError`] when value counts do not match the schema or
1033    /// the builder already contains weighted events.
1034    pub fn push<P, S>(&mut self, p4s: P, scalars: S) -> LadduDataResult<&mut Self>
1035    where
1036        P: IntoIterator<Item = RealVec4>,
1037        S: IntoIterator<Item = f64>,
1038    {
1039        self.push_event(OwnedEvent::new(
1040            p4s.into_iter().collect(),
1041            scalars.into_iter().collect(),
1042        ))
1043    }
1044
1045    /// Appends a weighted event from ordered values.
1046    ///
1047    /// # Errors
1048    ///
1049    /// Returns [`LadduDataError`] when value counts do not match the schema or
1050    /// the builder already contains unweighted events.
1051    pub fn push_weighted<P, S>(
1052        &mut self,
1053        p4s: P,
1054        scalars: S,
1055        weight: f64,
1056    ) -> LadduDataResult<&mut Self>
1057    where
1058        P: IntoIterator<Item = RealVec4>,
1059        S: IntoIterator<Item = f64>,
1060    {
1061        self.push_event(OwnedEvent::weighted(
1062            p4s.into_iter().collect(),
1063            scalars.into_iter().collect(),
1064            weight,
1065        ))
1066    }
1067
1068    /// Validates and appends one owned event.
1069    ///
1070    /// # Errors
1071    ///
1072    /// Returns [`LadduDataError`] when the event shape does not match the
1073    /// schema or its weight presence differs from prior events.
1074    pub fn push_event(&mut self, event: OwnedEvent) -> LadduDataResult<&mut Self> {
1075        self.assembler.push_owned(event)?;
1076        Ok(self)
1077    }
1078
1079    /// Appends all owned events from an iterator.
1080    ///
1081    /// # Errors
1082    ///
1083    /// Returns the first [`LadduDataError`] produced by an event whose shape or
1084    /// weight presence is incompatible with the builder.
1085    pub fn extend<I>(&mut self, events: I) -> LadduDataResult<&mut Self>
1086    where
1087        I: IntoIterator<Item = OwnedEvent>,
1088    {
1089        for event in events {
1090            self.push_event(event)?;
1091        }
1092
1093        Ok(self)
1094    }
1095
1096    /// Finalizes the builder into an immutable batch.
1097    ///
1098    /// # Errors
1099    ///
1100    /// Returns [`LadduDataError`] if the accumulated columns or weights have
1101    /// inconsistent lengths.
1102    pub fn finish(self) -> LadduDataResult<EventBatch> {
1103        self.assembler.finish()
1104    }
1105}
1106
1107#[cfg(test)]
1108mod tests {
1109    use super::*;
1110
1111    #[test]
1112    fn column_free_batches_preserve_counts_and_global_identities() -> LadduDataResult<()> {
1113        let schema = Arc::new(Schema::new(
1114            Vec::<String>::new(),
1115            Vec::<String>::new(),
1116            false,
1117        )?);
1118        let batch = EventBatch::new_with_len(schema, vec![], vec![], None, 3)?
1119            .with_row_ids(Arc::from([2, 5, 8]))?;
1120        assert_eq!(batch.resident_bytes(), 3 * std::mem::size_of::<u64>());
1121        assert_eq!(batch.reweight(|_, w| w).len(), 3);
1122        assert_eq!(batch.select(&[2, 0]).row_ids(), Some([8, 2].as_slice()));
1123        let rebuilt = EventBatch::concat(&[batch.slice(0, 1), batch.slice(1, 3)])?;
1124        assert_eq!(rebuilt.row_ids(), Some([2, 5, 8].as_slice()));
1125        assert!(batch.clone().with_row_ids(Arc::from([1])).is_err());
1126        Ok(())
1127    }
1128
1129    fn v(x: f64) -> RealVec4 {
1130        RealVec4 {
1131            e: x + 0.3,
1132            px: x,
1133            py: x + 0.1,
1134            pz: x + 0.2,
1135        }
1136    }
1137
1138    fn schema_with_weight() -> Arc<Schema> {
1139        Arc::new(Schema::new(["p"], ["x"], true).unwrap())
1140    }
1141
1142    fn weighted_batch(start: usize, len: usize) -> EventBatch {
1143        let schema = schema_with_weight();
1144
1145        let events = (start..start + len)
1146            .map(|i| OwnedEvent::weighted(vec![v(i as f64)], vec![i as f64], 10.0 + i as f64));
1147
1148        EventBatch::from_events(schema, events).unwrap()
1149    }
1150
1151    fn scalar_values(batch: &EventBatch) -> Vec<f64> {
1152        batch.scalar_column(0).to_vec()
1153    }
1154
1155    #[test]
1156    fn bound_columns_reuse_schema_resolution_for_row_access() {
1157        let batch = weighted_batch(3, 2);
1158        let scalar = batch.schema().bind_scalar("x").unwrap();
1159        let p4 = batch.schema().bind_p4("p").unwrap();
1160
1161        assert_eq!(batch.scalar_column_bound(&scalar).unwrap(), &[3.0, 4.0]);
1162        assert_eq!(batch.p4_column_bound(&p4).unwrap(), &[v(3.0), v(4.0)]);
1163
1164        let other_schema = Arc::new(Schema::new(["other"], ["x"], true).unwrap());
1165        let other = other_schema.bind_scalar("x").unwrap();
1166        assert!(matches!(
1167            batch.scalar_column_bound(&other),
1168            Err(LadduDataError::Schema(message)) if message == "column binding belongs to a different schema"
1169        ));
1170    }
1171
1172    #[test]
1173    fn event_batch_rejects_shape_mismatches_and_builder_rejects_mixed_weights() {
1174        let schema = schema_with_weight();
1175
1176        let bad_vec4_count = EventBatch::new(
1177            Arc::clone(&schema),
1178            vec![],
1179            vec![Arc::from([1.0, 2.0])],
1180            Some(Arc::from([1.0, 2.0])),
1181        );
1182
1183        assert!(matches!(bad_vec4_count, Err(LadduDataError::Schema(_))));
1184
1185        let bad_lengths = EventBatch::new(
1186            Arc::clone(&schema),
1187            vec![Arc::from([v(1.0), v(2.0)])],
1188            vec![Arc::from([1.0])],
1189            Some(Arc::from([1.0, 2.0])),
1190        );
1191
1192        assert!(matches!(bad_lengths, Err(LadduDataError::Schema(_))));
1193
1194        let mut builder = EventBatchBuilder::new(schema);
1195        builder.push([v(1.0)], [1.0]).unwrap();
1196
1197        let mixed = builder.push_weighted([v(2.0)], [2.0], 2.0);
1198        assert!(matches!(mixed, Err(LadduDataError::InvalidArgument(_))));
1199    }
1200
1201    #[test]
1202    fn select_slice_filter_reweight_and_concat_preserve_columns_and_weight_semantics() {
1203        let weighted = weighted_batch(0, 4);
1204        let selected = weighted.select(&[3, 1]);
1205
1206        assert_eq!(scalar_values(&selected), vec![3.0, 1.0]);
1207        assert_eq!(selected.weights_column().unwrap(), &[13.0, 11.0]);
1208        assert_eq!(selected.p4_at(0, 0).px, 3.0);
1209        assert_eq!(selected.p4_at(0, 1).e, 1.3);
1210
1211        let sliced = weighted.slice(1, 3);
1212        assert_eq!(scalar_values(&sliced), vec![1.0, 2.0]);
1213        assert_eq!(sliced.weights_column().unwrap(), &[11.0, 12.0]);
1214
1215        let filtered = weighted.filter(|ev| ev.scalar(0) >= 2.0);
1216        assert_eq!(scalar_values(&filtered), vec![2.0, 3.0]);
1217
1218        let reweighted = filtered.reweight(|i, w| w + 100.0 + i as f64);
1219        assert_eq!(reweighted.weights_column().unwrap(), &[112.0, 114.0]);
1220
1221        let schema = schema_with_weight();
1222
1223        let unweighted_with_weight_schema = EventBatch::from_events(
1224            Arc::clone(&schema),
1225            [
1226                OwnedEvent::new(vec![v(100.0)], vec![100.0]),
1227                OwnedEvent::new(vec![v(101.0)], vec![101.0]),
1228            ],
1229        )
1230        .unwrap();
1231
1232        let weighted_tail = EventBatch::from_events(
1233            schema,
1234            [
1235                OwnedEvent::weighted(vec![v(200.0)], vec![200.0], 5.0),
1236                OwnedEvent::weighted(vec![v(201.0)], vec![201.0], 6.0),
1237            ],
1238        )
1239        .unwrap();
1240
1241        let concatenated =
1242            EventBatch::concat(&[unweighted_with_weight_schema, weighted_tail]).unwrap();
1243
1244        assert_eq!(
1245            scalar_values(&concatenated),
1246            vec![100.0, 101.0, 200.0, 201.0]
1247        );
1248        assert_eq!(
1249            concatenated.weights_column().unwrap(),
1250            &[1.0, 1.0, 5.0, 6.0]
1251        );
1252    }
1253
1254    #[test]
1255    fn implicit_unit_weights_survive_assembly_and_row_transforms() {
1256        let schema = Arc::new(Schema::new(["p"], ["x"], false).unwrap());
1257        let batch = EventBatch::from_events(
1258            Arc::clone(&schema),
1259            (0..3).map(|i| OwnedEvent::new(vec![v(i as f64)], vec![i as f64])),
1260        )
1261        .unwrap();
1262
1263        assert!(batch.weights_column().is_none());
1264        assert_eq!(batch.weights_at(2), 1.0);
1265
1266        let selected = batch.select(&[2, 0]);
1267        let sliced = batch.slice(1, 3);
1268        let filtered = batch.filter(|event| event.scalar(0) > 0.0);
1269        let concatenated = EventBatch::concat(&[selected, sliced]).unwrap();
1270
1271        assert!(filtered.weights_column().is_none());
1272        assert!(concatenated.weights_column().is_none());
1273        assert_eq!(concatenated.weights_at(3), 1.0);
1274
1275        let reweighted = batch.reweight(|row, weight| weight + row as f64);
1276        assert_eq!(reweighted.weights_column().unwrap(), &[1.0, 2.0, 3.0]);
1277    }
1278
1279    #[test]
1280    fn shared_assembler_preserves_weight_mode_and_rejects_transitions() {
1281        let schema = Arc::new(Schema::new(["p"], ["x"], true).unwrap());
1282        let source = EventBatch::from_events(
1283            Arc::clone(&schema),
1284            [
1285                OwnedEvent::weighted(vec![v(1.0)], vec![1.0], 2.0),
1286                OwnedEvent::weighted(vec![v(2.0)], vec![2.0], 3.0),
1287            ],
1288        )
1289        .unwrap();
1290
1291        let mut explicit = BatchAssembler::new(Arc::clone(&schema), 2);
1292        let first = Event {
1293            batch: &source,
1294            row: 0,
1295            weight: source.weights_at(0),
1296        };
1297        explicit.push_borrowed(first, true).unwrap();
1298        let second = Event {
1299            batch: &source,
1300            row: 1,
1301            weight: source.weights_at(1),
1302        };
1303        let transition = explicit.push_borrowed(second, false);
1304        assert!(matches!(
1305            transition,
1306            Err(LadduDataError::InvalidArgument(_))
1307        ));
1308        let explicit = explicit.finish().unwrap();
1309        assert_eq!(explicit.len(), 1);
1310        assert_eq!(explicit.weights_column().unwrap(), &[2.0]);
1311
1312        let unweighted = EventBatch::from_events(
1313            Arc::clone(&schema),
1314            [OwnedEvent::new(vec![v(3.0)], vec![3.0])],
1315        )
1316        .unwrap();
1317        let mut implicit = BatchAssembler::new(schema, 1);
1318        let event = Event {
1319            batch: &unweighted,
1320            row: 0,
1321            weight: unweighted.weights_at(0),
1322        };
1323        implicit.push_borrowed(event, false).unwrap();
1324        let implicit = implicit.finish().unwrap();
1325        assert!(implicit.weights_column().is_none());
1326        assert_eq!(implicit.weights_at(0), 1.0);
1327    }
1328
1329    #[test]
1330    fn assembly_table_covers_column_shapes_and_observable_sharing() {
1331        let cases = [
1332            (
1333                Arc::new(Schema::new(Vec::<&str>::new(), Vec::<&str>::new(), false).unwrap()),
1334                false,
1335            ),
1336            (
1337                Arc::new(Schema::new(["p"], Vec::<&str>::new(), false).unwrap()),
1338                false,
1339            ),
1340            (
1341                Arc::new(Schema::new(Vec::<&str>::new(), ["x"], false).unwrap()),
1342                false,
1343            ),
1344            (
1345                Arc::new(Schema::new(Vec::<&str>::new(), Vec::<&str>::new(), true).unwrap()),
1346                true,
1347            ),
1348            (Arc::new(Schema::new(["p"], ["x"], true).unwrap()), true),
1349        ];
1350
1351        for (schema, weighted) in cases {
1352            let events = (0..2).map(|i| {
1353                let p4s = if schema.n_p4s() == 0 {
1354                    Vec::new()
1355                } else {
1356                    vec![v(i as f64)]
1357                };
1358                let scalars = if schema.n_scalars() == 0 {
1359                    Vec::new()
1360                } else {
1361                    vec![i as f64]
1362                };
1363                if weighted {
1364                    OwnedEvent::weighted(p4s, scalars, 2.0 + i as f64)
1365                } else {
1366                    OwnedEvent::new(p4s, scalars)
1367                }
1368            });
1369            let batch = EventBatch::from_events(Arc::clone(&schema), events).unwrap();
1370            assert_eq!(batch.len(), 2);
1371            assert_eq!(batch.weights_column().is_some(), weighted);
1372
1373            let selected = batch.select(&(0..batch.len()).collect::<Vec<_>>());
1374            assert_eq!(selected.len(), batch.len());
1375            for col in 0..schema.n_p4s() {
1376                assert_eq!(selected.vec4_column(col), batch.vec4_column(col));
1377            }
1378            for col in 0..schema.n_scalars() {
1379                assert_eq!(selected.scalar_column(col), batch.scalar_column(col));
1380            }
1381            assert_eq!(selected.weights_column(), batch.weights_column());
1382
1383            let concatenated = EventBatch::concat(&[batch.slice(0, 1), batch.slice(1, 2)]).unwrap();
1384            assert_eq!(concatenated.len(), batch.len());
1385            for col in 0..schema.n_p4s() {
1386                assert_eq!(concatenated.vec4_column(col), batch.vec4_column(col));
1387            }
1388            for col in 0..schema.n_scalars() {
1389                assert_eq!(concatenated.scalar_column(col), batch.scalar_column(col));
1390            }
1391            assert_eq!(concatenated.weights_column(), batch.weights_column());
1392
1393            if schema.n_p4s() > 0 || schema.n_scalars() > 0 {
1394                let reweighted = batch.reweight(|_, weight| weight + 1.0);
1395                if schema.n_p4s() > 0 {
1396                    assert_eq!(
1397                        reweighted.vec4_column(0).as_ptr(),
1398                        batch.vec4_column(0).as_ptr()
1399                    );
1400                }
1401                if schema.n_scalars() > 0 {
1402                    assert_eq!(
1403                        reweighted.scalar_column(0).as_ptr(),
1404                        batch.scalar_column(0).as_ptr()
1405                    );
1406                }
1407            }
1408        }
1409    }
1410}