Skip to main content

vortex_array/expr/traversal/
mod.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright the Vortex contributors
3
4//! Datafusion inspired tree traversal logic.
5//!
6//! Users should want to implement [`Node`] and potentially [`NodeContainer`].
7
8mod fold;
9mod references;
10mod visitor;
11
12use std::marker::PhantomData;
13use std::sync::Arc;
14
15pub use fold::FoldDown;
16pub use fold::FoldDownContext;
17pub use fold::FoldUp;
18pub use fold::NodeFolder;
19pub use fold::NodeFolderContext;
20use itertools::Itertools;
21pub use references::ReferenceCollector;
22pub use visitor::pre_order_visit_down;
23pub use visitor::pre_order_visit_up;
24use vortex_error::VortexResult;
25
26use crate::expr::BoundExpression;
27use crate::expr::Expression;
28use crate::expr::bound_expression::BoundKind;
29use crate::expr::traversal::fold::NodeFolderContextWrapper;
30
31/// Signal to control a traversal's flow
32#[derive(Debug, Clone, PartialEq, Eq)]
33pub enum TraversalOrder {
34    /// In a top-down traversal, skip visiting the children of the current node.
35    /// In the bottom-up phase of the traversal, skip the next step. Either skipping the children of the node,
36    /// moving to its next sibling, or skipping its parent once the children are traversed.
37    Skip,
38    /// Stop visiting any more nodes in the traversal.
39    Stop,
40    /// Continue with the traversal as expected.
41    Continue,
42}
43
44impl TraversalOrder {
45    /// If directed to, continue to visit nodes by running `f`, which should apply on the node's children.
46    pub fn visit_children<F: FnOnce() -> VortexResult<TraversalOrder>>(
47        self,
48        f: F,
49    ) -> VortexResult<TraversalOrder> {
50        match self {
51            Self::Skip => Ok(TraversalOrder::Continue),
52            Self::Stop => Ok(self),
53            Self::Continue => f(),
54        }
55    }
56
57    /// If directed to, continue to visit nodes by running `f`, which should apply on the node's parent.
58    pub fn visit_parent<F: FnOnce() -> VortexResult<TraversalOrder>>(
59        self,
60        f: F,
61    ) -> VortexResult<TraversalOrder> {
62        match self {
63            Self::Continue => f(),
64            Self::Skip | Self::Stop => Ok(self),
65        }
66    }
67}
68
69#[derive(Debug, Clone)]
70pub struct Transformed<T> {
71    /// Value that was being rewritten.
72    pub value: T,
73    /// Controls the flow of rewriting, see [`TraversalOrder`] for more details.
74    pub order: TraversalOrder,
75    /// Was the value changed during rewriting.
76    pub changed: bool,
77}
78
79impl<T> Transformed<T> {
80    pub fn yes(value: T) -> Self {
81        Self {
82            value,
83            order: TraversalOrder::Continue,
84            changed: true,
85        }
86    }
87
88    pub fn no(value: T) -> Self {
89        Self {
90            value,
91            order: TraversalOrder::Continue,
92            changed: false,
93        }
94    }
95
96    pub fn into_inner(self) -> T {
97        self.value
98    }
99
100    /// Apply a function to `value`, changing it without changing the `changed` field.
101    pub fn map<O, F: FnOnce(T) -> O>(self, f: F) -> Transformed<O> {
102        Transformed {
103            value: f(self.value),
104            order: self.order,
105            changed: self.changed,
106        }
107    }
108}
109
110pub trait NodeVisitor<'a> {
111    type NodeTy: Node;
112
113    fn visit_down(&mut self, node: &'a Self::NodeTy) -> VortexResult<TraversalOrder> {
114        _ = node;
115        Ok(TraversalOrder::Continue)
116    }
117
118    fn visit_up(&mut self, node: &'a Self::NodeTy) -> VortexResult<TraversalOrder> {
119        _ = node;
120        Ok(TraversalOrder::Continue)
121    }
122}
123
124pub trait NodeRewriter: Sized {
125    type NodeTy: Node;
126
127    fn visit_down(&mut self, node: Self::NodeTy) -> VortexResult<Transformed<Self::NodeTy>> {
128        Ok(Transformed::no(node))
129    }
130
131    fn visit_up(&mut self, node: Self::NodeTy) -> VortexResult<Transformed<Self::NodeTy>> {
132        Ok(Transformed::no(node))
133    }
134}
135
136pub trait Node: Sized + Clone {
137    /// Walk the node's children by applying `f` to them.
138    ///
139    /// This is a lower level API that other functions rely on for their implementation.
140    fn apply_children<'a, F: FnMut(&'a Self) -> VortexResult<TraversalOrder>>(
141        &'a self,
142        f: F,
143    ) -> VortexResult<TraversalOrder>;
144
145    /// Rewrite the node's children by applying `f` to them.
146    ///
147    /// This is a lower level API that other functions rely on for their implementation.
148    fn map_children<F: FnMut(Self) -> VortexResult<Transformed<Self>>>(
149        self,
150        f: F,
151    ) -> VortexResult<Transformed<Self>>;
152
153    /// This is a lower level API that other functions rely on for their implementation.
154    fn iter_children<T>(&self, f: impl FnOnce(&mut dyn Iterator<Item = &Self>) -> T) -> T;
155
156    /// This is a lower level API that other functions rely on for their implementation.
157    fn children_count(&self) -> usize;
158}
159
160pub trait NodeExt: Node {
161    /// Walk the tree in pre-order (top-down) way, rewriting it as it goes.
162    fn rewrite<R: NodeRewriter<NodeTy = Self>>(
163        self,
164        rewriter: &mut R,
165    ) -> VortexResult<Transformed<Self>> {
166        let mut transformed = rewriter.visit_down(self)?;
167
168        let transformed = match transformed.order {
169            TraversalOrder::Stop => Ok(transformed),
170            TraversalOrder::Skip => {
171                transformed.order = TraversalOrder::Continue;
172                Ok(transformed)
173            }
174            TraversalOrder::Continue => transformed
175                .value
176                .map_children(|c| c.rewrite(rewriter))
177                .map(|mut t| {
178                    t.changed |= transformed.changed;
179                    t
180                }),
181        }?;
182
183        match transformed.order {
184            TraversalOrder::Stop | TraversalOrder::Skip => Ok(transformed),
185            TraversalOrder::Continue => {
186                let mut up_rewrite = rewriter.visit_up(transformed.value)?;
187                up_rewrite.changed |= transformed.changed;
188                Ok(up_rewrite)
189            }
190        }
191    }
192
193    /// A pre-order (top-down) traversal.
194    fn accept<'a, V: NodeVisitor<'a, NodeTy = Self>>(
195        &'a self,
196        visitor: &mut V,
197    ) -> VortexResult<TraversalOrder> {
198        visitor
199            .visit_down(self)?
200            .visit_children(|| self.apply_children(|c| c.accept(visitor)))?
201            .visit_parent(|| visitor.visit_up(self))
202    }
203
204    /// A pre-order transformation
205    fn transform_down<F: FnMut(Self) -> VortexResult<Transformed<Self>>>(
206        self,
207        f: F,
208    ) -> VortexResult<Transformed<Self>> {
209        let mut rewriter = FnRewriter::<F, F, _> {
210            f_down: Some(f),
211            f_up: None,
212            _data: PhantomData,
213        };
214
215        self.rewrite(&mut rewriter)
216    }
217
218    fn transform<F, G>(self, down: F, up: G) -> VortexResult<Transformed<Self>>
219    where
220        F: FnMut(Self) -> VortexResult<Transformed<Self>>,
221        G: FnMut(Self) -> VortexResult<Transformed<Self>>,
222    {
223        let mut rewriter = FnRewriter {
224            f_down: Some(down),
225            f_up: Some(up),
226            _data: PhantomData,
227        };
228
229        self.rewrite(&mut rewriter)
230    }
231
232    /// A post-order transform
233    fn transform_up<F: FnMut(Self) -> VortexResult<Transformed<Self>>>(
234        self,
235        f: F,
236    ) -> VortexResult<Transformed<Self>> {
237        let mut rewriter = FnRewriter::<F, F, _> {
238            f_down: None,
239            f_up: Some(f),
240            _data: PhantomData,
241        };
242
243        self.rewrite(&mut rewriter)
244    }
245
246    /// applies the `NodeFolderContext` to the Node tree, with an initial `Context`.
247    fn fold_context<R, F: NodeFolderContext<NodeTy = Self, Result = R>>(
248        self,
249        ctx: &F::Context,
250        folder: &mut F,
251    ) -> VortexResult<FoldUp<R>> {
252        let transformed = folder.visit_down(ctx, &self)?;
253        let ctx = match transformed {
254            FoldDownContext::Continue(ctx) => ctx,
255            FoldDownContext::Skip(r) => return Ok(FoldUp::Continue(r)),
256            FoldDownContext::Stop(r) => return Ok(FoldUp::Stop(r)),
257        };
258
259        let mut children = Vec::with_capacity(self.children_count());
260        let mut stop_result = None;
261        self.iter_children(|children_iter| -> VortexResult<()> {
262            for c in children_iter {
263                let t = c.clone().fold_context(&ctx, folder)?;
264                match t {
265                    FoldUp::Stop(r) => {
266                        stop_result = Some(r);
267                        return Ok(());
268                    }
269                    FoldUp::Continue(r) => {
270                        children.push(r);
271                    }
272                }
273            }
274            Ok(())
275        })?;
276
277        if let Some(result) = stop_result {
278            return Ok(FoldUp::Stop(result));
279        }
280
281        folder.visit_up(self, &ctx, children)
282    }
283
284    /// applies the `NodeFolder` to the Node tree
285    fn fold<R, F: NodeFolder<NodeTy = Self, Result = R>>(
286        self,
287        folder: &mut F,
288    ) -> VortexResult<FoldUp<R>> {
289        let mut folder = NodeFolderContextWrapper { inner: folder };
290        self.fold_context(&(), &mut folder)
291    }
292}
293
294impl<T: Node> NodeExt for T {}
295
296struct FnRewriter<F, G, T> {
297    f_down: Option<F>,
298    f_up: Option<G>,
299    _data: PhantomData<T>,
300}
301
302impl<F, G, T> NodeRewriter for FnRewriter<F, G, T>
303where
304    T: Node,
305    F: FnMut(T) -> VortexResult<Transformed<T>>,
306    G: FnMut(T) -> VortexResult<Transformed<T>>,
307{
308    type NodeTy = T;
309
310    fn visit_down(&mut self, node: Self::NodeTy) -> VortexResult<Transformed<Self::NodeTy>> {
311        if let Some(f) = self.f_down.as_mut() {
312            f(node)
313        } else {
314            Ok(Transformed::no(node))
315        }
316    }
317
318    fn visit_up(&mut self, node: Self::NodeTy) -> VortexResult<Transformed<Self::NodeTy>> {
319        if let Some(f) = self.f_up.as_mut() {
320            f(node)
321        } else {
322            Ok(Transformed::no(node))
323        }
324    }
325}
326
327/// A container holding a [`Node`]'s children, which a function can be applied (or mapped) to.
328///
329/// The trait is also implemented to container types in order to make implementing [`Node::map_children`]
330/// and [`Node::apply_children`] easier.
331pub trait NodeContainer<'a, T: 'a>: Sized {
332    /// Applies `f` to all elements of the container, accepting them by reference
333    fn apply_elements<F: FnMut(&'a T) -> VortexResult<TraversalOrder>>(
334        &'a self,
335        f: F,
336    ) -> VortexResult<TraversalOrder>;
337
338    /// Consumes all the children of the node, replacing them with the result of `f`.
339    fn map_elements<F: FnMut(T) -> VortexResult<Transformed<T>>>(
340        self,
341        f: F,
342    ) -> VortexResult<Transformed<Self>>;
343}
344
345pub trait NodeRefContainer<'a, T: 'a>: Sized {
346    fn apply_ref_elements<F: FnMut(&'a T) -> VortexResult<TraversalOrder>>(
347        &self,
348        f: F,
349    ) -> VortexResult<TraversalOrder>;
350}
351
352impl<'a, T: 'a, C: NodeContainer<'a, T>> NodeRefContainer<'a, T> for &'a [C] {
353    fn apply_ref_elements<F: FnMut(&'a T) -> VortexResult<TraversalOrder>>(
354        &self,
355        mut f: F,
356    ) -> VortexResult<TraversalOrder> {
357        let mut order = TraversalOrder::Continue;
358
359        for c in *self {
360            order = c.apply_elements(&mut f)?;
361            match order {
362                TraversalOrder::Continue | TraversalOrder::Skip => {}
363                TraversalOrder::Stop => return Ok(TraversalOrder::Stop),
364            }
365        }
366
367        Ok(order)
368    }
369}
370
371impl<'a, T: 'a, C: NodeContainer<'a, T>> NodeContainer<'a, T> for Box<C> {
372    fn apply_elements<F: FnMut(&'a T) -> VortexResult<TraversalOrder>>(
373        &'a self,
374        f: F,
375    ) -> VortexResult<TraversalOrder> {
376        self.as_ref().apply_elements(f)
377    }
378
379    fn map_elements<F: FnMut(T) -> VortexResult<Transformed<T>>>(
380        self,
381        f: F,
382    ) -> VortexResult<Transformed<Box<C>>> {
383        Ok((*self).map_elements(f)?.map(Box::new))
384    }
385}
386
387impl<'a, T, C> NodeContainer<'a, T> for Arc<C>
388where
389    T: 'a,
390    C: NodeContainer<'a, T> + Clone,
391{
392    fn apply_elements<F: FnMut(&'a T) -> VortexResult<TraversalOrder>>(
393        &'a self,
394        f: F,
395    ) -> VortexResult<TraversalOrder> {
396        self.as_ref().apply_elements(f)
397    }
398
399    fn map_elements<F: FnMut(T) -> VortexResult<Transformed<T>>>(
400        self,
401        f: F,
402    ) -> VortexResult<Transformed<Arc<C>>> {
403        Ok(Arc::unwrap_or_clone(self).map_elements(f)?.map(Arc::new))
404    }
405}
406
407impl<'a, T: 'a, C: NodeContainer<'a, T>> NodeContainer<'a, T> for [C; 2] {
408    fn apply_elements<F: FnMut(&'a T) -> VortexResult<TraversalOrder>>(
409        &'a self,
410        mut f: F,
411    ) -> VortexResult<TraversalOrder> {
412        let [lhs, rhs] = self;
413        match lhs.apply_elements(&mut f)? {
414            TraversalOrder::Skip | TraversalOrder::Continue => rhs.apply_elements(&mut f),
415            TraversalOrder::Stop => Ok(TraversalOrder::Stop),
416        }
417    }
418
419    fn map_elements<F: FnMut(T) -> VortexResult<Transformed<T>>>(
420        self,
421        mut f: F,
422    ) -> VortexResult<Transformed<[C; 2]>> {
423        let [lhs, rhs] = self;
424        let transformed = lhs.map_elements(&mut f)?;
425        match transformed.order {
426            TraversalOrder::Skip | TraversalOrder::Continue => {
427                let mut t = rhs.map_elements(&mut f)?;
428                t.changed |= transformed.changed;
429                Ok(t.map(|new_lhs| [new_lhs, transformed.value]))
430            }
431            TraversalOrder::Stop => Ok(transformed.map(|new_lhs| [new_lhs, rhs])),
432        }
433    }
434}
435
436impl<'a, T: 'a, C: NodeContainer<'a, T>> NodeContainer<'a, T> for Vec<C> {
437    fn apply_elements<F: FnMut(&'a T) -> VortexResult<TraversalOrder>>(
438        &'a self,
439        mut f: F,
440    ) -> VortexResult<TraversalOrder> {
441        let mut order = TraversalOrder::Continue;
442
443        for c in self {
444            order = c.apply_elements(&mut f)?;
445            match order {
446                TraversalOrder::Continue | TraversalOrder::Skip => {}
447                TraversalOrder::Stop => return Ok(TraversalOrder::Stop),
448            }
449        }
450
451        Ok(order)
452    }
453
454    fn map_elements<F: FnMut(T) -> VortexResult<Transformed<T>>>(
455        self,
456        mut f: F,
457    ) -> VortexResult<Transformed<Self>> {
458        let mut order = TraversalOrder::Continue;
459        let mut changed = false;
460
461        let value = self
462            .into_iter()
463            .map(|c| match order {
464                TraversalOrder::Continue | TraversalOrder::Skip => {
465                    c.map_elements(&mut f).map(|result| {
466                        order = result.order;
467                        changed |= result.changed;
468                        result.value
469                    })
470                }
471                TraversalOrder::Stop => Ok(c),
472            })
473            .collect::<VortexResult<Vec<_>>>()?;
474
475        Ok(Transformed {
476            value,
477            order,
478            changed,
479        })
480    }
481}
482
483impl<'a> NodeContainer<'a, Self> for Expression {
484    fn apply_elements<F: FnMut(&'a Self) -> VortexResult<TraversalOrder>>(
485        &'a self,
486        mut f: F,
487    ) -> VortexResult<TraversalOrder> {
488        f(self)
489    }
490
491    fn map_elements<F: FnMut(Self) -> VortexResult<Transformed<Self>>>(
492        self,
493        mut f: F,
494    ) -> VortexResult<Transformed<Self>> {
495        f(self)
496    }
497}
498
499impl Node for Expression {
500    fn apply_children<'a, F: FnMut(&'a Self) -> VortexResult<TraversalOrder>>(
501        &'a self,
502        mut f: F,
503    ) -> VortexResult<TraversalOrder> {
504        self.children().as_ref().apply_elements(&mut f)
505    }
506
507    fn map_children<F: FnMut(Self) -> VortexResult<Transformed<Self>>>(
508        self,
509        f: F,
510    ) -> VortexResult<Transformed<Self>> {
511        let transformed = self
512            .children()
513            .iter()
514            .cloned()
515            .collect_vec()
516            .map_elements(f)?;
517
518        if transformed.changed {
519            Ok(Transformed {
520                value: self.with_children(transformed.value)?,
521                order: transformed.order,
522                changed: true,
523            })
524        } else {
525            Ok(Transformed::no(self))
526        }
527    }
528
529    fn iter_children<T>(&self, f: impl FnOnce(&mut dyn Iterator<Item = &Self>) -> T) -> T {
530        f(&mut self.children().iter())
531    }
532
533    fn children_count(&self) -> usize {
534        self.children().len()
535    }
536}
537
538impl Node for BoundExpression {
539    fn apply_children<'a, F: FnMut(&'a Self) -> VortexResult<TraversalOrder>>(
540        &'a self,
541        mut f: F,
542    ) -> VortexResult<TraversalOrder> {
543        let BoundKind::Scalar { children, .. } = self.kind() else {
544            return Ok(TraversalOrder::Continue);
545        };
546
547        for child in children.iter() {
548            match f(child)? {
549                TraversalOrder::Continue | TraversalOrder::Skip => {}
550                TraversalOrder::Stop => return Ok(TraversalOrder::Stop),
551            }
552        }
553
554        Ok(TraversalOrder::Continue)
555    }
556
557    fn map_children<F: FnMut(Self) -> VortexResult<Transformed<Self>>>(
558        self,
559        mut f: F,
560    ) -> VortexResult<Transformed<Self>> {
561        let BoundKind::Scalar { children, .. } = self.kind() else {
562            return Ok(Transformed::no(self));
563        };
564
565        let mut order = TraversalOrder::Continue;
566        let mut changed = false;
567        let children = children
568            .iter()
569            .cloned()
570            .map(|child| match order {
571                TraversalOrder::Continue | TraversalOrder::Skip => f(child).map(|result| {
572                    order = result.order;
573                    changed |= result.changed;
574                    result.value
575                }),
576                TraversalOrder::Stop => Ok(child),
577            })
578            .collect::<VortexResult<Vec<_>>>()?;
579
580        if changed {
581            Ok(Transformed {
582                value: self.with_children(children)?,
583                order,
584                changed: true,
585            })
586        } else {
587            Ok(Transformed::no(self))
588        }
589    }
590
591    fn iter_children<T>(&self, f: impl FnOnce(&mut dyn Iterator<Item = &Self>) -> T) -> T {
592        match self.kind() {
593            BoundKind::Scalar { children, .. } => f(&mut children.iter()),
594            BoundKind::Root => f(&mut std::iter::empty()),
595        }
596    }
597
598    fn children_count(&self) -> usize {
599        match self.kind() {
600            BoundKind::Scalar { children, .. } => children.len(),
601            BoundKind::Root => 0,
602        }
603    }
604}
605
606#[cfg(test)]
607mod tests {
608    use vortex_error::VortexResult;
609    use vortex_utils::aliases::hash_set::HashSet;
610
611    use super::NodeExt;
612    use super::NodeRewriter;
613    use super::NodeVisitor;
614    use super::Transformed;
615    use super::TraversalOrder;
616    use super::visitor::pre_order_visit_down;
617    use crate::expr::Expression;
618    use crate::expr::and;
619    use crate::expr::col;
620    use crate::expr::eq;
621    use crate::expr::is_root;
622    use crate::expr::lit;
623    use crate::expr::not_eq;
624    use crate::expr::root;
625    use crate::scalar_fn::fns::binary::Binary;
626    use crate::scalar_fn::fns::get_item::GetItem;
627    use crate::scalar_fn::fns::literal::Literal;
628    use crate::scalar_fn::fns::operators::Operator;
629
630    #[derive(Default)]
631    pub struct ExprLitCollector<'a>(pub Vec<&'a Expression>);
632
633    impl<'a> NodeVisitor<'a> for ExprLitCollector<'a> {
634        type NodeTy = Expression;
635
636        fn visit_down(&mut self, node: &'a Expression) -> VortexResult<TraversalOrder> {
637            if node.is::<Literal>() {
638                self.0.push(node)
639            }
640            Ok(TraversalOrder::Continue)
641        }
642
643        fn visit_up(&mut self, _node: &'a Expression) -> VortexResult<TraversalOrder> {
644            Ok(TraversalOrder::Continue)
645        }
646    }
647
648    fn expr_col_to_lit_transform(
649        node: Expression,
650        idx: &mut i32,
651    ) -> VortexResult<Transformed<Expression>> {
652        if node.is::<GetItem>() {
653            let lit_id = *idx;
654            *idx += 1;
655            Ok(Transformed::yes(lit(lit_id)))
656        } else {
657            Ok(Transformed::no(node))
658        }
659    }
660
661    #[derive(Default)]
662    pub struct SkipDownRewriter;
663
664    impl NodeRewriter for SkipDownRewriter {
665        type NodeTy = Expression;
666
667        fn visit_down(&mut self, node: Self::NodeTy) -> VortexResult<Transformed<Self::NodeTy>> {
668            Ok(Transformed {
669                value: node,
670                order: TraversalOrder::Skip,
671                changed: false,
672            })
673        }
674
675        fn visit_up(&mut self, _node: Self::NodeTy) -> VortexResult<Transformed<Self::NodeTy>> {
676            Ok(Transformed::yes(root()))
677        }
678    }
679
680    #[test]
681    fn expr_deep_visitor_test() {
682        let col1: Expression = col("col1");
683        let lit1 = lit(1);
684        let expr = eq(col1, lit1);
685        let lit2 = lit(2);
686        let expr = and(expr, lit2);
687        let mut printer = ExprLitCollector::default();
688        expr.accept(&mut printer).unwrap();
689        assert_eq!(printer.0.len(), 2);
690    }
691
692    #[test]
693    fn expr_deep_mut_visitor_test() {
694        let col1: Expression = col("col1");
695        let col2: Expression = col("col2");
696        let expr = eq(col1, col2);
697        let lit2 = lit(2);
698        let expr = and(expr, lit2);
699
700        let mut idx = 0_i32;
701        let new = expr
702            .transform_up(|node| expr_col_to_lit_transform(node, &mut idx))
703            .unwrap();
704        assert!(new.changed);
705
706        let expr = new.value;
707
708        let mut printer = ExprLitCollector::default();
709        expr.accept(&mut printer).unwrap();
710        assert_eq!(printer.0.len(), 3);
711    }
712
713    #[test]
714    fn expr_skip_test() {
715        let col1: Expression = col("col1");
716        let col2: Expression = col("col2");
717        let expr1 = eq(col1, col2);
718        let col3: Expression = col("col3");
719        let col4: Expression = col("col4");
720        let expr2 = not_eq(col3, col4);
721        let expr = and(expr1, expr2);
722
723        let mut nodes = Vec::new();
724        pre_order_visit_down(&expr, |node: &Expression| {
725            if node.is::<GetItem>() {
726                nodes.push(node)
727            }
728            if let Some(operator) = node.as_opt::<Binary>()
729                && *operator == Operator::Eq
730            {
731                return Ok(TraversalOrder::Skip);
732            }
733            Ok(TraversalOrder::Continue)
734        })
735        .unwrap();
736
737        let nodes: HashSet<Expression> = HashSet::from_iter(nodes.into_iter().cloned());
738        assert_eq!(nodes, HashSet::from_iter([col("col3"), col("col4")]));
739    }
740
741    #[test]
742    fn expr_stop_test() {
743        let col1: Expression = col("col1");
744        let col2: Expression = col("col2");
745        let expr1 = eq(col1, col2);
746        let col3: Expression = col("col3");
747        let col4: Expression = col("col4");
748        let expr2 = not_eq(col3, col4);
749        let expr = and(expr1, expr2);
750
751        let mut nodes = Vec::new();
752        pre_order_visit_down(&expr, |node: &Expression| {
753            if node.is::<GetItem>() {
754                nodes.push(node)
755            }
756            if let Some(operator) = node.as_opt::<Binary>()
757                && *operator == Operator::Eq
758            {
759                return Ok(TraversalOrder::Stop);
760            }
761            Ok(TraversalOrder::Continue)
762        })
763        .unwrap();
764
765        assert!(nodes.is_empty());
766    }
767
768    #[test]
769    fn expr_skip_down_visit_up() {
770        let col = col("col");
771
772        let mut visitor = SkipDownRewriter;
773        let result = col.rewrite(&mut visitor).unwrap();
774
775        assert!(result.changed);
776        assert!(is_root(&result.value));
777    }
778}