1use super::{ScalarExpr, ScalarFrameBound};
10
11impl ScalarExpr {
12 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 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 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 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 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 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 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}