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