Skip to main content

promql/extension_plan/
scalar_calculate.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::collections::HashMap;
16use std::pin::Pin;
17use std::sync::Arc;
18use std::task::{Context, Poll};
19
20use datafusion::common::stats::Precision;
21use datafusion::common::tree_node::TreeNodeRecursion;
22use datafusion::common::{
23    DFSchema, DFSchemaRef, Result as DataFusionResult, Statistics, TableReference,
24};
25use datafusion::error::DataFusionError;
26use datafusion::execution::context::TaskContext;
27use datafusion::logical_expr::{EmptyRelation, LogicalPlan, UserDefinedLogicalNodeCore};
28use datafusion::physical_expr::EquivalenceProperties;
29use datafusion::physical_plan::metrics::{BaselineMetrics, ExecutionPlanMetricsSet, MetricsSet};
30use datafusion::physical_plan::{
31    ChildStats, DisplayAs, DisplayFormatType, Distribution, ExecutionPlan,
32    InputDistributionRequirements, Partitioning, PhysicalExpr, PlanProperties, RecordBatchStream,
33    SendableRecordBatchStream, StatisticsArgs,
34};
35use datafusion::prelude::Expr;
36use datafusion_expr::ident;
37use datatypes::arrow::array::{Array, ArrayRef, Float64Array, TimestampMillisecondArray};
38use datatypes::arrow::compute::{CastOptions, cast_with_options, concat_batches};
39use datatypes::arrow::datatypes::{DataType, Field, Schema, SchemaRef, TimeUnit};
40use datatypes::arrow::record_batch::RecordBatch;
41use datatypes::arrow_array::string_array_value_at_index;
42use futures::{Stream, StreamExt, ready};
43use greptime_proto::substrait_extension as pb;
44use prost::Message;
45use snafu::ResultExt;
46
47use crate::error::{ColumnNotFoundSnafu, DataFusionPlanningSnafu, DeserializeSnafu, Result};
48use crate::extension_plan::{Millisecond, resolve_column_name, serialize_column_index};
49
50/// `ScalarCalculate` is the custom logical plan to calculate
51/// [`scalar`](https://prometheus.io/docs/prometheus/latest/querying/functions/#scalar)
52/// in PromQL, return NaN when have multiple time series.
53///
54/// Return the time series as scalar value when only have one time series.
55#[derive(Debug, Clone, PartialEq, Eq, Hash)]
56pub struct ScalarCalculate {
57    start: Millisecond,
58    end: Millisecond,
59    interval: Millisecond,
60
61    time_index: String,
62    tag_columns: Vec<String>,
63    field_column: String,
64    input: LogicalPlan,
65    output_schema: DFSchemaRef,
66    unfix: Option<UnfixIndices>,
67}
68
69#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd)]
70struct UnfixIndices {
71    pub time_index_idx: u64,
72    pub tag_column_indices: Vec<u64>,
73    pub field_column_idx: u64,
74}
75
76impl ScalarCalculate {
77    /// create a new `ScalarCalculate` plan
78    #[allow(clippy::too_many_arguments)]
79    pub fn new(
80        start: Millisecond,
81        end: Millisecond,
82        interval: Millisecond,
83        input: LogicalPlan,
84        time_index: &str,
85        tag_columns: &[String],
86        field_column: &str,
87        table_name: Option<&str>,
88    ) -> Result<Self> {
89        let input_schema = input.schema();
90        let Ok(ts_field) = input_schema
91            .field_with_unqualified_name(time_index)
92            .cloned()
93        else {
94            return ColumnNotFoundSnafu { col: time_index }.fail();
95        };
96        let val_field = Field::new(format!("scalar({})", field_column), DataType::Float64, true);
97        let qualifier = table_name.map(TableReference::bare);
98        let schema = DFSchema::new_with_metadata(
99            vec![
100                (qualifier.clone(), ts_field),
101                (qualifier, Arc::new(val_field)),
102            ],
103            input_schema.metadata().clone(),
104        )
105        .context(DataFusionPlanningSnafu)?;
106
107        Ok(Self {
108            start,
109            end,
110            interval,
111            time_index: time_index.to_string(),
112            tag_columns: tag_columns.to_vec(),
113            field_column: field_column.to_string(),
114            input,
115            output_schema: Arc::new(schema),
116            unfix: None,
117        })
118    }
119
120    /// The name of this custom plan
121    pub const fn name() -> &'static str {
122        "ScalarCalculate"
123    }
124
125    /// Create a new execution plan from ScalarCalculate
126    pub fn to_execution_plan(
127        &self,
128        exec_input: Arc<dyn ExecutionPlan>,
129    ) -> DataFusionResult<Arc<dyn ExecutionPlan>> {
130        let fields: Vec<_> = self
131            .output_schema
132            .fields()
133            .iter()
134            .map(|field| {
135                Field::new(field.name(), field.data_type().clone(), field.is_nullable())
136                    .with_metadata(field.metadata().clone())
137            })
138            .collect();
139        let input_schema = exec_input.schema();
140        let ts_index = input_schema
141            .index_of(&self.time_index)
142            .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?;
143        let val_index = input_schema
144            .index_of(&self.field_column)
145            .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?;
146        let schema = Arc::new(Schema::new_with_metadata(
147            fields,
148            input_schema.metadata().clone(),
149        ));
150        let properties = exec_input.properties();
151        let properties = Arc::new(PlanProperties::new(
152            EquivalenceProperties::new(schema.clone()),
153            Partitioning::UnknownPartitioning(1),
154            properties.emission_type,
155            properties.boundedness,
156        ));
157        Ok(Arc::new(ScalarCalculateExec {
158            start: self.start,
159            end: self.end,
160            interval: self.interval,
161            schema,
162            input: exec_input,
163            project_index: (ts_index, val_index),
164            tag_columns: self.tag_columns.clone(),
165            metric: ExecutionPlanMetricsSet::new(),
166            properties,
167        }))
168    }
169
170    pub fn serialize(&self) -> Vec<u8> {
171        let time_index_idx = serialize_column_index(self.input.schema(), &self.time_index);
172
173        let tag_column_indices = self
174            .tag_columns
175            .iter()
176            .map(|name| serialize_column_index(self.input.schema(), name))
177            .collect::<Vec<u64>>();
178
179        let field_column_idx = serialize_column_index(self.input.schema(), &self.field_column);
180
181        pb::ScalarCalculate {
182            start: self.start,
183            end: self.end,
184            interval: self.interval,
185            time_index_idx,
186            tag_column_indices,
187            field_column_idx,
188            ..Default::default()
189        }
190        .encode_to_vec()
191    }
192
193    pub fn deserialize(bytes: &[u8]) -> Result<Self> {
194        let pb_scalar_calculate = pb::ScalarCalculate::decode(bytes).context(DeserializeSnafu)?;
195        let placeholder_plan = LogicalPlan::EmptyRelation(EmptyRelation {
196            produce_one_row: false,
197            schema: Arc::new(DFSchema::empty()),
198        });
199
200        let unfix = UnfixIndices {
201            time_index_idx: pb_scalar_calculate.time_index_idx,
202            tag_column_indices: pb_scalar_calculate.tag_column_indices.clone(),
203            field_column_idx: pb_scalar_calculate.field_column_idx,
204        };
205
206        // TODO(Taylor-lagrange): Supports timestamps of different precisions
207        let ts_field = Field::new(
208            "placeholder_time_index",
209            DataType::Timestamp(TimeUnit::Millisecond, None),
210            true,
211        );
212        let val_field = Field::new("placeholder_field", DataType::Float64, true);
213        // TODO(Taylor-lagrange): missing tablename in pb
214        let schema = DFSchema::new_with_metadata(
215            vec![(None, Arc::new(ts_field)), (None, Arc::new(val_field))],
216            HashMap::new(),
217        )
218        .context(DataFusionPlanningSnafu)?;
219
220        Ok(Self {
221            start: pb_scalar_calculate.start,
222            end: pb_scalar_calculate.end,
223            interval: pb_scalar_calculate.interval,
224            time_index: String::new(),
225            tag_columns: Vec::new(),
226            field_column: String::new(),
227            output_schema: Arc::new(schema),
228            input: placeholder_plan,
229            unfix: Some(unfix),
230        })
231    }
232}
233
234impl PartialOrd for ScalarCalculate {
235    fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
236        // Compare fields in order excluding output_schema
237        match self.start.partial_cmp(&other.start) {
238            Some(core::cmp::Ordering::Equal) => {}
239            ord => return ord,
240        }
241        match self.end.partial_cmp(&other.end) {
242            Some(core::cmp::Ordering::Equal) => {}
243            ord => return ord,
244        }
245        match self.interval.partial_cmp(&other.interval) {
246            Some(core::cmp::Ordering::Equal) => {}
247            ord => return ord,
248        }
249        match self.time_index.partial_cmp(&other.time_index) {
250            Some(core::cmp::Ordering::Equal) => {}
251            ord => return ord,
252        }
253        match self.tag_columns.partial_cmp(&other.tag_columns) {
254            Some(core::cmp::Ordering::Equal) => {}
255            ord => return ord,
256        }
257        match self.field_column.partial_cmp(&other.field_column) {
258            Some(core::cmp::Ordering::Equal) => {}
259            ord => return ord,
260        }
261        self.input.partial_cmp(&other.input)
262    }
263}
264
265impl UserDefinedLogicalNodeCore for ScalarCalculate {
266    fn name(&self) -> &str {
267        Self::name()
268    }
269
270    fn inputs(&self) -> Vec<&LogicalPlan> {
271        vec![&self.input]
272    }
273
274    fn schema(&self) -> &DFSchemaRef {
275        &self.output_schema
276    }
277
278    fn expressions(&self) -> Vec<Expr> {
279        if self.unfix.is_some() {
280            return vec![];
281        }
282
283        self.tag_columns
284            .iter()
285            .map(ident)
286            .chain(std::iter::once(ident(&self.time_index)))
287            .chain(std::iter::once(ident(&self.field_column)))
288            .collect()
289    }
290
291    fn necessary_children_exprs(&self, _output_columns: &[usize]) -> Option<Vec<Vec<usize>>> {
292        if self.unfix.is_some() {
293            return None;
294        }
295
296        let input_schema = self.input.schema();
297        let time_index_idx = input_schema.index_of_column_by_name(None, &self.time_index)?;
298        let field_column_idx = input_schema.index_of_column_by_name(None, &self.field_column)?;
299
300        let mut required = Vec::with_capacity(2 + self.tag_columns.len());
301        required.extend([time_index_idx, field_column_idx]);
302        for tag in &self.tag_columns {
303            required.push(input_schema.index_of_column_by_name(None, tag)?);
304        }
305
306        required.sort_unstable();
307        required.dedup();
308        Some(vec![required])
309    }
310
311    fn fmt_for_explain(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
312        write!(f, "ScalarCalculate: tags={:?}", self.tag_columns)
313    }
314
315    fn with_exprs_and_inputs(
316        &self,
317        _exprs: Vec<Expr>,
318        inputs: Vec<LogicalPlan>,
319    ) -> DataFusionResult<Self> {
320        let input: LogicalPlan = inputs.into_iter().next().unwrap();
321        let input_schema = input.schema();
322
323        if let Some(unfix) = &self.unfix {
324            // transform indices to names
325            let time_index = resolve_column_name(
326                unfix.time_index_idx,
327                input_schema,
328                "ScalarCalculate",
329                "time index",
330            )?;
331
332            let tag_columns = unfix
333                .tag_column_indices
334                .iter()
335                .map(|idx| resolve_column_name(*idx, input_schema, "ScalarCalculate", "tag"))
336                .collect::<DataFusionResult<Vec<String>>>()?;
337
338            let field_column = resolve_column_name(
339                unfix.field_column_idx,
340                input_schema,
341                "ScalarCalculate",
342                "field",
343            )?;
344
345            // Recreate output schema with actual field names
346            let ts_field = Field::new(
347                &time_index,
348                DataType::Timestamp(TimeUnit::Millisecond, None),
349                true,
350            );
351            let val_field =
352                Field::new(format!("scalar({})", field_column), DataType::Float64, true);
353            let schema = DFSchema::new_with_metadata(
354                vec![(None, Arc::new(ts_field)), (None, Arc::new(val_field))],
355                HashMap::new(),
356            )
357            .context(DataFusionPlanningSnafu)?;
358
359            Ok(ScalarCalculate {
360                start: self.start,
361                end: self.end,
362                interval: self.interval,
363                time_index,
364                tag_columns,
365                field_column,
366                input,
367                output_schema: Arc::new(schema),
368                unfix: None,
369            })
370        } else {
371            Ok(ScalarCalculate {
372                start: self.start,
373                end: self.end,
374                interval: self.interval,
375                time_index: self.time_index.clone(),
376                tag_columns: self.tag_columns.clone(),
377                field_column: self.field_column.clone(),
378                input,
379                output_schema: self.output_schema.clone(),
380                unfix: None,
381            })
382        }
383    }
384}
385
386#[derive(Debug, Clone)]
387struct ScalarCalculateExec {
388    start: Millisecond,
389    end: Millisecond,
390    interval: Millisecond,
391    schema: SchemaRef,
392    project_index: (usize, usize),
393    input: Arc<dyn ExecutionPlan>,
394    tag_columns: Vec<String>,
395    metric: ExecutionPlanMetricsSet,
396    properties: Arc<PlanProperties>,
397}
398
399impl ExecutionPlan for ScalarCalculateExec {
400    fn apply_expressions(
401        &self,
402        _f: &mut dyn FnMut(&Arc<dyn PhysicalExpr>) -> datafusion_common::Result<TreeNodeRecursion>,
403    ) -> DataFusionResult<TreeNodeRecursion> {
404        Ok(TreeNodeRecursion::Continue)
405    }
406
407    fn schema(&self) -> SchemaRef {
408        self.schema.clone()
409    }
410
411    fn properties(&self) -> &Arc<PlanProperties> {
412        &self.properties
413    }
414
415    fn maintains_input_order(&self) -> Vec<bool> {
416        vec![true; self.children().len()]
417    }
418
419    fn input_distribution_requirements(&self) -> InputDistributionRequirements {
420        InputDistributionRequirements::new(vec![Distribution::SinglePartition])
421    }
422
423    fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
424        vec![&self.input]
425    }
426
427    fn with_new_children(
428        self: Arc<Self>,
429        children: Vec<Arc<dyn ExecutionPlan>>,
430    ) -> DataFusionResult<Arc<dyn ExecutionPlan>> {
431        Ok(Arc::new(ScalarCalculateExec {
432            start: self.start,
433            end: self.end,
434            interval: self.interval,
435            schema: self.schema.clone(),
436            project_index: self.project_index,
437            tag_columns: self.tag_columns.clone(),
438            input: children[0].clone(),
439            metric: self.metric.clone(),
440            properties: self.properties.clone(),
441        }))
442    }
443
444    fn execute(
445        &self,
446        partition: usize,
447        context: Arc<TaskContext>,
448    ) -> DataFusionResult<SendableRecordBatchStream> {
449        let baseline_metric = BaselineMetrics::new(&self.metric, partition);
450        let input = self.input.execute(partition, context)?;
451        let schema = input.schema();
452        let tag_indices = self
453            .tag_columns
454            .iter()
455            .map(|tag| {
456                schema
457                    .column_with_name(tag)
458                    .unwrap_or_else(|| panic!("tag column not found {tag}"))
459                    .0
460            })
461            .collect();
462
463        Ok(Box::pin(ScalarCalculateStream {
464            start: self.start,
465            end: self.end,
466            interval: self.interval,
467            schema: self.schema.clone(),
468            project_index: self.project_index,
469            metric: baseline_metric,
470            tag_indices,
471            input,
472            have_multi_series: false,
473            done: false,
474            batch: None,
475            tag_value: None,
476        }))
477    }
478
479    fn metrics(&self) -> Option<MetricsSet> {
480        Some(self.metric.clone_inner())
481    }
482
483    fn child_stats_requests(&self, partition: Option<usize>) -> Vec<ChildStats> {
484        vec![ChildStats::At(partition)]
485    }
486
487    fn statistics_from_inputs(
488        &self,
489        input_stats: &[Arc<Statistics>],
490        _args: &StatisticsArgs,
491    ) -> DataFusionResult<Arc<Statistics>> {
492        let input_stats = &input_stats[0];
493
494        let estimated_row_num = (self.end - self.start) as f64 / self.interval as f64;
495        let estimated_total_bytes = input_stats
496            .total_byte_size
497            .get_value()
498            .zip(input_stats.num_rows.get_value())
499            .map(|(size, rows)| {
500                Precision::Inexact(((*size as f64 / *rows as f64) * estimated_row_num).floor() as _)
501            })
502            .unwrap_or_default();
503
504        Ok(Arc::new(Statistics {
505            num_rows: Precision::Inexact(estimated_row_num as _),
506            total_byte_size: estimated_total_bytes,
507            // TODO(ruihang): support this column statistics
508            column_statistics: Statistics::unknown_column(&self.schema()),
509        }))
510    }
511
512    fn name(&self) -> &str {
513        "ScalarCalculateExec"
514    }
515}
516
517impl DisplayAs for ScalarCalculateExec {
518    fn fmt_as(&self, t: DisplayFormatType, f: &mut std::fmt::Formatter) -> std::fmt::Result {
519        match t {
520            DisplayFormatType::Default
521            | DisplayFormatType::Verbose
522            | DisplayFormatType::TreeRender => {
523                write!(f, "ScalarCalculateExec: tags={:?}", self.tag_columns)
524            }
525        }
526    }
527}
528
529struct ScalarCalculateStream {
530    start: Millisecond,
531    end: Millisecond,
532    interval: Millisecond,
533    schema: SchemaRef,
534    input: SendableRecordBatchStream,
535    metric: BaselineMetrics,
536    tag_indices: Vec<usize>,
537    /// with format `(ts_index, field_index)`
538    project_index: (usize, usize),
539    have_multi_series: bool,
540    done: bool,
541    batch: Option<RecordBatch>,
542    tag_value: Option<Vec<Option<String>>>,
543}
544
545impl RecordBatchStream for ScalarCalculateStream {
546    fn schema(&self) -> SchemaRef {
547        self.schema.clone()
548    }
549}
550
551impl ScalarCalculateStream {
552    fn update_batch(&mut self, batch: RecordBatch) -> DataFusionResult<()> {
553        let _timer = self.metric.elapsed_compute();
554        // if have multi time series or empty batch, scalar will return NaN
555        if self.have_multi_series || batch.num_rows() == 0 {
556            return Ok(());
557        }
558        // fast path: no tag columns means all data belongs to the same series.
559        if self.tag_indices.is_empty() {
560            self.append_batch(batch)?;
561            return Ok(());
562        }
563        let all_same = |val: Option<&str>, array: &ArrayRef| -> bool {
564            (0..array.len()).all(|i| string_array_value_at_index(array, i) == val)
565        };
566        // assert the entire batch belong to the same series
567        let all_tag_columns_same = if let Some(tags) = &self.tag_value {
568            tags.iter()
569                .zip(self.tag_indices.iter())
570                .all(|(value, index)| {
571                    let array = batch.column(*index);
572                    all_same(value.as_deref(), array)
573                })
574        } else {
575            let mut tag_values = Vec::with_capacity(self.tag_indices.len());
576            let is_same = self.tag_indices.iter().all(|index| {
577                let array = batch.column(*index);
578                let value = string_array_value_at_index(array, 0).map(str::to_string);
579                let is_same = all_same(value.as_deref(), array);
580                tag_values.push(value);
581                is_same
582            });
583            self.tag_value = Some(tag_values);
584            is_same
585        };
586        if all_tag_columns_same {
587            self.append_batch(batch)?;
588        } else {
589            self.have_multi_series = true;
590        }
591        Ok(())
592    }
593
594    fn append_batch(&mut self, input_batch: RecordBatch) -> DataFusionResult<()> {
595        let ts_column = input_batch.column(self.project_index.0).clone();
596        let val_column = cast_with_options(
597            input_batch.column(self.project_index.1),
598            &DataType::Float64,
599            &CastOptions::default(),
600        )?;
601        let input_batch = RecordBatch::try_new(self.schema.clone(), vec![ts_column, val_column])?;
602        if let Some(batch) = &self.batch {
603            self.batch = Some(concat_batches(&self.schema, vec![batch, &input_batch])?);
604        } else {
605            self.batch = Some(input_batch);
606        }
607        Ok(())
608    }
609}
610
611impl Stream for ScalarCalculateStream {
612    type Item = DataFusionResult<RecordBatch>;
613
614    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
615        loop {
616            if self.done {
617                return Poll::Ready(None);
618            }
619            match ready!(self.input.poll_next_unpin(cx)) {
620                Some(Ok(batch)) => {
621                    self.update_batch(batch)?;
622                }
623                // inner had error, return to caller
624                Some(Err(e)) => return Poll::Ready(Some(Err(e))),
625                // inner is done, producing output
626                None => {
627                    self.done = true;
628                    return match self.batch.take() {
629                        Some(batch) if !self.have_multi_series => {
630                            self.metric.record_output(batch.num_rows());
631                            Poll::Ready(Some(Ok(batch)))
632                        }
633                        _ => {
634                            let time_array = (self.start..=self.end)
635                                .step_by(self.interval as _)
636                                .collect::<Vec<_>>();
637                            let nums = time_array.len();
638                            let nan_batch = RecordBatch::try_new(
639                                self.schema.clone(),
640                                vec![
641                                    Arc::new(TimestampMillisecondArray::from(time_array)),
642                                    Arc::new(Float64Array::from(vec![f64::NAN; nums])),
643                                ],
644                            )?;
645                            self.metric.record_output(nan_batch.num_rows());
646                            Poll::Ready(Some(Ok(nan_batch)))
647                        }
648                    };
649                }
650            };
651        }
652    }
653}
654
655#[cfg(test)]
656mod test {
657    use datafusion::arrow::datatypes::{DataType, Field, Schema};
658    use datafusion::datasource::memory::MemorySourceConfig;
659    use datafusion::datasource::source::DataSourceExec;
660    use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType};
661    use datafusion::prelude::SessionContext;
662    use datatypes::arrow::array::{
663        ArrayRef, DictionaryArray, Float64Array, StringArray, TimestampMillisecondArray,
664        UInt32Array,
665    };
666    use datatypes::arrow::datatypes::{TimeUnit, UInt32Type};
667
668    use super::*;
669
670    fn project_batch(batch: &RecordBatch, indices: &[usize]) -> RecordBatch {
671        let fields = indices
672            .iter()
673            .map(|&idx| batch.schema().field(idx).clone())
674            .collect::<Vec<_>>();
675        let columns = indices
676            .iter()
677            .map(|&idx| batch.column(idx).clone())
678            .collect::<Vec<_>>();
679        let schema = Arc::new(Schema::new(fields));
680        RecordBatch::try_new(schema, columns).unwrap()
681    }
682
683    #[test]
684    fn necessary_children_exprs_preserve_tag_columns() {
685        let schema = Arc::new(Schema::new(vec![
686            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
687            Field::new("tag1", DataType::Utf8, true),
688            Field::new("tag2", DataType::Utf8, true),
689            Field::new("val", DataType::Float64, true),
690            Field::new("extra", DataType::Utf8, true),
691        ]));
692        let schema = Arc::new(DFSchema::try_from(schema).unwrap());
693        let input = LogicalPlan::EmptyRelation(EmptyRelation {
694            produce_one_row: false,
695            schema,
696        });
697        let tag_columns = vec!["tag1".to_string(), "tag2".to_string()];
698        let plan = ScalarCalculate::new(0, 1, 1, input, "ts", &tag_columns, "val", None).unwrap();
699
700        let required = plan.necessary_children_exprs(&[0, 1]).unwrap();
701        assert_eq!(required, vec![vec![0, 1, 2, 3]]);
702    }
703
704    #[tokio::test]
705    async fn pruning_should_keep_tag_columns_for_exec() {
706        let schema = Arc::new(Schema::new(vec![
707            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
708            Field::new("tag1", DataType::Utf8, true),
709            Field::new("tag2", DataType::Utf8, true),
710            Field::new("val", DataType::Float64, true),
711            Field::new("extra", DataType::Utf8, true),
712        ]));
713        let df_schema = Arc::new(DFSchema::try_from(schema.clone()).unwrap());
714        let input = LogicalPlan::EmptyRelation(EmptyRelation {
715            produce_one_row: false,
716            schema: df_schema,
717        });
718        let tag_columns = vec!["tag1".to_string(), "tag2".to_string()];
719        let plan =
720            ScalarCalculate::new(0, 15_000, 5000, input, "ts", &tag_columns, "val", None).unwrap();
721
722        let required = plan.necessary_children_exprs(&[0, 1]).unwrap();
723        let required = &required[0];
724
725        let batch = RecordBatch::try_new(
726            schema,
727            vec![
728                Arc::new(TimestampMillisecondArray::from(vec![
729                    0, 5_000, 10_000, 15_000,
730                ])),
731                Arc::new(StringArray::from(vec!["foo", "foo", "foo", "foo"])),
732                Arc::new(StringArray::from(vec!["bar", "bar", "bar", "bar"])),
733                Arc::new(Float64Array::from(vec![1.0, 2.0, 3.0, 4.0])),
734                Arc::new(StringArray::from(vec!["x", "x", "x", "x"])),
735            ],
736        )
737        .unwrap();
738
739        let projected_batch = project_batch(&batch, required);
740        let projected_schema = projected_batch.schema();
741        let memory_exec = Arc::new(DataSourceExec::new(Arc::new(
742            MemorySourceConfig::try_new(&[vec![projected_batch]], projected_schema, None).unwrap(),
743        )));
744        let scalar_exec = plan.to_execution_plan(memory_exec).unwrap();
745
746        let session_context = SessionContext::default();
747        let result = datafusion::physical_plan::collect(scalar_exec, session_context.task_ctx())
748            .await
749            .unwrap();
750
751        assert_eq!(result.len(), 1);
752        let batch = &result[0];
753        assert_eq!(batch.num_columns(), 2);
754        assert_eq!(batch.num_rows(), 4);
755        assert_eq!(batch.schema().field(0).name(), "ts");
756        assert_eq!(batch.schema().field(1).name(), "scalar(val)");
757
758        let ts = batch
759            .column(0)
760            .as_any()
761            .downcast_ref::<TimestampMillisecondArray>()
762            .unwrap();
763        assert_eq!(ts.values(), &[0i64, 5_000, 10_000, 15_000]);
764
765        let values = batch
766            .column(1)
767            .as_any()
768            .downcast_ref::<Float64Array>()
769            .unwrap();
770        assert_eq!(values.values(), &[1.0f64, 2.0, 3.0, 4.0]);
771    }
772
773    fn prepare_test_data(series: Vec<RecordBatch>) -> DataSourceExec {
774        let schema = series.first().unwrap().schema();
775        DataSourceExec::new(Arc::new(
776            MemorySourceConfig::try_new(&[series], schema, None).unwrap(),
777        ))
778    }
779
780    fn dictionary(values: &[&str], keys: Vec<u32>) -> ArrayRef {
781        Arc::new(DictionaryArray::<UInt32Type>::new(
782            UInt32Array::from(keys),
783            Arc::new(StringArray::from(values.to_vec())),
784        ))
785    }
786
787    async fn run_test(series: Vec<RecordBatch>, expected: &str) {
788        let memory_exec = Arc::new(prepare_test_data(series));
789        let schema = Arc::new(Schema::new(vec![
790            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
791            Field::new("val", DataType::Float64, true),
792        ]));
793        let properties = Arc::new(PlanProperties::new(
794            EquivalenceProperties::new(schema.clone()),
795            Partitioning::UnknownPartitioning(1),
796            EmissionType::Incremental,
797            Boundedness::Bounded,
798        ));
799        let scalar_exec = Arc::new(ScalarCalculateExec {
800            start: 0,
801            end: 15_000,
802            interval: 5000,
803            tag_columns: vec!["tag1".to_string(), "tag2".to_string()],
804            input: memory_exec,
805            schema,
806            project_index: (0, 3),
807            metric: ExecutionPlanMetricsSet::new(),
808            properties,
809        });
810        let session_context = SessionContext::default();
811        let result = datafusion::physical_plan::collect(scalar_exec, session_context.task_ctx())
812            .await
813            .unwrap();
814        let result_literal = datatypes::arrow::util::pretty::pretty_format_batches(&result)
815            .unwrap()
816            .to_string();
817        assert_eq!(result_literal, expected);
818    }
819
820    #[tokio::test]
821    async fn same_series() {
822        let schema = Arc::new(Schema::new(vec![
823            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
824            Field::new("tag1", DataType::Utf8, true),
825            Field::new("tag2", DataType::Utf8, true),
826            Field::new("val", DataType::Float64, true),
827        ]));
828        run_test(
829            vec![
830                RecordBatch::try_new(
831                    schema.clone(),
832                    vec![
833                        Arc::new(TimestampMillisecondArray::from(vec![0, 5_000])),
834                        Arc::new(StringArray::from(vec!["foo", "foo"])),
835                        Arc::new(StringArray::from(vec!["🥺", "🥺"])),
836                        Arc::new(Float64Array::from(vec![1.0, 2.0])),
837                    ],
838                )
839                .unwrap(),
840                RecordBatch::try_new(
841                    schema,
842                    vec![
843                        Arc::new(TimestampMillisecondArray::from(vec![10_000, 15_000])),
844                        Arc::new(StringArray::from(vec!["foo", "foo"])),
845                        Arc::new(StringArray::from(vec!["🥺", "🥺"])),
846                        Arc::new(Float64Array::from(vec![3.0, 4.0])),
847                    ],
848                )
849                .unwrap(),
850            ],
851            "+---------------------+-----+\
852            \n| ts                  | val |\
853            \n+---------------------+-----+\
854            \n| 1970-01-01T00:00:00 | 1.0 |\
855            \n| 1970-01-01T00:00:05 | 2.0 |\
856            \n| 1970-01-01T00:00:10 | 3.0 |\
857            \n| 1970-01-01T00:00:15 | 4.0 |\
858            \n+---------------------+-----+",
859        )
860        .await
861    }
862
863    #[tokio::test]
864    async fn same_series_with_dictionary_tags() {
865        let dictionary_type =
866            DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8));
867        let schema = Arc::new(Schema::new(vec![
868            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
869            Field::new("tag1", dictionary_type.clone(), true),
870            Field::new("tag2", dictionary_type, true),
871            Field::new("val", DataType::Float64, true),
872        ]));
873        run_test(
874            vec![
875                RecordBatch::try_new(
876                    schema.clone(),
877                    vec![
878                        Arc::new(TimestampMillisecondArray::from(vec![0, 5_000])),
879                        dictionary(&["foo"], vec![0, 0]),
880                        dictionary(&["unused", "bar"], vec![1, 1]),
881                        Arc::new(Float64Array::from(vec![1.0, 2.0])),
882                    ],
883                )
884                .unwrap(),
885                RecordBatch::try_new(
886                    schema,
887                    vec![
888                        Arc::new(TimestampMillisecondArray::from(vec![10_000, 15_000])),
889                        dictionary(&["other", "foo"], vec![1, 1]),
890                        dictionary(&["bar"], vec![0, 0]),
891                        Arc::new(Float64Array::from(vec![3.0, 4.0])),
892                    ],
893                )
894                .unwrap(),
895            ],
896            "+---------------------+-----+\
897            \n| ts                  | val |\
898            \n+---------------------+-----+\
899            \n| 1970-01-01T00:00:00 | 1.0 |\
900            \n| 1970-01-01T00:00:05 | 2.0 |\
901            \n| 1970-01-01T00:00:10 | 3.0 |\
902            \n| 1970-01-01T00:00:15 | 4.0 |\
903            \n+---------------------+-----+",
904        )
905        .await
906    }
907
908    #[tokio::test]
909    async fn diff_series() {
910        let schema = Arc::new(Schema::new(vec![
911            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
912            Field::new("tag1", DataType::Utf8, true),
913            Field::new("tag2", DataType::Utf8, true),
914            Field::new("val", DataType::Float64, true),
915        ]));
916        run_test(
917            vec![
918                RecordBatch::try_new(
919                    schema.clone(),
920                    vec![
921                        Arc::new(TimestampMillisecondArray::from(vec![0, 5_000])),
922                        Arc::new(StringArray::from(vec!["foo", "foo"])),
923                        Arc::new(StringArray::from(vec!["🥺", "🥺"])),
924                        Arc::new(Float64Array::from(vec![1.0, 2.0])),
925                    ],
926                )
927                .unwrap(),
928                RecordBatch::try_new(
929                    schema,
930                    vec![
931                        Arc::new(TimestampMillisecondArray::from(vec![10_000, 15_000])),
932                        Arc::new(StringArray::from(vec!["foo", "foo"])),
933                        Arc::new(StringArray::from(vec!["🥺", "😝"])),
934                        Arc::new(Float64Array::from(vec![3.0, 4.0])),
935                    ],
936                )
937                .unwrap(),
938            ],
939            "+---------------------+-----+\
940            \n| ts                  | val |\
941            \n+---------------------+-----+\
942            \n| 1970-01-01T00:00:00 | NaN |\
943            \n| 1970-01-01T00:00:05 | NaN |\
944            \n| 1970-01-01T00:00:10 | NaN |\
945            \n| 1970-01-01T00:00:15 | NaN |\
946            \n+---------------------+-----+",
947        )
948        .await
949    }
950
951    #[tokio::test]
952    async fn empty_series() {
953        let schema = Arc::new(Schema::new(vec![
954            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
955            Field::new("tag1", DataType::Utf8, true),
956            Field::new("tag2", DataType::Utf8, true),
957            Field::new("val", DataType::Float64, true),
958        ]));
959        run_test(
960            vec![
961                RecordBatch::try_new(
962                    schema,
963                    vec![
964                        Arc::new(TimestampMillisecondArray::new_null(0)),
965                        Arc::new(StringArray::new_null(0)),
966                        Arc::new(StringArray::new_null(0)),
967                        Arc::new(Float64Array::new_null(0)),
968                    ],
969                )
970                .unwrap(),
971            ],
972            "+---------------------+-----+\
973            \n| ts                  | val |\
974            \n+---------------------+-----+\
975            \n| 1970-01-01T00:00:00 | NaN |\
976            \n| 1970-01-01T00:00:05 | NaN |\
977            \n| 1970-01-01T00:00:10 | NaN |\
978            \n| 1970-01-01T00:00:15 | NaN |\
979            \n+---------------------+-----+",
980        )
981        .await
982    }
983}