Skip to main content

polydat_grammar/comprehension/
ast.rs

1// Copyright 2024-2026 Jonathan Shook
2// SPDX-License-Identifier: Apache-2.0
3
4//! Operator-tree comprehension AST — spec §3.
5//!
6//! Six constructors closed under composition: one source
7//! (`clause`), three combinators (`cartesian`, `zip`, `union`),
8//! two modifiers (`filter`, `order`). Every comprehension AST is
9//! a tree whose nodes are one of these six variants.
10//!
11//! Closure (spec §4.1):
12//!
13//! - C1 — every constructor returns and consumes
14//!   `Comprehension`. There is no auxiliary value type at the
15//!   AST level.
16//! - C2 — well-formedness is decidable in one bottom-up pass
17//!   (the `validate` module of the runtime
18//!   implements the check).
19
20use serde::{Deserialize, Serialize};
21
22use super::source::Source;
23use super::strategy::{StrategyName, ZipMode};
24
25/// The six-variant operator-tree comprehension.
26///
27/// Closure under composition (spec §4.1 C1): every variant
28/// holds one or more `Comprehension` operands plus
29/// constructor-specific scalar parameters (predicate, strategy,
30/// truncation, zip mode, source).
31///
32/// `Box<Comprehension>` appears wherever a variant needs a
33/// single child operand; `Vec<Comprehension>` wherever a
34/// constructor takes N children (cartesian, zip, union).
35#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
36#[serde(tag = "op", rename_all = "snake_case")]
37pub enum Comprehension {
38    /// Leaf source per spec §3.1. Binds `name` to one value
39    /// per dispense, drawn from `source`.
40    Clause {
41        /// The name bound.
42        name: String,
43        /// Where the values come from.
44        source: Source,
45    },
46
47    /// Cross-product combinator per spec §3.2. Children must
48    /// have disjoint name sets (V1).
49    Cartesian {
50        /// The factors.
51        children: Vec<Comprehension>,
52    },
53
54    /// Lockstep combinator per spec §3.3. Children must be
55    /// discrete (V7) and have disjoint name sets (V1).
56    Zip {
57        /// The streams zipped.
58        children: Vec<Comprehension>,
59        /// The length policy.
60        mode: ZipMode,
61    },
62
63    /// Concatenation combinator per spec §3.4. Children must
64    /// share an identical tuple shape (V2) and all be discrete
65    /// (V9).
66    Union {
67        /// The streams concatenated, in order.
68        children: Vec<Comprehension>,
69    },
70
71    /// Selection modifier per spec §3.5. Predicate is a GK
72    /// boolean expression; names must close over the child's
73    /// coordinates plus the parent scope (V3).
74    Filter {
75        /// The stream filtered.
76        child: Box<Comprehension>,
77        /// The predicate, a boolean expression over the tuple and the parent scope.
78        predicate: String,
79    },
80
81    /// Permutation modifier per spec §3.6. `strategy` must
82    /// accept the child's IndexFn (V4); `truncation` limits the
83    /// dispensed count; `seed` is the authored seed a seeded
84    /// strategy (`Shuffle`, `Lhs`) derives its state from, with a
85    /// fixed default when absent.
86    Order {
87        /// The stream ordered.
88        child: Box<Comprehension>,
89        /// The strategy applied.
90        strategy: StrategyName,
91        /// The dispensed count cap, if any.
92        truncation: Option<u64>,
93        /// The authored seed, read by `Shuffle` and `Lhs`.
94        #[serde(default, skip_serializing_if = "Option::is_none")]
95        seed: Option<u64>,
96    },
97}
98
99impl Comprehension {
100    /// Construct a leaf clause.
101    pub fn clause<S: Into<String>>(name: S, source: Source) -> Self {
102        Comprehension::Clause {
103            name: name.into(),
104            source,
105        }
106    }
107
108    /// Construct a cartesian over the supplied children.
109    pub fn cartesian(children: Vec<Comprehension>) -> Self {
110        Comprehension::Cartesian { children }
111    }
112
113    /// Construct a zip over the supplied children with the
114    /// given mode.
115    pub fn zip(children: Vec<Comprehension>, mode: ZipMode) -> Self {
116        Comprehension::Zip { children, mode }
117    }
118
119    /// Construct a union over the supplied children.
120    pub fn union(children: Vec<Comprehension>) -> Self {
121        Comprehension::Union { children }
122    }
123
124    /// Construct a filter wrapping `child` with `predicate`.
125    pub fn filter<S: Into<String>>(child: Comprehension, predicate: S) -> Self {
126        Comprehension::Filter {
127            child: Box::new(child),
128            predicate: predicate.into(),
129        }
130    }
131
132    /// Construct an order node wrapping `child`, with no authored
133    /// seed.
134    pub fn order(child: Comprehension, strategy: StrategyName, truncation: Option<u64>) -> Self {
135        Self::order_seeded(child, strategy, truncation, None)
136    }
137
138    /// Construct an order node wrapping `child` with an authored
139    /// `seed`, which `Shuffle` and `Lhs` derive their state from.
140    pub fn order_seeded(
141        child: Comprehension,
142        strategy: StrategyName,
143        truncation: Option<u64>,
144        seed: Option<u64>,
145    ) -> Self {
146        Comprehension::Order {
147            child: Box::new(child),
148            strategy,
149            truncation,
150            seed,
151        }
152    }
153
154    /// Compute the comprehension's coordinate name set,
155    /// recursively. The result preserves declaration order
156    /// (per spec §3.2 + §3.4's "in declaration order" tuple
157    /// shape rules). Used by V1, V2, V3, and the predicate
158    /// analyzer's coord-set input.
159    pub fn coordinate_names(&self) -> Vec<String> {
160        let mut acc = Vec::new();
161        self.collect_coordinate_names(&mut acc);
162        acc
163    }
164
165    /// Compute `(coordinate_name, source_text)` pairs in
166    /// declaration order, deduplicated by name (first
167    /// occurrence wins). Source text is the round-trip-to-
168    /// legacy form — `IntRange { 1, 10, 1 }` → `"1..10"`,
169    /// `Literal { [10, 100] }` → `"10, 100"`, etc.
170    ///
171    /// Gives a consumer a `[(var, spec_expr)]` list for type
172    /// detection, the spec-text shape the runtime's probe
173    /// pre-evaluation (`pre_evaluate_clause`) takes.
174    pub fn coordinate_specs(&self) -> Vec<(String, String)> {
175        let mut acc = Vec::new();
176        let mut seen = std::collections::HashSet::new();
177        self.collect_coordinate_specs(&mut acc, &mut seen);
178        acc
179    }
180
181    /// Grammar-based extraction of the free names referenced by
182    /// every source spec in this comprehension subtree —
183    /// workload params, outer iter-vars, and wires that a
184    /// `Generator` spec (`concat(foo)`, bare `eh_values`)
185    /// consumes. Each spec is parsed with the canonical Polydat
186    /// expression grammar (`crate::refs::referenced_names`)
187    /// rather than byte-scanned, so a bare source reference is
188    /// recognised exactly as the kernel compiler would resolve
189    /// it. `WorkloadParamList { name }` contributes `name`
190    /// directly; literals / ranges / intervals contribute
191    /// nothing. Used by the workload validator's
192    /// declared-but-unreferenced check.
193    pub fn referenced_source_names(&self) -> std::collections::BTreeSet<String> {
194        let mut out = std::collections::BTreeSet::new();
195        self.walk_sources(&mut |source| out.extend(source.referenced_names()));
196        out
197    }
198
199    /// Visit every leaf [`Source`] in this comprehension subtree.
200    fn walk_sources(&self, visit: &mut impl FnMut(&super::source::Source)) {
201        match self {
202            Comprehension::Clause { source, .. } => visit(source),
203            Comprehension::Cartesian { children }
204            | Comprehension::Zip { children, .. }
205            | Comprehension::Union { children } => {
206                for c in children {
207                    c.walk_sources(visit);
208                }
209            }
210            Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
211                child.walk_sources(visit);
212            }
213        }
214    }
215
216    fn collect_coordinate_specs(
217        &self,
218        acc: &mut Vec<(String, String)>,
219        seen: &mut std::collections::HashSet<String>,
220    ) {
221        match self {
222            Comprehension::Clause { name, source } => {
223                if seen.insert(name.clone()) {
224                    let spec_text = source.to_text().unwrap_or_else(|| "<source>".to_string());
225                    acc.push((name.clone(), spec_text));
226                }
227            }
228            Comprehension::Cartesian { children } | Comprehension::Zip { children, .. } => {
229                for c in children {
230                    c.collect_coordinate_specs(acc, seen);
231                }
232            }
233            Comprehension::Union { children } => {
234                for c in children {
235                    c.collect_coordinate_specs(acc, seen);
236                }
237            }
238            Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
239                child.collect_coordinate_specs(acc, seen);
240            }
241        }
242    }
243
244    fn collect_coordinate_names(&self, acc: &mut Vec<String>) {
245        match self {
246            Comprehension::Clause { name, .. } => {
247                if !acc.contains(name) {
248                    acc.push(name.clone());
249                }
250            }
251            Comprehension::Cartesian { children } | Comprehension::Zip { children, .. } => {
252                for c in children {
253                    c.collect_coordinate_names(acc);
254                }
255            }
256            Comprehension::Union { children } => {
257                // V2 requires identical shape; take the first
258                // child's coordinates as canonical.
259                if let Some(first) = children.first() {
260                    first.collect_coordinate_names(acc);
261                }
262            }
263            Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
264                child.collect_coordinate_names(acc);
265            }
266        }
267    }
268
269    /// `true` if this node is a leaf clause.
270    pub fn is_clause(&self) -> bool {
271        matches!(self, Comprehension::Clause { .. })
272    }
273
274    /// `true` if this node is one of the three combinators.
275    pub fn is_combinator(&self) -> bool {
276        matches!(
277            self,
278            Comprehension::Cartesian { .. }
279                | Comprehension::Zip { .. }
280                | Comprehension::Union { .. }
281        )
282    }
283
284    /// `true` if this node is a modifier (`filter` or `order`).
285    pub fn is_modifier(&self) -> bool {
286        matches!(
287            self,
288            Comprehension::Filter { .. } | Comprehension::Order { .. }
289        )
290    }
291
292    /// Iterate this node's direct operand children. Returns
293    /// an empty iterator for leaf clauses.
294    pub fn children(&self) -> Box<dyn Iterator<Item = &Comprehension> + '_> {
295        match self {
296            Comprehension::Clause { .. } => Box::new(std::iter::empty()),
297            Comprehension::Cartesian { children }
298            | Comprehension::Zip { children, .. }
299            | Comprehension::Union { children } => Box::new(children.iter()),
300            Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
301                Box::new(std::iter::once(child.as_ref()))
302            }
303        }
304    }
305
306    /// Count of nodes in the AST (this node + all descendants).
307    /// Used by the optimizer's well-founded measure for
308    /// termination (spec §10.6.3).
309    pub fn node_count(&self) -> usize {
310        1 + self.children().map(|c| c.node_count()).sum::<usize>()
311    }
312
313    /// Maximum depth of the AST. Constant for flat composition,
314    /// O(log N) for balanced trees. Bounds the operator stack
315    /// per spec §9.3.
316    pub fn depth(&self) -> usize {
317        1 + self.children().map(|c| c.depth()).max().unwrap_or(0)
318    }
319}
320
321#[cfg(test)]
322mod tests {
323    use super::*;
324    use crate::comprehension::source::{LiteralValue, Source};
325
326    fn lit_int_clause(name: &str, values: &[i64]) -> Comprehension {
327        Comprehension::clause(
328            name,
329            Source::Literal {
330                values: values.iter().map(|n| LiteralValue::Int(*n)).collect(),
331            },
332        )
333    }
334
335    #[test]
336    fn clause_coordinates() {
337        let c = lit_int_clause("k", &[1, 2, 3]);
338        assert_eq!(c.coordinate_names(), vec!["k"]);
339        assert!(c.is_clause());
340        assert!(!c.is_combinator());
341        assert!(!c.is_modifier());
342    }
343
344    #[test]
345    fn continuous_interval_spec_text_round_trips_as_float() {
346        // A continuous `[1.0, 5.0)` must reconstruct to a spec the
347        // source parser re-classifies as CONTINUOUS, not an integer
348        // range. `{}` on `1.0f64` prints "1", so the reconstruction
349        // must use a float-preserving format ("1.0..5.0"), else a
350        // downstream type-probe re-parses "1..5" and types the
351        // iter-var `U64` — silently corrupting a float optimize axis.
352        use crate::comprehension::cardinality::{Interval, ProductMeasure};
353        let c = Comprehension::clause(
354            "ef",
355            Source::ContinuousInterval {
356                interval: Interval {
357                    lo: 1.0,
358                    hi: 5.0,
359                    lo_open: false,
360                    hi_open: true,
361                },
362                measure: ProductMeasure::Uniform,
363            },
364        );
365        let (var, spec_text) = c.coordinate_specs().into_iter().next().unwrap();
366        assert_eq!(var, "ef");
367        // Re-parsing the reconstructed text must yield a continuous
368        // interval again — the round-trip the kernel-type probe relies on.
369        let reparsed = crate::comprehension::spec::parse_source(&spec_text).unwrap();
370        assert!(
371            matches!(reparsed, Source::ContinuousInterval { .. }),
372            "reconstructed '{spec_text}' re-parsed to {reparsed:?}, expected ContinuousInterval"
373        );
374    }
375
376    #[test]
377    fn referenced_source_names_grammar_based() {
378        // `eh in eh_values` — a bare source reference parses to
379        // a Generator whose free name is the workload param.
380        let bare = Comprehension::clause(
381            "eh",
382            Source::Generator {
383                expr: "eh_values".into(),
384                cardinality_hint: None,
385            },
386        );
387        let got: Vec<String> = bare.referenced_source_names().into_iter().collect();
388        assert_eq!(got, vec!["eh_values"]);
389
390        // `(nbo) in (concat(nbo_v_values))` — the source is a
391        // function call; the callee `concat` is NOT a reference
392        // but its argument IS.
393        let call = Comprehension::clause(
394            "nbo",
395            Source::Generator {
396                expr: "concat(nbo_v_values)".into(),
397                cardinality_hint: None,
398            },
399        );
400        let got: Vec<String> = call.referenced_source_names().into_iter().collect();
401        assert_eq!(got, vec!["nbo_v_values"]);
402
403        // `{profiles}` — an explicit WorkloadParamList contributes
404        // its name directly.
405        let wpl = Comprehension::clause(
406            "p",
407            Source::WorkloadParamList {
408                name: "profiles".into(),
409                len_hint: None,
410            },
411        );
412        let got: Vec<String> = wpl.referenced_source_names().into_iter().collect();
413        assert_eq!(got, vec!["profiles"]);
414
415        // Literal sources contribute nothing.
416        let lit = lit_int_clause("k", &[1, 2, 3]);
417        assert!(lit.referenced_source_names().is_empty());
418
419        // Cartesian unions the per-clause references.
420        let cart = Comprehension::cartesian(vec![bare, call]);
421        let got: Vec<String> = cart.referenced_source_names().into_iter().collect();
422        assert_eq!(got, vec!["eh_values", "nbo_v_values"]);
423    }
424
425    #[test]
426    fn cartesian_coordinates_in_declaration_order() {
427        let c = Comprehension::cartesian(vec![
428            lit_int_clause("k", &[1, 2]),
429            lit_int_clause("limit", &[10, 20, 30]),
430        ]);
431        assert_eq!(c.coordinate_names(), vec!["k", "limit"]);
432        assert!(c.is_combinator());
433    }
434
435    #[test]
436    fn zip_coordinates() {
437        let c = Comprehension::zip(
438            vec![
439                lit_int_clause("x", &[1, 2, 3]),
440                lit_int_clause("y", &[10, 20, 30]),
441            ],
442            ZipMode::Strict,
443        );
444        assert_eq!(c.coordinate_names(), vec!["x", "y"]);
445    }
446
447    #[test]
448    fn union_takes_first_childs_shape() {
449        let a = Comprehension::cartesian(vec![
450            lit_int_clause("k", &[10]),
451            lit_int_clause("limit", &[10, 20]),
452        ]);
453        let b = Comprehension::cartesian(vec![
454            lit_int_clause("k", &[100]),
455            lit_int_clause("limit", &[100, 200]),
456        ]);
457        let u = Comprehension::union(vec![a, b]);
458        assert_eq!(u.coordinate_names(), vec!["k", "limit"]);
459    }
460
461    #[test]
462    fn filter_and_order_pass_through_coordinates() {
463        let inner = Comprehension::cartesian(vec![
464            lit_int_clause("k", &[1, 2]),
465            lit_int_clause("limit", &[10]),
466        ]);
467        let filtered = Comprehension::filter(inner.clone(), "{k} > 0");
468        assert_eq!(filtered.coordinate_names(), vec!["k", "limit"]);
469        assert!(filtered.is_modifier());
470
471        let ordered = Comprehension::order(inner, StrategyName::Lex, Some(5));
472        assert_eq!(ordered.coordinate_names(), vec!["k", "limit"]);
473        assert!(ordered.is_modifier());
474    }
475
476    #[test]
477    fn node_count_and_depth() {
478        let inner = Comprehension::cartesian(vec![
479            lit_int_clause("k", &[1, 2]),
480            lit_int_clause("limit", &[10]),
481        ]);
482        // inner: 1 (cartesian) + 2 (clauses) = 3 nodes; depth 2
483        assert_eq!(inner.node_count(), 3);
484        assert_eq!(inner.depth(), 2);
485
486        let filtered = Comprehension::filter(inner, "{k} > 0");
487        // filtered: 1 (filter) + 3 = 4 nodes; depth 3
488        assert_eq!(filtered.node_count(), 4);
489        assert_eq!(filtered.depth(), 3);
490    }
491
492    #[test]
493    fn round_trip_serde() {
494        let c = Comprehension::order(
495            Comprehension::filter(
496                Comprehension::cartesian(vec![
497                    lit_int_clause("k", &[1, 2, 3]),
498                    lit_int_clause("limit", &[10, 20]),
499                ]),
500                "{k} * {limit} > 5",
501            ),
502            StrategyName::Halton,
503            Some(10),
504        );
505        let json = serde_json::to_string(&c).unwrap();
506        let back: Comprehension = serde_json::from_str(&json).unwrap();
507        assert_eq!(c, back);
508    }
509
510    /// An order's authored seed rides through serde, and an order
511    /// written without one reads back as unseeded.
512    #[test]
513    fn an_orders_seed_round_trips_and_defaults_to_none() {
514        let c = Comprehension::order_seeded(
515            Comprehension::clause(
516                "k",
517                Source::IntRange {
518                    lo: 1,
519                    hi: 4,
520                    step: 1,
521                },
522            ),
523            StrategyName::Shuffle,
524            Some(2),
525            Some(42),
526        );
527        let json = serde_json::to_string(&c).unwrap();
528        assert!(json.contains("\"seed\":42"), "{json}");
529        let back: Comprehension = serde_json::from_str(&json).unwrap();
530        assert_eq!(back, c);
531        let unseeded = json.replace(",\"seed\":42", "");
532        let back: Comprehension = serde_json::from_str(&unseeded).unwrap();
533        assert!(
534            matches!(back, Comprehension::Order { seed: None, .. }),
535            "{back:?}"
536        );
537    }
538}