Skip to main content

promql/extension_plan/
histogram_fold.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::any::Any;
16use std::borrow::Cow;
17use std::collections::{HashMap, HashSet};
18use std::sync::Arc;
19use std::task::Poll;
20use std::time::Instant;
21
22use common_telemetry::warn;
23use datafusion::arrow::array::{Array, ArrayRef, AsArray};
24use datafusion::arrow::compute::{SortOptions, concat_batches};
25use datafusion::arrow::datatypes::{DataType, Float64Type, SchemaRef};
26use datafusion::arrow::record_batch::RecordBatch;
27use datafusion::common::stats::Precision;
28use datafusion::common::{DFSchema, DFSchemaRef, Statistics};
29use datafusion::error::{DataFusionError, Result as DataFusionResult};
30use datafusion::execution::TaskContext;
31use datafusion::logical_expr::{LogicalPlan, UserDefinedLogicalNodeCore};
32use datafusion::physical_expr::{
33    EquivalenceProperties, LexRequirement, OrderingRequirements, PhysicalSortRequirement,
34};
35use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType};
36use datafusion::physical_plan::expressions::{Column as PhyColumn, TryCastExpr as PhyTryCast};
37use datafusion::physical_plan::metrics::{BaselineMetrics, ExecutionPlanMetricsSet, MetricsSet};
38use datafusion::physical_plan::{
39    DisplayAs, DisplayFormatType, Distribution, ExecutionPlan, ExecutionPlanProperties,
40    Partitioning, PhysicalExpr, PlanProperties, RecordBatchStream, SendableRecordBatchStream,
41};
42use datafusion::prelude::{Column, Expr};
43use datafusion_expr::{EmptyRelation, col};
44use datatypes::arrow_array::string_array_value_at_index;
45use datatypes::prelude::{ConcreteDataType, DataType as GtDataType};
46use datatypes::value::{OrderedF64, Value, ValueRef};
47use datatypes::vectors::{Helper, MutableVector, VectorRef};
48use futures::{Stream, StreamExt, ready};
49use greptime_proto::substrait_extension as pb;
50use prost::Message;
51use snafu::ResultExt;
52
53use crate::error::{DeserializeSnafu, Result};
54use crate::extension_plan::{resolve_column_name, serialize_column_index};
55
56/// `HistogramFold` will fold the conventional (non-native) histogram ([1]) for later
57/// computing.
58///
59/// Specifically, it will transform the `le` and `field` column into a complex
60/// type, and samples on other tag columns:
61/// - `le` will become a [ListArray] of [f64]. With each bucket bound parsed
62/// - `field` will become a [ListArray] of [f64]
63/// - other columns will be sampled every `bucket_num` element, but their types won't change.
64///
65/// Due to the folding or sampling, the output rows number will become `input_rows` / `bucket_num`.
66///
67/// # Requirement
68/// - Input should be sorted on `<tag list>, ts, le ASC`.
69/// - The value set of `le` should be same. I.e., buckets of every series should be same.
70///
71/// [1]: https://prometheus.io/docs/concepts/metric_types/#histogram
72#[derive(Debug, PartialEq, Hash, Eq)]
73pub struct HistogramFold {
74    /// Name of the `le` column. It's a special column in prometheus
75    /// for implementing conventional histogram. It's a string column
76    /// with "literal" float value, like "+Inf", "0.001" etc.
77    le_column: String,
78    ts_column: String,
79    input: LogicalPlan,
80    field_column: String,
81    /// Native histogram companion column for a mixed classic/native input.
82    histogram_column: Option<String>,
83    operation: HistogramFoldOperation,
84    output_schema: DFSchemaRef,
85    unfix: Option<UnfixIndices>,
86}
87
88#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd)]
89pub enum HistogramFoldOperation {
90    Quantile(OrderedF64),
91    Fraction {
92        lower: OrderedF64,
93        upper: OrderedF64,
94    },
95}
96
97impl HistogramFoldOperation {
98    pub const fn function_name(self) -> &'static str {
99        match self {
100            Self::Quantile(_) => "histogram_quantile",
101            Self::Fraction { .. } => "histogram_fraction",
102        }
103    }
104}
105
106impl std::fmt::Display for HistogramFoldOperation {
107    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
108        match self {
109            Self::Quantile(quantile) => write!(f, "quantile={quantile}"),
110            Self::Fraction { lower, upper } => {
111                write!(f, "fraction=[{lower}, {upper}]")
112            }
113        }
114    }
115}
116
117#[derive(Debug, PartialEq, Eq, Hash, PartialOrd)]
118struct UnfixIndices {
119    pub le_column_idx: u64,
120    pub ts_column_idx: u64,
121    pub field_column_idx: u64,
122}
123
124impl UserDefinedLogicalNodeCore for HistogramFold {
125    fn name(&self) -> &str {
126        Self::name()
127    }
128
129    fn inputs(&self) -> Vec<&LogicalPlan> {
130        vec![&self.input]
131    }
132
133    fn schema(&self) -> &DFSchemaRef {
134        &self.output_schema
135    }
136
137    fn expressions(&self) -> Vec<Expr> {
138        if self.unfix.is_some() {
139            return vec![];
140        }
141
142        let mut exprs = vec![
143            col(&self.le_column),
144            col(&self.ts_column),
145            col(&self.field_column),
146        ];
147        exprs.extend(self.input.schema().fields().iter().filter_map(|f| {
148            let name = f.name();
149            if name != &self.le_column && name != &self.ts_column && name != &self.field_column {
150                Some(col(name))
151            } else {
152                None
153            }
154        }));
155        exprs
156    }
157
158    fn necessary_children_exprs(&self, output_columns: &[usize]) -> Option<Vec<Vec<usize>>> {
159        if self.unfix.is_some() {
160            return None;
161        }
162
163        let input_schema = self.input.schema();
164        let le_column_index = input_schema.index_of_column_by_name(None, &self.le_column)?;
165
166        if output_columns.is_empty() {
167            let indices = (0..input_schema.fields().len()).collect::<Vec<_>>();
168            return Some(vec![indices]);
169        }
170
171        if let Some(histogram_column) = &self.histogram_column {
172            let mut necessary_indices = output_columns.to_vec();
173            for column in [
174                &self.le_column,
175                &self.ts_column,
176                &self.field_column,
177                histogram_column,
178            ] {
179                necessary_indices.push(input_schema.index_of_column_by_name(None, column)?);
180            }
181            necessary_indices.sort_unstable();
182            necessary_indices.dedup();
183            return Some(vec![necessary_indices]);
184        }
185
186        let mut necessary_indices = output_columns
187            .iter()
188            .map(|&output_column| {
189                if output_column < le_column_index {
190                    output_column
191                } else {
192                    output_column + 1
193                }
194            })
195            .collect::<Vec<_>>();
196        necessary_indices.push(le_column_index);
197        necessary_indices.sort_unstable();
198        necessary_indices.dedup();
199        Some(vec![necessary_indices])
200    }
201
202    fn fmt_for_explain(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
203        write!(
204            f,
205            "HistogramFold: le={}, field={}",
206            self.le_column, self.field_column
207        )?;
208        if let Some(histogram) = &self.histogram_column {
209            write!(f, ", histogram={histogram}")?;
210        }
211        write!(f, ", {}", self.operation)
212    }
213
214    fn with_exprs_and_inputs(
215        &self,
216        _exprs: Vec<Expr>,
217        inputs: Vec<LogicalPlan>,
218    ) -> DataFusionResult<Self> {
219        if inputs.is_empty() {
220            return Err(DataFusionError::Internal(
221                "HistogramFold must have at least one input".to_string(),
222            ));
223        }
224
225        let input: LogicalPlan = inputs.into_iter().next().unwrap();
226        let input_schema = input.schema();
227
228        if let Some(unfix) = &self.unfix {
229            let le_column =
230                resolve_column_name(unfix.le_column_idx, input_schema, "HistogramFold", "le")?;
231            let ts_column =
232                resolve_column_name(unfix.ts_column_idx, input_schema, "HistogramFold", "ts")?;
233            let field_column = resolve_column_name(
234                unfix.field_column_idx,
235                input_schema,
236                "HistogramFold",
237                "field",
238            )?;
239
240            let output_schema = Self::convert_schema(input_schema, &le_column)?;
241
242            Ok(Self {
243                le_column,
244                ts_column,
245                input,
246                field_column,
247                histogram_column: None,
248                operation: self.operation,
249                output_schema,
250                unfix: None,
251            })
252        } else {
253            Ok(Self {
254                le_column: self.le_column.clone(),
255                ts_column: self.ts_column.clone(),
256                input,
257                field_column: self.field_column.clone(),
258                histogram_column: self.histogram_column.clone(),
259                operation: self.operation,
260                output_schema: self.output_schema.clone(),
261                unfix: None,
262            })
263        }
264    }
265}
266
267impl HistogramFold {
268    pub fn new(
269        le_column: String,
270        field_column: String,
271        ts_column: String,
272        quantile: f64,
273        input: LogicalPlan,
274    ) -> DataFusionResult<Self> {
275        Self::new_with_operation(
276            le_column,
277            field_column,
278            ts_column,
279            HistogramFoldOperation::Quantile(quantile.into()),
280            None,
281            input,
282        )
283    }
284
285    pub fn new_with_operation(
286        le_column: String,
287        field_column: String,
288        ts_column: String,
289        operation: HistogramFoldOperation,
290        histogram_column: Option<String>,
291        input: LogicalPlan,
292    ) -> DataFusionResult<Self> {
293        let input_schema = input.schema();
294        Self::check_schema(input_schema, &le_column, &field_column, &ts_column)?;
295        if let Some(histogram_column) = &histogram_column {
296            Self::check_column(input_schema, histogram_column)?;
297        }
298        let output_schema = if histogram_column.is_some() {
299            input_schema.clone()
300        } else {
301            Self::convert_schema(input_schema, &le_column)?
302        };
303        Ok(Self {
304            le_column,
305            ts_column,
306            input,
307            field_column,
308            histogram_column,
309            operation,
310            output_schema,
311            unfix: None,
312        })
313    }
314
315    pub const fn name() -> &'static str {
316        "HistogramFold"
317    }
318
319    fn check_schema(
320        input_schema: &DFSchemaRef,
321        le_column: &str,
322        field_column: &str,
323        ts_column: &str,
324    ) -> DataFusionResult<()> {
325        Self::check_column(input_schema, le_column)?;
326        Self::check_column(input_schema, ts_column)?;
327        Self::check_column(input_schema, field_column)
328    }
329
330    fn check_column(input_schema: &DFSchemaRef, column: &str) -> DataFusionResult<()> {
331        if !input_schema.has_column_with_unqualified_name(column) {
332            return Err(DataFusionError::SchemaError(
333                Box::new(datafusion::common::SchemaError::FieldNotFound {
334                    field: Box::new(Column::new(None::<String>, column)),
335                    valid_fields: input_schema.columns(),
336                }),
337                Box::new(None),
338            ));
339        }
340        Ok(())
341    }
342
343    pub fn to_execution_plan(&self, exec_input: Arc<dyn ExecutionPlan>) -> Arc<dyn ExecutionPlan> {
344        let input_schema = self.input.schema();
345        // safety: those fields are checked in `check_schema()`
346        let le_column_index = input_schema
347            .index_of_column_by_name(None, &self.le_column)
348            .unwrap();
349        let field_column_index = input_schema
350            .index_of_column_by_name(None, &self.field_column)
351            .unwrap();
352        let histogram_column_index = self
353            .histogram_column
354            .as_ref()
355            .map(|column| input_schema.index_of_column_by_name(None, column).unwrap());
356        let ts_column_index = input_schema
357            .index_of_column_by_name(None, &self.ts_column)
358            .unwrap();
359
360        let tag_columns = exec_input
361            .schema()
362            .fields()
363            .iter()
364            .enumerate()
365            .filter_map(|(idx, field)| {
366                if idx == le_column_index
367                    || idx == field_column_index
368                    || Some(idx) == histogram_column_index
369                    || idx == ts_column_index
370                {
371                    None
372                } else {
373                    Some(Arc::new(PhyColumn::new(field.name(), idx)) as _)
374                }
375            })
376            .collect::<Vec<_>>();
377
378        let mut partition_exprs = tag_columns.clone();
379        partition_exprs.push(Arc::new(PhyColumn::new(
380            self.input.schema().field(ts_column_index).name(),
381            ts_column_index,
382        )) as _);
383
384        let output_schema: SchemaRef = self.output_schema.inner().clone();
385        let properties = Arc::new(PlanProperties::new(
386            EquivalenceProperties::new(output_schema.clone()),
387            Partitioning::Hash(
388                partition_exprs.clone(),
389                exec_input.output_partitioning().partition_count(),
390            ),
391            EmissionType::Incremental,
392            Boundedness::Bounded,
393        ));
394        Arc::new(HistogramFoldExec {
395            le_column_index,
396            field_column_index,
397            histogram_column_index,
398            ts_column_index,
399            input: exec_input,
400            tag_columns,
401            partition_exprs,
402            operation: self.operation,
403            output_schema,
404            metric: ExecutionPlanMetricsSet::new(),
405            properties,
406        })
407    }
408
409    /// Transform the schema
410    ///
411    /// - `le` will be removed
412    ///
413    /// Column qualifiers are preserved so downstream plan nodes can keep
414    /// referencing the columns by their original qualified names.
415    fn convert_schema(
416        input_schema: &DFSchemaRef,
417        le_column: &str,
418    ) -> DataFusionResult<DFSchemaRef> {
419        // safety: those fields are checked in `check_schema()`
420        let mut new_fields = Vec::with_capacity(input_schema.fields().len() - 1);
421        for (qualifier, field) in input_schema.iter() {
422            if field.name() != le_column {
423                new_fields.push((qualifier.cloned(), field.clone()));
424            }
425        }
426        Ok(Arc::new(DFSchema::new_with_metadata(
427            new_fields,
428            HashMap::new(),
429        )?))
430    }
431
432    pub fn serialize(&self) -> DataFusionResult<Vec<u8>> {
433        if self.histogram_column.is_some() {
434            return Err(DataFusionError::NotImplemented(
435                "mixed HistogramFold is frontend-only".to_string(),
436            ));
437        }
438        let HistogramFoldOperation::Quantile(quantile) = self.operation else {
439            return Err(DataFusionError::NotImplemented(
440                "HistogramFold fraction is frontend-only".to_string(),
441            ));
442        };
443        let le_column_idx = serialize_column_index(self.input.schema(), &self.le_column);
444        let ts_column_idx = serialize_column_index(self.input.schema(), &self.ts_column);
445        let field_column_idx = serialize_column_index(self.input.schema(), &self.field_column);
446
447        Ok(pb::HistogramFold {
448            le_column_idx,
449            ts_column_idx,
450            field_column_idx,
451            quantile: quantile.into(),
452        }
453        .encode_to_vec())
454    }
455
456    pub fn deserialize(bytes: &[u8]) -> Result<Self> {
457        let pb_histogram_fold = pb::HistogramFold::decode(bytes).context(DeserializeSnafu)?;
458        let placeholder_plan = LogicalPlan::EmptyRelation(EmptyRelation {
459            produce_one_row: false,
460            schema: Arc::new(DFSchema::empty()),
461        });
462
463        let unfix = UnfixIndices {
464            le_column_idx: pb_histogram_fold.le_column_idx,
465            ts_column_idx: pb_histogram_fold.ts_column_idx,
466            field_column_idx: pb_histogram_fold.field_column_idx,
467        };
468
469        Ok(Self {
470            le_column: String::new(),
471            ts_column: String::new(),
472            input: placeholder_plan,
473            field_column: String::new(),
474            histogram_column: None,
475            operation: HistogramFoldOperation::Quantile(pb_histogram_fold.quantile.into()),
476            output_schema: Arc::new(DFSchema::empty()),
477            unfix: Some(unfix),
478        })
479    }
480}
481
482impl PartialOrd for HistogramFold {
483    fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
484        // Compare fields in order excluding output_schema
485        match self.le_column.partial_cmp(&other.le_column) {
486            Some(core::cmp::Ordering::Equal) => {}
487            ord => return ord,
488        }
489        match self.ts_column.partial_cmp(&other.ts_column) {
490            Some(core::cmp::Ordering::Equal) => {}
491            ord => return ord,
492        }
493        match self.input.partial_cmp(&other.input) {
494            Some(core::cmp::Ordering::Equal) => {}
495            ord => return ord,
496        }
497        match self.field_column.partial_cmp(&other.field_column) {
498            Some(core::cmp::Ordering::Equal) => {}
499            ord => return ord,
500        }
501        match self.histogram_column.partial_cmp(&other.histogram_column) {
502            Some(core::cmp::Ordering::Equal) => {}
503            ord => return ord,
504        }
505        self.operation.partial_cmp(&other.operation)
506    }
507}
508
509#[derive(Debug)]
510pub struct HistogramFoldExec {
511    /// Index for `le` column in the schema of input.
512    le_column_index: usize,
513    input: Arc<dyn ExecutionPlan>,
514    output_schema: SchemaRef,
515    /// Index for field column in the schema of input.
516    field_column_index: usize,
517    /// Index for the native histogram companion column on mixed inputs.
518    histogram_column_index: Option<usize>,
519    ts_column_index: usize,
520    /// Tag columns are all columns except `le`, `field` and `ts` columns.
521    tag_columns: Vec<Arc<dyn PhysicalExpr>>,
522    partition_exprs: Vec<Arc<dyn PhysicalExpr>>,
523    operation: HistogramFoldOperation,
524    metric: ExecutionPlanMetricsSet,
525    properties: Arc<PlanProperties>,
526}
527
528impl ExecutionPlan for HistogramFoldExec {
529    fn as_any(&self) -> &dyn Any {
530        self
531    }
532
533    fn properties(&self) -> &Arc<PlanProperties> {
534        &self.properties
535    }
536
537    fn required_input_ordering(&self) -> Vec<Option<OrderingRequirements>> {
538        let mut cols = self
539            .tag_columns
540            .iter()
541            .map(|expr| PhysicalSortRequirement {
542                expr: expr.clone(),
543                options: None,
544            })
545            .collect::<Vec<PhysicalSortRequirement>>();
546        // add ts
547        cols.push(PhysicalSortRequirement {
548            expr: Arc::new(PhyColumn::new(
549                self.input.schema().field(self.ts_column_index).name(),
550                self.ts_column_index,
551            )),
552            options: None,
553        });
554        if self.histogram_column_index.is_none() {
555            // add le ASC
556            cols.push(PhysicalSortRequirement {
557                expr: Arc::new(PhyTryCast::new(
558                    Arc::new(PhyColumn::new(
559                        self.input.schema().field(self.le_column_index).name(),
560                        self.le_column_index,
561                    )),
562                    DataType::Float64,
563                )),
564                options: Some(SortOptions {
565                    descending: false,  // numeric bounds ascending
566                    nulls_first: false, // unparsable bounds last
567                }),
568            });
569        }
570
571        // Safety: `cols` is not empty
572        let requirement = LexRequirement::new(cols).unwrap();
573
574        vec![Some(OrderingRequirements::Hard(vec![requirement]))]
575    }
576
577    fn required_input_distribution(&self) -> Vec<Distribution> {
578        vec![Distribution::HashPartitioned(self.partition_exprs.clone())]
579    }
580
581    fn maintains_input_order(&self) -> Vec<bool> {
582        vec![self.histogram_column_index.is_none(); self.children().len()]
583    }
584
585    fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
586        vec![&self.input]
587    }
588
589    // cannot change schema with this method
590    fn with_new_children(
591        self: Arc<Self>,
592        children: Vec<Arc<dyn ExecutionPlan>>,
593    ) -> DataFusionResult<Arc<dyn ExecutionPlan>> {
594        assert!(!children.is_empty());
595        let new_input = children[0].clone();
596        let properties = Arc::new(PlanProperties::new(
597            EquivalenceProperties::new(self.output_schema.clone()),
598            Partitioning::Hash(
599                self.partition_exprs.clone(),
600                new_input.output_partitioning().partition_count(),
601            ),
602            EmissionType::Incremental,
603            Boundedness::Bounded,
604        ));
605        Ok(Arc::new(Self {
606            input: new_input,
607            metric: self.metric.clone(),
608            le_column_index: self.le_column_index,
609            ts_column_index: self.ts_column_index,
610            tag_columns: self.tag_columns.clone(),
611            partition_exprs: self.partition_exprs.clone(),
612            operation: self.operation,
613            output_schema: self.output_schema.clone(),
614            field_column_index: self.field_column_index,
615            histogram_column_index: self.histogram_column_index,
616            properties,
617        }))
618    }
619
620    fn execute(
621        &self,
622        partition: usize,
623        context: Arc<TaskContext>,
624    ) -> DataFusionResult<SendableRecordBatchStream> {
625        let baseline_metric = BaselineMetrics::new(&self.metric, partition);
626
627        let batch_size = context.session_config().batch_size();
628        let input = self.input.execute(partition, context)?;
629        let output_schema = self.output_schema.clone();
630
631        let mut normal_indices = (0..input.schema().fields().len()).collect::<HashSet<_>>();
632        normal_indices.remove(&self.field_column_index);
633        normal_indices.remove(&self.le_column_index);
634        if let Some(histogram_column_index) = self.histogram_column_index {
635            normal_indices.remove(&histogram_column_index);
636        }
637        let mode = if self.histogram_column_index.is_some() {
638            FoldMode::Safe
639        } else {
640            FoldMode::Optimistic
641        };
642        Ok(Box::pin(HistogramFoldStream {
643            le_column_index: self.le_column_index,
644            field_column_index: self.field_column_index,
645            histogram_column_index: self.histogram_column_index,
646            operation: self.operation,
647            normal_indices: normal_indices.into_iter().collect(),
648            bucket_size: None,
649            input_buffer: vec![],
650            input,
651            output_schema,
652            input_schema: self.input.schema(),
653            mode,
654            safe_group: None,
655            metric: baseline_metric,
656            batch_size,
657            input_buffered_rows: 0,
658            output_buffer: HistogramFoldStream::empty_output_buffer(&self.input.schema())?,
659            output_buffered_rows: 0,
660        }))
661    }
662
663    fn metrics(&self) -> Option<MetricsSet> {
664        Some(self.metric.clone_inner())
665    }
666
667    fn partition_statistics(&self, _: Option<usize>) -> DataFusionResult<Statistics> {
668        Ok(Statistics {
669            num_rows: Precision::Absent,
670            total_byte_size: Precision::Absent,
671            column_statistics: Statistics::unknown_column(&self.schema()),
672        })
673    }
674
675    fn name(&self) -> &str {
676        "HistogramFoldExec"
677    }
678}
679
680impl DisplayAs for HistogramFoldExec {
681    fn fmt_as(&self, t: DisplayFormatType, f: &mut std::fmt::Formatter) -> std::fmt::Result {
682        match t {
683            DisplayFormatType::Default
684            | DisplayFormatType::Verbose
685            | DisplayFormatType::TreeRender => {
686                write!(
687                    f,
688                    "HistogramFoldExec: le=@{}, field=@{}",
689                    self.le_column_index, self.field_column_index
690                )?;
691                if let Some(histogram) = self.histogram_column_index {
692                    write!(f, ", histogram=@{histogram}")?;
693                }
694                write!(f, ", {}", self.operation)
695            }
696        }
697    }
698}
699
700#[derive(Debug, Clone, Copy, PartialEq, Eq)]
701enum FoldMode {
702    Optimistic,
703    Safe,
704}
705
706pub struct HistogramFoldStream {
707    // internal states
708    le_column_index: usize,
709    field_column_index: usize,
710    histogram_column_index: Option<usize>,
711    operation: HistogramFoldOperation,
712    /// Columns need not folding. This indices is based on input schema
713    normal_indices: Vec<usize>,
714    bucket_size: Option<usize>,
715    /// Expected output batch size
716    batch_size: usize,
717    output_schema: SchemaRef,
718    input_schema: SchemaRef,
719    mode: FoldMode,
720    safe_group: Option<SafeGroup>,
721
722    // buffers
723    input_buffer: Vec<RecordBatch>,
724    input_buffered_rows: usize,
725    output_buffer: Vec<Box<dyn MutableVector>>,
726    output_buffered_rows: usize,
727
728    // runtime things
729    input: SendableRecordBatchStream,
730    metric: BaselineMetrics,
731}
732
733#[derive(Debug, Default)]
734struct SafeGroup {
735    tag_values: Vec<Value>,
736    buckets: Vec<f64>,
737    counters: Vec<f64>,
738    native_samples: Vec<(Value, Value)>,
739}
740
741impl RecordBatchStream for HistogramFoldStream {
742    fn schema(&self) -> SchemaRef {
743        self.output_schema.clone()
744    }
745}
746
747impl Stream for HistogramFoldStream {
748    type Item = DataFusionResult<RecordBatch>;
749
750    fn poll_next(
751        mut self: std::pin::Pin<&mut Self>,
752        cx: &mut std::task::Context<'_>,
753    ) -> Poll<Option<Self::Item>> {
754        let poll = loop {
755            match ready!(self.input.poll_next_unpin(cx)) {
756                Some(batch) => {
757                    let batch = batch?;
758                    let timer = Instant::now();
759                    let Some(result) = self.fold_input(batch)? else {
760                        self.metric.elapsed_compute().add_elapsed(timer);
761                        continue;
762                    };
763                    self.metric.elapsed_compute().add_elapsed(timer);
764                    break Poll::Ready(Some(result));
765                }
766                None => {
767                    self.flush_remaining()?;
768                    break Poll::Ready(self.take_output_buf()?.map(Ok));
769                }
770            }
771        };
772        self.metric.record_poll(poll)
773    }
774}
775
776impl HistogramFoldStream {
777    /// The inner most `Result` is for `poll_next()`
778    pub fn fold_input(
779        &mut self,
780        input: RecordBatch,
781    ) -> DataFusionResult<Option<DataFusionResult<RecordBatch>>> {
782        match self.mode {
783            FoldMode::Safe => {
784                self.push_input_buf(input);
785                self.process_safe_mode_buffer()?;
786            }
787            FoldMode::Optimistic => {
788                self.push_input_buf(input);
789                let Some(bucket_num) = self.calculate_bucket_num_from_buffer()? else {
790                    return Ok(None);
791                };
792                self.bucket_size = Some(bucket_num);
793
794                if self.input_buffered_rows < bucket_num {
795                    // not enough rows to fold
796                    return Ok(None);
797                }
798
799                self.fold_buf(bucket_num)?;
800            }
801        }
802
803        self.maybe_take_output()
804    }
805
806    /// Generate input-aligned output builders. Classic-only output drops `le` later.
807    pub fn empty_output_buffer(
808        schema: &SchemaRef,
809    ) -> DataFusionResult<Vec<Box<dyn MutableVector>>> {
810        let mut builders = Vec::with_capacity(schema.fields().len());
811        for field in schema.fields() {
812            let concrete_datatype = ConcreteDataType::try_from(field.data_type()).unwrap();
813            let mutable_vector = concrete_datatype.create_mutable_vector(0);
814            builders.push(mutable_vector);
815        }
816
817        Ok(builders)
818    }
819
820    /// Determines bucket count using buffered batches, concatenating them to
821    /// detect the first complete bucket that may span batch boundaries.
822    fn calculate_bucket_num_from_buffer(&mut self) -> DataFusionResult<Option<usize>> {
823        if let Some(size) = self.bucket_size {
824            return Ok(Some(size));
825        }
826
827        if self.input_buffer.is_empty() {
828            return Ok(None);
829        }
830
831        let batch_refs: Vec<&RecordBatch> = self.input_buffer.iter().collect();
832        let batch = concat_batches(&self.input_schema, batch_refs)?;
833        self.find_first_complete_bucket(&batch)
834    }
835
836    fn find_first_complete_bucket(&self, batch: &RecordBatch) -> DataFusionResult<Option<usize>> {
837        if batch.num_rows() == 0 {
838            return Ok(None);
839        }
840
841        let vectors = Helper::try_into_vectors(batch.columns())
842            .map_err(|e| DataFusionError::Execution(e.to_string()))?;
843        let le_array = batch.column(self.le_column_index);
844
845        let mut tag_values_buf = Vec::with_capacity(self.normal_indices.len());
846        self.collect_tag_values(&vectors, 0, &mut tag_values_buf);
847        let mut group_start = 0usize;
848
849        for row in 0..batch.num_rows() {
850            if !self.is_same_group(&vectors, row, &tag_values_buf) {
851                // new group begins
852                self.collect_tag_values(&vectors, row, &mut tag_values_buf);
853                group_start = row;
854            }
855
856            if Self::is_positive_infinity(le_array, row) {
857                return Ok(Some(row - group_start + 1));
858            }
859        }
860
861        Ok(None)
862    }
863
864    /// Fold record batches from input buffer and put to output buffer
865    fn fold_buf(&mut self, bucket_num: usize) -> DataFusionResult<()> {
866        let batch = concat_batches(&self.input_schema, self.input_buffer.drain(..).as_ref())?;
867        let mut remaining_rows = self.input_buffered_rows;
868        let mut cursor = 0;
869
870        // TODO(LFC): Try to get rid of the Arrow array to vector conversion here.
871        let vectors = Helper::try_into_vectors(batch.columns())
872            .map_err(|e| DataFusionError::Execution(e.to_string()))?;
873        let le_array = batch.column(self.le_column_index);
874        let field_array = batch.column(self.field_column_index);
875        let field_array = field_array.as_primitive::<Float64Type>();
876        let mut tag_values_buf = Vec::with_capacity(self.normal_indices.len());
877
878        while remaining_rows >= bucket_num && self.mode == FoldMode::Optimistic {
879            self.collect_tag_values(&vectors, cursor, &mut tag_values_buf);
880            if !self.validate_optimistic_group(
881                &vectors,
882                le_array,
883                cursor,
884                bucket_num,
885                &tag_values_buf,
886            ) {
887                let remaining_input_batch = batch.slice(cursor, remaining_rows);
888                self.switch_to_safe_mode(remaining_input_batch)?;
889                return Ok(());
890            }
891
892            // "sample" normal columns
893            for (idx, value) in self.normal_indices.iter().zip(tag_values_buf.iter()) {
894                self.output_buffer[*idx].push_value_ref(value);
895            }
896            // "fold" `le` and field columns
897            let mut bucket = Vec::with_capacity(bucket_num);
898            let mut counters = Vec::with_capacity(bucket_num);
899            for bias in 0..bucket_num {
900                let position = cursor + bias;
901                let le = string_array_value_at_index(le_array, position)
902                    .and_then(|value| value.parse::<f64>().ok())
903                    .unwrap_or(f64::NAN);
904                bucket.push(le);
905
906                let counter = if field_array.is_valid(position) {
907                    field_array.value(position)
908                } else {
909                    f64::NAN
910                };
911                counters.push(counter);
912            }
913            // ignore invalid data
914            let result = Self::evaluate_row(self.operation, &bucket, &counters).unwrap_or(f64::NAN);
915            self.output_buffer[self.field_column_index].push_value_ref(&ValueRef::from(result));
916            cursor += bucket_num;
917            remaining_rows -= bucket_num;
918            self.output_buffered_rows += 1;
919        }
920
921        let remaining_input_batch = batch.slice(cursor, remaining_rows);
922        self.input_buffered_rows = remaining_input_batch.num_rows();
923        if self.input_buffered_rows > 0 {
924            self.input_buffer.push(remaining_input_batch);
925        }
926
927        Ok(())
928    }
929
930    fn push_input_buf(&mut self, batch: RecordBatch) {
931        self.input_buffered_rows += batch.num_rows();
932        self.input_buffer.push(batch);
933    }
934
935    fn maybe_take_output(&mut self) -> DataFusionResult<Option<DataFusionResult<RecordBatch>>> {
936        if self.output_buffered_rows >= self.batch_size {
937            return Ok(self.take_output_buf()?.map(Ok));
938        }
939        Ok(None)
940    }
941
942    fn switch_to_safe_mode(&mut self, remaining_batch: RecordBatch) -> DataFusionResult<()> {
943        self.mode = FoldMode::Safe;
944        self.bucket_size = None;
945        self.input_buffer.clear();
946        self.input_buffered_rows = remaining_batch.num_rows();
947
948        if self.input_buffered_rows > 0 {
949            self.input_buffer.push(remaining_batch);
950            self.process_safe_mode_buffer()?;
951        }
952
953        Ok(())
954    }
955
956    fn collect_tag_values<'a>(
957        &self,
958        vectors: &'a [VectorRef],
959        row: usize,
960        tag_values: &mut Vec<ValueRef<'a>>,
961    ) {
962        tag_values.clear();
963        for idx in self.normal_indices.iter() {
964            tag_values.push(vectors[*idx].get_ref(row));
965        }
966    }
967
968    fn validate_optimistic_group(
969        &self,
970        vectors: &[VectorRef],
971        le_array: &ArrayRef,
972        cursor: usize,
973        bucket_num: usize,
974        tag_values: &[ValueRef<'_>],
975    ) -> bool {
976        let inf_index = cursor + bucket_num - 1;
977        if !Self::is_positive_infinity(le_array, inf_index) {
978            return false;
979        }
980        if (cursor..=inf_index).any(|row| {
981            string_array_value_at_index(le_array, row)
982                .and_then(|value| value.parse::<f64>().ok())
983                .is_none()
984        }) {
985            return false;
986        }
987
988        for offset in 1..bucket_num {
989            let row = cursor + offset;
990            for (idx, expected) in self.normal_indices.iter().zip(tag_values.iter()) {
991                if vectors[*idx].get_ref(row) != *expected {
992                    return false;
993                }
994            }
995        }
996        true
997    }
998
999    /// Checks whether a row belongs to the current group (same series).
1000    fn is_same_group(
1001        &self,
1002        vectors: &[VectorRef],
1003        row: usize,
1004        tag_values: &[ValueRef<'_>],
1005    ) -> bool {
1006        self.normal_indices
1007            .iter()
1008            .zip(tag_values.iter())
1009            .all(|(idx, expected)| vectors[*idx].get_ref(row) == *expected)
1010    }
1011
1012    fn push_output_row(&mut self, tag_values: &[ValueRef<'_>], result: f64) {
1013        debug_assert_eq!(self.normal_indices.len(), tag_values.len());
1014        for (idx, value) in self.normal_indices.iter().zip(tag_values.iter()) {
1015            self.output_buffer[*idx].push_value_ref(value);
1016        }
1017        self.output_buffer[self.field_column_index].push_value_ref(&ValueRef::from(result));
1018        self.output_buffered_rows += 1;
1019    }
1020
1021    fn push_mixed_output_row(
1022        &mut self,
1023        tag_values: &[Value],
1024        le: &Value,
1025        result: Option<f64>,
1026        histogram: &Value,
1027    ) {
1028        let histogram_column_index = self.histogram_column_index.unwrap();
1029        for (idx, value) in self.normal_indices.iter().zip(tag_values) {
1030            self.output_buffer[*idx].push_value_ref(&value.as_value_ref());
1031        }
1032        self.output_buffer[self.le_column_index].push_value_ref(&le.as_value_ref());
1033        self.output_buffer[self.field_column_index]
1034            .push_value_ref(&result.map_or(ValueRef::Null, ValueRef::from));
1035        self.output_buffer[histogram_column_index].push_value_ref(&histogram.as_value_ref());
1036        self.output_buffered_rows += 1;
1037    }
1038
1039    fn finalize_safe_group(&mut self) -> DataFusionResult<()> {
1040        let Some(group) = self.safe_group.take() else {
1041            return Ok(());
1042        };
1043        if group.tag_values.is_empty() {
1044            return Ok(());
1045        }
1046
1047        if self.histogram_column_index.is_some() {
1048            let classic_result = if group.buckets.is_empty() {
1049                None
1050            } else {
1051                let mut buckets = group
1052                    .buckets
1053                    .into_iter()
1054                    .zip(group.counters)
1055                    .collect::<Vec<_>>();
1056                buckets.sort_by(|lhs, rhs| lhs.0.total_cmp(&rhs.0));
1057                let (bounds, counters): (Vec<_>, Vec<_>) = buckets.into_iter().unzip();
1058                let has_inf = bounds
1059                    .last()
1060                    .is_some_and(|value| value.is_infinite() && value.is_sign_positive());
1061                Some(if has_inf {
1062                    Self::evaluate_row(self.operation, &bounds, &counters).unwrap_or(f64::NAN)
1063                } else {
1064                    f64::NAN
1065                })
1066            };
1067            let has_null_native = group.native_samples.iter().any(|(le, _)| le.is_null());
1068            for (le, histogram) in group.native_samples.iter().filter(|(le, _)| !le.is_null()) {
1069                self.push_mixed_output_row(&group.tag_values, le, None, histogram);
1070            }
1071            for (le, histogram) in group.native_samples.iter().filter(|(le, _)| le.is_null()) {
1072                self.push_mixed_output_row(&group.tag_values, le, classic_result, histogram);
1073            }
1074            if !has_null_native && let Some(result) = classic_result {
1075                self.push_mixed_output_row(
1076                    &group.tag_values,
1077                    &Value::Null,
1078                    Some(result),
1079                    &Value::Null,
1080                );
1081            }
1082            return Ok(());
1083        }
1084        if group.buckets.is_empty() {
1085            return Ok(());
1086        }
1087
1088        let has_inf = group
1089            .buckets
1090            .last()
1091            .is_some_and(|value| value.is_infinite() && value.is_sign_positive());
1092        let result = if has_inf {
1093            Self::evaluate_row(self.operation, &group.buckets, &group.counters).unwrap_or(f64::NAN)
1094        } else {
1095            f64::NAN
1096        };
1097        let tag_value_refs = group
1098            .tag_values
1099            .iter()
1100            .map(Value::as_value_ref)
1101            .collect::<Vec<_>>();
1102        self.push_output_row(&tag_value_refs, result);
1103        Ok(())
1104    }
1105
1106    fn process_safe_mode_buffer(&mut self) -> DataFusionResult<()> {
1107        if self.input_buffer.is_empty() {
1108            self.input_buffered_rows = 0;
1109            return Ok(());
1110        }
1111
1112        let batch = concat_batches(&self.input_schema, self.input_buffer.drain(..).as_ref())?;
1113        self.input_buffered_rows = 0;
1114        let vectors = Helper::try_into_vectors(batch.columns())
1115            .map_err(|e| DataFusionError::Execution(e.to_string()))?;
1116        let le_array = batch.column(self.le_column_index);
1117        let field_array = batch
1118            .column(self.field_column_index)
1119            .as_primitive::<Float64Type>();
1120        let mut tag_values_buf = Vec::with_capacity(self.normal_indices.len());
1121
1122        for row in 0..batch.num_rows() {
1123            self.collect_tag_values(&vectors, row, &mut tag_values_buf);
1124            let should_start_new_group = self
1125                .safe_group
1126                .as_ref()
1127                .is_none_or(|group| !Self::tag_values_equal(&group.tag_values, &tag_values_buf));
1128            if should_start_new_group {
1129                self.finalize_safe_group()?;
1130                self.safe_group = Some(SafeGroup {
1131                    tag_values: tag_values_buf.iter().cloned().map(Value::from).collect(),
1132                    buckets: Vec::new(),
1133                    counters: Vec::new(),
1134                    native_samples: Vec::new(),
1135                });
1136            }
1137
1138            let Some(group) = self.safe_group.as_mut() else {
1139                continue;
1140            };
1141
1142            let mixed = self.histogram_column_index.is_some();
1143            let bucket = string_array_value_at_index(le_array, row)
1144                .and_then(|value| value.parse::<f64>().ok());
1145            if let Some(bucket) = bucket
1146                && (!mixed || field_array.is_valid(row))
1147            {
1148                let counter = if field_array.is_valid(row) {
1149                    field_array.value(row)
1150                } else {
1151                    f64::NAN
1152                };
1153                group.buckets.push(bucket);
1154                group.counters.push(counter);
1155            }
1156            if let Some(histogram_column_index) = self.histogram_column_index {
1157                let histogram = vectors[histogram_column_index].get(row);
1158                if !histogram.is_null() {
1159                    group
1160                        .native_samples
1161                        .push((vectors[self.le_column_index].get(row), histogram));
1162                }
1163            }
1164        }
1165
1166        Ok(())
1167    }
1168
1169    fn tag_values_equal(group_values: &[Value], current: &[ValueRef<'_>]) -> bool {
1170        group_values.len() == current.len()
1171            && group_values
1172                .iter()
1173                .zip(current.iter())
1174                .all(|(group, now)| group.as_value_ref() == *now)
1175    }
1176
1177    /// Compute result from output buffer
1178    fn take_output_buf(&mut self) -> DataFusionResult<Option<RecordBatch>> {
1179        if self.output_buffered_rows == 0 {
1180            if self.input_buffered_rows != 0 {
1181                warn!(
1182                    "input buffer is not empty, {} rows remaining",
1183                    self.input_buffered_rows
1184                );
1185            }
1186            return Ok(None);
1187        }
1188
1189        let mut output_buf = Self::empty_output_buffer(&self.input_schema)?;
1190        std::mem::swap(&mut self.output_buffer, &mut output_buf);
1191        let mut columns = Vec::with_capacity(output_buf.len());
1192        for builder in output_buf.iter_mut() {
1193            columns.push(builder.to_vector().to_arrow_array());
1194        }
1195        if self.histogram_column_index.is_none() {
1196            columns.remove(self.le_column_index);
1197        }
1198
1199        self.output_buffered_rows = 0;
1200        RecordBatch::try_new(self.output_schema.clone(), columns)
1201            .map(Some)
1202            .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))
1203    }
1204
1205    fn flush_remaining(&mut self) -> DataFusionResult<()> {
1206        if self.mode == FoldMode::Optimistic && self.input_buffered_rows > 0 {
1207            let buffered_batches: Vec<_> = self.input_buffer.drain(..).collect();
1208            if !buffered_batches.is_empty() {
1209                let batch = concat_batches(&self.input_schema, buffered_batches.as_slice())?;
1210                self.switch_to_safe_mode(batch)?;
1211            } else {
1212                self.input_buffered_rows = 0;
1213            }
1214        }
1215
1216        if self.mode == FoldMode::Safe {
1217            self.process_safe_mode_buffer()?;
1218            self.finalize_safe_group()?;
1219        }
1220
1221        Ok(())
1222    }
1223
1224    fn is_positive_infinity(le_array: &ArrayRef, index: usize) -> bool {
1225        matches!(
1226            string_array_value_at_index(le_array, index).and_then(|value| value.parse::<f64>().ok()),
1227            Some(value) if value.is_infinite() && value.is_sign_positive()
1228        )
1229    }
1230
1231    /// Evaluate the field column and return the result
1232    fn evaluate_row(
1233        operation: HistogramFoldOperation,
1234        bucket: &[f64],
1235        counter: &[f64],
1236    ) -> DataFusionResult<f64> {
1237        // check bucket
1238        if bucket.is_empty()
1239            || matches!(operation, HistogramFoldOperation::Quantile(_)) && bucket.len() == 1
1240        {
1241            return Ok(f64::NAN);
1242        }
1243        if bucket.last() != Some(&f64::INFINITY) {
1244            return Err(DataFusionError::Execution(
1245                "last bucket should be +Inf".to_string(),
1246            ));
1247        }
1248        if bucket.len() != counter.len() {
1249            return Err(DataFusionError::Execution(
1250                "bucket and counter should have the same length".to_string(),
1251            ));
1252        }
1253        if let HistogramFoldOperation::Quantile(quantile) = operation {
1254            let quantile = f64::from(quantile);
1255            if quantile < 0.0 {
1256                return Ok(f64::NEG_INFINITY);
1257            } else if quantile > 1.0 {
1258                return Ok(f64::INFINITY);
1259            } else if quantile.is_nan() {
1260                return Ok(f64::NAN);
1261            }
1262        }
1263
1264        // check input value
1265        if !bucket.windows(2).all(|w| w[0] <= w[1]) {
1266            return Ok(f64::NAN);
1267        }
1268        let counter = match operation {
1269            HistogramFoldOperation::Quantile(_) => {
1270                let needs_fix = counter.iter().any(|v| !v.is_finite())
1271                    || !counter.windows(2).all(|w| w[0] <= w[1]);
1272                if !needs_fix {
1273                    Cow::Borrowed(counter)
1274                } else {
1275                    let mut fixed = Vec::with_capacity(counter.len());
1276                    let mut prev = 0.0;
1277                    for (idx, &v) in counter.iter().enumerate() {
1278                        let mut val = if v.is_finite() { v } else { prev };
1279                        if idx > 0 && val < prev {
1280                            val = prev;
1281                        }
1282                        fixed.push(val);
1283                        prev = val;
1284                    }
1285                    Cow::Owned(fixed)
1286                }
1287            }
1288            HistogramFoldOperation::Fraction { .. } => Cow::Borrowed(counter),
1289        };
1290
1291        Ok(match operation {
1292            HistogramFoldOperation::Quantile(quantile) => {
1293                Self::evaluate_quantile(quantile.into(), bucket, &counter)
1294            }
1295            HistogramFoldOperation::Fraction { lower, upper } => {
1296                Self::evaluate_fraction(lower.into(), upper.into(), bucket, &counter)
1297            }
1298        })
1299    }
1300
1301    fn evaluate_quantile(quantile: f64, bucket: &[f64], counter: &[f64]) -> f64 {
1302        let total = *counter.last().unwrap();
1303        let expected_pos = total * quantile;
1304        let mut fit_bucket_pos = 0;
1305        while fit_bucket_pos < bucket.len() && counter[fit_bucket_pos] < expected_pos {
1306            fit_bucket_pos += 1;
1307        }
1308        if fit_bucket_pos >= bucket.len() - 1 {
1309            bucket[bucket.len() - 2]
1310        } else {
1311            let upper_bound = bucket[fit_bucket_pos];
1312            let upper_count = counter[fit_bucket_pos];
1313            let mut lower_bound = bucket[0].min(0.0);
1314            let mut lower_count = 0.0;
1315            if fit_bucket_pos > 0 {
1316                lower_bound = bucket[fit_bucket_pos - 1];
1317                lower_count = counter[fit_bucket_pos - 1];
1318            }
1319            if (upper_count - lower_count).abs() < 1e-10 {
1320                return f64::NAN;
1321            }
1322            lower_bound
1323                + (upper_bound - lower_bound) / (upper_count - lower_count)
1324                    * (expected_pos - lower_count)
1325        }
1326    }
1327
1328    fn evaluate_fraction(lower: f64, upper: f64, bucket: &[f64], counter: &[f64]) -> f64 {
1329        let coalesced = bucket
1330            .windows(2)
1331            .any(|bounds| bounds[0] == bounds[1])
1332            .then(|| {
1333                let mut bounds = Vec::with_capacity(bucket.len());
1334                let mut counts = Vec::with_capacity(counter.len());
1335                for (&bound, &count) in bucket.iter().zip(counter) {
1336                    if bounds.last() == Some(&bound) {
1337                        *counts.last_mut().unwrap() += count;
1338                    } else {
1339                        bounds.push(bound);
1340                        counts.push(count);
1341                    }
1342                }
1343                (bounds, counts)
1344            });
1345        let (bucket, counter) = match &coalesced {
1346            Some((bounds, counts)) => (bounds.as_slice(), counts.as_slice()),
1347            None => (bucket, counter),
1348        };
1349        let total = *counter.last().unwrap();
1350        if total == 0.0 || lower.is_nan() || upper.is_nan() {
1351            return f64::NAN;
1352        }
1353        if lower >= upper {
1354            return 0.0;
1355        }
1356
1357        let mut rank = 0.0;
1358        let mut lower_rank = 0.0;
1359        let mut upper_rank = 0.0;
1360        let mut lower_set = false;
1361        let mut upper_set = false;
1362        let mut lower_bound = if bucket[0] > 0.0 {
1363            0.0
1364        } else {
1365            f64::NEG_INFINITY
1366        };
1367
1368        for (idx, (&upper_bound, &upper_count)) in bucket.iter().zip(counter).enumerate() {
1369            if idx > 0 {
1370                lower_bound = bucket[idx - 1];
1371            }
1372            let interpolate = |value: f64| {
1373                if lower_bound == f64::NEG_INFINITY {
1374                    upper_count
1375                } else {
1376                    rank + (upper_count - rank) * (value - lower_bound)
1377                        / (upper_bound - lower_bound)
1378                }
1379            };
1380
1381            if !lower_set && lower_bound >= lower {
1382                lower_rank = rank;
1383                lower_set = true;
1384            }
1385            if !upper_set && lower_bound >= upper {
1386                upper_rank = rank;
1387                upper_set = true;
1388            }
1389            if lower_set && upper_set {
1390                break;
1391            }
1392            if !lower_set && lower_bound < lower && upper_bound > lower {
1393                lower_rank = interpolate(lower);
1394                lower_set = true;
1395            }
1396            if !upper_set && lower_bound < upper && upper_bound > upper {
1397                upper_rank = interpolate(upper);
1398                upper_set = true;
1399            }
1400            if lower_set && upper_set {
1401                break;
1402            }
1403            rank = upper_count;
1404        }
1405
1406        if !lower_set || lower_rank > total {
1407            lower_rank = total;
1408        }
1409        if !upper_set || upper_rank > total {
1410            upper_rank = total;
1411        }
1412        (upper_rank - lower_rank) / total
1413    }
1414}
1415
1416#[cfg(test)]
1417mod test {
1418    use std::sync::Arc;
1419
1420    use datafusion::arrow::array::{
1421        DictionaryArray, Float64Array, StringDictionaryBuilder, TimestampMillisecondArray,
1422    };
1423    use datafusion::arrow::datatypes::{Field, Schema, SchemaRef, TimeUnit, UInt32Type};
1424    use datafusion::common::ToDFSchema;
1425    use datafusion::datasource::memory::MemorySourceConfig;
1426    use datafusion::datasource::source::DataSourceExec;
1427    use datafusion::logical_expr::EmptyRelation;
1428    use datafusion::prelude::SessionContext;
1429    use datatypes::arrow_array::StringArray;
1430    use futures::FutureExt;
1431
1432    use super::*;
1433
1434    fn project_batch(batch: &RecordBatch, indices: &[usize]) -> RecordBatch {
1435        let fields = indices
1436            .iter()
1437            .map(|&idx| batch.schema().field(idx).clone())
1438            .collect::<Vec<_>>();
1439        let columns = indices
1440            .iter()
1441            .map(|&idx| batch.column(idx).clone())
1442            .collect::<Vec<_>>();
1443        let schema = Arc::new(Schema::new(fields));
1444        RecordBatch::try_new(schema, columns).unwrap()
1445    }
1446
1447    fn prepare_test_data() -> DataSourceExec {
1448        let schema = Arc::new(Schema::new(vec![
1449            Field::new("host", DataType::Utf8, true),
1450            Field::new("le", DataType::Utf8, true),
1451            Field::new("val", DataType::Float64, true),
1452        ]));
1453
1454        // 12 items
1455        let host_column_1 = Arc::new(StringArray::from(vec![
1456            "host_1", "host_1", "host_1", "host_1", "host_1", "host_1", "host_1", "host_1",
1457            "host_1", "host_1", "host_1", "host_1",
1458        ])) as _;
1459        let le_column_1 = Arc::new(StringArray::from(vec![
1460            "0.001", "0.1", "10", "1000", "+Inf", "0.001", "0.1", "10", "1000", "+inf", "0.001",
1461            "0.1",
1462        ])) as _;
1463        let val_column_1 = Arc::new(Float64Array::from(vec![
1464            0_0.0, 1.0, 1.0, 5.0, 5.0, 0_0.0, 20.0, 60.0, 70.0, 100.0, 0_1.0, 1.0,
1465        ])) as _;
1466
1467        // 2 items
1468        let host_column_2 = Arc::new(StringArray::from(vec!["host_1", "host_1"])) as _;
1469        let le_column_2 = Arc::new(StringArray::from(vec!["10", "1000"])) as _;
1470        let val_column_2 = Arc::new(Float64Array::from(vec![1.0, 1.0])) as _;
1471
1472        // 11 items
1473        let host_column_3 = Arc::new(StringArray::from(vec![
1474            "host_1", "host_2", "host_2", "host_2", "host_2", "host_2", "host_2", "host_2",
1475            "host_2", "host_2", "host_2",
1476        ])) as _;
1477        let le_column_3 = Arc::new(StringArray::from(vec![
1478            "+INF", "0.001", "0.1", "10", "1000", "+iNf", "0.001", "0.1", "10", "1000", "+Inf",
1479        ])) as _;
1480        let val_column_3 = Arc::new(Float64Array::from(vec![
1481            1.0, 0_0.0, 0.0, 0.0, 0.0, 0.0, 0_0.0, 1.0, 2.0, 3.0, 4.0,
1482        ])) as _;
1483
1484        let data_1 = RecordBatch::try_new(
1485            schema.clone(),
1486            vec![host_column_1, le_column_1, val_column_1],
1487        )
1488        .unwrap();
1489        let data_2 = RecordBatch::try_new(
1490            schema.clone(),
1491            vec![host_column_2, le_column_2, val_column_2],
1492        )
1493        .unwrap();
1494        let data_3 = RecordBatch::try_new(
1495            schema.clone(),
1496            vec![host_column_3, le_column_3, val_column_3],
1497        )
1498        .unwrap();
1499
1500        DataSourceExec::new(Arc::new(
1501            MemorySourceConfig::try_new(&[vec![data_1, data_2, data_3]], schema, None).unwrap(),
1502        ))
1503    }
1504
1505    fn build_fold_exec_from_batches(
1506        batches: Vec<RecordBatch>,
1507        schema: SchemaRef,
1508        quantile: f64,
1509        ts_column_index: usize,
1510    ) -> Arc<HistogramFoldExec> {
1511        build_fold_exec_from_batches_with_operation(
1512            batches,
1513            schema,
1514            HistogramFoldOperation::Quantile(quantile.into()),
1515            ts_column_index,
1516        )
1517    }
1518
1519    fn build_fold_exec_from_batches_with_operation(
1520        batches: Vec<RecordBatch>,
1521        schema: SchemaRef,
1522        operation: HistogramFoldOperation,
1523        ts_column_index: usize,
1524    ) -> Arc<HistogramFoldExec> {
1525        let input: Arc<dyn ExecutionPlan> = Arc::new(DataSourceExec::new(Arc::new(
1526            MemorySourceConfig::try_new(&[batches], schema.clone(), None).unwrap(),
1527        )));
1528        let output_schema: SchemaRef = Arc::new(
1529            HistogramFold::convert_schema(&Arc::new(input.schema().to_dfschema().unwrap()), "le")
1530                .unwrap()
1531                .as_arrow()
1532                .clone(),
1533        );
1534
1535        let (tag_columns, partition_exprs, properties) =
1536            build_test_plan_properties(&input, output_schema.clone(), ts_column_index);
1537
1538        Arc::new(HistogramFoldExec {
1539            le_column_index: 1,
1540            field_column_index: 2,
1541            histogram_column_index: None,
1542            operation,
1543            ts_column_index,
1544            input,
1545            output_schema,
1546            tag_columns,
1547            partition_exprs,
1548            metric: ExecutionPlanMetricsSet::new(),
1549            properties,
1550        })
1551    }
1552
1553    type PlanPropsResult = (
1554        Vec<Arc<dyn PhysicalExpr>>,
1555        Vec<Arc<dyn PhysicalExpr>>,
1556        Arc<PlanProperties>,
1557    );
1558
1559    fn build_test_plan_properties(
1560        input: &Arc<dyn ExecutionPlan>,
1561        output_schema: SchemaRef,
1562        ts_column_index: usize,
1563    ) -> PlanPropsResult {
1564        let tag_columns = input
1565            .schema()
1566            .fields()
1567            .iter()
1568            .enumerate()
1569            .filter_map(|(idx, field)| {
1570                if idx == 1 || idx == 2 || idx == ts_column_index {
1571                    None
1572                } else {
1573                    Some(Arc::new(PhyColumn::new(field.name(), idx)) as _)
1574                }
1575            })
1576            .collect::<Vec<_>>();
1577
1578        let partition_exprs = if tag_columns.is_empty() {
1579            vec![Arc::new(PhyColumn::new(
1580                input.schema().field(ts_column_index).name(),
1581                ts_column_index,
1582            )) as _]
1583        } else {
1584            tag_columns.clone()
1585        };
1586
1587        let properties = PlanProperties::new(
1588            EquivalenceProperties::new(output_schema.clone()),
1589            Partitioning::Hash(
1590                partition_exprs.clone(),
1591                input.output_partitioning().partition_count(),
1592            ),
1593            EmissionType::Incremental,
1594            Boundedness::Bounded,
1595        );
1596
1597        (tag_columns, partition_exprs, Arc::new(properties))
1598    }
1599
1600    #[tokio::test]
1601    async fn fold_overall() {
1602        let memory_exec: Arc<dyn ExecutionPlan> = Arc::new(prepare_test_data());
1603        let output_schema: SchemaRef = Arc::new(
1604            HistogramFold::convert_schema(
1605                &Arc::new(memory_exec.schema().to_dfschema().unwrap()),
1606                "le",
1607            )
1608            .unwrap()
1609            .as_arrow()
1610            .clone(),
1611        );
1612        let (tag_columns, partition_exprs, properties) =
1613            build_test_plan_properties(&memory_exec, output_schema.clone(), 0);
1614        let fold_exec = Arc::new(HistogramFoldExec {
1615            le_column_index: 1,
1616            field_column_index: 2,
1617            histogram_column_index: None,
1618            operation: HistogramFoldOperation::Quantile(0.4.into()),
1619            ts_column_index: 0,
1620            input: memory_exec,
1621            output_schema,
1622            tag_columns,
1623            partition_exprs,
1624            metric: ExecutionPlanMetricsSet::new(),
1625            properties,
1626        });
1627
1628        let session_context = SessionContext::default();
1629        let result = datafusion::physical_plan::collect(fold_exec, session_context.task_ctx())
1630            .await
1631            .unwrap();
1632        let result_literal = datatypes::arrow::util::pretty::pretty_format_batches(&result)
1633            .unwrap()
1634            .to_string();
1635
1636        let expected = String::from(
1637            "+--------+-------------------+
1638| host   | val               |
1639+--------+-------------------+
1640| host_1 | 257.5             |
1641| host_1 | 5.05              |
1642| host_1 | 0.0004            |
1643| host_2 | NaN               |
1644| host_2 | 6.040000000000001 |
1645+--------+-------------------+",
1646        );
1647        assert_eq!(result_literal, expected);
1648    }
1649
1650    #[tokio::test]
1651    async fn fold_dictionary_encoded_labels() {
1652        let dictionary_type =
1653            DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8));
1654        let schema = Arc::new(Schema::new(vec![
1655            Field::new("host", dictionary_type.clone(), true),
1656            Field::new("le", dictionary_type, true),
1657            Field::new("val", DataType::Float64, true),
1658        ]));
1659
1660        let mut host = StringDictionaryBuilder::<UInt32Type>::new();
1661        let mut le = StringDictionaryBuilder::<UInt32Type>::new();
1662        for value in ["host_1", "host_1", "host_1"] {
1663            host.append_value(value);
1664        }
1665        for value in ["0.1", "1", "+Inf"] {
1666            le.append_value(value);
1667        }
1668        let batch = RecordBatch::try_new(
1669            schema.clone(),
1670            vec![
1671                Arc::new(host.finish()),
1672                Arc::new(le.finish()),
1673                Arc::new(Float64Array::from(vec![1.0, 2.0, 3.0])),
1674            ],
1675        )
1676        .unwrap();
1677
1678        let fold_exec = build_fold_exec_from_batches(vec![batch], schema, 0.5, 0);
1679        let result =
1680            datafusion::physical_plan::collect(fold_exec, SessionContext::default().task_ctx())
1681                .await
1682                .unwrap();
1683
1684        assert_eq!(result.len(), 1);
1685        assert_eq!(result[0].num_rows(), 1);
1686        let host = result[0]
1687            .column(0)
1688            .as_any()
1689            .downcast_ref::<DictionaryArray<UInt32Type>>()
1690            .unwrap();
1691        assert_eq!(host.values().len(), 1);
1692        assert_eq!(
1693            string_array_value_at_index(result[0].column(0), 0),
1694            Some("host_1")
1695        );
1696        let value = result[0].column(1).as_primitive::<Float64Type>().value(0);
1697        assert!((value - 0.55).abs() < 1e-12);
1698    }
1699
1700    #[tokio::test]
1701    async fn pruning_should_keep_le_column_for_exec() {
1702        let schema = Arc::new(Schema::new(vec![
1703            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
1704            Field::new("le", DataType::Utf8, true),
1705            Field::new("val", DataType::Float64, true),
1706        ]));
1707        let df_schema = schema.clone().to_dfschema_ref().unwrap();
1708        let input = LogicalPlan::EmptyRelation(EmptyRelation {
1709            produce_one_row: false,
1710            schema: df_schema,
1711        });
1712        let plan = HistogramFold::new(
1713            "le".to_string(),
1714            "val".to_string(),
1715            "ts".to_string(),
1716            0.5,
1717            input,
1718        )
1719        .unwrap();
1720
1721        let output_columns = [0usize, 1usize];
1722        let required = plan.necessary_children_exprs(&output_columns).unwrap();
1723        let required = &required[0];
1724        assert_eq!(required.as_slice(), &[0, 1, 2]);
1725
1726        let input_batch = RecordBatch::try_new(
1727            schema,
1728            vec![
1729                Arc::new(TimestampMillisecondArray::from(vec![0, 0])),
1730                Arc::new(StringArray::from(vec!["0.1", "+Inf"])),
1731                Arc::new(Float64Array::from(vec![1.0, 2.0])),
1732            ],
1733        )
1734        .unwrap();
1735        let projected = project_batch(&input_batch, required);
1736        let projected_schema = projected.schema();
1737        let memory_exec = Arc::new(DataSourceExec::new(Arc::new(
1738            MemorySourceConfig::try_new(&[vec![projected]], projected_schema, None).unwrap(),
1739        )));
1740
1741        let fold_exec = plan.to_execution_plan(memory_exec);
1742        let session_context = SessionContext::default();
1743        let output_batches =
1744            datafusion::physical_plan::collect(fold_exec, session_context.task_ctx())
1745                .await
1746                .unwrap();
1747        assert_eq!(output_batches.len(), 1);
1748
1749        let output_batch = &output_batches[0];
1750        assert_eq!(output_batch.num_rows(), 1);
1751
1752        let ts = output_batch
1753            .column(0)
1754            .as_any()
1755            .downcast_ref::<TimestampMillisecondArray>()
1756            .unwrap();
1757        assert_eq!(ts.values(), &[0i64]);
1758
1759        let values = output_batch
1760            .column(1)
1761            .as_any()
1762            .downcast_ref::<Float64Array>()
1763            .unwrap();
1764        assert!((values.value(0) - 0.1).abs() < 1e-12);
1765
1766        // Simulate the pre-fix pruning behavior: omit the `le` column from the child input.
1767        let le_index = 1usize;
1768        let broken_required = output_columns
1769            .iter()
1770            .map(|&output_column| {
1771                if output_column < le_index {
1772                    output_column
1773                } else {
1774                    output_column + 1
1775                }
1776            })
1777            .collect::<Vec<_>>();
1778
1779        let broken = project_batch(&input_batch, &broken_required);
1780        let broken_schema = broken.schema();
1781        let broken_exec = Arc::new(DataSourceExec::new(Arc::new(
1782            MemorySourceConfig::try_new(&[vec![broken]], broken_schema, None).unwrap(),
1783        )));
1784        let broken_fold_exec = plan.to_execution_plan(broken_exec);
1785        let session_context = SessionContext::default();
1786        let broken_result = std::panic::AssertUnwindSafe(async {
1787            datafusion::physical_plan::collect(broken_fold_exec, session_context.task_ctx()).await
1788        })
1789        .catch_unwind()
1790        .await;
1791        assert!(broken_result.is_err());
1792    }
1793
1794    #[test]
1795    fn confirm_schema() {
1796        let input_schema = Schema::new(vec![
1797            Field::new("host", DataType::Utf8, true),
1798            Field::new("le", DataType::Utf8, true),
1799            Field::new("val", DataType::Float64, true),
1800        ])
1801        .to_dfschema_ref()
1802        .unwrap();
1803        let expected_output_schema = Schema::new(vec![
1804            Field::new("host", DataType::Utf8, true),
1805            Field::new("val", DataType::Float64, true),
1806        ])
1807        .to_dfschema_ref()
1808        .unwrap();
1809
1810        let actual = HistogramFold::convert_schema(&input_schema, "le").unwrap();
1811        assert_eq!(actual, expected_output_schema)
1812    }
1813
1814    #[tokio::test]
1815    async fn fallback_to_safe_mode_on_missing_inf() {
1816        let schema = Arc::new(Schema::new(vec![
1817            Field::new("host", DataType::Utf8, true),
1818            Field::new("le", DataType::Utf8, true),
1819            Field::new("val", DataType::Float64, true),
1820        ]));
1821        let host_column = Arc::new(StringArray::from(vec!["a", "a", "a", "a", "b", "b"])) as _;
1822        let le_column = Arc::new(StringArray::from(vec![
1823            "0.1", "+Inf", "0.1", "1.0", "0.1", "+Inf",
1824        ])) as _;
1825        let val_column = Arc::new(Float64Array::from(vec![1.0, 2.0, 3.0, 3.0, 1.0, 5.0])) as _;
1826        let batch =
1827            RecordBatch::try_new(schema.clone(), vec![host_column, le_column, val_column]).unwrap();
1828        let fold_exec = build_fold_exec_from_batches(vec![batch], schema, 0.5, 0);
1829        let session_context = SessionContext::default();
1830        let result = datafusion::physical_plan::collect(fold_exec, session_context.task_ctx())
1831            .await
1832            .unwrap();
1833        let result_literal = datatypes::arrow::util::pretty::pretty_format_batches(&result)
1834            .unwrap()
1835            .to_string();
1836
1837        let expected = String::from(
1838            "+------+-----+
1839| host | val |
1840+------+-----+
1841| a    | 0.1 |
1842| a    | NaN |
1843| b    | 0.1 |
1844+------+-----+",
1845        );
1846        assert_eq!(result_literal, expected);
1847    }
1848
1849    #[tokio::test]
1850    async fn emit_nan_when_no_inf_present() {
1851        let schema = Arc::new(Schema::new(vec![
1852            Field::new("host", DataType::Utf8, true),
1853            Field::new("le", DataType::Utf8, true),
1854            Field::new("val", DataType::Float64, true),
1855        ]));
1856        let host_column = Arc::new(StringArray::from(vec!["c", "c"])) as _;
1857        let le_column = Arc::new(StringArray::from(vec!["0.1", "1.0"])) as _;
1858        let val_column = Arc::new(Float64Array::from(vec![1.0, 2.0])) as _;
1859        let batch =
1860            RecordBatch::try_new(schema.clone(), vec![host_column, le_column, val_column]).unwrap();
1861        let fold_exec = build_fold_exec_from_batches(vec![batch], schema, 0.9, 0);
1862        let session_context = SessionContext::default();
1863        let result = datafusion::physical_plan::collect(fold_exec, session_context.task_ctx())
1864            .await
1865            .unwrap();
1866        let result_literal = datatypes::arrow::util::pretty::pretty_format_batches(&result)
1867            .unwrap()
1868            .to_string();
1869
1870        let expected = String::from(
1871            "+------+-----+
1872| host | val |
1873+------+-----+
1874| c    | NaN |
1875+------+-----+",
1876        );
1877        assert_eq!(result_literal, expected);
1878    }
1879
1880    #[tokio::test]
1881    async fn ignore_unparsable_bucket_bounds() {
1882        let schema = Arc::new(Schema::new(vec![
1883            Field::new("host", DataType::Utf8, true),
1884            Field::new("le", DataType::Utf8, true),
1885            Field::new("val", DataType::Float64, true),
1886        ]));
1887        let batch = RecordBatch::try_new(
1888            schema.clone(),
1889            vec![
1890                Arc::new(StringArray::from(vec!["a", "a", "a", "b"])),
1891                Arc::new(StringArray::from(vec![
1892                    Some("bad"),
1893                    Some("1"),
1894                    Some("+Inf"),
1895                    None,
1896                ])),
1897                Arc::new(Float64Array::from(vec![1.0, 2.0, 4.0, 1.0])),
1898            ],
1899        )
1900        .unwrap();
1901        let fold_exec = build_fold_exec_from_batches(vec![batch], schema, 0.5, 0);
1902
1903        let result =
1904            datafusion::physical_plan::collect(fold_exec, SessionContext::default().task_ctx())
1905                .await
1906                .unwrap();
1907
1908        assert_eq!(result.len(), 1);
1909        assert_eq!(result[0].num_rows(), 1);
1910        assert_eq!(
1911            string_array_value_at_index(result[0].column(0), 0),
1912            Some("a")
1913        );
1914        assert_eq!(
1915            result[0].column(1).as_primitive::<Float64Type>().value(0),
1916            1.0
1917        );
1918    }
1919
1920    #[tokio::test]
1921    async fn safe_mode_handles_misaligned_groups() {
1922        let schema = Arc::new(Schema::new(vec![
1923            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
1924            Field::new("le", DataType::Utf8, true),
1925            Field::new("val", DataType::Float64, true),
1926        ]));
1927
1928        let ts_column = Arc::new(TimestampMillisecondArray::from(vec![
1929            2900000, 2900000, 2900000, 3000000, 3000000, 3000000, 3000000, 3005000, 3005000,
1930            3010000, 3010000, 3010000, 3010000, 3010000,
1931        ])) as _;
1932        let le_column = Arc::new(StringArray::from(vec![
1933            "0.1", "1", "5", "0.1", "1", "5", "+Inf", "0.1", "+Inf", "0.1", "1", "3", "5", "+Inf",
1934        ])) as _;
1935        let val_column = Arc::new(Float64Array::from(vec![
1936            0.0, 0.0, 0.0, 50.0, 70.0, 110.0, 120.0, 10.0, 30.0, 10.0, 20.0, 30.0, 40.0, 50.0,
1937        ])) as _;
1938        let batch =
1939            RecordBatch::try_new(schema.clone(), vec![ts_column, le_column, val_column]).unwrap();
1940        let fold_exec = build_fold_exec_from_batches(vec![batch], schema, 0.5, 0);
1941        let session_context = SessionContext::default();
1942        let result = datafusion::physical_plan::collect(fold_exec, session_context.task_ctx())
1943            .await
1944            .unwrap();
1945
1946        let mut values = Vec::new();
1947        for batch in result {
1948            let array = batch.column(1).as_primitive::<Float64Type>();
1949            values.extend(array.iter().map(|v| v.unwrap()));
1950        }
1951
1952        assert_eq!(values.len(), 4);
1953        assert!(values[0].is_nan());
1954        assert!((values[1] - 0.55).abs() < 1e-10);
1955        assert!((values[2] - 0.1).abs() < 1e-10);
1956        assert!((values[3] - 2.0).abs() < 1e-10);
1957    }
1958
1959    #[tokio::test]
1960    async fn missing_buckets_at_first_timestamp() {
1961        let schema = Arc::new(Schema::new(vec![
1962            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
1963            Field::new("le", DataType::Utf8, true),
1964            Field::new("val", DataType::Float64, true),
1965        ]));
1966
1967        let ts_column = Arc::new(TimestampMillisecondArray::from(vec![
1968            2_900_000, 3_000_000, 3_000_000, 3_000_000, 3_000_000, 3_005_000, 3_005_000, 3_010_000,
1969            3_010_000, 3_010_000, 3_010_000, 3_010_000,
1970        ])) as _;
1971        let le_column = Arc::new(StringArray::from(vec![
1972            "0.1", "0.1", "1", "5", "+Inf", "0.1", "+Inf", "0.1", "1", "3", "5", "+Inf",
1973        ])) as _;
1974        let val_column = Arc::new(Float64Array::from(vec![
1975            0.0, 50.0, 70.0, 110.0, 120.0, 10.0, 30.0, 10.0, 20.0, 30.0, 40.0, 50.0,
1976        ])) as _;
1977
1978        let batch =
1979            RecordBatch::try_new(schema.clone(), vec![ts_column, le_column, val_column]).unwrap();
1980        let fold_exec = build_fold_exec_from_batches(vec![batch], schema, 0.5, 0);
1981        let session_context = SessionContext::default();
1982        let result = datafusion::physical_plan::collect(fold_exec, session_context.task_ctx())
1983            .await
1984            .unwrap();
1985
1986        let mut values = Vec::new();
1987        for batch in result {
1988            let array = batch.column(1).as_primitive::<Float64Type>();
1989            values.extend(array.iter().map(|v| v.unwrap()));
1990        }
1991
1992        assert_eq!(values.len(), 4);
1993        assert!(values[0].is_nan());
1994        assert!((values[1] - 0.55).abs() < 1e-10);
1995        assert!((values[2] - 0.1).abs() < 1e-10);
1996        assert!((values[3] - 2.0).abs() < 1e-10);
1997    }
1998
1999    #[tokio::test]
2000    async fn missing_inf_in_first_group() {
2001        let schema = Arc::new(Schema::new(vec![
2002            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
2003            Field::new("le", DataType::Utf8, true),
2004            Field::new("val", DataType::Float64, true),
2005        ]));
2006
2007        let ts_column = Arc::new(TimestampMillisecondArray::from(vec![
2008            1000, 1000, 1000, 2000, 2000, 2000, 2000,
2009        ])) as _;
2010        let le_column = Arc::new(StringArray::from(vec![
2011            "0.1", "1", "5", "0.1", "1", "5", "+Inf",
2012        ])) as _;
2013        let val_column = Arc::new(Float64Array::from(vec![
2014            0.0, 0.0, 0.0, 10.0, 20.0, 30.0, 30.0,
2015        ])) as _;
2016        let batch =
2017            RecordBatch::try_new(schema.clone(), vec![ts_column, le_column, val_column]).unwrap();
2018        let fold_exec = build_fold_exec_from_batches(vec![batch], schema, 0.5, 0);
2019        let session_context = SessionContext::default();
2020        let result = datafusion::physical_plan::collect(fold_exec, session_context.task_ctx())
2021            .await
2022            .unwrap();
2023
2024        let mut values = Vec::new();
2025        for batch in result {
2026            let array = batch.column(1).as_primitive::<Float64Type>();
2027            values.extend(array.iter().map(|v| v.unwrap()));
2028        }
2029
2030        assert_eq!(values.len(), 2);
2031        assert!(values[0].is_nan());
2032        assert!((values[1] - 0.55).abs() < 1e-10, "{values:?}");
2033    }
2034
2035    #[test]
2036    fn evaluate_row_normal_case() {
2037        let bucket = [0.0, 1.0, 2.0, 3.0, 4.0, f64::INFINITY];
2038
2039        #[derive(Debug)]
2040        struct Case {
2041            quantile: f64,
2042            counters: Vec<f64>,
2043            expected: f64,
2044        }
2045
2046        let cases = [
2047            Case {
2048                quantile: 0.9,
2049                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
2050                expected: 4.0,
2051            },
2052            Case {
2053                quantile: 0.89,
2054                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
2055                expected: 4.0,
2056            },
2057            Case {
2058                quantile: 0.78,
2059                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
2060                expected: 3.9,
2061            },
2062            Case {
2063                quantile: 0.5,
2064                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
2065                expected: 2.5,
2066            },
2067            Case {
2068                quantile: 0.5,
2069                counters: vec![0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
2070                expected: f64::NAN,
2071            },
2072            Case {
2073                quantile: 1.0,
2074                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
2075                expected: 4.0,
2076            },
2077            Case {
2078                quantile: 0.0,
2079                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
2080                expected: f64::NAN,
2081            },
2082            Case {
2083                quantile: 1.1,
2084                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
2085                expected: f64::INFINITY,
2086            },
2087            Case {
2088                quantile: -1.0,
2089                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
2090                expected: f64::NEG_INFINITY,
2091            },
2092        ];
2093
2094        for case in cases {
2095            let actual = HistogramFoldStream::evaluate_row(
2096                HistogramFoldOperation::Quantile(case.quantile.into()),
2097                &bucket,
2098                &case.counters,
2099            )
2100            .unwrap();
2101            assert_eq!(
2102                format!("{actual}"),
2103                format!("{}", case.expected),
2104                "{:?}",
2105                case
2106            );
2107        }
2108    }
2109
2110    #[test]
2111    fn evaluate_out_of_order_input() {
2112        let bucket = [0.0, 1.0, 2.0, 3.0, 4.0, f64::INFINITY];
2113        let counters = [5.0, 4.0, 3.0, 2.0, 1.0, 0.0];
2114        let result = HistogramFoldStream::evaluate_row(
2115            HistogramFoldOperation::Quantile(0.5.into()),
2116            &bucket,
2117            &counters,
2118        )
2119        .unwrap();
2120        assert_eq!(0.0, result);
2121    }
2122
2123    #[test]
2124    fn evaluate_wrong_bucket() {
2125        let bucket = [0.0, 1.0, 2.0, 3.0, 4.0, f64::INFINITY, 5.0];
2126        let counters = [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0];
2127        let result = HistogramFoldStream::evaluate_row(
2128            HistogramFoldOperation::Quantile(0.5.into()),
2129            &bucket,
2130            &counters,
2131        );
2132        assert!(result.is_err());
2133    }
2134
2135    #[test]
2136    fn evaluate_small_fraction() {
2137        let bucket = [0.0, 2.0, 4.0, 6.0, f64::INFINITY];
2138        let counters = [0.0, 1.0 / 300.0, 2.0 / 300.0, 0.01, 0.01];
2139        let result = HistogramFoldStream::evaluate_row(
2140            HistogramFoldOperation::Quantile(0.5.into()),
2141            &bucket,
2142            &counters,
2143        )
2144        .unwrap();
2145        assert_eq!(3.0, result);
2146    }
2147
2148    #[test]
2149    fn evaluate_non_monotonic_counter() {
2150        let bucket = [0.0, 1.0, 2.0, 3.0, f64::INFINITY];
2151        let counters = [0.1, 0.2, 0.4, 0.17, 0.5];
2152        let result = HistogramFoldStream::evaluate_row(
2153            HistogramFoldOperation::Quantile(0.5.into()),
2154            &bucket,
2155            &counters,
2156        )
2157        .unwrap();
2158        assert!((result - 1.25).abs() < 1e-10, "{result}");
2159    }
2160
2161    #[test]
2162    fn evaluate_nan_counter() {
2163        let bucket = [0.0, 1.0, 2.0, 3.0, f64::INFINITY];
2164        let counters = [f64::NAN, 1.0, 2.0, 3.0, 3.0];
2165        let result = HistogramFoldStream::evaluate_row(
2166            HistogramFoldOperation::Quantile(0.5.into()),
2167            &bucket,
2168            &counters,
2169        )
2170        .unwrap();
2171        assert!((result - 1.5).abs() < 1e-10, "{result}");
2172    }
2173
2174    #[test]
2175    fn evaluate_classic_histogram_fraction() {
2176        let buckets = [1.0, 2.0, f64::INFINITY];
2177        let counters = [2.0, 4.0, 4.0];
2178        let fraction = |lower, upper| {
2179            HistogramFoldStream::evaluate_row(
2180                HistogramFoldOperation::Fraction {
2181                    lower: OrderedF64::from(lower),
2182                    upper: OrderedF64::from(upper),
2183                },
2184                &buckets,
2185                &counters,
2186            )
2187            .unwrap()
2188        };
2189
2190        assert_eq!(fraction(0.0, 1.0), 0.5);
2191        assert_eq!(fraction(f64::NEG_INFINITY, f64::INFINITY), 1.0);
2192        assert_eq!(fraction(2.0, 1.0), 0.0);
2193
2194        assert_eq!(
2195            HistogramFoldStream::evaluate_row(
2196                HistogramFoldOperation::Fraction {
2197                    lower: 0.0.into(),
2198                    upper: 1.0.into(),
2199                },
2200                &[1.0, 1.0, f64::INFINITY],
2201                &[1.0, 2.0, 4.0],
2202            )
2203            .unwrap(),
2204            0.75
2205        );
2206
2207        assert_eq!(
2208            HistogramFoldStream::evaluate_row(
2209                HistogramFoldOperation::Fraction {
2210                    lower: f64::NEG_INFINITY.into(),
2211                    upper: f64::INFINITY.into(),
2212                },
2213                &[f64::INFINITY],
2214                &[4.0],
2215            )
2216            .unwrap(),
2217            1.0
2218        );
2219    }
2220
2221    #[tokio::test]
2222    async fn fraction_handles_single_inf_bucket_after_safe_fallback() {
2223        let schema = Arc::new(Schema::new(vec![
2224            Field::new("host", DataType::Utf8, false),
2225            Field::new("le", DataType::Utf8, false),
2226            Field::new("val", DataType::Float64, false),
2227        ]));
2228        let batch = RecordBatch::try_new(
2229            schema.clone(),
2230            vec![
2231                Arc::new(StringArray::from(vec!["a", "a", "b"])),
2232                Arc::new(StringArray::from(vec!["1", "+Inf", "+Inf"])),
2233                Arc::new(Float64Array::from(vec![2.0, 4.0, 4.0])),
2234            ],
2235        )
2236        .unwrap();
2237        let fold = build_fold_exec_from_batches_with_operation(
2238            vec![batch],
2239            schema,
2240            HistogramFoldOperation::Fraction {
2241                lower: f64::NEG_INFINITY.into(),
2242                upper: f64::INFINITY.into(),
2243            },
2244            0,
2245        );
2246
2247        let batches =
2248            datafusion::physical_plan::collect(fold, SessionContext::default().task_ctx())
2249                .await
2250                .unwrap();
2251        let values = batches[0].column(1).as_primitive::<Float64Type>();
2252        assert_eq!(values.values(), &[1.0, 1.0]);
2253    }
2254
2255    fn build_empty_relation(schema: &Arc<Schema>) -> LogicalPlan {
2256        LogicalPlan::EmptyRelation(EmptyRelation {
2257            produce_one_row: false,
2258            schema: schema.clone().to_dfschema_ref().unwrap(),
2259        })
2260    }
2261
2262    #[tokio::test]
2263    async fn encode_decode_histogram_fold() {
2264        let schema = Arc::new(Schema::new(vec![
2265            Field::new("ts", DataType::Int64, false),
2266            Field::new("le", DataType::Utf8, false),
2267            Field::new("val", DataType::Float64, false),
2268        ]));
2269        let input_plan = build_empty_relation(&schema);
2270        let plan_node = HistogramFold::new(
2271            "le".to_string(),
2272            "val".to_string(),
2273            "ts".to_string(),
2274            0.8,
2275            input_plan.clone(),
2276        )
2277        .unwrap();
2278        let fraction_node = HistogramFold::new_with_operation(
2279            "le".to_string(),
2280            "val".to_string(),
2281            "ts".to_string(),
2282            HistogramFoldOperation::Fraction {
2283                lower: 0.0.into(),
2284                upper: 1.0.into(),
2285            },
2286            None,
2287            input_plan.clone(),
2288        )
2289        .unwrap();
2290        assert!(fraction_node.serialize().is_err());
2291
2292        let bytes = plan_node.serialize().unwrap();
2293
2294        let histogram_fold = HistogramFold::deserialize(&bytes).unwrap();
2295        // need fix
2296        let histogram_fold = histogram_fold
2297            .with_exprs_and_inputs(vec![], vec![input_plan])
2298            .unwrap();
2299
2300        assert_eq!(histogram_fold.le_column, "le");
2301        assert_eq!(histogram_fold.ts_column, "ts");
2302        assert_eq!(histogram_fold.field_column, "val");
2303        assert_eq!(
2304            histogram_fold.operation,
2305            HistogramFoldOperation::Quantile(OrderedF64::from(0.8))
2306        );
2307        assert_eq!(histogram_fold.output_schema.fields().len(), 2);
2308        assert_eq!(histogram_fold.output_schema.field(0).name(), "ts");
2309        assert_eq!(histogram_fold.output_schema.field(1).name(), "val");
2310    }
2311}