Skip to main content

laddu_data/data/dataset/
mod.rs

1use std::sync::{
2    Arc, Mutex,
3    atomic::{AtomicU64, Ordering},
4};
5
6use crate::{
7    LadduDataError, LadduDataResult,
8    data::event::{Event, EventBatch, OwnedEvent},
9    io::{EventSink, EventSource, ReadPlan, SourceCapabilities, WritePlan, memory::MemorySource},
10    schema::Schema,
11};
12use laddu_memory::{MemoryBudget, MemoryDecision};
13use num::complex::Complex64;
14
15#[cfg(feature = "parallel")]
16pub mod accurate;
17mod execution;
18mod ops;
19
20use execution::{DatasetExecutionPlan, DatasetExecutor, visit_events};
21use ops::DatasetOp;
22#[cfg(test)]
23use ops::{poisson1_from_hash, uniform_hash_01};
24
25static NEXT_DATASET_IDENTITY: AtomicU64 = AtomicU64::new(1);
26
27fn next_dataset_identity() -> u64 {
28    NEXT_DATASET_IDENTITY.fetch_add(1, Ordering::Relaxed)
29}
30
31/// Cached statistics for an immutable dataset view.
32#[derive(Copy, Clone, Debug, PartialEq)]
33pub struct DatasetStats {
34    events: u64,
35    sum_weights: f64,
36    sum_squared_weights: f64,
37    positive_weights: f64,
38    negative_weights: f64,
39}
40
41impl DatasetStats {
42    /// Returns the number of transformed events.
43    pub fn events(&self) -> u64 {
44        self.events
45    }
46
47    /// Returns the accurately accumulated event-weight sum.
48    pub fn sum_weights(&self) -> f64 {
49        self.sum_weights
50    }
51
52    /// Returns the accurately accumulated sum of squared event weights.
53    pub fn sum_squared_weights(&self) -> f64 {
54        self.sum_squared_weights
55    }
56
57    /// Returns the effective number of entries, when the squared-weight sum is positive.
58    ///
59    /// Empty, all-zero-weight, and other datasets with a non-positive squared-weight sum have
60    /// no defined effective entries and return `None`.
61    pub fn effective_entries(&self) -> Option<f64> {
62        (self.sum_squared_weights > 0.0)
63            .then(|| self.sum_weights * self.sum_weights / self.sum_squared_weights)
64    }
65
66    /// Returns the sum of positive event weights.
67    pub fn positive_weights(&self) -> f64 {
68        self.positive_weights
69    }
70
71    /// Returns the signed sum of negative event weights.
72    pub fn negative_weights(&self) -> f64 {
73        self.negative_weights
74    }
75}
76
77#[derive(Default)]
78struct DatasetStatsCache {
79    events: Option<u64>,
80    sum_weights: Option<f64>,
81    sum_squared_weights: Option<f64>,
82    positive_weights: Option<f64>,
83    negative_weights: Option<f64>,
84}
85
86/// Lazy event dataset combining a source, read plan, and row transformations.
87#[derive(Clone)]
88pub struct Dataset {
89    identity: u64,
90    source: Arc<dyn EventSource>,
91    plan: ReadPlan,
92    ops: Arc<[DatasetOp]>,
93    cache_storage: CacheStorage,
94    memory_policy: MemoryPolicy,
95    memory_budget: MemoryBudget,
96    last_memory_decision: Arc<Mutex<Option<MemoryDecision>>>,
97    stats: Arc<Mutex<DatasetStatsCache>>,
98    source_traversals: Arc<AtomicU64>,
99}
100
101/// Memory policy for compiled event-dependent model caches.
102#[derive(Copy, Clone, Debug, Default, PartialEq, Eq)]
103pub enum CacheStorage {
104    /// Materialize all event-dependent cache values once and retain them for repeated evaluations.
105    #[default]
106    Resident,
107    /// Retain only dataset statistics and rebuild each batch cache during every evaluation.
108    Streaming,
109}
110
111/// Strategy used to trade retained memory for execution speed.
112#[derive(Copy, Clone, Debug, Default, PartialEq, Eq)]
113pub enum MemoryPolicy {
114    /// Select the fastest supported resident or streaming strategy that fits.
115    #[default]
116    Fastest,
117    /// Require the complete compiled event cache to remain resident.
118    Resident,
119    /// Retain no compiled event cache between traversals.
120    Streaming,
121}
122
123impl Dataset {
124    /// Creates a dataset from an event source.
125    pub fn new<S>(source: S) -> Self
126    where
127        S: EventSource + 'static,
128    {
129        Self {
130            identity: next_dataset_identity(),
131            source: Arc::new(source),
132            plan: ReadPlan::default(),
133            ops: Arc::from([]),
134            cache_storage: CacheStorage::Resident,
135            memory_policy: MemoryPolicy::Fastest,
136            memory_budget: MemoryBudget::Auto,
137            last_memory_decision: Default::default(),
138            stats: Default::default(),
139            source_traversals: Default::default(),
140        }
141    }
142
143    /// Creates a dataset from a shared dynamically dispatched source.
144    pub fn from_arc(source: Arc<dyn EventSource>) -> Self {
145        Self {
146            identity: next_dataset_identity(),
147            source,
148            plan: ReadPlan::default(),
149            ops: Arc::from([]),
150            cache_storage: CacheStorage::Resident,
151            memory_policy: MemoryPolicy::Fastest,
152            memory_budget: MemoryBudget::Auto,
153            last_memory_decision: Default::default(),
154            stats: Default::default(),
155            source_traversals: Default::default(),
156        }
157    }
158
159    /// Builds a derived dataset while preserving this dataset's read and cache policy.
160    ///
161    /// The derived source owns its row lifetime; this method only carries over
162    /// execution policy and schema-independent dataset settings.
163    pub fn with_derived_source<S>(&self, source: S) -> Self
164    where
165        S: EventSource + 'static,
166    {
167        Self {
168            identity: next_dataset_identity(),
169            source: Arc::new(source),
170            plan: self.plan,
171            ops: Arc::from([]),
172            cache_storage: self.cache_storage,
173            memory_policy: self.memory_policy,
174            memory_budget: self.memory_budget,
175            last_memory_decision: Default::default(),
176            stats: Default::default(),
177            source_traversals: Default::default(),
178        }
179    }
180
181    /// Builds an empty derived dataset with this dataset's schema and policy.
182    ///
183    /// # Errors
184    ///
185    /// Returns [`LadduDataError`] when the source schema cannot be read.
186    pub fn empty_derived(&self) -> LadduDataResult<Self> {
187        Ok(self.with_derived_source(MemorySource::empty(self.schema()?)))
188    }
189
190    /// Creates an in-memory dataset from one batch.
191    pub fn from_batch(batch: EventBatch) -> Self {
192        Self::new(MemorySource::new(batch))
193    }
194
195    /// Creates an in-memory dataset from schema-compatible batches.
196    ///
197    /// # Errors
198    ///
199    /// Returns [`LadduDataError`] when batch schemas are incompatible.
200    pub fn from_batches(batches: Vec<EventBatch>) -> LadduDataResult<Self> {
201        Ok(Self::new(MemorySource::from_batches(batches)?))
202    }
203
204    /// Collects owned events into an in-memory dataset.
205    ///
206    /// # Errors
207    ///
208    /// Returns [`LadduDataError`] when an event does not match `schema`.
209    pub fn from_events<I>(schema: Arc<Schema>, events: I) -> LadduDataResult<Self>
210    where
211        I: IntoIterator<Item = OwnedEvent>,
212    {
213        Ok(Self::new(MemorySource::from_events(schema, events)?))
214    }
215
216    /// Returns the source schema.
217    ///
218    /// # Errors
219    ///
220    /// Returns [`LadduDataError`] when the underlying source cannot determine
221    /// or load its schema.
222    pub fn schema(&self) -> LadduDataResult<Arc<Schema>> {
223        self.source.schema()
224    }
225
226    /// Returns source planning capabilities.
227    pub fn capabilities(&self) -> SourceCapabilities {
228        self.source.capabilities()
229    }
230
231    /// Returns the source event count when cheaply available.
232    ///
233    /// # Errors
234    ///
235    /// Returns an error when source metadata cannot be read.
236    pub fn num_events(&self) -> LadduDataResult<Option<u64>> {
237        {
238            let stats = self.stats.lock().unwrap_or_else(|error| error.into_inner());
239            if let Some(events) = stats.events {
240                return Ok(Some(events));
241            }
242        }
243
244        if self
245            .ops
246            .iter()
247            .any(|op| !matches!(op, DatasetOp::Bootstrap { .. }))
248        {
249            return Ok(None);
250        }
251
252        let events = self.source.num_events()?;
253        if let Some(events) = events {
254            self.stats
255                .lock()
256                .unwrap_or_else(|error| error.into_inner())
257                .events = Some(events);
258        }
259        Ok(events)
260    }
261
262    /// Returns cached event-count and weight-sum statistics, computing them once if needed.
263    ///
264    /// Clones of a dataset view share this cache. Failed traversals are not cached and may be
265    /// retried.
266    ///
267    /// # Errors
268    ///
269    /// Returns [`LadduDataError`] when reading or transforming the dataset fails.
270    pub fn stats(&self) -> LadduDataResult<DatasetStats> {
271        {
272            let cache = self.stats.lock().unwrap_or_else(|error| error.into_inner());
273            if let (
274                Some(events),
275                Some(sum_weights),
276                Some(sum_squared_weights),
277                Some(positive_weights),
278                Some(negative_weights),
279            ) = (
280                cache.events,
281                cache.sum_weights,
282                cache.sum_squared_weights,
283                cache.positive_weights,
284                cache.negative_weights,
285            ) {
286                return Ok(DatasetStats {
287                    events,
288                    sum_weights,
289                    sum_squared_weights,
290                    positive_weights,
291                    negative_weights,
292                });
293            }
294        }
295
296        let mut executor = self.executor_with_plan(self.plan)?;
297        for batch in &mut executor {
298            batch?;
299        }
300
301        Ok(executor.stats())
302    }
303
304    /// Returns the current read plan.
305    pub fn read_plan(&self) -> ReadPlan {
306        self.plan
307    }
308
309    /// Returns the compiled-cache memory policy.
310    pub fn cache_storage(&self) -> CacheStorage {
311        self.cache_storage
312    }
313
314    /// Returns the memory-first cache selection policy.
315    pub fn memory_policy(&self) -> MemoryPolicy {
316        self.memory_policy
317    }
318
319    /// Returns the dataset's host-memory budget.
320    pub fn memory_budget(&self) -> MemoryBudget {
321        self.memory_budget
322    }
323
324    /// Returns the most recent memory-derived read decision.
325    pub fn last_memory_decision(&self) -> Option<MemoryDecision> {
326        self.last_memory_decision
327            .lock()
328            .unwrap_or_else(|error| error.into_inner())
329            .clone()
330    }
331
332    /// Returns the number of transformed source iterators opened by this dataset view.
333    pub fn source_traversals(&self) -> u64 {
334        self.source_traversals.load(Ordering::Relaxed)
335    }
336
337    /// Returns the immutable identity of this dataset view.
338    ///
339    /// Clones retain identity, while row, weight, and source transformations
340    /// create a new identity. This is intended for execution-scoped caches.
341    #[doc(hidden)]
342    pub fn identity(&self) -> u64 {
343        self.identity
344    }
345
346    /// Returns this dataset with a host-memory budget.
347    pub fn with_memory_budget(mut self, budget: MemoryBudget) -> Self {
348        self.memory_budget = budget;
349        self
350    }
351
352    /// Select the fastest strategy allowed by the active memory budget.
353    pub fn fastest(mut self) -> Self {
354        self.memory_policy = MemoryPolicy::Fastest;
355        self.cache_storage = CacheStorage::Resident;
356        self
357    }
358
359    /// Retain compiled event-dependent values for every local event.
360    ///
361    /// This strict policy is intended for repeatedly evaluating a likelihood
362    /// and returns an error if the resident cache cannot fit. [`Dataset::fastest`]
363    /// is the default.
364    pub fn resident(mut self) -> Self {
365        self.memory_policy = MemoryPolicy::Resident;
366        self.cache_storage = CacheStorage::Resident;
367        self
368    }
369
370    /// Re-read the source and rebuild one batch cache during each parameter evaluation.
371    ///
372    /// Only fixed dataset statistics are retained. This minimizes memory use but requires a
373    /// repeatable source and is expected to be slower than [`Dataset::resident`].
374    pub fn streaming(mut self) -> Self {
375        self.memory_policy = MemoryPolicy::Streaming;
376        self.cache_storage = CacheStorage::Streaming;
377        self
378    }
379
380    /// Returns this dataset with a low-level nonzero maximum event count.
381    ///
382    /// Prefer [`Dataset::with_memory_budget`] for portable application code.
383    /// This explicit read-plan override remains available for source debugging
384    /// and reproducibility and is always capped by execution memory planning.
385    ///
386    /// # Errors
387    ///
388    /// Returns [`LadduDataError::InvalidArgument`] when `chunk_size` is zero.
389    pub fn chunked(mut self, chunk_size: usize) -> LadduDataResult<Self> {
390        if chunk_size == 0 {
391            return Err(LadduDataError::InvalidArgument(
392                "chunk_size must be nonzero",
393            ));
394        }
395        self.plan.chunk_size = Some(chunk_size);
396        Ok(self)
397    }
398
399    /// Returns this dataset with source-native batch sizes.
400    pub fn unchunked(mut self) -> Self {
401        self.plan.chunk_size = None;
402        self
403    }
404
405    /// Lazily retains events satisfying `f`.
406    pub fn filter<F>(self, f: F) -> Self
407    where
408        F: Fn(Event<'_>) -> bool + Send + Sync + 'static,
409    {
410        self.push_op(DatasetOp::Filter(Arc::new(f)))
411    }
412
413    /// Lazily retains a deterministic fraction of events.
414    ///
415    /// # Errors
416    ///
417    /// Returns [`LadduDataError::InvalidArgument`] when `fraction` is outside
418    /// `[0, 1]` or is NaN.
419    pub fn subsample(self, fraction: f64, seed: u64) -> LadduDataResult<Self> {
420        if !(0.0..=1.0).contains(&fraction) {
421            return Err(LadduDataError::InvalidArgument(
422                "fraction must be in [0, 1]",
423            ));
424        }
425
426        Ok(self.push_op(DatasetOp::Subsample { fraction, seed }))
427    }
428
429    /// Applies deterministic Poisson bootstrap multiplicities to event weights.
430    pub fn bootstrap(self, seed: u64) -> Self {
431        self.push_op(DatasetOp::Bootstrap { seed })
432    }
433
434    /// Visits each transformed event.
435    ///
436    /// # Errors
437    ///
438    /// Returns [`LadduDataError`] when reading or transforming the source
439    /// fails.
440    pub fn for_each_event<F>(&self, mut f: F) -> LadduDataResult<()>
441    where
442        F: FnMut(Event<'_>),
443    {
444        self.try_for_each_event(|ev| {
445            f(ev);
446            Ok(())
447        })
448    }
449
450    /// Visits each transformed event and stops at the first error.
451    ///
452    /// # Errors
453    ///
454    /// Returns the first [`LadduDataError`] produced by the source,
455    /// transformations, or callback.
456    pub fn try_for_each_event<F>(&self, mut f: F) -> LadduDataResult<()>
457    where
458        F: FnMut(Event<'_>) -> LadduDataResult<()>,
459    {
460        visit_events(
461            self,
462            DatasetExecutionPlan::resolve(self, self.plan)?,
463            &mut f,
464        )
465    }
466
467    /// Fallibly maps transformed events into a vector.
468    ///
469    /// # Errors
470    ///
471    /// Returns the first [`LadduDataError`] produced while reading,
472    /// transforming, or mapping an event.
473    pub fn try_map_events<T, F>(&self, mut f: F) -> LadduDataResult<Vec<T>>
474    where
475        F: FnMut(Event<'_>) -> LadduDataResult<T>,
476    {
477        let mut out = Vec::new();
478        self.try_for_each_event(|ev| {
479            out.push(f(ev)?);
480            Ok(())
481        })?;
482
483        Ok(out)
484    }
485
486    /// Maps transformed events into a vector.
487    ///
488    /// # Errors
489    ///
490    /// Returns [`LadduDataError`] when reading or transforming the source
491    /// fails.
492    pub fn map_events<T, F>(&self, mut f: F) -> LadduDataResult<Vec<T>>
493    where
494        F: FnMut(Event<'_>) -> T,
495    {
496        let mut out = Vec::new();
497        self.try_for_each_event(|ev| {
498            out.push(f(ev));
499            Ok(())
500        })?;
501
502        Ok(out)
503    }
504
505    /// Fallibly folds transformed events into an owned accumulator.
506    ///
507    /// # Errors
508    ///
509    /// Returns the first [`LadduDataError`] produced while reading,
510    /// transforming, or folding an event.
511    pub fn try_fold_events<T, F>(&self, init: T, mut f: F) -> LadduDataResult<T>
512    where
513        F: FnMut(T, Event<'_>) -> LadduDataResult<T>,
514    {
515        let mut acc = Some(init);
516
517        self.try_for_each_event(|ev| {
518            let current = acc.take().ok_or_else(|| {
519                LadduDataError::Source("dataset fold accumulator was consumed".into())
520            })?;
521            acc = Some(f(current, ev)?);
522            Ok(())
523        })?;
524
525        acc.ok_or_else(|| LadduDataError::Source("dataset fold produced no accumulator".into()))
526    }
527
528    /// Folds transformed events into an owned accumulator.
529    ///
530    /// # Errors
531    ///
532    /// Returns [`LadduDataError`] when reading or transforming the source
533    /// fails.
534    pub fn fold_events<T, F>(&self, init: T, mut f: F) -> LadduDataResult<T>
535    where
536        F: FnMut(T, Event<'_>) -> T,
537    {
538        self.try_fold_events(init, |acc, ev| Ok(f(acc, ev)))
539    }
540
541    /// Fallibly folds transformed event batches into an owned accumulator.
542    ///
543    /// The callback receives each batch in source order and may stop the
544    /// traversal by returning a data error.
545    ///
546    /// # Errors
547    ///
548    /// Returns the first source or callback error.
549    pub fn try_fold_batches<T, F>(&self, init: T, mut f: F) -> LadduDataResult<T>
550    where
551        F: FnMut(T, EventBatch) -> LadduDataResult<T>,
552    {
553        let mut acc = init;
554        for batch in self.batches()? {
555            acc = f(acc, batch?)?;
556        }
557        Ok(acc)
558    }
559
560    /// Fallibly updates a mutable accumulator for every transformed event.
561    ///
562    /// # Errors
563    ///
564    /// Returns the first [`LadduDataError`] produced while reading,
565    /// transforming, or accumulating an event.
566    pub fn try_accumulate_events<T, F>(&self, mut acc: T, mut f: F) -> LadduDataResult<T>
567    where
568        F: FnMut(&mut T, Event<'_>) -> LadduDataResult<()>,
569    {
570        self.try_for_each_event(|ev| f(&mut acc, ev))?;
571        Ok(acc)
572    }
573
574    /// Updates a mutable accumulator for every transformed event.
575    ///
576    /// # Errors
577    ///
578    /// Returns [`LadduDataError`] when reading or transforming the source
579    /// fails.
580    pub fn accumulate_events<T, F>(&self, acc: T, mut f: F) -> LadduDataResult<T>
581    where
582        F: FnMut(&mut T, Event<'_>),
583    {
584        self.try_accumulate_events(acc, |acc, ev| {
585            f(acc, ev);
586            Ok(())
587        })
588    }
589
590    /// Sums effective event weights.
591    ///
592    /// # Errors
593    ///
594    /// Returns [`LadduDataError`] when reading or transforming the source
595    /// fails.
596    pub fn sum_weights(&self) -> LadduDataResult<f64> {
597        Ok(self.stats()?.sum_weights())
598    }
599
600    /// Sums `weight * f(event)` over transformed events.
601    ///
602    /// # Errors
603    ///
604    /// Returns [`LadduDataError`] when reading or transforming the source
605    /// fails.
606    pub fn weighted_sum<F>(&self, mut f: F) -> LadduDataResult<f64>
607    where
608        F: FnMut(Event<'_>) -> f64,
609    {
610        self.fold_events(0.0, |sum, ev| sum + ev.weight() * f(ev))
611    }
612
613    /// Sums complex `weight * f(event)` contributions.
614    ///
615    /// # Errors
616    ///
617    /// Returns [`LadduDataError`] when reading or transforming the source
618    /// fails.
619    pub fn weighted_complex_sum<F>(&self, mut f: F) -> LadduDataResult<Complex64>
620    where
621        F: FnMut(Event<'_>) -> Complex64,
622    {
623        self.fold_events(0.0.into(), |sum, ev| sum + ev.weight() * f(ev))
624    }
625
626    /// Opens an iterator of fully transformed event batches.
627    ///
628    /// # Errors
629    ///
630    /// Returns [`LadduDataError`] when the underlying source cannot initialize
631    /// a batch stream for the current plan.
632    pub fn batches(
633        &self,
634    ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
635        self.stream_with_plan(self.plan)
636    }
637
638    #[doc(hidden)]
639    /// Opens the shared transformed batch stream using an explicit read plan.
640    ///
641    /// # Errors
642    ///
643    /// Returns [`LadduDataError`] when the underlying source cannot initialize
644    /// the requested batch stream.
645    pub fn stream_with_plan(
646        &self,
647        plan: ReadPlan,
648    ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
649        Ok(Box::new(self.executor_with_plan(plan)?))
650    }
651
652    #[doc(hidden)]
653    /// Compatibility alias for the shared transformed batch stream.
654    pub fn batches_with_plan(
655        &self,
656        plan: ReadPlan,
657    ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
658        self.stream_with_plan(plan)
659    }
660
661    fn executor_with_plan(&self, plan: ReadPlan) -> LadduDataResult<DatasetExecutor> {
662        DatasetExecutor::new(self, DatasetExecutionPlan::resolve(self, plan)?)
663    }
664
665    /// Visits each transformed batch and stops at the first error.
666    ///
667    /// # Errors
668    ///
669    /// Returns the first [`LadduDataError`] produced by the source,
670    /// transformations, or callback.
671    pub fn try_for_each_batch<F>(&self, mut f: F) -> LadduDataResult<()>
672    where
673        F: FnMut(EventBatch) -> LadduDataResult<()>,
674    {
675        for batch in self.batches()? {
676            f(batch?)?;
677        }
678
679        Ok(())
680    }
681
682    /// Maps transformed batches into a vector.
683    ///
684    /// # Errors
685    ///
686    /// Returns [`LadduDataError`] when reading or transforming a batch fails.
687    pub fn map_batches<T, F>(&self, mut f: F) -> LadduDataResult<Vec<T>>
688    where
689        F: FnMut(EventBatch) -> T,
690    {
691        let mut out = Vec::new();
692
693        self.try_for_each_batch(|batch| {
694            out.push(f(batch));
695            Ok(())
696        })?;
697
698        Ok(out)
699    }
700
701    /// Streams the transformed dataset into an event sink.
702    ///
703    /// # Errors
704    ///
705    /// Returns [`LadduDataError`] when reading, transforming, or writing a
706    /// batch fails, or the sink cannot begin or finish the stream.
707    pub fn write_to<S: EventSink>(&self, sink: &mut S) -> LadduDataResult<()> {
708        sink.begin(self.schema()?, WritePlan::from(self.plan))?;
709
710        let result = (|| {
711            for batch in self.batches()? {
712                sink.write_batch(&batch?)?;
713            }
714
715            sink.finish()
716        })();
717
718        if result.is_err() {
719            // Preserve the operation error; abort is best-effort cleanup and
720            // may itself report a backend failure.
721            let _ = sink.abort();
722        }
723
724        result
725    }
726
727    fn push_op(self, op: DatasetOp) -> Self {
728        let preserved_events = if matches!(&op, DatasetOp::Bootstrap { .. }) {
729            self.num_events().ok().flatten()
730        } else {
731            None
732        };
733        let mut ops = self.ops.to_vec();
734        ops.push(op);
735
736        Self {
737            identity: next_dataset_identity(),
738            source: self.source,
739            plan: self.plan,
740            ops: ops.into(),
741            cache_storage: self.cache_storage,
742            memory_policy: self.memory_policy,
743            memory_budget: self.memory_budget,
744            last_memory_decision: Default::default(),
745            stats: Arc::new(Mutex::new(DatasetStatsCache {
746                events: preserved_events,
747                sum_weights: None,
748                sum_squared_weights: None,
749                positive_weights: None,
750                negative_weights: None,
751            })),
752            source_traversals: Default::default(),
753        }
754    }
755}
756
757#[cfg(test)]
758mod tests {
759    use super::ops::materialize_batch;
760    use super::*;
761    use crate::io::{EventBatchIter, EventSource, ReadPlan, memory::MemorySink};
762    use laddu_physics::vectors::RealVec4;
763    use std::sync::atomic::{AtomicUsize, Ordering};
764
765    #[derive(Clone)]
766    struct CountingSource {
767        batch: EventBatch,
768        reads: Arc<AtomicUsize>,
769    }
770
771    impl EventSource for CountingSource {
772        fn schema(&self) -> LadduDataResult<Arc<Schema>> {
773            Ok(Arc::clone(self.batch.schema()))
774        }
775
776        fn batches(
777            &self,
778            _plan: ReadPlan,
779        ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
780            self.reads.fetch_add(1, Ordering::Relaxed);
781            Ok(Box::new(std::iter::once(Ok(self.batch.clone()))))
782        }
783    }
784
785    #[derive(Clone)]
786    struct ErrorSource {
787        schema: Arc<Schema>,
788        items: Arc<[LadduDataResult<EventBatch>]>,
789    }
790
791    impl EventSource for ErrorSource {
792        fn schema(&self) -> LadduDataResult<Arc<Schema>> {
793            Ok(Arc::clone(&self.schema))
794        }
795
796        fn batches(&self, _plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
797            let items = Arc::clone(&self.items);
798            Ok(Box::new(
799                (0..items.len()).map(move |index| items[index].clone()),
800            ))
801        }
802    }
803
804    fn v(x: f64) -> RealVec4 {
805        RealVec4 {
806            e: x + 0.3,
807            px: x,
808            py: x + 0.1,
809            pz: x + 0.2,
810        }
811    }
812
813    fn schema_with_weight() -> Arc<Schema> {
814        Arc::new(Schema::new(["p"], ["x"], true).unwrap())
815    }
816
817    fn schema_without_weight() -> Arc<Schema> {
818        Arc::new(Schema::new(["p"], ["x"], false).unwrap())
819    }
820
821    fn weighted_batch(start: usize, len: usize) -> EventBatch {
822        let schema = schema_with_weight();
823
824        let events = (start..start + len)
825            .map(|i| OwnedEvent::weighted(vec![v(i as f64)], vec![i as f64], 10.0 + i as f64));
826
827        EventBatch::from_events(schema, events).unwrap()
828    }
829
830    fn unweighted_batch(start: usize, len: usize) -> EventBatch {
831        let schema = schema_without_weight();
832
833        let events =
834            (start..start + len).map(|i| OwnedEvent::new(vec![v(i as f64)], vec![i as f64]));
835
836        EventBatch::from_events(schema, events).unwrap()
837    }
838
839    fn error_source(error_after: Option<EventBatch>) -> ErrorSource {
840        let schema = schema_with_weight();
841        let mut items = Vec::new();
842        if let Some(batch) = error_after {
843            items.push(Ok(batch));
844        }
845        items.push(Err(LadduDataError::Unsupported("source")));
846        ErrorSource {
847            schema,
848            items: items.into(),
849        }
850    }
851
852    fn scalar_values(batch: &EventBatch) -> Vec<f64> {
853        batch.scalar_column(0).to_vec()
854    }
855
856    fn stats_dataset(weights: &[f64]) -> Dataset {
857        let schema = schema_with_weight();
858        let events = weights.iter().enumerate().map(|(index, &weight)| {
859            OwnedEvent::weighted(vec![v(index as f64)], vec![index as f64], weight)
860        });
861        Dataset::from_batch(EventBatch::from_events(schema, events).unwrap())
862    }
863
864    #[test]
865    fn dataset_statistics_are_shared_per_view_and_invalidated_by_selection() {
866        let reads = Arc::new(AtomicUsize::new(0));
867        let dataset = Dataset::new(CountingSource {
868            batch: weighted_batch(0, 5),
869            reads: Arc::clone(&reads),
870        });
871        let clone = dataset.clone();
872
873        assert_eq!(dataset.num_events().unwrap(), None);
874        assert_eq!(dataset.stats().unwrap().events(), 5);
875        assert_eq!(clone.sum_weights().unwrap(), 60.0);
876        let stats = clone.stats().unwrap();
877        assert_eq!(stats.sum_squared_weights(), 730.0);
878        assert_eq!(stats.positive_weights(), 60.0);
879        assert_eq!(stats.negative_weights(), 0.0);
880        assert_eq!(stats.effective_entries(), Some(3600.0 / 730.0));
881        assert_eq!(clone.num_events().unwrap(), Some(5));
882        assert_eq!(reads.load(Ordering::Relaxed), 1);
883
884        let bootstrapped = dataset.clone().bootstrap(7);
885        assert_eq!(bootstrapped.num_events().unwrap(), Some(5));
886        assert_eq!(reads.load(Ordering::Relaxed), 1);
887
888        let filtered = dataset.filter(|event| event.scalar(0) >= 2.0);
889        assert_eq!(filtered.num_events().unwrap(), None);
890        assert_eq!(
891            filtered.map_events(|event| event.scalar(0)).unwrap(),
892            [2.0, 3.0, 4.0]
893        );
894        assert_eq!(reads.load(Ordering::Relaxed), 2);
895        assert_eq!(filtered.stats().unwrap().events(), 3);
896        assert_eq!(reads.load(Ordering::Relaxed), 2);
897    }
898
899    #[test]
900    fn dataset_statistics_preserve_signed_and_squared_weight_diagnostics() {
901        let stats = stats_dataset(&[2.0, -1.0, 3.0, -4.0]).stats().unwrap();
902
903        assert_eq!(stats.events(), 4);
904        assert_eq!(stats.sum_weights(), 0.0);
905        assert_eq!(stats.sum_squared_weights(), 30.0);
906        assert_eq!(stats.positive_weights(), 5.0);
907        assert_eq!(stats.negative_weights(), -5.0);
908        assert_eq!(stats.effective_entries(), Some(0.0));
909    }
910
911    #[test]
912    fn dataset_statistics_make_effective_entries_unavailable_without_squared_weight() {
913        let zero = stats_dataset(&[0.0, 0.0]).stats().unwrap();
914        assert_eq!(zero.events(), 2);
915        assert_eq!(zero.sum_weights(), 0.0);
916        assert_eq!(zero.sum_squared_weights(), 0.0);
917        assert_eq!(zero.positive_weights(), 0.0);
918        assert_eq!(zero.negative_weights(), 0.0);
919        assert_eq!(zero.effective_entries(), None);
920
921        let empty = stats_dataset(&[1.0])
922            .empty_derived()
923            .unwrap()
924            .stats()
925            .unwrap();
926        assert_eq!(empty.events(), 0);
927        assert_eq!(empty.sum_weights(), 0.0);
928        assert_eq!(empty.sum_squared_weights(), 0.0);
929        assert_eq!(empty.effective_entries(), None);
930    }
931
932    #[test]
933    fn dataset_statistics_match_across_resident_streaming_and_chunked_views() {
934        let resident = Dataset::from_batch(weighted_batch(0, 10));
935        let streaming = Dataset::new(CountingSource {
936            batch: weighted_batch(0, 10),
937            reads: Arc::new(AtomicUsize::new(0)),
938        });
939        let chunked = Dataset::from_batches(
940            (0..10)
941                .map(|index| weighted_batch(index, 1))
942                .collect::<Vec<_>>(),
943        )
944        .unwrap()
945        .chunked(3)
946        .unwrap();
947
948        let expected = resident.stats().unwrap();
949        assert_eq!(streaming.stats().unwrap(), expected);
950        assert_eq!(chunked.stats().unwrap(), expected);
951    }
952
953    #[test]
954    fn transformed_fragments_are_coalesced_to_the_read_chunk_size() {
955        let fragments = (0..10)
956            .map(|index| weighted_batch(index, 1))
957            .collect::<Vec<_>>();
958        let dataset = Dataset::from_batches(fragments)
959            .unwrap()
960            .filter(|_| true)
961            .chunked(4)
962            .unwrap();
963
964        let batches = dataset
965            .batches()
966            .unwrap()
967            .collect::<LadduDataResult<Vec<_>>>()
968            .unwrap();
969        assert_eq!(
970            batches.iter().map(EventBatch::len).collect::<Vec<_>>(),
971            [4, 4, 2]
972        );
973        assert_eq!(
974            batches
975                .iter()
976                .flat_map(|batch| batch.scalar_column(0).iter().copied())
977                .collect::<Vec<_>>(),
978            (0..10).map(|value| value as f64).collect::<Vec<_>>()
979        );
980    }
981
982    #[test]
983    fn shared_stream_preserves_pending_batches_before_source_errors() {
984        let dataset = Dataset::new(error_source(Some(weighted_batch(0, 2))))
985            .chunked(4)
986            .unwrap();
987        let mut batches = dataset.batches().unwrap();
988
989        assert_eq!(batches.next().unwrap().unwrap().len(), 2);
990        assert!(matches!(
991            batches.next().unwrap(),
992            Err(LadduDataError::Unsupported("source"))
993        ));
994        assert!(batches.next().is_none());
995
996        assert!(matches!(
997            dataset.stats(),
998            Err(LadduDataError::Unsupported("source"))
999        ));
1000        assert_eq!(dataset.source_traversals(), 2);
1001    }
1002
1003    #[test]
1004    fn event_visitors_share_the_execution_plan_without_changing_source_rows() {
1005        let dataset = Dataset::new(error_source(Some(weighted_batch(0, 2))))
1006            .chunked(4)
1007            .unwrap()
1008            .filter(|event| event.scalar(0) >= 0.0);
1009        let mut rows = Vec::new();
1010
1011        let error = dataset
1012            .try_for_each_event(|event| {
1013                rows.push(event.row());
1014                Ok(())
1015            })
1016            .unwrap_err();
1017
1018        assert!(matches!(error, LadduDataError::Unsupported("source")));
1019        assert_eq!(rows, [0, 1]);
1020    }
1021
1022    #[test]
1023    fn dataset_map_fold_accumulate_complex_sum_and_error_paths_use_transformed_events() {
1024        let dataset =
1025            Dataset::from_batch(weighted_batch(0, 5)).filter(|ev| ev.scalar(0) % 2.0 == 0.0);
1026
1027        let rows = dataset
1028            .map_events(|ev| (ev.row(), ev.scalar(0), ev.weight()))
1029            .unwrap();
1030
1031        assert_eq!(rows, vec![(0, 0.0, 10.0), (2, 2.0, 12.0), (4, 4.0, 14.0)]);
1032
1033        let folded = dataset
1034            .fold_events(String::new(), |mut out, ev| {
1035                out.push_str(&format!("{};", ev.scalar(0)));
1036                out
1037            })
1038            .unwrap();
1039
1040        assert_eq!(folded, "0;2;4;");
1041
1042        let accumulated = dataset
1043            .accumulate_events(Vec::<f64>::new(), |values, ev| values.push(ev.weight()))
1044            .unwrap();
1045
1046        assert_eq!(accumulated, vec![10.0, 12.0, 14.0]);
1047
1048        let weighted_sum = dataset.weighted_sum(|ev| ev.scalar(0)).unwrap();
1049        assert_eq!(weighted_sum, 0.0 * 10.0 + 2.0 * 12.0 + 4.0 * 14.0);
1050
1051        let complex_sum = dataset
1052            .weighted_complex_sum(|ev| Complex64::new(ev.scalar(0), 1.0))
1053            .unwrap();
1054
1055        assert_eq!(complex_sum.re, weighted_sum);
1056        assert_eq!(complex_sum.im, 10.0 + 12.0 + 14.0);
1057
1058        let err = dataset
1059            .try_map_events(|ev| {
1060                if ev.scalar(0) == 2.0 {
1061                    Err(LadduDataError::Unsupported("stop"))
1062                } else {
1063                    Ok(ev.scalar(0))
1064                }
1065            })
1066            .unwrap_err();
1067
1068        assert!(matches!(err, LadduDataError::Unsupported("stop")));
1069    }
1070
1071    #[test]
1072    fn batch_folds_and_empty_derived_sources_preserve_schema_and_errors() {
1073        let dataset = Dataset::from_batches(vec![weighted_batch(0, 2), weighted_batch(2, 2)])
1074            .unwrap()
1075            .chunked(2)
1076            .unwrap();
1077        let event_count = dataset
1078            .try_fold_batches(0usize, |count, batch| Ok(count + batch.len()))
1079            .unwrap();
1080        assert_eq!(event_count, 4);
1081        let rows = dataset
1082            .try_fold_batches(Vec::new(), |mut rows, batch| {
1083                rows.extend((0..batch.len()).map(|row| batch.scalar_at(0, row)));
1084                Ok(rows)
1085            })
1086            .unwrap();
1087        assert_eq!(rows, [0.0, 1.0, 2.0, 3.0]);
1088
1089        let error = dataset
1090            .try_fold_batches(0usize, |_count, _batch| {
1091                Err(LadduDataError::Unsupported("stop"))
1092            })
1093            .unwrap_err();
1094        assert!(matches!(error, LadduDataError::Unsupported("stop")));
1095
1096        let empty = dataset.empty_derived().unwrap();
1097        assert_eq!(
1098            empty.schema().unwrap().as_ref(),
1099            dataset.schema().unwrap().as_ref()
1100        );
1101        assert_eq!(empty.num_events().unwrap(), Some(0));
1102        assert!(empty.batches().unwrap().next().is_none());
1103    }
1104
1105    #[test]
1106    fn deterministic_subsample_and_bootstrap_use_global_event_ids_across_batches() {
1107        let seed = 0x0BAD_5EED;
1108        let bootstrap_seed = 0xB007_57A9;
1109
1110        let dataset = Dataset::from_batches(vec![weighted_batch(0, 3), weighted_batch(3, 3)])
1111            .unwrap()
1112            .subsample(0.5, seed)
1113            .unwrap()
1114            .bootstrap(bootstrap_seed);
1115
1116        let observed = dataset
1117            .map_events(|ev| (ev.scalar(0) as u64, ev.weight()))
1118            .unwrap();
1119
1120        let expected: Vec<(u64, f64)> = (0_u64..6)
1121            .filter(|&event_id| uniform_hash_01(seed, event_id) < 0.5)
1122            .map(|event_id| {
1123                let original_weight = 10.0 + event_id as f64;
1124                let bootstrap_weight =
1125                    poisson1_from_hash(bootstrap_seed, event_id) as f64 * original_weight;
1126                (event_id, bootstrap_weight)
1127            })
1128            .collect();
1129
1130        assert_eq!(observed, expected);
1131    }
1132
1133    #[test]
1134    fn materialized_batches_store_weights_only_when_needed() {
1135        let unweighted = unweighted_batch(0, 4);
1136
1137        let filtered = Dataset::from_batch(unweighted.clone())
1138            .filter(|ev| ev.scalar(0) >= 1.0)
1139            .subsample(1.0, 123)
1140            .unwrap();
1141
1142        let filtered_batch = filtered.batches().unwrap().next().unwrap().unwrap();
1143
1144        assert_eq!(scalar_values(&filtered_batch), vec![1.0, 2.0, 3.0]);
1145        assert!(filtered_batch.weights_column().is_none());
1146
1147        let bootstrapped = Dataset::from_batch(unweighted).bootstrap(999);
1148        let bootstrapped_batch = bootstrapped.batches().unwrap().next().unwrap().unwrap();
1149
1150        assert!(bootstrapped_batch.weights_column().is_some());
1151
1152        let source = weighted_batch(0, 2);
1153        let empty_weighted =
1154            materialize_batch(&source, &[DatasetOp::Filter(Arc::new(|_| false))], 0).unwrap();
1155        assert!(empty_weighted.is_empty());
1156        assert_eq!(empty_weighted.weights_column(), Some([].as_slice()));
1157    }
1158
1159    #[test]
1160    fn write_to_memory_sink_captures_transformed_dataset() {
1161        let dataset = Dataset::from_batch(weighted_batch(0, 5)).filter(|ev| ev.scalar(0) >= 2.0);
1162
1163        let mut sink = MemorySink::new();
1164        dataset.write_to(&mut sink).unwrap();
1165
1166        let captured = sink.into_batch().unwrap();
1167
1168        assert_eq!(scalar_values(&captured), vec![2.0, 3.0, 4.0]);
1169        assert_eq!(captured.weights_column().unwrap(), &[12.0, 13.0, 14.0]);
1170    }
1171
1172    #[test]
1173    fn immutable_dataset_identity_tracks_semantic_views() {
1174        let dataset = Dataset::from_batch(weighted_batch(0, 3));
1175        assert_eq!(dataset.identity(), dataset.clone().identity());
1176        assert_eq!(dataset.identity(), dataset.clone().streaming().identity());
1177        assert_ne!(
1178            dataset.identity(),
1179            dataset.clone().subsample(1.0, 7).unwrap().identity()
1180        );
1181        assert_ne!(dataset.identity(), dataset.clone().bootstrap(7).identity());
1182        assert_ne!(
1183            dataset.identity(),
1184            dataset.clone().filter(|_| true).identity()
1185        );
1186    }
1187}