Skip to main content

polydat_core/iteration/comprehension/optimize/
mod.rs

1// Copyright 2024-2026 Jonathan Shook
2// SPDX-License-Identifier: Apache-2.0
3
4//! Post-parse optimizer — comprehension_forms.md §10.
5//!
6//! Required pass upstream of compilation. Takes an AST and
7//! produces a canonical, push-down form with these properties
8//! (§10.6):
9//!
10//! 1. **Semantic-preserving.** Output produces the same
11//!    dispense sequence (per §9.2).
12//! 2. **Idempotent.** `optimize(optimize(C)) == optimize(C)`.
13//! 3. **Decidable termination.** Each rewrite strictly
14//!    decreases a metadata-derived measure or leaves the AST
15//!    unchanged.
16//! 4. **Bounds-improving.** Peak memory never grows.
17//! 5. **No rejections.** Validity is decided on the tree as written,
18//!    before any rewrite, and every rewrite keeps a valid tree valid.
19//!
20//! ## R-rule catalog
21//!
22//! Priority order: R0a → R0b → R1 → R2 → R3 → R4 → R5 → R6 →
23//! R7 (§10.10.5).
24//!
25//! - **R0a — identity elimination** (I1–I5): singleton
26//!   combinators, trivially-true filter, `order(Lex, None)`.
27//! - **R0b — associativity flattening** (A1, A2): nested
28//!   union / cartesian collapse to n-ary form.
29//! - **R1 — `order(Lex)` → `ORDER_STREAMING`**: the IR
30//!   compiler emits `ORDER_STREAMING` for `order(Lex, _)` and
31//!   `ORDER_MATERIALIZE` otherwise, carrying the input's
32//!   `metadata.index_addressable`. Not an AST rewrite; recorded
33//!   in the reducibility catalog as an IR-compilation
34//!   eligibility.
35//! - **R2 — `order(c, strategy, Some(n))` → `indexed_order`**:
36//!   metadata-driven. The working set is the selection, sized by
37//!   `strategy_working_set` in `metadata.rs`'s propagation rule,
38//!   and `ORDER_MATERIALIZE` selects positions from its input's
39//!   shape when it evaluates.
40//! - **R3 — `order(filter, Lex, None)` → `filter(order, Lex, None)`**:
41//!   AST rewrite. Commute when un-truncated.
42//! - **R4 — `filter(union(...), p)` → `union(filter(...))`**:
43//!   AST rewrite. Distribute filter into each union child, except
44//!   under a non-`Lex` order that ranks the filter's survivors.
45//! - **R5 — per-axis filter pushdown**: AST rewrite. Consults
46//!   the predicate analyzer (§10.9) for factorization; when
47//!   `factorization = PerAxis` and every per-axis sub-predicate is
48//!   total over its axis, splits the filter into per-axis filters
49//!   wrapping each cartesian child, except under a non-`Lex` order
50//!   that ranks the filter's survivors.
51//! - **R6 — chained filter folding** (F1): AST rewrite.
52//!   `filter(filter(c, p), q)` → `filter(c, p && q)`.
53//! - **R7 — order chain folding** (O1): AST rewrite.
54//!   `order(order(c, s1, None), s2, t)` → `order(c, s2, t)` when `s2`
55//!   selects from its input's shape.
56//!
57//! ## Module layout
58//!
59//! - [`finding`] — `ReducibilityFinding`, `Reduction`,
60//!   `ComplexityDelta`.
61//! - [`r0a_identity`] — I1–I5 elimination.
62//! - [`r0b_flatten`] — A1, A2 flattening.
63//! - [`r3_commute`] — Lex/filter commute.
64//! - [`r4_distribute`] — filter over union.
65//! - [`r5_factorize`] — per-axis filter pushdown.
66//! - [`r6_filter_fold`] — chained filter folding.
67//! - [`r7_order_fold`] — order chain folding.
68
69use super::ast::Comprehension;
70use super::predicate::CoordSet;
71use super::strategies::for_name;
72use super::strategy::StrategyName;
73use crate::iteration::comprehension::metadata::Metadata;
74
75pub mod finding;
76pub mod r0a_identity;
77pub mod r0b_flatten;
78pub mod r3_commute;
79pub mod r4_distribute;
80pub mod r5_factorize;
81pub mod r6_filter_fold;
82pub mod r7_order_fold;
83
84pub use finding::{
85    ComplexityDelta, Ordering as ComplexityOrdering, ReducibilityFinding, Reduction, RuleId,
86};
87
88/// Top-level optimizer entry. Applies the R-rule catalog to a
89/// fixed point and returns the optimized AST.
90///
91/// Per comprehension_forms.md §10.6 the function is total — it
92/// never rejects. Validation (V1–V9) runs on the tree as written,
93/// before this; the optimizer assumes its input is well-formed and
94/// keeps it so.
95///
96/// The optimizer is a thin loop over the reducibility analyzer
97/// (§10.10): ask `analyze_reducibility` for a finding; apply
98/// its witness if non-empty; repeat. The empty finding ends
99/// the loop.
100pub fn optimize(ast: Comprehension) -> Comprehension {
101    let mut current = ast;
102    let mut steps_remaining = max_steps(&current);
103    while steps_remaining > 0 {
104        match analyze_reducibility(&current) {
105            ReducibilityFinding {
106                reduction: Some(Reduction::Rewrite { witness, .. }),
107                ..
108            } => {
109                current = witness;
110            }
111            ReducibilityFinding {
112                reduction: Some(Reduction::Replace { with }),
113                ..
114            } => {
115                current = with;
116            }
117            _ => break,
118        }
119        steps_remaining -= 1;
120    }
121    current
122}
123
124/// Reducibility analyzer entry — comprehension_forms.md §10.10.
125///
126/// Walks the AST bottom-up trying each R-rule in priority
127/// order. Returns the first non-empty finding; returns
128/// the empty finding when no rule fires.
129pub fn analyze_reducibility(ast: &Comprehension) -> ReducibilityFinding {
130    analyze_at(ast, Place::Free)
131}
132
133/// Where a node sits relative to the orders above it.
134#[derive(Debug, Clone, Copy, PartialEq, Eq)]
135enum Place {
136    /// No non-`Lex` order ranks this node's tuples by position.
137    Free,
138    /// A non-`Lex` order ranks this node's tuples by their positions in
139    /// its input (comprehension_forms.md §5 V5): a filter here keeps
140    /// its input's shape, so R4 and R5, which reshape that input, do
141    /// not fire. `through_orders` holds when the ranking order selects
142    /// from the shape and so reads through an untruncated order here
143    /// (§7.4 O1).
144    Ranked { through_orders: bool },
145}
146
147impl Place {
148    /// The place of a child of `node`, which sits at `self`.
149    fn of_child(self, node: &Comprehension) -> Place {
150        // The input of a ranked filter keeps its shape too: an order
151        // that selects from the shape ranks the survivors by their
152        // positions beneath the untruncated orders there.
153        if let (Comprehension::Filter { .. }, Place::Ranked { through_orders }) = (node, self) {
154            return Place::Ranked { through_orders };
155        }
156        let Comprehension::Order {
157            strategy,
158            truncation,
159            ..
160        } = node
161        else {
162            return Place::Free;
163        };
164        let inherited = truncation.is_none()
165            && self
166                == Place::Ranked {
167                    through_orders: true,
168                };
169        if *strategy != StrategyName::Lex {
170            Place::Ranked {
171                through_orders: inherited || for_name(*strategy).selects_from_shape(),
172            }
173        } else if inherited {
174            self
175        } else {
176            Place::Free
177        }
178    }
179}
180
181fn analyze_at(ast: &Comprehension, place: Place) -> ReducibilityFinding {
182    // Bottom-up: try to rewrite each child first.
183    // Rewriting a child returns a new parent that wraps the
184    // rewritten child; subsequent rule attempts then see the
185    // updated subtree on the next outer-loop iteration.
186    if let Some(finding) = try_rewrite_child_first(ast, place) {
187        return finding;
188    }
189    // No rewrite in a child — try rules at this node.
190    try_rules_at_node(ast, place)
191}
192
193/// Attempt to rewrite a child; return a finding that wraps
194/// the rewritten subtree in this node's variant.
195fn try_rewrite_child_first(ast: &Comprehension, place: Place) -> Option<ReducibilityFinding> {
196    let children: Vec<Comprehension> = ast.children().cloned().collect();
197    for (i, child) in children.iter().enumerate() {
198        let child_finding = analyze_at(child, place.of_child(ast));
199        let rewritten = match child_finding.reduction {
200            Some(Reduction::Rewrite { witness, .. }) => witness,
201            Some(Reduction::Replace { with }) => with,
202            None => continue,
203        };
204        // Re-build this node with the rewritten child at position i.
205        let new_ast = replace_child_at(ast, i, rewritten);
206        return Some(ReducibilityFinding {
207            reduction: Some(Reduction::Rewrite {
208                rule: child_finding.rule.unwrap_or(RuleId::R0a),
209                witness: new_ast,
210            }),
211            rule: child_finding.rule,
212            improvement: child_finding.improvement,
213        });
214    }
215    None
216}
217
218/// Try every R-rule at this node in priority order.
219/// First fire wins.
220fn try_rules_at_node(ast: &Comprehension, place: Place) -> ReducibilityFinding {
221    let reshapes = place == Place::Free;
222    // R0a — identity elimination
223    if let Some(witness) = r0a_identity::apply(ast) {
224        return ReducibilityFinding {
225            reduction: Some(Reduction::Rewrite {
226                rule: RuleId::R0a,
227                witness,
228            }),
229            rule: Some(RuleId::R0a),
230            improvement: ComplexityDelta::less_compute(),
231        };
232    }
233    // R0b — associativity flattening
234    if let Some(witness) = r0b_flatten::apply(ast) {
235        return ReducibilityFinding {
236            reduction: Some(Reduction::Rewrite {
237                rule: RuleId::R0b,
238                witness,
239            }),
240            rule: Some(RuleId::R0b),
241            improvement: ComplexityDelta::less_compute(),
242        };
243    }
244    // R3 — Lex/filter commute
245    if let Some(witness) = r3_commute::apply(ast) {
246        return ReducibilityFinding {
247            reduction: Some(Reduction::Rewrite {
248                rule: RuleId::R3,
249                witness,
250            }),
251            rule: Some(RuleId::R3),
252            improvement: ComplexityDelta::less_memory(),
253        };
254    }
255    // R4 — filter distributes over union
256    if let Some(witness) = r4_distribute::apply(ast).filter(|_| reshapes) {
257        return ReducibilityFinding {
258            reduction: Some(Reduction::Rewrite {
259                rule: RuleId::R4,
260                witness,
261            }),
262            rule: Some(RuleId::R4),
263            improvement: ComplexityDelta::less_memory(),
264        };
265    }
266    // R5 — per-axis filter pushdown
267    if let Some(witness) =
268        r5_factorize::apply(ast, &|p, c| super::predicate::analyze(p, c)).filter(|_| reshapes)
269    {
270        return ReducibilityFinding {
271            reduction: Some(Reduction::Rewrite {
272                rule: RuleId::R5,
273                witness,
274            }),
275            rule: Some(RuleId::R5),
276            improvement: ComplexityDelta::less_both(),
277        };
278    }
279    // R6 — chained filter folding
280    if let Some(witness) = r6_filter_fold::apply(ast) {
281        return ReducibilityFinding {
282            reduction: Some(Reduction::Rewrite {
283                rule: RuleId::R6,
284                witness,
285            }),
286            rule: Some(RuleId::R6),
287            improvement: ComplexityDelta::less_compute(),
288        };
289    }
290    // R7 — order chain folding
291    if let Some(witness) = r7_order_fold::apply(ast) {
292        return ReducibilityFinding {
293            reduction: Some(Reduction::Rewrite {
294                rule: RuleId::R7,
295                witness,
296            }),
297            rule: Some(RuleId::R7),
298            improvement: ComplexityDelta::less_both(),
299        };
300    }
301    // No rule fires.
302    ReducibilityFinding {
303        reduction: None,
304        rule: None,
305        improvement: ComplexityDelta::equal(),
306    }
307}
308
309/// Replace the i-th child of `ast` with `replacement`. Used by
310/// the bottom-up walker to plumb child rewrites back into the
311/// parent node.
312fn replace_child_at(ast: &Comprehension, i: usize, replacement: Comprehension) -> Comprehension {
313    match ast {
314        Comprehension::Clause { .. } => unreachable!("clause has no children"),
315        Comprehension::Cartesian { children } => {
316            let mut new_children = children.clone();
317            new_children[i] = replacement;
318            Comprehension::Cartesian {
319                children: new_children,
320            }
321        }
322        Comprehension::Zip { children, mode } => {
323            let mut new_children = children.clone();
324            new_children[i] = replacement;
325            Comprehension::Zip {
326                children: new_children,
327                mode: *mode,
328            }
329        }
330        Comprehension::Union { children } => {
331            let mut new_children = children.clone();
332            new_children[i] = replacement;
333            Comprehension::Union {
334                children: new_children,
335            }
336        }
337        Comprehension::Filter { predicate, .. } => Comprehension::Filter {
338            child: Box::new(replacement),
339            predicate: predicate.clone(),
340        },
341        Comprehension::Order {
342            strategy,
343            truncation,
344            seed,
345            ..
346        } => Comprehension::Order {
347            child: Box::new(replacement),
348            strategy: *strategy,
349            truncation: *truncation,
350            seed: *seed,
351        },
352    }
353}
354
355/// Bound on optimizer iterations. Per comprehension_forms.md §10.6
356/// (property 3) the optimizer halts because each rewrite strictly
357/// decreases a well-founded measure. Iterations are bounded as
358/// `node_count^2` to guard against any bug in a rule that
359/// would otherwise loop.
360fn max_steps(ast: &Comprehension) -> usize {
361    let n = ast.node_count();
362    n.saturating_mul(n).saturating_add(16)
363}
364
365/// Convenience: build a `CoordSet` from a comprehension's
366/// coordinate names and its computed metadata. R5 uses this
367/// when invoking the predicate analyzer.
368pub fn coord_set_for(ast: &Comprehension) -> CoordSet {
369    let names = ast.coordinate_names();
370    let metadata = ast.metadata();
371    coord_set_from(&names, &metadata)
372}
373
374fn coord_set_from(names: &[String], metadata: &Metadata) -> CoordSet {
375    CoordSet::from_metadata(names, metadata)
376}
377
378#[cfg(test)]
379mod tests {
380    use super::*;
381    use crate::iteration::comprehension::source::{LiteralValue, Source};
382    use crate::iteration::comprehension::strategy::StrategyName;
383
384    fn clause(name: &str, vs: &[i64]) -> Comprehension {
385        Comprehension::clause(
386            name,
387            Source::Literal {
388                values: vs.iter().map(|n| LiteralValue::Int(*n)).collect(),
389            },
390        )
391    }
392
393    #[test]
394    fn optimize_well_formed_ast_does_not_panic() {
395        let ast = Comprehension::cartesian(vec![clause("k", &[1, 2]), clause("limit", &[10, 20])]);
396        let _ = optimize(ast);
397    }
398
399    #[test]
400    fn optimize_singleton_cartesian_eliminates() {
401        // R0a I2: singleton cartesian → its only child.
402        let ast = Comprehension::cartesian(vec![clause("k", &[1, 2, 3])]);
403        let optimized = optimize(ast);
404        assert!(matches!(optimized, Comprehension::Clause { .. }));
405    }
406
407    #[test]
408    fn optimize_lex_none_eliminates() {
409        // R0a I5: order(c, Lex, None) → c.
410        let inner = clause("k", &[1, 2, 3]);
411        let ast = Comprehension::order(inner.clone(), StrategyName::Lex, None);
412        let optimized = optimize(ast);
413        assert_eq!(optimized, inner);
414    }
415
416    #[test]
417    fn optimize_is_idempotent() {
418        let ast = Comprehension::cartesian(vec![
419            Comprehension::cartesian(vec![clause("a", &[1])]),
420            clause("b", &[2]),
421        ]);
422        let once = optimize(ast.clone());
423        let twice = optimize(once.clone());
424        assert_eq!(once, twice);
425    }
426}