Skip to main content

vortex_layout/layouts/
file_stats.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright the Vortex contributors
3
4use std::future;
5use std::marker::PhantomData;
6use std::sync::Arc;
7
8use futures::StreamExt;
9use itertools::Itertools;
10use parking_lot::Mutex;
11use vortex_array::ArrayRef;
12use vortex_array::ExecutionCtx;
13use vortex_array::IntoArray;
14use vortex_array::VortexSessionExecute;
15use vortex_array::aggregate_fn::fns::sum::sum;
16use vortex_array::arrays::ConstantArray;
17use vortex_array::arrays::StructArray;
18use vortex_array::arrays::struct_::StructArrayExt;
19use vortex_array::builders::ArrayBuilder;
20use vortex_array::builders::BoolBuilder;
21use vortex_array::builders::builder_with_capacity;
22use vortex_array::dtype::DType;
23use vortex_array::dtype::FieldName;
24use vortex_array::dtype::Nullability;
25use vortex_array::dtype::PType;
26use vortex_array::expr::stats::Precision;
27use vortex_array::expr::stats::Stat;
28use vortex_array::scalar::Scalar;
29use vortex_array::scalar::ScalarTruncation;
30use vortex_array::scalar::lower_bound;
31use vortex_array::scalar::upper_bound;
32use vortex_array::stats::StatsSet;
33use vortex_array::validity::Validity;
34use vortex_buffer::BufferString;
35use vortex_buffer::ByteBuffer;
36use vortex_error::VortexExpect;
37use vortex_error::VortexResult;
38use vortex_error::vortex_panic;
39use vortex_session::VortexSession;
40
41use crate::layouts::zoned::MAX_IS_TRUNCATED;
42use crate::layouts::zoned::MIN_IS_TRUNCATED;
43use crate::sequence::SendableSequentialStream;
44use crate::sequence::SequenceId;
45use crate::sequence::SequentialStreamAdapter;
46use crate::sequence::SequentialStreamExt;
47
48pub fn accumulate_stats(
49    stream: SendableSequentialStream,
50    stats: Arc<[Stat]>,
51    max_variable_length_statistics_size: usize,
52    session: &VortexSession,
53) -> (FileStatsAccumulator, SendableSequentialStream) {
54    let accumulator = FileStatsAccumulator::new(
55        stream.dtype(),
56        stats,
57        max_variable_length_statistics_size,
58        session,
59    );
60    let stream = SequentialStreamAdapter::new(
61        stream.dtype().clone(),
62        stream.scan(accumulator.clone(), |acc, item| {
63            future::ready(Some(acc.process(item)))
64        }),
65    )
66    .sendable();
67    (accumulator, stream)
68}
69
70/// Accumulates write-time statistics for a single file column.
71struct StatsAccumulator {
72    builders: Vec<Box<dyn StatsArrayBuilder>>,
73    length: usize,
74}
75
76impl StatsAccumulator {
77    fn new(dtype: &DType, stats: &[Stat], max_variable_length_statistics_size: usize) -> Self {
78        if !supports_file_stats(dtype) {
79            return Self {
80                builders: Vec::new(),
81                length: 0,
82            };
83        }
84
85        let builders = stats
86            .iter()
87            .filter_map(|&stat| {
88                stat.dtype(dtype).map(|stat_dtype| {
89                    stats_builder_with_capacity(
90                        stat,
91                        &stat_dtype.as_nullable(),
92                        1024,
93                        max_variable_length_statistics_size,
94                    )
95                })
96            })
97            .collect::<Vec<_>>();
98
99        Self {
100            builders,
101            length: 0,
102        }
103    }
104
105    fn push_chunk(&mut self, array: &ArrayRef, ctx: &mut ExecutionCtx) -> VortexResult<()> {
106        for builder in &mut self.builders {
107            if let Some(value) = array.statistics().compute_stat(builder.stat(), ctx)? {
108                builder.append_scalar(value.cast(&value.dtype().as_nullable())?)?;
109            } else {
110                builder.append_null();
111            }
112        }
113        self.length += 1;
114        Ok(())
115    }
116
117    fn as_array(&mut self, ctx: &mut ExecutionCtx) -> VortexResult<Option<StructArray>> {
118        let mut names = Vec::new();
119        let mut fields = Vec::new();
120
121        for builder in self
122            .builders
123            .iter_mut()
124            // We sort the stats so the DType is deterministic based on which stats are present.
125            .sorted_unstable_by_key(|builder| builder.stat())
126        {
127            let values = builder.finish();
128
129            // We drop any all-null stats columns.
130            if values.all_invalid(ctx)? {
131                continue;
132            }
133
134            names.extend(values.names);
135            fields.extend(values.arrays);
136        }
137
138        if names.is_empty() {
139            return Ok(None);
140        }
141
142        StructArray::try_new(names.into(), fields, self.length, Validity::NonNullable).map(Some)
143    }
144
145    /// Returns an aggregated stats set for the table.
146    fn as_stats_set(&mut self, stats: &[Stat], ctx: &mut ExecutionCtx) -> VortexResult<StatsSet> {
147        let mut stats_set = StatsSet::default();
148        let Some(stats_table) = self.as_array(ctx)? else {
149            return Ok(stats_set);
150        };
151
152        for &stat in stats {
153            let Some(values) = stats_table.unmasked_field_by_name_opt(stat.name()) else {
154                continue;
155            };
156
157            match stat {
158                Stat::Max if is_varlen_dtype(values.dtype()) && !values.all_valid(ctx)? => {
159                    // A null truncated varlen max can mean either an empty chunk or no finite
160                    // upper bound, so aggregating by skipping nulls would be unsound.
161                    continue;
162                }
163                Stat::Min | Stat::Max | Stat::Sum => {
164                    if let Some(s) = values.statistics().compute_stat(stat, ctx)?
165                        && let Some(v) = s.into_value()
166                    {
167                        let precision = if stat_was_truncated(&stats_table, stat, ctx)? {
168                            Precision::inexact(v)
169                        } else {
170                            Precision::exact(v)
171                        };
172                        stats_set.set(stat, precision)
173                    }
174                }
175                Stat::NullCount | Stat::NaNCount | Stat::UncompressedSizeInBytes => {
176                    if let Some(sum_value) = sum(values, ctx)?
177                        .cast(&DType::Primitive(PType::U64, Nullability::Nullable))?
178                        .into_value()
179                    {
180                        stats_set.set(stat, Precision::exact(sum_value));
181                    }
182                }
183                Stat::IsConstant | Stat::IsSorted | Stat::IsStrictSorted => {}
184            }
185        }
186        Ok(stats_set)
187    }
188}
189
190fn stat_was_truncated(
191    stats_table: &StructArray,
192    stat: Stat,
193    ctx: &mut ExecutionCtx,
194) -> VortexResult<bool> {
195    let field_name = match stat {
196        Stat::Min => MIN_IS_TRUNCATED,
197        Stat::Max => MAX_IS_TRUNCATED,
198        _ => return Ok(false),
199    };
200    let Some(is_truncated) = stats_table.unmasked_field_by_name_opt(field_name) else {
201        return Ok(false);
202    };
203
204    Ok(is_truncated
205        .statistics()
206        .compute_stat(Stat::Max, ctx)?
207        .is_some_and(|max| max.as_bool().value() == Some(true)))
208}
209
210fn supports_file_stats(dtype: &DType) -> bool {
211    !matches!(dtype, DType::Variant(_))
212}
213
214fn is_varlen_dtype(dtype: &DType) -> bool {
215    matches!(dtype, DType::Utf8(_) | DType::Binary(_))
216}
217
218fn stats_builder_with_capacity(
219    stat: Stat,
220    dtype: &DType,
221    capacity: usize,
222    max_length: usize,
223) -> Box<dyn StatsArrayBuilder> {
224    let values_builder = builder_with_capacity(dtype, capacity);
225    match stat {
226        Stat::Max => match dtype {
227            DType::Utf8(_) => Box::new(TruncatedMaxBinaryStatsBuilder::<BufferString>::new(
228                values_builder,
229                BoolBuilder::with_capacity(Nullability::NonNullable, capacity),
230                max_length,
231            )),
232            DType::Binary(_) => Box::new(TruncatedMaxBinaryStatsBuilder::<ByteBuffer>::new(
233                values_builder,
234                BoolBuilder::with_capacity(Nullability::NonNullable, capacity),
235                max_length,
236            )),
237            _ => Box::new(StatNameArrayBuilder::new(stat, values_builder)),
238        },
239        Stat::Min => match dtype {
240            DType::Utf8(_) => Box::new(TruncatedMinBinaryStatsBuilder::<BufferString>::new(
241                values_builder,
242                BoolBuilder::with_capacity(Nullability::NonNullable, capacity),
243                max_length,
244            )),
245            DType::Binary(_) => Box::new(TruncatedMinBinaryStatsBuilder::<ByteBuffer>::new(
246                values_builder,
247                BoolBuilder::with_capacity(Nullability::NonNullable, capacity),
248                max_length,
249            )),
250            _ => Box::new(StatNameArrayBuilder::new(stat, values_builder)),
251        },
252        _ => Box::new(StatNameArrayBuilder::new(stat, values_builder)),
253    }
254}
255
256/// Arrays with their associated names, reduced version of a `StructArray`.
257struct NamedArrays {
258    names: Vec<FieldName>,
259    arrays: Vec<ArrayRef>,
260}
261
262impl NamedArrays {
263    fn all_invalid(&self, ctx: &mut ExecutionCtx) -> VortexResult<bool> {
264        self.arrays[0].all_invalid(ctx)
265    }
266}
267
268trait StatsArrayBuilder: Send {
269    fn stat(&self) -> Stat;
270
271    fn append_scalar(&mut self, value: Scalar) -> VortexResult<()>;
272
273    fn append_null(&mut self);
274
275    fn finish(&mut self) -> NamedArrays;
276}
277
278struct StatNameArrayBuilder {
279    stat: Stat,
280    builder: Box<dyn ArrayBuilder>,
281}
282
283impl StatNameArrayBuilder {
284    fn new(stat: Stat, builder: Box<dyn ArrayBuilder>) -> Self {
285        Self { stat, builder }
286    }
287}
288
289impl StatsArrayBuilder for StatNameArrayBuilder {
290    fn stat(&self) -> Stat {
291        self.stat
292    }
293
294    fn append_scalar(&mut self, value: Scalar) -> VortexResult<()> {
295        self.builder.append_scalar(&value)
296    }
297
298    fn append_null(&mut self) {
299        self.builder.append_null()
300    }
301
302    fn finish(&mut self) -> NamedArrays {
303        let array = self.builder.finish();
304        let len = array.len();
305        match self.stat {
306            Stat::Max => NamedArrays {
307                names: vec![self.stat.name().into(), MAX_IS_TRUNCATED.into()],
308                arrays: vec![array, ConstantArray::new(false, len).into_array()],
309            },
310            Stat::Min => NamedArrays {
311                names: vec![self.stat.name().into(), MIN_IS_TRUNCATED.into()],
312                arrays: vec![array, ConstantArray::new(false, len).into_array()],
313            },
314            _ => NamedArrays {
315                names: vec![self.stat.name().into()],
316                arrays: vec![array],
317            },
318        }
319    }
320}
321
322struct TruncatedMaxBinaryStatsBuilder<T: ScalarTruncation> {
323    values: Box<dyn ArrayBuilder>,
324    is_truncated: BoolBuilder,
325    max_value_length: usize,
326    _marker: PhantomData<T>,
327}
328
329impl<T: ScalarTruncation> TruncatedMaxBinaryStatsBuilder<T> {
330    fn new(
331        values: Box<dyn ArrayBuilder>,
332        is_truncated: BoolBuilder,
333        max_value_length: usize,
334    ) -> Self {
335        Self {
336            values,
337            is_truncated,
338            max_value_length,
339            _marker: PhantomData,
340        }
341    }
342}
343
344struct TruncatedMinBinaryStatsBuilder<T: ScalarTruncation> {
345    values: Box<dyn ArrayBuilder>,
346    is_truncated: BoolBuilder,
347    max_value_length: usize,
348    _marker: PhantomData<T>,
349}
350
351impl<T: ScalarTruncation> TruncatedMinBinaryStatsBuilder<T> {
352    fn new(
353        values: Box<dyn ArrayBuilder>,
354        is_truncated: BoolBuilder,
355        max_value_length: usize,
356    ) -> Self {
357        Self {
358            values,
359            is_truncated,
360            max_value_length,
361            _marker: PhantomData,
362        }
363    }
364}
365
366impl<T: ScalarTruncation> StatsArrayBuilder for TruncatedMaxBinaryStatsBuilder<T> {
367    fn stat(&self) -> Stat {
368        Stat::Max
369    }
370
371    fn append_scalar(&mut self, value: Scalar) -> VortexResult<()> {
372        let nullability = value.dtype().nullability();
373        if let Some((upper_bound, truncated)) =
374            upper_bound(T::from_scalar(value)?, self.max_value_length, nullability)
375        {
376            self.values.append_scalar(&upper_bound)?;
377            self.is_truncated.append_value(truncated);
378        } else {
379            self.append_null()
380        }
381        Ok(())
382    }
383
384    fn append_null(&mut self) {
385        ArrayBuilder::append_null(self.values.as_mut());
386        self.is_truncated.append_value(false);
387    }
388
389    fn finish(&mut self) -> NamedArrays {
390        NamedArrays {
391            names: vec![Stat::Max.name().into(), MAX_IS_TRUNCATED.into()],
392            arrays: vec![
393                ArrayBuilder::finish(self.values.as_mut()),
394                ArrayBuilder::finish(&mut self.is_truncated),
395            ],
396        }
397    }
398}
399
400impl<T: ScalarTruncation> StatsArrayBuilder for TruncatedMinBinaryStatsBuilder<T> {
401    fn stat(&self) -> Stat {
402        Stat::Min
403    }
404
405    fn append_scalar(&mut self, value: Scalar) -> VortexResult<()> {
406        let nullability = value.dtype().nullability();
407        if let Some((lower_bound, truncated)) =
408            lower_bound(T::from_scalar(value)?, self.max_value_length, nullability)
409        {
410            self.values.append_scalar(&lower_bound)?;
411            self.is_truncated.append_value(truncated);
412        } else {
413            self.append_null()
414        }
415        Ok(())
416    }
417
418    fn append_null(&mut self) {
419        ArrayBuilder::append_null(self.values.as_mut());
420        self.is_truncated.append_value(false);
421    }
422
423    fn finish(&mut self) -> NamedArrays {
424        NamedArrays {
425            names: vec![Stat::Min.name().into(), MIN_IS_TRUNCATED.into()],
426            arrays: vec![
427                ArrayBuilder::finish(self.values.as_mut()),
428                ArrayBuilder::finish(&mut self.is_truncated),
429            ],
430        }
431    }
432}
433
434/// An array stream processor that computes aggregate statistics for all fields.
435///
436/// Note: for now this only collects top-level struct fields.
437#[derive(Clone)]
438pub struct FileStatsAccumulator {
439    stats: Arc<[Stat]>,
440    accumulators: Arc<Mutex<Vec<StatsAccumulator>>>,
441    ctx: Arc<Mutex<ExecutionCtx>>,
442}
443
444impl FileStatsAccumulator {
445    fn new(
446        dtype: &DType,
447        stats: Arc<[Stat]>,
448        max_variable_length_statistics_size: usize,
449        session: &VortexSession,
450    ) -> Self {
451        let accumulators = Arc::new(Mutex::new(match dtype.as_struct_fields_opt() {
452            Some(struct_dtype) => {
453                if dtype.nullability() == Nullability::Nullable {
454                    // top level dtype could be nullable, but we don't support it yet
455                    vortex_panic!(
456                        "FileStatsAccumulator temporarily does not support nullable top-level structs, got: {}. Use Validity::NonNullable",
457                        dtype
458                    );
459                }
460
461                struct_dtype
462                    .fields()
463                    .map(|field_dtype| {
464                        StatsAccumulator::new(
465                            &field_dtype,
466                            &stats,
467                            max_variable_length_statistics_size,
468                        )
469                    })
470                    .collect()
471            }
472            None => [StatsAccumulator::new(
473                dtype,
474                &stats,
475                max_variable_length_statistics_size,
476            )]
477            .into(),
478        }));
479
480        Self {
481            stats,
482            accumulators,
483            ctx: Arc::new(Mutex::new(session.create_execution_ctx())),
484        }
485    }
486
487    fn process(
488        &self,
489        chunk: VortexResult<(SequenceId, ArrayRef)>,
490    ) -> VortexResult<(SequenceId, ArrayRef)> {
491        let (sequence_id, chunk) = chunk?;
492        let mut ctx = self.ctx.lock();
493        if chunk.dtype().is_struct() {
494            let struct_chunk = chunk.clone().execute::<StructArray>(&mut ctx)?;
495            for (acc, field) in self
496                .accumulators
497                .lock()
498                .iter_mut()
499                .zip_eq(struct_chunk.iter_unmasked_fields())
500            {
501                acc.push_chunk(field, &mut ctx)?;
502            }
503        } else {
504            self.accumulators.lock()[0].push_chunk(&chunk, &mut ctx)?;
505        }
506        Ok((sequence_id, chunk))
507    }
508
509    pub fn stats_sets(&self) -> Vec<StatsSet> {
510        let mut ctx = self.ctx.lock();
511        self.accumulators
512            .lock()
513            .iter_mut()
514            .map(|acc| {
515                acc.as_stats_set(&self.stats, &mut ctx)
516                    .vortex_expect("as_stats_table should not fail")
517            })
518            .collect()
519    }
520}
521
522#[cfg(test)]
523mod tests {
524    use rstest::rstest;
525    use vortex_array::array_session;
526    use vortex_array::arrays::BoolArray;
527    use vortex_array::arrays::bool::BoolArrayExt;
528    use vortex_array::builders::VarBinViewBuilder;
529    use vortex_buffer::BitBuffer;
530    use vortex_buffer::buffer;
531
532    use super::*;
533
534    #[rstest]
535    #[case(DType::Utf8(Nullability::NonNullable))]
536    #[case(DType::Binary(Nullability::NonNullable))]
537    fn truncates_accumulated_stats(#[case] dtype: DType) {
538        let mut ctx = array_session().create_execution_ctx();
539        let mut builder = VarBinViewBuilder::with_capacity(dtype.clone(), 2);
540        builder.append_value("Value to be truncated");
541        builder.append_value("untruncated");
542        let mut builder2 = VarBinViewBuilder::with_capacity(dtype, 2);
543        builder2.append_value("Another");
544        builder2.append_value("wait a minute");
545        let mut acc =
546            StatsAccumulator::new(builder.dtype(), &[Stat::Max, Stat::Min, Stat::Sum], 12);
547        acc.push_chunk(&builder.finish(), &mut ctx)
548            .vortex_expect("push_chunk should succeed for test data");
549        acc.push_chunk(&builder2.finish(), &mut ctx)
550            .vortex_expect("push_chunk should succeed for test data");
551        let stats_table = acc
552            .as_array(&mut ctx)
553            .unwrap()
554            .expect("Must have stats table");
555        assert_eq!(
556            stats_table.names().as_ref(),
557            &[
558                Stat::Max.name(),
559                MAX_IS_TRUNCATED,
560                Stat::Min.name(),
561                MIN_IS_TRUNCATED,
562            ]
563        );
564        let field1_bool = stats_table
565            .unmasked_field(1)
566            .clone()
567            .execute::<BoolArray>(&mut ctx)
568            .unwrap();
569        assert_eq!(
570            field1_bool.to_bit_buffer(),
571            BitBuffer::from(vec![false, true])
572        );
573        let field3_bool = stats_table
574            .unmasked_field(3)
575            .clone()
576            .execute::<BoolArray>(&mut ctx)
577            .unwrap();
578        assert_eq!(
579            field3_bool.to_bit_buffer(),
580            BitBuffer::from(vec![true, false])
581        );
582    }
583
584    #[rstest]
585    #[case(DType::Utf8(Nullability::NonNullable))]
586    #[case(DType::Binary(Nullability::NonNullable))]
587    fn truncated_accumulated_stats_are_inexact(#[case] dtype: DType) {
588        let mut ctx = array_session().create_execution_ctx();
589        let mut builder = VarBinViewBuilder::with_capacity(dtype, 2);
590        builder.append_value("Value to be truncated");
591        builder.append_value("Another truncated value");
592        let mut acc = StatsAccumulator::new(builder.dtype(), &[Stat::Max, Stat::Min], 12);
593        acc.push_chunk(&builder.finish(), &mut ctx)
594            .vortex_expect("push_chunk should succeed for test data");
595
596        let stats = acc
597            .as_stats_set(&[Stat::Max, Stat::Min], &mut ctx)
598            .vortex_expect("as_stats_set should succeed for test data");
599
600        assert!(matches!(stats.get(Stat::Min), Precision::Inexact(_)));
601        assert!(matches!(stats.get(Stat::Max), Precision::Inexact(_)));
602    }
603
604    #[test]
605    fn always_adds_is_truncated_column() {
606        let mut ctx = array_session().create_execution_ctx();
607        let array = buffer![0, 1, 2].into_array();
608        let mut acc = StatsAccumulator::new(array.dtype(), &[Stat::Max, Stat::Min, Stat::Sum], 12);
609        acc.push_chunk(&array, &mut ctx)
610            .vortex_expect("push_chunk should succeed for test array");
611        let stats_table = acc
612            .as_array(&mut ctx)
613            .unwrap()
614            .expect("Must have stats table");
615        assert_eq!(
616            stats_table.names().as_ref(),
617            &[
618                Stat::Max.name(),
619                MAX_IS_TRUNCATED,
620                Stat::Min.name(),
621                MIN_IS_TRUNCATED,
622                Stat::Sum.name(),
623            ]
624        );
625        let field1_bool = stats_table
626            .unmasked_field(1)
627            .clone()
628            .execute::<BoolArray>(&mut ctx)
629            .unwrap();
630        assert_eq!(field1_bool.to_bit_buffer(), BitBuffer::from(vec![false]));
631        let field3_bool = stats_table
632            .unmasked_field(3)
633            .clone()
634            .execute::<BoolArray>(&mut ctx)
635            .unwrap();
636        assert_eq!(field3_bool.to_bit_buffer(), BitBuffer::from(vec![false]));
637    }
638}