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