Skip to main content

laddu_runtime/
query.rs

1use std::sync::Arc;
2
3use laddu_compile::CompiledQuery;
4use laddu_data::{
5    LadduDataError, LadduDataResult,
6    data::{Dataset, EventBatch},
7    io::{EventBatchIter, EventSource, ReadPlan, SourceCapabilities, memory::MemorySource},
8    schema::Schema,
9};
10use laddu_expr::{Expr, ValueKind};
11use laddu_physics::{
12    binning::{BinningAxis, FinalUpperEdge},
13    histogram::Histogram,
14    joint_histogram::JointHistogram,
15};
16use num::complex::Complex64;
17use serde::{Deserialize, Deserializer, Serialize};
18
19use crate::{Execution, PreparedModel, RuntimeError, RuntimeResult};
20
21/// Comparison operation used by a dataset predicate.
22#[derive(Copy, Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
23pub enum Comparison {
24    /// Less than.
25    Lt,
26    /// Less than or equal to.
27    Le,
28    /// Greater than.
29    Gt,
30    /// Greater than or equal to.
31    Ge,
32    /// Equal to.
33    Eq,
34    /// Not equal to.
35    Ne,
36}
37
38/// Determines which endpoints are included by an interval predicate.
39#[derive(Copy, Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
40pub enum IntervalClosure {
41    /// Exclude both endpoints.
42    Open,
43    /// Include only the lower endpoint.
44    LeftClosed,
45    /// Include only the upper endpoint.
46    RightClosed,
47    /// Include both endpoints.
48    #[default]
49    Closed,
50}
51
52/// Boolean expression used to select events from a dataset.
53#[derive(Clone, Debug)]
54pub enum Predicate {
55    /// Compare two scalar expressions.
56    Compare {
57        /// Left-hand expression.
58        lhs: Expr,
59        /// Comparison operation.
60        op: Comparison,
61        /// Right-hand expression.
62        rhs: Expr,
63    },
64    /// Require both child predicates to hold.
65    And(Box<Self>, Box<Self>),
66    /// Require either child predicate to hold.
67    Or(Box<Self>, Box<Self>),
68    /// Negate a predicate.
69    Not(Box<Self>),
70    /// Test whether a value lies between two bounds.
71    Between {
72        /// Expression whose value is tested.
73        value: Expr,
74        /// Lower bound expression.
75        lower: Expr,
76        /// Upper bound expression.
77        upper: Expr,
78        /// Endpoint inclusion policy.
79        closure: IntervalClosure,
80    },
81}
82
83impl Predicate {
84    /// Creates a comparison predicate.
85    pub fn compare(lhs: impl Into<Expr>, op: Comparison, rhs: impl Into<Expr>) -> Self {
86        Self::Compare {
87            lhs: lhs.into(),
88            op,
89            rhs: rhs.into(),
90        }
91    }
92
93    /// Creates a less-than predicate.
94    pub fn lt(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Self {
95        Self::compare(lhs, Comparison::Lt, rhs)
96    }
97    /// Creates a less-than-or-equal predicate.
98    pub fn le(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Self {
99        Self::compare(lhs, Comparison::Le, rhs)
100    }
101    /// Creates a greater-than predicate.
102    pub fn gt(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Self {
103        Self::compare(lhs, Comparison::Gt, rhs)
104    }
105    /// Creates a greater-than-or-equal predicate.
106    pub fn ge(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Self {
107        Self::compare(lhs, Comparison::Ge, rhs)
108    }
109    /// Creates an equality predicate.
110    pub fn eq(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Self {
111        Self::compare(lhs, Comparison::Eq, rhs)
112    }
113    /// Creates an inequality predicate.
114    pub fn ne(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Self {
115        Self::compare(lhs, Comparison::Ne, rhs)
116    }
117    /// Combines this predicate with `rhs` using logical AND.
118    pub fn and(self, rhs: Self) -> Self {
119        Self::And(Box::new(self), Box::new(rhs))
120    }
121    /// Combines this predicate with `rhs` using logical OR.
122    pub fn or(self, rhs: Self) -> Self {
123        Self::Or(Box::new(self), Box::new(rhs))
124    }
125    /// Creates a closed-interval predicate.
126    pub fn between(value: impl Into<Expr>, lower: impl Into<Expr>, upper: impl Into<Expr>) -> Self {
127        Self::between_with(value, lower, upper, IntervalClosure::Closed)
128    }
129    /// Creates an interval predicate with an explicit endpoint policy.
130    pub fn between_with(
131        value: impl Into<Expr>,
132        lower: impl Into<Expr>,
133        upper: impl Into<Expr>,
134        closure: IntervalClosure,
135    ) -> Self {
136        Self::Between {
137            value: value.into(),
138            lower: lower.into(),
139            upper: upper.into(),
140            closure,
141        }
142    }
143}
144
145impl std::ops::Not for Predicate {
146    type Output = Self;
147    fn not(self) -> Self::Output {
148        Self::Not(Box::new(self))
149    }
150}
151
152/// Validated, monotonically increasing bin edges.
153#[derive(Clone, Debug, PartialEq, Serialize)]
154#[serde(transparent)]
155pub struct BinSpec {
156    axis: BinningAxis,
157}
158
159impl<'de> Deserialize<'de> for BinSpec {
160    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
161    where
162        D: Deserializer<'de>,
163    {
164        let edges = Vec::<f64>::deserialize(deserializer)?;
165        Self::edges(edges).map_err(serde::de::Error::custom)
166    }
167}
168
169impl BinSpec {
170    /// Creates `count` uniformly spaced bins spanning `[min, max]`.
171    ///
172    /// # Errors
173    ///
174    /// Returns [`RuntimeError`] when `count` is zero or the bounds are
175    /// non-finite or not increasing.
176    pub fn uniform(count: usize, min: f64, max: f64) -> RuntimeResult<Self> {
177        Ok(Self {
178            axis: BinningAxis::uniform(count, min, max)
179                .map_err(|error| query_error(error.to_string()))?,
180        })
181    }
182
183    /// Creates bins from explicit, strictly increasing finite edges.
184    ///
185    /// # Errors
186    ///
187    /// Returns [`RuntimeError`] when fewer than two edges are supplied or an
188    /// edge is non-finite or not strictly increasing.
189    pub fn edges(edges: impl IntoIterator<Item = f64>) -> RuntimeResult<Self> {
190        Ok(Self {
191            axis: BinningAxis::new(edges).map_err(|error| query_error(error.to_string()))?,
192        })
193    }
194
195    /// Returns the number of bins.
196    pub fn bin_count(&self) -> usize {
197        self.axis.bin_count()
198    }
199    /// Returns the validated bin edges.
200    pub fn edges_slice(&self) -> &[f64] {
201        self.axis.edges()
202    }
203
204    fn index(&self, value: f64) -> Option<usize> {
205        self.axis.index(value, FinalUpperEdge::Inclusive)
206    }
207}
208
209/// A lazily filtered dataset corresponding to one bin.
210#[derive(Clone)]
211pub struct DatasetBin {
212    index: usize,
213    lower: f64,
214    upper: f64,
215    dataset: Dataset,
216}
217
218impl DatasetBin {
219    /// Returns the zero-based bin index.
220    pub fn index(&self) -> usize {
221        self.index
222    }
223    /// Returns the bin's lower edge.
224    pub fn lower(&self) -> f64 {
225        self.lower
226    }
227    /// Returns the bin's upper edge.
228    pub fn upper(&self) -> f64 {
229        self.upper
230    }
231    /// Returns the dataset containing events in this bin.
232    pub fn dataset(&self) -> &Dataset {
233        &self.dataset
234    }
235    /// Consumes the bin and returns its dataset.
236    pub fn into_dataset(self) -> Dataset {
237        self.dataset
238    }
239}
240
241/// Expression-based query operations for datasets.
242pub trait DatasetExprExt {
243    /// Validates real scalar query expressions without reading dataset events.
244    ///
245    /// # Errors
246    /// Returns an error for empty, parameter-dependent, or non-real expressions,
247    /// or preparation failure.
248    fn validate_real_expressions(
249        &self,
250        expressions: &[Expr],
251        execution: &Execution,
252    ) -> RuntimeResult<()>;
253    /// Evaluates a scalar expression for every event.
254    ///
255    /// # Errors
256    ///
257    /// Returns [`RuntimeError`] when compilation, dataset reading, or
258    /// evaluation fails, or the expression is not scalar.
259    fn evaluate_expr(&self, expr: &Expr, execution: &Execution) -> RuntimeResult<Vec<Complex64>>;
260    /// Evaluates a real scalar expression for every event.
261    ///
262    /// # Errors
263    ///
264    /// Returns [`RuntimeError`] when compilation, dataset reading, or
265    /// evaluation fails, or the expression is not real scalar-valued.
266    fn evaluate_real(&self, expr: &Expr, execution: &Execution) -> RuntimeResult<Vec<f64>>;
267    /// Visits real scalar expression values in bounded event chunks.
268    /// The offset is the first event's global row number and expressions retain
269    /// their requested order within each callback.
270    ///
271    /// # Errors
272    /// Returns an error for zero chunk size, invalid expressions, dataset access,
273    /// evaluation, or a callback failure.
274    fn visit_real_chunks(
275        &self,
276        expressions: &[Expr],
277        execution: &Execution,
278        chunk_size: usize,
279        consume: impl FnMut(usize, &[Vec<f64>]) -> RuntimeResult<()>,
280    ) -> RuntimeResult<()>;
281    /// Evaluates and fills one weighted histogram in a bounded dataset traversal.
282    ///
283    /// Dataset event weights are used when `event_weights` is true. When
284    /// `weight` is supplied, its real scalar value multiplies the selected
285    /// event weight (or unit weight when event weights are disabled).
286    ///
287    /// # Errors
288    ///
289    /// Returns [`RuntimeError`] when bin validation, expression preparation,
290    /// dataset reading, evaluation, or histogram filling fails.
291    fn histogram(
292        &self,
293        expr: &Expr,
294        bins: BinSpec,
295        event_weights: bool,
296        weight: Option<&Expr>,
297        execution: &Execution,
298    ) -> RuntimeResult<Histogram>;
299    /// Evaluates ordered scalar axes and fills one bounded row-major joint histogram.
300    ///
301    /// # Errors
302    ///
303    /// Returns [`RuntimeError`] for missing or mismatched axes, invalid edges,
304    /// non-scalar expressions, evaluation failures, or source read failures.
305    fn joint_histogram(
306        &self,
307        axes: &[Expr],
308        bins: Vec<BinSpec>,
309        event_weights: bool,
310        weight: Option<&Expr>,
311        execution: &Execution,
312    ) -> RuntimeResult<JointHistogram>;
313    /// Creates a lazily filtered dataset containing events that satisfy `predicate`.
314    ///
315    /// # Errors
316    ///
317    /// Returns [`RuntimeError`] when predicate compilation or evaluation
318    /// fails, or its expression is not real scalar-valued.
319    fn select(&self, predicate: &Predicate, execution: &Execution) -> RuntimeResult<Dataset>;
320    /// Partitions the dataset in one pass according to an expression and bin specification.
321    ///
322    /// # Errors
323    ///
324    /// Returns [`RuntimeError`] when expression compilation or evaluation
325    /// fails, or the expression is not real scalar-valued.
326    fn bin_by(
327        &self,
328        expr: &Expr,
329        bins: BinSpec,
330        execution: &Execution,
331    ) -> RuntimeResult<Vec<DatasetBin>>;
332}
333
334impl DatasetExprExt for Dataset {
335    fn validate_real_expressions(
336        &self,
337        expressions: &[Expr],
338        execution: &Execution,
339    ) -> RuntimeResult<()> {
340        if expressions.is_empty() {
341            return Err(query_error(
342                "real expression validation needs at least one expression",
343            ));
344        }
345        QueryExprSet::prepare(expressions.to_vec(), execution, true)?;
346        Ok(())
347    }
348
349    fn evaluate_expr(&self, expr: &Expr, execution: &Execution) -> RuntimeResult<Vec<Complex64>> {
350        let query = QueryExprSet::prepare(vec![expr.clone()], execution, false)?;
351        let mut output = Vec::new();
352        for batch in self.batches().map_err(data_error)? {
353            output.extend(
354                query.evaluate_batch(&batch.map_err(data_error)?)?[0]
355                    .iter()
356                    .copied(),
357            );
358        }
359        Ok(output)
360    }
361
362    fn evaluate_real(&self, expr: &Expr, execution: &Execution) -> RuntimeResult<Vec<f64>> {
363        let query = QueryExprSet::prepare(vec![expr.clone()], execution, true)?;
364        let mut output = Vec::new();
365        for batch in self.batches().map_err(data_error)? {
366            output.extend(
367                query.evaluate_batch(&batch.map_err(data_error)?)?[0]
368                    .iter()
369                    .copied()
370                    .map(|v| v.re),
371            );
372        }
373        Ok(output)
374    }
375
376    fn visit_real_chunks(
377        &self,
378        expressions: &[Expr],
379        execution: &Execution,
380        chunk_size: usize,
381        mut consume: impl FnMut(usize, &[Vec<f64>]) -> RuntimeResult<()>,
382    ) -> RuntimeResult<()> {
383        if chunk_size == 0 || expressions.is_empty() {
384            return Err(query_error(
385                "real chunk evaluation needs expressions and positive chunk size",
386            ));
387        }
388        let query = QueryExprSet::prepare(expressions.to_vec(), execution, true)?;
389        let mut offset = 0;
390        for batch in self.batches().map_err(data_error)? {
391            let batch = batch.map_err(data_error)?;
392            for start in (0..batch.len()).step_by(chunk_size) {
393                let end = (start + chunk_size).min(batch.len());
394                let values = query
395                    .evaluate_batch(&batch.slice(start, end))?
396                    .into_iter()
397                    .map(|column| column.into_iter().map(|value| value.re).collect::<Vec<_>>())
398                    .collect::<Vec<_>>();
399                consume(offset, &values)?;
400                offset += end - start;
401            }
402        }
403        Ok(())
404    }
405
406    fn histogram(
407        &self,
408        expr: &Expr,
409        bins: BinSpec,
410        event_weights: bool,
411        weight: Option<&Expr>,
412        execution: &Execution,
413    ) -> RuntimeResult<Histogram> {
414        let mut expressions = vec![expr.clone()];
415        expressions.extend(weight.cloned());
416        let query = QueryExprSet::prepare(expressions, execution, true)?;
417        let mut histogram = Histogram::empty_with_edges(bins.edges_slice().to_vec())
418            .map_err(|error| query_error(error.to_string()))?;
419
420        for batch in self.batches().map_err(data_error)? {
421            let batch = batch.map_err(data_error)?;
422            let values = query.evaluate_batch(&batch)?;
423            for row in 0..batch.len() {
424                let base_weight = if event_weights {
425                    batch.weights_at(row)
426                } else {
427                    1.0
428                };
429                let custom_weight = values.get(1).map_or(1.0, |weights| weights[row].re);
430                histogram
431                    .fill_weighted(values[0][row].re, base_weight * custom_weight)
432                    .map_err(|error| query_error(error.to_string()))?;
433            }
434        }
435        Ok(histogram)
436    }
437
438    fn joint_histogram(
439        &self,
440        axes: &[Expr],
441        bins: Vec<BinSpec>,
442        event_weights: bool,
443        weight: Option<&Expr>,
444        execution: &Execution,
445    ) -> RuntimeResult<JointHistogram> {
446        if axes.is_empty() || axes.len() != bins.len() {
447            return Err(query_error(
448                "joint histogram requires one bin specification per non-empty ordered axis",
449            ));
450        }
451        let edge_vectors = bins
452            .iter()
453            .map(|bins| bins.edges_slice().to_vec())
454            .collect();
455        let mut histogram =
456            JointHistogram::empty(edge_vectors).map_err(|error| query_error(error.to_string()))?;
457        let mut expressions = axes.to_vec();
458        expressions.extend(weight.cloned());
459        let query = QueryExprSet::prepare(expressions, execution, true)?;
460        let mut coordinates = vec![0.0; axes.len()];
461        for batch in self.batches().map_err(data_error)? {
462            let batch = batch.map_err(data_error)?;
463            let values = query.evaluate_batch(&batch)?;
464            for row in 0..batch.len() {
465                for (axis, coordinate) in coordinates.iter_mut().enumerate() {
466                    *coordinate = values[axis][row].re;
467                }
468                let base_weight = if event_weights {
469                    batch.weights_at(row)
470                } else {
471                    1.0
472                };
473                let custom_weight = values
474                    .get(axes.len())
475                    .map_or(1.0, |weights| weights[row].re);
476                histogram
477                    .fill_weighted(&coordinates, base_weight * custom_weight)
478                    .map_err(|error| query_error(error.to_string()))?;
479            }
480        }
481        Ok(histogram)
482    }
483
484    fn select(&self, predicate: &Predicate, execution: &Execution) -> RuntimeResult<Dataset> {
485        let compiled = CompiledPredicate::prepare(predicate, execution)?;
486        Ok(self.with_derived_source(QuerySource {
487            source: self.clone(),
488            filter: QueryFilter::Predicate(Arc::new(compiled)),
489        }))
490    }
491
492    fn bin_by(
493        &self,
494        expr: &Expr,
495        bins: BinSpec,
496        execution: &Execution,
497    ) -> RuntimeResult<Vec<DatasetBin>> {
498        let query = QueryExprSet::prepare(vec![expr.clone()], execution, true)?;
499        let schema = self.schema().map_err(data_error)?;
500        let mut partitions = vec![Vec::new(); bins.bin_count()];
501        for batch in self.batches().map_err(data_error)? {
502            let batch = batch.map_err(data_error)?;
503            let mut rows = vec![Vec::new(); bins.bin_count()];
504            for (row, value) in query.evaluate_batch(&batch)?[0].iter().copied().enumerate() {
505                if let Some(index) = bins.index(value.re) {
506                    rows[index].push(row);
507                }
508            }
509            for (partition, rows) in partitions.iter_mut().zip(rows) {
510                if !rows.is_empty() {
511                    partition.push(batch.select(&rows));
512                }
513            }
514        }
515
516        partitions
517            .into_iter()
518            .enumerate()
519            .map(|(index, batches)| {
520                let source = if batches.is_empty() {
521                    MemorySource::empty(Arc::clone(&schema))
522                } else {
523                    MemorySource::from_batches(batches).map_err(data_error)?
524                };
525                Ok(DatasetBin {
526                    index,
527                    lower: bins.edges_slice()[index],
528                    upper: bins.edges_slice()[index + 1],
529                    dataset: self.with_derived_source(source),
530                })
531            })
532            .collect()
533    }
534}
535
536struct QueryExpr {
537    model: PreparedModel,
538    params: laddu_expr::parameters::ParamValues,
539    outputs: Vec<laddu_expr::ExprId>,
540}
541
542struct QueryExprSet {
543    shared: QueryExprStorage,
544    outputs: usize,
545}
546
547enum QueryExprStorage {
548    Shared(QueryExpr),
549    Separate(Vec<QueryExpr>),
550}
551
552impl QueryExprSet {
553    fn prepare(
554        expressions: Vec<Expr>,
555        execution: &Execution,
556        require_real: bool,
557    ) -> RuntimeResult<Self> {
558        let expression_count = expressions.len();
559        let compiled = CompiledQuery::from_exprs(expressions.clone())
560            .map_err(|error| query_error(error.to_string()))?;
561        let model = compiled.model();
562        let outputs = compiled.outputs();
563        if outputs.len() != expression_count {
564            return Err(query_error(
565                "compiled query output count changed during lowering",
566            ));
567        }
568        for element in outputs {
569            let value_kind = model
570                .node_facts(*element)
571                .map(|facts| facts.value_kind)
572                .ok_or_else(|| query_error("compiled expression facts are incomplete"))?;
573            if value_kind != ValueKind::Real
574                && (require_real || !matches!(value_kind, ValueKind::Complex))
575            {
576                return Err(query_error(if require_real {
577                    "this dataset operation requires a real-valued expression"
578                } else {
579                    "dataset expressions must be scalar"
580                }));
581            }
582        }
583        if model.params().n_free() != 0 {
584            return Err(query_error(
585                "dataset expressions cannot contain free parameters",
586            ));
587        }
588        let params = model.params().default_values();
589        let plan = match PreparedModel::prepare(model, execution) {
590            Ok(plan) => QueryExprStorage::Shared(QueryExpr {
591                model: plan,
592                params,
593                outputs: outputs.to_vec(),
594            }),
595            Err(shared_error) => {
596                if !may_fallback_to_scalar(execution, &shared_error) {
597                    return Err(shared_error);
598                }
599                let separate = expressions
600                    .iter()
601                    .map(|expr| QueryExpr::prepare(expr, execution, require_real))
602                    .collect::<RuntimeResult<Vec<_>>>();
603                QueryExprStorage::Separate(separate?)
604            }
605        };
606        Ok(Self {
607            shared: plan,
608            outputs: expression_count,
609        })
610    }
611
612    fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<Vec<Complex64>>> {
613        let values = match &self.shared {
614            QueryExprStorage::Shared(query) => {
615                query
616                    .model
617                    .evaluate_batch_outputs(&query.params, batch, &query.outputs)?
618            }
619            QueryExprStorage::Separate(queries) => queries
620                .iter()
621                .map(|query| query.evaluate_batch(batch))
622                .collect::<RuntimeResult<Vec<_>>>()?,
623        };
624        if values.len() != self.outputs {
625            return Err(query_error(
626                "compiled query returned an unexpected output count",
627            ));
628        }
629        Ok(values)
630    }
631}
632
633impl QueryExpr {
634    fn prepare(expr: &Expr, execution: &Execution, require_real: bool) -> RuntimeResult<Self> {
635        if expr.shape().map_err(|e| query_error(e.to_string()))? != laddu_expr::ExprShape::Scalar {
636            return Err(query_error("dataset expressions must be scalar"));
637        }
638        let compiled = laddu_compile::CompiledModel::from_expr(expr)
639            .map_err(|error| query_error(error.to_string()))?;
640        let value_kind = compiled
641            .node_facts(compiled.graph().root())
642            .map(|facts| facts.value_kind)
643            .ok_or_else(|| query_error("compiled expression facts are incomplete"))?;
644        if require_real && value_kind == ValueKind::Complex {
645            return Err(query_error(
646                "this dataset operation requires a real-valued expression",
647            ));
648        }
649        if compiled.params().n_free() != 0 {
650            return Err(query_error(
651                "dataset expressions cannot contain free parameters",
652            ));
653        }
654        let params = compiled.params().default_values();
655        let model = PreparedModel::prepare(&compiled, execution)?;
656        Ok(Self {
657            model,
658            params,
659            outputs: Vec::new(),
660        })
661    }
662
663    fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<Complex64>> {
664        self.model.evaluate_batch(&self.params, batch)
665    }
666}
667
668struct CompiledPredicate {
669    expressions: QueryExprSet,
670    program: PredicateProgram,
671}
672
673enum PredicateProgram {
674    Compare {
675        lhs: usize,
676        op: Comparison,
677        rhs: usize,
678    },
679    And(Box<Self>, Box<Self>),
680    Or(Box<Self>, Box<Self>),
681    Not(Box<Self>),
682    Between {
683        value: usize,
684        lower: usize,
685        upper: usize,
686        closure: IntervalClosure,
687    },
688}
689
690impl CompiledPredicate {
691    fn prepare(predicate: &Predicate, execution: &Execution) -> RuntimeResult<Self> {
692        let mut expressions = Vec::new();
693        let program = Self::compile_program(predicate, &mut expressions);
694        Ok(Self {
695            expressions: QueryExprSet::prepare(expressions, execution, true)?,
696            program,
697        })
698    }
699
700    fn compile_program(predicate: &Predicate, expressions: &mut Vec<Expr>) -> PredicateProgram {
701        let leaf = |expr: &Expr, expressions: &mut Vec<Expr>| {
702            let index = expressions.len();
703            expressions.push(expr.clone());
704            index
705        };
706        match predicate {
707            Predicate::Compare { lhs, op, rhs } => PredicateProgram::Compare {
708                lhs: leaf(lhs, expressions),
709                op: *op,
710                rhs: leaf(rhs, expressions),
711            },
712            Predicate::And(lhs, rhs) => PredicateProgram::And(
713                Box::new(Self::compile_program(lhs, expressions)),
714                Box::new(Self::compile_program(rhs, expressions)),
715            ),
716            Predicate::Or(lhs, rhs) => PredicateProgram::Or(
717                Box::new(Self::compile_program(lhs, expressions)),
718                Box::new(Self::compile_program(rhs, expressions)),
719            ),
720            Predicate::Not(inner) => {
721                PredicateProgram::Not(Box::new(Self::compile_program(inner, expressions)))
722            }
723            Predicate::Between {
724                value,
725                lower,
726                upper,
727                closure,
728            } => PredicateProgram::Between {
729                value: leaf(value, expressions),
730                lower: leaf(lower, expressions),
731                upper: leaf(upper, expressions),
732                closure: *closure,
733            },
734        }
735    }
736
737    fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<usize>> {
738        let values = self.expressions.evaluate_batch(batch)?;
739        Ok((0..batch.len())
740            .filter(|row| Self::evaluate_row(&self.program, &values, *row))
741            .collect())
742    }
743
744    fn evaluate_row(program: &PredicateProgram, values: &[Vec<Complex64>], row: usize) -> bool {
745        match program {
746            PredicateProgram::Compare { lhs, op, rhs } => {
747                compare(values[*lhs][row].re, *op, values[*rhs][row].re)
748            }
749            PredicateProgram::And(lhs, rhs) => {
750                Self::evaluate_row(lhs, values, row) && Self::evaluate_row(rhs, values, row)
751            }
752            PredicateProgram::Or(lhs, rhs) => {
753                Self::evaluate_row(lhs, values, row) || Self::evaluate_row(rhs, values, row)
754            }
755            PredicateProgram::Not(inner) => !Self::evaluate_row(inner, values, row),
756            PredicateProgram::Between {
757                value,
758                lower,
759                upper,
760                closure,
761            } => {
762                let lower_op = match closure {
763                    IntervalClosure::Open | IntervalClosure::RightClosed => Comparison::Gt,
764                    IntervalClosure::LeftClosed | IntervalClosure::Closed => Comparison::Ge,
765                };
766                let upper_op = match closure {
767                    IntervalClosure::Open | IntervalClosure::LeftClosed => Comparison::Lt,
768                    IntervalClosure::RightClosed | IntervalClosure::Closed => Comparison::Le,
769                };
770                compare(values[*value][row].re, lower_op, values[*lower][row].re)
771                    && compare(values[*value][row].re, upper_op, values[*upper][row].re)
772            }
773        }
774    }
775}
776
777fn compare(lhs: f64, op: Comparison, rhs: f64) -> bool {
778    if lhs.is_nan() || rhs.is_nan() {
779        return false;
780    }
781    match op {
782        Comparison::Lt => lhs < rhs,
783        Comparison::Le => lhs <= rhs,
784        Comparison::Gt => lhs > rhs,
785        Comparison::Ge => lhs >= rhs,
786        Comparison::Eq => lhs == rhs,
787        Comparison::Ne => lhs != rhs,
788    }
789}
790
791#[derive(Clone)]
792struct QuerySource {
793    source: Dataset,
794    filter: QueryFilter,
795}
796
797#[derive(Clone)]
798enum QueryFilter {
799    Predicate(Arc<CompiledPredicate>),
800}
801
802impl EventSource for QuerySource {
803    fn schema(&self) -> LadduDataResult<Arc<Schema>> {
804        self.source.schema()
805    }
806
807    fn capabilities(&self) -> SourceCapabilities {
808        let source = self.source.capabilities();
809        SourceCapabilities {
810            exact_len: false,
811            exact_weighted_total: false,
812            random_access: false,
813            deterministic_partitioning: source.deterministic_partitioning,
814            predicate_pushdown: false,
815            projection_pushdown: false,
816            streaming: true,
817        }
818    }
819
820    fn batches(&self, plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
821        let batches = self.source.stream_with_plan(plan)?;
822        let filter = self.filter.clone();
823        Ok(Box::new(batches.filter_map(move |batch| {
824            let batch = match batch {
825                Ok(batch) => batch,
826                Err(error) => return Some(Err(error)),
827            };
828            let rows = match filter.rows(&batch) {
829                Ok(rows) => rows,
830                Err(error) => return Some(Err(LadduDataError::Source(error.to_string()))),
831            };
832            (!rows.is_empty()).then(|| Ok(batch.select(&rows)))
833        })))
834    }
835}
836
837impl QueryFilter {
838    fn rows(&self, batch: &EventBatch) -> RuntimeResult<Vec<usize>> {
839        match self {
840            Self::Predicate(predicate) => predicate.evaluate_batch(batch),
841        }
842    }
843}
844
845fn query_error(message: impl Into<String>) -> RuntimeError {
846    RuntimeError::InvalidShape {
847        index: 0,
848        message: message.into(),
849    }
850}
851
852fn may_fallback_to_scalar(execution: &Execution, error: &RuntimeError) -> bool {
853    let cpu_f32 = matches!(
854        error,
855        RuntimeError::Execution(crate::ExecutionError::UnsupportedCpuF32Model)
856    );
857    #[cfg(feature = "wgpu")]
858    {
859        cpu_f32 || (execution.wgpu_context().is_some() && matches!(error, RuntimeError::Wgpu(_)))
860    }
861    #[cfg(not(feature = "wgpu"))]
862    {
863        let _ = execution;
864        cpu_f32
865    }
866}
867
868fn data_error(error: impl ToString) -> RuntimeError {
869    RuntimeError::Data(error.to_string())
870}
871
872#[cfg(test)]
873mod tests {
874    use super::*;
875    use crate::{CpuOptions, Device, ExecutionOptions, Precision};
876    #[cfg(feature = "jit")]
877    use crate::{JitPolicy, ThreadPolicy};
878    use laddu_compile::CompiledModel;
879    use laddu_data::{
880        data::{EventBatch, OwnedEvent},
881        io::{EventSource, ReadPlan, SourceCapabilities, memory::MemorySource},
882        schema::Schema,
883    };
884    use laddu_expr::{complex, event_scalar};
885    use std::sync::atomic::{AtomicUsize, Ordering};
886
887    #[test]
888    fn bin_spec_roundtrip_preserves_validation() {
889        let bins = BinSpec::edges([-1.0, 0.0, 2.0]).unwrap();
890        let json = serde_json::to_string(&bins).unwrap();
891        assert_eq!(serde_json::from_str::<BinSpec>(&json).unwrap(), bins);
892        assert!(serde_json::from_str::<BinSpec>("[0.0,0.0]").is_err());
893    }
894
895    #[derive(Clone)]
896    struct CountingSource {
897        inner: MemorySource,
898        reads: Arc<AtomicUsize>,
899    }
900
901    impl EventSource for CountingSource {
902        fn schema(&self) -> LadduDataResult<Arc<Schema>> {
903            EventSource::schema(&self.inner)
904        }
905
906        fn capabilities(&self) -> SourceCapabilities {
907            self.inner.capabilities()
908        }
909
910        fn batches(&self, plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
911            self.reads.fetch_add(1, Ordering::Relaxed);
912            self.inner.batches(plan)
913        }
914    }
915
916    #[derive(Clone)]
917    struct FailingSource {
918        schema: Arc<Schema>,
919    }
920
921    impl EventSource for FailingSource {
922        fn schema(&self) -> LadduDataResult<Arc<Schema>> {
923            Ok(Arc::clone(&self.schema))
924        }
925
926        fn capabilities(&self) -> SourceCapabilities {
927            SourceCapabilities {
928                exact_len: false,
929                exact_weighted_total: false,
930                random_access: false,
931                deterministic_partitioning: true,
932                predicate_pushdown: false,
933                projection_pushdown: false,
934                streaming: true,
935            }
936        }
937
938        fn batches(&self, _plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
939            Ok(Box::new(std::iter::once(Err(LadduDataError::Source(
940                "query source failed".into(),
941            )))))
942        }
943    }
944
945    fn capability_tuple(
946        capabilities: SourceCapabilities,
947    ) -> (bool, bool, bool, bool, bool, bool, bool) {
948        (
949            capabilities.exact_len,
950            capabilities.exact_weighted_total,
951            capabilities.random_access,
952            capabilities.deterministic_partitioning,
953            capabilities.predicate_pushdown,
954            capabilities.projection_pushdown,
955            capabilities.streaming,
956        )
957    }
958
959    fn dataset() -> Dataset {
960        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
961        Dataset::from_events(
962            schema,
963            [
964                OwnedEvent::weighted(vec![], vec![-1.0], 0.5),
965                OwnedEvent::weighted(vec![], vec![0.0], 1.0),
966                OwnedEvent::weighted(vec![], vec![1.0], 1.5),
967                OwnedEvent::weighted(vec![], vec![2.0], 2.0),
968            ],
969        )
970        .unwrap()
971    }
972
973    #[test]
974    fn evaluates_selects_and_bins_dataset_expressions() {
975        let dataset = dataset().chunked(1).unwrap();
976        let execution = Execution::default();
977        let x = event_scalar("x");
978        assert_eq!(
979            dataset.evaluate_real(&x, &execution).unwrap(),
980            vec![-1.0, 0.0, 1.0, 2.0]
981        );
982
983        let selected = dataset
984            .select(
985                &Predicate::ge(x.clone(), 0.0).and(Predicate::lt(x.clone(), 2.0)),
986                &execution,
987            )
988            .unwrap();
989        assert_eq!(
990            selected.map_events(|event| event.scalar(0)).unwrap(),
991            vec![0.0, 1.0]
992        );
993        assert_eq!(selected.sum_weights().unwrap(), 2.5);
994
995        let bins = dataset
996            .bin_by(&x, BinSpec::uniform(2, 0.0, 2.0).unwrap(), &execution)
997            .unwrap();
998        assert_eq!(bins.len(), 2);
999        assert_eq!(
1000            bins[0]
1001                .dataset()
1002                .map_events(|event| event.scalar(0))
1003                .unwrap(),
1004            vec![0.0]
1005        );
1006        assert_eq!(
1007            bins[1]
1008                .dataset()
1009                .map_events(|event| event.scalar(0))
1010                .unwrap(),
1011            vec![1.0, 2.0]
1012        );
1013    }
1014
1015    #[test]
1016    fn dataset_histogram_uses_event_weights_and_excludes_the_final_upper_edge() {
1017        let histogram = dataset()
1018            .histogram(
1019                &event_scalar("x"),
1020                BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1021                true,
1022                None,
1023                &Execution::default(),
1024            )
1025            .unwrap();
1026
1027        assert_eq!(histogram.counts(), [1.0, 1.5]);
1028        assert_eq!(histogram.sum_squared_weights(), [1.0, 2.25]);
1029        assert_eq!(histogram.underflow(), 0.5);
1030        assert_eq!(histogram.underflow_sum_squared_weights(), Some(0.25));
1031        assert_eq!(histogram.overflow(), 2.0);
1032        assert_eq!(histogram.overflow_sum_squared_weights(), Some(4.0));
1033    }
1034
1035    #[test]
1036    fn dataset_joint_histogram_uses_row_major_bins_and_aggregates_invalid_events() {
1037        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x", "y"], true).unwrap());
1038        let dataset = Dataset::from_events(
1039            schema,
1040            [
1041                OwnedEvent::weighted(vec![], vec![0.0, 10.0], 1.0),
1042                OwnedEvent::weighted(vec![], vec![1.0, 10.0], -2.0),
1043                OwnedEvent::weighted(vec![], vec![0.0, 20.0], 3.0),
1044                OwnedEvent::weighted(vec![], vec![f64::NAN, 10.0], 4.0),
1045            ],
1046        )
1047        .unwrap();
1048
1049        let histogram = dataset
1050            .joint_histogram(
1051                &[event_scalar("x"), event_scalar("y")],
1052                vec![
1053                    BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1054                    BinSpec::edges([0.0, 15.0, 25.0]).unwrap(),
1055                ],
1056                true,
1057                None,
1058                &Execution::default(),
1059            )
1060            .unwrap();
1061
1062        assert_eq!(histogram.shape(), [2, 2]);
1063        assert_eq!(histogram.values(), [1.0, 3.0, -2.0, 0.0]);
1064        assert_eq!(histogram.sum_squared_weights(), [1.0, 9.0, 4.0, 0.0]);
1065        assert_eq!(histogram.diagnostics().nonfinite_count(), 1);
1066        assert_eq!(histogram.diagnostics().out_of_range_count(), 0);
1067    }
1068
1069    #[test]
1070    fn dataset_histogram_multiplies_the_selected_base_and_custom_weights() {
1071        let x = event_scalar("x");
1072        let bins = BinSpec::edges([-2.0, 0.0, 2.0]).unwrap();
1073        let execution = Execution::default();
1074
1075        let weighted = dataset()
1076            .histogram(&x, bins.clone(), true, Some(&(x.clone() + 2.0)), &execution)
1077            .unwrap();
1078        assert_eq!(weighted.counts(), [0.5, 6.5]);
1079        assert_eq!(weighted.sum_squared_weights(), [0.25, 24.25]);
1080        assert_eq!(weighted.overflow(), 8.0);
1081
1082        let custom_only = dataset()
1083            .histogram(&x, bins, false, Some(&(x.clone() + 2.0)), &execution)
1084            .unwrap();
1085        assert_eq!(custom_only.counts(), [1.0, 5.0]);
1086        assert_eq!(custom_only.sum_squared_weights(), [1.0, 13.0]);
1087        assert_eq!(custom_only.overflow(), 4.0);
1088    }
1089
1090    #[test]
1091    fn dataset_histogram_preserves_view_source_and_memory_semantics() {
1092        let source = dataset();
1093        let batch = source.batches().unwrap().next().unwrap().unwrap();
1094        let reads = Arc::new(AtomicUsize::new(0));
1095        let counted = Dataset::new(CountingSource {
1096            inner: MemorySource::new(batch),
1097            reads: Arc::clone(&reads),
1098        })
1099        .streaming()
1100        .chunked(1)
1101        .unwrap();
1102        let selected = counted
1103            .select(
1104                &Predicate::ge(event_scalar("x"), 0.0),
1105                &Execution::default(),
1106            )
1107            .unwrap();
1108
1109        let histogram = selected
1110            .histogram(
1111                &event_scalar("x"),
1112                BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1113                true,
1114                None,
1115                &Execution::default(),
1116            )
1117            .unwrap();
1118
1119        assert_eq!(histogram.counts(), [1.0, 1.5]);
1120        assert_eq!(histogram.overflow(), 2.0);
1121        assert_eq!(reads.load(Ordering::Relaxed), 1);
1122    }
1123
1124    #[test]
1125    fn dataset_histogram_matches_across_memory_policies_and_chunking() {
1126        let source = dataset();
1127        let bins = BinSpec::edges([0.0, 1.0, 2.0]).unwrap();
1128        let execution = Execution::default();
1129        let expected = source
1130            .clone()
1131            .resident()
1132            .histogram(&event_scalar("x"), bins.clone(), true, None, &execution)
1133            .unwrap();
1134
1135        for candidate in [
1136            source.clone().streaming(),
1137            source.clone().resident().chunked(1).unwrap(),
1138            source.streaming().chunked(2).unwrap(),
1139        ] {
1140            let actual = candidate
1141                .histogram(&event_scalar("x"), bins.clone(), true, None, &execution)
1142                .unwrap();
1143            assert_eq!(actual, expected);
1144        }
1145    }
1146
1147    #[cfg(feature = "jit")]
1148    #[test]
1149    fn dataset_histogram_matches_cpu_interpreter_and_jit_backends() {
1150        let source = dataset();
1151        let bins = BinSpec::edges([0.0, 1.0, 2.0]).unwrap();
1152        let execution = |jit| {
1153            Execution::local(ExecutionOptions {
1154                device: Device::Cpu(CpuOptions {
1155                    threads: ThreadPolicy::Serial,
1156                    jit,
1157                }),
1158                precision: Precision::F64,
1159                ..ExecutionOptions::default()
1160            })
1161            .unwrap()
1162        };
1163        let interpreted = source
1164            .histogram(
1165                &event_scalar("x"),
1166                bins.clone(),
1167                true,
1168                None,
1169                &execution(JitPolicy::Disabled),
1170            )
1171            .unwrap();
1172        let compiled = source
1173            .histogram(
1174                &event_scalar("x"),
1175                bins,
1176                true,
1177                None,
1178                &execution(JitPolicy::Enabled),
1179            )
1180            .unwrap();
1181
1182        assert_eq!(compiled, interpreted);
1183    }
1184
1185    #[test]
1186    fn dataset_histogram_reports_invalid_values_and_source_failures() {
1187        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1188        let nonfinite = Dataset::from_events(
1189            Arc::clone(&schema),
1190            [OwnedEvent::weighted(vec![], vec![f64::NAN], 1.0)],
1191        )
1192        .unwrap();
1193        let error = nonfinite
1194            .histogram(
1195                &event_scalar("x"),
1196                BinSpec::edges([0.0, 1.0]).unwrap(),
1197                true,
1198                None,
1199                &Execution::default(),
1200            )
1201            .unwrap_err();
1202        assert!(
1203            error.to_string().contains("expected finite, got NaN"),
1204            "unexpected error: {error}",
1205        );
1206
1207        let source_error = Dataset::new(FailingSource { schema })
1208            .histogram(
1209                &event_scalar("x"),
1210                BinSpec::edges([0.0, 1.0]).unwrap(),
1211                true,
1212                None,
1213                &Execution::default(),
1214            )
1215            .unwrap_err();
1216        assert!(source_error.to_string().contains("query source failed"));
1217    }
1218
1219    #[test]
1220    fn empty_dataset_histogram_is_a_valid_empirical_histogram() {
1221        let histogram = dataset()
1222            .empty_derived()
1223            .unwrap()
1224            .histogram(
1225                &event_scalar("x"),
1226                BinSpec::edges([0.0, 1.0, 3.0]).unwrap(),
1227                true,
1228                None,
1229                &Execution::default(),
1230            )
1231            .unwrap();
1232
1233        assert_eq!(histogram.counts(), [0.0, 0.0]);
1234        assert_eq!(histogram.sum_squared_weights(), [0.0, 0.0]);
1235        assert_eq!(histogram.underflow_sum_squared_weights(), Some(0.0));
1236        assert_eq!(histogram.overflow_sum_squared_weights(), Some(0.0));
1237    }
1238
1239    #[test]
1240    fn cancelling_dataset_weights_keep_their_squared_weight_uncertainty() {
1241        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1242        let cancelling = Dataset::from_events(
1243            schema,
1244            [
1245                OwnedEvent::weighted(vec![], vec![0.5], 1.0),
1246                OwnedEvent::weighted(vec![], vec![0.5], -1.0),
1247            ],
1248        )
1249        .unwrap();
1250        let histogram = cancelling
1251            .histogram(
1252                &event_scalar("x"),
1253                BinSpec::edges([0.0, 1.0]).unwrap(),
1254                true,
1255                None,
1256                &Execution::default(),
1257            )
1258            .unwrap();
1259
1260        assert_eq!(histogram.counts(), [0.0]);
1261        assert_eq!(histogram.sum_squared_weights(), [2.0]);
1262        assert_eq!(histogram.errors(), [2.0_f64.sqrt()]);
1263    }
1264
1265    #[test]
1266    fn empty_batches_are_valid_query_inputs() {
1267        let execution = Execution::default();
1268        let x = event_scalar("x");
1269
1270        let empty_batch_schema =
1271            Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1272        let empty_batch = Dataset::from_batch(
1273            EventBatch::from_events(empty_batch_schema, std::iter::empty::<OwnedEvent>()).unwrap(),
1274        );
1275        for empty in [empty_batch, dataset().empty_derived().unwrap()] {
1276            assert!(empty.evaluate_real(&x, &execution).unwrap().is_empty());
1277            assert!(
1278                empty
1279                    .select(&Predicate::ge(x.clone(), 0.0), &execution)
1280                    .unwrap()
1281                    .map_events(|event| event.scalar(0))
1282                    .unwrap()
1283                    .is_empty()
1284            );
1285        }
1286    }
1287
1288    #[test]
1289    fn event_column_nan_comparisons_are_false() {
1290        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1291        let dataset = Dataset::from_events(
1292            schema,
1293            [
1294                OwnedEvent::weighted(vec![], vec![f64::NAN], 1.0),
1295                OwnedEvent::weighted(vec![], vec![0.0], 1.0),
1296                OwnedEvent::weighted(vec![], vec![1.0], 1.0),
1297            ],
1298        )
1299        .unwrap();
1300        let x = event_scalar("x");
1301        let selected = dataset
1302            .select(&Predicate::ne(x, 0.0), &Execution::default())
1303            .unwrap();
1304
1305        assert_eq!(
1306            selected.map_events(|event| event.scalar(0)).unwrap(),
1307            vec![1.0]
1308        );
1309    }
1310
1311    #[test]
1312    fn query_propagates_source_batch_errors() {
1313        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1314        let dataset = Dataset::new(FailingSource { schema });
1315
1316        let error = dataset
1317            .evaluate_real(&event_scalar("x"), &Execution::default())
1318            .unwrap_err();
1319        assert!(
1320            matches!(error, RuntimeError::Data(message) if message.contains("query source failed"))
1321        );
1322    }
1323
1324    #[test]
1325    fn all_empty_bins_retain_valid_empty_derived_sources() {
1326        let source = dataset();
1327        let before = capability_tuple(source.capabilities());
1328        let bins = source
1329            .bin_by(
1330                &event_scalar("x"),
1331                BinSpec::edges([10.0, 20.0, 30.0]).unwrap(),
1332                &Execution::default(),
1333            )
1334            .unwrap();
1335
1336        assert_eq!(capability_tuple(source.capabilities()), before);
1337        assert_eq!(bins.len(), 2);
1338        for bin in bins {
1339            assert_eq!(bin.dataset().num_events().unwrap(), Some(0));
1340            assert!(
1341                bin.dataset()
1342                    .evaluate_real(&event_scalar("x"), &Execution::default())
1343                    .unwrap()
1344                    .is_empty()
1345            );
1346        }
1347    }
1348
1349    #[test]
1350    fn traversing_all_bins_reads_the_source_once() {
1351        let reads = Arc::new(AtomicUsize::new(0));
1352        let source = CountingSource {
1353            inner: match dataset().batches().unwrap().next().unwrap() {
1354                Ok(batch) => MemorySource::new(batch),
1355                Err(error) => panic!("unexpected source error: {error}"),
1356            },
1357            reads: Arc::clone(&reads),
1358        };
1359        let dataset = Dataset::new(source).chunked(1).unwrap();
1360        let bins = dataset
1361            .bin_by(
1362                &event_scalar("x"),
1363                BinSpec::uniform(4, -1.0, 3.0).unwrap(),
1364                &Execution::default(),
1365            )
1366            .unwrap();
1367
1368        let values = bins
1369            .into_iter()
1370            .map(|bin| {
1371                bin.into_dataset()
1372                    .map_events(|event| event.scalar(0))
1373                    .unwrap()
1374            })
1375            .collect::<Vec<_>>();
1376        assert_eq!(values, [vec![-1.0], vec![0.0], vec![1.0], vec![2.0]]);
1377        assert_eq!(reads.load(Ordering::Relaxed), 1);
1378    }
1379
1380    #[test]
1381    fn between_predicates_have_explicit_endpoint_semantics() {
1382        let dataset = dataset();
1383        let execution = Execution::default();
1384        let x = event_scalar("x");
1385
1386        let closed = dataset
1387            .select(&Predicate::between(x.clone(), 0.0, 1.0), &execution)
1388            .unwrap();
1389        assert_eq!(
1390            closed.map_events(|event| event.scalar(0)).unwrap(),
1391            vec![0.0, 1.0]
1392        );
1393
1394        let open = dataset
1395            .select(
1396                &Predicate::between_with(x, -1.0, 1.0, IntervalClosure::Open),
1397                &execution,
1398            )
1399            .unwrap();
1400        assert_eq!(open.map_events(|event| event.scalar(0)).unwrap(), vec![0.0]);
1401    }
1402
1403    #[test]
1404    fn real_queries_reject_complex_and_free_parameter_expressions() {
1405        let dataset = dataset();
1406        let execution = Execution::default();
1407        assert!(
1408            dataset
1409                .evaluate_real(&complex(1.0, 1.0), &execution)
1410                .is_err()
1411        );
1412        let parameter = Expr::from(laddu_expr::parameters::Parameter::free("p"));
1413        assert!(dataset.evaluate_expr(&parameter, &execution).is_err());
1414    }
1415
1416    #[test]
1417    fn compiled_query_outputs_preserve_order_and_values() {
1418        let source = dataset();
1419        let batch = source.batches().unwrap().next().unwrap().unwrap();
1420        let x = event_scalar("x");
1421        let query = QueryExprSet::prepare(
1422            vec![x.clone() + 1.0, x.clone() * 2.0, x],
1423            &Execution::default(),
1424            false,
1425        )
1426        .unwrap();
1427        let values = query.evaluate_batch(&batch).unwrap();
1428        assert_eq!(
1429            values[0].iter().map(|v| v.re).collect::<Vec<_>>(),
1430            [0.0, 1.0, 2.0, 3.0]
1431        );
1432        assert_eq!(
1433            values[1].iter().map(|v| v.re).collect::<Vec<_>>(),
1434            [-2.0, 0.0, 2.0, 4.0]
1435        );
1436        assert_eq!(
1437            values[2].iter().map(|v| v.re).collect::<Vec<_>>(),
1438            [-1.0, 0.0, 1.0, 2.0]
1439        );
1440    }
1441
1442    #[test]
1443    fn repeated_predicate_leaves_are_evaluated_once() {
1444        let x = event_scalar("x");
1445        let selected = dataset()
1446            .select(
1447                &Predicate::ge(x.clone() + 1.0, 0.0).and(Predicate::lt(x + 1.0, 2.0)),
1448                &Execution::default(),
1449            )
1450            .unwrap();
1451        assert_eq!(
1452            selected.map_events(|event| event.scalar(0)).unwrap(),
1453            [-1.0, 0.0]
1454        );
1455    }
1456
1457    #[test]
1458    fn f32_queries_match_f64_query_results() {
1459        let x = event_scalar("x");
1460        let f64_values = dataset().evaluate_real(&x, &Execution::default()).unwrap();
1461        let f32_execution = Execution::local(ExecutionOptions {
1462            device: Device::Cpu(CpuOptions::default()),
1463            precision: Precision::F32,
1464            ..ExecutionOptions::default()
1465        })
1466        .unwrap();
1467        let f32_values = dataset().evaluate_real(&x, &f32_execution).unwrap();
1468        assert_eq!(f32_values, f64_values);
1469
1470        let bins = BinSpec::edges([-2.0, 0.0, 2.0]).unwrap();
1471        let f64_histogram = dataset()
1472            .histogram(&x, bins.clone(), true, None, &Execution::default())
1473            .unwrap();
1474        let f32_histogram = dataset()
1475            .histogram(&x, bins, true, None, &f32_execution)
1476            .unwrap();
1477        assert_eq!(f32_histogram, f64_histogram);
1478    }
1479
1480    #[test]
1481    fn bin_edges_validate_and_nan_predicates_are_false() {
1482        assert!(BinSpec::edges([0.0, 0.0]).is_err());
1483        assert!(!compare(f64::NAN, Comparison::Ne, 0.0));
1484    }
1485
1486    #[test]
1487    fn selection_is_lazy_and_one_pass_binning_preserves_streaming_policy() {
1488        let source = dataset();
1489        let batch = source.batches().unwrap().next().unwrap().unwrap();
1490        let reads = Arc::new(AtomicUsize::new(0));
1491        let dataset = Dataset::new(CountingSource {
1492            inner: MemorySource::new(batch),
1493            reads: Arc::clone(&reads),
1494        })
1495        .streaming();
1496        let execution = Execution::default();
1497        let x = event_scalar("x");
1498
1499        let selected = dataset
1500            .select(&Predicate::ge(x.clone(), 0.0), &execution)
1501            .unwrap();
1502        let bins = dataset
1503            .bin_by(&x, BinSpec::uniform(2, 0.0, 2.0).unwrap(), &execution)
1504            .unwrap();
1505        assert_eq!(reads.load(Ordering::Relaxed), 1);
1506        assert_eq!(
1507            selected.cache_storage(),
1508            laddu_data::data::CacheStorage::Streaming
1509        );
1510
1511        assert_eq!(
1512            selected.map_events(|event| event.scalar(0)).unwrap(),
1513            vec![0.0, 1.0, 2.0]
1514        );
1515        assert_eq!(reads.load(Ordering::Relaxed), 2);
1516        assert_eq!(
1517            bins[0]
1518                .dataset()
1519                .map_events(|event| event.scalar(0))
1520                .unwrap(),
1521            vec![0.0]
1522        );
1523        assert_eq!(reads.load(Ordering::Relaxed), 2);
1524    }
1525
1526    #[test]
1527    fn unknown_cardinality_fastest_discovers_and_retains_small_selection() {
1528        let source = dataset();
1529        let batch = source.batches().unwrap().next().unwrap().unwrap();
1530        let reads = Arc::new(AtomicUsize::new(0));
1531        let dataset = Dataset::new(CountingSource {
1532            inner: MemorySource::new(batch),
1533            reads: Arc::clone(&reads),
1534        });
1535        let execution = Execution::default();
1536        let x = event_scalar("x");
1537        let selected = dataset
1538            .select(&Predicate::ge(x.clone(), 0.0), &execution)
1539            .unwrap();
1540        let compiled = CompiledModel::from_expr(&x).unwrap();
1541        let params = compiled.params().default_values();
1542        let model = PreparedModel::prepare(&compiled, &execution).unwrap();
1543        let prepared = model.prepare_dataset(&execution, &selected).unwrap();
1544
1545        #[cfg(not(feature = "wgpu"))]
1546        let crate::PreparedDataset::Cpu(prepared_cpu) = &prepared;
1547        #[cfg(feature = "wgpu")]
1548        let crate::PreparedDataset::Cpu(prepared_cpu) = &prepared else {
1549            panic!("default execution prepares CPU datasets");
1550        };
1551        assert_eq!(
1552            prepared_cpu.stats().storage(),
1553            laddu_data::data::CacheStorage::Resident
1554        );
1555        assert_eq!(prepared_cpu.stats().local_events(), 3);
1556        assert_eq!(reads.load(Ordering::Relaxed), 2);
1557
1558        for _ in 0..2 {
1559            assert_eq!(
1560                model
1561                    .reduce(
1562                        &execution,
1563                        &params,
1564                        &prepared,
1565                        laddu_compile::ReductionPlan::weighted_real(),
1566                    )
1567                    .unwrap(),
1568                5.5
1569            );
1570        }
1571        assert_eq!(reads.load(Ordering::Relaxed), 2);
1572    }
1573}