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    /// Evaluates ordered scalar expressions in one dataset traversal.
268    ///
269    /// Each expression retains the preparation and execution behavior of an
270    /// independent evaluation. Empty input returns no columns without reading.
271    ///
272    /// # Errors
273    /// Returns an error for invalid expressions, failed preparation, source
274    /// reads, or evaluation. `require_real` rejects complex-valued expressions.
275    fn evaluate_exprs(
276        &self,
277        expressions: &[Expr],
278        execution: &Execution,
279        require_real: bool,
280    ) -> RuntimeResult<Vec<Vec<Complex64>>>;
281    /// Visits real scalar expression values in bounded event chunks.
282    /// The offset is the first event's global row number and expressions retain
283    /// their requested order within each callback.
284    ///
285    /// # Errors
286    /// Returns an error for zero chunk size, invalid expressions, dataset access,
287    /// evaluation, or a callback failure.
288    fn visit_real_chunks(
289        &self,
290        expressions: &[Expr],
291        execution: &Execution,
292        chunk_size: usize,
293        consume: impl FnMut(usize, &[Vec<f64>]) -> RuntimeResult<()>,
294    ) -> RuntimeResult<()>;
295    /// Evaluates and fills one weighted histogram in a bounded dataset traversal.
296    ///
297    /// Dataset event weights are used when `event_weights` is true. When
298    /// `weight` is supplied, its real scalar value multiplies the selected
299    /// event weight (or unit weight when event weights are disabled).
300    ///
301    /// # Errors
302    ///
303    /// Returns [`RuntimeError`] when bin validation, expression preparation,
304    /// dataset reading, evaluation, or histogram filling fails.
305    fn histogram(
306        &self,
307        expr: &Expr,
308        bins: BinSpec,
309        event_weights: bool,
310        weight: Option<&Expr>,
311        execution: &Execution,
312    ) -> RuntimeResult<Histogram>;
313    /// Evaluates ordered scalar axes and fills one bounded row-major joint histogram.
314    ///
315    /// # Errors
316    ///
317    /// Returns [`RuntimeError`] for missing or mismatched axes, invalid edges,
318    /// non-scalar expressions, evaluation failures, or source read failures.
319    fn joint_histogram(
320        &self,
321        axes: &[Expr],
322        bins: Vec<BinSpec>,
323        event_weights: bool,
324        weight: Option<&Expr>,
325        execution: &Execution,
326    ) -> RuntimeResult<JointHistogram>;
327    /// Creates a lazily filtered dataset containing events that satisfy `predicate`.
328    ///
329    /// # Errors
330    ///
331    /// Returns [`RuntimeError`] when predicate compilation or evaluation
332    /// fails, or its expression is not real scalar-valued.
333    fn select(&self, predicate: &Predicate, execution: &Execution) -> RuntimeResult<Dataset>;
334    /// Partitions the dataset in one pass according to an expression and bin specification.
335    ///
336    /// # Errors
337    ///
338    /// Returns [`RuntimeError`] when expression compilation or evaluation
339    /// fails, or the expression is not real scalar-valued.
340    fn bin_by(
341        &self,
342        expr: &Expr,
343        bins: BinSpec,
344        execution: &Execution,
345    ) -> RuntimeResult<Vec<DatasetBin>>;
346}
347
348impl DatasetExprExt for Dataset {
349    fn validate_real_expressions(
350        &self,
351        expressions: &[Expr],
352        execution: &Execution,
353    ) -> RuntimeResult<()> {
354        if expressions.is_empty() {
355            return Err(query_error(
356                "real expression validation needs at least one expression",
357            ));
358        }
359        PreparedQuery::prepare(expressions.to_vec(), execution, true)?;
360        Ok(())
361    }
362
363    fn evaluate_expr(&self, expr: &Expr, execution: &Execution) -> RuntimeResult<Vec<Complex64>> {
364        let query = PreparedQuery::prepare(vec![expr.clone()], execution, false)?;
365        let mut output = Vec::new();
366        for batch in self.batches().map_err(data_error)? {
367            output.extend(
368                query.evaluate_batch(&batch.map_err(data_error)?)?[0]
369                    .iter()
370                    .copied(),
371            );
372        }
373        Ok(output)
374    }
375
376    fn evaluate_real(&self, expr: &Expr, execution: &Execution) -> RuntimeResult<Vec<f64>> {
377        let query = PreparedQuery::prepare(vec![expr.clone()], execution, true)?;
378        let mut output = Vec::new();
379        for batch in self.batches().map_err(data_error)? {
380            output.extend(
381                query.evaluate_batch(&batch.map_err(data_error)?)?[0]
382                    .iter()
383                    .copied()
384                    .map(|v| v.re),
385            );
386        }
387        Ok(output)
388    }
389
390    fn evaluate_exprs(
391        &self,
392        expressions: &[Expr],
393        execution: &Execution,
394        require_real: bool,
395    ) -> RuntimeResult<Vec<Vec<Complex64>>> {
396        if expressions.is_empty() {
397            return Ok(Vec::new());
398        }
399        let queries = expressions
400            .iter()
401            .map(|expr| PreparedQuery::prepare(vec![expr.clone()], execution, require_real))
402            .collect::<RuntimeResult<Vec<_>>>()?;
403        let mut outputs = vec![Vec::new(); expressions.len()];
404        for batch in self.batches().map_err(data_error)? {
405            let batch = batch.map_err(data_error)?;
406            for (output, query) in outputs.iter_mut().zip(&queries) {
407                output.extend(query.evaluate_batch(&batch)?[0].iter().copied());
408            }
409        }
410        Ok(outputs)
411    }
412
413    fn visit_real_chunks(
414        &self,
415        expressions: &[Expr],
416        execution: &Execution,
417        chunk_size: usize,
418        mut consume: impl FnMut(usize, &[Vec<f64>]) -> RuntimeResult<()>,
419    ) -> RuntimeResult<()> {
420        if chunk_size == 0 || expressions.is_empty() {
421            return Err(query_error(
422                "real chunk evaluation needs expressions and positive chunk size",
423            ));
424        }
425        let query = PreparedQuery::prepare(expressions.to_vec(), execution, true)?;
426        let mut offset = 0;
427        for batch in self.batches().map_err(data_error)? {
428            let batch = batch.map_err(data_error)?;
429            for start in (0..batch.len()).step_by(chunk_size) {
430                let end = (start + chunk_size).min(batch.len());
431                let values = query
432                    .evaluate_batch(&batch.slice(start, end))?
433                    .into_iter()
434                    .map(|column| column.into_iter().map(|value| value.re).collect::<Vec<_>>())
435                    .collect::<Vec<_>>();
436                consume(offset, &values)?;
437                offset += end - start;
438            }
439        }
440        Ok(())
441    }
442
443    fn histogram(
444        &self,
445        expr: &Expr,
446        bins: BinSpec,
447        event_weights: bool,
448        weight: Option<&Expr>,
449        execution: &Execution,
450    ) -> RuntimeResult<Histogram> {
451        let mut expressions = vec![expr.clone()];
452        expressions.extend(weight.cloned());
453        let query = PreparedQuery::prepare(expressions, execution, true)?;
454        let mut histogram = Histogram::empty_with_edges(bins.edges_slice().to_vec())
455            .map_err(|error| query_error(error.to_string()))?;
456
457        for batch in self.batches().map_err(data_error)? {
458            let batch = batch.map_err(data_error)?;
459            let values = query.evaluate_batch(&batch)?;
460            for row in 0..batch.len() {
461                let base_weight = if event_weights {
462                    batch.weights_at(row)
463                } else {
464                    1.0
465                };
466                let custom_weight = values.get(1).map_or(1.0, |weights| weights[row].re);
467                histogram
468                    .fill_weighted(values[0][row].re, base_weight * custom_weight)
469                    .map_err(|error| query_error(error.to_string()))?;
470            }
471        }
472        Ok(histogram)
473    }
474
475    fn joint_histogram(
476        &self,
477        axes: &[Expr],
478        bins: Vec<BinSpec>,
479        event_weights: bool,
480        weight: Option<&Expr>,
481        execution: &Execution,
482    ) -> RuntimeResult<JointHistogram> {
483        if axes.is_empty() || axes.len() != bins.len() {
484            return Err(query_error(
485                "joint histogram requires one bin specification per non-empty ordered axis",
486            ));
487        }
488        let edge_vectors = bins
489            .iter()
490            .map(|bins| bins.edges_slice().to_vec())
491            .collect();
492        let mut histogram =
493            JointHistogram::empty(edge_vectors).map_err(|error| query_error(error.to_string()))?;
494        let mut expressions = axes.to_vec();
495        expressions.extend(weight.cloned());
496        let query = PreparedQuery::prepare(expressions, execution, true)?;
497        let mut coordinates = vec![0.0; axes.len()];
498        for batch in self.batches().map_err(data_error)? {
499            let batch = batch.map_err(data_error)?;
500            let values = query.evaluate_batch(&batch)?;
501            for row in 0..batch.len() {
502                for (axis, coordinate) in coordinates.iter_mut().enumerate() {
503                    *coordinate = values[axis][row].re;
504                }
505                let base_weight = if event_weights {
506                    batch.weights_at(row)
507                } else {
508                    1.0
509                };
510                let custom_weight = values
511                    .get(axes.len())
512                    .map_or(1.0, |weights| weights[row].re);
513                histogram
514                    .fill_weighted(&coordinates, base_weight * custom_weight)
515                    .map_err(|error| query_error(error.to_string()))?;
516            }
517        }
518        Ok(histogram)
519    }
520
521    fn select(&self, predicate: &Predicate, execution: &Execution) -> RuntimeResult<Dataset> {
522        let compiled = CompiledPredicate::prepare(predicate, execution)?;
523        Ok(self.with_derived_source(QuerySource {
524            source: self.clone(),
525            filter: QueryFilter::Predicate(Arc::new(compiled)),
526        }))
527    }
528
529    fn bin_by(
530        &self,
531        expr: &Expr,
532        bins: BinSpec,
533        execution: &Execution,
534    ) -> RuntimeResult<Vec<DatasetBin>> {
535        let query = PreparedQuery::prepare(vec![expr.clone()], execution, true)?;
536        let schema = self.schema().map_err(data_error)?;
537        let mut partitions = vec![Vec::new(); bins.bin_count()];
538        for batch in self.batches().map_err(data_error)? {
539            let batch = batch.map_err(data_error)?;
540            let mut rows = vec![Vec::new(); bins.bin_count()];
541            for (row, value) in query.evaluate_batch(&batch)?[0].iter().copied().enumerate() {
542                if let Some(index) = bins.index(value.re) {
543                    rows[index].push(row);
544                }
545            }
546            for (partition, rows) in partitions.iter_mut().zip(rows) {
547                if !rows.is_empty() {
548                    partition.push(batch.select(&rows));
549                }
550            }
551        }
552
553        partitions
554            .into_iter()
555            .enumerate()
556            .map(|(index, batches)| {
557                let source = if batches.is_empty() {
558                    MemorySource::empty(Arc::clone(&schema))
559                } else {
560                    MemorySource::from_batches(batches).map_err(data_error)?
561                };
562                Ok(DatasetBin {
563                    index,
564                    lower: bins.edges_slice()[index],
565                    upper: bins.edges_slice()[index + 1],
566                    dataset: self.with_derived_source(source),
567                })
568            })
569            .collect()
570    }
571}
572
573struct QueryExpr {
574    model: PreparedModel,
575    params: laddu_expr::parameters::ParamValues,
576    outputs: Vec<laddu_expr::ExprId>,
577}
578
579/// Scalar expressions compiled and prepared once for repeated batch evaluation.
580/// This object retains no event data and owns no cross-request compilation cache.
581pub struct PreparedQuery {
582    shared: QueryExprStorage,
583    outputs: usize,
584}
585
586enum QueryExprStorage {
587    Shared(QueryExpr),
588    Separate(Vec<QueryExpr>),
589}
590
591impl PreparedQuery {
592    /// Compile scalar outputs in caller order for the selected backend.
593    /// Set `require_real` to reject complex outputs. Free parameters are rejected.
594    ///
595    /// # Errors
596    /// Returns an error for invalid expressions or backend preparation failure.
597    pub fn prepare(
598        expressions: Vec<Expr>,
599        execution: &Execution,
600        require_real: bool,
601    ) -> RuntimeResult<Self> {
602        let expression_count = expressions.len();
603        let compiled = CompiledQuery::from_exprs(expressions.clone())
604            .map_err(|error| query_error(error.to_string()))?;
605        let model = compiled.model();
606        let outputs = compiled.outputs();
607        if outputs.len() != expression_count {
608            return Err(query_error(
609                "compiled query output count changed during lowering",
610            ));
611        }
612        for element in outputs {
613            let value_kind = model
614                .node_facts(*element)
615                .map(|facts| facts.value_kind)
616                .ok_or_else(|| query_error("compiled expression facts are incomplete"))?;
617            if value_kind != ValueKind::Real
618                && (require_real || !matches!(value_kind, ValueKind::Complex))
619            {
620                return Err(query_error(if require_real {
621                    "this dataset operation requires a real-valued expression"
622                } else {
623                    "dataset expressions must be scalar"
624                }));
625            }
626        }
627        if model.params().n_free() != 0 {
628            return Err(query_error(
629                "dataset expressions cannot contain free parameters",
630            ));
631        }
632        let params = model.params().default_values();
633        let plan = match PreparedModel::prepare(model, execution) {
634            Ok(plan) => QueryExprStorage::Shared(QueryExpr {
635                model: plan,
636                params,
637                outputs: outputs.to_vec(),
638            }),
639            Err(shared_error) => {
640                if !may_fallback_to_scalar(execution, &shared_error) {
641                    return Err(shared_error);
642                }
643                let separate = expressions
644                    .iter()
645                    .map(|expr| QueryExpr::prepare(expr, execution, require_real))
646                    .collect::<RuntimeResult<Vec<_>>>();
647                QueryExprStorage::Separate(separate?)
648            }
649        };
650        Ok(Self {
651            shared: plan,
652            outputs: expression_count,
653        })
654    }
655
656    /// Estimate peak host workspace for one batch, including output columns.
657    #[doc(hidden)]
658    pub fn batch_memory_estimate(&self, events: usize) -> usize {
659        let model = match &self.shared {
660            QueryExprStorage::Shared(query) => query.model.batch_memory_estimate(events),
661            QueryExprStorage::Separate(queries) => queries
662                .iter()
663                .map(|query| query.model.batch_memory_estimate(events))
664                .max()
665                .unwrap_or(0),
666        };
667        model.saturating_add(events.saturating_mul(self.outputs).saturating_mul(32))
668    }
669
670    /// Evaluate every output on a batch without retaining event values.
671    ///
672    /// # Errors
673    /// Returns an error for incompatible event columns or failed evaluation.
674    pub fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<Vec<Complex64>>> {
675        let values = match &self.shared {
676            QueryExprStorage::Shared(query) => {
677                query
678                    .model
679                    .evaluate_batch_outputs(&query.params, batch, &query.outputs)?
680            }
681            QueryExprStorage::Separate(queries) => queries
682                .iter()
683                .map(|query| query.evaluate_batch(batch))
684                .collect::<RuntimeResult<Vec<_>>>()?,
685        };
686        if values.len() != self.outputs {
687            return Err(query_error(
688                "compiled query returned an unexpected output count",
689            ));
690        }
691        Ok(values)
692    }
693}
694
695impl QueryExpr {
696    fn prepare(expr: &Expr, execution: &Execution, require_real: bool) -> RuntimeResult<Self> {
697        if expr.shape().map_err(|e| query_error(e.to_string()))? != laddu_expr::ExprShape::Scalar {
698            return Err(query_error("dataset expressions must be scalar"));
699        }
700        let compiled = laddu_compile::CompiledModel::from_expr(expr)
701            .map_err(|error| query_error(error.to_string()))?;
702        let value_kind = compiled
703            .node_facts(compiled.graph().root())
704            .map(|facts| facts.value_kind)
705            .ok_or_else(|| query_error("compiled expression facts are incomplete"))?;
706        if require_real && value_kind == ValueKind::Complex {
707            return Err(query_error(
708                "this dataset operation requires a real-valued expression",
709            ));
710        }
711        if compiled.params().n_free() != 0 {
712            return Err(query_error(
713                "dataset expressions cannot contain free parameters",
714            ));
715        }
716        let params = compiled.params().default_values();
717        let model = PreparedModel::prepare(&compiled, execution)?;
718        Ok(Self {
719            model,
720            params,
721            outputs: Vec::new(),
722        })
723    }
724
725    fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<Complex64>> {
726        self.model.evaluate_batch(&self.params, batch)
727    }
728}
729
730struct CompiledPredicate {
731    expressions: PreparedQuery,
732    program: PredicateProgram,
733}
734
735enum PredicateProgram {
736    Compare {
737        lhs: usize,
738        op: Comparison,
739        rhs: usize,
740    },
741    And(Box<Self>, Box<Self>),
742    Or(Box<Self>, Box<Self>),
743    Not(Box<Self>),
744    Between {
745        value: usize,
746        lower: usize,
747        upper: usize,
748        closure: IntervalClosure,
749    },
750}
751
752impl CompiledPredicate {
753    fn prepare(predicate: &Predicate, execution: &Execution) -> RuntimeResult<Self> {
754        let mut expressions = Vec::new();
755        let program = Self::compile_program(predicate, &mut expressions);
756        Ok(Self {
757            expressions: PreparedQuery::prepare(expressions, execution, true)?,
758            program,
759        })
760    }
761
762    fn compile_program(predicate: &Predicate, expressions: &mut Vec<Expr>) -> PredicateProgram {
763        let leaf = |expr: &Expr, expressions: &mut Vec<Expr>| {
764            let index = expressions.len();
765            expressions.push(expr.clone());
766            index
767        };
768        match predicate {
769            Predicate::Compare { lhs, op, rhs } => PredicateProgram::Compare {
770                lhs: leaf(lhs, expressions),
771                op: *op,
772                rhs: leaf(rhs, expressions),
773            },
774            Predicate::And(lhs, rhs) => PredicateProgram::And(
775                Box::new(Self::compile_program(lhs, expressions)),
776                Box::new(Self::compile_program(rhs, expressions)),
777            ),
778            Predicate::Or(lhs, rhs) => PredicateProgram::Or(
779                Box::new(Self::compile_program(lhs, expressions)),
780                Box::new(Self::compile_program(rhs, expressions)),
781            ),
782            Predicate::Not(inner) => {
783                PredicateProgram::Not(Box::new(Self::compile_program(inner, expressions)))
784            }
785            Predicate::Between {
786                value,
787                lower,
788                upper,
789                closure,
790            } => PredicateProgram::Between {
791                value: leaf(value, expressions),
792                lower: leaf(lower, expressions),
793                upper: leaf(upper, expressions),
794                closure: *closure,
795            },
796        }
797    }
798
799    fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<usize>> {
800        let values = self.expressions.evaluate_batch(batch)?;
801        Ok((0..batch.len())
802            .filter(|row| Self::evaluate_row(&self.program, &values, *row))
803            .collect())
804    }
805
806    fn evaluate_row(program: &PredicateProgram, values: &[Vec<Complex64>], row: usize) -> bool {
807        match program {
808            PredicateProgram::Compare { lhs, op, rhs } => {
809                compare(values[*lhs][row].re, *op, values[*rhs][row].re)
810            }
811            PredicateProgram::And(lhs, rhs) => {
812                Self::evaluate_row(lhs, values, row) && Self::evaluate_row(rhs, values, row)
813            }
814            PredicateProgram::Or(lhs, rhs) => {
815                Self::evaluate_row(lhs, values, row) || Self::evaluate_row(rhs, values, row)
816            }
817            PredicateProgram::Not(inner) => !Self::evaluate_row(inner, values, row),
818            PredicateProgram::Between {
819                value,
820                lower,
821                upper,
822                closure,
823            } => {
824                let lower_op = match closure {
825                    IntervalClosure::Open | IntervalClosure::RightClosed => Comparison::Gt,
826                    IntervalClosure::LeftClosed | IntervalClosure::Closed => Comparison::Ge,
827                };
828                let upper_op = match closure {
829                    IntervalClosure::Open | IntervalClosure::LeftClosed => Comparison::Lt,
830                    IntervalClosure::RightClosed | IntervalClosure::Closed => Comparison::Le,
831                };
832                compare(values[*value][row].re, lower_op, values[*lower][row].re)
833                    && compare(values[*value][row].re, upper_op, values[*upper][row].re)
834            }
835        }
836    }
837}
838
839fn compare(lhs: f64, op: Comparison, rhs: f64) -> bool {
840    if lhs.is_nan() || rhs.is_nan() {
841        return false;
842    }
843    match op {
844        Comparison::Lt => lhs < rhs,
845        Comparison::Le => lhs <= rhs,
846        Comparison::Gt => lhs > rhs,
847        Comparison::Ge => lhs >= rhs,
848        Comparison::Eq => lhs == rhs,
849        Comparison::Ne => lhs != rhs,
850    }
851}
852
853#[derive(Clone)]
854struct QuerySource {
855    source: Dataset,
856    filter: QueryFilter,
857}
858
859#[derive(Clone)]
860enum QueryFilter {
861    Predicate(Arc<CompiledPredicate>),
862}
863
864impl EventSource for QuerySource {
865    fn schema(&self) -> LadduDataResult<Arc<Schema>> {
866        self.source.schema()
867    }
868
869    fn capabilities(&self) -> SourceCapabilities {
870        let source = self.source.capabilities();
871        SourceCapabilities {
872            exact_len: false,
873            exact_weighted_total: false,
874            random_access: false,
875            deterministic_partitioning: source.deterministic_partitioning,
876            predicate_pushdown: false,
877            projection_pushdown: false,
878            streaming: true,
879        }
880    }
881
882    fn batches(&self, plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
883        let batches = self.source.stream_with_plan(plan)?;
884        let filter = self.filter.clone();
885        Ok(Box::new(batches.filter_map(move |batch| {
886            let batch = match batch {
887                Ok(batch) => batch,
888                Err(error) => return Some(Err(error)),
889            };
890            let rows = match filter.rows(&batch) {
891                Ok(rows) => rows,
892                Err(error) => return Some(Err(LadduDataError::Source(error.to_string()))),
893            };
894            (!rows.is_empty()).then(|| Ok(batch.select(&rows)))
895        })))
896    }
897}
898
899impl QueryFilter {
900    fn rows(&self, batch: &EventBatch) -> RuntimeResult<Vec<usize>> {
901        match self {
902            Self::Predicate(predicate) => predicate.evaluate_batch(batch),
903        }
904    }
905}
906
907fn query_error(message: impl Into<String>) -> RuntimeError {
908    RuntimeError::InvalidShape {
909        index: 0,
910        message: message.into(),
911    }
912}
913
914fn may_fallback_to_scalar(execution: &Execution, error: &RuntimeError) -> bool {
915    let cpu_f32 = matches!(
916        error,
917        RuntimeError::Execution(crate::ExecutionError::UnsupportedCpuF32Model)
918    );
919    #[cfg(feature = "wgpu")]
920    {
921        cpu_f32 || (execution.wgpu_context().is_some() && matches!(error, RuntimeError::Wgpu(_)))
922    }
923    #[cfg(not(feature = "wgpu"))]
924    {
925        let _ = execution;
926        cpu_f32
927    }
928}
929
930fn data_error(error: impl ToString) -> RuntimeError {
931    RuntimeError::Data(error.to_string())
932}
933
934#[cfg(test)]
935mod tests {
936    use super::*;
937    use crate::{CpuOptions, Device, ExecutionOptions, Precision};
938    #[cfg(feature = "jit")]
939    use crate::{JitPolicy, ThreadPolicy};
940    use laddu_compile::CompiledModel;
941    use laddu_data::{
942        data::{EventBatch, OwnedEvent},
943        io::{EventSource, ReadPlan, SourceCapabilities, memory::MemorySource},
944        schema::Schema,
945    };
946    use laddu_expr::{complex, event_scalar};
947    use std::sync::atomic::{AtomicUsize, Ordering};
948
949    #[test]
950    fn bin_spec_roundtrip_preserves_validation() {
951        let bins = BinSpec::edges([-1.0, 0.0, 2.0]).unwrap();
952        let json = serde_json::to_string(&bins).unwrap();
953        assert_eq!(serde_json::from_str::<BinSpec>(&json).unwrap(), bins);
954        assert!(serde_json::from_str::<BinSpec>("[0.0,0.0]").is_err());
955    }
956
957    #[derive(Clone)]
958    struct CountingSource {
959        inner: MemorySource,
960        reads: Arc<AtomicUsize>,
961    }
962
963    impl EventSource for CountingSource {
964        fn schema(&self) -> LadduDataResult<Arc<Schema>> {
965            EventSource::schema(&self.inner)
966        }
967
968        fn capabilities(&self) -> SourceCapabilities {
969            self.inner.capabilities()
970        }
971
972        fn batches(&self, plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
973            self.reads.fetch_add(1, Ordering::Relaxed);
974            self.inner.batches(plan)
975        }
976    }
977
978    #[derive(Clone)]
979    struct FailingSource {
980        schema: Arc<Schema>,
981    }
982
983    impl EventSource for FailingSource {
984        fn schema(&self) -> LadduDataResult<Arc<Schema>> {
985            Ok(Arc::clone(&self.schema))
986        }
987
988        fn capabilities(&self) -> SourceCapabilities {
989            SourceCapabilities {
990                exact_len: false,
991                exact_weighted_total: false,
992                random_access: false,
993                deterministic_partitioning: true,
994                predicate_pushdown: false,
995                projection_pushdown: false,
996                streaming: true,
997            }
998        }
999
1000        fn batches(&self, _plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
1001            Ok(Box::new(std::iter::once(Err(LadduDataError::Source(
1002                "query source failed".into(),
1003            )))))
1004        }
1005    }
1006
1007    fn capability_tuple(
1008        capabilities: SourceCapabilities,
1009    ) -> (bool, bool, bool, bool, bool, bool, bool) {
1010        (
1011            capabilities.exact_len,
1012            capabilities.exact_weighted_total,
1013            capabilities.random_access,
1014            capabilities.deterministic_partitioning,
1015            capabilities.predicate_pushdown,
1016            capabilities.projection_pushdown,
1017            capabilities.streaming,
1018        )
1019    }
1020
1021    fn dataset() -> Dataset {
1022        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1023        Dataset::from_events(
1024            schema,
1025            [
1026                OwnedEvent::weighted(vec![], vec![-1.0], 0.5),
1027                OwnedEvent::weighted(vec![], vec![0.0], 1.0),
1028                OwnedEvent::weighted(vec![], vec![1.0], 1.5),
1029                OwnedEvent::weighted(vec![], vec![2.0], 2.0),
1030            ],
1031        )
1032        .unwrap()
1033    }
1034
1035    #[test]
1036    fn evaluates_selects_and_bins_dataset_expressions() {
1037        let dataset = dataset().chunked(1).unwrap();
1038        let execution = Execution::default();
1039        let x = event_scalar("x");
1040        assert_eq!(
1041            dataset.evaluate_real(&x, &execution).unwrap(),
1042            vec![-1.0, 0.0, 1.0, 2.0]
1043        );
1044
1045        let selected = dataset
1046            .select(
1047                &Predicate::ge(x.clone(), 0.0).and(Predicate::lt(x.clone(), 2.0)),
1048                &execution,
1049            )
1050            .unwrap();
1051        assert_eq!(
1052            selected.map_events(|event| event.scalar(0)).unwrap(),
1053            vec![0.0, 1.0]
1054        );
1055        assert_eq!(selected.sum_weights().unwrap(), 2.5);
1056
1057        let bins = dataset
1058            .bin_by(&x, BinSpec::uniform(2, 0.0, 2.0).unwrap(), &execution)
1059            .unwrap();
1060        assert_eq!(bins.len(), 2);
1061        assert_eq!(
1062            bins[0]
1063                .dataset()
1064                .map_events(|event| event.scalar(0))
1065                .unwrap(),
1066            vec![0.0]
1067        );
1068        assert_eq!(
1069            bins[1]
1070                .dataset()
1071                .map_events(|event| event.scalar(0))
1072                .unwrap(),
1073            vec![1.0, 2.0]
1074        );
1075    }
1076
1077    #[test]
1078    fn dataset_histogram_uses_event_weights_and_excludes_the_final_upper_edge() {
1079        let histogram = dataset()
1080            .histogram(
1081                &event_scalar("x"),
1082                BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1083                true,
1084                None,
1085                &Execution::default(),
1086            )
1087            .unwrap();
1088
1089        assert_eq!(histogram.counts(), [1.0, 1.5]);
1090        assert_eq!(histogram.sum_squared_weights(), [1.0, 2.25]);
1091        assert_eq!(histogram.underflow(), 0.5);
1092        assert_eq!(histogram.underflow_sum_squared_weights(), Some(0.25));
1093        assert_eq!(histogram.overflow(), 2.0);
1094        assert_eq!(histogram.overflow_sum_squared_weights(), Some(4.0));
1095    }
1096
1097    #[test]
1098    fn dataset_joint_histogram_uses_row_major_bins_and_aggregates_invalid_events() {
1099        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x", "y"], true).unwrap());
1100        let dataset = Dataset::from_events(
1101            schema,
1102            [
1103                OwnedEvent::weighted(vec![], vec![0.0, 10.0], 1.0),
1104                OwnedEvent::weighted(vec![], vec![1.0, 10.0], -2.0),
1105                OwnedEvent::weighted(vec![], vec![0.0, 20.0], 3.0),
1106                OwnedEvent::weighted(vec![], vec![f64::NAN, 10.0], 4.0),
1107            ],
1108        )
1109        .unwrap();
1110
1111        let histogram = dataset
1112            .joint_histogram(
1113                &[event_scalar("x"), event_scalar("y")],
1114                vec![
1115                    BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1116                    BinSpec::edges([0.0, 15.0, 25.0]).unwrap(),
1117                ],
1118                true,
1119                None,
1120                &Execution::default(),
1121            )
1122            .unwrap();
1123
1124        assert_eq!(histogram.shape(), [2, 2]);
1125        assert_eq!(histogram.values(), [1.0, 3.0, -2.0, 0.0]);
1126        assert_eq!(histogram.sum_squared_weights(), [1.0, 9.0, 4.0, 0.0]);
1127        assert_eq!(histogram.diagnostics().nonfinite_count(), 1);
1128        assert_eq!(histogram.diagnostics().out_of_range_count(), 0);
1129    }
1130
1131    #[test]
1132    fn dataset_histogram_multiplies_the_selected_base_and_custom_weights() {
1133        let x = event_scalar("x");
1134        let bins = BinSpec::edges([-2.0, 0.0, 2.0]).unwrap();
1135        let execution = Execution::default();
1136
1137        let weighted = dataset()
1138            .histogram(&x, bins.clone(), true, Some(&(x.clone() + 2.0)), &execution)
1139            .unwrap();
1140        assert_eq!(weighted.counts(), [0.5, 6.5]);
1141        assert_eq!(weighted.sum_squared_weights(), [0.25, 24.25]);
1142        assert_eq!(weighted.overflow(), 8.0);
1143
1144        let custom_only = dataset()
1145            .histogram(&x, bins, false, Some(&(x.clone() + 2.0)), &execution)
1146            .unwrap();
1147        assert_eq!(custom_only.counts(), [1.0, 5.0]);
1148        assert_eq!(custom_only.sum_squared_weights(), [1.0, 13.0]);
1149        assert_eq!(custom_only.overflow(), 4.0);
1150    }
1151
1152    #[test]
1153    fn dataset_histogram_preserves_view_source_and_memory_semantics() {
1154        let source = dataset();
1155        let batch = source.batches().unwrap().next().unwrap().unwrap();
1156        let reads = Arc::new(AtomicUsize::new(0));
1157        let counted = Dataset::new(CountingSource {
1158            inner: MemorySource::new(batch),
1159            reads: Arc::clone(&reads),
1160        })
1161        .streaming()
1162        .chunked(1)
1163        .unwrap();
1164        let selected = counted
1165            .select(
1166                &Predicate::ge(event_scalar("x"), 0.0),
1167                &Execution::default(),
1168            )
1169            .unwrap();
1170
1171        let histogram = selected
1172            .histogram(
1173                &event_scalar("x"),
1174                BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1175                true,
1176                None,
1177                &Execution::default(),
1178            )
1179            .unwrap();
1180
1181        assert_eq!(histogram.counts(), [1.0, 1.5]);
1182        assert_eq!(histogram.overflow(), 2.0);
1183        assert_eq!(reads.load(Ordering::Relaxed), 1);
1184    }
1185
1186    #[test]
1187    fn dataset_histogram_matches_across_memory_policies_and_chunking() {
1188        let source = dataset();
1189        let bins = BinSpec::edges([0.0, 1.0, 2.0]).unwrap();
1190        let execution = Execution::default();
1191        let expected = source
1192            .clone()
1193            .resident()
1194            .histogram(&event_scalar("x"), bins.clone(), true, None, &execution)
1195            .unwrap();
1196
1197        for candidate in [
1198            source.clone().streaming(),
1199            source.clone().resident().chunked(1).unwrap(),
1200            source.streaming().chunked(2).unwrap(),
1201        ] {
1202            let actual = candidate
1203                .histogram(&event_scalar("x"), bins.clone(), true, None, &execution)
1204                .unwrap();
1205            assert_eq!(actual, expected);
1206        }
1207    }
1208
1209    #[cfg(feature = "jit")]
1210    #[test]
1211    fn dataset_histogram_matches_cpu_interpreter_and_jit_backends() {
1212        let source = dataset();
1213        let bins = BinSpec::edges([0.0, 1.0, 2.0]).unwrap();
1214        let execution = |jit| {
1215            Execution::local(ExecutionOptions {
1216                device: Device::Cpu(CpuOptions {
1217                    threads: ThreadPolicy::Serial,
1218                    jit,
1219                }),
1220                precision: Precision::F64,
1221                ..ExecutionOptions::default()
1222            })
1223            .unwrap()
1224        };
1225        let interpreted = source
1226            .histogram(
1227                &event_scalar("x"),
1228                bins.clone(),
1229                true,
1230                None,
1231                &execution(JitPolicy::Disabled),
1232            )
1233            .unwrap();
1234        let compiled = source
1235            .histogram(
1236                &event_scalar("x"),
1237                bins,
1238                true,
1239                None,
1240                &execution(JitPolicy::Enabled),
1241            )
1242            .unwrap();
1243
1244        assert_eq!(compiled, interpreted);
1245    }
1246
1247    #[test]
1248    fn dataset_histogram_reports_invalid_values_and_source_failures() {
1249        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1250        let nonfinite = Dataset::from_events(
1251            Arc::clone(&schema),
1252            [OwnedEvent::weighted(vec![], vec![f64::NAN], 1.0)],
1253        )
1254        .unwrap();
1255        let error = nonfinite
1256            .histogram(
1257                &event_scalar("x"),
1258                BinSpec::edges([0.0, 1.0]).unwrap(),
1259                true,
1260                None,
1261                &Execution::default(),
1262            )
1263            .unwrap_err();
1264        assert!(
1265            error.to_string().contains("expected finite, got NaN"),
1266            "unexpected error: {error}",
1267        );
1268
1269        let source_error = Dataset::new(FailingSource { schema })
1270            .histogram(
1271                &event_scalar("x"),
1272                BinSpec::edges([0.0, 1.0]).unwrap(),
1273                true,
1274                None,
1275                &Execution::default(),
1276            )
1277            .unwrap_err();
1278        assert!(source_error.to_string().contains("query source failed"));
1279    }
1280
1281    #[test]
1282    fn empty_dataset_histogram_is_a_valid_empirical_histogram() {
1283        let histogram = dataset()
1284            .empty_derived()
1285            .unwrap()
1286            .histogram(
1287                &event_scalar("x"),
1288                BinSpec::edges([0.0, 1.0, 3.0]).unwrap(),
1289                true,
1290                None,
1291                &Execution::default(),
1292            )
1293            .unwrap();
1294
1295        assert_eq!(histogram.counts(), [0.0, 0.0]);
1296        assert_eq!(histogram.sum_squared_weights(), [0.0, 0.0]);
1297        assert_eq!(histogram.underflow_sum_squared_weights(), Some(0.0));
1298        assert_eq!(histogram.overflow_sum_squared_weights(), Some(0.0));
1299    }
1300
1301    #[test]
1302    fn cancelling_dataset_weights_keep_their_squared_weight_uncertainty() {
1303        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1304        let cancelling = Dataset::from_events(
1305            schema,
1306            [
1307                OwnedEvent::weighted(vec![], vec![0.5], 1.0),
1308                OwnedEvent::weighted(vec![], vec![0.5], -1.0),
1309            ],
1310        )
1311        .unwrap();
1312        let histogram = cancelling
1313            .histogram(
1314                &event_scalar("x"),
1315                BinSpec::edges([0.0, 1.0]).unwrap(),
1316                true,
1317                None,
1318                &Execution::default(),
1319            )
1320            .unwrap();
1321
1322        assert_eq!(histogram.counts(), [0.0]);
1323        assert_eq!(histogram.sum_squared_weights(), [2.0]);
1324        assert_eq!(histogram.errors(), [2.0_f64.sqrt()]);
1325    }
1326
1327    #[test]
1328    fn empty_batches_are_valid_query_inputs() {
1329        let execution = Execution::default();
1330        let x = event_scalar("x");
1331
1332        let empty_batch_schema =
1333            Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1334        let empty_batch = Dataset::from_batch(
1335            EventBatch::from_events(empty_batch_schema, std::iter::empty::<OwnedEvent>()).unwrap(),
1336        );
1337        for empty in [empty_batch, dataset().empty_derived().unwrap()] {
1338            assert!(empty.evaluate_real(&x, &execution).unwrap().is_empty());
1339            assert!(
1340                empty
1341                    .select(&Predicate::ge(x.clone(), 0.0), &execution)
1342                    .unwrap()
1343                    .map_events(|event| event.scalar(0))
1344                    .unwrap()
1345                    .is_empty()
1346            );
1347        }
1348    }
1349
1350    #[test]
1351    fn event_column_nan_comparisons_are_false() {
1352        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1353        let dataset = Dataset::from_events(
1354            schema,
1355            [
1356                OwnedEvent::weighted(vec![], vec![f64::NAN], 1.0),
1357                OwnedEvent::weighted(vec![], vec![0.0], 1.0),
1358                OwnedEvent::weighted(vec![], vec![1.0], 1.0),
1359            ],
1360        )
1361        .unwrap();
1362        let x = event_scalar("x");
1363        let selected = dataset
1364            .select(&Predicate::ne(x, 0.0), &Execution::default())
1365            .unwrap();
1366
1367        assert_eq!(
1368            selected.map_events(|event| event.scalar(0)).unwrap(),
1369            vec![1.0]
1370        );
1371    }
1372
1373    #[test]
1374    fn query_propagates_source_batch_errors() {
1375        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1376        let dataset = Dataset::new(FailingSource { schema });
1377
1378        let error = dataset
1379            .evaluate_real(&event_scalar("x"), &Execution::default())
1380            .unwrap_err();
1381        assert!(
1382            matches!(error, RuntimeError::Data(message) if message.contains("query source failed"))
1383        );
1384    }
1385
1386    #[test]
1387    fn all_empty_bins_retain_valid_empty_derived_sources() {
1388        let source = dataset();
1389        let before = capability_tuple(source.capabilities());
1390        let bins = source
1391            .bin_by(
1392                &event_scalar("x"),
1393                BinSpec::edges([10.0, 20.0, 30.0]).unwrap(),
1394                &Execution::default(),
1395            )
1396            .unwrap();
1397
1398        assert_eq!(capability_tuple(source.capabilities()), before);
1399        assert_eq!(bins.len(), 2);
1400        for bin in bins {
1401            assert_eq!(bin.dataset().num_events().unwrap(), Some(0));
1402            assert!(
1403                bin.dataset()
1404                    .evaluate_real(&event_scalar("x"), &Execution::default())
1405                    .unwrap()
1406                    .is_empty()
1407            );
1408        }
1409    }
1410
1411    #[test]
1412    fn traversing_all_bins_reads_the_source_once() {
1413        let reads = Arc::new(AtomicUsize::new(0));
1414        let source = CountingSource {
1415            inner: match dataset().batches().unwrap().next().unwrap() {
1416                Ok(batch) => MemorySource::new(batch),
1417                Err(error) => panic!("unexpected source error: {error}"),
1418            },
1419            reads: Arc::clone(&reads),
1420        };
1421        let dataset = Dataset::new(source).chunked(1).unwrap();
1422        let bins = dataset
1423            .bin_by(
1424                &event_scalar("x"),
1425                BinSpec::uniform(4, -1.0, 3.0).unwrap(),
1426                &Execution::default(),
1427            )
1428            .unwrap();
1429
1430        let values = bins
1431            .into_iter()
1432            .map(|bin| {
1433                bin.into_dataset()
1434                    .map_events(|event| event.scalar(0))
1435                    .unwrap()
1436            })
1437            .collect::<Vec<_>>();
1438        assert_eq!(values, [vec![-1.0], vec![0.0], vec![1.0], vec![2.0]]);
1439        assert_eq!(reads.load(Ordering::Relaxed), 1);
1440    }
1441
1442    #[test]
1443    fn between_predicates_have_explicit_endpoint_semantics() {
1444        let dataset = dataset();
1445        let execution = Execution::default();
1446        let x = event_scalar("x");
1447
1448        let closed = dataset
1449            .select(&Predicate::between(x.clone(), 0.0, 1.0), &execution)
1450            .unwrap();
1451        assert_eq!(
1452            closed.map_events(|event| event.scalar(0)).unwrap(),
1453            vec![0.0, 1.0]
1454        );
1455
1456        let open = dataset
1457            .select(
1458                &Predicate::between_with(x, -1.0, 1.0, IntervalClosure::Open),
1459                &execution,
1460            )
1461            .unwrap();
1462        assert_eq!(open.map_events(|event| event.scalar(0)).unwrap(), vec![0.0]);
1463    }
1464
1465    #[test]
1466    fn real_queries_reject_complex_and_free_parameter_expressions() {
1467        let dataset = dataset();
1468        let execution = Execution::default();
1469        assert!(
1470            dataset
1471                .evaluate_real(&complex(1.0, 1.0), &execution)
1472                .is_err()
1473        );
1474        let parameter = Expr::from(laddu_expr::parameters::Parameter::free("p"));
1475        assert!(dataset.evaluate_expr(&parameter, &execution).is_err());
1476    }
1477
1478    #[test]
1479    fn compiled_query_outputs_preserve_order_and_values() {
1480        let source = dataset();
1481        let batch = source.batches().unwrap().next().unwrap().unwrap();
1482        let x = event_scalar("x");
1483        let query = PreparedQuery::prepare(
1484            vec![x.clone() + 1.0, x.clone() * 2.0, x],
1485            &Execution::default(),
1486            false,
1487        )
1488        .unwrap();
1489        let values = query.evaluate_batch(&batch).unwrap();
1490        assert_eq!(
1491            values[0].iter().map(|v| v.re).collect::<Vec<_>>(),
1492            [0.0, 1.0, 2.0, 3.0]
1493        );
1494        assert_eq!(
1495            values[1].iter().map(|v| v.re).collect::<Vec<_>>(),
1496            [-2.0, 0.0, 2.0, 4.0]
1497        );
1498        assert_eq!(
1499            values[2].iter().map(|v| v.re).collect::<Vec<_>>(),
1500            [-1.0, 0.0, 1.0, 2.0]
1501        );
1502    }
1503
1504    #[test]
1505    fn repeated_predicate_leaves_are_evaluated_once() {
1506        let x = event_scalar("x");
1507        let selected = dataset()
1508            .select(
1509                &Predicate::ge(x.clone() + 1.0, 0.0).and(Predicate::lt(x + 1.0, 2.0)),
1510                &Execution::default(),
1511            )
1512            .unwrap();
1513        assert_eq!(
1514            selected.map_events(|event| event.scalar(0)).unwrap(),
1515            [-1.0, 0.0]
1516        );
1517    }
1518
1519    #[test]
1520    fn f32_queries_match_f64_query_results() {
1521        let x = event_scalar("x");
1522        let f64_values = dataset().evaluate_real(&x, &Execution::default()).unwrap();
1523        let f32_execution = Execution::local(ExecutionOptions {
1524            device: Device::Cpu(CpuOptions::default()),
1525            precision: Precision::F32,
1526            ..ExecutionOptions::default()
1527        })
1528        .unwrap();
1529        let f32_values = dataset().evaluate_real(&x, &f32_execution).unwrap();
1530        assert_eq!(f32_values, f64_values);
1531
1532        let bins = BinSpec::edges([-2.0, 0.0, 2.0]).unwrap();
1533        let f64_histogram = dataset()
1534            .histogram(&x, bins.clone(), true, None, &Execution::default())
1535            .unwrap();
1536        let f32_histogram = dataset()
1537            .histogram(&x, bins, true, None, &f32_execution)
1538            .unwrap();
1539        assert_eq!(f32_histogram, f64_histogram);
1540    }
1541
1542    #[test]
1543    fn bin_edges_validate_and_nan_predicates_are_false() {
1544        assert!(BinSpec::edges([0.0, 0.0]).is_err());
1545        assert!(!compare(f64::NAN, Comparison::Ne, 0.0));
1546    }
1547
1548    #[test]
1549    fn selection_is_lazy_and_one_pass_binning_preserves_streaming_policy() {
1550        let source = dataset();
1551        let batch = source.batches().unwrap().next().unwrap().unwrap();
1552        let reads = Arc::new(AtomicUsize::new(0));
1553        let dataset = Dataset::new(CountingSource {
1554            inner: MemorySource::new(batch),
1555            reads: Arc::clone(&reads),
1556        })
1557        .streaming();
1558        let execution = Execution::default();
1559        let x = event_scalar("x");
1560
1561        let selected = dataset
1562            .select(&Predicate::ge(x.clone(), 0.0), &execution)
1563            .unwrap();
1564        let bins = dataset
1565            .bin_by(&x, BinSpec::uniform(2, 0.0, 2.0).unwrap(), &execution)
1566            .unwrap();
1567        assert_eq!(reads.load(Ordering::Relaxed), 1);
1568        assert_eq!(
1569            selected.cache_storage(),
1570            laddu_data::data::CacheStorage::Streaming
1571        );
1572
1573        assert_eq!(
1574            selected.map_events(|event| event.scalar(0)).unwrap(),
1575            vec![0.0, 1.0, 2.0]
1576        );
1577        assert_eq!(reads.load(Ordering::Relaxed), 2);
1578        assert_eq!(
1579            bins[0]
1580                .dataset()
1581                .map_events(|event| event.scalar(0))
1582                .unwrap(),
1583            vec![0.0]
1584        );
1585        assert_eq!(reads.load(Ordering::Relaxed), 2);
1586    }
1587
1588    #[test]
1589    fn unknown_cardinality_fastest_discovers_and_retains_small_selection() {
1590        let source = dataset();
1591        let batch = source.batches().unwrap().next().unwrap().unwrap();
1592        let reads = Arc::new(AtomicUsize::new(0));
1593        let dataset = Dataset::new(CountingSource {
1594            inner: MemorySource::new(batch),
1595            reads: Arc::clone(&reads),
1596        });
1597        let execution = Execution::default();
1598        let x = event_scalar("x");
1599        let selected = dataset
1600            .select(&Predicate::ge(x.clone(), 0.0), &execution)
1601            .unwrap();
1602        let compiled = CompiledModel::from_expr(&x).unwrap();
1603        let params = compiled.params().default_values();
1604        let model = PreparedModel::prepare(&compiled, &execution).unwrap();
1605        let prepared = model.prepare_dataset(&execution, &selected).unwrap();
1606
1607        #[cfg(not(feature = "wgpu"))]
1608        let crate::PreparedDataset::Cpu(prepared_cpu) = &prepared;
1609        #[cfg(feature = "wgpu")]
1610        let crate::PreparedDataset::Cpu(prepared_cpu) = &prepared else {
1611            panic!("default execution prepares CPU datasets");
1612        };
1613        assert_eq!(
1614            prepared_cpu.stats().storage(),
1615            laddu_data::data::CacheStorage::Resident
1616        );
1617        assert_eq!(prepared_cpu.stats().local_events(), 3);
1618        assert_eq!(reads.load(Ordering::Relaxed), 2);
1619
1620        for _ in 0..2 {
1621            assert_eq!(
1622                model
1623                    .reduce(
1624                        &execution,
1625                        &params,
1626                        &prepared,
1627                        laddu_compile::ReductionPlan::weighted_real(),
1628                    )
1629                    .unwrap(),
1630                5.5
1631            );
1632        }
1633        assert_eq!(reads.load(Ordering::Relaxed), 2);
1634    }
1635}