Skip to main content

mz_sql/plan/
hir.rs

1// Copyright Materialize, Inc. and contributors. All rights reserved.
2//
3// Use of this software is governed by the Business Source License
4// included in the LICENSE file.
5//
6// As of the Change Date specified in that file, in accordance with
7// the Business Source License, use of this software will be governed
8// by the Apache License, Version 2.0.
9
10//! This file houses HIR, a representation of a SQL plan that is parallel to MIR, but represents
11//! an earlier phase of planning. It's structurally very similar to MIR, with some differences
12//! which are noted below. It gets turned into MIR via a call to lower().
13
14use std::collections::{BTreeMap, BTreeSet};
15use std::fmt::{Display, Formatter};
16use std::sync::Arc;
17use std::{fmt, mem};
18
19use itertools::Itertools;
20use mz_expr::virtual_syntax::{AlgExcept, Except, IR};
21use mz_expr::visit::{Visit, VisitChildren};
22use mz_expr::{CollectionPlan, Id, LetRecLimit, RowSetFinishing, func};
23// these happen to be unchanged at the moment, but there might be additions later
24use mz_expr::AggregateFunc::{FusedWindowAggregate, WindowAggregate};
25use mz_expr::func::variadic::{And, Or};
26pub use mz_expr::{
27    BinaryFunc, ColumnOrder, TableFunc, UnaryFunc, UnmaterializableFunc, VariadicFunc, WindowFrame,
28};
29use mz_ore::collections::CollectionExt;
30use mz_ore::error::ErrorExt;
31use mz_ore::str::separated;
32use mz_ore::treat_as_equal::TreatAsEqual;
33use mz_ore::{soft_assert_or_log, stack};
34use mz_repr::adt::array::ArrayDimension;
35use mz_repr::adt::numeric::NumericMaxScale;
36use mz_repr::*;
37use serde::{Deserialize, Serialize};
38
39use crate::plan::error::PlanError;
40use crate::plan::query::{
41    EXECUTE_CAST_CONTEXT, ExprContext, execute_expr_context, offset_into_value,
42};
43use crate::plan::typeconv::{self, CastContext, plan_cast};
44use crate::plan::{Params, QueryContext, QueryLifetime, StatementContext};
45
46use super::plan_utils::GroupSizeHints;
47
48#[allow(missing_debug_implementations)]
49pub struct Hir;
50
51impl IR for Hir {
52    type Relation = HirRelationExpr;
53    type Scalar = HirScalarExpr;
54}
55
56impl AlgExcept for Hir {
57    fn except(all: &bool, lhs: Self::Relation, rhs: Self::Relation) -> Self::Relation {
58        if *all {
59            let rhs = rhs.negate();
60            HirRelationExpr::union(lhs, rhs).threshold()
61        } else {
62            let lhs = lhs.distinct();
63            let rhs = rhs.distinct().negate();
64            HirRelationExpr::union(lhs, rhs).threshold()
65        }
66    }
67
68    fn un_except<'a>(expr: &'a Self::Relation) -> Option<Except<'a, Self>> {
69        let mut result = None;
70
71        use HirRelationExpr::*;
72        if let Threshold { input } = expr {
73            if let Union { base: lhs, inputs } = input.as_ref() {
74                if let [rhs] = &inputs[..] {
75                    if let Negate { input: rhs } = rhs {
76                        match (lhs.as_ref(), rhs.as_ref()) {
77                            (Distinct { input: lhs }, Distinct { input: rhs }) => {
78                                let all = false;
79                                let lhs = lhs.as_ref();
80                                let rhs = rhs.as_ref();
81                                result = Some(Except { all, lhs, rhs })
82                            }
83                            (lhs, rhs) => {
84                                let all = true;
85                                result = Some(Except { all, lhs, rhs })
86                            }
87                        }
88                    }
89                }
90            }
91        }
92
93        result
94    }
95}
96
97#[derive(
98    Debug,
99    Clone,
100    PartialEq,
101    Eq,
102    PartialOrd,
103    Ord,
104    Hash,
105    Serialize,
106    Deserialize
107)]
108/// Just like [`mz_expr::MirRelationExpr`], except where otherwise noted below.
109pub enum HirRelationExpr {
110    Constant {
111        rows: Vec<Row>,
112        typ: SqlRelationType,
113    },
114    Get {
115        id: mz_expr::Id,
116        typ: SqlRelationType,
117    },
118    /// Mutually recursive CTE
119    LetRec {
120        /// Maximum number of iterations to evaluate. If None, then there is no limit.
121        limit: Option<LetRecLimit>,
122        /// List of bindings all of which are in scope of each other.
123        bindings: Vec<(String, mz_expr::LocalId, HirRelationExpr, SqlRelationType)>,
124        /// Result of the AST node.
125        body: Box<HirRelationExpr>,
126    },
127    /// CTE
128    Let {
129        name: String,
130        /// The identifier to be used in `Get` variants to retrieve `value`.
131        id: mz_expr::LocalId,
132        /// The collection to be bound to `name`.
133        value: Box<HirRelationExpr>,
134        /// The result of the `Let`, evaluated with `name` bound to `value`.
135        body: Box<HirRelationExpr>,
136    },
137    Project {
138        input: Box<HirRelationExpr>,
139        outputs: Vec<usize>,
140    },
141    Map {
142        input: Box<HirRelationExpr>,
143        scalars: Vec<HirScalarExpr>,
144    },
145    CallTable {
146        func: TableFunc,
147        exprs: Vec<HirScalarExpr>,
148    },
149    Filter {
150        input: Box<HirRelationExpr>,
151        predicates: Vec<HirScalarExpr>,
152    },
153    /// Unlike MirRelationExpr, we haven't yet compiled LeftOuter/RightOuter/FullOuter
154    /// joins away into more primitive exprs
155    Join {
156        left: Box<HirRelationExpr>,
157        right: Box<HirRelationExpr>,
158        on: HirScalarExpr,
159        kind: JoinKind,
160    },
161    /// Unlike MirRelationExpr, when `key` is empty AND `input` is empty this returns
162    /// a single row with the aggregates evaluated over empty groups, rather than returning zero
163    /// rows
164    Reduce {
165        input: Box<HirRelationExpr>,
166        group_key: Vec<usize>,
167        aggregates: Vec<AggregateExpr>,
168        expected_group_size: Option<u64>,
169    },
170    Distinct {
171        input: Box<HirRelationExpr>,
172    },
173    /// Groups and orders within each group, limiting output.
174    TopK {
175        /// The source collection.
176        input: Box<HirRelationExpr>,
177        /// Column indices used to form groups.
178        group_key: Vec<usize>,
179        /// Column indices used to order rows within groups.
180        order_key: Vec<ColumnOrder>,
181        /// Number of records to retain.
182        /// It is of SqlScalarType::Int64.
183        /// (UInt64 would make sense in theory: Then we wouldn't need to manually check
184        /// non-negativity, but would just get this for free when casting to UInt64. However, Int64
185        /// is better for Postgres compat. This is because if there is a $1 here, then when external
186        /// tools `describe` the prepared statement, they discover this type. If what they find
187        /// were UInt64, then they might have trouble calling the prepared statement, because the
188        /// unsigned types are non-standard, and also don't exist even in Postgres.)
189        limit: Option<HirScalarExpr>,
190        /// Number of records to skip.
191        /// It is of SqlScalarType::Int64.
192        /// This can contain parameters at first, but by the time we reach lowering, this should
193        /// already be simply a Literal.
194        offset: HirScalarExpr,
195        /// User-supplied hint: how many rows will have the same group key.
196        expected_group_size: Option<u64>,
197    },
198    Negate {
199        input: Box<HirRelationExpr>,
200    },
201    /// Keep rows from a dataflow where the row counts are positive.
202    Threshold {
203        input: Box<HirRelationExpr>,
204    },
205    Union {
206        base: Box<HirRelationExpr>,
207        inputs: Vec<HirRelationExpr>,
208    },
209}
210
211/// Stored column metadata.
212pub type NameMetadata = TreatAsEqual<Option<Arc<str>>>;
213
214#[derive(
215    Debug,
216    Clone,
217    PartialEq,
218    Eq,
219    PartialOrd,
220    Ord,
221    Hash,
222    Serialize,
223    Deserialize
224)]
225/// Just like [`mz_expr::MirScalarExpr`], except where otherwise noted below.
226pub enum HirScalarExpr {
227    /// Unlike mz_expr::MirScalarExpr, we can nest HirRelationExprs via eg Exists. This means that a
228    /// variable could refer to a column of the current input, or to a column of an outer relation.
229    /// We use ColumnRef to denote the difference.
230    Column(ColumnRef, NameMetadata),
231    Parameter(usize, NameMetadata),
232    Literal(Row, SqlColumnType, NameMetadata),
233    CallUnmaterializable(UnmaterializableFunc, NameMetadata),
234    CallUnary {
235        func: UnaryFunc,
236        expr: Box<HirScalarExpr>,
237        name: NameMetadata,
238    },
239    CallBinary {
240        func: BinaryFunc,
241        expr1: Box<HirScalarExpr>,
242        expr2: Box<HirScalarExpr>,
243        name: NameMetadata,
244    },
245    CallVariadic {
246        func: VariadicFunc,
247        exprs: Vec<HirScalarExpr>,
248        name: NameMetadata,
249    },
250    If {
251        cond: Box<HirScalarExpr>,
252        then: Box<HirScalarExpr>,
253        els: Box<HirScalarExpr>,
254        name: NameMetadata,
255    },
256    /// Returns true if `expr` returns any rows
257    Exists(Box<HirRelationExpr>, NameMetadata),
258    /// Given `expr` with arity 1. If expr returns:
259    /// * 0 rows, return NULL
260    /// * 1 row, return the value of that row
261    /// * >1 rows, we return an error
262    Select(Box<HirRelationExpr>, NameMetadata),
263    Windowing(WindowExpr, NameMetadata),
264}
265
266#[derive(
267    Debug,
268    Clone,
269    PartialEq,
270    Eq,
271    PartialOrd,
272    Ord,
273    Hash,
274    Serialize,
275    Deserialize
276)]
277/// Represents the invocation of a window function over an optional partitioning with an optional
278/// order.
279pub struct WindowExpr {
280    pub func: WindowExprType,
281    pub partition_by: Vec<HirScalarExpr>,
282    /// ORDER BY is represented in a complicated way: `plan_function_order_by` gave us two things:
283    ///  - the `ColumnOrder`s we have put in the `order_by` fields in the `WindowExprType` in `func`
284    ///    above,
285    ///  - the `HirScalarExpr`s we have put in the following `order_by` field.
286    /// These are separated because they are used in different places: the outer `order_by` is used
287    /// in the lowering: based on it, we create a Row constructor that collects the scalar exprs;
288    /// the inner `order_by` is used in the rendering to actually execute the ordering on these Rows.
289    /// (`WindowExpr` exists only in HIR, but not in MIR.)
290    /// Note that the `column` field in the `ColumnOrder`s point into the Row constructed in the
291    /// lowering, and not to original input columns.
292    pub order_by: Vec<HirScalarExpr>,
293}
294
295impl WindowExpr {
296    pub fn visit_expressions<'a, F, E>(&'a self, f: &mut F) -> Result<(), E>
297    where
298        F: FnMut(&'a HirScalarExpr) -> Result<(), E>,
299    {
300        #[allow(deprecated)]
301        self.func.visit_expressions(f)?;
302        for expr in self.partition_by.iter() {
303            f(expr)?;
304        }
305        for expr in self.order_by.iter() {
306            f(expr)?;
307        }
308        Ok(())
309    }
310
311    pub fn visit_expressions_mut<'a, F, E>(&'a mut self, f: &mut F) -> Result<(), E>
312    where
313        F: FnMut(&'a mut HirScalarExpr) -> Result<(), E>,
314    {
315        #[allow(deprecated)]
316        self.func.visit_expressions_mut(f)?;
317        for expr in self.partition_by.iter_mut() {
318            f(expr)?;
319        }
320        for expr in self.order_by.iter_mut() {
321            f(expr)?;
322        }
323        Ok(())
324    }
325}
326
327/// Yields the scalars in `func`'s arguments, plus those in `partition_by`
328/// and `order_by`; does not descend into them.
329impl VisitChildren<HirScalarExpr> for WindowExpr {
330    fn visit_children<F>(&self, mut f: F)
331    where
332        F: FnMut(&HirScalarExpr),
333    {
334        self.func.visit_children(&mut f);
335        for expr in self.partition_by.iter() {
336            f(expr);
337        }
338        for expr in self.order_by.iter() {
339            f(expr);
340        }
341    }
342
343    fn visit_mut_children<F>(&mut self, mut f: F)
344    where
345        F: FnMut(&mut HirScalarExpr),
346    {
347        self.func.visit_mut_children(&mut f);
348        for expr in self.partition_by.iter_mut() {
349            f(expr);
350        }
351        for expr in self.order_by.iter_mut() {
352            f(expr);
353        }
354    }
355
356    fn try_visit_children<F, E>(&self, mut f: F) -> Result<(), E>
357    where
358        F: FnMut(&HirScalarExpr) -> Result<(), E>,
359    {
360        self.func.try_visit_children(&mut f)?;
361        for expr in self.partition_by.iter() {
362            f(expr)?;
363        }
364        for expr in self.order_by.iter() {
365            f(expr)?;
366        }
367        Ok(())
368    }
369
370    fn try_visit_mut_children<F, E>(&mut self, mut f: F) -> Result<(), E>
371    where
372        F: FnMut(&mut HirScalarExpr) -> Result<(), E>,
373    {
374        self.func.try_visit_mut_children(&mut f)?;
375        for expr in self.partition_by.iter_mut() {
376            f(expr)?;
377        }
378        for expr in self.order_by.iter_mut() {
379            f(expr)?;
380        }
381        Ok(())
382    }
383
384    fn children<'a>(&'a self) -> impl DoubleEndedIterator<Item = &'a HirScalarExpr>
385    where
386        HirScalarExpr: 'a,
387    {
388        self.func
389            .children()
390            .chain(self.partition_by.iter())
391            .chain(self.order_by.iter())
392    }
393
394    fn children_mut<'a>(&'a mut self) -> impl DoubleEndedIterator<Item = &'a mut HirScalarExpr>
395    where
396        HirScalarExpr: 'a,
397    {
398        self.func
399            .children_mut()
400            .chain(self.partition_by.iter_mut())
401            .chain(self.order_by.iter_mut())
402    }
403}
404
405#[derive(
406    Debug,
407    Clone,
408    PartialEq,
409    Eq,
410    PartialOrd,
411    Ord,
412    Hash,
413    Serialize,
414    Deserialize
415)]
416/// A window function with its parameters.
417///
418/// There are three types of window functions:
419/// - scalar window functions, which return a different scalar value for each
420///   row within a partition that depends exclusively on the position of the row
421///   within the partition;
422/// - value window functions, which return a scalar value for each row within a
423///   partition that might be computed based on a single row, which is usually not
424///   the current row (e.g., previous or following row; first or last row of the
425///   partition);
426/// - aggregate window functions, which compute a traditional aggregation as a
427///   window function (e.g. `sum(x) OVER (...)`).
428///   (Aggregate window  functions can in some cases be computed by joining the
429///   input relation with a reduction over the same relation that computes the
430///   aggregation using the partition key as its grouping key, but we don't
431///   automatically do this currently.)
432pub enum WindowExprType {
433    Scalar(ScalarWindowExpr),
434    Value(ValueWindowExpr),
435    Aggregate(AggregateWindowExpr),
436}
437
438impl WindowExprType {
439    #[deprecated = "Use `VisitChildren<HirScalarExpr>::visit_children` instead."]
440    pub fn visit_expressions<'a, F, E>(&'a self, f: &mut F) -> Result<(), E>
441    where
442        F: FnMut(&'a HirScalarExpr) -> Result<(), E>,
443    {
444        #[allow(deprecated)]
445        match self {
446            Self::Scalar(expr) => expr.visit_expressions(f),
447            Self::Value(expr) => expr.visit_expressions(f),
448            Self::Aggregate(expr) => expr.visit_expressions(f),
449        }
450    }
451
452    #[deprecated = "Use `VisitChildren<HirScalarExpr>::visit_mut_children` instead."]
453    pub fn visit_expressions_mut<'a, F, E>(&'a mut self, f: &mut F) -> Result<(), E>
454    where
455        F: FnMut(&'a mut HirScalarExpr) -> Result<(), E>,
456    {
457        #[allow(deprecated)]
458        match self {
459            Self::Scalar(expr) => expr.visit_expressions_mut(f),
460            Self::Value(expr) => expr.visit_expressions_mut(f),
461            Self::Aggregate(expr) => expr.visit_expressions_mut(f),
462        }
463    }
464
465    fn typ(
466        &self,
467        outers: &[SqlRelationType],
468        inner: &SqlRelationType,
469        params: &BTreeMap<usize, SqlScalarType>,
470    ) -> SqlColumnType {
471        match self {
472            Self::Scalar(expr) => expr.typ(outers, inner, params),
473            Self::Value(expr) => expr.typ(outers, inner, params),
474            Self::Aggregate(expr) => expr.typ(outers, inner, params),
475        }
476    }
477}
478
479/// Dispatches to the inner `Value` / `Aggregate` variant's scalar children;
480/// `Scalar` window functions have no scalar children.
481impl VisitChildren<HirScalarExpr> for WindowExprType {
482    fn visit_children<F>(&self, f: F)
483    where
484        F: FnMut(&HirScalarExpr),
485    {
486        match self {
487            Self::Scalar(_) => (),
488            Self::Value(expr) => expr.visit_children(f),
489            Self::Aggregate(expr) => expr.visit_children(f),
490        }
491    }
492
493    fn visit_mut_children<F>(&mut self, f: F)
494    where
495        F: FnMut(&mut HirScalarExpr),
496    {
497        match self {
498            Self::Scalar(_) => (),
499            Self::Value(expr) => expr.visit_mut_children(f),
500            Self::Aggregate(expr) => expr.visit_mut_children(f),
501        }
502    }
503
504    fn try_visit_children<F, E>(&self, f: F) -> Result<(), E>
505    where
506        F: FnMut(&HirScalarExpr) -> Result<(), E>,
507    {
508        match self {
509            Self::Scalar(_) => Ok(()),
510            Self::Value(expr) => expr.try_visit_children(f),
511            Self::Aggregate(expr) => expr.try_visit_children(f),
512        }
513    }
514
515    fn try_visit_mut_children<F, E>(&mut self, f: F) -> Result<(), E>
516    where
517        F: FnMut(&mut HirScalarExpr) -> Result<(), E>,
518    {
519        match self {
520            Self::Scalar(_) => Ok(()),
521            Self::Value(expr) => expr.try_visit_mut_children(f),
522            Self::Aggregate(expr) => expr.try_visit_mut_children(f),
523        }
524    }
525
526    fn children<'a>(&'a self) -> impl DoubleEndedIterator<Item = &'a HirScalarExpr>
527    where
528        HirScalarExpr: 'a,
529    {
530        match self {
531            Self::Scalar(_) => vec![],
532            Self::Value(expr) => expr.children().collect(),
533            Self::Aggregate(expr) => expr.children().collect(),
534        }
535        .into_iter()
536    }
537
538    fn children_mut<'a>(&'a mut self) -> impl DoubleEndedIterator<Item = &'a mut HirScalarExpr>
539    where
540        HirScalarExpr: 'a,
541    {
542        match self {
543            Self::Scalar(_) => vec![],
544            Self::Value(expr) => expr.children_mut().collect(),
545            Self::Aggregate(expr) => expr.children_mut().collect(),
546        }
547        .into_iter()
548    }
549}
550
551#[derive(
552    Debug,
553    Clone,
554    PartialEq,
555    Eq,
556    PartialOrd,
557    Ord,
558    Hash,
559    Serialize,
560    Deserialize
561)]
562pub struct ScalarWindowExpr {
563    pub func: ScalarWindowFunc,
564    pub order_by: Vec<ColumnOrder>,
565}
566
567impl ScalarWindowExpr {
568    #[deprecated = "Implement `VisitChildren<HirScalarExpr>` if needed."]
569    pub fn visit_expressions<'a, F, E>(&'a self, _f: &mut F) -> Result<(), E>
570    where
571        F: FnMut(&'a HirScalarExpr) -> Result<(), E>,
572    {
573        match self.func {
574            ScalarWindowFunc::RowNumber => {}
575            ScalarWindowFunc::Rank => {}
576            ScalarWindowFunc::DenseRank => {}
577        }
578        Ok(())
579    }
580
581    #[deprecated = "Implement `VisitChildren<HirScalarExpr>` if needed."]
582    pub fn visit_expressions_mut<'a, F, E>(&'a self, _f: &mut F) -> Result<(), E>
583    where
584        F: FnMut(&'a mut HirScalarExpr) -> Result<(), E>,
585    {
586        match self.func {
587            ScalarWindowFunc::RowNumber => {}
588            ScalarWindowFunc::Rank => {}
589            ScalarWindowFunc::DenseRank => {}
590        }
591        Ok(())
592    }
593
594    fn typ(
595        &self,
596        _outers: &[SqlRelationType],
597        _inner: &SqlRelationType,
598        _params: &BTreeMap<usize, SqlScalarType>,
599    ) -> SqlColumnType {
600        self.func.output_sql_type()
601    }
602
603    pub fn into_expr(self) -> mz_expr::AggregateFunc {
604        match self.func {
605            ScalarWindowFunc::RowNumber => mz_expr::AggregateFunc::RowNumber {
606                order_by: self.order_by,
607            },
608            ScalarWindowFunc::Rank => mz_expr::AggregateFunc::Rank {
609                order_by: self.order_by,
610            },
611            ScalarWindowFunc::DenseRank => mz_expr::AggregateFunc::DenseRank {
612                order_by: self.order_by,
613            },
614        }
615    }
616}
617
618#[derive(
619    Debug,
620    Clone,
621    PartialEq,
622    Eq,
623    PartialOrd,
624    Ord,
625    Hash,
626    Serialize,
627    Deserialize
628)]
629/// Scalar Window functions
630pub enum ScalarWindowFunc {
631    RowNumber,
632    Rank,
633    DenseRank,
634}
635
636impl Display for ScalarWindowFunc {
637    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
638        match self {
639            ScalarWindowFunc::RowNumber => write!(f, "row_number"),
640            ScalarWindowFunc::Rank => write!(f, "rank"),
641            ScalarWindowFunc::DenseRank => write!(f, "dense_rank"),
642        }
643    }
644}
645
646impl ScalarWindowFunc {
647    pub fn output_sql_type(&self) -> SqlColumnType {
648        match self {
649            ScalarWindowFunc::RowNumber => SqlScalarType::Int64.nullable(false),
650            ScalarWindowFunc::Rank => SqlScalarType::Int64.nullable(false),
651            ScalarWindowFunc::DenseRank => SqlScalarType::Int64.nullable(false),
652        }
653    }
654}
655
656#[derive(
657    Debug,
658    Clone,
659    PartialEq,
660    Eq,
661    PartialOrd,
662    Ord,
663    Hash,
664    Serialize,
665    Deserialize
666)]
667pub struct ValueWindowExpr {
668    pub func: ValueWindowFunc,
669    /// If the argument list has a single element (e.g., for `first_value`), then it's that element.
670    /// If the argument list has multiple elements (e.g., for `lag`), then it's encoded in a record,
671    /// e.g., `row(#1, 3, null)`.
672    /// If it's a fused window function, then the arguments of each of the constituent function
673    /// calls are wrapped in an outer record.
674    pub args: Box<HirScalarExpr>,
675    /// See comment on `WindowExpr::order_by`.
676    pub order_by: Vec<ColumnOrder>,
677    pub window_frame: WindowFrame,
678    pub ignore_nulls: bool,
679}
680
681impl Display for ValueWindowFunc {
682    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
683        match self {
684            ValueWindowFunc::Lag => write!(f, "lag"),
685            ValueWindowFunc::Lead => write!(f, "lead"),
686            ValueWindowFunc::FirstValue => write!(f, "first_value"),
687            ValueWindowFunc::LastValue => write!(f, "last_value"),
688            ValueWindowFunc::Fused(funcs) => write!(f, "fused[{}]", separated(", ", funcs)),
689        }
690    }
691}
692
693impl ValueWindowExpr {
694    #[deprecated = "Use `VisitChildren<HirScalarExpr>::visit_children` instead."]
695    pub fn visit_expressions<'a, F, E>(&'a self, f: &mut F) -> Result<(), E>
696    where
697        F: FnMut(&'a HirScalarExpr) -> Result<(), E>,
698    {
699        f(&self.args)
700    }
701
702    #[deprecated = "Use `VisitChildren<HirScalarExpr>::visit_mut_children` instead."]
703    pub fn visit_expressions_mut<'a, F, E>(&'a mut self, f: &mut F) -> Result<(), E>
704    where
705        F: FnMut(&'a mut HirScalarExpr) -> Result<(), E>,
706    {
707        f(&mut self.args)
708    }
709
710    fn typ(
711        &self,
712        outers: &[SqlRelationType],
713        inner: &SqlRelationType,
714        params: &BTreeMap<usize, SqlScalarType>,
715    ) -> SqlColumnType {
716        self.func
717            .output_sql_type(self.args.typ(outers, inner, params))
718    }
719
720    /// Converts into `mz_expr::AggregateFunc`.
721    pub fn into_expr(self) -> (Box<HirScalarExpr>, mz_expr::AggregateFunc) {
722        (
723            self.args,
724            self.func
725                .into_expr(self.order_by, self.window_frame, self.ignore_nulls),
726        )
727    }
728}
729
730/// Yields `args` (the value-window function's argument expression);
731/// does not descend into it.
732impl VisitChildren<HirScalarExpr> for ValueWindowExpr {
733    // `visit_children` and friends are not implemented explicitly: the trait
734    // defaults delegate to `children`/`children_mut` below, which yield the
735    // single argument expression.
736
737    fn children<'a>(&'a self) -> impl DoubleEndedIterator<Item = &'a HirScalarExpr>
738    where
739        HirScalarExpr: 'a,
740    {
741        // Yield `args` itself; we must not descend into it (see the impl-level
742        // doc comment). Descending would skip `args` when it is a leaf node
743        // (e.g. `first_value(mz_now())`), breaking `contains_temporal`.
744        std::iter::once(&*self.args)
745    }
746
747    fn children_mut<'a>(&'a mut self) -> impl DoubleEndedIterator<Item = &'a mut HirScalarExpr>
748    where
749        HirScalarExpr: 'a,
750    {
751        std::iter::once(&mut *self.args)
752    }
753}
754
755#[derive(
756    Debug,
757    Clone,
758    PartialEq,
759    Eq,
760    PartialOrd,
761    Ord,
762    Hash,
763    Serialize,
764    Deserialize
765)]
766/// Value Window functions
767pub enum ValueWindowFunc {
768    Lag,
769    Lead,
770    FirstValue,
771    LastValue,
772    Fused(Vec<ValueWindowFunc>),
773}
774
775impl ValueWindowFunc {
776    pub fn output_sql_type(&self, input_type: SqlColumnType) -> SqlColumnType {
777        match self {
778            ValueWindowFunc::Lag | ValueWindowFunc::Lead => {
779                // The input is a (value, offset, default) record, so extract the type of the first arg
780                input_type.scalar_type.unwrap_record_element_type()[0]
781                    .clone()
782                    .nullable(true)
783            }
784            ValueWindowFunc::FirstValue | ValueWindowFunc::LastValue => {
785                input_type.scalar_type.nullable(true)
786            }
787            ValueWindowFunc::Fused(funcs) => {
788                let input_types = input_type.scalar_type.unwrap_record_element_column_type();
789                SqlScalarType::Record {
790                    fields: funcs
791                        .iter()
792                        .zip_eq(input_types)
793                        .map(|(f, t)| (ColumnName::from(""), f.output_sql_type(t.clone())))
794                        .collect(),
795                    custom_id: None,
796                }
797                .nullable(false)
798            }
799        }
800    }
801
802    pub fn into_expr(
803        self,
804        order_by: Vec<ColumnOrder>,
805        window_frame: WindowFrame,
806        ignore_nulls: bool,
807    ) -> mz_expr::AggregateFunc {
808        match self {
809            // Lag and Lead are fundamentally the same function, just with opposite directions
810            ValueWindowFunc::Lag => mz_expr::AggregateFunc::LagLead {
811                order_by,
812                lag_lead: mz_expr::LagLeadType::Lag,
813                ignore_nulls,
814            },
815            ValueWindowFunc::Lead => mz_expr::AggregateFunc::LagLead {
816                order_by,
817                lag_lead: mz_expr::LagLeadType::Lead,
818                ignore_nulls,
819            },
820            ValueWindowFunc::FirstValue => mz_expr::AggregateFunc::FirstValue {
821                order_by,
822                window_frame,
823            },
824            ValueWindowFunc::LastValue => mz_expr::AggregateFunc::LastValue {
825                order_by,
826                window_frame,
827            },
828            ValueWindowFunc::Fused(funcs) => mz_expr::AggregateFunc::FusedValueWindowFunc {
829                funcs: funcs
830                    .into_iter()
831                    .map(|func| {
832                        func.into_expr(order_by.clone(), window_frame.clone(), ignore_nulls)
833                    })
834                    .collect(),
835                order_by,
836            },
837        }
838    }
839}
840
841#[derive(
842    Debug,
843    Clone,
844    PartialEq,
845    Eq,
846    PartialOrd,
847    Ord,
848    Hash,
849    Serialize,
850    Deserialize
851)]
852pub struct AggregateWindowExpr {
853    pub aggregate_expr: AggregateExpr,
854    pub order_by: Vec<ColumnOrder>,
855    pub window_frame: WindowFrame,
856}
857
858impl AggregateWindowExpr {
859    #[deprecated = "Use `VisitChildren<HirScalarExpr>::visit_children` instead."]
860    pub fn visit_expressions<'a, F, E>(&'a self, f: &mut F) -> Result<(), E>
861    where
862        F: FnMut(&'a HirScalarExpr) -> Result<(), E>,
863    {
864        f(&self.aggregate_expr.expr)
865    }
866
867    #[deprecated = "Use `VisitChildren<HirScalarExpr>::visit_mut_children` instead."]
868    pub fn visit_expressions_mut<'a, F, E>(&'a mut self, f: &mut F) -> Result<(), E>
869    where
870        F: FnMut(&'a mut HirScalarExpr) -> Result<(), E>,
871    {
872        f(&mut self.aggregate_expr.expr)
873    }
874
875    fn typ(
876        &self,
877        outers: &[SqlRelationType],
878        inner: &SqlRelationType,
879        params: &BTreeMap<usize, SqlScalarType>,
880    ) -> SqlColumnType {
881        self.aggregate_expr
882            .func
883            .output_sql_type(self.aggregate_expr.expr.typ(outers, inner, params))
884    }
885
886    pub fn into_expr(self) -> (Box<HirScalarExpr>, mz_expr::AggregateFunc) {
887        if let AggregateFunc::FusedWindowAgg { funcs } = &self.aggregate_expr.func {
888            (
889                self.aggregate_expr.expr,
890                FusedWindowAggregate {
891                    wrapped_aggregates: funcs.iter().map(|f| f.clone().into_expr()).collect(),
892                    order_by: self.order_by,
893                    window_frame: self.window_frame,
894                },
895            )
896        } else {
897            (
898                self.aggregate_expr.expr,
899                WindowAggregate {
900                    wrapped_aggregate: Box::new(self.aggregate_expr.func.into_expr()),
901                    order_by: self.order_by,
902                    window_frame: self.window_frame,
903                },
904            )
905        }
906    }
907}
908
909/// Yields the aggregate's argument expression (`aggregate_expr.expr`);
910/// does not descend into it.
911impl VisitChildren<HirScalarExpr> for AggregateWindowExpr {
912    // `visit_children` and friends are not implemented explicitly: the trait
913    // defaults delegate to `children`/`children_mut` below, which yield the
914    // single aggregate argument expression.
915
916    fn children<'a>(&'a self) -> impl DoubleEndedIterator<Item = &'a HirScalarExpr>
917    where
918        HirScalarExpr: 'a,
919    {
920        // Yield the aggregate's argument expression itself; we must not descend
921        // into it. Descending would skip the argument when it is a leaf node
922        // (e.g. `max(mz_now())`), breaking `contains_temporal`.
923        std::iter::once(&*self.aggregate_expr.expr)
924    }
925
926    fn children_mut<'a>(&'a mut self) -> impl DoubleEndedIterator<Item = &'a mut HirScalarExpr>
927    where
928        HirScalarExpr: 'a,
929    {
930        std::iter::once(&mut *self.aggregate_expr.expr)
931    }
932}
933
934/// A `CoercibleScalarExpr` is a [`HirScalarExpr`] whose type is not fully
935/// determined. Several SQL expressions can be freely coerced based upon where
936/// in the expression tree they appear. For example, the string literal '42'
937/// will be automatically coerced to the integer 42 if used in a numeric
938/// context:
939///
940/// ```sql
941/// SELECT '42' + 42
942/// ```
943///
944/// This separate type gives the code that needs to interact with coercions very
945/// fine-grained control over what coercions happen and when.
946///
947/// The primary driver of coercion is function and operator selection, as
948/// choosing the correct function or operator implementation depends on the type
949/// of the provided arguments. Coercion also occurs at the very root of the
950/// scalar expression tree. For example in
951///
952/// ```sql
953/// SELECT ... WHERE $1
954/// ```
955///
956/// the `WHERE` clause will coerce the contained unconstrained type parameter
957/// `$1` to have type bool.
958#[derive(Clone, Debug)]
959pub enum CoercibleScalarExpr {
960    Coerced(HirScalarExpr),
961    Parameter(usize),
962    LiteralNull,
963    LiteralString(String),
964    LiteralRecord(Vec<CoercibleScalarExpr>),
965}
966
967impl CoercibleScalarExpr {
968    pub fn type_as(
969        self,
970        ecx: &ExprContext,
971        ty: &SqlScalarType,
972    ) -> Result<HirScalarExpr, PlanError> {
973        let expr = typeconv::plan_coerce(ecx, self, ty)?;
974        let expr_ty = ecx.scalar_type(&expr);
975        if ty != &expr_ty {
976            sql_bail!(
977                "{} must have type {}, not type {}",
978                ecx.name,
979                ecx.humanize_sql_scalar_type(ty, false),
980                ecx.humanize_sql_scalar_type(&expr_ty, false),
981            );
982        }
983        Ok(expr)
984    }
985
986    pub fn type_as_any(self, ecx: &ExprContext) -> Result<HirScalarExpr, PlanError> {
987        typeconv::plan_coerce(ecx, self, &SqlScalarType::String)
988    }
989
990    pub fn cast_to(
991        self,
992        ecx: &ExprContext,
993        ccx: CastContext,
994        ty: &SqlScalarType,
995    ) -> Result<HirScalarExpr, PlanError> {
996        let expr = typeconv::plan_coerce(ecx, self, ty)?;
997        typeconv::plan_cast(ecx, ccx, expr, ty)
998    }
999}
1000
1001/// The column type for a [`CoercibleScalarExpr`].
1002#[derive(Clone, Debug)]
1003pub enum CoercibleColumnType {
1004    Coerced(SqlColumnType),
1005    Record(Vec<CoercibleColumnType>),
1006    Uncoerced,
1007}
1008
1009impl CoercibleColumnType {
1010    /// Reports the nullability of the type.
1011    pub fn nullable(&self) -> bool {
1012        match self {
1013            // A coerced value's nullability is known.
1014            CoercibleColumnType::Coerced(ct) => ct.nullable,
1015
1016            // A literal record can never be null.
1017            CoercibleColumnType::Record(_) => false,
1018
1019            // An uncoerced literal may be the literal `NULL`, so we have
1020            // to conservatively assume it is nullable.
1021            CoercibleColumnType::Uncoerced => true,
1022        }
1023    }
1024}
1025
1026/// The scalar type for a [`CoercibleScalarExpr`].
1027#[derive(Clone, Debug)]
1028pub enum CoercibleScalarType {
1029    Coerced(SqlScalarType),
1030    Record(Vec<CoercibleColumnType>),
1031    Uncoerced,
1032}
1033
1034impl CoercibleScalarType {
1035    /// Reports whether the scalar type has been coerced.
1036    pub fn is_coerced(&self) -> bool {
1037        matches!(self, CoercibleScalarType::Coerced(_))
1038    }
1039
1040    /// Returns the coerced scalar type, if the type is coerced.
1041    pub fn as_coerced(&self) -> Option<&SqlScalarType> {
1042        match self {
1043            CoercibleScalarType::Coerced(t) => Some(t),
1044            _ => None,
1045        }
1046    }
1047
1048    /// If the type is coerced, apply the mapping function to the contained
1049    /// scalar type.
1050    pub fn map_coerced<F>(self, f: F) -> CoercibleScalarType
1051    where
1052        F: FnOnce(SqlScalarType) -> SqlScalarType,
1053    {
1054        match self {
1055            CoercibleScalarType::Coerced(t) => CoercibleScalarType::Coerced(f(t)),
1056            _ => self,
1057        }
1058    }
1059
1060    /// If the type is an coercible record, forcibly converts to a coerced
1061    /// record type. Any uncoerced field types are assumed to be of type text.
1062    ///
1063    /// Generally you should prefer to use [`typeconv::plan_coerce`], which
1064    /// accepts a type hint that can indicate the types of uncoerced field
1065    /// types.
1066    pub fn force_coerced_if_record(&mut self) {
1067        fn convert(uncoerced_fields: impl Iterator<Item = CoercibleColumnType>) -> SqlScalarType {
1068            let mut fields = vec![];
1069            for (i, uf) in uncoerced_fields.enumerate() {
1070                let name = ColumnName::from(format!("f{}", i + 1));
1071                let ty = match uf {
1072                    CoercibleColumnType::Coerced(ty) => ty,
1073                    CoercibleColumnType::Record(mut fields) => {
1074                        convert(fields.drain(..)).nullable(false)
1075                    }
1076                    CoercibleColumnType::Uncoerced => SqlScalarType::String.nullable(true),
1077                };
1078                fields.push((name, ty))
1079            }
1080            SqlScalarType::Record {
1081                fields: fields.into(),
1082                custom_id: None,
1083            }
1084        }
1085
1086        if let CoercibleScalarType::Record(fields) = self {
1087            *self = CoercibleScalarType::Coerced(convert(fields.drain(..)));
1088        }
1089    }
1090}
1091
1092/// An expression whose type can be ascertained.
1093///
1094/// Abstracts over `ScalarExpr` and `CoercibleScalarExpr`.
1095pub trait AbstractExpr {
1096    type Type: AbstractColumnType;
1097
1098    /// Computes the type of the expression.
1099    fn typ(
1100        &self,
1101        outers: &[SqlRelationType],
1102        inner: &SqlRelationType,
1103        params: &BTreeMap<usize, SqlScalarType>,
1104    ) -> Self::Type;
1105}
1106
1107impl AbstractExpr for CoercibleScalarExpr {
1108    type Type = CoercibleColumnType;
1109
1110    fn typ(
1111        &self,
1112        outers: &[SqlRelationType],
1113        inner: &SqlRelationType,
1114        params: &BTreeMap<usize, SqlScalarType>,
1115    ) -> Self::Type {
1116        match self {
1117            CoercibleScalarExpr::Coerced(expr) => {
1118                CoercibleColumnType::Coerced(expr.typ(outers, inner, params))
1119            }
1120            CoercibleScalarExpr::LiteralRecord(scalars) => {
1121                let fields = scalars
1122                    .iter()
1123                    .map(|s| s.typ(outers, inner, params))
1124                    .collect();
1125                CoercibleColumnType::Record(fields)
1126            }
1127            _ => CoercibleColumnType::Uncoerced,
1128        }
1129    }
1130}
1131
1132/// A column type-like object whose underlying scalar type-like object can be
1133/// ascertained.
1134///
1135/// Abstracts over `SqlColumnType` and `CoercibleColumnType`.
1136pub trait AbstractColumnType {
1137    type AbstractScalarType;
1138
1139    /// Converts the column type-like object into its inner scalar type-like
1140    /// object.
1141    fn scalar_type(self) -> Self::AbstractScalarType;
1142}
1143
1144impl AbstractColumnType for SqlColumnType {
1145    type AbstractScalarType = SqlScalarType;
1146
1147    fn scalar_type(self) -> Self::AbstractScalarType {
1148        self.scalar_type
1149    }
1150}
1151
1152impl AbstractColumnType for CoercibleColumnType {
1153    type AbstractScalarType = CoercibleScalarType;
1154
1155    fn scalar_type(self) -> Self::AbstractScalarType {
1156        match self {
1157            CoercibleColumnType::Coerced(t) => CoercibleScalarType::Coerced(t.scalar_type),
1158            CoercibleColumnType::Record(t) => CoercibleScalarType::Record(t),
1159            CoercibleColumnType::Uncoerced => CoercibleScalarType::Uncoerced,
1160        }
1161    }
1162}
1163
1164impl From<HirScalarExpr> for CoercibleScalarExpr {
1165    fn from(expr: HirScalarExpr) -> CoercibleScalarExpr {
1166        CoercibleScalarExpr::Coerced(expr)
1167    }
1168}
1169
1170/// A leveled column reference.
1171///
1172/// In the course of decorrelation, multiple levels of nested subqueries are
1173/// traversed, and references to columns may correspond to different levels
1174/// of containing outer subqueries.
1175///
1176/// A `ColumnRef` allows expressions to refer to columns while being clear
1177/// about which level the column references without manually performing the
1178/// bookkeeping tracking their actual column locations.
1179///
1180/// Specifically, a `ColumnRef` refers to a column `level` subquery level *out*
1181/// from the reference, using `column` as a unique identifier in that subquery level.
1182/// A `level` of zero corresponds to the current scope, and levels increase to
1183/// indicate subqueries further "outwards".
1184#[derive(
1185    Debug,
1186    Clone,
1187    Copy,
1188    PartialEq,
1189    Eq,
1190    Hash,
1191    Ord,
1192    PartialOrd,
1193    Serialize,
1194    Deserialize
1195)]
1196pub struct ColumnRef {
1197    // scope level, where 0 is the current scope and 1+ are outer scopes.
1198    pub level: usize,
1199    // level-local column identifier used.
1200    pub column: usize,
1201}
1202
1203#[derive(
1204    Debug,
1205    Clone,
1206    PartialEq,
1207    Eq,
1208    PartialOrd,
1209    Ord,
1210    Hash,
1211    Serialize,
1212    Deserialize
1213)]
1214pub enum JoinKind {
1215    Inner,
1216    LeftOuter,
1217    RightOuter,
1218    FullOuter,
1219}
1220
1221impl fmt::Display for JoinKind {
1222    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
1223        write!(
1224            f,
1225            "{}",
1226            match self {
1227                JoinKind::Inner => "Inner",
1228                JoinKind::LeftOuter => "LeftOuter",
1229                JoinKind::RightOuter => "RightOuter",
1230                JoinKind::FullOuter => "FullOuter",
1231            }
1232        )
1233    }
1234}
1235
1236impl JoinKind {
1237    pub fn can_be_correlated(&self) -> bool {
1238        match self {
1239            JoinKind::Inner | JoinKind::LeftOuter => true,
1240            JoinKind::RightOuter | JoinKind::FullOuter => false,
1241        }
1242    }
1243
1244    pub fn can_elide_identity_left_join(&self) -> bool {
1245        match self {
1246            JoinKind::Inner | JoinKind::RightOuter => true,
1247            JoinKind::LeftOuter | JoinKind::FullOuter => false,
1248        }
1249    }
1250
1251    pub fn can_elide_identity_right_join(&self) -> bool {
1252        match self {
1253            JoinKind::Inner | JoinKind::LeftOuter => true,
1254            JoinKind::RightOuter | JoinKind::FullOuter => false,
1255        }
1256    }
1257}
1258
1259#[derive(
1260    Debug,
1261    Clone,
1262    PartialEq,
1263    Eq,
1264    PartialOrd,
1265    Ord,
1266    Hash,
1267    Serialize,
1268    Deserialize
1269)]
1270pub struct AggregateExpr {
1271    pub func: AggregateFunc,
1272    pub expr: Box<HirScalarExpr>,
1273    pub distinct: bool,
1274}
1275
1276/// Aggregate functions analogous to `mz_expr::AggregateFunc`, but whose
1277/// types may be different.
1278///
1279/// Specifically, the nullability of the aggregate columns is more common
1280/// here than in `expr`, as these aggregates may be applied over empty
1281/// result sets and should be null in those cases, whereas `expr` variants
1282/// only return null values when supplied nulls as input.
1283#[derive(
1284    Clone,
1285    Debug,
1286    Eq,
1287    PartialEq,
1288    PartialOrd,
1289    Ord,
1290    Hash,
1291    Serialize,
1292    Deserialize
1293)]
1294pub enum AggregateFunc {
1295    MaxNumeric,
1296    MaxInt16,
1297    MaxInt32,
1298    MaxInt64,
1299    MaxUInt16,
1300    MaxUInt32,
1301    MaxUInt64,
1302    MaxMzTimestamp,
1303    MaxFloat32,
1304    MaxFloat64,
1305    MaxBool,
1306    MaxString,
1307    MaxDate,
1308    MaxTimestamp,
1309    MaxTimestampTz,
1310    MaxInterval,
1311    MaxTime,
1312    MinNumeric,
1313    MinInt16,
1314    MinInt32,
1315    MinInt64,
1316    MinUInt16,
1317    MinUInt32,
1318    MinUInt64,
1319    MinMzTimestamp,
1320    MinFloat32,
1321    MinFloat64,
1322    MinBool,
1323    MinString,
1324    MinDate,
1325    MinTimestamp,
1326    MinTimestampTz,
1327    MinInterval,
1328    MinTime,
1329    SumInt16,
1330    SumInt32,
1331    SumInt64,
1332    SumUInt16,
1333    SumUInt32,
1334    SumUInt64,
1335    SumFloat32,
1336    SumFloat64,
1337    SumNumeric,
1338    SumInterval,
1339    Count,
1340    Any,
1341    All,
1342    /// Accumulates `Datum::List`s whose first element is a JSON-typed `Datum`s
1343    /// into a JSON list. The other elements are columns used by `order_by`.
1344    ///
1345    /// WARNING: Unlike the `jsonb_agg` function that is exposed by the SQL
1346    /// layer, this function filters out `Datum::Null`, for consistency with
1347    /// the other aggregate functions.
1348    JsonbAgg {
1349        order_by: Vec<ColumnOrder>,
1350    },
1351    /// Zips `Datum::List`s whose first element is a JSON-typed `Datum`s into a
1352    /// JSON map. The other elements are columns used by `order_by`.
1353    JsonbObjectAgg {
1354        order_by: Vec<ColumnOrder>,
1355    },
1356    /// Zips a `Datum::List` whose first element is a `Datum::List` guaranteed
1357    /// to be non-empty and whose len % 2 == 0 into a `Datum::Map`. The other
1358    /// elements are columns used by `order_by`.
1359    MapAgg {
1360        order_by: Vec<ColumnOrder>,
1361        value_type: SqlScalarType,
1362    },
1363    /// Accumulates `Datum::List`s whose first element is a `Datum::Array` into a
1364    /// single `Datum::Array`. The other elements are columns used by `order_by`.
1365    ArrayConcat {
1366        order_by: Vec<ColumnOrder>,
1367    },
1368    /// Accumulates `Datum::List`s whose first element is a `Datum::List` into a
1369    /// single `Datum::List`. The other elements are columns used by `order_by`.
1370    ListConcat {
1371        order_by: Vec<ColumnOrder>,
1372    },
1373    StringAgg {
1374        order_by: Vec<ColumnOrder>,
1375    },
1376    /// A bundle of fused window aggregations: its input is a record, whose each
1377    /// component will be the input to one of the `AggregateFunc`s.
1378    ///
1379    /// Importantly, this aggregation can only be present inside a `WindowExpr`,
1380    /// more specifically an `AggregateWindowExpr`.
1381    FusedWindowAgg {
1382        funcs: Vec<AggregateFunc>,
1383    },
1384    /// Accumulates any number of `Datum::Dummy`s into `Datum::Dummy`.
1385    ///
1386    /// Useful for removing an expensive aggregation while maintaining the shape
1387    /// of a reduce operator.
1388    Dummy,
1389}
1390
1391impl AggregateFunc {
1392    /// Converts the `sql::AggregateFunc` to a corresponding `mz_expr::AggregateFunc`.
1393    pub fn into_expr(self) -> mz_expr::AggregateFunc {
1394        match self {
1395            AggregateFunc::MaxNumeric => mz_expr::AggregateFunc::MaxNumeric,
1396            AggregateFunc::MaxInt16 => mz_expr::AggregateFunc::MaxInt16,
1397            AggregateFunc::MaxInt32 => mz_expr::AggregateFunc::MaxInt32,
1398            AggregateFunc::MaxInt64 => mz_expr::AggregateFunc::MaxInt64,
1399            AggregateFunc::MaxUInt16 => mz_expr::AggregateFunc::MaxUInt16,
1400            AggregateFunc::MaxUInt32 => mz_expr::AggregateFunc::MaxUInt32,
1401            AggregateFunc::MaxUInt64 => mz_expr::AggregateFunc::MaxUInt64,
1402            AggregateFunc::MaxMzTimestamp => mz_expr::AggregateFunc::MaxMzTimestamp,
1403            AggregateFunc::MaxFloat32 => mz_expr::AggregateFunc::MaxFloat32,
1404            AggregateFunc::MaxFloat64 => mz_expr::AggregateFunc::MaxFloat64,
1405            AggregateFunc::MaxBool => mz_expr::AggregateFunc::MaxBool,
1406            AggregateFunc::MaxString => mz_expr::AggregateFunc::MaxString,
1407            AggregateFunc::MaxDate => mz_expr::AggregateFunc::MaxDate,
1408            AggregateFunc::MaxTimestamp => mz_expr::AggregateFunc::MaxTimestamp,
1409            AggregateFunc::MaxTimestampTz => mz_expr::AggregateFunc::MaxTimestampTz,
1410            AggregateFunc::MaxInterval => mz_expr::AggregateFunc::MaxInterval,
1411            AggregateFunc::MaxTime => mz_expr::AggregateFunc::MaxTime,
1412            AggregateFunc::MinNumeric => mz_expr::AggregateFunc::MinNumeric,
1413            AggregateFunc::MinInt16 => mz_expr::AggregateFunc::MinInt16,
1414            AggregateFunc::MinInt32 => mz_expr::AggregateFunc::MinInt32,
1415            AggregateFunc::MinInt64 => mz_expr::AggregateFunc::MinInt64,
1416            AggregateFunc::MinUInt16 => mz_expr::AggregateFunc::MinUInt16,
1417            AggregateFunc::MinUInt32 => mz_expr::AggregateFunc::MinUInt32,
1418            AggregateFunc::MinUInt64 => mz_expr::AggregateFunc::MinUInt64,
1419            AggregateFunc::MinMzTimestamp => mz_expr::AggregateFunc::MinMzTimestamp,
1420            AggregateFunc::MinFloat32 => mz_expr::AggregateFunc::MinFloat32,
1421            AggregateFunc::MinFloat64 => mz_expr::AggregateFunc::MinFloat64,
1422            AggregateFunc::MinBool => mz_expr::AggregateFunc::MinBool,
1423            AggregateFunc::MinString => mz_expr::AggregateFunc::MinString,
1424            AggregateFunc::MinDate => mz_expr::AggregateFunc::MinDate,
1425            AggregateFunc::MinTimestamp => mz_expr::AggregateFunc::MinTimestamp,
1426            AggregateFunc::MinTimestampTz => mz_expr::AggregateFunc::MinTimestampTz,
1427            AggregateFunc::MinInterval => mz_expr::AggregateFunc::MinInterval,
1428            AggregateFunc::MinTime => mz_expr::AggregateFunc::MinTime,
1429            AggregateFunc::SumInt16 => mz_expr::AggregateFunc::SumInt16,
1430            AggregateFunc::SumInt32 => mz_expr::AggregateFunc::SumInt32,
1431            AggregateFunc::SumInt64 => mz_expr::AggregateFunc::SumInt64,
1432            AggregateFunc::SumUInt16 => mz_expr::AggregateFunc::SumUInt16,
1433            AggregateFunc::SumUInt32 => mz_expr::AggregateFunc::SumUInt32,
1434            AggregateFunc::SumUInt64 => mz_expr::AggregateFunc::SumUInt64,
1435            AggregateFunc::SumFloat32 => mz_expr::AggregateFunc::SumFloat32,
1436            AggregateFunc::SumFloat64 => mz_expr::AggregateFunc::SumFloat64,
1437            AggregateFunc::SumNumeric => mz_expr::AggregateFunc::SumNumeric,
1438            AggregateFunc::SumInterval => mz_expr::AggregateFunc::SumInterval,
1439            AggregateFunc::Count => mz_expr::AggregateFunc::Count,
1440            AggregateFunc::Any => mz_expr::AggregateFunc::Any,
1441            AggregateFunc::All => mz_expr::AggregateFunc::All,
1442            AggregateFunc::JsonbAgg { order_by } => mz_expr::AggregateFunc::JsonbAgg { order_by },
1443            AggregateFunc::JsonbObjectAgg { order_by } => {
1444                mz_expr::AggregateFunc::JsonbObjectAgg { order_by }
1445            }
1446            AggregateFunc::MapAgg {
1447                order_by,
1448                value_type,
1449            } => mz_expr::AggregateFunc::MapAgg {
1450                order_by,
1451                value_type,
1452            },
1453            AggregateFunc::ArrayConcat { order_by } => {
1454                mz_expr::AggregateFunc::ArrayConcat { order_by }
1455            }
1456            AggregateFunc::ListConcat { order_by } => {
1457                mz_expr::AggregateFunc::ListConcat { order_by }
1458            }
1459            AggregateFunc::StringAgg { order_by } => mz_expr::AggregateFunc::StringAgg { order_by },
1460            // `AggregateFunc::FusedWindowAgg` should be specially handled in
1461            // `AggregateWindowExpr::into_expr`.
1462            AggregateFunc::FusedWindowAgg { funcs: _ } => {
1463                panic!("into_expr called on FusedWindowAgg")
1464            }
1465            AggregateFunc::Dummy => mz_expr::AggregateFunc::Dummy,
1466        }
1467    }
1468
1469    /// Returns a datum whose inclusion in the aggregation will not change its
1470    /// result.
1471    ///
1472    /// # Panics
1473    ///
1474    /// Panics if called on a `FusedWindowAgg`.
1475    pub fn identity_datum(&self) -> Datum<'static> {
1476        match self {
1477            AggregateFunc::Any => Datum::False,
1478            AggregateFunc::All => Datum::True,
1479            AggregateFunc::Dummy => Datum::Dummy,
1480            AggregateFunc::ArrayConcat { .. } => Datum::empty_array(),
1481            AggregateFunc::ListConcat { .. } => Datum::empty_list(),
1482            AggregateFunc::MaxNumeric
1483            | AggregateFunc::MaxInt16
1484            | AggregateFunc::MaxInt32
1485            | AggregateFunc::MaxInt64
1486            | AggregateFunc::MaxUInt16
1487            | AggregateFunc::MaxUInt32
1488            | AggregateFunc::MaxUInt64
1489            | AggregateFunc::MaxMzTimestamp
1490            | AggregateFunc::MaxFloat32
1491            | AggregateFunc::MaxFloat64
1492            | AggregateFunc::MaxBool
1493            | AggregateFunc::MaxString
1494            | AggregateFunc::MaxDate
1495            | AggregateFunc::MaxTimestamp
1496            | AggregateFunc::MaxTimestampTz
1497            | AggregateFunc::MaxInterval
1498            | AggregateFunc::MaxTime
1499            | AggregateFunc::MinNumeric
1500            | AggregateFunc::MinInt16
1501            | AggregateFunc::MinInt32
1502            | AggregateFunc::MinInt64
1503            | AggregateFunc::MinUInt16
1504            | AggregateFunc::MinUInt32
1505            | AggregateFunc::MinUInt64
1506            | AggregateFunc::MinMzTimestamp
1507            | AggregateFunc::MinFloat32
1508            | AggregateFunc::MinFloat64
1509            | AggregateFunc::MinBool
1510            | AggregateFunc::MinString
1511            | AggregateFunc::MinDate
1512            | AggregateFunc::MinTimestamp
1513            | AggregateFunc::MinTimestampTz
1514            | AggregateFunc::MinInterval
1515            | AggregateFunc::MinTime
1516            | AggregateFunc::SumInt16
1517            | AggregateFunc::SumInt32
1518            | AggregateFunc::SumInt64
1519            | AggregateFunc::SumUInt16
1520            | AggregateFunc::SumUInt32
1521            | AggregateFunc::SumUInt64
1522            | AggregateFunc::SumFloat32
1523            | AggregateFunc::SumFloat64
1524            | AggregateFunc::SumNumeric
1525            | AggregateFunc::SumInterval
1526            | AggregateFunc::Count
1527            | AggregateFunc::JsonbAgg { .. }
1528            | AggregateFunc::JsonbObjectAgg { .. }
1529            | AggregateFunc::MapAgg { .. }
1530            | AggregateFunc::StringAgg { .. } => Datum::Null,
1531            AggregateFunc::FusedWindowAgg { funcs: _ } => {
1532                // `identity_datum` is used only in HIR planning, and `FusedWindowAgg` can't occur
1533                // in HIR planning, because it is introduced only during HIR transformation.
1534                //
1535                // The implementation could be something like the following, except that we need to
1536                // return a `Datum<'static>`, so we can't actually dynamically compute this.
1537                // ```
1538                // let temp_storage = RowArena::new();
1539                // temp_storage.make_datum(|packer| packer.push_list(funcs.iter().map(|f| f.identity_datum())))
1540                // ```
1541                panic!("FusedWindowAgg doesn't have an identity_datum")
1542            }
1543        }
1544    }
1545
1546    /// The output column type for the result of an aggregation.
1547    ///
1548    /// The output column type also contains nullability information, which
1549    /// is (without further information) true for aggregations that are not
1550    /// counts.
1551    pub fn output_sql_type(&self, input_type: SqlColumnType) -> SqlColumnType {
1552        let scalar_type = match self {
1553            AggregateFunc::Count => SqlScalarType::Int64,
1554            AggregateFunc::Any => SqlScalarType::Bool,
1555            AggregateFunc::All => SqlScalarType::Bool,
1556            AggregateFunc::JsonbAgg { .. } => SqlScalarType::Jsonb,
1557            AggregateFunc::JsonbObjectAgg { .. } => SqlScalarType::Jsonb,
1558            AggregateFunc::StringAgg { .. } => SqlScalarType::String,
1559            AggregateFunc::SumInt16 | AggregateFunc::SumInt32 => SqlScalarType::Int64,
1560            AggregateFunc::SumInt64 => SqlScalarType::Numeric {
1561                max_scale: Some(NumericMaxScale::ZERO),
1562            },
1563            AggregateFunc::SumUInt16 | AggregateFunc::SumUInt32 => SqlScalarType::UInt64,
1564            AggregateFunc::SumUInt64 => SqlScalarType::Numeric {
1565                max_scale: Some(NumericMaxScale::ZERO),
1566            },
1567            AggregateFunc::MapAgg { value_type, .. } => SqlScalarType::Map {
1568                value_type: Box::new(value_type.clone()),
1569                custom_id: None,
1570            },
1571            AggregateFunc::ArrayConcat { .. } | AggregateFunc::ListConcat { .. } => {
1572                match input_type.scalar_type {
1573                    // The input is wrapped in a Record if there's an ORDER BY, so extract it out.
1574                    SqlScalarType::Record { fields, .. } => fields[0].1.scalar_type.clone(),
1575                    _ => unreachable!(),
1576                }
1577            }
1578            AggregateFunc::MaxNumeric
1579            | AggregateFunc::MaxInt16
1580            | AggregateFunc::MaxInt32
1581            | AggregateFunc::MaxInt64
1582            | AggregateFunc::MaxUInt16
1583            | AggregateFunc::MaxUInt32
1584            | AggregateFunc::MaxUInt64
1585            | AggregateFunc::MaxMzTimestamp
1586            | AggregateFunc::MaxFloat32
1587            | AggregateFunc::MaxFloat64
1588            | AggregateFunc::MaxBool
1589            | AggregateFunc::MaxString
1590            | AggregateFunc::MaxDate
1591            | AggregateFunc::MaxTimestamp
1592            | AggregateFunc::MaxTimestampTz
1593            | AggregateFunc::MaxInterval
1594            | AggregateFunc::MaxTime
1595            | AggregateFunc::MinNumeric
1596            | AggregateFunc::MinInt16
1597            | AggregateFunc::MinInt32
1598            | AggregateFunc::MinInt64
1599            | AggregateFunc::MinUInt16
1600            | AggregateFunc::MinUInt32
1601            | AggregateFunc::MinUInt64
1602            | AggregateFunc::MinMzTimestamp
1603            | AggregateFunc::MinFloat32
1604            | AggregateFunc::MinFloat64
1605            | AggregateFunc::MinBool
1606            | AggregateFunc::MinString
1607            | AggregateFunc::MinDate
1608            | AggregateFunc::MinTimestamp
1609            | AggregateFunc::MinTimestampTz
1610            | AggregateFunc::MinInterval
1611            | AggregateFunc::MinTime
1612            | AggregateFunc::SumFloat32
1613            | AggregateFunc::SumFloat64
1614            | AggregateFunc::SumNumeric
1615            | AggregateFunc::SumInterval
1616            | AggregateFunc::Dummy => input_type.scalar_type,
1617            AggregateFunc::FusedWindowAgg { funcs } => {
1618                let input_types = input_type.scalar_type.unwrap_record_element_column_type();
1619                SqlScalarType::Record {
1620                    fields: funcs
1621                        .iter()
1622                        .zip_eq(input_types)
1623                        .map(|(f, t)| (ColumnName::from(""), f.output_sql_type(t.clone())))
1624                        .collect(),
1625                    custom_id: None,
1626                }
1627            }
1628        };
1629        // max/min/sum return null on empty sets
1630        let nullable = !matches!(self, AggregateFunc::Count);
1631        scalar_type.nullable(nullable)
1632    }
1633
1634    pub fn is_order_sensitive(&self) -> bool {
1635        use AggregateFunc::*;
1636        matches!(
1637            self,
1638            JsonbAgg { .. }
1639                | JsonbObjectAgg { .. }
1640                | MapAgg { .. }
1641                | ArrayConcat { .. }
1642                | ListConcat { .. }
1643                | StringAgg { .. }
1644        )
1645    }
1646}
1647
1648impl HirRelationExpr {
1649    /// Gets the SQL type of a self-contained, top-level expression.
1650    pub fn top_level_typ(&self) -> SqlRelationType {
1651        self.typ(&[], &BTreeMap::new())
1652    }
1653
1654    /// Gets the SQL type of the expression.
1655    ///
1656    /// `outers` gives types for outer relations.
1657    /// `params` gives types for parameters.
1658    pub fn typ(
1659        &self,
1660        outers: &[SqlRelationType],
1661        params: &BTreeMap<usize, SqlScalarType>,
1662    ) -> SqlRelationType {
1663        stack::maybe_grow(|| match self {
1664            HirRelationExpr::Constant { typ, .. } => typ.clone(),
1665            HirRelationExpr::Get { typ, .. } => typ.clone(),
1666            HirRelationExpr::Let { body, .. } => body.typ(outers, params),
1667            HirRelationExpr::LetRec { body, .. } => body.typ(outers, params),
1668            HirRelationExpr::Project { input, outputs } => {
1669                let input_typ = input.typ(outers, params);
1670                SqlRelationType::new(
1671                    outputs
1672                        .iter()
1673                        .map(|&i| input_typ.column_types[i].clone())
1674                        .collect(),
1675                )
1676            }
1677            HirRelationExpr::Map { input, scalars } => {
1678                let mut typ = input.typ(outers, params);
1679                for scalar in scalars {
1680                    typ.column_types.push(scalar.typ(outers, &typ, params));
1681                }
1682                typ
1683            }
1684            HirRelationExpr::CallTable { func, exprs: _ } => func.output_sql_type(),
1685            HirRelationExpr::Filter { input, .. } | HirRelationExpr::TopK { input, .. } => {
1686                input.typ(outers, params)
1687            }
1688            HirRelationExpr::Join {
1689                left, right, kind, ..
1690            } => {
1691                let left_nullable = matches!(kind, JoinKind::RightOuter | JoinKind::FullOuter);
1692                let right_nullable =
1693                    matches!(kind, JoinKind::LeftOuter { .. } | JoinKind::FullOuter);
1694                let lt = left.typ(outers, params).column_types.into_iter().map(|t| {
1695                    let nullable = t.nullable || left_nullable;
1696                    t.nullable(nullable)
1697                });
1698                let mut outers = outers.to_vec();
1699                outers.insert(0, SqlRelationType::new(lt.clone().collect()));
1700                let rt = right
1701                    .typ(&outers, params)
1702                    .column_types
1703                    .into_iter()
1704                    .map(|t| {
1705                        let nullable = t.nullable || right_nullable;
1706                        t.nullable(nullable)
1707                    });
1708                SqlRelationType::new(lt.chain(rt).collect())
1709            }
1710            HirRelationExpr::Reduce {
1711                input,
1712                group_key,
1713                aggregates,
1714                expected_group_size: _,
1715            } => {
1716                let input_typ = input.typ(outers, params);
1717                let mut column_types = group_key
1718                    .iter()
1719                    .map(|&i| input_typ.column_types[i].clone())
1720                    .collect::<Vec<_>>();
1721                for agg in aggregates {
1722                    column_types.push(agg.typ(outers, &input_typ, params));
1723                }
1724                // TODO(frank): add primary key information.
1725                SqlRelationType::new(column_types)
1726            }
1727            // TODO(frank): check for removal; add primary key information.
1728            HirRelationExpr::Distinct { input }
1729            | HirRelationExpr::Negate { input }
1730            | HirRelationExpr::Threshold { input } => input.typ(outers, params),
1731            HirRelationExpr::Union { base, inputs } => {
1732                let mut base_cols = base.typ(outers, params).column_types;
1733                for input in inputs {
1734                    for (base_col, col) in base_cols
1735                        .iter_mut()
1736                        .zip_eq(input.typ(outers, params).column_types)
1737                    {
1738                        *base_col = base_col.sql_union(&col).unwrap(); // HIR deliberately not using `union`
1739                    }
1740                }
1741                SqlRelationType::new(base_cols)
1742            }
1743        })
1744    }
1745
1746    pub fn arity(&self) -> usize {
1747        match self {
1748            HirRelationExpr::Constant { typ, .. } => typ.column_types.len(),
1749            HirRelationExpr::Get { typ, .. } => typ.column_types.len(),
1750            HirRelationExpr::Let { body, .. } => body.arity(),
1751            HirRelationExpr::LetRec { body, .. } => body.arity(),
1752            HirRelationExpr::Project { outputs, .. } => outputs.len(),
1753            HirRelationExpr::Map { input, scalars } => input.arity() + scalars.len(),
1754            HirRelationExpr::CallTable { func, exprs: _ } => func.output_arity(),
1755            HirRelationExpr::Filter { input, .. }
1756            | HirRelationExpr::TopK { input, .. }
1757            | HirRelationExpr::Distinct { input }
1758            | HirRelationExpr::Negate { input }
1759            | HirRelationExpr::Threshold { input } => input.arity(),
1760            HirRelationExpr::Join { left, right, .. } => left.arity() + right.arity(),
1761            HirRelationExpr::Union { base, .. } => base.arity(),
1762            HirRelationExpr::Reduce {
1763                group_key,
1764                aggregates,
1765                ..
1766            } => group_key.len() + aggregates.len(),
1767        }
1768    }
1769
1770    /// The number of relation nodes in this expression.
1771    ///
1772    /// Relations reached through scalar subqueries are included. The scalar
1773    /// nodes themselves are not, so a large predicate over a small input still
1774    /// counts as small. This is a structural size for comparing two
1775    /// expressions against each other, not a cost estimate.
1776    pub fn relation_node_count(&self) -> usize {
1777        let mut count = 0;
1778        self.visit_post(&mut |_| count += 1);
1779        count
1780    }
1781
1782    /// If self is a constant, return the value and the type, otherwise `None`.
1783    pub fn as_const(&self) -> Option<(&Vec<Row>, &SqlRelationType)> {
1784        match self {
1785            Self::Constant { rows, typ } => Some((rows, typ)),
1786            _ => None,
1787        }
1788    }
1789
1790    /// Reports whether this expression contains a column reference to its
1791    /// direct parent scope.
1792    pub fn is_correlated(&self) -> bool {
1793        let mut correlated = false;
1794        #[allow(deprecated)]
1795        self.visit_columns(0, &mut |depth, col| {
1796            if col.level > depth && col.level - depth == 1 {
1797                correlated = true;
1798            }
1799        });
1800        correlated
1801    }
1802
1803    pub fn is_join_identity(&self) -> bool {
1804        match self {
1805            HirRelationExpr::Constant { rows, .. } => rows.len() == 1 && self.arity() == 0,
1806            _ => false,
1807        }
1808    }
1809
1810    pub fn project(self, outputs: Vec<usize>) -> Self {
1811        if outputs.iter().copied().eq(0..self.arity()) {
1812            // The projection is trivial. Suppress it.
1813            self
1814        } else {
1815            HirRelationExpr::Project {
1816                input: Box::new(self),
1817                outputs,
1818            }
1819        }
1820    }
1821
1822    pub fn map(mut self, scalars: Vec<HirScalarExpr>) -> Self {
1823        if scalars.is_empty() {
1824            // The map is trivial. Suppress it.
1825            self
1826        } else if let HirRelationExpr::Map {
1827            scalars: old_scalars,
1828            input: _,
1829        } = &mut self
1830        {
1831            // Map applied to a map. Fuse the maps.
1832            old_scalars.extend(scalars);
1833            self
1834        } else {
1835            HirRelationExpr::Map {
1836                input: Box::new(self),
1837                scalars,
1838            }
1839        }
1840    }
1841
1842    pub fn filter(mut self, mut preds: Vec<HirScalarExpr>) -> Self {
1843        if let HirRelationExpr::Filter {
1844            input: _,
1845            predicates,
1846        } = &mut self
1847        {
1848            predicates.extend(preds);
1849            predicates.sort();
1850            predicates.dedup();
1851            self
1852        } else {
1853            preds.sort();
1854            preds.dedup();
1855            HirRelationExpr::Filter {
1856                input: Box::new(self),
1857                predicates: preds,
1858            }
1859        }
1860    }
1861
1862    pub fn reduce(
1863        self,
1864        group_key: Vec<usize>,
1865        aggregates: Vec<AggregateExpr>,
1866        expected_group_size: Option<u64>,
1867    ) -> Self {
1868        HirRelationExpr::Reduce {
1869            input: Box::new(self),
1870            group_key,
1871            aggregates,
1872            expected_group_size,
1873        }
1874    }
1875
1876    pub fn top_k(
1877        self,
1878        group_key: Vec<usize>,
1879        order_key: Vec<ColumnOrder>,
1880        limit: Option<HirScalarExpr>,
1881        offset: HirScalarExpr,
1882        expected_group_size: Option<u64>,
1883    ) -> Self {
1884        HirRelationExpr::TopK {
1885            input: Box::new(self),
1886            group_key,
1887            order_key,
1888            limit,
1889            offset,
1890            expected_group_size,
1891        }
1892    }
1893
1894    pub fn negate(self) -> Self {
1895        if let HirRelationExpr::Negate { input } = self {
1896            *input
1897        } else {
1898            HirRelationExpr::Negate {
1899                input: Box::new(self),
1900            }
1901        }
1902    }
1903
1904    pub fn distinct(self) -> Self {
1905        if let HirRelationExpr::Distinct { .. } = self {
1906            self
1907        } else {
1908            HirRelationExpr::Distinct {
1909                input: Box::new(self),
1910            }
1911        }
1912    }
1913
1914    pub fn threshold(self) -> Self {
1915        if let HirRelationExpr::Threshold { .. } = self {
1916            self
1917        } else {
1918            HirRelationExpr::Threshold {
1919                input: Box::new(self),
1920            }
1921        }
1922    }
1923
1924    pub fn union(self, other: Self) -> Self {
1925        let mut terms = Vec::new();
1926        if let HirRelationExpr::Union { base, inputs } = self {
1927            terms.push(*base);
1928            terms.extend(inputs);
1929        } else {
1930            terms.push(self);
1931        }
1932        if let HirRelationExpr::Union { base, inputs } = other {
1933            terms.push(*base);
1934            terms.extend(inputs);
1935        } else {
1936            terms.push(other);
1937        }
1938        HirRelationExpr::Union {
1939            base: Box::new(terms.remove(0)),
1940            inputs: terms,
1941        }
1942    }
1943
1944    pub fn exists(self) -> HirScalarExpr {
1945        HirScalarExpr::Exists(Box::new(self), NameMetadata::default())
1946    }
1947
1948    pub fn select(self) -> HirScalarExpr {
1949        HirScalarExpr::Select(Box::new(self), NameMetadata::default())
1950    }
1951
1952    pub fn join(
1953        self,
1954        mut right: HirRelationExpr,
1955        on: HirScalarExpr,
1956        kind: JoinKind,
1957    ) -> HirRelationExpr {
1958        if self.is_join_identity()
1959            && !right.is_correlated()
1960            && on == HirScalarExpr::literal_true()
1961            && kind.can_elide_identity_left_join()
1962        {
1963            // The join can be elided, but we need to adjust column references
1964            // on the right-hand side to account for the removal of the scope
1965            // introduced by the join.
1966            #[allow(deprecated)]
1967            right.visit_columns_mut(0, &mut |depth, col| {
1968                if col.level > depth {
1969                    col.level -= 1;
1970                }
1971            });
1972            right
1973        } else if right.is_join_identity()
1974            && on == HirScalarExpr::literal_true()
1975            && kind.can_elide_identity_right_join()
1976        {
1977            self
1978        } else {
1979            HirRelationExpr::Join {
1980                left: Box::new(self),
1981                right: Box::new(right),
1982                on,
1983                kind,
1984            }
1985        }
1986    }
1987
1988    pub fn take(&mut self) -> HirRelationExpr {
1989        mem::replace(
1990            self,
1991            HirRelationExpr::constant(vec![], SqlRelationType::new(Vec::new())),
1992        )
1993    }
1994
1995    #[deprecated = "Use `Visit::visit_post`."]
1996    pub fn visit<'a, F>(&'a self, depth: usize, f: &mut F)
1997    where
1998        F: FnMut(&'a Self, usize),
1999    {
2000        #[allow(deprecated)]
2001        let _ = self.visit_fallible(depth, &mut |e: &HirRelationExpr,
2002                                                 depth: usize|
2003         -> Result<(), ()> {
2004            f(e, depth);
2005            Ok(())
2006        });
2007    }
2008
2009    #[deprecated = "Use `Visit::try_visit_post`."]
2010    pub fn visit_fallible<'a, F, E>(&'a self, depth: usize, f: &mut F) -> Result<(), E>
2011    where
2012        F: FnMut(&'a Self, usize) -> Result<(), E>,
2013    {
2014        #[allow(deprecated)]
2015        // Grow the stack: this recurses over the relation tree, whose depth is
2016        // user-controlled (e.g. a long JOIN chain or a chain of CTEs).
2017        stack::maybe_grow(|| {
2018            self.visit1(depth, |e: &HirRelationExpr, depth: usize| {
2019                e.visit_fallible(depth, f)
2020            })
2021        })?;
2022        f(self, depth)
2023    }
2024
2025    /// WARNING: `VisitChildren<HirRelationExpr>::try_visit_children` is NOT a
2026    /// drop-in replacement: in addition to the relation children visited here,
2027    /// it also descends into every subquery (`Exists`/`Select`) reachable
2028    /// through `self`'s scalar children — at any depth of scalar nesting.
2029    #[deprecated = "Use `VisitChildren<HirRelationExpr>::try_visit_children` instead."]
2030    pub fn visit1<'a, F, E>(&'a self, depth: usize, mut f: F) -> Result<(), E>
2031    where
2032        F: FnMut(&'a Self, usize) -> Result<(), E>,
2033    {
2034        match self {
2035            HirRelationExpr::Constant { .. }
2036            | HirRelationExpr::Get { .. }
2037            | HirRelationExpr::CallTable { .. } => (),
2038            HirRelationExpr::Let { body, value, .. } => {
2039                f(value, depth)?;
2040                f(body, depth)?;
2041            }
2042            HirRelationExpr::LetRec {
2043                limit: _,
2044                bindings,
2045                body,
2046            } => {
2047                for (_, _, value, _) in bindings.iter() {
2048                    f(value, depth)?;
2049                }
2050                f(body, depth)?;
2051            }
2052            HirRelationExpr::Project { input, .. } => {
2053                f(input, depth)?;
2054            }
2055            HirRelationExpr::Map { input, .. } => {
2056                f(input, depth)?;
2057            }
2058            HirRelationExpr::Filter { input, .. } => {
2059                f(input, depth)?;
2060            }
2061            HirRelationExpr::Join { left, right, .. } => {
2062                f(left, depth)?;
2063                f(right, depth + 1)?;
2064            }
2065            HirRelationExpr::Reduce { input, .. } => {
2066                f(input, depth)?;
2067            }
2068            HirRelationExpr::Distinct { input } => {
2069                f(input, depth)?;
2070            }
2071            HirRelationExpr::TopK { input, .. } => {
2072                f(input, depth)?;
2073            }
2074            HirRelationExpr::Negate { input } => {
2075                f(input, depth)?;
2076            }
2077            HirRelationExpr::Threshold { input } => {
2078                f(input, depth)?;
2079            }
2080            HirRelationExpr::Union { base, inputs } => {
2081                f(base, depth)?;
2082                for input in inputs {
2083                    f(input, depth)?;
2084                }
2085            }
2086        }
2087        Ok(())
2088    }
2089
2090    #[deprecated = "Use `Visit::visit_mut_post` instead."]
2091    pub fn visit_mut<F>(&mut self, depth: usize, f: &mut F)
2092    where
2093        F: FnMut(&mut Self, usize),
2094    {
2095        #[allow(deprecated)]
2096        let _ = self.visit_mut_fallible(depth, &mut |e: &mut HirRelationExpr,
2097                                                     depth: usize|
2098         -> Result<(), ()> {
2099            f(e, depth);
2100            Ok(())
2101        });
2102    }
2103
2104    #[deprecated = "Use `Visit::try_visit_mut_post` instead."]
2105    pub fn visit_mut_fallible<F, E>(&mut self, depth: usize, f: &mut F) -> Result<(), E>
2106    where
2107        F: FnMut(&mut Self, usize) -> Result<(), E>,
2108    {
2109        #[allow(deprecated)]
2110        // Grow the stack: this recurses over the relation tree, whose depth is
2111        // user-controlled (e.g. a long JOIN chain or a chain of CTEs).
2112        stack::maybe_grow(|| {
2113            self.visit1_mut(depth, |e: &mut HirRelationExpr, depth: usize| {
2114                e.visit_mut_fallible(depth, f)
2115            })
2116        })?;
2117        f(self, depth)
2118    }
2119
2120    /// WARNING: `VisitChildren<HirRelationExpr>::try_visit_mut_children` is NOT a
2121    /// drop-in replacement: in addition to the relation children visited here,
2122    /// it also descends into every subquery (`Exists`/`Select`) reachable
2123    /// through `self`'s scalar children — at any depth of scalar nesting.
2124    #[deprecated = "Use `VisitChildren<HirRelationExpr>::try_visit_mut_children` instead."]
2125    pub fn visit1_mut<'a, F, E>(&'a mut self, depth: usize, mut f: F) -> Result<(), E>
2126    where
2127        F: FnMut(&'a mut Self, usize) -> Result<(), E>,
2128    {
2129        match self {
2130            HirRelationExpr::Constant { .. }
2131            | HirRelationExpr::Get { .. }
2132            | HirRelationExpr::CallTable { .. } => (),
2133            HirRelationExpr::Let { body, value, .. } => {
2134                f(value, depth)?;
2135                f(body, depth)?;
2136            }
2137            HirRelationExpr::LetRec {
2138                limit: _,
2139                bindings,
2140                body,
2141            } => {
2142                for (_, _, value, _) in bindings.iter_mut() {
2143                    f(value, depth)?;
2144                }
2145                f(body, depth)?;
2146            }
2147            HirRelationExpr::Project { input, .. } => {
2148                f(input, depth)?;
2149            }
2150            HirRelationExpr::Map { input, .. } => {
2151                f(input, depth)?;
2152            }
2153            HirRelationExpr::Filter { input, .. } => {
2154                f(input, depth)?;
2155            }
2156            HirRelationExpr::Join { left, right, .. } => {
2157                f(left, depth)?;
2158                f(right, depth + 1)?;
2159            }
2160            HirRelationExpr::Reduce { input, .. } => {
2161                f(input, depth)?;
2162            }
2163            HirRelationExpr::Distinct { input } => {
2164                f(input, depth)?;
2165            }
2166            HirRelationExpr::TopK { input, .. } => {
2167                f(input, depth)?;
2168            }
2169            HirRelationExpr::Negate { input } => {
2170                f(input, depth)?;
2171            }
2172            HirRelationExpr::Threshold { input } => {
2173                f(input, depth)?;
2174            }
2175            HirRelationExpr::Union { base, inputs } => {
2176                f(base, depth)?;
2177                for input in inputs {
2178                    f(input, depth)?;
2179                }
2180            }
2181        }
2182        Ok(())
2183    }
2184
2185    #[deprecated = "Use a combination of `Visit` and `VisitChildren` methods."]
2186    /// Visits all scalar expressions directly held by relation nodes within the sub-tree of `self`.
2187    ///
2188    /// Note: this does NOT descend into subqueries that may appear inside the visited
2189    /// `HirScalarExpr`s (i.e., `HirScalarExpr::Exists` / `HirScalarExpr::Select`). The closure
2190    /// `f` is invoked once per top-level scalar expression attached to a relation node, and it
2191    /// is the closure's responsibility to recurse into any subqueries if desired.
2192    ///
2193    /// The `depth` argument is just a seed: it is the value passed to `f` for scalar expressions
2194    /// at the root of `self`, and it is incremented by 1 when descending into the RHS of a
2195    /// `Join` node (the only place this function increments it). It does NOT control or limit
2196    /// recursion; passing `0` is the usual choice.
2197    pub fn visit_scalar_expressions<F, E>(&self, depth: usize, f: &mut F) -> Result<(), E>
2198    where
2199        F: FnMut(&HirScalarExpr, usize) -> Result<(), E>,
2200    {
2201        #[allow(deprecated)]
2202        self.visit_fallible(depth, &mut |e: &HirRelationExpr,
2203                                         depth: usize|
2204         -> Result<(), E> {
2205            match e {
2206                HirRelationExpr::Join { on, .. } => {
2207                    f(on, depth)?;
2208                }
2209                HirRelationExpr::Map { scalars, .. } => {
2210                    for scalar in scalars {
2211                        f(scalar, depth)?;
2212                    }
2213                }
2214                HirRelationExpr::CallTable { exprs, .. } => {
2215                    for expr in exprs {
2216                        f(expr, depth)?;
2217                    }
2218                }
2219                HirRelationExpr::Filter { predicates, .. } => {
2220                    for predicate in predicates {
2221                        f(predicate, depth)?;
2222                    }
2223                }
2224                HirRelationExpr::Reduce { aggregates, .. } => {
2225                    for aggregate in aggregates {
2226                        f(&aggregate.expr, depth)?;
2227                    }
2228                }
2229                HirRelationExpr::TopK { limit, offset, .. } => {
2230                    if let Some(limit) = limit {
2231                        f(limit, depth)?;
2232                    }
2233                    f(offset, depth)?;
2234                }
2235                HirRelationExpr::Union { .. }
2236                | HirRelationExpr::Let { .. }
2237                | HirRelationExpr::LetRec { .. }
2238                | HirRelationExpr::Project { .. }
2239                | HirRelationExpr::Distinct { .. }
2240                | HirRelationExpr::Negate { .. }
2241                | HirRelationExpr::Threshold { .. }
2242                | HirRelationExpr::Constant { .. }
2243                | HirRelationExpr::Get { .. } => (),
2244            }
2245            Ok(())
2246        })
2247    }
2248
2249    #[deprecated = "Use a combination of `Visit` and `VisitChildren` methods."]
2250    /// Like `visit_scalar_expressions`, but permits mutating the expressions.
2251    ///
2252    /// In particular, this also does NOT descend into subqueries inside the visited scalar
2253    /// expressions, and `depth` is just a seed for the value passed to `f` (see
2254    /// [`HirRelationExpr::visit_scalar_expressions`]).
2255    pub fn visit_scalar_expressions_mut<F, E>(&mut self, depth: usize, f: &mut F) -> Result<(), E>
2256    where
2257        F: FnMut(&mut HirScalarExpr, usize) -> Result<(), E>,
2258    {
2259        #[allow(deprecated)]
2260        self.visit_mut_fallible(depth, &mut |e: &mut HirRelationExpr,
2261                                             depth: usize|
2262         -> Result<(), E> {
2263            match e {
2264                HirRelationExpr::Join { on, .. } => {
2265                    f(on, depth)?;
2266                }
2267                HirRelationExpr::Map { scalars, .. } => {
2268                    for scalar in scalars.iter_mut() {
2269                        f(scalar, depth)?;
2270                    }
2271                }
2272                HirRelationExpr::CallTable { exprs, .. } => {
2273                    for expr in exprs.iter_mut() {
2274                        f(expr, depth)?;
2275                    }
2276                }
2277                HirRelationExpr::Filter { predicates, .. } => {
2278                    for predicate in predicates.iter_mut() {
2279                        f(predicate, depth)?;
2280                    }
2281                }
2282                HirRelationExpr::Reduce { aggregates, .. } => {
2283                    for aggregate in aggregates.iter_mut() {
2284                        f(&mut aggregate.expr, depth)?;
2285                    }
2286                }
2287                HirRelationExpr::TopK { limit, offset, .. } => {
2288                    if let Some(limit) = limit {
2289                        f(limit, depth)?;
2290                    }
2291                    f(offset, depth)?;
2292                }
2293                HirRelationExpr::Union { .. }
2294                | HirRelationExpr::Let { .. }
2295                | HirRelationExpr::LetRec { .. }
2296                | HirRelationExpr::Project { .. }
2297                | HirRelationExpr::Distinct { .. }
2298                | HirRelationExpr::Negate { .. }
2299                | HirRelationExpr::Threshold { .. }
2300                | HirRelationExpr::Constant { .. }
2301                | HirRelationExpr::Get { .. } => (),
2302            }
2303            Ok(())
2304        })
2305    }
2306
2307    #[deprecated = "Redefine this based on the `Visit` and `VisitChildren` methods."]
2308    /// Visits the column references in this relation expression.
2309    ///
2310    /// The `depth` argument should indicate the subquery nesting depth of the expression,
2311    /// which will be incremented when entering the RHS of a join or a subquery and
2312    /// presented to the supplied function `f`.
2313    pub fn visit_columns<F>(&self, depth: usize, f: &mut F)
2314    where
2315        F: FnMut(usize, &ColumnRef),
2316    {
2317        #[allow(deprecated)]
2318        let _ = self.visit_scalar_expressions(depth, &mut |e: &HirScalarExpr,
2319                                                           depth: usize|
2320         -> Result<(), ()> {
2321            e.visit_columns(depth, f);
2322            Ok(())
2323        });
2324    }
2325
2326    #[deprecated = "Redefine this based on the `Visit` and `VisitChildren` methods."]
2327    /// Like `visit_columns`, but permits mutating the column references.
2328    pub fn visit_columns_mut<F>(&mut self, depth: usize, f: &mut F)
2329    where
2330        F: FnMut(usize, &mut ColumnRef),
2331    {
2332        #[allow(deprecated)]
2333        let _ = self.visit_scalar_expressions_mut(depth, &mut |e: &mut HirScalarExpr,
2334                                                               depth: usize|
2335         -> Result<(), ()> {
2336            e.visit_columns_mut(depth, f);
2337            Ok(())
2338        });
2339    }
2340
2341    /// Replaces parameter references in the expression with the corresponding datum from `params`.
2342    /// Additionally, it simplifies OFFSET clauses to constants after parameter binding, and checks
2343    /// them for non-negativity.
2344    ///
2345    /// This walks the entire `HirRelationExpr` tree, including subqueries nested inside scalar
2346    /// expressions: `visit_scalar_expressions_mut` covers the scalar expressions held directly
2347    /// by relation nodes, and the corecursive call into
2348    /// [`HirScalarExpr::bind_parameters_and_simplify_offset`] (made by the closure below) is
2349    /// what descends into subqueries occurring inside those scalar expressions.
2350    pub fn bind_parameters_and_simplify_offset(
2351        &mut self,
2352        scx: &StatementContext,
2353        lifetime: QueryLifetime,
2354        params: &Params,
2355    ) -> Result<(), PlanError> {
2356        #[allow(deprecated)]
2357        self.visit_scalar_expressions_mut(0, &mut |e: &mut HirScalarExpr, _: usize| {
2358            e.bind_parameters_and_simplify_offset(scx, lifetime, params)
2359        })?;
2360
2361        // OFFSET clauses in `expr` should become constants with the above binding of parameters.
2362        // Let's check this and simplify them to literals.
2363        self.try_visit_mut_pre(&mut |expr| {
2364            if let HirRelationExpr::TopK { offset, .. } = expr {
2365                let offset_value = offset_into_value(offset.take())?;
2366                *offset = HirScalarExpr::literal(Datum::Int64(offset_value), SqlScalarType::Int64);
2367            }
2368            Ok::<(), PlanError>(())
2369        })
2370        // (We don't need to simplify LIMIT clauses in `expr`, because we can handle non-constant
2371        // expressions there. If they happen to be simplifiable to literals, then the optimizer will do
2372        // so later.)
2373    }
2374
2375    pub fn contains_parameters(&self) -> Result<bool, PlanError> {
2376        let mut contains_parameters = false;
2377        #[allow(deprecated)]
2378        self.visit_scalar_expressions(0, &mut |e: &HirScalarExpr, _: usize| {
2379            if e.contains_parameters() {
2380                contains_parameters = true;
2381            }
2382            Ok::<(), PlanError>(())
2383        })?;
2384        Ok(contains_parameters)
2385    }
2386
2387    /// See the documentation for [`HirScalarExpr::splice_parameters`].
2388    pub fn splice_parameters(&mut self, params: &[HirScalarExpr], depth: usize) {
2389        #[allow(deprecated)]
2390        let _ = self.visit_scalar_expressions_mut(depth, &mut |e: &mut HirScalarExpr,
2391                                                               depth: usize|
2392         -> Result<(), ()> {
2393            e.splice_parameters(params, depth);
2394            Ok(())
2395        });
2396    }
2397
2398    /// Constructs a constant collection from specific rows and schema.
2399    pub fn constant(rows: Vec<Vec<Datum>>, typ: SqlRelationType) -> Self {
2400        let rows = rows
2401            .into_iter()
2402            .map(move |datums| Row::pack_slice(&datums))
2403            .collect();
2404        HirRelationExpr::Constant { rows, typ }
2405    }
2406
2407    /// A `RowSetFinishing` can only be directly applied to the result of a one-shot select.
2408    /// This function is concerned with maintained queries, e.g., an index or materialized view.
2409    /// Instead of directly applying the given `RowSetFinishing`, it converts the `RowSetFinishing`
2410    /// to a `TopK`, which it then places at the top of `self`. Additionally, it turns the given
2411    /// finishing into a trivial finishing.
2412    pub fn finish_maintained(
2413        &mut self,
2414        finishing: &mut RowSetFinishing<HirScalarExpr, HirScalarExpr>,
2415        group_size_hints: GroupSizeHints,
2416    ) {
2417        if !HirRelationExpr::is_trivial_row_set_finishing_hir(finishing, self.arity()) {
2418            let old_finishing = mem::replace(
2419                finishing,
2420                HirRelationExpr::trivial_row_set_finishing_hir(finishing.project.len()),
2421            );
2422            *self = HirRelationExpr::top_k(
2423                std::mem::replace(
2424                    self,
2425                    HirRelationExpr::Constant {
2426                        rows: vec![],
2427                        typ: SqlRelationType::new(Vec::new()),
2428                    },
2429                ),
2430                vec![],
2431                old_finishing.order_by,
2432                old_finishing.limit,
2433                old_finishing.offset,
2434                group_size_hints.limit_input_group_size,
2435            )
2436            .project(old_finishing.project);
2437        }
2438    }
2439
2440    /// Returns a trivial finishing, i.e., that does nothing to the result set.
2441    ///
2442    /// (There is also `RowSetFinishing::trivial`, but that is specialized for when the O generic
2443    /// parameter is not an HirScalarExpr anymore.)
2444    pub fn trivial_row_set_finishing_hir(
2445        arity: usize,
2446    ) -> RowSetFinishing<HirScalarExpr, HirScalarExpr> {
2447        RowSetFinishing {
2448            order_by: Vec::new(),
2449            limit: None,
2450            offset: HirScalarExpr::literal(Datum::Int64(0), SqlScalarType::Int64),
2451            project: (0..arity).collect(),
2452        }
2453    }
2454
2455    /// True if the finishing does nothing to any result set.
2456    ///
2457    /// (There is also `RowSetFinishing::is_trivial`, but that is specialized for when the O generic
2458    /// parameter is not an HirScalarExpr anymore.)
2459    pub fn is_trivial_row_set_finishing_hir(
2460        rsf: &RowSetFinishing<HirScalarExpr, HirScalarExpr>,
2461        arity: usize,
2462    ) -> bool {
2463        rsf.limit.is_none()
2464            && rsf.order_by.is_empty()
2465            && rsf
2466                .offset
2467                .clone()
2468                .try_into_literal_int64()
2469                .is_ok_and(|o| o == 0)
2470            && rsf.project.iter().copied().eq(0..arity)
2471    }
2472
2473    /// The HirRelationExpr is considered potentially expensive if and only if
2474    /// at least one of the following conditions is true:
2475    ///
2476    ///  - It contains at least one HirScalarExpr with a function call.
2477    ///  - It contains at least one CallTable or a Reduce operator.
2478    ///  - We run into a RecursionLimitError while analyzing the expression.
2479    ///
2480    /// !!!WARNING!!!: this method has an MirRelationExpr counterpart. The two
2481    /// should be kept in sync w.r.t. HIR ⇒ MIR lowering!
2482    pub fn could_run_expensive_function(&self) -> bool {
2483        let mut result = false;
2484        self.visit_pre(&mut |e: &HirRelationExpr| {
2485            use HirRelationExpr::*;
2486            use HirScalarExpr::*;
2487
2488            e.visit_children(|scalar: &HirScalarExpr| {
2489                scalar.visit_pre(&mut |scalar: &HirScalarExpr| {
2490                    result |= match scalar {
2491                        Column(..)
2492                        | Literal(..)
2493                        | CallUnmaterializable(..)
2494                        | If { .. }
2495                        | Parameter(..)
2496                        | Select(..)
2497                        | Exists(..) => false,
2498                        // Function calls are considered expensive
2499                        CallUnary { .. }
2500                        | CallBinary { .. }
2501                        | CallVariadic { .. }
2502                        | Windowing(..) => true,
2503                    };
2504                })
2505            });
2506
2507            // CallTable has a table function; Reduce has an aggregate function.
2508            // Other constructs use MirScalarExpr to run a function
2509            result |= matches!(e, CallTable { .. } | Reduce { .. });
2510        });
2511
2512        result
2513    }
2514
2515    /// Whether the expression contains an [`UnmaterializableFunc::MzNow`] call.
2516    pub fn contains_temporal(&self) -> bool {
2517        let mut contains = false;
2518        self.visit_post(&mut |expr| {
2519            expr.visit_children(|expr: &HirScalarExpr| {
2520                contains = contains || expr.contains_temporal()
2521            })
2522        });
2523        contains
2524    }
2525
2526    /// Whether the expression contains any [`UnmaterializableFunc`] call.
2527    pub fn contains_unmaterializable(&self) -> bool {
2528        let mut contains = false;
2529        self.visit_post(&mut |expr| {
2530            expr.visit_children(|expr: &HirScalarExpr| {
2531                contains = contains || expr.contains_unmaterializable()
2532            })
2533        });
2534        contains
2535    }
2536
2537    /// Whether the expression contains any [`UnmaterializableFunc`] call other than
2538    /// [`UnmaterializableFunc::MzNow`].
2539    pub fn contains_unmaterializable_except_temporal(&self) -> bool {
2540        let mut contains = false;
2541        self.visit_post(&mut |expr| {
2542            expr.visit_children(|expr: &HirScalarExpr| {
2543                contains = contains || expr.contains_unmaterializable_except_temporal()
2544            })
2545        });
2546        contains
2547    }
2548}
2549
2550impl CollectionPlan for HirRelationExpr {
2551    /// Collects the global collections that this HIR expression directly depends on, i.e., that it
2552    /// has a `Get` for. (It does _not_ traverse view definitions transitively.)
2553    /// (It does explore inside subqueries.)
2554    ///
2555    /// !!!WARNING!!!: this method has an MirRelationExpr counterpart. The two
2556    /// should be kept in sync w.r.t. HIR ⇒ MIR lowering!
2557    fn depends_on_into(&self, out: &mut BTreeSet<GlobalId>) {
2558        if let Self::Get {
2559            id: Id::Global(id), ..
2560        } = self
2561        {
2562            out.insert(*id);
2563        }
2564        self.visit_children(|expr: &HirRelationExpr| expr.depends_on_into(out))
2565    }
2566}
2567
2568/// In addition to direct relation children of `self`, this also yields every
2569/// subquery (`Exists` / `Select`) reachable through `self`'s scalar children,
2570/// at any depth of scalar nesting. This is the asymmetry warned about on the
2571/// [`VisitChildren`] trait; the matching impl on `HirScalarExpr` deliberately
2572/// does not do the symmetric thing, to avoid mutual recursion.
2573impl VisitChildren<Self> for HirRelationExpr {
2574    fn visit_children<F>(&self, mut f: F)
2575    where
2576        F: FnMut(&Self),
2577    {
2578        // subqueries of type HirRelationExpr might be wrapped in
2579        // Exists or Select variants within HirScalarExpr trees
2580        // attached at the current node, and we want to visit them as well
2581        VisitChildren::visit_children(self, |expr: &HirScalarExpr| {
2582            expr.visit_direct_subqueries(&mut f);
2583        });
2584
2585        use HirRelationExpr::*;
2586        match self {
2587            Constant { rows: _, typ: _ } | Get { id: _, typ: _ } => (),
2588            Let {
2589                name: _,
2590                id: _,
2591                value,
2592                body,
2593            } => {
2594                f(value);
2595                f(body);
2596            }
2597            LetRec {
2598                limit: _,
2599                bindings,
2600                body,
2601            } => {
2602                for (_, _, value, _) in bindings.iter() {
2603                    f(value);
2604                }
2605                f(body);
2606            }
2607            Project { input, outputs: _ } => f(input),
2608            Map { input, scalars: _ } => {
2609                f(input);
2610            }
2611            CallTable { func: _, exprs: _ } => (),
2612            Filter {
2613                input,
2614                predicates: _,
2615            } => {
2616                f(input);
2617            }
2618            Join {
2619                left,
2620                right,
2621                on: _,
2622                kind: _,
2623            } => {
2624                f(left);
2625                f(right);
2626            }
2627            Reduce {
2628                input,
2629                group_key: _,
2630                aggregates: _,
2631                expected_group_size: _,
2632            } => {
2633                f(input);
2634            }
2635            Distinct { input }
2636            | TopK {
2637                input,
2638                group_key: _,
2639                order_key: _,
2640                limit: _,
2641                offset: _,
2642                expected_group_size: _,
2643            }
2644            | Negate { input }
2645            | Threshold { input } => {
2646                f(input);
2647            }
2648            Union { base, inputs } => {
2649                f(base);
2650                for input in inputs {
2651                    f(input);
2652                }
2653            }
2654        }
2655    }
2656
2657    fn visit_mut_children<F>(&mut self, mut f: F)
2658    where
2659        F: FnMut(&mut Self),
2660    {
2661        // subqueries of type HirRelationExpr might be wrapped in
2662        // Exists or Select variants within HirScalarExpr trees
2663        // attached at the current node, and we want to visit them as well
2664        VisitChildren::visit_mut_children(self, |expr: &mut HirScalarExpr| {
2665            expr.visit_direct_subqueries_mut(&mut f);
2666        });
2667
2668        use HirRelationExpr::*;
2669        match self {
2670            Constant { rows: _, typ: _ } | Get { id: _, typ: _ } => (),
2671            Let {
2672                name: _,
2673                id: _,
2674                value,
2675                body,
2676            } => {
2677                f(value);
2678                f(body);
2679            }
2680            LetRec {
2681                limit: _,
2682                bindings,
2683                body,
2684            } => {
2685                for (_, _, value, _) in bindings.iter_mut() {
2686                    f(value);
2687                }
2688                f(body);
2689            }
2690            Project { input, outputs: _ } => f(input),
2691            Map { input, scalars: _ } => {
2692                f(input);
2693            }
2694            CallTable { func: _, exprs: _ } => (),
2695            Filter {
2696                input,
2697                predicates: _,
2698            } => {
2699                f(input);
2700            }
2701            Join {
2702                left,
2703                right,
2704                on: _,
2705                kind: _,
2706            } => {
2707                f(left);
2708                f(right);
2709            }
2710            Reduce {
2711                input,
2712                group_key: _,
2713                aggregates: _,
2714                expected_group_size: _,
2715            } => {
2716                f(input);
2717            }
2718            Distinct { input }
2719            | TopK {
2720                input,
2721                group_key: _,
2722                order_key: _,
2723                limit: _,
2724                offset: _,
2725                expected_group_size: _,
2726            }
2727            | Negate { input }
2728            | Threshold { input } => {
2729                f(input);
2730            }
2731            Union { base, inputs } => {
2732                f(base);
2733                for input in inputs {
2734                    f(input);
2735                }
2736            }
2737        }
2738    }
2739
2740    fn try_visit_children<F, E>(&self, mut f: F) -> Result<(), E>
2741    where
2742        F: FnMut(&Self) -> Result<(), E>,
2743    {
2744        // subqueries of type HirRelationExpr might be wrapped in
2745        // Exists or Select variants within HirScalarExpr trees
2746        // attached at the current node, and we want to visit them as well
2747        VisitChildren::try_visit_children(self, |expr: &HirScalarExpr| {
2748            expr.try_visit_direct_subqueries(&mut f)
2749        })?;
2750
2751        use HirRelationExpr::*;
2752        match self {
2753            Constant { rows: _, typ: _ } | Get { id: _, typ: _ } => (),
2754            Let {
2755                name: _,
2756                id: _,
2757                value,
2758                body,
2759            } => {
2760                f(value)?;
2761                f(body)?;
2762            }
2763            LetRec {
2764                limit: _,
2765                bindings,
2766                body,
2767            } => {
2768                for (_, _, value, _) in bindings.iter() {
2769                    f(value)?;
2770                }
2771                f(body)?;
2772            }
2773            Project { input, outputs: _ } => f(input)?,
2774            Map { input, scalars: _ } => {
2775                f(input)?;
2776            }
2777            CallTable { func: _, exprs: _ } => (),
2778            Filter {
2779                input,
2780                predicates: _,
2781            } => {
2782                f(input)?;
2783            }
2784            Join {
2785                left,
2786                right,
2787                on: _,
2788                kind: _,
2789            } => {
2790                f(left)?;
2791                f(right)?;
2792            }
2793            Reduce {
2794                input,
2795                group_key: _,
2796                aggregates: _,
2797                expected_group_size: _,
2798            } => {
2799                f(input)?;
2800            }
2801            Distinct { input }
2802            | TopK {
2803                input,
2804                group_key: _,
2805                order_key: _,
2806                limit: _,
2807                offset: _,
2808                expected_group_size: _,
2809            }
2810            | Negate { input }
2811            | Threshold { input } => {
2812                f(input)?;
2813            }
2814            Union { base, inputs } => {
2815                f(base)?;
2816                for input in inputs {
2817                    f(input)?;
2818                }
2819            }
2820        }
2821        Ok(())
2822    }
2823
2824    fn try_visit_mut_children<F, E>(&mut self, mut f: F) -> Result<(), E>
2825    where
2826        F: FnMut(&mut Self) -> Result<(), E>,
2827    {
2828        // subqueries of type HirRelationExpr might be wrapped in
2829        // Exists or Select variants within HirScalarExpr trees
2830        // attached at the current node, and we want to visit them as well
2831        VisitChildren::try_visit_mut_children(self, |expr: &mut HirScalarExpr| {
2832            expr.try_visit_direct_subqueries_mut(&mut f)
2833        })?;
2834
2835        use HirRelationExpr::*;
2836        match self {
2837            Constant { rows: _, typ: _ } | Get { id: _, typ: _ } => (),
2838            Let {
2839                name: _,
2840                id: _,
2841                value,
2842                body,
2843            } => {
2844                f(value)?;
2845                f(body)?;
2846            }
2847            LetRec {
2848                limit: _,
2849                bindings,
2850                body,
2851            } => {
2852                for (_, _, value, _) in bindings.iter_mut() {
2853                    f(value)?;
2854                }
2855                f(body)?;
2856            }
2857            Project { input, outputs: _ } => f(input)?,
2858            Map { input, scalars: _ } => {
2859                f(input)?;
2860            }
2861            CallTable { func: _, exprs: _ } => (),
2862            Filter {
2863                input,
2864                predicates: _,
2865            } => {
2866                f(input)?;
2867            }
2868            Join {
2869                left,
2870                right,
2871                on: _,
2872                kind: _,
2873            } => {
2874                f(left)?;
2875                f(right)?;
2876            }
2877            Reduce {
2878                input,
2879                group_key: _,
2880                aggregates: _,
2881                expected_group_size: _,
2882            } => {
2883                f(input)?;
2884            }
2885            Distinct { input }
2886            | TopK {
2887                input,
2888                group_key: _,
2889                order_key: _,
2890                limit: _,
2891                offset: _,
2892                expected_group_size: _,
2893            }
2894            | Negate { input }
2895            | Threshold { input } => {
2896                f(input)?;
2897            }
2898            Union { base, inputs } => {
2899                f(base)?;
2900                for input in inputs {
2901                    f(input)?;
2902                }
2903            }
2904        }
2905        Ok(())
2906    }
2907
2908    fn children<'a>(&'a self) -> impl DoubleEndedIterator<Item = &'a Self>
2909    where
2910        Self: 'a,
2911    {
2912        // we visit subqueries _first_, then the input
2913        let mut v: Vec<&HirRelationExpr> = vec![];
2914        use HirRelationExpr::*;
2915        match self {
2916            Constant { rows: _, typ: _ } | Get { id: _, typ: _ } => (),
2917            Let {
2918                name: _,
2919                id: _,
2920                value,
2921                body,
2922            } => {
2923                v.push(&*value);
2924                v.push(&*body);
2925            }
2926            LetRec {
2927                limit: _,
2928                bindings,
2929                body,
2930            } => {
2931                v.extend(bindings.iter().map(|(_, _, value, _)| value));
2932                v.push(&*body);
2933            }
2934            Map { input, scalars }
2935            | Filter {
2936                input,
2937                predicates: scalars,
2938            } => {
2939                for scalar in scalars {
2940                    v.append(&mut scalar.direct_subqueries());
2941                }
2942                v.push(&*input);
2943            }
2944            Reduce {
2945                input,
2946                group_key: _,
2947                aggregates,
2948                expected_group_size: _,
2949            } => {
2950                for agg in aggregates {
2951                    v.append(&mut agg.expr.direct_subqueries());
2952                }
2953                v.push(&*input);
2954            }
2955            TopK {
2956                input,
2957                group_key: _,
2958                order_key: _,
2959                limit,
2960                offset,
2961                expected_group_size: _,
2962            } => {
2963                if let Some(limit) = limit {
2964                    v.append(&mut limit.direct_subqueries());
2965                }
2966                v.append(&mut offset.direct_subqueries());
2967                v.push(&*input);
2968            }
2969            Project { input, outputs: _ }
2970            | Distinct { input }
2971            | Negate { input }
2972            | Threshold { input } => v.push(&*input),
2973            CallTable { func: _, exprs } => v.extend(
2974                exprs
2975                    .iter()
2976                    .map(|scalar| scalar.direct_subqueries())
2977                    .flatten(),
2978            ),
2979            Join {
2980                left,
2981                right,
2982                on,
2983                kind: _,
2984            } => {
2985                v.append(&mut on.direct_subqueries());
2986                v.push(&*left);
2987                v.push(&*right);
2988            }
2989            Union { base, inputs } => {
2990                v.push(&*base);
2991                v.extend(inputs.iter());
2992            }
2993        }
2994
2995        v.into_iter()
2996    }
2997
2998    fn children_mut<'a>(&'a mut self) -> impl DoubleEndedIterator<Item = &'a mut Self>
2999    where
3000        Self: 'a,
3001    {
3002        // we visit subqueries _first_, then the input
3003        let mut v = vec![];
3004        use HirRelationExpr::*;
3005        match self {
3006            Constant { rows: _, typ: _ } | Get { id: _, typ: _ } => (),
3007            Let {
3008                name: _,
3009                id: _,
3010                value,
3011                body,
3012            } => {
3013                v.push(&mut **value);
3014                v.push(&mut **body);
3015            }
3016            LetRec {
3017                limit: _,
3018                bindings,
3019                body,
3020            } => {
3021                v.extend(bindings.iter_mut().map(|(_, _, value, _)| value));
3022                v.push(&mut **body);
3023            }
3024            Map { input, scalars }
3025            | Filter {
3026                input,
3027                predicates: scalars,
3028            } => {
3029                for scalar in scalars {
3030                    v.append(&mut scalar.direct_subqueries_mut());
3031                }
3032                v.push(&mut **input);
3033            }
3034            Reduce {
3035                input,
3036                group_key: _,
3037                aggregates,
3038                expected_group_size: _,
3039            } => {
3040                for agg in aggregates {
3041                    v.append(&mut agg.expr.direct_subqueries_mut());
3042                }
3043                v.push(&mut **input);
3044            }
3045            TopK {
3046                input,
3047                group_key: _,
3048                order_key: _,
3049                limit,
3050                offset,
3051                expected_group_size: _,
3052            } => {
3053                if let Some(limit) = limit {
3054                    v.append(&mut limit.direct_subqueries_mut());
3055                }
3056                v.append(&mut offset.direct_subqueries_mut());
3057                v.push(&mut **input);
3058            }
3059            Project { input, outputs: _ }
3060            | Distinct { input }
3061            | Negate { input }
3062            | Threshold { input } => v.push(&mut **input),
3063            CallTable { func: _, exprs } => v.extend(
3064                exprs
3065                    .iter_mut()
3066                    .map(|scalar| scalar.direct_subqueries_mut())
3067                    .flatten(),
3068            ),
3069            Join {
3070                left,
3071                right,
3072                on,
3073                kind: _,
3074            } => {
3075                v.append(&mut on.direct_subqueries_mut());
3076                v.push(&mut **left);
3077                v.push(&mut **right);
3078            }
3079            Union { base, inputs } => {
3080                v.push(&mut **base);
3081                v.extend(inputs.iter_mut());
3082            }
3083        }
3084
3085        v.into_iter()
3086    }
3087}
3088
3089/// Yields the scalars directly attached to relation nodes (e.g. `Map.scalars`,
3090/// `Filter.predicates`, `Join.on`, `Reduce` aggregate args, `TopK.{limit,
3091/// offset}`, `CallTable.exprs`); does not descend into them.
3092impl VisitChildren<HirScalarExpr> for HirRelationExpr {
3093    fn visit_children<F>(&self, mut f: F)
3094    where
3095        F: FnMut(&HirScalarExpr),
3096    {
3097        use HirRelationExpr::*;
3098        match self {
3099            Constant { rows: _, typ: _ }
3100            | Get { id: _, typ: _ }
3101            | Let {
3102                name: _,
3103                id: _,
3104                value: _,
3105                body: _,
3106            }
3107            | LetRec {
3108                limit: _,
3109                bindings: _,
3110                body: _,
3111            }
3112            | Project {
3113                input: _,
3114                outputs: _,
3115            } => (),
3116            Map { input: _, scalars } => {
3117                for scalar in scalars {
3118                    f(scalar);
3119                }
3120            }
3121            CallTable { func: _, exprs } => {
3122                for expr in exprs {
3123                    f(expr);
3124                }
3125            }
3126            Filter {
3127                input: _,
3128                predicates,
3129            } => {
3130                for predicate in predicates {
3131                    f(predicate);
3132                }
3133            }
3134            Join {
3135                left: _,
3136                right: _,
3137                on,
3138                kind: _,
3139            } => f(on),
3140            Reduce {
3141                input: _,
3142                group_key: _,
3143                aggregates,
3144                expected_group_size: _,
3145            } => {
3146                for aggregate in aggregates {
3147                    f(aggregate.expr.as_ref());
3148                }
3149            }
3150            TopK {
3151                input: _,
3152                group_key: _,
3153                order_key: _,
3154                limit,
3155                offset,
3156                expected_group_size: _,
3157            } => {
3158                if let Some(limit) = limit {
3159                    f(limit)
3160                }
3161                f(offset)
3162            }
3163            Distinct { input: _ }
3164            | Negate { input: _ }
3165            | Threshold { input: _ }
3166            | Union { base: _, inputs: _ } => (),
3167        }
3168    }
3169
3170    fn visit_mut_children<F>(&mut self, mut f: F)
3171    where
3172        F: FnMut(&mut HirScalarExpr),
3173    {
3174        use HirRelationExpr::*;
3175        match self {
3176            Constant { rows: _, typ: _ }
3177            | Get { id: _, typ: _ }
3178            | Let {
3179                name: _,
3180                id: _,
3181                value: _,
3182                body: _,
3183            }
3184            | LetRec {
3185                limit: _,
3186                bindings: _,
3187                body: _,
3188            }
3189            | Project {
3190                input: _,
3191                outputs: _,
3192            } => (),
3193            Map { input: _, scalars } => {
3194                for scalar in scalars {
3195                    f(scalar);
3196                }
3197            }
3198            CallTable { func: _, exprs } => {
3199                for expr in exprs {
3200                    f(expr);
3201                }
3202            }
3203            Filter {
3204                input: _,
3205                predicates,
3206            } => {
3207                for predicate in predicates {
3208                    f(predicate);
3209                }
3210            }
3211            Join {
3212                left: _,
3213                right: _,
3214                on,
3215                kind: _,
3216            } => f(on),
3217            Reduce {
3218                input: _,
3219                group_key: _,
3220                aggregates,
3221                expected_group_size: _,
3222            } => {
3223                for aggregate in aggregates {
3224                    f(aggregate.expr.as_mut());
3225                }
3226            }
3227            TopK {
3228                input: _,
3229                group_key: _,
3230                order_key: _,
3231                limit,
3232                offset,
3233                expected_group_size: _,
3234            } => {
3235                if let Some(limit) = limit {
3236                    f(limit)
3237                }
3238                f(offset)
3239            }
3240            Distinct { input: _ }
3241            | Negate { input: _ }
3242            | Threshold { input: _ }
3243            | Union { base: _, inputs: _ } => (),
3244        }
3245    }
3246
3247    fn try_visit_children<F, E>(&self, mut f: F) -> Result<(), E>
3248    where
3249        F: FnMut(&HirScalarExpr) -> Result<(), E>,
3250    {
3251        use HirRelationExpr::*;
3252        match self {
3253            Constant { rows: _, typ: _ }
3254            | Get { id: _, typ: _ }
3255            | Let {
3256                name: _,
3257                id: _,
3258                value: _,
3259                body: _,
3260            }
3261            | LetRec {
3262                limit: _,
3263                bindings: _,
3264                body: _,
3265            }
3266            | Project {
3267                input: _,
3268                outputs: _,
3269            } => (),
3270            Map { input: _, scalars } => {
3271                for scalar in scalars {
3272                    f(scalar)?;
3273                }
3274            }
3275            CallTable { func: _, exprs } => {
3276                for expr in exprs {
3277                    f(expr)?;
3278                }
3279            }
3280            Filter {
3281                input: _,
3282                predicates,
3283            } => {
3284                for predicate in predicates {
3285                    f(predicate)?;
3286                }
3287            }
3288            Join {
3289                left: _,
3290                right: _,
3291                on,
3292                kind: _,
3293            } => f(on)?,
3294            Reduce {
3295                input: _,
3296                group_key: _,
3297                aggregates,
3298                expected_group_size: _,
3299            } => {
3300                for aggregate in aggregates {
3301                    f(aggregate.expr.as_ref())?;
3302                }
3303            }
3304            TopK {
3305                input: _,
3306                group_key: _,
3307                order_key: _,
3308                limit,
3309                offset,
3310                expected_group_size: _,
3311            } => {
3312                if let Some(limit) = limit {
3313                    f(limit)?
3314                }
3315                f(offset)?
3316            }
3317            Distinct { input: _ }
3318            | Negate { input: _ }
3319            | Threshold { input: _ }
3320            | Union { base: _, inputs: _ } => (),
3321        }
3322        Ok(())
3323    }
3324
3325    fn try_visit_mut_children<F, E>(&mut self, mut f: F) -> Result<(), E>
3326    where
3327        F: FnMut(&mut HirScalarExpr) -> Result<(), E>,
3328    {
3329        use HirRelationExpr::*;
3330        match self {
3331            Constant { rows: _, typ: _ }
3332            | Get { id: _, typ: _ }
3333            | Let {
3334                name: _,
3335                id: _,
3336                value: _,
3337                body: _,
3338            }
3339            | LetRec {
3340                limit: _,
3341                bindings: _,
3342                body: _,
3343            }
3344            | Project {
3345                input: _,
3346                outputs: _,
3347            } => (),
3348            Map { input: _, scalars } => {
3349                for scalar in scalars {
3350                    f(scalar)?;
3351                }
3352            }
3353            CallTable { func: _, exprs } => {
3354                for expr in exprs {
3355                    f(expr)?;
3356                }
3357            }
3358            Filter {
3359                input: _,
3360                predicates,
3361            } => {
3362                for predicate in predicates {
3363                    f(predicate)?;
3364                }
3365            }
3366            Join {
3367                left: _,
3368                right: _,
3369                on,
3370                kind: _,
3371            } => f(on)?,
3372            Reduce {
3373                input: _,
3374                group_key: _,
3375                aggregates,
3376                expected_group_size: _,
3377            } => {
3378                for aggregate in aggregates {
3379                    f(aggregate.expr.as_mut())?;
3380                }
3381            }
3382            TopK {
3383                input: _,
3384                group_key: _,
3385                order_key: _,
3386                limit,
3387                offset,
3388                expected_group_size: _,
3389            } => {
3390                if let Some(limit) = limit {
3391                    f(limit)?
3392                }
3393                f(offset)?
3394            }
3395            Distinct { input: _ }
3396            | Negate { input: _ }
3397            | Threshold { input: _ }
3398            | Union { base: _, inputs: _ } => (),
3399        }
3400        Ok(())
3401    }
3402
3403    fn children<'a>(&'a self) -> impl DoubleEndedIterator<Item = &'a HirScalarExpr>
3404    where
3405        HirScalarExpr: 'a,
3406    {
3407        use HirRelationExpr::*;
3408        match self {
3409            Constant { rows: _, typ: _ }
3410            | Get { id: _, typ: _ }
3411            | Let {
3412                name: _,
3413                id: _,
3414                value: _,
3415                body: _,
3416            }
3417            | LetRec {
3418                limit: _,
3419                bindings: _,
3420                body: _,
3421            }
3422            | Project {
3423                input: _,
3424                outputs: _,
3425            }
3426            | Distinct { input: _ }
3427            | Negate { input: _ }
3428            | Threshold { input: _ }
3429            | Union { base: _, inputs: _ } => vec![],
3430            Map { input: _, scalars }
3431            | CallTable {
3432                func: _,
3433                exprs: scalars,
3434            }
3435            | Filter {
3436                input: _,
3437                predicates: scalars,
3438            } => scalars.iter().collect(),
3439            Join {
3440                left: _,
3441                right: _,
3442                on,
3443                kind: _,
3444            } => vec![on],
3445            Reduce {
3446                input: _,
3447                group_key: _,
3448                aggregates,
3449                expected_group_size: _,
3450            } => aggregates.iter().map(|agg| &*agg.expr).collect(),
3451            TopK {
3452                input: _,
3453                group_key: _,
3454                order_key: _,
3455                limit,
3456                offset,
3457                expected_group_size: _,
3458            } => limit.iter().chain(std::iter::once(offset)).collect(),
3459        }
3460        .into_iter()
3461    }
3462
3463    fn children_mut<'a>(&'a mut self) -> impl DoubleEndedIterator<Item = &'a mut HirScalarExpr>
3464    where
3465        HirScalarExpr: 'a,
3466    {
3467        use HirRelationExpr::*;
3468        match self {
3469            Constant { rows: _, typ: _ }
3470            | Get { id: _, typ: _ }
3471            | Let {
3472                name: _,
3473                id: _,
3474                value: _,
3475                body: _,
3476            }
3477            | LetRec {
3478                limit: _,
3479                bindings: _,
3480                body: _,
3481            }
3482            | Project {
3483                input: _,
3484                outputs: _,
3485            }
3486            | Distinct { input: _ }
3487            | Negate { input: _ }
3488            | Threshold { input: _ }
3489            | Union { base: _, inputs: _ } => vec![],
3490            Map { input: _, scalars }
3491            | CallTable {
3492                func: _,
3493                exprs: scalars,
3494            }
3495            | Filter {
3496                input: _,
3497                predicates: scalars,
3498            } => scalars.iter_mut().collect(),
3499            Join {
3500                left: _,
3501                right: _,
3502                on,
3503                kind: _,
3504            } => vec![on],
3505            Reduce {
3506                input: _,
3507                group_key: _,
3508                aggregates,
3509                expected_group_size: _,
3510            } => aggregates.iter_mut().map(|agg| &mut *agg.expr).collect(),
3511            TopK {
3512                input: _,
3513                group_key: _,
3514                order_key: _,
3515                limit,
3516                offset,
3517                expected_group_size: _,
3518            } => limit.iter_mut().chain(std::iter::once(offset)).collect(),
3519        }
3520        .into_iter()
3521    }
3522}
3523
3524impl HirScalarExpr {
3525    pub fn name(&self) -> Option<Arc<str>> {
3526        use HirScalarExpr::*;
3527        match self {
3528            Column(_, name)
3529            | Parameter(_, name)
3530            | Literal(_, _, name)
3531            | CallUnmaterializable(_, name)
3532            | CallUnary { name, .. }
3533            | CallBinary { name, .. }
3534            | CallVariadic { name, .. }
3535            | If { name, .. }
3536            | Exists(_, name)
3537            | Select(_, name)
3538            | Windowing(_, name) => name.0.clone(),
3539        }
3540    }
3541
3542    /// Visit every subquery, without descending into subqueries.
3543    pub fn visit_direct_subqueries<F>(&self, mut f: F)
3544    where
3545        F: FnMut(&HirRelationExpr),
3546    {
3547        self.visit_post(&mut |e| {
3548            VisitChildren::<HirRelationExpr>::visit_children(e, &mut f);
3549        });
3550    }
3551
3552    /// Mutable counterpart of [`HirScalarExpr::visit_direct_subqueries`].
3553    pub fn visit_direct_subqueries_mut<F>(&mut self, mut f: F)
3554    where
3555        F: FnMut(&mut HirRelationExpr),
3556    {
3557        self.visit_mut_post(&mut |e| {
3558            VisitChildren::<HirRelationExpr>::visit_mut_children(e, &mut f);
3559        });
3560    }
3561
3562    /// Fallible counterpart of [`HirScalarExpr::visit_direct_subqueries`].
3563    pub fn try_visit_direct_subqueries<F, E>(&self, mut f: F) -> Result<(), E>
3564    where
3565        F: FnMut(&HirRelationExpr) -> Result<(), E>,
3566    {
3567        self.try_visit_post(&mut |e| {
3568            VisitChildren::<HirRelationExpr>::try_visit_children(e, &mut f)
3569        })
3570    }
3571
3572    /// Fallible mutable counterpart of [`HirScalarExpr::visit_direct_subqueries`].
3573    pub fn try_visit_direct_subqueries_mut<F, E>(&mut self, mut f: F) -> Result<(), E>
3574    where
3575        F: FnMut(&mut HirRelationExpr) -> Result<(), E>,
3576    {
3577        self.try_visit_mut_post(&mut |e| {
3578            VisitChildren::<HirRelationExpr>::try_visit_mut_children(e, &mut f)
3579        })
3580    }
3581
3582    /// Replaces parameter references in the expression with the corresponding datum from `params`.
3583    /// Additionally, it simplifies OFFSET clauses to constants after parameter binding, and checks
3584    /// them for non-negativity.
3585    ///
3586    /// This handles subqueries nested inside `self` by calling back into
3587    /// [`HirRelationExpr::bind_parameters_and_simplify_offset`] on each direct subquery; that
3588    /// relation-level function in turn calls back here for the scalar expressions it holds, so
3589    /// the two functions corecursively cover the entire HIR tree.
3590    pub fn bind_parameters_and_simplify_offset(
3591        &mut self,
3592        scx: &StatementContext,
3593        lifetime: QueryLifetime,
3594        params: &Params,
3595    ) -> Result<(), PlanError> {
3596        // First, rewrite each `Parameter` node to a literal. This walks only the scalar tree;
3597        // it does not yet descend into subqueries.
3598        self.try_visit_mut_post(&mut |e: &mut HirScalarExpr| {
3599            if let HirScalarExpr::Parameter(n, name) = e {
3600                let datum = match params.datums.iter().nth(*n - 1) {
3601                    None => return Err(PlanError::UnknownParameter(*n)),
3602                    Some(datum) => datum,
3603                };
3604                let scalar_type = &params.execute_types[*n - 1];
3605                let row = Row::pack([datum]);
3606                let column_type = scalar_type.clone().nullable(datum.is_null());
3607
3608                let name = if let Some(name) = &name.0 {
3609                    Some(Arc::clone(name))
3610                } else {
3611                    Some(Arc::from(format!("${n}")))
3612                };
3613
3614                let qcx = QueryContext::root(scx, lifetime);
3615                let ecx = execute_expr_context(&qcx);
3616
3617                *e = plan_cast(
3618                    &ecx,
3619                    *EXECUTE_CAST_CONTEXT,
3620                    HirScalarExpr::Literal(row, column_type, TreatAsEqual(name)),
3621                    &params.expected_types[*n - 1],
3622                )
3623                .expect("checked in plan_params");
3624            }
3625            Ok(())
3626        })?;
3627        // Then descend into any subqueries; the relation-side `bind_parameters_and_simplify_offset`
3628        // handles corecursion back into scalars.
3629        self.try_visit_direct_subqueries_mut(|r: &mut HirRelationExpr| {
3630            r.bind_parameters_and_simplify_offset(scx, lifetime, params)
3631        })
3632    }
3633
3634    /// Like [`HirScalarExpr::bind_parameters_and_simplify_offset`], except that parameters are
3635    /// replaced with the corresponding expression fragment from `params` rather
3636    /// than a datum.
3637    ///
3638    /// Specifically, the parameter `$1` will be replaced with `params[0]`, the
3639    /// parameter `$2` will be replaced with `params[1]`, and so on. Parameters
3640    /// in `self` that refer to invalid indices of `params` will cause a panic.
3641    ///
3642    /// Column references in parameters will be corrected to account for the
3643    /// depth at which they are spliced.
3644    pub fn splice_parameters(&mut self, params: &[HirScalarExpr], depth: usize) {
3645        #[allow(deprecated)]
3646        let _ = self.visit_recursively_mut(depth, &mut |depth: usize,
3647                                                        e: &mut HirScalarExpr|
3648         -> Result<(), ()> {
3649            if let HirScalarExpr::Parameter(i, _name) = e {
3650                *e = params[*i - 1].clone();
3651                // Correct any column references in the parameter expression for
3652                // its new depth.
3653                e.visit_columns_mut(0, &mut |d: usize, col: &mut ColumnRef| {
3654                    if col.level >= d {
3655                        col.level += depth
3656                    }
3657                });
3658            }
3659            Ok(())
3660        });
3661    }
3662
3663    /// Whether the expression contains an [`UnmaterializableFunc::MzNow`] call.
3664    pub fn contains_temporal(&self) -> bool {
3665        let mut contains = false;
3666        self.visit_post(&mut |e| {
3667            if let Self::CallUnmaterializable(UnmaterializableFunc::MzNow, _name) = e {
3668                contains = true;
3669            }
3670        });
3671        contains
3672    }
3673
3674    /// Whether the expression contains any [`UnmaterializableFunc`] call.
3675    pub fn contains_unmaterializable(&self) -> bool {
3676        let mut contains = false;
3677        self.visit_post(&mut |e| {
3678            if let Self::CallUnmaterializable(_, _) = e {
3679                contains = true;
3680            }
3681        });
3682        contains
3683    }
3684
3685    /// Whether the expression contains any [`UnmaterializableFunc`] call other than
3686    /// [`UnmaterializableFunc::MzNow`].
3687    pub fn contains_unmaterializable_except_temporal(&self) -> bool {
3688        let mut contains = false;
3689        self.visit_post(&mut |e| {
3690            if let Self::CallUnmaterializable(f, _) = e {
3691                if *f != UnmaterializableFunc::MzNow {
3692                    contains = true;
3693                }
3694            }
3695        });
3696        contains
3697    }
3698
3699    /// Constructs an unnamed column reference in the current scope.
3700    /// Use [`HirScalarExpr::named_column`] when a name is known.
3701    /// Use [`HirScalarExpr::unnamed_column`] for a `ColumnRef`.
3702    pub fn column(index: usize) -> HirScalarExpr {
3703        HirScalarExpr::Column(
3704            ColumnRef {
3705                level: 0,
3706                column: index,
3707            },
3708            TreatAsEqual(None),
3709        )
3710    }
3711
3712    /// Constructs an unnamed column reference.
3713    pub fn unnamed_column(cr: ColumnRef) -> HirScalarExpr {
3714        HirScalarExpr::Column(cr, TreatAsEqual(None))
3715    }
3716
3717    /// Constructs a named column reference.
3718    /// Names are interned by a `NameManager`.
3719    pub fn named_column(cr: ColumnRef, name: Arc<str>) -> HirScalarExpr {
3720        HirScalarExpr::Column(cr, TreatAsEqual(Some(name)))
3721    }
3722
3723    pub fn parameter(n: usize) -> HirScalarExpr {
3724        HirScalarExpr::Parameter(n, TreatAsEqual(None))
3725    }
3726
3727    pub fn literal(datum: Datum, scalar_type: SqlScalarType) -> HirScalarExpr {
3728        let col_type = scalar_type.nullable(datum.is_null());
3729        soft_assert_or_log!(datum.is_instance_of_sql(&col_type), "type is correct");
3730        let row = Row::pack([datum]);
3731        HirScalarExpr::Literal(row, col_type, TreatAsEqual(None))
3732    }
3733
3734    pub fn literal_true() -> HirScalarExpr {
3735        HirScalarExpr::literal(Datum::True, SqlScalarType::Bool)
3736    }
3737
3738    pub fn literal_false() -> HirScalarExpr {
3739        HirScalarExpr::literal(Datum::False, SqlScalarType::Bool)
3740    }
3741
3742    pub fn literal_null(scalar_type: SqlScalarType) -> HirScalarExpr {
3743        HirScalarExpr::literal(Datum::Null, scalar_type)
3744    }
3745
3746    pub fn literal_1d_array(
3747        datums: Vec<Datum>,
3748        element_scalar_type: SqlScalarType,
3749    ) -> Result<HirScalarExpr, PlanError> {
3750        let scalar_type = match element_scalar_type {
3751            SqlScalarType::Array(_) => {
3752                sql_bail!("cannot build array from array type");
3753            }
3754            typ => SqlScalarType::Array(Box::new(typ)).nullable(false),
3755        };
3756
3757        let mut row = Row::default();
3758        row.packer()
3759            .try_push_array(
3760                &[ArrayDimension {
3761                    lower_bound: 1,
3762                    length: datums.len(),
3763                }],
3764                datums,
3765            )
3766            .expect("array constructed to be valid");
3767
3768        Ok(HirScalarExpr::Literal(row, scalar_type, TreatAsEqual(None)))
3769    }
3770
3771    pub fn as_literal(&self) -> Option<Datum<'_>> {
3772        if let HirScalarExpr::Literal(row, _column_type, _name) = self {
3773            Some(row.unpack_first())
3774        } else {
3775            None
3776        }
3777    }
3778
3779    pub fn is_literal_true(&self) -> bool {
3780        Some(Datum::True) == self.as_literal()
3781    }
3782
3783    pub fn is_literal_false(&self) -> bool {
3784        Some(Datum::False) == self.as_literal()
3785    }
3786
3787    pub fn is_literal_null(&self) -> bool {
3788        Some(Datum::Null) == self.as_literal()
3789    }
3790
3791    /// Return true iff `self` consists only of literals, materializable function calls, and
3792    /// if-else statements.
3793    pub fn is_constant(&self) -> bool {
3794        let mut worklist = vec![self];
3795        while let Some(expr) = worklist.pop() {
3796            match expr {
3797                Self::Literal(..) => {
3798                    // leaf node, do nothing
3799                }
3800                Self::CallUnary { expr, .. } => {
3801                    worklist.push(expr);
3802                }
3803                Self::CallBinary {
3804                    func: _,
3805                    expr1,
3806                    expr2,
3807                    name: _,
3808                } => {
3809                    worklist.push(expr1);
3810                    worklist.push(expr2);
3811                }
3812                Self::CallVariadic {
3813                    func: _,
3814                    exprs,
3815                    name: _,
3816                } => {
3817                    worklist.extend(exprs.iter());
3818                }
3819                // (CallUnmaterializable is not allowed)
3820                Self::If {
3821                    cond,
3822                    then,
3823                    els,
3824                    name: _,
3825                } => {
3826                    worklist.push(cond);
3827                    worklist.push(then);
3828                    worklist.push(els);
3829                }
3830                _ => {
3831                    return false; // Any other node makes `self` non-constant.
3832                }
3833            }
3834        }
3835        true
3836    }
3837
3838    pub fn call_unary(self, func: UnaryFunc) -> Self {
3839        HirScalarExpr::CallUnary {
3840            func,
3841            expr: Box::new(self),
3842            name: NameMetadata::default(),
3843        }
3844    }
3845
3846    pub fn call_binary<B: Into<BinaryFunc>>(self, other: Self, func: B) -> Self {
3847        HirScalarExpr::CallBinary {
3848            func: func.into(),
3849            expr1: Box::new(self),
3850            expr2: Box::new(other),
3851            name: NameMetadata::default(),
3852        }
3853    }
3854
3855    pub fn call_unmaterializable(func: UnmaterializableFunc) -> Self {
3856        HirScalarExpr::CallUnmaterializable(func, NameMetadata::default())
3857    }
3858
3859    pub fn call_variadic<V: Into<VariadicFunc>>(func: V, exprs: Vec<Self>) -> Self {
3860        HirScalarExpr::CallVariadic {
3861            func: func.into(),
3862            exprs,
3863            name: NameMetadata::default(),
3864        }
3865    }
3866
3867    pub fn if_then_else(cond: Self, then: Self, els: Self) -> Self {
3868        HirScalarExpr::If {
3869            cond: Box::new(cond),
3870            then: Box::new(then),
3871            els: Box::new(els),
3872            name: NameMetadata::default(),
3873        }
3874    }
3875
3876    pub fn windowing(expr: WindowExpr) -> Self {
3877        HirScalarExpr::Windowing(expr, TreatAsEqual(None))
3878    }
3879
3880    pub fn or(self, other: Self) -> Self {
3881        HirScalarExpr::call_variadic(Or, vec![self, other])
3882    }
3883
3884    pub fn and(self, other: Self) -> Self {
3885        HirScalarExpr::call_variadic(And, vec![self, other])
3886    }
3887
3888    pub fn not(self) -> Self {
3889        self.call_unary(UnaryFunc::Not(func::Not))
3890    }
3891
3892    pub fn call_is_null(self) -> Self {
3893        self.call_unary(UnaryFunc::IsNull(func::IsNull))
3894    }
3895
3896    /// Calls AND with the given arguments. Simplifies if 0 or 1 args.
3897    pub fn variadic_and(mut args: Vec<HirScalarExpr>) -> HirScalarExpr {
3898        match args.len() {
3899            0 => HirScalarExpr::literal_true(), // Same as unit_of_and_or, but that's MirScalarExpr
3900            1 => args.swap_remove(0),
3901            _ => HirScalarExpr::call_variadic(And, args),
3902        }
3903    }
3904
3905    /// Calls OR with the given arguments. Simplifies if 0 or 1 args.
3906    pub fn variadic_or(mut args: Vec<HirScalarExpr>) -> HirScalarExpr {
3907        match args.len() {
3908            0 => HirScalarExpr::literal_false(), // Same as unit_of_and_or, but that's MirScalarExpr
3909            1 => args.swap_remove(0),
3910            _ => HirScalarExpr::call_variadic(Or, args),
3911        }
3912    }
3913
3914    pub fn take(&mut self) -> Self {
3915        mem::replace(self, HirScalarExpr::literal_null(SqlScalarType::String))
3916    }
3917
3918    #[deprecated = "Redefine this based on the `Visit` and `VisitChildren` methods."]
3919    /// Visits the column references in this scalar expression.
3920    ///
3921    /// The `depth` argument should indicate the subquery nesting depth of the expression,
3922    /// which will be incremented with each subquery entered and presented to the supplied
3923    /// function `f`.
3924    pub fn visit_columns<F>(&self, depth: usize, f: &mut F)
3925    where
3926        F: FnMut(usize, &ColumnRef),
3927    {
3928        #[allow(deprecated)]
3929        let _ = self.visit_recursively(depth, &mut |depth: usize,
3930                                                    e: &HirScalarExpr|
3931         -> Result<(), ()> {
3932            if let HirScalarExpr::Column(col, _name) = e {
3933                f(depth, col)
3934            }
3935            Ok(())
3936        });
3937    }
3938
3939    #[deprecated = "Redefine this based on the `Visit` and `VisitChildren` methods."]
3940    /// Like `visit_columns`, but permits mutating the column references.
3941    pub fn visit_columns_mut<F>(&mut self, depth: usize, f: &mut F)
3942    where
3943        F: FnMut(usize, &mut ColumnRef),
3944    {
3945        #[allow(deprecated)]
3946        let _ = self.visit_recursively_mut(depth, &mut |depth: usize,
3947                                                        e: &mut HirScalarExpr|
3948         -> Result<(), ()> {
3949            if let HirScalarExpr::Column(col, _name) = e {
3950                f(depth, col)
3951            }
3952            Ok(())
3953        });
3954    }
3955
3956    /// Visits those column references in this scalar expression that refer to the root
3957    /// level. These include column references that are at the root level, as well as column
3958    /// references that are at a deeper subquery nesting depth, but refer back to the root level.
3959    /// (Note that even if `self` is embedded inside a larger expression, we consider the
3960    /// "root level" to be `self`'s level.)
3961    pub fn visit_columns_referring_to_root_level<F>(&self, f: &mut F)
3962    where
3963        F: FnMut(usize),
3964    {
3965        #[allow(deprecated)]
3966        let _ = self.visit_recursively(0, &mut |depth: usize,
3967                                                e: &HirScalarExpr|
3968         -> Result<(), ()> {
3969            if let HirScalarExpr::Column(col, _name) = e {
3970                if col.level == depth {
3971                    f(col.column)
3972                }
3973            }
3974            Ok(())
3975        });
3976    }
3977
3978    /// Like `visit_columns_referring_to_root_level`, but permits mutating the column references.
3979    pub fn visit_columns_referring_to_root_level_mut<F>(&mut self, f: &mut F)
3980    where
3981        F: FnMut(&mut usize),
3982    {
3983        #[allow(deprecated)]
3984        let _ = self.visit_recursively_mut(0, &mut |depth: usize,
3985                                                    e: &mut HirScalarExpr|
3986         -> Result<(), ()> {
3987            if let HirScalarExpr::Column(col, _name) = e {
3988                if col.level == depth {
3989                    f(&mut col.column)
3990                }
3991            }
3992            Ok(())
3993        });
3994    }
3995
3996    #[deprecated = "Redefine this based on the `Visit` and `VisitChildren` methods."]
3997    /// Like `visit` but it enters the subqueries visiting the scalar expressions contained
3998    /// in them. It takes the current depth of the expression and increases it when
3999    /// entering a subquery.
4000    pub fn visit_recursively<F, E>(&self, depth: usize, f: &mut F) -> Result<(), E>
4001    where
4002        F: FnMut(usize, &HirScalarExpr) -> Result<(), E>,
4003    {
4004        match self {
4005            HirScalarExpr::Literal(..)
4006            | HirScalarExpr::Parameter(..)
4007            | HirScalarExpr::CallUnmaterializable(..)
4008            | HirScalarExpr::Column(..) => (),
4009            HirScalarExpr::CallUnary { expr, .. } => expr.visit_recursively(depth, f)?,
4010            HirScalarExpr::CallBinary { expr1, expr2, .. } => {
4011                expr1.visit_recursively(depth, f)?;
4012                expr2.visit_recursively(depth, f)?;
4013            }
4014            HirScalarExpr::CallVariadic { exprs, .. } => {
4015                for expr in exprs {
4016                    expr.visit_recursively(depth, f)?;
4017                }
4018            }
4019            HirScalarExpr::If {
4020                cond,
4021                then,
4022                els,
4023                name: _,
4024            } => {
4025                cond.visit_recursively(depth, f)?;
4026                then.visit_recursively(depth, f)?;
4027                els.visit_recursively(depth, f)?;
4028            }
4029            HirScalarExpr::Exists(expr, _name) | HirScalarExpr::Select(expr, _name) => {
4030                #[allow(deprecated)]
4031                expr.visit_scalar_expressions(depth + 1, &mut |e, depth| {
4032                    e.visit_recursively(depth, f)
4033                })?;
4034            }
4035            HirScalarExpr::Windowing(expr, _name) => {
4036                expr.visit_expressions(&mut |e| e.visit_recursively(depth, f))?;
4037            }
4038        }
4039        f(depth, self)
4040    }
4041
4042    #[deprecated = "Redefine this based on the `Visit` and `VisitChildren` methods."]
4043    /// Like `visit_recursively`, but permits mutating the scalar expressions.
4044    pub fn visit_recursively_mut<F, E>(&mut self, depth: usize, f: &mut F) -> Result<(), E>
4045    where
4046        F: FnMut(usize, &mut HirScalarExpr) -> Result<(), E>,
4047    {
4048        match self {
4049            HirScalarExpr::Literal(..)
4050            | HirScalarExpr::Parameter(..)
4051            | HirScalarExpr::CallUnmaterializable(..)
4052            | HirScalarExpr::Column(..) => (),
4053            HirScalarExpr::CallUnary { expr, .. } => expr.visit_recursively_mut(depth, f)?,
4054            HirScalarExpr::CallBinary { expr1, expr2, .. } => {
4055                expr1.visit_recursively_mut(depth, f)?;
4056                expr2.visit_recursively_mut(depth, f)?;
4057            }
4058            HirScalarExpr::CallVariadic { exprs, .. } => {
4059                for expr in exprs {
4060                    expr.visit_recursively_mut(depth, f)?;
4061                }
4062            }
4063            HirScalarExpr::If {
4064                cond,
4065                then,
4066                els,
4067                name: _,
4068            } => {
4069                cond.visit_recursively_mut(depth, f)?;
4070                then.visit_recursively_mut(depth, f)?;
4071                els.visit_recursively_mut(depth, f)?;
4072            }
4073            HirScalarExpr::Exists(expr, _name) | HirScalarExpr::Select(expr, _name) => {
4074                #[allow(deprecated)]
4075                expr.visit_scalar_expressions_mut(depth + 1, &mut |e, depth| {
4076                    e.visit_recursively_mut(depth, f)
4077                })?;
4078            }
4079            HirScalarExpr::Windowing(expr, _name) => {
4080                expr.visit_expressions_mut(&mut |e| e.visit_recursively_mut(depth, f))?;
4081            }
4082        }
4083        f(depth, self)
4084    }
4085
4086    /// Attempts to simplify self into a literal.
4087    ///
4088    /// Returns None if self is not constant and therefore can't be simplified to a literal, or if
4089    /// an evaluation error occurs during simplification, or if self contains
4090    /// - a subquery
4091    /// - a column reference to an outer level
4092    /// - a parameter
4093    /// - a window function call
4094    fn simplify_to_literal(self) -> Option<Row> {
4095        let mut expr = self
4096            .lower_uncorrelated(crate::plan::lowering::Config::default())
4097            .ok()?;
4098        // Using MIR evaluation with repr types is fine here: the
4099        // result is an untyped Row, so any intermediate type
4100        // canonicalization is discarded.
4101        expr.reduce(&[]);
4102        match expr {
4103            mz_expr::MirScalarExpr::Literal(Ok(row), _) => Some(row),
4104            _ => None,
4105        }
4106    }
4107
4108    /// Simplifies self into a literal. If this is not possible (e.g., because self is not constant
4109    /// or an evaluation error occurs during simplification), it returns
4110    /// [`PlanError::ConstantExpressionSimplificationFailed`].
4111    ///
4112    /// The returned error is an _internal_ error if the expression contains
4113    /// - a subquery
4114    /// - a column reference to an outer level
4115    /// - a parameter
4116    /// - a window function call
4117    ///
4118    /// TODO: use this everywhere instead of `simplify_to_literal`, so that we don't hide the error
4119    /// msg.
4120    fn simplify_to_literal_with_result(self) -> Result<Row, PlanError> {
4121        let mut expr = self
4122            .lower_uncorrelated(crate::plan::lowering::Config::default())
4123            .map_err(|err| {
4124                PlanError::ConstantExpressionSimplificationFailed(err.to_string_with_causes())
4125            })?;
4126        // Using MIR evaluation with repr types is fine here: the
4127        // result is an untyped Row, so any intermediate type
4128        // canonicalization is discarded.
4129        expr.reduce(&[]);
4130        match expr {
4131            mz_expr::MirScalarExpr::Literal(Ok(row), _) => Ok(row),
4132            mz_expr::MirScalarExpr::Literal(Err(err), _) => Err(
4133                PlanError::ConstantExpressionSimplificationFailed(err.to_string_with_causes()),
4134            ),
4135            _ => Err(PlanError::ConstantExpressionSimplificationFailed(
4136                "Not a constant".to_string(),
4137            )),
4138        }
4139    }
4140
4141    /// Attempts to simplify this expression to a literal 64-bit integer.
4142    ///
4143    /// Returns `None` if this expression cannot be simplified, e.g. because it
4144    /// contains non-literal values.
4145    ///
4146    /// # Panics
4147    ///
4148    /// Panics if this expression does not have type [`SqlScalarType::Int64`].
4149    pub fn into_literal_int64(self) -> Option<i64> {
4150        self.simplify_to_literal().and_then(|row| {
4151            let datum = row.unpack_first();
4152            if datum.is_null() {
4153                None
4154            } else {
4155                Some(datum.unwrap_int64())
4156            }
4157        })
4158    }
4159
4160    /// Attempts to simplify this expression to a literal string.
4161    ///
4162    /// Returns `None` if this expression cannot be simplified, e.g. because it
4163    /// contains non-literal values.
4164    ///
4165    /// # Panics
4166    ///
4167    /// Panics if this expression does not have type [`SqlScalarType::String`].
4168    pub fn into_literal_string(self) -> Option<String> {
4169        self.simplify_to_literal().and_then(|row| {
4170            let datum = row.unpack_first();
4171            if datum.is_null() {
4172                None
4173            } else {
4174                Some(datum.unwrap_str().to_owned())
4175            }
4176        })
4177    }
4178
4179    /// Attempts to simplify this expression to a literal MzTimestamp.
4180    ///
4181    /// Returns `None` if the expression simplifies to `null` or if the expression cannot be
4182    /// simplified, e.g. because it contains non-literal values or a cast fails.
4183    ///
4184    /// TODO: Make this (and the other similar fns above) return Result, so that we can show the
4185    /// error when it fails. (E.g., there can be non-trivial cast errors.)
4186    /// See `try_into_literal_int64` as an example.
4187    ///
4188    /// # Panics
4189    ///
4190    /// Panics if this expression does not have type [`SqlScalarType::MzTimestamp`].
4191    pub fn into_literal_mz_timestamp(self) -> Option<Timestamp> {
4192        self.simplify_to_literal().and_then(|row| {
4193            let datum = row.unpack_first();
4194            if datum.is_null() {
4195                None
4196            } else {
4197                Some(datum.unwrap_mz_timestamp())
4198            }
4199        })
4200    }
4201
4202    /// Attempts to simplify this expression of [`SqlScalarType::Int64`] to a literal Int64 and
4203    /// returns it as an i64.
4204    ///
4205    /// Returns `PlanError::ConstantExpressionSimplificationFailed` if
4206    /// - it's not a constant expression (as determined by `is_constant`)
4207    /// - evaluates to null
4208    /// - an EvalError occurs during evaluation (e.g., a cast fails)
4209    ///
4210    /// # Panics
4211    ///
4212    /// Panics if this expression does not have type [`SqlScalarType::Int64`].
4213    pub fn try_into_literal_int64(self) -> Result<i64, PlanError> {
4214        match self.clone().try_into_nullable_literal_int64()? {
4215            Some(value) => Ok(value),
4216            None => Err(PlanError::ConstantExpressionSimplificationFailed(format!(
4217                "Expected an expression that evaluates to a non-null value, got {}",
4218                self
4219            ))),
4220        }
4221    }
4222
4223    /// Like [`try_into_literal_int64`](Self::try_into_literal_int64), but an expression that
4224    /// evaluates to null is `Ok(None)` instead of an error.
4225    pub fn try_into_nullable_literal_int64(self) -> Result<Option<i64>, PlanError> {
4226        // TODO: add the `is_constant` check also to all the other into_literal_... (by adding it to
4227        // `simplify_to_literal`), but those should be just soft_asserts at first that it doesn't
4228        // actually happen that it's weaker than `reduce`, and then add them for real after 1 week.
4229        // (Without the is_constant check, lower_uncorrelated's preconditions spill out to be
4230        // preconditions also of all the other into_literal_... functions.)
4231        if !self.is_constant() {
4232            return Err(PlanError::ConstantExpressionSimplificationFailed(format!(
4233                "Expected a constant expression, got {}",
4234                self
4235            )));
4236        }
4237        self.simplify_to_literal_with_result().map(|row| {
4238            let datum = row.unpack_first();
4239            if datum.is_null() {
4240                None
4241            } else {
4242                Some(datum.unwrap_int64())
4243            }
4244        })
4245    }
4246
4247    pub fn contains_parameters(&self) -> bool {
4248        let mut contains_parameters = false;
4249        #[allow(deprecated)]
4250        let _ = self.visit_recursively(0, &mut |_depth: usize,
4251                                                expr: &HirScalarExpr|
4252         -> Result<(), ()> {
4253            if let HirScalarExpr::Parameter(..) = expr {
4254                contains_parameters = true;
4255            }
4256            Ok(())
4257        });
4258        contains_parameters
4259    }
4260
4261    fn direct_subqueries(&self) -> Vec<&HirRelationExpr> {
4262        let mut subqueries: Vec<&HirRelationExpr> = vec![];
4263
4264        let mut worklist = vec![self];
4265        while let Some(elt) = worklist.pop() {
4266            match elt {
4267                HirScalarExpr::Column(_, _)
4268                | HirScalarExpr::Parameter(_, _)
4269                | HirScalarExpr::Literal(_, _, _)
4270                | HirScalarExpr::CallUnmaterializable(_, _) => (),
4271                HirScalarExpr::CallUnary {
4272                    func: _,
4273                    expr,
4274                    name: _,
4275                } => worklist.push(&*expr),
4276                HirScalarExpr::CallBinary {
4277                    func: _,
4278                    expr1,
4279                    expr2,
4280                    name: _name,
4281                } => {
4282                    // Push in reverse so children pop (and are visited) left-to-right.
4283                    worklist.push(&*expr2);
4284                    worklist.push(&*expr1);
4285                }
4286                HirScalarExpr::CallVariadic {
4287                    func: _,
4288                    exprs,
4289                    name: _name,
4290                } => {
4291                    worklist.extend(exprs.iter().rev());
4292                }
4293                HirScalarExpr::If {
4294                    cond,
4295                    then,
4296                    els,
4297                    name: _,
4298                } => {
4299                    worklist.push(&*els);
4300                    worklist.push(&*then);
4301                    worklist.push(&*cond);
4302                }
4303                HirScalarExpr::Exists(hir, _) | HirScalarExpr::Select(hir, _) => {
4304                    subqueries.push(&*hir);
4305                }
4306                HirScalarExpr::Windowing(
4307                    WindowExpr {
4308                        func,
4309                        partition_by,
4310                        order_by,
4311                    },
4312                    _,
4313                ) => {
4314                    // Push in reverse so children pop (and are visited) left-to-right:
4315                    // func args, then partition_by, then order_by.
4316                    worklist.extend(order_by.iter().rev());
4317                    worklist.extend(partition_by.iter().rev());
4318                    match func {
4319                        WindowExprType::Scalar(_) => (),
4320                        WindowExprType::Value(val) => worklist.push(&*val.args),
4321                        WindowExprType::Aggregate(agg) => worklist.push(&*agg.aggregate_expr.expr),
4322                    }
4323                }
4324            }
4325        }
4326
4327        subqueries
4328    }
4329
4330    fn direct_subqueries_mut(&mut self) -> Vec<&mut HirRelationExpr> {
4331        let mut subqueries: Vec<&mut HirRelationExpr> = vec![];
4332
4333        let mut worklist = vec![self];
4334        while let Some(elt) = worklist.pop() {
4335            match elt {
4336                HirScalarExpr::Column(_, _)
4337                | HirScalarExpr::Parameter(_, _)
4338                | HirScalarExpr::Literal(_, _, _)
4339                | HirScalarExpr::CallUnmaterializable(_, _) => (),
4340                HirScalarExpr::CallUnary {
4341                    func: _,
4342                    expr,
4343                    name: _,
4344                } => worklist.push(&mut **expr),
4345                HirScalarExpr::CallBinary {
4346                    func: _,
4347                    expr1,
4348                    expr2,
4349                    name: _name,
4350                } => {
4351                    // Push in reverse so children pop (and are visited) left-to-right.
4352                    worklist.push(&mut **expr2);
4353                    worklist.push(&mut **expr1);
4354                }
4355                HirScalarExpr::CallVariadic {
4356                    func: _,
4357                    exprs,
4358                    name: _name,
4359                } => {
4360                    worklist.extend(exprs.iter_mut().rev());
4361                }
4362                HirScalarExpr::If {
4363                    cond,
4364                    then,
4365                    els,
4366                    name: _,
4367                } => {
4368                    worklist.push(&mut **els);
4369                    worklist.push(&mut **then);
4370                    worklist.push(&mut **cond);
4371                }
4372                HirScalarExpr::Exists(hir, _) | HirScalarExpr::Select(hir, _) => {
4373                    subqueries.push(&mut **hir);
4374                }
4375                HirScalarExpr::Windowing(
4376                    WindowExpr {
4377                        func,
4378                        partition_by,
4379                        order_by,
4380                    },
4381                    _,
4382                ) => {
4383                    // Push in reverse so children pop (and are visited) left-to-right:
4384                    // func args, then partition_by, then order_by.
4385                    worklist.extend(order_by.iter_mut().rev());
4386                    worklist.extend(partition_by.iter_mut().rev());
4387                    match func {
4388                        WindowExprType::Scalar(_) => (),
4389                        WindowExprType::Value(val) => worklist.push(&mut val.args),
4390                        WindowExprType::Aggregate(agg) => {
4391                            worklist.push(&mut agg.aggregate_expr.expr)
4392                        }
4393                    }
4394                }
4395            }
4396        }
4397
4398        subqueries
4399    }
4400}
4401
4402/// Yields the direct scalar children of `self`. Stops at `Exists` / `Select`:
4403/// scalars inside subquery bodies are not surfaced. The asymmetry with
4404/// `VisitChildren<Self> for HirRelationExpr` (which does see through scalars
4405/// into subqueries) is what avoids the mutual recursion warned about on the
4406/// [`VisitChildren`] trait.
4407impl VisitChildren<Self> for HirScalarExpr {
4408    fn visit_children<F>(&self, mut f: F)
4409    where
4410        F: FnMut(&Self),
4411    {
4412        use HirScalarExpr::*;
4413        match self {
4414            Column(..) | Parameter(..) | Literal(..) | CallUnmaterializable(..) => (),
4415            CallUnary { expr, .. } => f(expr),
4416            CallBinary { expr1, expr2, .. } => {
4417                f(expr1);
4418                f(expr2);
4419            }
4420            CallVariadic { exprs, .. } => {
4421                for expr in exprs {
4422                    f(expr);
4423                }
4424            }
4425            If {
4426                cond,
4427                then,
4428                els,
4429                name: _,
4430            } => {
4431                f(cond);
4432                f(then);
4433                f(els);
4434            }
4435            Exists(..) | Select(..) => (),
4436            Windowing(expr, _name) => expr.visit_children(f),
4437        }
4438    }
4439
4440    fn visit_mut_children<F>(&mut self, mut f: F)
4441    where
4442        F: FnMut(&mut Self),
4443    {
4444        use HirScalarExpr::*;
4445        match self {
4446            Column(..) | Parameter(..) | Literal(..) | CallUnmaterializable(..) => (),
4447            CallUnary { expr, .. } => f(expr),
4448            CallBinary { expr1, expr2, .. } => {
4449                f(expr1);
4450                f(expr2);
4451            }
4452            CallVariadic { exprs, .. } => {
4453                for expr in exprs {
4454                    f(expr);
4455                }
4456            }
4457            If {
4458                cond,
4459                then,
4460                els,
4461                name: _,
4462            } => {
4463                f(cond);
4464                f(then);
4465                f(els);
4466            }
4467            Exists(..) | Select(..) => (),
4468            Windowing(expr, _name) => expr.visit_mut_children(f),
4469        }
4470    }
4471
4472    fn try_visit_children<F, E>(&self, mut f: F) -> Result<(), E>
4473    where
4474        F: FnMut(&Self) -> Result<(), E>,
4475    {
4476        use HirScalarExpr::*;
4477        match self {
4478            Column(..) | Parameter(..) | Literal(..) | CallUnmaterializable(..) => (),
4479            CallUnary { expr, .. } => f(expr)?,
4480            CallBinary { expr1, expr2, .. } => {
4481                f(expr1)?;
4482                f(expr2)?;
4483            }
4484            CallVariadic { exprs, .. } => {
4485                for expr in exprs {
4486                    f(expr)?;
4487                }
4488            }
4489            If {
4490                cond,
4491                then,
4492                els,
4493                name: _,
4494            } => {
4495                f(cond)?;
4496                f(then)?;
4497                f(els)?;
4498            }
4499            Exists(..) | Select(..) => (),
4500            Windowing(expr, _name) => expr.try_visit_children(f)?,
4501        }
4502        Ok(())
4503    }
4504
4505    fn try_visit_mut_children<F, E>(&mut self, mut f: F) -> Result<(), E>
4506    where
4507        F: FnMut(&mut Self) -> Result<(), E>,
4508    {
4509        use HirScalarExpr::*;
4510        match self {
4511            Column(..) | Parameter(..) | Literal(..) | CallUnmaterializable(..) => (),
4512            CallUnary { expr, .. } => f(expr)?,
4513            CallBinary { expr1, expr2, .. } => {
4514                f(expr1)?;
4515                f(expr2)?;
4516            }
4517            CallVariadic { exprs, .. } => {
4518                for expr in exprs {
4519                    f(expr)?;
4520                }
4521            }
4522            If {
4523                cond,
4524                then,
4525                els,
4526                name: _,
4527            } => {
4528                f(cond)?;
4529                f(then)?;
4530                f(els)?;
4531            }
4532            Exists(..) | Select(..) => (),
4533            Windowing(expr, _name) => expr.try_visit_mut_children(f)?,
4534        }
4535        Ok(())
4536    }
4537
4538    fn children<'a>(&'a self) -> impl DoubleEndedIterator<Item = &'a Self>
4539    where
4540        Self: 'a,
4541    {
4542        use HirScalarExpr::*;
4543        let v: Vec<&Self> = match self {
4544            Column(..) | Parameter(..) | Literal(..) | CallUnmaterializable(..) => vec![],
4545            CallUnary { expr, .. } => vec![&*expr],
4546            CallBinary { expr1, expr2, .. } => {
4547                vec![&*expr1, &*expr2]
4548            }
4549            CallVariadic { exprs, .. } => exprs.iter().collect(),
4550            If {
4551                cond,
4552                then,
4553                els,
4554                name: _,
4555            } => {
4556                vec![&*cond, &*then, &*els]
4557            }
4558            Exists(..) | Select(..) => vec![],
4559            Windowing(expr, _name) => expr.children().collect(),
4560        };
4561        v.into_iter()
4562    }
4563
4564    fn children_mut<'a>(&'a mut self) -> impl DoubleEndedIterator<Item = &'a mut Self>
4565    where
4566        Self: 'a,
4567    {
4568        use HirScalarExpr::*;
4569        let v: Vec<&mut Self> = match self {
4570            Column(..) | Parameter(..) | Literal(..) | CallUnmaterializable(..) => vec![],
4571            CallUnary { expr, .. } => vec![&mut **expr],
4572            CallBinary { expr1, expr2, .. } => {
4573                vec![&mut **expr1, &mut **expr2]
4574            }
4575            CallVariadic { exprs, .. } => exprs.iter_mut().collect(),
4576            If {
4577                cond,
4578                then,
4579                els,
4580                name: _,
4581            } => {
4582                vec![&mut **cond, &mut **then, &mut **els]
4583            }
4584            Exists(..) | Select(..) => vec![],
4585            Windowing(expr, _name) => expr.children_mut().collect(),
4586        };
4587        v.into_iter()
4588    }
4589}
4590
4591/// Yields the immediate `HirRelationExpr` children of `self` (the bodies of
4592/// `Exists` / `Select`); does not descend into them.
4593impl VisitChildren<HirRelationExpr> for HirScalarExpr {
4594    fn visit_children<F>(&self, mut f: F)
4595    where
4596        F: FnMut(&HirRelationExpr),
4597    {
4598        use HirScalarExpr::*;
4599        match self {
4600            Column(..)
4601            | Parameter(..)
4602            | Literal(..)
4603            | CallUnmaterializable(..)
4604            | CallUnary { .. }
4605            | CallBinary { .. }
4606            | CallVariadic { .. }
4607            | If { .. }
4608            | Windowing(..) => (),
4609            Exists(expr, _name) | Select(expr, _name) => f(expr),
4610        }
4611    }
4612
4613    fn visit_mut_children<F>(&mut self, mut f: F)
4614    where
4615        F: FnMut(&mut HirRelationExpr),
4616    {
4617        use HirScalarExpr::*;
4618        match self {
4619            Column(..)
4620            | Parameter(..)
4621            | Literal(..)
4622            | CallUnmaterializable(..)
4623            | CallUnary { .. }
4624            | CallBinary { .. }
4625            | CallVariadic { .. }
4626            | If { .. }
4627            | Windowing(..) => (),
4628            Exists(expr, _name) | Select(expr, _name) => f(expr),
4629        }
4630    }
4631
4632    fn try_visit_children<F, E>(&self, mut f: F) -> Result<(), E>
4633    where
4634        F: FnMut(&HirRelationExpr) -> Result<(), E>,
4635    {
4636        use HirScalarExpr::*;
4637        match self {
4638            Column(..)
4639            | Parameter(..)
4640            | Literal(..)
4641            | CallUnmaterializable(..)
4642            | CallUnary { .. }
4643            | CallBinary { .. }
4644            | CallVariadic { .. }
4645            | If { .. }
4646            | Windowing(..) => (),
4647            Exists(expr, _name) | Select(expr, _name) => f(expr)?,
4648        }
4649        Ok(())
4650    }
4651
4652    fn try_visit_mut_children<F, E>(&mut self, mut f: F) -> Result<(), E>
4653    where
4654        F: FnMut(&mut HirRelationExpr) -> Result<(), E>,
4655    {
4656        use HirScalarExpr::*;
4657        match self {
4658            Column(..)
4659            | Parameter(..)
4660            | Literal(..)
4661            | CallUnmaterializable(..)
4662            | CallUnary { .. }
4663            | CallBinary { .. }
4664            | CallVariadic { .. }
4665            | If { .. }
4666            | Windowing(..) => (),
4667            Exists(expr, _name) | Select(expr, _name) => f(expr)?,
4668        }
4669        Ok(())
4670    }
4671
4672    fn children<'a>(&'a self) -> impl DoubleEndedIterator<Item = &'a HirRelationExpr>
4673    where
4674        HirRelationExpr: 'a,
4675    {
4676        let mut child: Option<&HirRelationExpr> = None;
4677        use HirScalarExpr::*;
4678        match self {
4679            Column(..)
4680            | Parameter(..)
4681            | Literal(..)
4682            | CallUnmaterializable(..)
4683            | CallUnary { .. }
4684            | CallBinary { .. }
4685            | CallVariadic { .. }
4686            | If { .. }
4687            | Windowing(..) => (),
4688            Exists(expr, _name) | Select(expr, _name) => child = Some(&*expr),
4689        }
4690
4691        child.into_iter()
4692    }
4693
4694    fn children_mut<'a>(&'a mut self) -> impl DoubleEndedIterator<Item = &'a mut HirRelationExpr>
4695    where
4696        HirRelationExpr: 'a,
4697    {
4698        let mut child: Option<&mut HirRelationExpr> = None;
4699        use HirScalarExpr::*;
4700        match self {
4701            Column(..)
4702            | Parameter(..)
4703            | Literal(..)
4704            | CallUnmaterializable(..)
4705            | CallUnary { .. }
4706            | CallBinary { .. }
4707            | CallVariadic { .. }
4708            | If { .. }
4709            | Windowing(..) => (),
4710            Exists(expr, _name) | Select(expr, _name) => child = Some(&mut **expr),
4711        }
4712
4713        child.into_iter()
4714    }
4715}
4716
4717impl AbstractExpr for HirScalarExpr {
4718    type Type = SqlColumnType;
4719
4720    fn typ(
4721        &self,
4722        outers: &[SqlRelationType],
4723        inner: &SqlRelationType,
4724        params: &BTreeMap<usize, SqlScalarType>,
4725    ) -> Self::Type {
4726        stack::maybe_grow(|| match self {
4727            HirScalarExpr::Column(ColumnRef { level, column }, _name) => {
4728                if *level == 0 {
4729                    inner.column_types[*column].clone()
4730                } else {
4731                    outers[*level - 1].column_types[*column].clone()
4732                }
4733            }
4734            HirScalarExpr::Parameter(n, _name) => params[n].clone().nullable(true),
4735            HirScalarExpr::Literal(_, typ, _name) => typ.clone(),
4736            HirScalarExpr::CallUnmaterializable(func, _name) => func.output_sql_type(),
4737            HirScalarExpr::CallUnary {
4738                expr,
4739                func,
4740                name: _,
4741            } => func.output_sql_type(expr.typ(outers, inner, params)),
4742            HirScalarExpr::CallBinary {
4743                expr1,
4744                expr2,
4745                func,
4746                name: _,
4747            } => func.output_sql_type(&[
4748                expr1.typ(outers, inner, params),
4749                expr2.typ(outers, inner, params),
4750            ]),
4751            HirScalarExpr::CallVariadic {
4752                exprs,
4753                func,
4754                name: _,
4755            } => func.output_sql_type(exprs.iter().map(|e| e.typ(outers, inner, params)).collect()),
4756            HirScalarExpr::If {
4757                cond: _,
4758                then,
4759                els,
4760                name: _,
4761            } => {
4762                let then_type = then.typ(outers, inner, params);
4763                let else_type = els.typ(outers, inner, params);
4764                then_type.sql_union(&else_type).unwrap() // HIR deliberately not using `union`
4765            }
4766            HirScalarExpr::Exists(_, _name) => SqlScalarType::Bool.nullable(true),
4767            HirScalarExpr::Select(expr, _name) => {
4768                let mut outers = outers.to_vec();
4769                outers.insert(0, inner.clone());
4770                expr.typ(&outers, params)
4771                    .column_types
4772                    .into_element()
4773                    .nullable(true)
4774            }
4775            HirScalarExpr::Windowing(expr, _name) => expr.func.typ(outers, inner, params),
4776        })
4777    }
4778}
4779
4780impl AggregateExpr {
4781    pub fn typ(
4782        &self,
4783        outers: &[SqlRelationType],
4784        inner: &SqlRelationType,
4785        params: &BTreeMap<usize, SqlScalarType>,
4786    ) -> SqlColumnType {
4787        self.func
4788            .output_sql_type(self.expr.typ(outers, inner, params))
4789    }
4790
4791    /// Returns whether the expression is COUNT(*) or not.  Note that
4792    /// when we define the count builtin in sql::func, we convert
4793    /// COUNT(*) to COUNT(true), making it indistinguishable from
4794    /// literal COUNT(true), but we prefer to consider this as the
4795    /// former.
4796    ///
4797    /// (MIR has the same `is_count_asterisk`.)
4798    pub fn is_count_asterisk(&self) -> bool {
4799        self.func == AggregateFunc::Count && self.expr.is_literal_true() && !self.distinct
4800    }
4801}