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