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        PreparedQuery::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 = PreparedQuery::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 = PreparedQuery::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 = PreparedQuery::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 = PreparedQuery::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 = PreparedQuery::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 = PreparedQuery::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
542/// Scalar expressions compiled and prepared once for repeated batch evaluation.
543/// This object retains no event data and owns no cross-request compilation cache.
544pub struct PreparedQuery {
545    shared: QueryExprStorage,
546    outputs: usize,
547}
548
549enum QueryExprStorage {
550    Shared(QueryExpr),
551    Separate(Vec<QueryExpr>),
552}
553
554impl PreparedQuery {
555    /// Compile scalar outputs in caller order for the selected backend.
556    /// Set `require_real` to reject complex outputs. Free parameters are rejected.
557    ///
558    /// # Errors
559    /// Returns an error for invalid expressions or backend preparation failure.
560    pub fn prepare(
561        expressions: Vec<Expr>,
562        execution: &Execution,
563        require_real: bool,
564    ) -> RuntimeResult<Self> {
565        let expression_count = expressions.len();
566        let compiled = CompiledQuery::from_exprs(expressions.clone())
567            .map_err(|error| query_error(error.to_string()))?;
568        let model = compiled.model();
569        let outputs = compiled.outputs();
570        if outputs.len() != expression_count {
571            return Err(query_error(
572                "compiled query output count changed during lowering",
573            ));
574        }
575        for element in outputs {
576            let value_kind = model
577                .node_facts(*element)
578                .map(|facts| facts.value_kind)
579                .ok_or_else(|| query_error("compiled expression facts are incomplete"))?;
580            if value_kind != ValueKind::Real
581                && (require_real || !matches!(value_kind, ValueKind::Complex))
582            {
583                return Err(query_error(if require_real {
584                    "this dataset operation requires a real-valued expression"
585                } else {
586                    "dataset expressions must be scalar"
587                }));
588            }
589        }
590        if model.params().n_free() != 0 {
591            return Err(query_error(
592                "dataset expressions cannot contain free parameters",
593            ));
594        }
595        let params = model.params().default_values();
596        let plan = match PreparedModel::prepare(model, execution) {
597            Ok(plan) => QueryExprStorage::Shared(QueryExpr {
598                model: plan,
599                params,
600                outputs: outputs.to_vec(),
601            }),
602            Err(shared_error) => {
603                if !may_fallback_to_scalar(execution, &shared_error) {
604                    return Err(shared_error);
605                }
606                let separate = expressions
607                    .iter()
608                    .map(|expr| QueryExpr::prepare(expr, execution, require_real))
609                    .collect::<RuntimeResult<Vec<_>>>();
610                QueryExprStorage::Separate(separate?)
611            }
612        };
613        Ok(Self {
614            shared: plan,
615            outputs: expression_count,
616        })
617    }
618
619    /// Estimate peak host workspace for one batch, including output columns.
620    #[doc(hidden)]
621    pub fn batch_memory_estimate(&self, events: usize) -> usize {
622        let model = match &self.shared {
623            QueryExprStorage::Shared(query) => query.model.batch_memory_estimate(events),
624            QueryExprStorage::Separate(queries) => queries
625                .iter()
626                .map(|query| query.model.batch_memory_estimate(events))
627                .max()
628                .unwrap_or(0),
629        };
630        model.saturating_add(events.saturating_mul(self.outputs).saturating_mul(32))
631    }
632
633    /// Evaluate every output on a batch without retaining event values.
634    ///
635    /// # Errors
636    /// Returns an error for incompatible event columns or failed evaluation.
637    pub fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<Vec<Complex64>>> {
638        let values = match &self.shared {
639            QueryExprStorage::Shared(query) => {
640                query
641                    .model
642                    .evaluate_batch_outputs(&query.params, batch, &query.outputs)?
643            }
644            QueryExprStorage::Separate(queries) => queries
645                .iter()
646                .map(|query| query.evaluate_batch(batch))
647                .collect::<RuntimeResult<Vec<_>>>()?,
648        };
649        if values.len() != self.outputs {
650            return Err(query_error(
651                "compiled query returned an unexpected output count",
652            ));
653        }
654        Ok(values)
655    }
656}
657
658impl QueryExpr {
659    fn prepare(expr: &Expr, execution: &Execution, require_real: bool) -> RuntimeResult<Self> {
660        if expr.shape().map_err(|e| query_error(e.to_string()))? != laddu_expr::ExprShape::Scalar {
661            return Err(query_error("dataset expressions must be scalar"));
662        }
663        let compiled = laddu_compile::CompiledModel::from_expr(expr)
664            .map_err(|error| query_error(error.to_string()))?;
665        let value_kind = compiled
666            .node_facts(compiled.graph().root())
667            .map(|facts| facts.value_kind)
668            .ok_or_else(|| query_error("compiled expression facts are incomplete"))?;
669        if require_real && value_kind == ValueKind::Complex {
670            return Err(query_error(
671                "this dataset operation requires a real-valued expression",
672            ));
673        }
674        if compiled.params().n_free() != 0 {
675            return Err(query_error(
676                "dataset expressions cannot contain free parameters",
677            ));
678        }
679        let params = compiled.params().default_values();
680        let model = PreparedModel::prepare(&compiled, execution)?;
681        Ok(Self {
682            model,
683            params,
684            outputs: Vec::new(),
685        })
686    }
687
688    fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<Complex64>> {
689        self.model.evaluate_batch(&self.params, batch)
690    }
691}
692
693struct CompiledPredicate {
694    expressions: PreparedQuery,
695    program: PredicateProgram,
696}
697
698enum PredicateProgram {
699    Compare {
700        lhs: usize,
701        op: Comparison,
702        rhs: usize,
703    },
704    And(Box<Self>, Box<Self>),
705    Or(Box<Self>, Box<Self>),
706    Not(Box<Self>),
707    Between {
708        value: usize,
709        lower: usize,
710        upper: usize,
711        closure: IntervalClosure,
712    },
713}
714
715impl CompiledPredicate {
716    fn prepare(predicate: &Predicate, execution: &Execution) -> RuntimeResult<Self> {
717        let mut expressions = Vec::new();
718        let program = Self::compile_program(predicate, &mut expressions);
719        Ok(Self {
720            expressions: PreparedQuery::prepare(expressions, execution, true)?,
721            program,
722        })
723    }
724
725    fn compile_program(predicate: &Predicate, expressions: &mut Vec<Expr>) -> PredicateProgram {
726        let leaf = |expr: &Expr, expressions: &mut Vec<Expr>| {
727            let index = expressions.len();
728            expressions.push(expr.clone());
729            index
730        };
731        match predicate {
732            Predicate::Compare { lhs, op, rhs } => PredicateProgram::Compare {
733                lhs: leaf(lhs, expressions),
734                op: *op,
735                rhs: leaf(rhs, expressions),
736            },
737            Predicate::And(lhs, rhs) => PredicateProgram::And(
738                Box::new(Self::compile_program(lhs, expressions)),
739                Box::new(Self::compile_program(rhs, expressions)),
740            ),
741            Predicate::Or(lhs, rhs) => PredicateProgram::Or(
742                Box::new(Self::compile_program(lhs, expressions)),
743                Box::new(Self::compile_program(rhs, expressions)),
744            ),
745            Predicate::Not(inner) => {
746                PredicateProgram::Not(Box::new(Self::compile_program(inner, expressions)))
747            }
748            Predicate::Between {
749                value,
750                lower,
751                upper,
752                closure,
753            } => PredicateProgram::Between {
754                value: leaf(value, expressions),
755                lower: leaf(lower, expressions),
756                upper: leaf(upper, expressions),
757                closure: *closure,
758            },
759        }
760    }
761
762    fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<usize>> {
763        let values = self.expressions.evaluate_batch(batch)?;
764        Ok((0..batch.len())
765            .filter(|row| Self::evaluate_row(&self.program, &values, *row))
766            .collect())
767    }
768
769    fn evaluate_row(program: &PredicateProgram, values: &[Vec<Complex64>], row: usize) -> bool {
770        match program {
771            PredicateProgram::Compare { lhs, op, rhs } => {
772                compare(values[*lhs][row].re, *op, values[*rhs][row].re)
773            }
774            PredicateProgram::And(lhs, rhs) => {
775                Self::evaluate_row(lhs, values, row) && Self::evaluate_row(rhs, values, row)
776            }
777            PredicateProgram::Or(lhs, rhs) => {
778                Self::evaluate_row(lhs, values, row) || Self::evaluate_row(rhs, values, row)
779            }
780            PredicateProgram::Not(inner) => !Self::evaluate_row(inner, values, row),
781            PredicateProgram::Between {
782                value,
783                lower,
784                upper,
785                closure,
786            } => {
787                let lower_op = match closure {
788                    IntervalClosure::Open | IntervalClosure::RightClosed => Comparison::Gt,
789                    IntervalClosure::LeftClosed | IntervalClosure::Closed => Comparison::Ge,
790                };
791                let upper_op = match closure {
792                    IntervalClosure::Open | IntervalClosure::LeftClosed => Comparison::Lt,
793                    IntervalClosure::RightClosed | IntervalClosure::Closed => Comparison::Le,
794                };
795                compare(values[*value][row].re, lower_op, values[*lower][row].re)
796                    && compare(values[*value][row].re, upper_op, values[*upper][row].re)
797            }
798        }
799    }
800}
801
802fn compare(lhs: f64, op: Comparison, rhs: f64) -> bool {
803    if lhs.is_nan() || rhs.is_nan() {
804        return false;
805    }
806    match op {
807        Comparison::Lt => lhs < rhs,
808        Comparison::Le => lhs <= rhs,
809        Comparison::Gt => lhs > rhs,
810        Comparison::Ge => lhs >= rhs,
811        Comparison::Eq => lhs == rhs,
812        Comparison::Ne => lhs != rhs,
813    }
814}
815
816#[derive(Clone)]
817struct QuerySource {
818    source: Dataset,
819    filter: QueryFilter,
820}
821
822#[derive(Clone)]
823enum QueryFilter {
824    Predicate(Arc<CompiledPredicate>),
825}
826
827impl EventSource for QuerySource {
828    fn schema(&self) -> LadduDataResult<Arc<Schema>> {
829        self.source.schema()
830    }
831
832    fn capabilities(&self) -> SourceCapabilities {
833        let source = self.source.capabilities();
834        SourceCapabilities {
835            exact_len: false,
836            exact_weighted_total: false,
837            random_access: false,
838            deterministic_partitioning: source.deterministic_partitioning,
839            predicate_pushdown: false,
840            projection_pushdown: false,
841            streaming: true,
842        }
843    }
844
845    fn batches(&self, plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
846        let batches = self.source.stream_with_plan(plan)?;
847        let filter = self.filter.clone();
848        Ok(Box::new(batches.filter_map(move |batch| {
849            let batch = match batch {
850                Ok(batch) => batch,
851                Err(error) => return Some(Err(error)),
852            };
853            let rows = match filter.rows(&batch) {
854                Ok(rows) => rows,
855                Err(error) => return Some(Err(LadduDataError::Source(error.to_string()))),
856            };
857            (!rows.is_empty()).then(|| Ok(batch.select(&rows)))
858        })))
859    }
860}
861
862impl QueryFilter {
863    fn rows(&self, batch: &EventBatch) -> RuntimeResult<Vec<usize>> {
864        match self {
865            Self::Predicate(predicate) => predicate.evaluate_batch(batch),
866        }
867    }
868}
869
870fn query_error(message: impl Into<String>) -> RuntimeError {
871    RuntimeError::InvalidShape {
872        index: 0,
873        message: message.into(),
874    }
875}
876
877fn may_fallback_to_scalar(execution: &Execution, error: &RuntimeError) -> bool {
878    let cpu_f32 = matches!(
879        error,
880        RuntimeError::Execution(crate::ExecutionError::UnsupportedCpuF32Model)
881    );
882    #[cfg(feature = "wgpu")]
883    {
884        cpu_f32 || (execution.wgpu_context().is_some() && matches!(error, RuntimeError::Wgpu(_)))
885    }
886    #[cfg(not(feature = "wgpu"))]
887    {
888        let _ = execution;
889        cpu_f32
890    }
891}
892
893fn data_error(error: impl ToString) -> RuntimeError {
894    RuntimeError::Data(error.to_string())
895}
896
897#[cfg(test)]
898mod tests {
899    use super::*;
900    use crate::{CpuOptions, Device, ExecutionOptions, Precision};
901    #[cfg(feature = "jit")]
902    use crate::{JitPolicy, ThreadPolicy};
903    use laddu_compile::CompiledModel;
904    use laddu_data::{
905        data::{EventBatch, OwnedEvent},
906        io::{EventSource, ReadPlan, SourceCapabilities, memory::MemorySource},
907        schema::Schema,
908    };
909    use laddu_expr::{complex, event_scalar};
910    use std::sync::atomic::{AtomicUsize, Ordering};
911
912    #[test]
913    fn bin_spec_roundtrip_preserves_validation() {
914        let bins = BinSpec::edges([-1.0, 0.0, 2.0]).unwrap();
915        let json = serde_json::to_string(&bins).unwrap();
916        assert_eq!(serde_json::from_str::<BinSpec>(&json).unwrap(), bins);
917        assert!(serde_json::from_str::<BinSpec>("[0.0,0.0]").is_err());
918    }
919
920    #[derive(Clone)]
921    struct CountingSource {
922        inner: MemorySource,
923        reads: Arc<AtomicUsize>,
924    }
925
926    impl EventSource for CountingSource {
927        fn schema(&self) -> LadduDataResult<Arc<Schema>> {
928            EventSource::schema(&self.inner)
929        }
930
931        fn capabilities(&self) -> SourceCapabilities {
932            self.inner.capabilities()
933        }
934
935        fn batches(&self, plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
936            self.reads.fetch_add(1, Ordering::Relaxed);
937            self.inner.batches(plan)
938        }
939    }
940
941    #[derive(Clone)]
942    struct FailingSource {
943        schema: Arc<Schema>,
944    }
945
946    impl EventSource for FailingSource {
947        fn schema(&self) -> LadduDataResult<Arc<Schema>> {
948            Ok(Arc::clone(&self.schema))
949        }
950
951        fn capabilities(&self) -> SourceCapabilities {
952            SourceCapabilities {
953                exact_len: false,
954                exact_weighted_total: false,
955                random_access: false,
956                deterministic_partitioning: true,
957                predicate_pushdown: false,
958                projection_pushdown: false,
959                streaming: true,
960            }
961        }
962
963        fn batches(&self, _plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
964            Ok(Box::new(std::iter::once(Err(LadduDataError::Source(
965                "query source failed".into(),
966            )))))
967        }
968    }
969
970    fn capability_tuple(
971        capabilities: SourceCapabilities,
972    ) -> (bool, bool, bool, bool, bool, bool, bool) {
973        (
974            capabilities.exact_len,
975            capabilities.exact_weighted_total,
976            capabilities.random_access,
977            capabilities.deterministic_partitioning,
978            capabilities.predicate_pushdown,
979            capabilities.projection_pushdown,
980            capabilities.streaming,
981        )
982    }
983
984    fn dataset() -> Dataset {
985        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
986        Dataset::from_events(
987            schema,
988            [
989                OwnedEvent::weighted(vec![], vec![-1.0], 0.5),
990                OwnedEvent::weighted(vec![], vec![0.0], 1.0),
991                OwnedEvent::weighted(vec![], vec![1.0], 1.5),
992                OwnedEvent::weighted(vec![], vec![2.0], 2.0),
993            ],
994        )
995        .unwrap()
996    }
997
998    #[test]
999    fn evaluates_selects_and_bins_dataset_expressions() {
1000        let dataset = dataset().chunked(1).unwrap();
1001        let execution = Execution::default();
1002        let x = event_scalar("x");
1003        assert_eq!(
1004            dataset.evaluate_real(&x, &execution).unwrap(),
1005            vec![-1.0, 0.0, 1.0, 2.0]
1006        );
1007
1008        let selected = dataset
1009            .select(
1010                &Predicate::ge(x.clone(), 0.0).and(Predicate::lt(x.clone(), 2.0)),
1011                &execution,
1012            )
1013            .unwrap();
1014        assert_eq!(
1015            selected.map_events(|event| event.scalar(0)).unwrap(),
1016            vec![0.0, 1.0]
1017        );
1018        assert_eq!(selected.sum_weights().unwrap(), 2.5);
1019
1020        let bins = dataset
1021            .bin_by(&x, BinSpec::uniform(2, 0.0, 2.0).unwrap(), &execution)
1022            .unwrap();
1023        assert_eq!(bins.len(), 2);
1024        assert_eq!(
1025            bins[0]
1026                .dataset()
1027                .map_events(|event| event.scalar(0))
1028                .unwrap(),
1029            vec![0.0]
1030        );
1031        assert_eq!(
1032            bins[1]
1033                .dataset()
1034                .map_events(|event| event.scalar(0))
1035                .unwrap(),
1036            vec![1.0, 2.0]
1037        );
1038    }
1039
1040    #[test]
1041    fn dataset_histogram_uses_event_weights_and_excludes_the_final_upper_edge() {
1042        let histogram = dataset()
1043            .histogram(
1044                &event_scalar("x"),
1045                BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1046                true,
1047                None,
1048                &Execution::default(),
1049            )
1050            .unwrap();
1051
1052        assert_eq!(histogram.counts(), [1.0, 1.5]);
1053        assert_eq!(histogram.sum_squared_weights(), [1.0, 2.25]);
1054        assert_eq!(histogram.underflow(), 0.5);
1055        assert_eq!(histogram.underflow_sum_squared_weights(), Some(0.25));
1056        assert_eq!(histogram.overflow(), 2.0);
1057        assert_eq!(histogram.overflow_sum_squared_weights(), Some(4.0));
1058    }
1059
1060    #[test]
1061    fn dataset_joint_histogram_uses_row_major_bins_and_aggregates_invalid_events() {
1062        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x", "y"], true).unwrap());
1063        let dataset = Dataset::from_events(
1064            schema,
1065            [
1066                OwnedEvent::weighted(vec![], vec![0.0, 10.0], 1.0),
1067                OwnedEvent::weighted(vec![], vec![1.0, 10.0], -2.0),
1068                OwnedEvent::weighted(vec![], vec![0.0, 20.0], 3.0),
1069                OwnedEvent::weighted(vec![], vec![f64::NAN, 10.0], 4.0),
1070            ],
1071        )
1072        .unwrap();
1073
1074        let histogram = dataset
1075            .joint_histogram(
1076                &[event_scalar("x"), event_scalar("y")],
1077                vec![
1078                    BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1079                    BinSpec::edges([0.0, 15.0, 25.0]).unwrap(),
1080                ],
1081                true,
1082                None,
1083                &Execution::default(),
1084            )
1085            .unwrap();
1086
1087        assert_eq!(histogram.shape(), [2, 2]);
1088        assert_eq!(histogram.values(), [1.0, 3.0, -2.0, 0.0]);
1089        assert_eq!(histogram.sum_squared_weights(), [1.0, 9.0, 4.0, 0.0]);
1090        assert_eq!(histogram.diagnostics().nonfinite_count(), 1);
1091        assert_eq!(histogram.diagnostics().out_of_range_count(), 0);
1092    }
1093
1094    #[test]
1095    fn dataset_histogram_multiplies_the_selected_base_and_custom_weights() {
1096        let x = event_scalar("x");
1097        let bins = BinSpec::edges([-2.0, 0.0, 2.0]).unwrap();
1098        let execution = Execution::default();
1099
1100        let weighted = dataset()
1101            .histogram(&x, bins.clone(), true, Some(&(x.clone() + 2.0)), &execution)
1102            .unwrap();
1103        assert_eq!(weighted.counts(), [0.5, 6.5]);
1104        assert_eq!(weighted.sum_squared_weights(), [0.25, 24.25]);
1105        assert_eq!(weighted.overflow(), 8.0);
1106
1107        let custom_only = dataset()
1108            .histogram(&x, bins, false, Some(&(x.clone() + 2.0)), &execution)
1109            .unwrap();
1110        assert_eq!(custom_only.counts(), [1.0, 5.0]);
1111        assert_eq!(custom_only.sum_squared_weights(), [1.0, 13.0]);
1112        assert_eq!(custom_only.overflow(), 4.0);
1113    }
1114
1115    #[test]
1116    fn dataset_histogram_preserves_view_source_and_memory_semantics() {
1117        let source = dataset();
1118        let batch = source.batches().unwrap().next().unwrap().unwrap();
1119        let reads = Arc::new(AtomicUsize::new(0));
1120        let counted = Dataset::new(CountingSource {
1121            inner: MemorySource::new(batch),
1122            reads: Arc::clone(&reads),
1123        })
1124        .streaming()
1125        .chunked(1)
1126        .unwrap();
1127        let selected = counted
1128            .select(
1129                &Predicate::ge(event_scalar("x"), 0.0),
1130                &Execution::default(),
1131            )
1132            .unwrap();
1133
1134        let histogram = selected
1135            .histogram(
1136                &event_scalar("x"),
1137                BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1138                true,
1139                None,
1140                &Execution::default(),
1141            )
1142            .unwrap();
1143
1144        assert_eq!(histogram.counts(), [1.0, 1.5]);
1145        assert_eq!(histogram.overflow(), 2.0);
1146        assert_eq!(reads.load(Ordering::Relaxed), 1);
1147    }
1148
1149    #[test]
1150    fn dataset_histogram_matches_across_memory_policies_and_chunking() {
1151        let source = dataset();
1152        let bins = BinSpec::edges([0.0, 1.0, 2.0]).unwrap();
1153        let execution = Execution::default();
1154        let expected = source
1155            .clone()
1156            .resident()
1157            .histogram(&event_scalar("x"), bins.clone(), true, None, &execution)
1158            .unwrap();
1159
1160        for candidate in [
1161            source.clone().streaming(),
1162            source.clone().resident().chunked(1).unwrap(),
1163            source.streaming().chunked(2).unwrap(),
1164        ] {
1165            let actual = candidate
1166                .histogram(&event_scalar("x"), bins.clone(), true, None, &execution)
1167                .unwrap();
1168            assert_eq!(actual, expected);
1169        }
1170    }
1171
1172    #[cfg(feature = "jit")]
1173    #[test]
1174    fn dataset_histogram_matches_cpu_interpreter_and_jit_backends() {
1175        let source = dataset();
1176        let bins = BinSpec::edges([0.0, 1.0, 2.0]).unwrap();
1177        let execution = |jit| {
1178            Execution::local(ExecutionOptions {
1179                device: Device::Cpu(CpuOptions {
1180                    threads: ThreadPolicy::Serial,
1181                    jit,
1182                }),
1183                precision: Precision::F64,
1184                ..ExecutionOptions::default()
1185            })
1186            .unwrap()
1187        };
1188        let interpreted = source
1189            .histogram(
1190                &event_scalar("x"),
1191                bins.clone(),
1192                true,
1193                None,
1194                &execution(JitPolicy::Disabled),
1195            )
1196            .unwrap();
1197        let compiled = source
1198            .histogram(
1199                &event_scalar("x"),
1200                bins,
1201                true,
1202                None,
1203                &execution(JitPolicy::Enabled),
1204            )
1205            .unwrap();
1206
1207        assert_eq!(compiled, interpreted);
1208    }
1209
1210    #[test]
1211    fn dataset_histogram_reports_invalid_values_and_source_failures() {
1212        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1213        let nonfinite = Dataset::from_events(
1214            Arc::clone(&schema),
1215            [OwnedEvent::weighted(vec![], vec![f64::NAN], 1.0)],
1216        )
1217        .unwrap();
1218        let error = nonfinite
1219            .histogram(
1220                &event_scalar("x"),
1221                BinSpec::edges([0.0, 1.0]).unwrap(),
1222                true,
1223                None,
1224                &Execution::default(),
1225            )
1226            .unwrap_err();
1227        assert!(
1228            error.to_string().contains("expected finite, got NaN"),
1229            "unexpected error: {error}",
1230        );
1231
1232        let source_error = Dataset::new(FailingSource { schema })
1233            .histogram(
1234                &event_scalar("x"),
1235                BinSpec::edges([0.0, 1.0]).unwrap(),
1236                true,
1237                None,
1238                &Execution::default(),
1239            )
1240            .unwrap_err();
1241        assert!(source_error.to_string().contains("query source failed"));
1242    }
1243
1244    #[test]
1245    fn empty_dataset_histogram_is_a_valid_empirical_histogram() {
1246        let histogram = dataset()
1247            .empty_derived()
1248            .unwrap()
1249            .histogram(
1250                &event_scalar("x"),
1251                BinSpec::edges([0.0, 1.0, 3.0]).unwrap(),
1252                true,
1253                None,
1254                &Execution::default(),
1255            )
1256            .unwrap();
1257
1258        assert_eq!(histogram.counts(), [0.0, 0.0]);
1259        assert_eq!(histogram.sum_squared_weights(), [0.0, 0.0]);
1260        assert_eq!(histogram.underflow_sum_squared_weights(), Some(0.0));
1261        assert_eq!(histogram.overflow_sum_squared_weights(), Some(0.0));
1262    }
1263
1264    #[test]
1265    fn cancelling_dataset_weights_keep_their_squared_weight_uncertainty() {
1266        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1267        let cancelling = Dataset::from_events(
1268            schema,
1269            [
1270                OwnedEvent::weighted(vec![], vec![0.5], 1.0),
1271                OwnedEvent::weighted(vec![], vec![0.5], -1.0),
1272            ],
1273        )
1274        .unwrap();
1275        let histogram = cancelling
1276            .histogram(
1277                &event_scalar("x"),
1278                BinSpec::edges([0.0, 1.0]).unwrap(),
1279                true,
1280                None,
1281                &Execution::default(),
1282            )
1283            .unwrap();
1284
1285        assert_eq!(histogram.counts(), [0.0]);
1286        assert_eq!(histogram.sum_squared_weights(), [2.0]);
1287        assert_eq!(histogram.errors(), [2.0_f64.sqrt()]);
1288    }
1289
1290    #[test]
1291    fn empty_batches_are_valid_query_inputs() {
1292        let execution = Execution::default();
1293        let x = event_scalar("x");
1294
1295        let empty_batch_schema =
1296            Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1297        let empty_batch = Dataset::from_batch(
1298            EventBatch::from_events(empty_batch_schema, std::iter::empty::<OwnedEvent>()).unwrap(),
1299        );
1300        for empty in [empty_batch, dataset().empty_derived().unwrap()] {
1301            assert!(empty.evaluate_real(&x, &execution).unwrap().is_empty());
1302            assert!(
1303                empty
1304                    .select(&Predicate::ge(x.clone(), 0.0), &execution)
1305                    .unwrap()
1306                    .map_events(|event| event.scalar(0))
1307                    .unwrap()
1308                    .is_empty()
1309            );
1310        }
1311    }
1312
1313    #[test]
1314    fn event_column_nan_comparisons_are_false() {
1315        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1316        let dataset = Dataset::from_events(
1317            schema,
1318            [
1319                OwnedEvent::weighted(vec![], vec![f64::NAN], 1.0),
1320                OwnedEvent::weighted(vec![], vec![0.0], 1.0),
1321                OwnedEvent::weighted(vec![], vec![1.0], 1.0),
1322            ],
1323        )
1324        .unwrap();
1325        let x = event_scalar("x");
1326        let selected = dataset
1327            .select(&Predicate::ne(x, 0.0), &Execution::default())
1328            .unwrap();
1329
1330        assert_eq!(
1331            selected.map_events(|event| event.scalar(0)).unwrap(),
1332            vec![1.0]
1333        );
1334    }
1335
1336    #[test]
1337    fn query_propagates_source_batch_errors() {
1338        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1339        let dataset = Dataset::new(FailingSource { schema });
1340
1341        let error = dataset
1342            .evaluate_real(&event_scalar("x"), &Execution::default())
1343            .unwrap_err();
1344        assert!(
1345            matches!(error, RuntimeError::Data(message) if message.contains("query source failed"))
1346        );
1347    }
1348
1349    #[test]
1350    fn all_empty_bins_retain_valid_empty_derived_sources() {
1351        let source = dataset();
1352        let before = capability_tuple(source.capabilities());
1353        let bins = source
1354            .bin_by(
1355                &event_scalar("x"),
1356                BinSpec::edges([10.0, 20.0, 30.0]).unwrap(),
1357                &Execution::default(),
1358            )
1359            .unwrap();
1360
1361        assert_eq!(capability_tuple(source.capabilities()), before);
1362        assert_eq!(bins.len(), 2);
1363        for bin in bins {
1364            assert_eq!(bin.dataset().num_events().unwrap(), Some(0));
1365            assert!(
1366                bin.dataset()
1367                    .evaluate_real(&event_scalar("x"), &Execution::default())
1368                    .unwrap()
1369                    .is_empty()
1370            );
1371        }
1372    }
1373
1374    #[test]
1375    fn traversing_all_bins_reads_the_source_once() {
1376        let reads = Arc::new(AtomicUsize::new(0));
1377        let source = CountingSource {
1378            inner: match dataset().batches().unwrap().next().unwrap() {
1379                Ok(batch) => MemorySource::new(batch),
1380                Err(error) => panic!("unexpected source error: {error}"),
1381            },
1382            reads: Arc::clone(&reads),
1383        };
1384        let dataset = Dataset::new(source).chunked(1).unwrap();
1385        let bins = dataset
1386            .bin_by(
1387                &event_scalar("x"),
1388                BinSpec::uniform(4, -1.0, 3.0).unwrap(),
1389                &Execution::default(),
1390            )
1391            .unwrap();
1392
1393        let values = bins
1394            .into_iter()
1395            .map(|bin| {
1396                bin.into_dataset()
1397                    .map_events(|event| event.scalar(0))
1398                    .unwrap()
1399            })
1400            .collect::<Vec<_>>();
1401        assert_eq!(values, [vec![-1.0], vec![0.0], vec![1.0], vec![2.0]]);
1402        assert_eq!(reads.load(Ordering::Relaxed), 1);
1403    }
1404
1405    #[test]
1406    fn between_predicates_have_explicit_endpoint_semantics() {
1407        let dataset = dataset();
1408        let execution = Execution::default();
1409        let x = event_scalar("x");
1410
1411        let closed = dataset
1412            .select(&Predicate::between(x.clone(), 0.0, 1.0), &execution)
1413            .unwrap();
1414        assert_eq!(
1415            closed.map_events(|event| event.scalar(0)).unwrap(),
1416            vec![0.0, 1.0]
1417        );
1418
1419        let open = dataset
1420            .select(
1421                &Predicate::between_with(x, -1.0, 1.0, IntervalClosure::Open),
1422                &execution,
1423            )
1424            .unwrap();
1425        assert_eq!(open.map_events(|event| event.scalar(0)).unwrap(), vec![0.0]);
1426    }
1427
1428    #[test]
1429    fn real_queries_reject_complex_and_free_parameter_expressions() {
1430        let dataset = dataset();
1431        let execution = Execution::default();
1432        assert!(
1433            dataset
1434                .evaluate_real(&complex(1.0, 1.0), &execution)
1435                .is_err()
1436        );
1437        let parameter = Expr::from(laddu_expr::parameters::Parameter::free("p"));
1438        assert!(dataset.evaluate_expr(&parameter, &execution).is_err());
1439    }
1440
1441    #[test]
1442    fn compiled_query_outputs_preserve_order_and_values() {
1443        let source = dataset();
1444        let batch = source.batches().unwrap().next().unwrap().unwrap();
1445        let x = event_scalar("x");
1446        let query = PreparedQuery::prepare(
1447            vec![x.clone() + 1.0, x.clone() * 2.0, x],
1448            &Execution::default(),
1449            false,
1450        )
1451        .unwrap();
1452        let values = query.evaluate_batch(&batch).unwrap();
1453        assert_eq!(
1454            values[0].iter().map(|v| v.re).collect::<Vec<_>>(),
1455            [0.0, 1.0, 2.0, 3.0]
1456        );
1457        assert_eq!(
1458            values[1].iter().map(|v| v.re).collect::<Vec<_>>(),
1459            [-2.0, 0.0, 2.0, 4.0]
1460        );
1461        assert_eq!(
1462            values[2].iter().map(|v| v.re).collect::<Vec<_>>(),
1463            [-1.0, 0.0, 1.0, 2.0]
1464        );
1465    }
1466
1467    #[test]
1468    fn repeated_predicate_leaves_are_evaluated_once() {
1469        let x = event_scalar("x");
1470        let selected = dataset()
1471            .select(
1472                &Predicate::ge(x.clone() + 1.0, 0.0).and(Predicate::lt(x + 1.0, 2.0)),
1473                &Execution::default(),
1474            )
1475            .unwrap();
1476        assert_eq!(
1477            selected.map_events(|event| event.scalar(0)).unwrap(),
1478            [-1.0, 0.0]
1479        );
1480    }
1481
1482    #[test]
1483    fn f32_queries_match_f64_query_results() {
1484        let x = event_scalar("x");
1485        let f64_values = dataset().evaluate_real(&x, &Execution::default()).unwrap();
1486        let f32_execution = Execution::local(ExecutionOptions {
1487            device: Device::Cpu(CpuOptions::default()),
1488            precision: Precision::F32,
1489            ..ExecutionOptions::default()
1490        })
1491        .unwrap();
1492        let f32_values = dataset().evaluate_real(&x, &f32_execution).unwrap();
1493        assert_eq!(f32_values, f64_values);
1494
1495        let bins = BinSpec::edges([-2.0, 0.0, 2.0]).unwrap();
1496        let f64_histogram = dataset()
1497            .histogram(&x, bins.clone(), true, None, &Execution::default())
1498            .unwrap();
1499        let f32_histogram = dataset()
1500            .histogram(&x, bins, true, None, &f32_execution)
1501            .unwrap();
1502        assert_eq!(f32_histogram, f64_histogram);
1503    }
1504
1505    #[test]
1506    fn bin_edges_validate_and_nan_predicates_are_false() {
1507        assert!(BinSpec::edges([0.0, 0.0]).is_err());
1508        assert!(!compare(f64::NAN, Comparison::Ne, 0.0));
1509    }
1510
1511    #[test]
1512    fn selection_is_lazy_and_one_pass_binning_preserves_streaming_policy() {
1513        let source = dataset();
1514        let batch = source.batches().unwrap().next().unwrap().unwrap();
1515        let reads = Arc::new(AtomicUsize::new(0));
1516        let dataset = Dataset::new(CountingSource {
1517            inner: MemorySource::new(batch),
1518            reads: Arc::clone(&reads),
1519        })
1520        .streaming();
1521        let execution = Execution::default();
1522        let x = event_scalar("x");
1523
1524        let selected = dataset
1525            .select(&Predicate::ge(x.clone(), 0.0), &execution)
1526            .unwrap();
1527        let bins = dataset
1528            .bin_by(&x, BinSpec::uniform(2, 0.0, 2.0).unwrap(), &execution)
1529            .unwrap();
1530        assert_eq!(reads.load(Ordering::Relaxed), 1);
1531        assert_eq!(
1532            selected.cache_storage(),
1533            laddu_data::data::CacheStorage::Streaming
1534        );
1535
1536        assert_eq!(
1537            selected.map_events(|event| event.scalar(0)).unwrap(),
1538            vec![0.0, 1.0, 2.0]
1539        );
1540        assert_eq!(reads.load(Ordering::Relaxed), 2);
1541        assert_eq!(
1542            bins[0]
1543                .dataset()
1544                .map_events(|event| event.scalar(0))
1545                .unwrap(),
1546            vec![0.0]
1547        );
1548        assert_eq!(reads.load(Ordering::Relaxed), 2);
1549    }
1550
1551    #[test]
1552    fn unknown_cardinality_fastest_discovers_and_retains_small_selection() {
1553        let source = dataset();
1554        let batch = source.batches().unwrap().next().unwrap().unwrap();
1555        let reads = Arc::new(AtomicUsize::new(0));
1556        let dataset = Dataset::new(CountingSource {
1557            inner: MemorySource::new(batch),
1558            reads: Arc::clone(&reads),
1559        });
1560        let execution = Execution::default();
1561        let x = event_scalar("x");
1562        let selected = dataset
1563            .select(&Predicate::ge(x.clone(), 0.0), &execution)
1564            .unwrap();
1565        let compiled = CompiledModel::from_expr(&x).unwrap();
1566        let params = compiled.params().default_values();
1567        let model = PreparedModel::prepare(&compiled, &execution).unwrap();
1568        let prepared = model.prepare_dataset(&execution, &selected).unwrap();
1569
1570        #[cfg(not(feature = "wgpu"))]
1571        let crate::PreparedDataset::Cpu(prepared_cpu) = &prepared;
1572        #[cfg(feature = "wgpu")]
1573        let crate::PreparedDataset::Cpu(prepared_cpu) = &prepared else {
1574            panic!("default execution prepares CPU datasets");
1575        };
1576        assert_eq!(
1577            prepared_cpu.stats().storage(),
1578            laddu_data::data::CacheStorage::Resident
1579        );
1580        assert_eq!(prepared_cpu.stats().local_events(), 3);
1581        assert_eq!(reads.load(Ordering::Relaxed), 2);
1582
1583        for _ in 0..2 {
1584            assert_eq!(
1585                model
1586                    .reduce(
1587                        &execution,
1588                        &params,
1589                        &prepared,
1590                        laddu_compile::ReductionPlan::weighted_real(),
1591                    )
1592                    .unwrap(),
1593                5.5
1594            );
1595        }
1596        assert_eq!(reads.load(Ordering::Relaxed), 2);
1597    }
1598}