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