1use 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#[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 plan.visit(&mut finder).unwrap();
86 let col = finder.into_column();
87
88 if let Some(col) = &col {
93 if plan.inputs().len() > 1 {
95 return None;
96 }
97 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
113impl 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 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}