Skip to main content

query/dist_plan/
commutativity.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 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
37/// generate the upper aggregation plan that will execute on the frontend.
38/// Basically a logical plan resembling the following:
39/// Projection:
40///     Aggregate:
41///
42/// from Aggregate
43///
44/// The upper Projection exists sole to make sure parent plan can recognize the output
45/// of the upper aggregation plan.
46pub 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    // TODO(discord9): remove duplication
69    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        /// Return plans from parent to child order
83        transformer: Option<StageTransformer>,
84    },
85    NonCommutative,
86    Unimplemented,
87    /// For unrelated plans like DDL
88    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        // Subquery is treated separately in `inspect_plan_with_subquery`. To avoid rewrite the
99        // "maybe rewritten" plan, stop the check here.
100        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            // TODO(ruihang): Change this to Commutative once Like is supported in substrait
117            LogicalPlan::Filter(filter) => Self::check_expr(&filter.predicate),
118            LogicalPlan::Window(_) => Commutativity::Unimplemented,
119            LogicalPlan::Aggregate(aggr) => {
120                // The state/merge split maps each group expression to one output column,
121                // which doesn't hold for grouping sets.
122                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                // all group by expressions are partition columns can push down, unless
155                // another push down(including `Limit` or `Sort`) is already in progress(which will then prevent next cond commutative node from being push down).
156                // TODO(discord9): This is a temporary solution(that works), a better description of
157                // commutativity is needed under this situation.
158                Commutativity::ConditionalCommutative(None)
159            }
160            LogicalPlan::Sort(_sort) => {
161                if partition_cols.is_empty() {
162                    return Ok(Commutativity::Commutative);
163                }
164
165                // sort plan needs to consider column priority
166                // Change Sort to MergeSort which assumes the input streams are already sorted hence can be more efficient.
167                Commutativity::ConditionalCommutative(Some(Arc::new(merge_sort_transformer)))
168            }
169            LogicalPlan::Join(_) => Commutativity::NonCommutative,
170            LogicalPlan::Repartition(_) => {
171                // unsupported? or non-commutative
172                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                // Only execute `fetch` on remote nodes.
181                // wait for https://github.com/apache/arrow-datafusion/pull/7669
182                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                // Metric engine `__tsid` uniquely identifies a time-series. Treat a series divide
223                // that keys by `__tsid` as commutative across regions so it can be pushed down.
224                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                // They should always follows Series Divide.
248                // Either all commutative or all non-commutative (which will be blocked by SeriesDivide).
249                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            // Annotation collection must not affect distribution. Keep scalar UDFs pushdownable;
281            // TODO: Preserve datanode PromQL annotations through plan decoding, return them in
282            // Flight stream metadata, and merge them on the frontend.
283            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    /// Return true if the given expr and partition cols satisfied the rule.
311    /// In this case the plan can be treated as fully commutative.
312    ///
313    /// So only if every partition column is itself one of `exprs`, return true.
314    /// Otherwise return false.
315    ///
316    /// An expression that only references a partition column, like `substr(host, 3, 1)`
317    /// or `k % 2`, doesn't count: it can put rows from different partitions into the same
318    /// group.
319    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            // check if group columns intersect with all alias of partition columns
339            // is empty, if it's empty, not all partition columns show up in `exprs`
340            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
400/// Returns transformer action that need to be applied
401pub type StageTransformer = Arc<dyn Fn(&LogicalPlan) -> DfResult<TransformerAction>>;
402
403/// The Action that a transformer should take on the plan.
404pub struct TransformerAction {
405    /// list of plans that need to be applied to parent plans, in the order of parent to child.
406    /// i.e. if this returns `[Projection, Aggregate]`, then the parent plan should be transformed to
407    /// ```ignore
408    /// Original Parent Plan:
409    ///     Projection:
410    ///         Aggregate:
411    ///             MergeScan: ...
412    /// ```
413    pub extra_parent_plans: Vec<LogicalPlan>,
414    /// new child plan, if None, use the original plan.
415    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}