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