Skip to main content

vortex_array/stats/
expr.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright the Vortex contributors
3
4//! Expression constructors for statistics backed by aggregate functions.
5
6use vortex_error::VortexExpect;
7
8use crate::aggregate_fn::AggregateFnRef;
9use crate::aggregate_fn::AggregateFnVTableExt;
10use crate::aggregate_fn::EmptyOptions;
11use crate::aggregate_fn::NumericalAggregateOpts;
12use crate::aggregate_fn::fns::all_nan::AllNan;
13use crate::aggregate_fn::fns::all_non_nan::AllNonNan;
14use crate::aggregate_fn::fns::all_non_null::AllNonNull;
15use crate::aggregate_fn::fns::all_null::AllNull;
16use crate::aggregate_fn::fns::min_max::MinMax;
17use crate::aggregate_fn::fns::nan_count::NanCount;
18use crate::aggregate_fn::fns::null_count::NullCount;
19use crate::aggregate_fn::fns::sum::Sum;
20use crate::expr::BoundExpression;
21use crate::expr::Expression;
22use crate::scalar_fn::ScalarFnVTableExt;
23pub use crate::scalar_fn::fns::stat::StatFn;
24pub use crate::scalar_fn::fns::stat::StatOptions;
25
26/// Creates an expression that reads a stored aggregate statistic for `expr`.
27///
28/// If the statistic is not available in the current stats scope, evaluating the expression returns
29/// a nullable all-null array with the aggregate return type.
30pub fn stat(expr: Expression, aggregate_fn: AggregateFnRef) -> Expression {
31    StatFn.new_expr(StatOptions::new(aggregate_fn), [expr])
32}
33
34fn bound_stat(expr: BoundExpression, aggregate_fn: AggregateFnRef) -> BoundExpression {
35    StatFn
36        .try_new_bound_expr(StatOptions::new(aggregate_fn), [expr])
37        .vortex_expect("stat expressions must use an aggregate supported by the child dtype")
38}
39
40/// Creates `stat(expr, min_max)`, returning a nullable `{ min, max }` struct statistic.
41pub fn min_max(expr: Expression) -> Expression {
42    // Statistics follow NaN-skipping semantics; request it explicitly rather than via the default.
43    stat(expr, MinMax.bind(NumericalAggregateOpts::skip_nans()))
44}
45
46fn bound_min_max(expr: BoundExpression) -> BoundExpression {
47    bound_stat(expr, MinMax.bind(NumericalAggregateOpts::skip_nans()))
48}
49
50/// Creates `stat(expr, sum)`, returning a nullable sum statistic.
51pub fn sum(expr: Expression) -> Expression {
52    // Statistics follow NaN-skipping semantics; request it explicitly rather than via the default.
53    stat(expr, Sum.bind(NumericalAggregateOpts::skip_nans()))
54}
55
56fn bound_sum(expr: BoundExpression) -> BoundExpression {
57    bound_stat(expr, Sum.bind(NumericalAggregateOpts::skip_nans()))
58}
59
60/// Creates `stat(expr, null_count)`, returning a nullable null-count statistic.
61pub fn null_count(expr: Expression) -> Expression {
62    stat(expr, NullCount.bind(EmptyOptions))
63}
64
65fn bound_null_count(expr: BoundExpression) -> BoundExpression {
66    bound_stat(expr, NullCount.bind(EmptyOptions))
67}
68
69/// Creates `stat(expr, all_null)`, returning a nullable all-null statistic.
70pub fn all_null(expr: Expression) -> Expression {
71    stat(expr, AllNull.bind(EmptyOptions))
72}
73
74fn bound_all_null(expr: BoundExpression) -> BoundExpression {
75    bound_stat(expr, AllNull.bind(EmptyOptions))
76}
77
78/// Creates `stat(expr, all_nan)`, returning a nullable all-NaN statistic.
79pub fn all_nan(expr: Expression) -> Expression {
80    stat(expr, AllNan.bind(EmptyOptions))
81}
82
83fn bound_all_nan(expr: BoundExpression) -> BoundExpression {
84    bound_stat(expr, AllNan.bind(EmptyOptions))
85}
86
87/// Creates `stat(expr, all_non_null)`, returning a nullable all-non-null statistic.
88pub fn all_non_null(expr: Expression) -> Expression {
89    stat(expr, AllNonNull.bind(EmptyOptions))
90}
91
92fn bound_all_non_null(expr: BoundExpression) -> BoundExpression {
93    bound_stat(expr, AllNonNull.bind(EmptyOptions))
94}
95
96/// Creates `stat(expr, all_non_nan)`, returning a nullable all-non-NaN statistic.
97pub fn all_non_nan(expr: Expression) -> Expression {
98    stat(expr, AllNonNan.bind(EmptyOptions))
99}
100
101fn bound_all_non_nan(expr: BoundExpression) -> BoundExpression {
102    bound_stat(expr, AllNonNan.bind(EmptyOptions))
103}
104
105/// Creates `stat(expr, nan_count)`, returning a nullable NaN-count statistic.
106pub fn nan_count(expr: Expression) -> Expression {
107    stat(expr, NanCount.bind(EmptyOptions))
108}
109
110fn bound_nan_count(expr: BoundExpression) -> BoundExpression {
111    bound_stat(expr, NanCount.bind(EmptyOptions))
112}
113
114/// Constructors for statistic expressions whose input has already been bound.
115///
116/// These mirror the constructors in [`crate::stats`] and panic when the aggregate does not support
117/// the input dtype.
118pub mod bound {
119    use crate::aggregate_fn::AggregateFnRef;
120    use crate::expr::BoundExpression;
121
122    /// Creates a bound expression that reads a stored aggregate statistic.
123    pub fn stat(expr: BoundExpression, aggregate_fn: AggregateFnRef) -> BoundExpression {
124        super::bound_stat(expr, aggregate_fn)
125    }
126
127    /// Creates a bound nullable `{ min, max }` statistic expression.
128    pub fn min_max(expr: BoundExpression) -> BoundExpression {
129        super::bound_min_max(expr)
130    }
131
132    /// Creates a bound nullable sum statistic expression.
133    pub fn sum(expr: BoundExpression) -> BoundExpression {
134        super::bound_sum(expr)
135    }
136
137    /// Creates a bound nullable null-count statistic expression.
138    pub fn null_count(expr: BoundExpression) -> BoundExpression {
139        super::bound_null_count(expr)
140    }
141
142    /// Creates a bound nullable all-null statistic expression.
143    pub fn all_null(expr: BoundExpression) -> BoundExpression {
144        super::bound_all_null(expr)
145    }
146
147    /// Creates a bound nullable all-NaN statistic expression.
148    pub fn all_nan(expr: BoundExpression) -> BoundExpression {
149        super::bound_all_nan(expr)
150    }
151
152    /// Creates a bound nullable all-non-null statistic expression.
153    pub fn all_non_null(expr: BoundExpression) -> BoundExpression {
154        super::bound_all_non_null(expr)
155    }
156
157    /// Creates a bound nullable all-non-NaN statistic expression.
158    pub fn all_non_nan(expr: BoundExpression) -> BoundExpression {
159        super::bound_all_non_nan(expr)
160    }
161
162    /// Creates a bound nullable NaN-count statistic expression.
163    pub fn nan_count(expr: BoundExpression) -> BoundExpression {
164        super::bound_nan_count(expr)
165    }
166}
167
168#[cfg(test)]
169mod tests {
170    use std::sync::LazyLock;
171
172    use vortex_buffer::buffer;
173    use vortex_error::VortexExpect;
174    use vortex_error::VortexResult;
175    use vortex_session::VortexSession;
176
177    use super::all_nan;
178    use super::all_non_nan;
179    use super::all_non_null;
180    use super::all_null;
181    use super::bound as bound_stats;
182    use super::null_count;
183    use super::stat;
184    use super::sum;
185    use crate::Canonical;
186    use crate::IntoArray;
187    use crate::VortexSessionExecute;
188    use crate::array_session;
189    use crate::arrays::Chunked;
190    use crate::arrays::ChunkedArray;
191    use crate::arrays::ConstantArray;
192    use crate::arrays::PrimitiveArray;
193    use crate::arrays::chunked::ChunkedArrayExt;
194    use crate::assert_arrays_eq;
195    use crate::dtype::DType;
196    use crate::dtype::Nullability;
197    use crate::dtype::PType;
198    use crate::expr::bound as bound_expr;
199    use crate::expr::root;
200    use crate::expr::stats::Precision;
201    use crate::expr::stats::Stat;
202    use crate::scalar::Scalar;
203    use crate::scalar::ScalarValue;
204    use crate::validity::Validity;
205
206    static SESSION: LazyLock<VortexSession> = LazyLock::new(array_session);
207
208    #[test]
209    fn bound_stats_constructor_preserves_child_and_dtype() -> VortexResult<()> {
210        let input_dtype = DType::Primitive(PType::I32, Nullability::NonNullable);
211        let root = bound_expr::root(input_dtype.clone());
212        let bound = bound_stats::sum(root.clone());
213
214        assert_eq!(bound.children(), &[root]);
215        assert_eq!(
216            bound.dtype(),
217            &DType::Primitive(PType::I64, Nullability::Nullable)
218        );
219        assert_eq!(bound, sum(crate::expr::root()).bind(&input_dtype)?);
220        Ok(())
221    }
222
223    #[test]
224    fn stat_expr_reads_cached_sum() -> VortexResult<()> {
225        let array = buffer![1i32, 2, 3].into_array();
226        let sum_scalar = Scalar::primitive(6i64, Nullability::Nullable);
227        array.statistics().set(
228            Stat::Sum,
229            Precision::exact(sum_scalar.into_value().vortex_expect("non-null sum")),
230        );
231
232        let result = array
233            .apply(&sum(root()))?
234            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
235            .into_array();
236
237        let expected =
238            ConstantArray::new(Scalar::primitive(6i64, Nullability::Nullable), 3).into_array();
239        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
240
241        Ok(())
242    }
243
244    #[test]
245    fn stat_expr_returns_null_when_sum_is_missing() -> VortexResult<()> {
246        let array = buffer![1i32, 2, 3].into_array();
247
248        let result = array
249            .apply(&sum(root()))?
250            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
251            .into_array();
252
253        let expected = ConstantArray::new(
254            Scalar::null(DType::Primitive(PType::I64, Nullability::Nullable)),
255            3,
256        )
257        .into_array();
258        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
259
260        Ok(())
261    }
262
263    #[test]
264    fn stat_expr_reads_cached_sum_per_chunk() -> VortexResult<()> {
265        let chunk0 = buffer![1i32, 2].into_array();
266        let sum_scalar = Scalar::primitive(3i64, Nullability::Nullable);
267        chunk0.statistics().set(
268            Stat::Sum,
269            Precision::exact(sum_scalar.into_value().vortex_expect("non-null sum")),
270        );
271        let chunk1 = buffer![4i32, 5, 6].into_array();
272        let chunked = ChunkedArray::try_new(
273            vec![chunk0, chunk1],
274            DType::Primitive(PType::I32, Nullability::NonNullable),
275        )?
276        .into_array();
277
278        let result = chunked.apply(&sum(root()))?;
279
280        let chunked_result = result
281            .as_opt::<Chunked>()
282            .vortex_expect("stat expression should preserve chunked alignment");
283        assert_eq!(chunked_result.nchunks(), 2);
284
285        let result = result
286            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
287            .into_array();
288        let expected = PrimitiveArray::new(
289            buffer![3i64, 3, 0, 0, 0],
290            Validity::from_iter([true, true, false, false, false]),
291        )
292        .into_array();
293        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
294
295        Ok(())
296    }
297
298    #[test]
299    fn stat_expr_reads_cached_null_count() -> VortexResult<()> {
300        let array =
301            PrimitiveArray::from_option_iter([Some(1i32), None, Some(3), None]).into_array();
302        let null_count_scalar = Scalar::primitive(2u64, Nullability::NonNullable);
303        array.statistics().set(
304            Stat::NullCount,
305            Precision::exact(
306                null_count_scalar
307                    .into_value()
308                    .vortex_expect("non-null null_count"),
309            ),
310        );
311
312        let result = array
313            .apply(&null_count(root()))?
314            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
315            .into_array();
316
317        let expected =
318            ConstantArray::new(Scalar::primitive(2u64, Nullability::Nullable), 4).into_array();
319        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
320
321        Ok(())
322    }
323
324    #[test]
325    fn stat_expr_reads_cached_all_null_from_null_count() -> VortexResult<()> {
326        let array = PrimitiveArray::from_option_iter::<i32, _>([None, None, None]).into_array();
327        array
328            .statistics()
329            .set(Stat::NullCount, Precision::exact(ScalarValue::from(3u64)));
330
331        let result = array
332            .apply(&all_null(root()))?
333            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
334            .into_array();
335
336        let expected =
337            ConstantArray::new(Scalar::bool(true, Nullability::Nullable), 3).into_array();
338        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
339
340        Ok(())
341    }
342
343    #[test]
344    fn stat_expr_reads_cached_all_null_false_from_inexact_low_null_count() -> VortexResult<()> {
345        let array = PrimitiveArray::from_option_iter::<i32, _>([None, Some(2), None]).into_array();
346        array
347            .statistics()
348            .set(Stat::NullCount, Precision::inexact(ScalarValue::from(2u64)));
349
350        let result = array
351            .apply(&all_null(root()))?
352            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
353            .into_array();
354
355        let expected =
356            ConstantArray::new(Scalar::bool(false, Nullability::Nullable), 3).into_array();
357        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
358
359        Ok(())
360    }
361
362    #[test]
363    fn stat_expr_returns_null_for_inexact_full_null_count_as_all_null() -> VortexResult<()> {
364        let array = PrimitiveArray::from_option_iter::<i32, _>([None, Some(2), None]).into_array();
365        array
366            .statistics()
367            .set(Stat::NullCount, Precision::inexact(ScalarValue::from(3u64)));
368
369        let result = array
370            .apply(&all_null(root()))?
371            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
372            .into_array();
373
374        let expected =
375            ConstantArray::new(Scalar::null(DType::Bool(Nullability::Nullable)), 3).into_array();
376        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
377
378        Ok(())
379    }
380
381    #[test]
382    fn stat_expr_reads_cached_all_non_null_from_null_count() -> VortexResult<()> {
383        let array = buffer![1i32, 2, 3].into_array();
384        array
385            .statistics()
386            .set(Stat::NullCount, Precision::exact(ScalarValue::from(0u64)));
387
388        let result = array
389            .apply(&all_non_null(root()))?
390            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
391            .into_array();
392
393        let expected =
394            ConstantArray::new(Scalar::bool(true, Nullability::Nullable), 3).into_array();
395        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
396
397        Ok(())
398    }
399
400    #[test]
401    fn stat_expr_reads_cached_all_non_null_true_from_inexact_zero_null_count() -> VortexResult<()> {
402        let array = buffer![1i32, 2, 3].into_array();
403        array
404            .statistics()
405            .set(Stat::NullCount, Precision::inexact(ScalarValue::from(0u64)));
406
407        let result = array
408            .apply(&all_non_null(root()))?
409            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
410            .into_array();
411
412        let expected =
413            ConstantArray::new(Scalar::bool(true, Nullability::Nullable), 3).into_array();
414        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
415
416        Ok(())
417    }
418
419    #[test]
420    fn stat_expr_returns_null_for_inexact_nonzero_null_count_as_all_non_null() -> VortexResult<()> {
421        let array =
422            PrimitiveArray::from_option_iter([Some(1i32), None, Some(3), None]).into_array();
423        array
424            .statistics()
425            .set(Stat::NullCount, Precision::inexact(ScalarValue::from(2u64)));
426
427        let result = array
428            .apply(&all_non_null(root()))?
429            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
430            .into_array();
431
432        let expected =
433            ConstantArray::new(Scalar::null(DType::Bool(Nullability::Nullable)), 4).into_array();
434        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
435
436        Ok(())
437    }
438
439    #[test]
440    fn stat_expr_rejects_all_nan_for_non_float() -> VortexResult<()> {
441        let array = PrimitiveArray::empty::<i32>(Nullability::NonNullable).into_array();
442        let mut ctx = SESSION.create_execution_ctx();
443
444        let result = array
445            .apply(&all_nan(root()))
446            .and_then(|array| array.execute::<Canonical>(&mut ctx));
447
448        assert!(result.is_err());
449        Ok(())
450    }
451
452    #[test]
453    fn stat_expr_reads_cached_all_nan_from_nan_count() -> VortexResult<()> {
454        let array =
455            PrimitiveArray::from_option_iter([Some(f32::NAN), Some(f32::NAN), Some(f32::NAN)])
456                .into_array();
457        array
458            .statistics()
459            .set(Stat::NaNCount, Precision::exact(ScalarValue::from(3u64)));
460
461        let result = array
462            .apply(&all_nan(root()))?
463            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
464            .into_array();
465
466        let expected =
467            ConstantArray::new(Scalar::bool(true, Nullability::Nullable), 3).into_array();
468        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
469
470        Ok(())
471    }
472
473    #[test]
474    fn stat_expr_reads_cached_all_nan_false_from_inexact_low_nan_count() -> VortexResult<()> {
475        let array =
476            PrimitiveArray::from_option_iter([Some(f32::NAN), Some(1.0f32), Some(f32::NAN)])
477                .into_array();
478        array
479            .statistics()
480            .set(Stat::NaNCount, Precision::inexact(ScalarValue::from(2u64)));
481
482        let result = array
483            .apply(&all_nan(root()))?
484            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
485            .into_array();
486
487        let expected =
488            ConstantArray::new(Scalar::bool(false, Nullability::Nullable), 3).into_array();
489        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
490
491        Ok(())
492    }
493
494    #[test]
495    fn stat_expr_returns_null_for_inexact_full_nan_count_as_all_nan() -> VortexResult<()> {
496        let array =
497            PrimitiveArray::from_option_iter([Some(f32::NAN), Some(1.0f32), Some(f32::NAN)])
498                .into_array();
499        array
500            .statistics()
501            .set(Stat::NaNCount, Precision::inexact(ScalarValue::from(3u64)));
502
503        let result = array
504            .apply(&all_nan(root()))?
505            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
506            .into_array();
507
508        let expected =
509            ConstantArray::new(Scalar::null(DType::Bool(Nullability::Nullable)), 3).into_array();
510        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
511
512        Ok(())
513    }
514
515    #[test]
516    fn stat_expr_reads_cached_all_non_nan_true_from_inexact_zero_nan_count() -> VortexResult<()> {
517        let array = buffer![1.0f32, 2.0, 3.0].into_array();
518        array
519            .statistics()
520            .set(Stat::NaNCount, Precision::inexact(ScalarValue::from(0u64)));
521
522        let result = array
523            .apply(&all_non_nan(root()))?
524            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
525            .into_array();
526
527        let expected =
528            ConstantArray::new(Scalar::bool(true, Nullability::Nullable), 3).into_array();
529        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
530
531        Ok(())
532    }
533
534    #[test]
535    fn stat_expr_returns_null_for_inexact_nonzero_nan_count_as_all_non_nan() -> VortexResult<()> {
536        let array = PrimitiveArray::from_option_iter([Some(1.0f32), Some(f32::NAN), Some(3.0)])
537            .into_array();
538        array
539            .statistics()
540            .set(Stat::NaNCount, Precision::inexact(ScalarValue::from(1u64)));
541
542        let result = array
543            .apply(&all_non_nan(root()))?
544            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
545            .into_array();
546
547        let expected =
548            ConstantArray::new(Scalar::null(DType::Bool(Nullability::Nullable)), 3).into_array();
549        assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
550
551        Ok(())
552    }
553
554    #[test]
555    fn stat_expr_reads_cached_min_and_max() -> VortexResult<()> {
556        let array = buffer![3i32, 1, 2].into_array();
557        array
558            .statistics()
559            .set(Stat::Min, Precision::exact(ScalarValue::from(1i32)));
560        array
561            .statistics()
562            .set(Stat::Max, Precision::exact(ScalarValue::from(3i32)));
563
564        let min_result = array
565            .clone()
566            .apply(&stat(
567                root(),
568                Stat::Min
569                    .aggregate_fn()
570                    .vortex_expect("min should have an aggregate function"),
571            ))?
572            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
573            .into_array();
574        let expected_min =
575            ConstantArray::new(Scalar::primitive(1i32, Nullability::Nullable), 3).into_array();
576        assert_arrays_eq!(
577            min_result,
578            expected_min,
579            &mut SESSION.create_execution_ctx()
580        );
581
582        let max_result = array
583            .apply(&stat(
584                root(),
585                Stat::Max
586                    .aggregate_fn()
587                    .vortex_expect("max should have an aggregate function"),
588            ))?
589            .execute::<Canonical>(&mut SESSION.create_execution_ctx())?
590            .into_array();
591        let expected_max =
592            ConstantArray::new(Scalar::primitive(3i32, Nullability::Nullable), 3).into_array();
593        assert_arrays_eq!(
594            max_result,
595            expected_max,
596            &mut SESSION.create_execution_ctx()
597        );
598
599        Ok(())
600    }
601}