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