diff --git a/src/query/src/optimizer/count_wildcard.rs b/src/query/src/optimizer/count_wildcard.rs index affdaa18f7..1c6a5b814b 100644 --- a/src/query/src/optimizer/count_wildcard.rs +++ b/src/query/src/optimizer/count_wildcard.rs @@ -85,20 +85,23 @@ impl CountWildcardToTimeIndexRule { plan.visit(&mut finder).unwrap(); let col = finder.into_column(); - // check if the time index is a valid column as for current plan + // The resolved time index must be present and non-nullable in the + // immediate input schema. Schema-changing nodes can otherwise expose + // a nullable field with the same name as the source time index, and + // `count()` would then count fewer rows than `count(*)`. if let Some(col) = &col { - let mut is_valid = false; // if more than one input, we give up and just use `count(1)` if plan.inputs().len() > 1 { return None; } - for input in plan.inputs() { - if input.schema().has_column(col) { - is_valid = true; - break; - } - } - if !is_valid { + // The guard above guarantees exactly one input here, so checking + // the first input is equivalent to checking all inputs as the rule + // used to: a plan with zero inputs also falls back to `count(1)`. + let input = plan.inputs().first().copied()?; + let Ok((_, field)) = input.schema().qualified_field_from_column(col) else { + return None; + }; + if field.is_nullable() { return None; } } @@ -147,7 +150,8 @@ impl TreeNodeVisitor<'_> for TimeIndexFinder { fn f_down(&mut self, node: &Self::Node) -> DataFusionResult { if let LogicalPlan::SubqueryAlias(subquery_alias) = node { - self.table_alias = Some(subquery_alias.alias.clone()); + self.table_alias + .get_or_insert_with(|| subquery_alias.alias.clone()); } if let LogicalPlan::TableScan(table_scan) = &node @@ -199,17 +203,20 @@ mod test { use common_catalog::consts::DEFAULT_CATALOG_NAME; use common_error::ext::{BoxedError, ErrorExt, StackError}; use common_error::status_code::StatusCode; - use common_recordbatch::SendableRecordBatchStream; + use common_recordbatch::{RecordBatch, SendableRecordBatchStream}; use datafusion::functions_aggregate::count::count_all; + use datafusion::functions_aggregate::min_max::max; use datafusion_common::Column; use datafusion_expr::LogicalPlanBuilder; use datafusion_sql::TableReference; use datatypes::data_type::ConcreteDataType; - use datatypes::schema::{ColumnSchema, SchemaBuilder}; + use datatypes::schema::{ColumnSchema, Schema, SchemaBuilder}; + use datatypes::vectors::{Int64Vector, TimestampMillisecondVector, VectorRef}; use store_api::data_source::DataSource; use store_api::storage::ScanRequest; use table::metadata::{FilterPushDownType, TableInfoBuilder, TableMetaBuilder, TableType}; use table::table::numbers::NumbersTable; + use table::test_util::MemTable; use table::{Table, TableRef}; use super::*; @@ -319,6 +326,216 @@ mod test { ); } + #[test] + fn count_wildcard_shape_matrix() { + let config = datafusion::config::ConfigOptions::default(); + + let direct = CountWildcardToTimeIndexRule + .analyze(count_star(source_plan("source")), &config) + .unwrap(); + assert_count_argument_column(&direct, "source", "ts"); + + let simple_alias = count_star( + LogicalPlanBuilder::from(source_plan("source")) + .alias("projected") + .unwrap() + .build() + .unwrap(), + ); + let simple_alias = CountWildcardToTimeIndexRule + .analyze(simple_alias, &config) + .unwrap(); + assert_count_argument_column(&simple_alias, "projected", "ts"); + + let nested_alias = count_star( + LogicalPlanBuilder::from(source_plan("source")) + .alias("inner") + .unwrap() + .alias("outer") + .unwrap() + .build() + .unwrap(), + ); + let nested_alias = CountWildcardToTimeIndexRule + .analyze(nested_alias, &config) + .unwrap(); + assert_count_argument_column(&nested_alias, "outer", "ts"); + + let nested_rename = count_star( + LogicalPlanBuilder::from(source_plan("source")) + .project(vec![col("ts").alias("renamed")]) + .unwrap() + .alias("projected") + .unwrap() + .build() + .unwrap(), + ); + let nested_rename = CountWildcardToTimeIndexRule + .analyze(nested_rename, &config) + .unwrap(); + assert_count_argument_literal_one(&nested_rename); + + let nested_rename_with_payload_reorder = count_star( + LogicalPlanBuilder::from(source_plan("source")) + .project(vec![col("payload"), col("ts").alias("renamed")]) + .unwrap() + .alias("projected") + .unwrap() + .build() + .unwrap(), + ); + let nested_rename_with_payload_reorder = CountWildcardToTimeIndexRule + .analyze(nested_rename_with_payload_reorder, &config) + .unwrap(); + assert_count_argument_literal_one(&nested_rename_with_payload_reorder); + + let multi_input = count_star( + LogicalPlanBuilder::from(source_plan("left")) + .cross_join(source_plan("right")) + .unwrap() + .build() + .unwrap(), + ); + let multi_input = CountWildcardToTimeIndexRule + .analyze(multi_input, &config) + .unwrap(); + assert_count_argument_literal_one(&multi_input); + } + + #[test] + fn projection_name_collision_falls_back_to_literal_one() { + let before = count_star( + LogicalPlanBuilder::from(source_plan("source")) + .project(vec![col("payload").alias("ts")]) + .unwrap() + .alias("projected") + .unwrap() + .build() + .unwrap(), + ); + + let aggregate = aggregate_plan(&before); + let field = aggregate + .input + .schema() + .qualified_field_with_name(Some(&TableReference::bare("projected")), "ts") + .unwrap(); + assert!(field.1.is_nullable()); + + let after = CountWildcardToTimeIndexRule + .analyze(before, &datafusion::config::ConfigOptions::default()) + .unwrap(); + assert_count_argument_literal_one(&after); + } + + #[test] + fn inner_aggregate_nullable_time_index_name_falls_back_to_literal_one() { + let before = count_star( + LogicalPlanBuilder::from(source_plan("source")) + .aggregate(Vec::::new(), vec![max(col("payload")).alias("ts")]) + .unwrap() + .alias("aggregated") + .unwrap() + .build() + .unwrap(), + ); + + let aggregate = aggregate_plan(&before); + let field = aggregate + .input + .schema() + .qualified_field_with_name(Some(&TableReference::bare("aggregated")), "ts") + .unwrap(); + assert!(field.1.is_nullable()); + + let after = CountWildcardToTimeIndexRule + .analyze(before, &datafusion::config::ConfigOptions::default()) + .unwrap(); + assert_count_argument_literal_one(&after); + } + + fn source_plan(table_name: &str) -> LogicalPlan { + let schema = Arc::new(Schema::new(vec![ + ColumnSchema::new( + "ts", + ConcreteDataType::timestamp_millisecond_datatype(), + false, + ) + .with_time_index(true), + ColumnSchema::new("payload", ConcreteDataType::int64_datatype(), true), + ])); + let columns: Vec = vec![ + Arc::new(TimestampMillisecondVector::from_slice([1, 2, 3])), + Arc::new(Int64Vector::from(vec![Some(10), None, Some(30)])), + ]; + let table = MemTable::table( + table_name, + RecordBatch::new(schema, columns).expect("test record batch must be valid"), + ); + let source = Arc::new(DefaultTableSource::new(Arc::new( + DfTableProviderAdapter::new(table), + ))); + LogicalPlanBuilder::scan_with_filters(table_name, source, None, vec![]) + .unwrap() + .build() + .unwrap() + } + + fn count_star(input: LogicalPlan) -> LogicalPlan { + LogicalPlanBuilder::from(input) + .aggregate(Vec::::new(), vec![count_all()]) + .unwrap() + .build() + .unwrap() + } + + fn count_aggregate(plan: &LogicalPlan) -> &AggregateFunction { + let LogicalPlan::Aggregate(aggregate) = plan else { + panic!("expected aggregate plan, got {plan:?}"); + }; + assert_eq!(1, aggregate.aggr_expr.len()); + let expr = unwrap_aliases(&aggregate.aggr_expr[0]); + let Expr::AggregateFunction(count) = expr else { + panic!("expected count aggregate, got {:?}", aggregate.aggr_expr[0]); + }; + assert_eq!("count", count.func.name()); + count + } + + fn unwrap_aliases(expr: &Expr) -> &Expr { + match expr { + Expr::Alias(alias) => unwrap_aliases(alias.expr.as_ref()), + expr => expr, + } + } + + fn assert_count_argument_column(plan: &LogicalPlan, relation: &str, name: &str) { + let count = count_aggregate(plan); + let [Expr::Column(column)] = count.params.args.as_slice() else { + panic!( + "expected one column count argument, got {:?}", + count.params.args + ); + }; + assert_eq!(Some(TableReference::bare(relation)), column.relation); + assert_eq!(name, column.name); + } + + fn assert_count_argument_literal_one(plan: &LogicalPlan) { + let count = count_aggregate(plan); + assert!(matches!( + count.params.args.as_slice(), + [Expr::Literal(ScalarValue::Int64(Some(1)), _)] + )); + } + + fn aggregate_plan(plan: &LogicalPlan) -> &datafusion_expr::logical_plan::Aggregate { + let LogicalPlan::Aggregate(aggregate) = plan else { + panic!("expected aggregate plan, got {plan:?}"); + }; + aggregate + } + fn build_time_index_table(table_name: &str, schema_name: &str, catalog_name: &str) -> TableRef { let column_schemas = vec![ ColumnSchema::new( diff --git a/tests/cases/standalone/common/tql/tql-cte.result b/tests/cases/standalone/common/tql/tql-cte.result index 62a81804fa..0547754b48 100644 --- a/tests/cases/standalone/common/tql/tql-cte.result +++ b/tests/cases/standalone/common/tql/tql-cte.result @@ -155,7 +155,7 @@ SELECT count(*) FROM filtered; +---------------+-----------------------------------------------------------------------------------------------------------------------------------------------------------+ | logical_plan | MergeScan [is_placeholder=false, remote_input=[ | | | Projection: count(Int64(1)) AS count(*) | -| | Aggregate: groupBy=[[]], aggr=[[count(Int64(1))]] | +| | Aggregate: groupBy=[[]], aggr=[[count(filtered.ts) AS count(Int64(1))]] | | | SubqueryAlias: filtered | | | Projection: tql_data.ts, tql_data.val | | | SubqueryAlias: tql_data | @@ -642,7 +642,7 @@ SELECT count(*) as high_values FROM final; +---------------+-------------------------------------------------------------------------------------------------------------------------------------------------------------------+ | logical_plan | MergeScan [is_placeholder=false, remote_input=[ | | | Projection: count(Int64(1)) AS count(*) AS high_values | -| | Aggregate: groupBy=[[]], aggr=[[count(Int64(1))]] | +| | Aggregate: groupBy=[[]], aggr=[[count(final.ts) AS count(Int64(1))]] | | | SubqueryAlias: final | | | Projection: processed.ts AS ts, processed.percent AS percent | | | Projection: processed.ts, processed.percent |