Skip to main content

uqa_sql/ir/
traversal.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Complete scalar IR traversal.
8
9use super::{ScalarExpr, ScalarFrameBound};
10
11impl ScalarExpr {
12    /// Visit this expression and every nested scalar expression in pre-order.
13    pub fn visit(&self, visitor: &mut impl FnMut(&Self)) {
14        self.try_visit(&mut |expression| {
15            visitor(expression);
16            Ok::<_, std::convert::Infallible>(true)
17        })
18        .unwrap_or_else(|never| match never {});
19    }
20
21    /// Visit a window call's arguments, `FILTER` condition, partition and ordering keys, and frame offsets.
22    fn try_visit_window<E>(
23        args: &[Self],
24        filter: Option<&Self>,
25        spec: &super::ScalarWindowSpec,
26        visitor: &mut impl FnMut(&Self) -> Result<bool, E>,
27    ) -> Result<(), E> {
28        for expression in args
29            .iter()
30            .chain(filter)
31            .chain(&spec.partition_by)
32            .chain(spec.order_by.iter().map(|order| &order.expr))
33        {
34            expression.try_visit(visitor)?;
35        }
36        for bound in spec
37            .frame
38            .iter()
39            .flat_map(|frame| [&frame.start, &frame.end])
40        {
41            if let ScalarFrameBound::Preceding(expression)
42            | ScalarFrameBound::Following(expression) = bound
43            {
44                expression.try_visit(visitor)?;
45            }
46        }
47        Ok(())
48    }
49
50    /// Visit in pre-order, skipping a subtree when the visitor returns false and stopping immediately on its first error.
51    pub fn try_visit<E>(
52        &self,
53        visitor: &mut impl FnMut(&Self) -> Result<bool, E>,
54    ) -> Result<(), E> {
55        if !visitor(self)? {
56            return Ok(());
57        }
58        match self {
59            Self::And(parts)
60            | Self::Or(parts)
61            | Self::Array(parts)
62            | Self::Row(parts)
63            | Self::CompositeRow { items: parts, .. } => {
64                for part in parts {
65                    part.try_visit(visitor)?;
66                }
67            }
68            Self::Not(inner)
69            | Self::UnaryMinus(inner)
70            | Self::Cast { expr: inner, .. }
71            | Self::IsNull { expr: inner, .. }
72            | Self::InSubquery { expr: inner, .. } => inner.try_visit(visitor)?,
73            Self::Binary { lhs, rhs, .. } => {
74                lhs.try_visit(visitor)?;
75                rhs.try_visit(visitor)?;
76            }
77            Self::Between { expr, low, high } => {
78                expr.try_visit(visitor)?;
79                low.try_visit(visitor)?;
80                high.try_visit(visitor)?;
81            }
82            Self::InList { expr, list, .. } => {
83                expr.try_visit(visitor)?;
84                for part in list {
85                    part.try_visit(visitor)?;
86                }
87            }
88            Self::Func {
89                args,
90                order_by,
91                filter,
92                ..
93            } => {
94                for argument in args {
95                    argument.try_visit(visitor)?;
96                }
97                for order in order_by {
98                    order.expr.try_visit(visitor)?;
99                }
100                if let Some(filter) = filter {
101                    filter.try_visit(visitor)?;
102                }
103            }
104            Self::WindowCall {
105                args, spec, filter, ..
106            } => Self::try_visit_window(args, filter.as_deref(), spec, visitor)?,
107            Self::Case {
108                base,
109                when,
110                else_branch,
111            } => {
112                if let Some(base) = base {
113                    base.try_visit(visitor)?;
114                }
115                for (condition, result) in when {
116                    condition.try_visit(visitor)?;
117                    result.try_visit(visitor)?;
118                }
119                if let Some(else_branch) = else_branch {
120                    else_branch.try_visit(visitor)?;
121                }
122            }
123            Self::Default
124            | Self::Star
125            | Self::QualifiedStar(_)
126            | Self::Column(_)
127            | Self::Position(_)
128            | Self::InternalColumn(_)
129            | Self::QualifiedColumn { .. }
130            | Self::Literal(_)
131            | Self::TypedLiteral { .. }
132            | Self::Param(_)
133            | Self::ScalarSubquery(_)
134            | Self::Exists { .. } => {}
135        }
136        Ok(())
137    }
138
139    /// Visit a window call's arguments, `FILTER` condition, partition and ordering keys, and frame offsets mutably.
140    fn visit_window_mut(
141        args: &mut [Self],
142        filter: Option<&mut Self>,
143        spec: &mut super::ScalarWindowSpec,
144        visitor: &mut impl FnMut(&mut Self),
145    ) {
146        for expression in args
147            .iter_mut()
148            .chain(filter)
149            .chain(&mut spec.partition_by)
150            .chain(spec.order_by.iter_mut().map(|order| &mut order.expr))
151        {
152            expression.visit_mut(visitor);
153        }
154        for bound in spec
155            .frame
156            .iter_mut()
157            .flat_map(|frame| [&mut frame.start, &mut frame.end])
158        {
159            if let ScalarFrameBound::Preceding(expression)
160            | ScalarFrameBound::Following(expression) = bound
161            {
162                expression.visit_mut(visitor);
163            }
164        }
165    }
166
167    /// Visit this expression and every nested scalar expression mutably in pre-order: a node is visited before its children, so the children of a node the visitor replaces are those of the replacement.
168    pub fn visit_mut(&mut self, visitor: &mut impl FnMut(&mut Self)) {
169        visitor(self);
170        match self {
171            Self::And(parts)
172            | Self::Or(parts)
173            | Self::Array(parts)
174            | Self::Row(parts)
175            | Self::CompositeRow { items: parts, .. } => {
176                for part in parts {
177                    part.visit_mut(visitor);
178                }
179            }
180            Self::Not(inner)
181            | Self::UnaryMinus(inner)
182            | Self::Cast { expr: inner, .. }
183            | Self::IsNull { expr: inner, .. }
184            | Self::InSubquery { expr: inner, .. } => inner.visit_mut(visitor),
185            Self::Binary { lhs, rhs, .. } => {
186                lhs.visit_mut(visitor);
187                rhs.visit_mut(visitor);
188            }
189            Self::Between { expr, low, high } => {
190                expr.visit_mut(visitor);
191                low.visit_mut(visitor);
192                high.visit_mut(visitor);
193            }
194            Self::InList { expr, list, .. } => {
195                expr.visit_mut(visitor);
196                for part in list {
197                    part.visit_mut(visitor);
198                }
199            }
200            Self::Func {
201                args,
202                order_by,
203                filter,
204                ..
205            } => {
206                for argument in args {
207                    argument.visit_mut(visitor);
208                }
209                for order in order_by {
210                    order.expr.visit_mut(visitor);
211                }
212                if let Some(filter) = filter {
213                    filter.visit_mut(visitor);
214                }
215            }
216            Self::WindowCall {
217                args, spec, filter, ..
218            } => Self::visit_window_mut(args, filter.as_deref_mut(), spec, visitor),
219            Self::Case {
220                base,
221                when,
222                else_branch,
223            } => {
224                if let Some(base) = base {
225                    base.visit_mut(visitor);
226                }
227                for (condition, result) in when {
228                    condition.visit_mut(visitor);
229                    result.visit_mut(visitor);
230                }
231                if let Some(else_branch) = else_branch {
232                    else_branch.visit_mut(visitor);
233                }
234            }
235            Self::Default
236            | Self::Star
237            | Self::QualifiedStar(_)
238            | Self::Column(_)
239            | Self::Position(_)
240            | Self::InternalColumn(_)
241            | Self::QualifiedColumn { .. }
242            | Self::Literal(_)
243            | Self::TypedLiteral { .. }
244            | Self::Param(_)
245            | Self::ScalarSubquery(_)
246            | Self::Exists { .. } => {}
247        }
248    }
249
250    /// Collect every column needed to evaluate this expression. Returns `false` when evaluation needs row shape or a relational child that a projected field scan cannot provide.
251    pub fn collect_columns(&self, output: &mut std::collections::BTreeSet<String>) -> bool {
252        match self.try_visit_columns(&mut |name| {
253            output.insert(name.to_owned());
254            Ok::<_, std::convert::Infallible>(())
255        }) {
256            Ok(projectable) => projectable,
257            Err(never) => match never {},
258        }
259    }
260
261    /// Borrow referenced column names in evaluation-tree order, stopping at the first visitor failure or unprojectable expression. Qualified references yield their column component, matching `collect_columns`; repeated references remain visible to the visitor. No name or result container is allocated by this traversal.
262    pub fn try_visit_columns<'a, E>(
263        &'a self,
264        visitor: &mut impl FnMut(&'a str) -> Result<(), E>,
265    ) -> Result<bool, E> {
266        match self {
267            Self::Column(name) | Self::QualifiedColumn { column: name, .. } => {
268                visitor(name)?;
269                Ok(true)
270            }
271            Self::Literal(_)
272            | Self::TypedLiteral { .. }
273            | Self::Param(_)
274            | Self::InternalColumn(_) => Ok(true),
275            Self::Func {
276                args,
277                order_by,
278                filter,
279                ..
280            } => {
281                for expression in args
282                    .iter()
283                    .chain(order_by.iter().map(|order| &order.expr))
284                    .chain(filter.as_deref())
285                {
286                    if !expression.try_visit_columns(visitor)? {
287                        return Ok(false);
288                    }
289                }
290                Ok(true)
291            }
292            Self::Array(items)
293            | Self::Row(items)
294            | Self::CompositeRow { items, .. }
295            | Self::And(items)
296            | Self::Or(items) => {
297                for item in items {
298                    if !item.try_visit_columns(visitor)? {
299                        return Ok(false);
300                    }
301                }
302                Ok(true)
303            }
304            Self::Binary { lhs, rhs, .. } => {
305                Ok(lhs.try_visit_columns(visitor)? && rhs.try_visit_columns(visitor)?)
306            }
307            Self::UnaryMinus(expr)
308            | Self::Not(expr)
309            | Self::IsNull { expr, .. }
310            | Self::Cast { expr, .. } => expr.try_visit_columns(visitor),
311            Self::Between { expr, low, high } => Ok(expr.try_visit_columns(visitor)?
312                && low.try_visit_columns(visitor)?
313                && high.try_visit_columns(visitor)?),
314            Self::InList { expr, list, .. } => {
315                for item in std::iter::once(expr.as_ref()).chain(list) {
316                    if !item.try_visit_columns(visitor)? {
317                        return Ok(false);
318                    }
319                }
320                Ok(true)
321            }
322            Self::Case {
323                base,
324                when,
325                else_branch,
326            } => {
327                for expression in base
328                    .as_deref()
329                    .into_iter()
330                    .chain(
331                        when.iter()
332                            .flat_map(|(condition, result)| [condition, result]),
333                    )
334                    .chain(else_branch.as_deref())
335                {
336                    if !expression.try_visit_columns(visitor)? {
337                        return Ok(false);
338                    }
339                }
340                Ok(true)
341            }
342            Self::Default
343            | Self::Star
344            | Self::QualifiedStar(_)
345            | Self::Position(_)
346            | Self::WindowCall { .. }
347            | Self::ScalarSubquery(_)
348            | Self::Exists { .. }
349            | Self::InSubquery { .. } => Ok(false),
350        }
351    }
352
353    #[must_use]
354    pub fn contains_window(&self) -> bool {
355        match self {
356            Self::WindowCall { .. } => true,
357            Self::Func {
358                args,
359                order_by,
360                filter,
361                ..
362            } => {
363                args.iter().any(Self::contains_window)
364                    || order_by.iter().any(|order| order.expr.contains_window())
365                    || filter.as_deref().is_some_and(Self::contains_window)
366            }
367            Self::Array(items)
368            | Self::Row(items)
369            | Self::CompositeRow { items, .. }
370            | Self::And(items)
371            | Self::Or(items) => items.iter().any(Self::contains_window),
372            Self::Binary { lhs, rhs, .. } => lhs.contains_window() || rhs.contains_window(),
373            Self::UnaryMinus(expr)
374            | Self::Not(expr)
375            | Self::IsNull { expr, .. }
376            | Self::Cast { expr, .. }
377            | Self::InSubquery { expr, .. } => expr.contains_window(),
378            Self::Between { expr, low, high } => {
379                expr.contains_window() || low.contains_window() || high.contains_window()
380            }
381            Self::InList { expr, list, .. } => {
382                expr.contains_window() || list.iter().any(Self::contains_window)
383            }
384            Self::Case {
385                base,
386                when,
387                else_branch,
388            } => {
389                base.as_deref().is_some_and(Self::contains_window)
390                    || when.iter().any(|(condition, result)| {
391                        condition.contains_window() || result.contains_window()
392                    })
393                    || else_branch.as_deref().is_some_and(Self::contains_window)
394            }
395            Self::Default
396            | Self::Star
397            | Self::QualifiedStar(_)
398            | Self::Column(_)
399            | Self::QualifiedColumn { .. }
400            | Self::Position(_)
401            | Self::InternalColumn(_)
402            | Self::Literal(_)
403            | Self::TypedLiteral { .. }
404            | Self::Param(_)
405            | Self::ScalarSubquery(_)
406            | Self::Exists { .. } => false,
407        }
408    }
409
410    #[must_use]
411    pub fn contains_subquery(&self) -> bool {
412        match self {
413            Self::ScalarSubquery(_) | Self::Exists { .. } | Self::InSubquery { .. } => true,
414            Self::Func {
415                args,
416                order_by,
417                filter,
418                ..
419            } => {
420                args.iter().any(Self::contains_subquery)
421                    || order_by.iter().any(|order| order.expr.contains_subquery())
422                    || filter.as_deref().is_some_and(Self::contains_subquery)
423            }
424            Self::Array(items)
425            | Self::Row(items)
426            | Self::CompositeRow { items, .. }
427            | Self::And(items)
428            | Self::Or(items) => items.iter().any(Self::contains_subquery),
429            Self::Binary { lhs, rhs, .. } => lhs.contains_subquery() || rhs.contains_subquery(),
430            Self::UnaryMinus(expr)
431            | Self::Not(expr)
432            | Self::IsNull { expr, .. }
433            | Self::Cast { expr, .. } => expr.contains_subquery(),
434            Self::Between { expr, low, high } => {
435                expr.contains_subquery() || low.contains_subquery() || high.contains_subquery()
436            }
437            Self::InList { expr, list, .. } => {
438                expr.contains_subquery() || list.iter().any(Self::contains_subquery)
439            }
440            Self::WindowCall {
441                args, spec, filter, ..
442            } => {
443                args.iter().any(Self::contains_subquery)
444                    || filter.as_deref().is_some_and(Self::contains_subquery)
445                    || spec.partition_by.iter().any(Self::contains_subquery)
446                    || spec
447                        .order_by
448                        .iter()
449                        .any(|order| order.expr.contains_subquery())
450                    || spec.frame.as_ref().is_some_and(|frame| {
451                        frame_has(&frame.start, Self::contains_subquery)
452                            || frame_has(&frame.end, Self::contains_subquery)
453                    })
454            }
455            Self::Case {
456                base,
457                when,
458                else_branch,
459            } => {
460                base.as_deref().is_some_and(Self::contains_subquery)
461                    || when.iter().any(|(condition, result)| {
462                        condition.contains_subquery() || result.contains_subquery()
463                    })
464                    || else_branch.as_deref().is_some_and(Self::contains_subquery)
465            }
466            Self::Default
467            | Self::Star
468            | Self::QualifiedStar(_)
469            | Self::Column(_)
470            | Self::QualifiedColumn { .. }
471            | Self::Position(_)
472            | Self::InternalColumn(_)
473            | Self::Literal(_)
474            | Self::TypedLiteral { .. }
475            | Self::Param(_) => false,
476        }
477    }
478
479    #[must_use]
480    pub fn contains_parameter(&self) -> bool {
481        match self {
482            Self::Param(_) => true,
483            Self::Func {
484                args,
485                order_by,
486                filter,
487                ..
488            } => {
489                args.iter().any(Self::contains_parameter)
490                    || order_by.iter().any(|order| order.expr.contains_parameter())
491                    || filter.as_deref().is_some_and(Self::contains_parameter)
492            }
493            Self::Array(items)
494            | Self::Row(items)
495            | Self::CompositeRow { items, .. }
496            | Self::And(items)
497            | Self::Or(items) => items.iter().any(Self::contains_parameter),
498            Self::Binary { lhs, rhs, .. } => lhs.contains_parameter() || rhs.contains_parameter(),
499            Self::UnaryMinus(expr)
500            | Self::Not(expr)
501            | Self::IsNull { expr, .. }
502            | Self::Cast { expr, .. }
503            | Self::InSubquery { expr, .. } => expr.contains_parameter(),
504            Self::Between { expr, low, high } => {
505                expr.contains_parameter() || low.contains_parameter() || high.contains_parameter()
506            }
507            Self::InList { expr, list, .. } => {
508                expr.contains_parameter() || list.iter().any(Self::contains_parameter)
509            }
510            Self::WindowCall {
511                args, spec, filter, ..
512            } => {
513                args.iter().any(Self::contains_parameter)
514                    || filter.as_deref().is_some_and(Self::contains_parameter)
515                    || spec.partition_by.iter().any(Self::contains_parameter)
516                    || spec
517                        .order_by
518                        .iter()
519                        .any(|order| order.expr.contains_parameter())
520                    || spec.frame.as_ref().is_some_and(|frame| {
521                        frame_has(&frame.start, Self::contains_parameter)
522                            || frame_has(&frame.end, Self::contains_parameter)
523                    })
524            }
525            Self::Case {
526                base,
527                when,
528                else_branch,
529            } => {
530                base.as_deref().is_some_and(Self::contains_parameter)
531                    || when.iter().any(|(condition, result)| {
532                        condition.contains_parameter() || result.contains_parameter()
533                    })
534                    || else_branch.as_deref().is_some_and(Self::contains_parameter)
535            }
536            Self::Default
537            | Self::Star
538            | Self::QualifiedStar(_)
539            | Self::Column(_)
540            | Self::QualifiedColumn { .. }
541            | Self::Position(_)
542            | Self::InternalColumn(_)
543            | Self::Literal(_)
544            | Self::TypedLiteral { .. }
545            | Self::ScalarSubquery(_)
546            | Self::Exists { .. } => false,
547        }
548    }
549
550    #[must_use]
551    pub fn contains_aggregate(&self, is_aggregate: &dyn Fn(&str) -> bool) -> bool {
552        match self {
553            Self::Func {
554                name,
555                args,
556                order_by,
557                filter,
558                ..
559            } => {
560                is_aggregate(name)
561                    || args
562                        .iter()
563                        .any(|expression| expression.contains_aggregate(is_aggregate))
564                    || order_by
565                        .iter()
566                        .any(|order| order.expr.contains_aggregate(is_aggregate))
567                    || filter
568                        .as_deref()
569                        .is_some_and(|expression| expression.contains_aggregate(is_aggregate))
570            }
571            Self::Array(items)
572            | Self::Row(items)
573            | Self::CompositeRow { items, .. }
574            | Self::And(items)
575            | Self::Or(items) => items
576                .iter()
577                .any(|expression| expression.contains_aggregate(is_aggregate)),
578            Self::Binary { lhs, rhs, .. } => {
579                lhs.contains_aggregate(is_aggregate) || rhs.contains_aggregate(is_aggregate)
580            }
581            Self::UnaryMinus(expr)
582            | Self::Not(expr)
583            | Self::IsNull { expr, .. }
584            | Self::Cast { expr, .. }
585            | Self::InSubquery { expr, .. } => expr.contains_aggregate(is_aggregate),
586            Self::Between { expr, low, high } => {
587                expr.contains_aggregate(is_aggregate)
588                    || low.contains_aggregate(is_aggregate)
589                    || high.contains_aggregate(is_aggregate)
590            }
591            Self::InList { expr, list, .. } => {
592                expr.contains_aggregate(is_aggregate)
593                    || list
594                        .iter()
595                        .any(|item| item.contains_aggregate(is_aggregate))
596            }
597            Self::Case {
598                base,
599                when,
600                else_branch,
601            } => {
602                base.as_deref()
603                    .is_some_and(|expression| expression.contains_aggregate(is_aggregate))
604                    || when.iter().any(|(condition, result)| {
605                        condition.contains_aggregate(is_aggregate)
606                            || result.contains_aggregate(is_aggregate)
607                    })
608                    || else_branch
609                        .as_deref()
610                        .is_some_and(|expression| expression.contains_aggregate(is_aggregate))
611            }
612            Self::Default
613            | Self::Star
614            | Self::QualifiedStar(_)
615            | Self::Column(_)
616            | Self::QualifiedColumn { .. }
617            | Self::Position(_)
618            | Self::InternalColumn(_)
619            | Self::Literal(_)
620            | Self::TypedLiteral { .. }
621            | Self::Param(_)
622            | Self::ScalarSubquery(_)
623            | Self::Exists { .. }
624            | Self::WindowCall { .. } => false,
625        }
626    }
627}
628
629fn frame_has(bound: &ScalarFrameBound, predicate: fn(&ScalarExpr) -> bool) -> bool {
630    match bound {
631        ScalarFrameBound::Preceding(expression) | ScalarFrameBound::Following(expression) => {
632            predicate(expression)
633        }
634        ScalarFrameBound::UnboundedPreceding
635        | ScalarFrameBound::UnboundedFollowing
636        | ScalarFrameBound::CurrentRow => false,
637    }
638}
639
640#[cfg(test)]
641mod tests {
642    use super::{ScalarExpr, ScalarFrameBound};
643    use crate::ast::{FrameExclusion, FrameMode};
644    use uqa_core::Value;
645
646    #[test]
647    fn visit_includes_root_and_nested_expressions() {
648        let expression = ScalarExpr::Binary {
649            op: crate::ast::BinaryOp::Add,
650            lhs: Box::new(ScalarExpr::Column("amount".into())),
651            rhs: Box::new(ScalarExpr::Literal(Value::Int(1))),
652        };
653        let mut visited = Vec::new();
654        expression.visit(&mut |part| visited.push(part.clone()));
655        assert_eq!(visited.len(), 3);
656        assert_eq!(visited[0], expression);
657    }
658
659    #[test]
660    fn fallible_visits_skip_selected_subtrees_and_stop_before_later_siblings() {
661        let expression = ScalarExpr::Row(vec![
662            ScalarExpr::Array(vec![ScalarExpr::Column("hidden".into())]),
663            ScalarExpr::Column("reject".into()),
664            ScalarExpr::Column("unvisited".into()),
665        ]);
666        let mut visited = Vec::new();
667        let result = expression.try_visit(&mut |part| {
668            visited.push(part.clone());
669            match part {
670                ScalarExpr::Array(_) => Ok(false),
671                ScalarExpr::Column(name) if name == "reject" => Err("grouping"),
672                _ => Ok(true),
673            }
674        });
675        assert_eq!(result, Err("grouping"));
676        assert_eq!(visited.len(), 3);
677        assert!(matches!(&visited[2], ScalarExpr::Column(name) if name == "reject"));
678    }
679
680    #[test]
681    fn traversal_includes_window_frame_expressions() {
682        let expression = ScalarExpr::WindowCall {
683            name: "sum".into(),
684            args: vec![ScalarExpr::Column("amount".into())],
685            spec: super::super::ScalarWindowSpec {
686                definition: None,
687                partition_by: vec![ScalarExpr::QualifiedColumn {
688                    qualifier: "orders".into(),
689                    column: "account_id".into(),
690                }],
691                order_by: Vec::new(),
692                frame: Some(super::super::ScalarWindowFrame {
693                    mode: FrameMode::Rows,
694                    start: ScalarFrameBound::Preceding(Box::new(ScalarExpr::Param(0))),
695                    end: ScalarFrameBound::CurrentRow,
696                    between: true,
697                    exclusion: FrameExclusion::NoOthers,
698                }),
699            },
700            filter: None,
701            modifiers: crate::ast::WindowCallModifiers::default(),
702        };
703        let mut visited_parameter = false;
704        expression.visit(&mut |part| {
705            visited_parameter |= matches!(part, ScalarExpr::Param(0));
706        });
707        assert!(visited_parameter);
708        assert!(expression.contains_window());
709        assert!(expression.contains_parameter());
710    }
711
712    #[test]
713    fn mutable_visits_rewrite_nested_expressions() {
714        let literal = || ScalarExpr::Literal(Value::Str("t".into()));
715        let mut expression = ScalarExpr::Cast {
716            implicit: false,
717            expr: Box::new(ScalarExpr::Case {
718                base: None,
719                when: vec![(
720                    literal(),
721                    ScalarExpr::Cast {
722                        implicit: false,
723                        expr: Box::new(literal()),
724                        ty: "regclass".into(),
725                    },
726                )],
727                else_branch: Some(Box::new(ScalarExpr::WindowCall {
728                    name: "sum".into(),
729                    args: vec![literal()],
730                    spec: super::super::ScalarWindowSpec {
731                        definition: None,
732                        partition_by: Vec::new(),
733                        order_by: Vec::new(),
734                        frame: Some(super::super::ScalarWindowFrame {
735                            mode: FrameMode::Rows,
736                            start: ScalarFrameBound::Preceding(Box::new(literal())),
737                            end: ScalarFrameBound::CurrentRow,
738                            between: true,
739                            exclusion: FrameExclusion::NoOthers,
740                        }),
741                    },
742                    filter: Some(Box::new(literal())),
743                    modifiers: crate::ast::WindowCallModifiers::default(),
744                })),
745            }),
746            ty: "text".into(),
747        };
748        let mut rewritten = 0;
749        expression.visit_mut(&mut |part| {
750            if matches!(part, ScalarExpr::Literal(Value::Str(_))) {
751                *part = ScalarExpr::Literal(Value::Int(1));
752                rewritten += 1;
753            }
754        });
755        assert_eq!(rewritten, 5);
756        let mut remaining = 0;
757        expression.visit(&mut |part| {
758            remaining += usize::from(matches!(part, ScalarExpr::Literal(Value::Str(_))));
759        });
760        assert_eq!(remaining, 0);
761    }
762
763    #[test]
764    fn owned_walkers_preserve_column_and_aggregate_policy() {
765        let expression = ScalarExpr::Func {
766            order_syntax: crate::ast::FunctionOrderSyntax::Ordinary,
767            name: "sum".into(),
768            binding: None,
769            args: vec![ScalarExpr::QualifiedColumn {
770                qualifier: "orders".into(),
771                column: "amount".into(),
772            }],
773            distinct: false,
774            order_by: Vec::new(),
775            filter: None,
776        };
777        let mut columns = std::collections::BTreeSet::new();
778        assert!(expression.collect_columns(&mut columns));
779        assert_eq!(columns, std::collections::BTreeSet::from(["amount".into()]));
780        assert!(expression.contains_aggregate(&|name| name == "sum"));
781        assert!(!expression.contains_subquery());
782    }
783
784    #[test]
785    fn borrowed_column_visits_keep_names_and_stop_at_the_first_failure() {
786        let expression = ScalarExpr::Row(vec![
787            ScalarExpr::Column("first".into()),
788            ScalarExpr::QualifiedColumn {
789                qualifier: "table".into(),
790                column: "second".into(),
791            },
792            ScalarExpr::Column("first".into()),
793        ]);
794        let mut borrowed = Vec::new();
795        assert!(expression
796            .try_visit_columns(&mut |name| {
797                borrowed.push(name);
798                Ok::<_, &str>(())
799            })
800            .unwrap());
801        assert_eq!(borrowed, ["first", "second", "first"]);
802        let ScalarExpr::Row(items) = &expression else {
803            unreachable!()
804        };
805        let ScalarExpr::Column(first) = &items[0] else {
806            unreachable!()
807        };
808        assert_eq!(borrowed[0].as_ptr(), first.as_ptr());
809        let mut visits = 0;
810        let result = expression.try_visit_columns(&mut |_| {
811            visits += 1;
812            if visits == 2 {
813                Err("quota")
814            } else {
815                Ok(())
816            }
817        });
818        assert_eq!(result, Err("quota"));
819        assert_eq!(visits, 2);
820    }
821
822    #[test]
823    fn borrowed_column_visits_preserve_unprojectable_prefix_semantics() {
824        let expression = ScalarExpr::Array(vec![
825            ScalarExpr::Column("before".into()),
826            ScalarExpr::Position(0),
827            ScalarExpr::Column("after".into()),
828        ]);
829        let mut borrowed = Vec::new();
830        assert!(!expression
831            .try_visit_columns(&mut |name| {
832                borrowed.push(name);
833                Ok::<_, &str>(())
834            })
835            .unwrap());
836        let mut owned = std::collections::BTreeSet::new();
837        assert!(!expression.collect_columns(&mut owned));
838        assert_eq!(borrowed, ["before"]);
839        assert_eq!(owned, std::collections::BTreeSet::from(["before".into()]));
840    }
841}