1mod 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#[derive(Debug, Clone, PartialEq, Eq)]
33pub enum TraversalOrder {
34 Skip,
38 Stop,
40 Continue,
42}
43
44impl TraversalOrder {
45 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 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 pub value: T,
73 pub order: TraversalOrder,
75 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 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 fn apply_children<'a, F: FnMut(&'a Self) -> VortexResult<TraversalOrder>>(
141 &'a self,
142 f: F,
143 ) -> VortexResult<TraversalOrder>;
144
145 fn map_children<F: FnMut(Self) -> VortexResult<Transformed<Self>>>(
149 self,
150 f: F,
151 ) -> VortexResult<Transformed<Self>>;
152
153 fn iter_children<T>(&self, f: impl FnOnce(&mut dyn Iterator<Item = &Self>) -> T) -> T;
155
156 fn children_count(&self) -> usize;
158}
159
160pub trait NodeExt: Node {
161 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 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 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 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 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 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
327pub trait NodeContainer<'a, T: 'a>: Sized {
332 fn apply_elements<F: FnMut(&'a T) -> VortexResult<TraversalOrder>>(
334 &'a self,
335 f: F,
336 ) -> VortexResult<TraversalOrder>;
337
338 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}