1use std::collections::HashSet;
16use std::sync::Arc;
17
18use common_function::aggrs::aggr_wrapper::{StateMergeHelper, is_all_aggr_exprs_steppable};
19use common_telemetry::debug;
20use datafusion::error::Result as DfResult;
21use datafusion_common::tree_node::{TreeNode, TreeNodeRecursion};
22use datafusion_expr::{Expr, LogicalPlan, UserDefinedLogicalNode};
23use promql::extension_plan::{
24 EmptyMetric, InstantManipulate, RangeManipulate, SeriesDivide, SeriesNormalize,
25};
26use store_api::metric_engine_consts::DATA_SCHEMA_TSID_COLUMN_NAME;
27
28use crate::dist_plan::MergeScanLogicalPlan;
29use crate::dist_plan::analyzer::AliasMapping;
30use crate::dist_plan::merge_sort::{MergeSortLogicalPlan, merge_sort_transformer};
31
32pub struct StepTransformAction {
33 extra_parent_plans: Vec<LogicalPlan>,
34 new_child_plan: Option<LogicalPlan>,
35}
36
37pub fn step_aggr_to_upper_aggr(
47 aggr_plan: &LogicalPlan,
48) -> datafusion_common::Result<StepTransformAction> {
49 let LogicalPlan::Aggregate(input_aggr) = aggr_plan else {
50 return Err(datafusion_common::DataFusionError::Plan(
51 "step_aggr_to_upper_aggr only accepts Aggregate plan".to_string(),
52 ));
53 };
54 if !is_all_aggr_exprs_steppable(&input_aggr.aggr_expr) {
55 return Err(datafusion_common::DataFusionError::NotImplemented(format!(
56 "Some aggregate expressions are not steppable in [{}]",
57 input_aggr
58 .aggr_expr
59 .iter()
60 .map(|e| e.to_string())
61 .collect::<Vec<_>>()
62 .join(", ")
63 )));
64 }
65
66 let step_aggr_plan = StateMergeHelper::split_aggr_node(input_aggr.clone())?;
67
68 let ret = StepTransformAction {
70 extra_parent_plans: vec![step_aggr_plan.upper_merge.clone()],
71 new_child_plan: Some(step_aggr_plan.lower_state.clone()),
72 };
73 Ok(ret)
74}
75
76#[allow(dead_code)]
77pub enum Commutativity {
78 Commutative,
79 PartialCommutative,
80 ConditionalCommutative(Option<Transformer>),
81 TransformedCommutative {
82 transformer: Option<StageTransformer>,
84 },
85 NonCommutative,
86 Unimplemented,
87 Unsupported,
89}
90
91pub struct Categorizer {}
92
93impl Categorizer {
94 pub fn check_plan(
95 plan: &LogicalPlan,
96 partition_cols: Option<AliasMapping>,
97 ) -> DfResult<Commutativity> {
98 if has_subquery(plan)? {
101 return Ok(Commutativity::Unimplemented);
102 }
103
104 let partition_cols = partition_cols.unwrap_or_default();
105
106 let comm = match plan {
107 LogicalPlan::Projection(proj) => {
108 for expr in &proj.expr {
109 let commutativity = Self::check_expr(expr);
110 if !matches!(commutativity, Commutativity::Commutative) {
111 return Ok(commutativity);
112 }
113 }
114 Commutativity::Commutative
115 }
116 LogicalPlan::Filter(filter) => Self::check_expr(&filter.predicate),
118 LogicalPlan::Window(_) => Commutativity::Unimplemented,
119 LogicalPlan::Aggregate(aggr) => {
120 let has_grouping_set = aggr
123 .group_expr
124 .iter()
125 .any(|expr| matches!(expr, Expr::GroupingSet(_)));
126 let is_all_steppable =
127 !has_grouping_set && is_all_aggr_exprs_steppable(&aggr.aggr_expr);
128 let matches_partition = Self::check_partition(&aggr.group_expr, &partition_cols);
129 if !matches_partition && is_all_steppable {
130 debug!("Plan is steppable: {plan}");
131 return Ok(Commutativity::TransformedCommutative {
132 transformer: Some(Arc::new(|plan: &LogicalPlan| {
133 debug!("Before Step optimize: {plan}");
134 let ret = step_aggr_to_upper_aggr(plan);
135 ret.inspect_err(|err| {
136 common_telemetry::error!("Failed to step aggregate plan: {err:?}");
137 })
138 .map(|s| TransformerAction {
139 extra_parent_plans: s.extra_parent_plans,
140 new_child_plan: s.new_child_plan,
141 })
142 })),
143 });
144 }
145 if !matches_partition {
146 return Ok(Commutativity::NonCommutative);
147 }
148 for expr in &aggr.aggr_expr {
149 let commutativity = Self::check_expr(expr);
150 if !matches!(commutativity, Commutativity::Commutative) {
151 return Ok(commutativity);
152 }
153 }
154 Commutativity::ConditionalCommutative(None)
159 }
160 LogicalPlan::Sort(_sort) => {
161 if partition_cols.is_empty() {
162 return Ok(Commutativity::Commutative);
163 }
164
165 Commutativity::ConditionalCommutative(Some(Arc::new(merge_sort_transformer)))
168 }
169 LogicalPlan::Join(_) => Commutativity::NonCommutative,
170 LogicalPlan::Repartition(_) => {
171 Commutativity::Unimplemented
173 }
174 LogicalPlan::Union(_) => Commutativity::Unimplemented,
175 LogicalPlan::TableScan(_) => Commutativity::Commutative,
176 LogicalPlan::EmptyRelation(_) => Commutativity::NonCommutative,
177 LogicalPlan::Subquery(_) => Commutativity::Unimplemented,
178 LogicalPlan::SubqueryAlias(_) => Commutativity::Commutative,
179 LogicalPlan::Limit(limit) => {
180 if partition_cols.is_empty() && limit.fetch.is_some() {
183 Commutativity::Commutative
184 } else if limit.skip.is_none() && limit.fetch.is_some() {
185 Commutativity::PartialCommutative
186 } else {
187 Commutativity::Unimplemented
188 }
189 }
190 LogicalPlan::Extension(extension) => {
191 Self::check_extension_plan(extension.node.as_ref() as _, &partition_cols)
192 }
193 LogicalPlan::Distinct(_) => {
194 if partition_cols.is_empty() {
195 Commutativity::Commutative
196 } else {
197 Commutativity::PartialCommutative
198 }
199 }
200 LogicalPlan::Unnest(_) => Commutativity::Commutative,
201 LogicalPlan::Statement(_) => Commutativity::Unsupported,
202 LogicalPlan::Values(_) => Commutativity::Unsupported,
203 LogicalPlan::Explain(_) => Commutativity::Unsupported,
204 LogicalPlan::Analyze(_) => Commutativity::Unsupported,
205 LogicalPlan::DescribeTable(_) => Commutativity::Unsupported,
206 LogicalPlan::Dml(_) => Commutativity::Unsupported,
207 LogicalPlan::Ddl(_) => Commutativity::Unsupported,
208 LogicalPlan::Copy(_) => Commutativity::Unsupported,
209 LogicalPlan::RecursiveQuery(_) => Commutativity::Unsupported,
210 };
211
212 Ok(comm)
213 }
214
215 pub fn check_extension_plan(
216 plan: &dyn UserDefinedLogicalNode,
217 partition_cols: &AliasMapping,
218 ) -> Commutativity {
219 match plan.name() {
220 name if name == SeriesDivide::name() => {
221 let series_divide = plan.as_any().downcast_ref::<SeriesDivide>().unwrap();
222 if series_divide
225 .tags()
226 .iter()
227 .any(|tag| tag == DATA_SCHEMA_TSID_COLUMN_NAME)
228 {
229 return Commutativity::Commutative;
230 }
231
232 let tags = series_divide.tags().iter().collect::<HashSet<_>>();
233
234 for all_alias in partition_cols.values() {
235 let all_alias = all_alias.iter().map(|c| &c.name).collect::<HashSet<_>>();
236 if tags.intersection(&all_alias).count() == 0 {
237 return Commutativity::NonCommutative;
238 }
239 }
240
241 Commutativity::Commutative
242 }
243 name if name == SeriesNormalize::name()
244 || name == InstantManipulate::name()
245 || name == RangeManipulate::name() =>
246 {
247 Commutativity::Commutative
250 }
251 name if name == EmptyMetric::name()
252 || name == MergeScanLogicalPlan::name()
253 || name == MergeSortLogicalPlan::name() =>
254 {
255 Commutativity::Unimplemented
256 }
257 _ => Commutativity::Unsupported,
258 }
259 }
260
261 pub fn check_expr(expr: &Expr) -> Commutativity {
262 #[allow(deprecated)]
263 match expr {
264 Expr::Column(_)
265 | Expr::ScalarVariable(_, _)
266 | Expr::Literal(_, _)
267 | Expr::BinaryExpr(_)
268 | Expr::Not(_)
269 | Expr::IsNotNull(_)
270 | Expr::IsNull(_)
271 | Expr::IsTrue(_)
272 | Expr::IsFalse(_)
273 | Expr::IsNotTrue(_)
274 | Expr::IsNotFalse(_)
275 | Expr::Negative(_)
276 | Expr::Between(_)
277 | Expr::Exists(_)
278 | Expr::InList(_)
279 | Expr::Case(_) => Commutativity::Commutative,
280 Expr::ScalarFunction(_) => Commutativity::Commutative,
284 Expr::AggregateFunction(_udaf) => Commutativity::Commutative,
285
286 Expr::Like(_)
287 | Expr::SimilarTo(_)
288 | Expr::IsUnknown(_)
289 | Expr::IsNotUnknown(_)
290 | Expr::WindowFunction(_)
291 | Expr::InSubquery(_)
292 | Expr::ScalarSubquery(_)
293 | Expr::HigherOrderFunction(_)
294 | Expr::Lambda(_)
295 | Expr::LambdaVariable(_)
296 | Expr::Wildcard { .. } => Commutativity::Unimplemented,
297
298 Expr::Alias(alias) => Self::check_expr(&alias.expr),
299 Expr::Cast(cast) => Self::check_expr(&cast.expr),
300 Expr::TryCast(try_cast) => Self::check_expr(&try_cast.expr),
301
302 Expr::Unnest(_)
303 | Expr::GroupingSet(_)
304 | Expr::Placeholder(_)
305 | Expr::OuterReferenceColumn(_, _)
306 | Expr::SetComparison(_) => Commutativity::Unimplemented,
307 }
308 }
309
310 fn check_partition(exprs: &[Expr], partition_cols: &AliasMapping) -> bool {
320 let group_cols = exprs
321 .iter()
322 .filter_map(|expr| {
323 let mut expr = expr;
324 while let Expr::Alias(alias) = expr {
325 expr = &alias.expr;
326 }
327 match expr {
328 Expr::Column(column) => Some(column.name.clone()),
329 _ => None,
330 }
331 })
332 .collect::<HashSet<_>>();
333 for all_alias in partition_cols.values() {
334 let all_alias = all_alias
335 .iter()
336 .map(|c| c.name.clone())
337 .collect::<HashSet<_>>();
338 if group_cols.intersection(&all_alias).count() == 0 {
341 return false;
342 }
343 }
344
345 true
346 }
347}
348
349#[cfg(test)]
350mod tests {
351 use std::collections::{BTreeMap, BTreeSet};
352
353 use datafusion_common::Column;
354 use datafusion_expr::LogicalPlanBuilder;
355 use datafusion_expr::expr::ScalarFunction;
356 use datafusion_functions::core::coalesce;
357 use promql::functions::{
358 NativeHistogramDrop, NativeHistogramFraction, NativeHistogramQuantile,
359 };
360
361 use super::*;
362
363 #[test]
364 fn series_divide_by_tsid_is_commutative() {
365 let input = LogicalPlanBuilder::empty(false).build().unwrap();
366 let series_divide = SeriesDivide::new(
367 vec![DATA_SCHEMA_TSID_COLUMN_NAME.to_string()],
368 "ts".to_string(),
369 input,
370 );
371
372 let partition_cols: AliasMapping = BTreeMap::from([(
373 "some_partition_col".to_string(),
374 BTreeSet::from([Column::from_name("some_partition_col")]),
375 )]);
376
377 let commutativity = Categorizer::check_extension_plan(&series_divide, &partition_cols);
378 assert!(matches!(commutativity, Commutativity::Commutative));
379 }
380
381 #[test]
382 fn annotated_histogram_helpers_do_not_block_pushdown() {
383 for udf in [
384 NativeHistogramQuantile::scalar_udf(),
385 NativeHistogramFraction::scalar_udf(),
386 NativeHistogramDrop::warning_bool_false_udf(String::new(), None),
387 ] {
388 let helper = Expr::ScalarFunction(ScalarFunction::new_udf(Arc::new(udf), vec![]));
389 let expr = Expr::ScalarFunction(ScalarFunction::new_udf(coalesce(), vec![helper]));
390 assert!(matches!(
391 Categorizer::check_expr(&expr),
392 Commutativity::Commutative
393 ));
394 }
395 }
396}
397
398pub type Transformer = Arc<dyn Fn(&LogicalPlan) -> Option<LogicalPlan>>;
399
400pub type StageTransformer = Arc<dyn Fn(&LogicalPlan) -> DfResult<TransformerAction>>;
402
403pub struct TransformerAction {
405 pub extra_parent_plans: Vec<LogicalPlan>,
414 pub new_child_plan: Option<LogicalPlan>,
416}
417
418pub fn partial_commutative_transformer(plan: &LogicalPlan) -> Option<LogicalPlan> {
419 Some(plan.clone())
420}
421
422fn has_subquery(plan: &LogicalPlan) -> DfResult<bool> {
423 let mut found = false;
424 plan.apply_expressions(|e| {
425 e.apply(|x| {
426 if matches!(
427 x,
428 Expr::Exists(_) | Expr::InSubquery(_) | Expr::ScalarSubquery(_)
429 ) {
430 found = true;
431 Ok(TreeNodeRecursion::Stop)
432 } else {
433 Ok(TreeNodeRecursion::Continue)
434 }
435 })
436 })?;
437 Ok(found)
438}
439
440#[cfg(test)]
441mod test {
442 use datafusion_expr::{LogicalPlanBuilder, Sort};
443
444 use super::*;
445
446 #[test]
447 fn sort_on_empty_partition() {
448 let plan = LogicalPlan::Sort(Sort {
449 expr: vec![],
450 input: Arc::new(LogicalPlanBuilder::empty(false).build().unwrap()),
451 fetch: None,
452 });
453 assert!(matches!(
454 Categorizer::check_plan(&plan, Some(Default::default())).unwrap(),
455 Commutativity::Commutative
456 ));
457 }
458}