Skip to main content

query/optimizer/
count_wildcard.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 datafusion::datasource::DefaultTableSource;
16use datafusion_common::tree_node::{
17    Transformed, TransformedResult, TreeNode, TreeNodeRecursion, TreeNodeVisitor,
18};
19use datafusion_common::{Column, Result as DataFusionResult, ScalarValue};
20use datafusion_expr::expr::{AggregateFunction, WindowFunction};
21use datafusion_expr::utils::COUNT_STAR_EXPANSION;
22use datafusion_expr::{Expr, LogicalPlan, WindowFunctionDefinition, col, lit};
23use datafusion_optimizer::AnalyzerRule;
24use datafusion_optimizer::utils::NamePreserver;
25use datafusion_sql::TableReference;
26use table::table::adapter::DfTableProviderAdapter;
27
28/// A replacement to DataFusion's [`CountWildcardRule`]. This rule
29/// would prefer to use TIME INDEX for counting wildcard as it's
30/// faster to read comparing to PRIMARY KEYs.
31///
32/// [`CountWildcardRule`]: datafusion::optimizer::analyzer::CountWildcardRule
33#[derive(Debug)]
34pub struct CountWildcardToTimeIndexRule;
35
36impl AnalyzerRule for CountWildcardToTimeIndexRule {
37    fn name(&self) -> &str {
38        "count_wildcard_to_time_index_rule"
39    }
40
41    fn analyze(
42        &self,
43        plan: LogicalPlan,
44        _config: &datafusion::config::ConfigOptions,
45    ) -> DataFusionResult<LogicalPlan> {
46        plan.transform_down_with_subqueries(&Self::analyze_internal)
47            .data()
48    }
49}
50
51impl CountWildcardToTimeIndexRule {
52    fn analyze_internal(plan: LogicalPlan) -> DataFusionResult<Transformed<LogicalPlan>> {
53        let name_preserver = NamePreserver::new(&plan);
54        let new_arg = if let Some(time_index) = Self::try_find_time_index_col(&plan) {
55            vec![col(time_index)]
56        } else {
57            vec![lit(COUNT_STAR_EXPANSION)]
58        };
59        plan.map_expressions(|expr| {
60            let original_name = name_preserver.save(&expr);
61            let transformed_expr = expr.transform_up(|expr| match expr {
62                Expr::WindowFunction(mut window_function)
63                    if Self::is_count_star_window_aggregate(&window_function) =>
64                {
65                    window_function.params.args.clone_from(&new_arg);
66                    Ok(Transformed::yes(Expr::WindowFunction(window_function)))
67                }
68                Expr::AggregateFunction(mut aggregate_function)
69                    if Self::is_count_star_aggregate(&aggregate_function) =>
70                {
71                    aggregate_function.params.args.clone_from(&new_arg);
72                    Ok(Transformed::yes(Expr::AggregateFunction(
73                        aggregate_function,
74                    )))
75                }
76                _ => Ok(Transformed::no(expr)),
77            })?;
78            Ok(transformed_expr.update_data(|data| original_name.restore(data)))
79        })
80    }
81
82    fn try_find_time_index_col(plan: &LogicalPlan) -> Option<Column> {
83        let mut finder = TimeIndexFinder::default();
84        // Safety: `TimeIndexFinder` won't throw error.
85        plan.visit(&mut finder).unwrap();
86        let col = finder.into_column();
87
88        // The resolved time index must be present and non-nullable in the
89        // immediate input schema. Schema-changing nodes can otherwise expose
90        // a nullable field with the same name as the source time index, and
91        // `count(<col>)` would then count fewer rows than `count(*)`.
92        if let Some(col) = &col {
93            // if more than one input, we give up and just use `count(1)`
94            if plan.inputs().len() > 1 {
95                return None;
96            }
97            // The guard above guarantees exactly one input here, so checking
98            // the first input is equivalent to checking all inputs as the rule
99            // used to: a plan with zero inputs also falls back to `count(1)`.
100            let input = plan.inputs().first().copied()?;
101            let Ok((_, field)) = input.schema().qualified_field_from_column(col) else {
102                return None;
103            };
104            if field.is_nullable() {
105                return None;
106            }
107        }
108
109        col
110    }
111}
112
113/// Utility functions from the original rule.
114impl CountWildcardToTimeIndexRule {
115    #[expect(deprecated)]
116    fn args_at_most_wildcard_or_literal_one(args: &[Expr]) -> bool {
117        match args {
118            [] => true,
119            [Expr::Literal(ScalarValue::Int64(Some(v)), _)] => *v == 1,
120            [Expr::Wildcard { .. }] => true,
121            _ => false,
122        }
123    }
124
125    fn is_count_star_aggregate(aggregate_function: &AggregateFunction) -> bool {
126        let args = &aggregate_function.params.args;
127        matches!(aggregate_function,
128            AggregateFunction {
129                func,
130                ..
131            } if func.name() == "count" && Self::args_at_most_wildcard_or_literal_one(args))
132    }
133
134    fn is_count_star_window_aggregate(window_function: &WindowFunction) -> bool {
135        let args = &window_function.params.args;
136        matches!(window_function.fun,
137                WindowFunctionDefinition::AggregateUDF(ref udaf)
138                    if udaf.name() == "count" && Self::args_at_most_wildcard_or_literal_one(args))
139    }
140}
141
142#[derive(Default)]
143struct TimeIndexFinder {
144    time_index_col: Option<String>,
145    table_alias: Option<TableReference>,
146}
147
148impl TreeNodeVisitor<'_> for TimeIndexFinder {
149    type Node = LogicalPlan;
150
151    fn f_down(&mut self, node: &Self::Node) -> DataFusionResult<TreeNodeRecursion> {
152        if let LogicalPlan::SubqueryAlias(subquery_alias) = node {
153            self.table_alias
154                .get_or_insert_with(|| subquery_alias.alias.clone());
155        }
156
157        if let LogicalPlan::TableScan(table_scan) = &node
158            && let Some(source) = table_scan
159                .source
160                .as_any()
161                .downcast_ref::<DefaultTableSource>()
162            && let Some(adapter) = source
163                .table_provider
164                .as_any()
165                .downcast_ref::<DfTableProviderAdapter>()
166        {
167            let table_info = adapter.table().table_info();
168            self.table_alias
169                .get_or_insert(table_scan.table_name.clone());
170            self.time_index_col = table_info
171                .meta
172                .schema
173                .timestamp_column()
174                .map(|c| c.name.clone());
175
176            return Ok(TreeNodeRecursion::Stop);
177        }
178
179        if node.inputs().len() > 1 {
180            // if more than one input, we give up and just use `count(1)`
181            return Ok(TreeNodeRecursion::Stop);
182        }
183
184        Ok(TreeNodeRecursion::Continue)
185    }
186
187    fn f_up(&mut self, _node: &Self::Node) -> DataFusionResult<TreeNodeRecursion> {
188        Ok(TreeNodeRecursion::Stop)
189    }
190}
191
192impl TimeIndexFinder {
193    fn into_column(self) -> Option<Column> {
194        self.time_index_col
195            .map(|c| Column::new(self.table_alias, c))
196    }
197}
198
199#[cfg(test)]
200mod test {
201    use std::sync::Arc;
202
203    use common_catalog::consts::DEFAULT_CATALOG_NAME;
204    use common_error::ext::{BoxedError, ErrorExt, StackError};
205    use common_error::status_code::StatusCode;
206    use common_recordbatch::{RecordBatch, SendableRecordBatchStream};
207    use datafusion::functions_aggregate::count::count_all;
208    use datafusion::functions_aggregate::min_max::max;
209    use datafusion_common::Column;
210    use datafusion_expr::LogicalPlanBuilder;
211    use datafusion_sql::TableReference;
212    use datatypes::data_type::ConcreteDataType;
213    use datatypes::schema::{ColumnSchema, Schema, SchemaBuilder};
214    use datatypes::vectors::{Int64Vector, TimestampMillisecondVector, VectorRef};
215    use store_api::data_source::DataSource;
216    use store_api::storage::ScanRequest;
217    use table::metadata::{FilterPushDownType, TableInfoBuilder, TableMetaBuilder, TableType};
218    use table::table::numbers::NumbersTable;
219    use table::test_util::MemTable;
220    use table::{Table, TableRef};
221
222    use super::*;
223
224    #[test]
225    fn uppercase_table_name() {
226        let numbers_table = NumbersTable::table_with_name(0, "AbCdE".to_string());
227        let table_source = Arc::new(DefaultTableSource::new(Arc::new(
228            DfTableProviderAdapter::new(numbers_table),
229        )));
230
231        let plan = LogicalPlanBuilder::scan_with_filters("t", table_source, None, vec![])
232            .unwrap()
233            .aggregate(Vec::<Expr>::new(), vec![count_all()])
234            .unwrap()
235            .alias(r#""FgHiJ""#)
236            .unwrap()
237            .build()
238            .unwrap();
239
240        let mut finder = TimeIndexFinder::default();
241        plan.visit(&mut finder).unwrap();
242
243        assert_eq!(finder.table_alias, Some(TableReference::bare("FgHiJ")));
244        assert!(finder.time_index_col.is_none());
245    }
246
247    #[test]
248    fn bare_table_name_time_index() {
249        let table_ref = TableReference::bare("multi_partitioned_test_1");
250        let table =
251            build_time_index_table("multi_partitioned_test_1", "public", DEFAULT_CATALOG_NAME);
252        let table_source = Arc::new(DefaultTableSource::new(Arc::new(
253            DfTableProviderAdapter::new(table),
254        )));
255
256        let plan =
257            LogicalPlanBuilder::scan_with_filters(table_ref.clone(), table_source, None, vec![])
258                .unwrap()
259                .aggregate(Vec::<Expr>::new(), vec![count_all()])
260                .unwrap()
261                .build()
262                .unwrap();
263
264        let time_index = CountWildcardToTimeIndexRule::try_find_time_index_col(&plan);
265        assert_eq!(
266            time_index,
267            Some(Column::new(Some(table_ref), "greptime_timestamp"))
268        );
269    }
270
271    #[test]
272    fn schema_qualified_table_name_time_index() {
273        let table_ref = TableReference::partial("telemetry_events", "multi_partitioned_test_1");
274        let table = build_time_index_table(
275            "multi_partitioned_test_1",
276            "telemetry_events",
277            DEFAULT_CATALOG_NAME,
278        );
279        let table_source = Arc::new(DefaultTableSource::new(Arc::new(
280            DfTableProviderAdapter::new(table),
281        )));
282
283        let plan =
284            LogicalPlanBuilder::scan_with_filters(table_ref.clone(), table_source, None, vec![])
285                .unwrap()
286                .aggregate(Vec::<Expr>::new(), vec![count_all()])
287                .unwrap()
288                .build()
289                .unwrap();
290
291        let time_index = CountWildcardToTimeIndexRule::try_find_time_index_col(&plan);
292        assert_eq!(
293            time_index,
294            Some(Column::new(Some(table_ref), "greptime_timestamp"))
295        );
296    }
297
298    #[test]
299    fn fully_qualified_table_name_time_index() {
300        let table_ref = TableReference::full(
301            "telemetry_catalog",
302            "telemetry_events",
303            "multi_partitioned_test_1",
304        );
305        let table = build_time_index_table(
306            "multi_partitioned_test_1",
307            "telemetry_events",
308            "telemetry_catalog",
309        );
310        let table_source = Arc::new(DefaultTableSource::new(Arc::new(
311            DfTableProviderAdapter::new(table),
312        )));
313
314        let plan =
315            LogicalPlanBuilder::scan_with_filters(table_ref.clone(), table_source, None, vec![])
316                .unwrap()
317                .aggregate(Vec::<Expr>::new(), vec![count_all()])
318                .unwrap()
319                .build()
320                .unwrap();
321
322        let time_index = CountWildcardToTimeIndexRule::try_find_time_index_col(&plan);
323        assert_eq!(
324            time_index,
325            Some(Column::new(Some(table_ref), "greptime_timestamp"))
326        );
327    }
328
329    #[test]
330    fn count_wildcard_shape_matrix() {
331        let config = datafusion::config::ConfigOptions::default();
332
333        let direct = CountWildcardToTimeIndexRule
334            .analyze(count_star(source_plan("source")), &config)
335            .unwrap();
336        assert_count_argument_column(&direct, "source", "ts");
337
338        let simple_alias = count_star(
339            LogicalPlanBuilder::from(source_plan("source"))
340                .alias("projected")
341                .unwrap()
342                .build()
343                .unwrap(),
344        );
345        let simple_alias = CountWildcardToTimeIndexRule
346            .analyze(simple_alias, &config)
347            .unwrap();
348        assert_count_argument_column(&simple_alias, "projected", "ts");
349
350        let nested_alias = count_star(
351            LogicalPlanBuilder::from(source_plan("source"))
352                .alias("inner")
353                .unwrap()
354                .alias("outer")
355                .unwrap()
356                .build()
357                .unwrap(),
358        );
359        let nested_alias = CountWildcardToTimeIndexRule
360            .analyze(nested_alias, &config)
361            .unwrap();
362        assert_count_argument_column(&nested_alias, "outer", "ts");
363
364        let nested_rename = count_star(
365            LogicalPlanBuilder::from(source_plan("source"))
366                .project(vec![col("ts").alias("renamed")])
367                .unwrap()
368                .alias("projected")
369                .unwrap()
370                .build()
371                .unwrap(),
372        );
373        let nested_rename = CountWildcardToTimeIndexRule
374            .analyze(nested_rename, &config)
375            .unwrap();
376        assert_count_argument_literal_one(&nested_rename);
377
378        let nested_rename_with_payload_reorder = count_star(
379            LogicalPlanBuilder::from(source_plan("source"))
380                .project(vec![col("payload"), col("ts").alias("renamed")])
381                .unwrap()
382                .alias("projected")
383                .unwrap()
384                .build()
385                .unwrap(),
386        );
387        let nested_rename_with_payload_reorder = CountWildcardToTimeIndexRule
388            .analyze(nested_rename_with_payload_reorder, &config)
389            .unwrap();
390        assert_count_argument_literal_one(&nested_rename_with_payload_reorder);
391
392        let multi_input = count_star(
393            LogicalPlanBuilder::from(source_plan("left"))
394                .cross_join(source_plan("right"))
395                .unwrap()
396                .build()
397                .unwrap(),
398        );
399        let multi_input = CountWildcardToTimeIndexRule
400            .analyze(multi_input, &config)
401            .unwrap();
402        assert_count_argument_literal_one(&multi_input);
403    }
404
405    #[test]
406    fn projection_name_collision_falls_back_to_literal_one() {
407        let before = count_star(
408            LogicalPlanBuilder::from(source_plan("source"))
409                .project(vec![col("payload").alias("ts")])
410                .unwrap()
411                .alias("projected")
412                .unwrap()
413                .build()
414                .unwrap(),
415        );
416
417        let aggregate = aggregate_plan(&before);
418        let field = aggregate
419            .input
420            .schema()
421            .qualified_field_with_name(Some(&TableReference::bare("projected")), "ts")
422            .unwrap();
423        assert!(field.1.is_nullable());
424
425        let after = CountWildcardToTimeIndexRule
426            .analyze(before, &datafusion::config::ConfigOptions::default())
427            .unwrap();
428        assert_count_argument_literal_one(&after);
429    }
430
431    #[test]
432    fn inner_aggregate_nullable_time_index_name_falls_back_to_literal_one() {
433        let before = count_star(
434            LogicalPlanBuilder::from(source_plan("source"))
435                .aggregate(Vec::<Expr>::new(), vec![max(col("payload")).alias("ts")])
436                .unwrap()
437                .alias("aggregated")
438                .unwrap()
439                .build()
440                .unwrap(),
441        );
442
443        let aggregate = aggregate_plan(&before);
444        let field = aggregate
445            .input
446            .schema()
447            .qualified_field_with_name(Some(&TableReference::bare("aggregated")), "ts")
448            .unwrap();
449        assert!(field.1.is_nullable());
450
451        let after = CountWildcardToTimeIndexRule
452            .analyze(before, &datafusion::config::ConfigOptions::default())
453            .unwrap();
454        assert_count_argument_literal_one(&after);
455    }
456
457    fn source_plan(table_name: &str) -> LogicalPlan {
458        let schema = Arc::new(Schema::new(vec![
459            ColumnSchema::new(
460                "ts",
461                ConcreteDataType::timestamp_millisecond_datatype(),
462                false,
463            )
464            .with_time_index(true),
465            ColumnSchema::new("payload", ConcreteDataType::int64_datatype(), true),
466        ]));
467        let columns: Vec<VectorRef> = vec![
468            Arc::new(TimestampMillisecondVector::from_slice([1, 2, 3])),
469            Arc::new(Int64Vector::from(vec![Some(10), None, Some(30)])),
470        ];
471        let table = MemTable::table(
472            table_name,
473            RecordBatch::new(schema, columns).expect("test record batch must be valid"),
474        );
475        let source = Arc::new(DefaultTableSource::new(Arc::new(
476            DfTableProviderAdapter::new(table),
477        )));
478        LogicalPlanBuilder::scan_with_filters(table_name, source, None, vec![])
479            .unwrap()
480            .build()
481            .unwrap()
482    }
483
484    fn count_star(input: LogicalPlan) -> LogicalPlan {
485        LogicalPlanBuilder::from(input)
486            .aggregate(Vec::<Expr>::new(), vec![count_all()])
487            .unwrap()
488            .build()
489            .unwrap()
490    }
491
492    fn count_aggregate(plan: &LogicalPlan) -> &AggregateFunction {
493        let LogicalPlan::Aggregate(aggregate) = plan else {
494            panic!("expected aggregate plan, got {plan:?}");
495        };
496        assert_eq!(1, aggregate.aggr_expr.len());
497        let expr = unwrap_aliases(&aggregate.aggr_expr[0]);
498        let Expr::AggregateFunction(count) = expr else {
499            panic!("expected count aggregate, got {:?}", aggregate.aggr_expr[0]);
500        };
501        assert_eq!("count", count.func.name());
502        count
503    }
504
505    fn unwrap_aliases(expr: &Expr) -> &Expr {
506        match expr {
507            Expr::Alias(alias) => unwrap_aliases(alias.expr.as_ref()),
508            expr => expr,
509        }
510    }
511
512    fn assert_count_argument_column(plan: &LogicalPlan, relation: &str, name: &str) {
513        let count = count_aggregate(plan);
514        let [Expr::Column(column)] = count.params.args.as_slice() else {
515            panic!(
516                "expected one column count argument, got {:?}",
517                count.params.args
518            );
519        };
520        assert_eq!(Some(TableReference::bare(relation)), column.relation);
521        assert_eq!(name, column.name);
522    }
523
524    fn assert_count_argument_literal_one(plan: &LogicalPlan) {
525        let count = count_aggregate(plan);
526        assert!(matches!(
527            count.params.args.as_slice(),
528            [Expr::Literal(ScalarValue::Int64(Some(1)), _)]
529        ));
530    }
531
532    fn aggregate_plan(plan: &LogicalPlan) -> &datafusion_expr::logical_plan::Aggregate {
533        let LogicalPlan::Aggregate(aggregate) = plan else {
534            panic!("expected aggregate plan, got {plan:?}");
535        };
536        aggregate
537    }
538
539    fn build_time_index_table(table_name: &str, schema_name: &str, catalog_name: &str) -> TableRef {
540        let column_schemas = vec![
541            ColumnSchema::new(
542                "greptime_timestamp",
543                ConcreteDataType::timestamp_nanosecond_datatype(),
544                false,
545            )
546            .with_time_index(true),
547        ];
548        let schema = SchemaBuilder::try_from_columns(column_schemas)
549            .unwrap()
550            .build()
551            .unwrap();
552        let meta = TableMetaBuilder::new_external_table()
553            .schema(Arc::new(schema))
554            .next_column_id(1)
555            .build()
556            .unwrap();
557        let info = TableInfoBuilder::new(table_name.to_string(), meta)
558            .table_id(1)
559            .table_version(0)
560            .catalog_name(catalog_name)
561            .schema_name(schema_name)
562            .table_type(TableType::Base)
563            .build()
564            .unwrap();
565        let data_source = Arc::new(DummyDataSource);
566        Arc::new(Table::new(
567            Arc::new(info),
568            FilterPushDownType::Unsupported,
569            data_source,
570        ))
571    }
572
573    struct DummyDataSource;
574
575    impl DataSource for DummyDataSource {
576        fn get_stream(
577            &self,
578            _request: ScanRequest,
579        ) -> Result<SendableRecordBatchStream, BoxedError> {
580            Err(BoxedError::new(DummyDataSourceError))
581        }
582    }
583
584    #[derive(Debug)]
585    struct DummyDataSourceError;
586
587    impl std::fmt::Display for DummyDataSourceError {
588        fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
589            write!(f, "dummy data source error")
590        }
591    }
592
593    impl std::error::Error for DummyDataSourceError {}
594
595    impl StackError for DummyDataSourceError {
596        fn debug_fmt(&self, _: usize, _: &mut Vec<String>) {}
597
598        fn next(&self) -> Option<&dyn StackError> {
599            None
600        }
601    }
602
603    impl ErrorExt for DummyDataSourceError {
604        fn status_code(&self) -> StatusCode {
605            StatusCode::Internal
606        }
607
608        fn as_any(&self) -> &dyn std::any::Any {
609            self
610        }
611    }
612}