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