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::{CastExpr as PhyCast, Column as PhyColumn};
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    quantile: OrderedF64,
82    output_schema: DFSchemaRef,
83    unfix: Option<UnfixIndices>,
84}
85
86#[derive(Debug, PartialEq, Eq, Hash, PartialOrd)]
87struct UnfixIndices {
88    pub le_column_idx: u64,
89    pub ts_column_idx: u64,
90    pub field_column_idx: u64,
91}
92
93impl UserDefinedLogicalNodeCore for HistogramFold {
94    fn name(&self) -> &str {
95        Self::name()
96    }
97
98    fn inputs(&self) -> Vec<&LogicalPlan> {
99        vec![&self.input]
100    }
101
102    fn schema(&self) -> &DFSchemaRef {
103        &self.output_schema
104    }
105
106    fn expressions(&self) -> Vec<Expr> {
107        if self.unfix.is_some() {
108            return vec![];
109        }
110
111        let mut exprs = vec![
112            col(&self.le_column),
113            col(&self.ts_column),
114            col(&self.field_column),
115        ];
116        exprs.extend(self.input.schema().fields().iter().filter_map(|f| {
117            let name = f.name();
118            if name != &self.le_column && name != &self.ts_column && name != &self.field_column {
119                Some(col(name))
120            } else {
121                None
122            }
123        }));
124        exprs
125    }
126
127    fn necessary_children_exprs(&self, output_columns: &[usize]) -> Option<Vec<Vec<usize>>> {
128        if self.unfix.is_some() {
129            return None;
130        }
131
132        let input_schema = self.input.schema();
133        let le_column_index = input_schema.index_of_column_by_name(None, &self.le_column)?;
134
135        if output_columns.is_empty() {
136            let indices = (0..input_schema.fields().len()).collect::<Vec<_>>();
137            return Some(vec![indices]);
138        }
139
140        let mut necessary_indices = output_columns
141            .iter()
142            .map(|&output_column| {
143                if output_column < le_column_index {
144                    output_column
145                } else {
146                    output_column + 1
147                }
148            })
149            .collect::<Vec<_>>();
150        necessary_indices.push(le_column_index);
151        necessary_indices.sort_unstable();
152        necessary_indices.dedup();
153        Some(vec![necessary_indices])
154    }
155
156    fn fmt_for_explain(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
157        write!(
158            f,
159            "HistogramFold: le={}, field={}, quantile={}",
160            self.le_column, self.field_column, self.quantile
161        )
162    }
163
164    fn with_exprs_and_inputs(
165        &self,
166        _exprs: Vec<Expr>,
167        inputs: Vec<LogicalPlan>,
168    ) -> DataFusionResult<Self> {
169        if inputs.is_empty() {
170            return Err(DataFusionError::Internal(
171                "HistogramFold must have at least one input".to_string(),
172            ));
173        }
174
175        let input: LogicalPlan = inputs.into_iter().next().unwrap();
176        let input_schema = input.schema();
177
178        if let Some(unfix) = &self.unfix {
179            let le_column =
180                resolve_column_name(unfix.le_column_idx, input_schema, "HistogramFold", "le")?;
181            let ts_column =
182                resolve_column_name(unfix.ts_column_idx, input_schema, "HistogramFold", "ts")?;
183            let field_column = resolve_column_name(
184                unfix.field_column_idx,
185                input_schema,
186                "HistogramFold",
187                "field",
188            )?;
189
190            let output_schema = Self::convert_schema(input_schema, &le_column)?;
191
192            Ok(Self {
193                le_column,
194                ts_column,
195                input,
196                field_column,
197                quantile: self.quantile,
198                output_schema,
199                unfix: None,
200            })
201        } else {
202            Ok(Self {
203                le_column: self.le_column.clone(),
204                ts_column: self.ts_column.clone(),
205                input,
206                field_column: self.field_column.clone(),
207                quantile: self.quantile,
208                output_schema: self.output_schema.clone(),
209                unfix: None,
210            })
211        }
212    }
213}
214
215impl HistogramFold {
216    pub fn new(
217        le_column: String,
218        field_column: String,
219        ts_column: String,
220        quantile: f64,
221        input: LogicalPlan,
222    ) -> DataFusionResult<Self> {
223        let input_schema = input.schema();
224        Self::check_schema(input_schema, &le_column, &field_column, &ts_column)?;
225        let output_schema = Self::convert_schema(input_schema, &le_column)?;
226        Ok(Self {
227            le_column,
228            ts_column,
229            input,
230            field_column,
231            quantile: quantile.into(),
232            output_schema,
233            unfix: None,
234        })
235    }
236
237    pub const fn name() -> &'static str {
238        "HistogramFold"
239    }
240
241    fn check_schema(
242        input_schema: &DFSchemaRef,
243        le_column: &str,
244        field_column: &str,
245        ts_column: &str,
246    ) -> DataFusionResult<()> {
247        let check_column = |col| {
248            if !input_schema.has_column_with_unqualified_name(col) {
249                Err(DataFusionError::SchemaError(
250                    Box::new(datafusion::common::SchemaError::FieldNotFound {
251                        field: Box::new(Column::new(None::<String>, col)),
252                        valid_fields: input_schema.columns(),
253                    }),
254                    Box::new(None),
255                ))
256            } else {
257                Ok(())
258            }
259        };
260
261        check_column(le_column)?;
262        check_column(ts_column)?;
263        check_column(field_column)
264    }
265
266    pub fn to_execution_plan(&self, exec_input: Arc<dyn ExecutionPlan>) -> Arc<dyn ExecutionPlan> {
267        let input_schema = self.input.schema();
268        // safety: those fields are checked in `check_schema()`
269        let le_column_index = input_schema
270            .index_of_column_by_name(None, &self.le_column)
271            .unwrap();
272        let field_column_index = input_schema
273            .index_of_column_by_name(None, &self.field_column)
274            .unwrap();
275        let ts_column_index = input_schema
276            .index_of_column_by_name(None, &self.ts_column)
277            .unwrap();
278
279        let tag_columns = exec_input
280            .schema()
281            .fields()
282            .iter()
283            .enumerate()
284            .filter_map(|(idx, field)| {
285                if idx == le_column_index || idx == field_column_index || idx == ts_column_index {
286                    None
287                } else {
288                    Some(Arc::new(PhyColumn::new(field.name(), idx)) as _)
289                }
290            })
291            .collect::<Vec<_>>();
292
293        let mut partition_exprs = tag_columns.clone();
294        partition_exprs.push(Arc::new(PhyColumn::new(
295            self.input.schema().field(ts_column_index).name(),
296            ts_column_index,
297        )) as _);
298
299        let output_schema: SchemaRef = self.output_schema.inner().clone();
300        let properties = Arc::new(PlanProperties::new(
301            EquivalenceProperties::new(output_schema.clone()),
302            Partitioning::Hash(
303                partition_exprs.clone(),
304                exec_input.output_partitioning().partition_count(),
305            ),
306            EmissionType::Incremental,
307            Boundedness::Bounded,
308        ));
309        Arc::new(HistogramFoldExec {
310            le_column_index,
311            field_column_index,
312            ts_column_index,
313            input: exec_input,
314            tag_columns,
315            partition_exprs,
316            quantile: self.quantile.into(),
317            output_schema,
318            metric: ExecutionPlanMetricsSet::new(),
319            properties,
320        })
321    }
322
323    /// Transform the schema
324    ///
325    /// - `le` will be removed
326    ///
327    /// Column qualifiers are preserved so downstream plan nodes can keep
328    /// referencing the columns by their original qualified names.
329    fn convert_schema(
330        input_schema: &DFSchemaRef,
331        le_column: &str,
332    ) -> DataFusionResult<DFSchemaRef> {
333        // safety: those fields are checked in `check_schema()`
334        let mut new_fields = Vec::with_capacity(input_schema.fields().len() - 1);
335        for (qualifier, field) in input_schema.iter() {
336            if field.name() != le_column {
337                new_fields.push((qualifier.cloned(), field.clone()));
338            }
339        }
340        Ok(Arc::new(DFSchema::new_with_metadata(
341            new_fields,
342            HashMap::new(),
343        )?))
344    }
345
346    pub fn serialize(&self) -> Vec<u8> {
347        let le_column_idx = serialize_column_index(self.input.schema(), &self.le_column);
348        let ts_column_idx = serialize_column_index(self.input.schema(), &self.ts_column);
349        let field_column_idx = serialize_column_index(self.input.schema(), &self.field_column);
350
351        pb::HistogramFold {
352            le_column_idx,
353            ts_column_idx,
354            field_column_idx,
355            quantile: self.quantile.into(),
356        }
357        .encode_to_vec()
358    }
359
360    pub fn deserialize(bytes: &[u8]) -> Result<Self> {
361        let pb_histogram_fold = pb::HistogramFold::decode(bytes).context(DeserializeSnafu)?;
362        let placeholder_plan = LogicalPlan::EmptyRelation(EmptyRelation {
363            produce_one_row: false,
364            schema: Arc::new(DFSchema::empty()),
365        });
366
367        let unfix = UnfixIndices {
368            le_column_idx: pb_histogram_fold.le_column_idx,
369            ts_column_idx: pb_histogram_fold.ts_column_idx,
370            field_column_idx: pb_histogram_fold.field_column_idx,
371        };
372
373        Ok(Self {
374            le_column: String::new(),
375            ts_column: String::new(),
376            input: placeholder_plan,
377            field_column: String::new(),
378            quantile: pb_histogram_fold.quantile.into(),
379            output_schema: Arc::new(DFSchema::empty()),
380            unfix: Some(unfix),
381        })
382    }
383}
384
385impl PartialOrd for HistogramFold {
386    fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
387        // Compare fields in order excluding output_schema
388        match self.le_column.partial_cmp(&other.le_column) {
389            Some(core::cmp::Ordering::Equal) => {}
390            ord => return ord,
391        }
392        match self.ts_column.partial_cmp(&other.ts_column) {
393            Some(core::cmp::Ordering::Equal) => {}
394            ord => return ord,
395        }
396        match self.input.partial_cmp(&other.input) {
397            Some(core::cmp::Ordering::Equal) => {}
398            ord => return ord,
399        }
400        match self.field_column.partial_cmp(&other.field_column) {
401            Some(core::cmp::Ordering::Equal) => {}
402            ord => return ord,
403        }
404        self.quantile.partial_cmp(&other.quantile)
405    }
406}
407
408#[derive(Debug)]
409pub struct HistogramFoldExec {
410    /// Index for `le` column in the schema of input.
411    le_column_index: usize,
412    input: Arc<dyn ExecutionPlan>,
413    output_schema: SchemaRef,
414    /// Index for field column in the schema of input.
415    field_column_index: usize,
416    ts_column_index: usize,
417    /// Tag columns are all columns except `le`, `field` and `ts` columns.
418    tag_columns: Vec<Arc<dyn PhysicalExpr>>,
419    partition_exprs: Vec<Arc<dyn PhysicalExpr>>,
420    quantile: f64,
421    metric: ExecutionPlanMetricsSet,
422    properties: Arc<PlanProperties>,
423}
424
425impl ExecutionPlan for HistogramFoldExec {
426    fn as_any(&self) -> &dyn Any {
427        self
428    }
429
430    fn properties(&self) -> &Arc<PlanProperties> {
431        &self.properties
432    }
433
434    fn required_input_ordering(&self) -> Vec<Option<OrderingRequirements>> {
435        let mut cols = self
436            .tag_columns
437            .iter()
438            .map(|expr| PhysicalSortRequirement {
439                expr: expr.clone(),
440                options: None,
441            })
442            .collect::<Vec<PhysicalSortRequirement>>();
443        // add ts
444        cols.push(PhysicalSortRequirement {
445            expr: Arc::new(PhyColumn::new(
446                self.input.schema().field(self.ts_column_index).name(),
447                self.ts_column_index,
448            )),
449            options: None,
450        });
451        // add le ASC
452        cols.push(PhysicalSortRequirement {
453            expr: Arc::new(PhyCast::new(
454                Arc::new(PhyColumn::new(
455                    self.input.schema().field(self.le_column_index).name(),
456                    self.le_column_index,
457                )),
458                DataType::Float64,
459                None,
460            )),
461            options: Some(SortOptions {
462                descending: false,  // +INF in the last
463                nulls_first: false, // not nullable
464            }),
465        });
466
467        // Safety: `cols` is not empty
468        let requirement = LexRequirement::new(cols).unwrap();
469
470        vec![Some(OrderingRequirements::Hard(vec![requirement]))]
471    }
472
473    fn required_input_distribution(&self) -> Vec<Distribution> {
474        vec![Distribution::HashPartitioned(self.partition_exprs.clone())]
475    }
476
477    fn maintains_input_order(&self) -> Vec<bool> {
478        vec![true; self.children().len()]
479    }
480
481    fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
482        vec![&self.input]
483    }
484
485    // cannot change schema with this method
486    fn with_new_children(
487        self: Arc<Self>,
488        children: Vec<Arc<dyn ExecutionPlan>>,
489    ) -> DataFusionResult<Arc<dyn ExecutionPlan>> {
490        assert!(!children.is_empty());
491        let new_input = children[0].clone();
492        let properties = Arc::new(PlanProperties::new(
493            EquivalenceProperties::new(self.output_schema.clone()),
494            Partitioning::Hash(
495                self.partition_exprs.clone(),
496                new_input.output_partitioning().partition_count(),
497            ),
498            EmissionType::Incremental,
499            Boundedness::Bounded,
500        ));
501        Ok(Arc::new(Self {
502            input: new_input,
503            metric: self.metric.clone(),
504            le_column_index: self.le_column_index,
505            ts_column_index: self.ts_column_index,
506            tag_columns: self.tag_columns.clone(),
507            partition_exprs: self.partition_exprs.clone(),
508            quantile: self.quantile,
509            output_schema: self.output_schema.clone(),
510            field_column_index: self.field_column_index,
511            properties,
512        }))
513    }
514
515    fn execute(
516        &self,
517        partition: usize,
518        context: Arc<TaskContext>,
519    ) -> DataFusionResult<SendableRecordBatchStream> {
520        let baseline_metric = BaselineMetrics::new(&self.metric, partition);
521
522        let batch_size = context.session_config().batch_size();
523        let input = self.input.execute(partition, context)?;
524        let output_schema = self.output_schema.clone();
525
526        let mut normal_indices = (0..input.schema().fields().len()).collect::<HashSet<_>>();
527        normal_indices.remove(&self.field_column_index);
528        normal_indices.remove(&self.le_column_index);
529        Ok(Box::pin(HistogramFoldStream {
530            le_column_index: self.le_column_index,
531            field_column_index: self.field_column_index,
532            quantile: self.quantile,
533            normal_indices: normal_indices.into_iter().collect(),
534            bucket_size: None,
535            input_buffer: vec![],
536            input,
537            output_schema,
538            input_schema: self.input.schema(),
539            mode: FoldMode::Optimistic,
540            safe_group: None,
541            metric: baseline_metric,
542            batch_size,
543            input_buffered_rows: 0,
544            output_buffer: HistogramFoldStream::empty_output_buffer(
545                &self.output_schema,
546                self.le_column_index,
547            )?,
548            output_buffered_rows: 0,
549        }))
550    }
551
552    fn metrics(&self) -> Option<MetricsSet> {
553        Some(self.metric.clone_inner())
554    }
555
556    fn partition_statistics(&self, _: Option<usize>) -> DataFusionResult<Statistics> {
557        Ok(Statistics {
558            num_rows: Precision::Absent,
559            total_byte_size: Precision::Absent,
560            column_statistics: Statistics::unknown_column(&self.schema()),
561        })
562    }
563
564    fn name(&self) -> &str {
565        "HistogramFoldExec"
566    }
567}
568
569impl DisplayAs for HistogramFoldExec {
570    fn fmt_as(&self, t: DisplayFormatType, f: &mut std::fmt::Formatter) -> std::fmt::Result {
571        match t {
572            DisplayFormatType::Default
573            | DisplayFormatType::Verbose
574            | DisplayFormatType::TreeRender => {
575                write!(
576                    f,
577                    "HistogramFoldExec: le=@{}, field=@{}, quantile={}",
578                    self.le_column_index, self.field_column_index, self.quantile
579                )
580            }
581        }
582    }
583}
584
585#[derive(Debug, Clone, Copy, PartialEq, Eq)]
586enum FoldMode {
587    Optimistic,
588    Safe,
589}
590
591pub struct HistogramFoldStream {
592    // internal states
593    le_column_index: usize,
594    field_column_index: usize,
595    quantile: f64,
596    /// Columns need not folding. This indices is based on input schema
597    normal_indices: Vec<usize>,
598    bucket_size: Option<usize>,
599    /// Expected output batch size
600    batch_size: usize,
601    output_schema: SchemaRef,
602    input_schema: SchemaRef,
603    mode: FoldMode,
604    safe_group: Option<SafeGroup>,
605
606    // buffers
607    input_buffer: Vec<RecordBatch>,
608    input_buffered_rows: usize,
609    output_buffer: Vec<Box<dyn MutableVector>>,
610    output_buffered_rows: usize,
611
612    // runtime things
613    input: SendableRecordBatchStream,
614    metric: BaselineMetrics,
615}
616
617#[derive(Debug, Default)]
618struct SafeGroup {
619    tag_values: Vec<Value>,
620    buckets: Vec<f64>,
621    counters: Vec<f64>,
622}
623
624impl RecordBatchStream for HistogramFoldStream {
625    fn schema(&self) -> SchemaRef {
626        self.output_schema.clone()
627    }
628}
629
630impl Stream for HistogramFoldStream {
631    type Item = DataFusionResult<RecordBatch>;
632
633    fn poll_next(
634        mut self: std::pin::Pin<&mut Self>,
635        cx: &mut std::task::Context<'_>,
636    ) -> Poll<Option<Self::Item>> {
637        let poll = loop {
638            match ready!(self.input.poll_next_unpin(cx)) {
639                Some(batch) => {
640                    let batch = batch?;
641                    let timer = Instant::now();
642                    let Some(result) = self.fold_input(batch)? else {
643                        self.metric.elapsed_compute().add_elapsed(timer);
644                        continue;
645                    };
646                    self.metric.elapsed_compute().add_elapsed(timer);
647                    break Poll::Ready(Some(result));
648                }
649                None => {
650                    self.flush_remaining()?;
651                    break Poll::Ready(self.take_output_buf()?.map(Ok));
652                }
653            }
654        };
655        self.metric.record_poll(poll)
656    }
657}
658
659impl HistogramFoldStream {
660    /// The inner most `Result` is for `poll_next()`
661    pub fn fold_input(
662        &mut self,
663        input: RecordBatch,
664    ) -> DataFusionResult<Option<DataFusionResult<RecordBatch>>> {
665        match self.mode {
666            FoldMode::Safe => {
667                self.push_input_buf(input);
668                self.process_safe_mode_buffer()?;
669            }
670            FoldMode::Optimistic => {
671                self.push_input_buf(input);
672                let Some(bucket_num) = self.calculate_bucket_num_from_buffer()? else {
673                    return Ok(None);
674                };
675                self.bucket_size = Some(bucket_num);
676
677                if self.input_buffered_rows < bucket_num {
678                    // not enough rows to fold
679                    return Ok(None);
680                }
681
682                self.fold_buf(bucket_num)?;
683            }
684        }
685
686        self.maybe_take_output()
687    }
688
689    /// Generate a group of empty [MutableVector]s from the output schema.
690    ///
691    /// For simplicity, this method will insert a placeholder for `le`. So that
692    /// the output buffers has the same schema with input. This placeholder needs
693    /// to be removed before returning the output batch.
694    pub fn empty_output_buffer(
695        schema: &SchemaRef,
696        le_column_index: usize,
697    ) -> DataFusionResult<Vec<Box<dyn MutableVector>>> {
698        let mut builders = Vec::with_capacity(schema.fields().len() + 1);
699        for field in schema.fields() {
700            let concrete_datatype = ConcreteDataType::try_from(field.data_type()).unwrap();
701            let mutable_vector = concrete_datatype.create_mutable_vector(0);
702            builders.push(mutable_vector);
703        }
704        builders.insert(
705            le_column_index,
706            ConcreteDataType::float64_datatype().create_mutable_vector(0),
707        );
708
709        Ok(builders)
710    }
711
712    /// Determines bucket count using buffered batches, concatenating them to
713    /// detect the first complete bucket that may span batch boundaries.
714    fn calculate_bucket_num_from_buffer(&mut self) -> DataFusionResult<Option<usize>> {
715        if let Some(size) = self.bucket_size {
716            return Ok(Some(size));
717        }
718
719        if self.input_buffer.is_empty() {
720            return Ok(None);
721        }
722
723        let batch_refs: Vec<&RecordBatch> = self.input_buffer.iter().collect();
724        let batch = concat_batches(&self.input_schema, batch_refs)?;
725        self.find_first_complete_bucket(&batch)
726    }
727
728    fn find_first_complete_bucket(&self, batch: &RecordBatch) -> DataFusionResult<Option<usize>> {
729        if batch.num_rows() == 0 {
730            return Ok(None);
731        }
732
733        let vectors = Helper::try_into_vectors(batch.columns())
734            .map_err(|e| DataFusionError::Execution(e.to_string()))?;
735        let le_array = batch.column(self.le_column_index);
736
737        let mut tag_values_buf = Vec::with_capacity(self.normal_indices.len());
738        self.collect_tag_values(&vectors, 0, &mut tag_values_buf);
739        let mut group_start = 0usize;
740
741        for row in 0..batch.num_rows() {
742            if !self.is_same_group(&vectors, row, &tag_values_buf) {
743                // new group begins
744                self.collect_tag_values(&vectors, row, &mut tag_values_buf);
745                group_start = row;
746            }
747
748            if Self::is_positive_infinity(le_array, row) {
749                return Ok(Some(row - group_start + 1));
750            }
751        }
752
753        Ok(None)
754    }
755
756    /// Fold record batches from input buffer and put to output buffer
757    fn fold_buf(&mut self, bucket_num: usize) -> DataFusionResult<()> {
758        let batch = concat_batches(&self.input_schema, self.input_buffer.drain(..).as_ref())?;
759        let mut remaining_rows = self.input_buffered_rows;
760        let mut cursor = 0;
761
762        // TODO(LFC): Try to get rid of the Arrow array to vector conversion here.
763        let vectors = Helper::try_into_vectors(batch.columns())
764            .map_err(|e| DataFusionError::Execution(e.to_string()))?;
765        let le_array = batch.column(self.le_column_index);
766        let field_array = batch.column(self.field_column_index);
767        let field_array = field_array.as_primitive::<Float64Type>();
768        let mut tag_values_buf = Vec::with_capacity(self.normal_indices.len());
769
770        while remaining_rows >= bucket_num && self.mode == FoldMode::Optimistic {
771            self.collect_tag_values(&vectors, cursor, &mut tag_values_buf);
772            if !self.validate_optimistic_group(
773                &vectors,
774                le_array,
775                cursor,
776                bucket_num,
777                &tag_values_buf,
778            ) {
779                let remaining_input_batch = batch.slice(cursor, remaining_rows);
780                self.switch_to_safe_mode(remaining_input_batch)?;
781                return Ok(());
782            }
783
784            // "sample" normal columns
785            for (idx, value) in self.normal_indices.iter().zip(tag_values_buf.iter()) {
786                self.output_buffer[*idx].push_value_ref(value);
787            }
788            // "fold" `le` and field columns
789            let mut bucket = Vec::with_capacity(bucket_num);
790            let mut counters = Vec::with_capacity(bucket_num);
791            for bias in 0..bucket_num {
792                let position = cursor + bias;
793                let le = string_array_value_at_index(le_array, position)
794                    .and_then(|value| value.parse::<f64>().ok())
795                    .unwrap_or(f64::NAN);
796                bucket.push(le);
797
798                let counter = if field_array.is_valid(position) {
799                    field_array.value(position)
800                } else {
801                    f64::NAN
802                };
803                counters.push(counter);
804            }
805            // ignore invalid data
806            let result = Self::evaluate_row(self.quantile, &bucket, &counters).unwrap_or(f64::NAN);
807            self.output_buffer[self.field_column_index].push_value_ref(&ValueRef::from(result));
808            cursor += bucket_num;
809            remaining_rows -= bucket_num;
810            self.output_buffered_rows += 1;
811        }
812
813        let remaining_input_batch = batch.slice(cursor, remaining_rows);
814        self.input_buffered_rows = remaining_input_batch.num_rows();
815        if self.input_buffered_rows > 0 {
816            self.input_buffer.push(remaining_input_batch);
817        }
818
819        Ok(())
820    }
821
822    fn push_input_buf(&mut self, batch: RecordBatch) {
823        self.input_buffered_rows += batch.num_rows();
824        self.input_buffer.push(batch);
825    }
826
827    fn maybe_take_output(&mut self) -> DataFusionResult<Option<DataFusionResult<RecordBatch>>> {
828        if self.output_buffered_rows >= self.batch_size {
829            return Ok(self.take_output_buf()?.map(Ok));
830        }
831        Ok(None)
832    }
833
834    fn switch_to_safe_mode(&mut self, remaining_batch: RecordBatch) -> DataFusionResult<()> {
835        self.mode = FoldMode::Safe;
836        self.bucket_size = None;
837        self.input_buffer.clear();
838        self.input_buffered_rows = remaining_batch.num_rows();
839
840        if self.input_buffered_rows > 0 {
841            self.input_buffer.push(remaining_batch);
842            self.process_safe_mode_buffer()?;
843        }
844
845        Ok(())
846    }
847
848    fn collect_tag_values<'a>(
849        &self,
850        vectors: &'a [VectorRef],
851        row: usize,
852        tag_values: &mut Vec<ValueRef<'a>>,
853    ) {
854        tag_values.clear();
855        for idx in self.normal_indices.iter() {
856            tag_values.push(vectors[*idx].get_ref(row));
857        }
858    }
859
860    fn validate_optimistic_group(
861        &self,
862        vectors: &[VectorRef],
863        le_array: &ArrayRef,
864        cursor: usize,
865        bucket_num: usize,
866        tag_values: &[ValueRef<'_>],
867    ) -> bool {
868        let inf_index = cursor + bucket_num - 1;
869        if !Self::is_positive_infinity(le_array, inf_index) {
870            return false;
871        }
872
873        for offset in 1..bucket_num {
874            let row = cursor + offset;
875            for (idx, expected) in self.normal_indices.iter().zip(tag_values.iter()) {
876                if vectors[*idx].get_ref(row) != *expected {
877                    return false;
878                }
879            }
880        }
881        true
882    }
883
884    /// Checks whether a row belongs to the current group (same series).
885    fn is_same_group(
886        &self,
887        vectors: &[VectorRef],
888        row: usize,
889        tag_values: &[ValueRef<'_>],
890    ) -> bool {
891        self.normal_indices
892            .iter()
893            .zip(tag_values.iter())
894            .all(|(idx, expected)| vectors[*idx].get_ref(row) == *expected)
895    }
896
897    fn push_output_row(&mut self, tag_values: &[ValueRef<'_>], result: f64) {
898        debug_assert_eq!(self.normal_indices.len(), tag_values.len());
899        for (idx, value) in self.normal_indices.iter().zip(tag_values.iter()) {
900            self.output_buffer[*idx].push_value_ref(value);
901        }
902        self.output_buffer[self.field_column_index].push_value_ref(&ValueRef::from(result));
903        self.output_buffered_rows += 1;
904    }
905
906    fn finalize_safe_group(&mut self) -> DataFusionResult<()> {
907        if let Some(group) = self.safe_group.take() {
908            if group.tag_values.is_empty() {
909                return Ok(());
910            }
911
912            let has_inf = group
913                .buckets
914                .last()
915                .map(|v| v.is_infinite() && v.is_sign_positive())
916                .unwrap_or(false);
917            let result = if group.buckets.len() < 2 || !has_inf {
918                f64::NAN
919            } else {
920                Self::evaluate_row(self.quantile, &group.buckets, &group.counters)
921                    .unwrap_or(f64::NAN)
922            };
923            let mut tag_value_refs = Vec::with_capacity(group.tag_values.len());
924            tag_value_refs.extend(group.tag_values.iter().map(|v| v.as_value_ref()));
925            self.push_output_row(&tag_value_refs, result);
926        }
927        Ok(())
928    }
929
930    fn process_safe_mode_buffer(&mut self) -> DataFusionResult<()> {
931        if self.input_buffer.is_empty() {
932            self.input_buffered_rows = 0;
933            return Ok(());
934        }
935
936        let batch = concat_batches(&self.input_schema, self.input_buffer.drain(..).as_ref())?;
937        self.input_buffered_rows = 0;
938        let vectors = Helper::try_into_vectors(batch.columns())
939            .map_err(|e| DataFusionError::Execution(e.to_string()))?;
940        let le_array = batch.column(self.le_column_index);
941        let field_array = batch
942            .column(self.field_column_index)
943            .as_primitive::<Float64Type>();
944        let mut tag_values_buf = Vec::with_capacity(self.normal_indices.len());
945
946        for row in 0..batch.num_rows() {
947            self.collect_tag_values(&vectors, row, &mut tag_values_buf);
948            let should_start_new_group = self
949                .safe_group
950                .as_ref()
951                .is_none_or(|group| !Self::tag_values_equal(&group.tag_values, &tag_values_buf));
952            if should_start_new_group {
953                self.finalize_safe_group()?;
954                self.safe_group = Some(SafeGroup {
955                    tag_values: tag_values_buf.iter().cloned().map(Value::from).collect(),
956                    buckets: Vec::new(),
957                    counters: Vec::new(),
958                });
959            }
960
961            let Some(group) = self.safe_group.as_mut() else {
962                continue;
963            };
964
965            let bucket = string_array_value_at_index(le_array, row)
966                .and_then(|value| value.parse::<f64>().ok())
967                .unwrap_or(f64::NAN);
968            let counter = if field_array.is_valid(row) {
969                field_array.value(row)
970            } else {
971                f64::NAN
972            };
973
974            group.buckets.push(bucket);
975            group.counters.push(counter);
976        }
977
978        Ok(())
979    }
980
981    fn tag_values_equal(group_values: &[Value], current: &[ValueRef<'_>]) -> bool {
982        group_values.len() == current.len()
983            && group_values
984                .iter()
985                .zip(current.iter())
986                .all(|(group, now)| group.as_value_ref() == *now)
987    }
988
989    /// Compute result from output buffer
990    fn take_output_buf(&mut self) -> DataFusionResult<Option<RecordBatch>> {
991        if self.output_buffered_rows == 0 {
992            if self.input_buffered_rows != 0 {
993                warn!(
994                    "input buffer is not empty, {} rows remaining",
995                    self.input_buffered_rows
996                );
997            }
998            return Ok(None);
999        }
1000
1001        let mut output_buf = Self::empty_output_buffer(&self.output_schema, self.le_column_index)?;
1002        std::mem::swap(&mut self.output_buffer, &mut output_buf);
1003        let mut columns = Vec::with_capacity(output_buf.len());
1004        for builder in output_buf.iter_mut() {
1005            columns.push(builder.to_vector().to_arrow_array());
1006        }
1007        // remove the placeholder column for `le`
1008        columns.remove(self.le_column_index);
1009
1010        self.output_buffered_rows = 0;
1011        RecordBatch::try_new(self.output_schema.clone(), columns)
1012            .map(Some)
1013            .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))
1014    }
1015
1016    fn flush_remaining(&mut self) -> DataFusionResult<()> {
1017        if self.mode == FoldMode::Optimistic && self.input_buffered_rows > 0 {
1018            let buffered_batches: Vec<_> = self.input_buffer.drain(..).collect();
1019            if !buffered_batches.is_empty() {
1020                let batch = concat_batches(&self.input_schema, buffered_batches.as_slice())?;
1021                self.switch_to_safe_mode(batch)?;
1022            } else {
1023                self.input_buffered_rows = 0;
1024            }
1025        }
1026
1027        if self.mode == FoldMode::Safe {
1028            self.process_safe_mode_buffer()?;
1029            self.finalize_safe_group()?;
1030        }
1031
1032        Ok(())
1033    }
1034
1035    fn is_positive_infinity(le_array: &ArrayRef, index: usize) -> bool {
1036        matches!(
1037            string_array_value_at_index(le_array, index).and_then(|value| value.parse::<f64>().ok()),
1038            Some(value) if value.is_infinite() && value.is_sign_positive()
1039        )
1040    }
1041
1042    /// Evaluate the field column and return the result
1043    fn evaluate_row(quantile: f64, bucket: &[f64], counter: &[f64]) -> DataFusionResult<f64> {
1044        // check bucket
1045        if bucket.len() <= 1 {
1046            return Ok(f64::NAN);
1047        }
1048        if bucket.last().unwrap().is_finite() {
1049            return Err(DataFusionError::Execution(
1050                "last bucket should be +Inf".to_string(),
1051            ));
1052        }
1053        if bucket.len() != counter.len() {
1054            return Err(DataFusionError::Execution(
1055                "bucket and counter should have the same length".to_string(),
1056            ));
1057        }
1058        // check quantile
1059        if quantile < 0.0 {
1060            return Ok(f64::NEG_INFINITY);
1061        } else if quantile > 1.0 {
1062            return Ok(f64::INFINITY);
1063        } else if quantile.is_nan() {
1064            return Ok(f64::NAN);
1065        }
1066
1067        // check input value
1068        if !bucket.windows(2).all(|w| w[0] <= w[1]) {
1069            return Ok(f64::NAN);
1070        }
1071        let counter = {
1072            let needs_fix =
1073                counter.iter().any(|v| !v.is_finite()) || !counter.windows(2).all(|w| w[0] <= w[1]);
1074            if !needs_fix {
1075                Cow::Borrowed(counter)
1076            } else {
1077                let mut fixed = Vec::with_capacity(counter.len());
1078                let mut prev = 0.0;
1079                for (idx, &v) in counter.iter().enumerate() {
1080                    let mut val = if v.is_finite() { v } else { prev };
1081                    if idx > 0 && val < prev {
1082                        val = prev;
1083                    }
1084                    fixed.push(val);
1085                    prev = val;
1086                }
1087                Cow::Owned(fixed)
1088            }
1089        };
1090
1091        let total = *counter.last().unwrap();
1092        let expected_pos = total * quantile;
1093        let mut fit_bucket_pos = 0;
1094        while fit_bucket_pos < bucket.len() && counter[fit_bucket_pos] < expected_pos {
1095            fit_bucket_pos += 1;
1096        }
1097        if fit_bucket_pos >= bucket.len() - 1 {
1098            Ok(bucket[bucket.len() - 2])
1099        } else {
1100            let upper_bound = bucket[fit_bucket_pos];
1101            let upper_count = counter[fit_bucket_pos];
1102            let mut lower_bound = bucket[0].min(0.0);
1103            let mut lower_count = 0.0;
1104            if fit_bucket_pos > 0 {
1105                lower_bound = bucket[fit_bucket_pos - 1];
1106                lower_count = counter[fit_bucket_pos - 1];
1107            }
1108            if (upper_count - lower_count).abs() < 1e-10 {
1109                return Ok(f64::NAN);
1110            }
1111            Ok(lower_bound
1112                + (upper_bound - lower_bound) / (upper_count - lower_count)
1113                    * (expected_pos - lower_count))
1114        }
1115    }
1116}
1117
1118#[cfg(test)]
1119mod test {
1120    use std::sync::Arc;
1121
1122    use datafusion::arrow::array::{
1123        DictionaryArray, Float64Array, StringDictionaryBuilder, TimestampMillisecondArray,
1124    };
1125    use datafusion::arrow::datatypes::{Field, Schema, SchemaRef, TimeUnit, UInt32Type};
1126    use datafusion::common::ToDFSchema;
1127    use datafusion::datasource::memory::MemorySourceConfig;
1128    use datafusion::datasource::source::DataSourceExec;
1129    use datafusion::logical_expr::EmptyRelation;
1130    use datafusion::prelude::SessionContext;
1131    use datatypes::arrow_array::StringArray;
1132    use futures::FutureExt;
1133
1134    use super::*;
1135
1136    fn project_batch(batch: &RecordBatch, indices: &[usize]) -> RecordBatch {
1137        let fields = indices
1138            .iter()
1139            .map(|&idx| batch.schema().field(idx).clone())
1140            .collect::<Vec<_>>();
1141        let columns = indices
1142            .iter()
1143            .map(|&idx| batch.column(idx).clone())
1144            .collect::<Vec<_>>();
1145        let schema = Arc::new(Schema::new(fields));
1146        RecordBatch::try_new(schema, columns).unwrap()
1147    }
1148
1149    fn prepare_test_data() -> DataSourceExec {
1150        let schema = Arc::new(Schema::new(vec![
1151            Field::new("host", DataType::Utf8, true),
1152            Field::new("le", DataType::Utf8, true),
1153            Field::new("val", DataType::Float64, true),
1154        ]));
1155
1156        // 12 items
1157        let host_column_1 = Arc::new(StringArray::from(vec![
1158            "host_1", "host_1", "host_1", "host_1", "host_1", "host_1", "host_1", "host_1",
1159            "host_1", "host_1", "host_1", "host_1",
1160        ])) as _;
1161        let le_column_1 = Arc::new(StringArray::from(vec![
1162            "0.001", "0.1", "10", "1000", "+Inf", "0.001", "0.1", "10", "1000", "+inf", "0.001",
1163            "0.1",
1164        ])) as _;
1165        let val_column_1 = Arc::new(Float64Array::from(vec![
1166            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,
1167        ])) as _;
1168
1169        // 2 items
1170        let host_column_2 = Arc::new(StringArray::from(vec!["host_1", "host_1"])) as _;
1171        let le_column_2 = Arc::new(StringArray::from(vec!["10", "1000"])) as _;
1172        let val_column_2 = Arc::new(Float64Array::from(vec![1.0, 1.0])) as _;
1173
1174        // 11 items
1175        let host_column_3 = Arc::new(StringArray::from(vec![
1176            "host_1", "host_2", "host_2", "host_2", "host_2", "host_2", "host_2", "host_2",
1177            "host_2", "host_2", "host_2",
1178        ])) as _;
1179        let le_column_3 = Arc::new(StringArray::from(vec![
1180            "+INF", "0.001", "0.1", "10", "1000", "+iNf", "0.001", "0.1", "10", "1000", "+Inf",
1181        ])) as _;
1182        let val_column_3 = Arc::new(Float64Array::from(vec![
1183            1.0, 0_0.0, 0.0, 0.0, 0.0, 0.0, 0_0.0, 1.0, 2.0, 3.0, 4.0,
1184        ])) as _;
1185
1186        let data_1 = RecordBatch::try_new(
1187            schema.clone(),
1188            vec![host_column_1, le_column_1, val_column_1],
1189        )
1190        .unwrap();
1191        let data_2 = RecordBatch::try_new(
1192            schema.clone(),
1193            vec![host_column_2, le_column_2, val_column_2],
1194        )
1195        .unwrap();
1196        let data_3 = RecordBatch::try_new(
1197            schema.clone(),
1198            vec![host_column_3, le_column_3, val_column_3],
1199        )
1200        .unwrap();
1201
1202        DataSourceExec::new(Arc::new(
1203            MemorySourceConfig::try_new(&[vec![data_1, data_2, data_3]], schema, None).unwrap(),
1204        ))
1205    }
1206
1207    fn build_fold_exec_from_batches(
1208        batches: Vec<RecordBatch>,
1209        schema: SchemaRef,
1210        quantile: f64,
1211        ts_column_index: usize,
1212    ) -> Arc<HistogramFoldExec> {
1213        let input: Arc<dyn ExecutionPlan> = Arc::new(DataSourceExec::new(Arc::new(
1214            MemorySourceConfig::try_new(&[batches], schema.clone(), None).unwrap(),
1215        )));
1216        let output_schema: SchemaRef = Arc::new(
1217            HistogramFold::convert_schema(&Arc::new(input.schema().to_dfschema().unwrap()), "le")
1218                .unwrap()
1219                .as_arrow()
1220                .clone(),
1221        );
1222
1223        let (tag_columns, partition_exprs, properties) =
1224            build_test_plan_properties(&input, output_schema.clone(), ts_column_index);
1225
1226        Arc::new(HistogramFoldExec {
1227            le_column_index: 1,
1228            field_column_index: 2,
1229            quantile,
1230            ts_column_index,
1231            input,
1232            output_schema,
1233            tag_columns,
1234            partition_exprs,
1235            metric: ExecutionPlanMetricsSet::new(),
1236            properties,
1237        })
1238    }
1239
1240    type PlanPropsResult = (
1241        Vec<Arc<dyn PhysicalExpr>>,
1242        Vec<Arc<dyn PhysicalExpr>>,
1243        Arc<PlanProperties>,
1244    );
1245
1246    fn build_test_plan_properties(
1247        input: &Arc<dyn ExecutionPlan>,
1248        output_schema: SchemaRef,
1249        ts_column_index: usize,
1250    ) -> PlanPropsResult {
1251        let tag_columns = input
1252            .schema()
1253            .fields()
1254            .iter()
1255            .enumerate()
1256            .filter_map(|(idx, field)| {
1257                if idx == 1 || idx == 2 || idx == ts_column_index {
1258                    None
1259                } else {
1260                    Some(Arc::new(PhyColumn::new(field.name(), idx)) as _)
1261                }
1262            })
1263            .collect::<Vec<_>>();
1264
1265        let partition_exprs = if tag_columns.is_empty() {
1266            vec![Arc::new(PhyColumn::new(
1267                input.schema().field(ts_column_index).name(),
1268                ts_column_index,
1269            )) as _]
1270        } else {
1271            tag_columns.clone()
1272        };
1273
1274        let properties = PlanProperties::new(
1275            EquivalenceProperties::new(output_schema.clone()),
1276            Partitioning::Hash(
1277                partition_exprs.clone(),
1278                input.output_partitioning().partition_count(),
1279            ),
1280            EmissionType::Incremental,
1281            Boundedness::Bounded,
1282        );
1283
1284        (tag_columns, partition_exprs, Arc::new(properties))
1285    }
1286
1287    #[tokio::test]
1288    async fn fold_overall() {
1289        let memory_exec: Arc<dyn ExecutionPlan> = Arc::new(prepare_test_data());
1290        let output_schema: SchemaRef = Arc::new(
1291            HistogramFold::convert_schema(
1292                &Arc::new(memory_exec.schema().to_dfschema().unwrap()),
1293                "le",
1294            )
1295            .unwrap()
1296            .as_arrow()
1297            .clone(),
1298        );
1299        let (tag_columns, partition_exprs, properties) =
1300            build_test_plan_properties(&memory_exec, output_schema.clone(), 0);
1301        let fold_exec = Arc::new(HistogramFoldExec {
1302            le_column_index: 1,
1303            field_column_index: 2,
1304            quantile: 0.4,
1305            ts_column_index: 0,
1306            input: memory_exec,
1307            output_schema,
1308            tag_columns,
1309            partition_exprs,
1310            metric: ExecutionPlanMetricsSet::new(),
1311            properties,
1312        });
1313
1314        let session_context = SessionContext::default();
1315        let result = datafusion::physical_plan::collect(fold_exec, session_context.task_ctx())
1316            .await
1317            .unwrap();
1318        let result_literal = datatypes::arrow::util::pretty::pretty_format_batches(&result)
1319            .unwrap()
1320            .to_string();
1321
1322        let expected = String::from(
1323            "+--------+-------------------+
1324| host   | val               |
1325+--------+-------------------+
1326| host_1 | 257.5             |
1327| host_1 | 5.05              |
1328| host_1 | 0.0004            |
1329| host_2 | NaN               |
1330| host_2 | 6.040000000000001 |
1331+--------+-------------------+",
1332        );
1333        assert_eq!(result_literal, expected);
1334    }
1335
1336    #[tokio::test]
1337    async fn fold_dictionary_encoded_labels() {
1338        let dictionary_type =
1339            DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8));
1340        let schema = Arc::new(Schema::new(vec![
1341            Field::new("host", dictionary_type.clone(), true),
1342            Field::new("le", dictionary_type, true),
1343            Field::new("val", DataType::Float64, true),
1344        ]));
1345
1346        let mut host = StringDictionaryBuilder::<UInt32Type>::new();
1347        let mut le = StringDictionaryBuilder::<UInt32Type>::new();
1348        for value in ["host_1", "host_1", "host_1"] {
1349            host.append_value(value);
1350        }
1351        for value in ["0.1", "1", "+Inf"] {
1352            le.append_value(value);
1353        }
1354        let batch = RecordBatch::try_new(
1355            schema.clone(),
1356            vec![
1357                Arc::new(host.finish()),
1358                Arc::new(le.finish()),
1359                Arc::new(Float64Array::from(vec![1.0, 2.0, 3.0])),
1360            ],
1361        )
1362        .unwrap();
1363
1364        let fold_exec = build_fold_exec_from_batches(vec![batch], schema, 0.5, 0);
1365        let result =
1366            datafusion::physical_plan::collect(fold_exec, SessionContext::default().task_ctx())
1367                .await
1368                .unwrap();
1369
1370        assert_eq!(result.len(), 1);
1371        assert_eq!(result[0].num_rows(), 1);
1372        let host = result[0]
1373            .column(0)
1374            .as_any()
1375            .downcast_ref::<DictionaryArray<UInt32Type>>()
1376            .unwrap();
1377        assert_eq!(host.values().len(), 1);
1378        assert_eq!(
1379            string_array_value_at_index(result[0].column(0), 0),
1380            Some("host_1")
1381        );
1382        let value = result[0].column(1).as_primitive::<Float64Type>().value(0);
1383        assert!((value - 0.55).abs() < 1e-12);
1384    }
1385
1386    #[tokio::test]
1387    async fn pruning_should_keep_le_column_for_exec() {
1388        let schema = Arc::new(Schema::new(vec![
1389            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
1390            Field::new("le", DataType::Utf8, true),
1391            Field::new("val", DataType::Float64, true),
1392        ]));
1393        let df_schema = schema.clone().to_dfschema_ref().unwrap();
1394        let input = LogicalPlan::EmptyRelation(EmptyRelation {
1395            produce_one_row: false,
1396            schema: df_schema,
1397        });
1398        let plan = HistogramFold::new(
1399            "le".to_string(),
1400            "val".to_string(),
1401            "ts".to_string(),
1402            0.5,
1403            input,
1404        )
1405        .unwrap();
1406
1407        let output_columns = [0usize, 1usize];
1408        let required = plan.necessary_children_exprs(&output_columns).unwrap();
1409        let required = &required[0];
1410        assert_eq!(required.as_slice(), &[0, 1, 2]);
1411
1412        let input_batch = RecordBatch::try_new(
1413            schema,
1414            vec![
1415                Arc::new(TimestampMillisecondArray::from(vec![0, 0])),
1416                Arc::new(StringArray::from(vec!["0.1", "+Inf"])),
1417                Arc::new(Float64Array::from(vec![1.0, 2.0])),
1418            ],
1419        )
1420        .unwrap();
1421        let projected = project_batch(&input_batch, required);
1422        let projected_schema = projected.schema();
1423        let memory_exec = Arc::new(DataSourceExec::new(Arc::new(
1424            MemorySourceConfig::try_new(&[vec![projected]], projected_schema, None).unwrap(),
1425        )));
1426
1427        let fold_exec = plan.to_execution_plan(memory_exec);
1428        let session_context = SessionContext::default();
1429        let output_batches =
1430            datafusion::physical_plan::collect(fold_exec, session_context.task_ctx())
1431                .await
1432                .unwrap();
1433        assert_eq!(output_batches.len(), 1);
1434
1435        let output_batch = &output_batches[0];
1436        assert_eq!(output_batch.num_rows(), 1);
1437
1438        let ts = output_batch
1439            .column(0)
1440            .as_any()
1441            .downcast_ref::<TimestampMillisecondArray>()
1442            .unwrap();
1443        assert_eq!(ts.values(), &[0i64]);
1444
1445        let values = output_batch
1446            .column(1)
1447            .as_any()
1448            .downcast_ref::<Float64Array>()
1449            .unwrap();
1450        assert!((values.value(0) - 0.1).abs() < 1e-12);
1451
1452        // Simulate the pre-fix pruning behavior: omit the `le` column from the child input.
1453        let le_index = 1usize;
1454        let broken_required = output_columns
1455            .iter()
1456            .map(|&output_column| {
1457                if output_column < le_index {
1458                    output_column
1459                } else {
1460                    output_column + 1
1461                }
1462            })
1463            .collect::<Vec<_>>();
1464
1465        let broken = project_batch(&input_batch, &broken_required);
1466        let broken_schema = broken.schema();
1467        let broken_exec = Arc::new(DataSourceExec::new(Arc::new(
1468            MemorySourceConfig::try_new(&[vec![broken]], broken_schema, None).unwrap(),
1469        )));
1470        let broken_fold_exec = plan.to_execution_plan(broken_exec);
1471        let session_context = SessionContext::default();
1472        let broken_result = std::panic::AssertUnwindSafe(async {
1473            datafusion::physical_plan::collect(broken_fold_exec, session_context.task_ctx()).await
1474        })
1475        .catch_unwind()
1476        .await;
1477        assert!(broken_result.is_err());
1478    }
1479
1480    #[test]
1481    fn confirm_schema() {
1482        let input_schema = Schema::new(vec![
1483            Field::new("host", DataType::Utf8, true),
1484            Field::new("le", DataType::Utf8, true),
1485            Field::new("val", DataType::Float64, true),
1486        ])
1487        .to_dfschema_ref()
1488        .unwrap();
1489        let expected_output_schema = Schema::new(vec![
1490            Field::new("host", DataType::Utf8, true),
1491            Field::new("val", DataType::Float64, true),
1492        ])
1493        .to_dfschema_ref()
1494        .unwrap();
1495
1496        let actual = HistogramFold::convert_schema(&input_schema, "le").unwrap();
1497        assert_eq!(actual, expected_output_schema)
1498    }
1499
1500    #[tokio::test]
1501    async fn fallback_to_safe_mode_on_missing_inf() {
1502        let schema = Arc::new(Schema::new(vec![
1503            Field::new("host", DataType::Utf8, true),
1504            Field::new("le", DataType::Utf8, true),
1505            Field::new("val", DataType::Float64, true),
1506        ]));
1507        let host_column = Arc::new(StringArray::from(vec!["a", "a", "a", "a", "b", "b"])) as _;
1508        let le_column = Arc::new(StringArray::from(vec![
1509            "0.1", "+Inf", "0.1", "1.0", "0.1", "+Inf",
1510        ])) as _;
1511        let val_column = Arc::new(Float64Array::from(vec![1.0, 2.0, 3.0, 3.0, 1.0, 5.0])) as _;
1512        let batch =
1513            RecordBatch::try_new(schema.clone(), vec![host_column, le_column, val_column]).unwrap();
1514        let fold_exec = build_fold_exec_from_batches(vec![batch], schema, 0.5, 0);
1515        let session_context = SessionContext::default();
1516        let result = datafusion::physical_plan::collect(fold_exec, session_context.task_ctx())
1517            .await
1518            .unwrap();
1519        let result_literal = datatypes::arrow::util::pretty::pretty_format_batches(&result)
1520            .unwrap()
1521            .to_string();
1522
1523        let expected = String::from(
1524            "+------+-----+
1525| host | val |
1526+------+-----+
1527| a    | 0.1 |
1528| a    | NaN |
1529| b    | 0.1 |
1530+------+-----+",
1531        );
1532        assert_eq!(result_literal, expected);
1533    }
1534
1535    #[tokio::test]
1536    async fn emit_nan_when_no_inf_present() {
1537        let schema = Arc::new(Schema::new(vec![
1538            Field::new("host", DataType::Utf8, true),
1539            Field::new("le", DataType::Utf8, true),
1540            Field::new("val", DataType::Float64, true),
1541        ]));
1542        let host_column = Arc::new(StringArray::from(vec!["c", "c"])) as _;
1543        let le_column = Arc::new(StringArray::from(vec!["0.1", "1.0"])) as _;
1544        let val_column = Arc::new(Float64Array::from(vec![1.0, 2.0])) as _;
1545        let batch =
1546            RecordBatch::try_new(schema.clone(), vec![host_column, le_column, val_column]).unwrap();
1547        let fold_exec = build_fold_exec_from_batches(vec![batch], schema, 0.9, 0);
1548        let session_context = SessionContext::default();
1549        let result = datafusion::physical_plan::collect(fold_exec, session_context.task_ctx())
1550            .await
1551            .unwrap();
1552        let result_literal = datatypes::arrow::util::pretty::pretty_format_batches(&result)
1553            .unwrap()
1554            .to_string();
1555
1556        let expected = String::from(
1557            "+------+-----+
1558| host | val |
1559+------+-----+
1560| c    | NaN |
1561+------+-----+",
1562        );
1563        assert_eq!(result_literal, expected);
1564    }
1565
1566    #[tokio::test]
1567    async fn safe_mode_handles_misaligned_groups() {
1568        let schema = Arc::new(Schema::new(vec![
1569            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
1570            Field::new("le", DataType::Utf8, true),
1571            Field::new("val", DataType::Float64, true),
1572        ]));
1573
1574        let ts_column = Arc::new(TimestampMillisecondArray::from(vec![
1575            2900000, 2900000, 2900000, 3000000, 3000000, 3000000, 3000000, 3005000, 3005000,
1576            3010000, 3010000, 3010000, 3010000, 3010000,
1577        ])) as _;
1578        let le_column = Arc::new(StringArray::from(vec![
1579            "0.1", "1", "5", "0.1", "1", "5", "+Inf", "0.1", "+Inf", "0.1", "1", "3", "5", "+Inf",
1580        ])) as _;
1581        let val_column = Arc::new(Float64Array::from(vec![
1582            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,
1583        ])) as _;
1584        let batch =
1585            RecordBatch::try_new(schema.clone(), vec![ts_column, le_column, val_column]).unwrap();
1586        let fold_exec = build_fold_exec_from_batches(vec![batch], schema, 0.5, 0);
1587        let session_context = SessionContext::default();
1588        let result = datafusion::physical_plan::collect(fold_exec, session_context.task_ctx())
1589            .await
1590            .unwrap();
1591
1592        let mut values = Vec::new();
1593        for batch in result {
1594            let array = batch.column(1).as_primitive::<Float64Type>();
1595            values.extend(array.iter().map(|v| v.unwrap()));
1596        }
1597
1598        assert_eq!(values.len(), 4);
1599        assert!(values[0].is_nan());
1600        assert!((values[1] - 0.55).abs() < 1e-10);
1601        assert!((values[2] - 0.1).abs() < 1e-10);
1602        assert!((values[3] - 2.0).abs() < 1e-10);
1603    }
1604
1605    #[tokio::test]
1606    async fn missing_buckets_at_first_timestamp() {
1607        let schema = Arc::new(Schema::new(vec![
1608            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
1609            Field::new("le", DataType::Utf8, true),
1610            Field::new("val", DataType::Float64, true),
1611        ]));
1612
1613        let ts_column = Arc::new(TimestampMillisecondArray::from(vec![
1614            2_900_000, 3_000_000, 3_000_000, 3_000_000, 3_000_000, 3_005_000, 3_005_000, 3_010_000,
1615            3_010_000, 3_010_000, 3_010_000, 3_010_000,
1616        ])) as _;
1617        let le_column = Arc::new(StringArray::from(vec![
1618            "0.1", "0.1", "1", "5", "+Inf", "0.1", "+Inf", "0.1", "1", "3", "5", "+Inf",
1619        ])) as _;
1620        let val_column = Arc::new(Float64Array::from(vec![
1621            0.0, 50.0, 70.0, 110.0, 120.0, 10.0, 30.0, 10.0, 20.0, 30.0, 40.0, 50.0,
1622        ])) as _;
1623
1624        let batch =
1625            RecordBatch::try_new(schema.clone(), vec![ts_column, le_column, val_column]).unwrap();
1626        let fold_exec = build_fold_exec_from_batches(vec![batch], schema, 0.5, 0);
1627        let session_context = SessionContext::default();
1628        let result = datafusion::physical_plan::collect(fold_exec, session_context.task_ctx())
1629            .await
1630            .unwrap();
1631
1632        let mut values = Vec::new();
1633        for batch in result {
1634            let array = batch.column(1).as_primitive::<Float64Type>();
1635            values.extend(array.iter().map(|v| v.unwrap()));
1636        }
1637
1638        assert_eq!(values.len(), 4);
1639        assert!(values[0].is_nan());
1640        assert!((values[1] - 0.55).abs() < 1e-10);
1641        assert!((values[2] - 0.1).abs() < 1e-10);
1642        assert!((values[3] - 2.0).abs() < 1e-10);
1643    }
1644
1645    #[tokio::test]
1646    async fn missing_inf_in_first_group() {
1647        let schema = Arc::new(Schema::new(vec![
1648            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
1649            Field::new("le", DataType::Utf8, true),
1650            Field::new("val", DataType::Float64, true),
1651        ]));
1652
1653        let ts_column = Arc::new(TimestampMillisecondArray::from(vec![
1654            1000, 1000, 1000, 2000, 2000, 2000, 2000,
1655        ])) as _;
1656        let le_column = Arc::new(StringArray::from(vec![
1657            "0.1", "1", "5", "0.1", "1", "5", "+Inf",
1658        ])) as _;
1659        let val_column = Arc::new(Float64Array::from(vec![
1660            0.0, 0.0, 0.0, 10.0, 20.0, 30.0, 30.0,
1661        ])) as _;
1662        let batch =
1663            RecordBatch::try_new(schema.clone(), vec![ts_column, le_column, val_column]).unwrap();
1664        let fold_exec = build_fold_exec_from_batches(vec![batch], schema, 0.5, 0);
1665        let session_context = SessionContext::default();
1666        let result = datafusion::physical_plan::collect(fold_exec, session_context.task_ctx())
1667            .await
1668            .unwrap();
1669
1670        let mut values = Vec::new();
1671        for batch in result {
1672            let array = batch.column(1).as_primitive::<Float64Type>();
1673            values.extend(array.iter().map(|v| v.unwrap()));
1674        }
1675
1676        assert_eq!(values.len(), 2);
1677        assert!(values[0].is_nan());
1678        assert!((values[1] - 0.55).abs() < 1e-10, "{values:?}");
1679    }
1680
1681    #[test]
1682    fn evaluate_row_normal_case() {
1683        let bucket = [0.0, 1.0, 2.0, 3.0, 4.0, f64::INFINITY];
1684
1685        #[derive(Debug)]
1686        struct Case {
1687            quantile: f64,
1688            counters: Vec<f64>,
1689            expected: f64,
1690        }
1691
1692        let cases = [
1693            Case {
1694                quantile: 0.9,
1695                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
1696                expected: 4.0,
1697            },
1698            Case {
1699                quantile: 0.89,
1700                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
1701                expected: 4.0,
1702            },
1703            Case {
1704                quantile: 0.78,
1705                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
1706                expected: 3.9,
1707            },
1708            Case {
1709                quantile: 0.5,
1710                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
1711                expected: 2.5,
1712            },
1713            Case {
1714                quantile: 0.5,
1715                counters: vec![0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
1716                expected: f64::NAN,
1717            },
1718            Case {
1719                quantile: 1.0,
1720                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
1721                expected: 4.0,
1722            },
1723            Case {
1724                quantile: 0.0,
1725                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
1726                expected: f64::NAN,
1727            },
1728            Case {
1729                quantile: 1.1,
1730                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
1731                expected: f64::INFINITY,
1732            },
1733            Case {
1734                quantile: -1.0,
1735                counters: vec![0.0, 10.0, 20.0, 30.0, 40.0, 50.0],
1736                expected: f64::NEG_INFINITY,
1737            },
1738        ];
1739
1740        for case in cases {
1741            let actual =
1742                HistogramFoldStream::evaluate_row(case.quantile, &bucket, &case.counters).unwrap();
1743            assert_eq!(
1744                format!("{actual}"),
1745                format!("{}", case.expected),
1746                "{:?}",
1747                case
1748            );
1749        }
1750    }
1751
1752    #[test]
1753    fn evaluate_out_of_order_input() {
1754        let bucket = [0.0, 1.0, 2.0, 3.0, 4.0, f64::INFINITY];
1755        let counters = [5.0, 4.0, 3.0, 2.0, 1.0, 0.0];
1756        let result = HistogramFoldStream::evaluate_row(0.5, &bucket, &counters).unwrap();
1757        assert_eq!(0.0, result);
1758    }
1759
1760    #[test]
1761    fn evaluate_wrong_bucket() {
1762        let bucket = [0.0, 1.0, 2.0, 3.0, 4.0, f64::INFINITY, 5.0];
1763        let counters = [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0];
1764        let result = HistogramFoldStream::evaluate_row(0.5, &bucket, &counters);
1765        assert!(result.is_err());
1766    }
1767
1768    #[test]
1769    fn evaluate_small_fraction() {
1770        let bucket = [0.0, 2.0, 4.0, 6.0, f64::INFINITY];
1771        let counters = [0.0, 1.0 / 300.0, 2.0 / 300.0, 0.01, 0.01];
1772        let result = HistogramFoldStream::evaluate_row(0.5, &bucket, &counters).unwrap();
1773        assert_eq!(3.0, result);
1774    }
1775
1776    #[test]
1777    fn evaluate_non_monotonic_counter() {
1778        let bucket = [0.0, 1.0, 2.0, 3.0, f64::INFINITY];
1779        let counters = [0.1, 0.2, 0.4, 0.17, 0.5];
1780        let result = HistogramFoldStream::evaluate_row(0.5, &bucket, &counters).unwrap();
1781        assert!((result - 1.25).abs() < 1e-10, "{result}");
1782    }
1783
1784    #[test]
1785    fn evaluate_nan_counter() {
1786        let bucket = [0.0, 1.0, 2.0, 3.0, f64::INFINITY];
1787        let counters = [f64::NAN, 1.0, 2.0, 3.0, 3.0];
1788        let result = HistogramFoldStream::evaluate_row(0.5, &bucket, &counters).unwrap();
1789        assert!((result - 1.5).abs() < 1e-10, "{result}");
1790    }
1791
1792    fn build_empty_relation(schema: &Arc<Schema>) -> LogicalPlan {
1793        LogicalPlan::EmptyRelation(EmptyRelation {
1794            produce_one_row: false,
1795            schema: schema.clone().to_dfschema_ref().unwrap(),
1796        })
1797    }
1798
1799    #[tokio::test]
1800    async fn encode_decode_histogram_fold() {
1801        let schema = Arc::new(Schema::new(vec![
1802            Field::new("ts", DataType::Int64, false),
1803            Field::new("le", DataType::Utf8, false),
1804            Field::new("val", DataType::Float64, false),
1805        ]));
1806        let input_plan = build_empty_relation(&schema);
1807        let plan_node = HistogramFold::new(
1808            "le".to_string(),
1809            "val".to_string(),
1810            "ts".to_string(),
1811            0.8,
1812            input_plan.clone(),
1813        )
1814        .unwrap();
1815
1816        let bytes = plan_node.serialize();
1817
1818        let histogram_fold = HistogramFold::deserialize(&bytes).unwrap();
1819        // need fix
1820        let histogram_fold = histogram_fold
1821            .with_exprs_and_inputs(vec![], vec![input_plan])
1822            .unwrap();
1823
1824        assert_eq!(histogram_fold.le_column, "le");
1825        assert_eq!(histogram_fold.ts_column, "ts");
1826        assert_eq!(histogram_fold.field_column, "val");
1827        assert_eq!(histogram_fold.quantile, OrderedF64::from(0.8));
1828        assert_eq!(histogram_fold.output_schema.fields().len(), 2);
1829        assert_eq!(histogram_fold.output_schema.field(0).name(), "ts");
1830        assert_eq!(histogram_fold.output_schema.field(1).name(), "val");
1831    }
1832}