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    /// The names every source in this subtree reads when it is evaluated
201    /// ([`Source::names_read`]): a composed name's leaves, and a cursor's
202    /// extent outputs for `all(<cursor>)`. What a clause's tuples depend
203    /// on, and what a traversal's scope supplies to its sources.
204    pub fn source_names_read(&self) -> std::collections::BTreeSet<String> {
205        let mut out = std::collections::BTreeSet::new();
206        self.walk_sources(&mut |source| out.extend(source.names_read()));
207        out
208    }
209
210    /// Visit every leaf [`Source`] in this comprehension subtree.
211    fn walk_sources(&self, visit: &mut impl FnMut(&super::source::Source)) {
212        match self {
213            Comprehension::Clause { source, .. } => visit(source),
214            Comprehension::Cartesian { children }
215            | Comprehension::Zip { children, .. }
216            | Comprehension::Union { children } => {
217                for c in children {
218                    c.walk_sources(visit);
219                }
220            }
221            Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
222                child.walk_sources(visit);
223            }
224        }
225    }
226
227    fn collect_coordinate_specs(
228        &self,
229        acc: &mut Vec<(String, String)>,
230        seen: &mut std::collections::HashSet<String>,
231    ) {
232        match self {
233            Comprehension::Clause { name, source } => {
234                if seen.insert(name.clone()) {
235                    let spec_text = source.to_text().unwrap_or_else(|| "<source>".to_string());
236                    acc.push((name.clone(), spec_text));
237                }
238            }
239            Comprehension::Cartesian { children } | Comprehension::Zip { children, .. } => {
240                for c in children {
241                    c.collect_coordinate_specs(acc, seen);
242                }
243            }
244            Comprehension::Union { children } => {
245                for c in children {
246                    c.collect_coordinate_specs(acc, seen);
247                }
248            }
249            Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
250                child.collect_coordinate_specs(acc, seen);
251            }
252        }
253    }
254
255    fn collect_coordinate_names(&self, acc: &mut Vec<String>) {
256        match self {
257            Comprehension::Clause { name, .. } => {
258                if !acc.contains(name) {
259                    acc.push(name.clone());
260                }
261            }
262            Comprehension::Cartesian { children } | Comprehension::Zip { children, .. } => {
263                for c in children {
264                    c.collect_coordinate_names(acc);
265                }
266            }
267            Comprehension::Union { children } => {
268                // V2 requires identical shape; take the first
269                // child's coordinates as canonical.
270                if let Some(first) = children.first() {
271                    first.collect_coordinate_names(acc);
272                }
273            }
274            Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
275                child.collect_coordinate_names(acc);
276            }
277        }
278    }
279
280    /// `true` if this node is a leaf clause.
281    pub fn is_clause(&self) -> bool {
282        matches!(self, Comprehension::Clause { .. })
283    }
284
285    /// `true` if this node is one of the three combinators.
286    pub fn is_combinator(&self) -> bool {
287        matches!(
288            self,
289            Comprehension::Cartesian { .. }
290                | Comprehension::Zip { .. }
291                | Comprehension::Union { .. }
292        )
293    }
294
295    /// `true` if this node is a modifier (`filter` or `order`).
296    pub fn is_modifier(&self) -> bool {
297        matches!(
298            self,
299            Comprehension::Filter { .. } | Comprehension::Order { .. }
300        )
301    }
302
303    /// Iterate this node's direct operand children. Returns
304    /// an empty iterator for leaf clauses.
305    pub fn children(&self) -> Box<dyn Iterator<Item = &Comprehension> + '_> {
306        match self {
307            Comprehension::Clause { .. } => Box::new(std::iter::empty()),
308            Comprehension::Cartesian { children }
309            | Comprehension::Zip { children, .. }
310            | Comprehension::Union { children } => Box::new(children.iter()),
311            Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
312                Box::new(std::iter::once(child.as_ref()))
313            }
314        }
315    }
316
317    /// Count of nodes in the AST (this node + all descendants).
318    /// Used by the optimizer's well-founded measure for
319    /// termination (comprehension_forms.md §10.6, property 3).
320    pub fn node_count(&self) -> usize {
321        1 + self.children().map(|c| c.node_count()).sum::<usize>()
322    }
323
324    /// Maximum depth of the AST. Constant for flat composition,
325    /// O(log N) for balanced trees. Bounds the operator stack
326    /// (comprehension_forms.md §9.3).
327    pub fn depth(&self) -> usize {
328        1 + self.children().map(|c| c.depth()).max().unwrap_or(0)
329    }
330}
331
332#[cfg(test)]
333mod tests {
334    use super::*;
335    use crate::comprehension::source::{LiteralValue, Source};
336
337    fn lit_int_clause(name: &str, values: &[i64]) -> Comprehension {
338        Comprehension::clause(
339            name,
340            Source::Literal {
341                values: values.iter().map(|n| LiteralValue::Int(*n)).collect(),
342            },
343        )
344    }
345
346    #[test]
347    fn clause_coordinates() {
348        let c = lit_int_clause("k", &[1, 2, 3]);
349        assert_eq!(c.coordinate_names(), vec!["k"]);
350        assert!(c.is_clause());
351        assert!(!c.is_combinator());
352        assert!(!c.is_modifier());
353    }
354
355    #[test]
356    fn continuous_interval_spec_text_round_trips_as_float() {
357        // A continuous `[1.0, 5.0)` must reconstruct to a spec the
358        // source parser re-classifies as CONTINUOUS, not an integer
359        // range. `{}` on `1.0f64` prints "1", so the reconstruction
360        // must use a float-preserving format ("1.0..5.0"), else a
361        // downstream type-probe re-parses "1..5" and types the
362        // iter-var `U64` — silently corrupting a float optimize axis.
363        use crate::comprehension::cardinality::{Interval, ProductMeasure};
364        let c = Comprehension::clause(
365            "ef",
366            Source::ContinuousInterval {
367                interval: Interval {
368                    lo: 1.0,
369                    hi: 5.0,
370                    lo_open: false,
371                    hi_open: true,
372                },
373                measure: ProductMeasure::Uniform,
374            },
375        );
376        let (var, spec_text) = c.coordinate_specs().into_iter().next().unwrap();
377        assert_eq!(var, "ef");
378        // Re-parsing the reconstructed text must yield a continuous
379        // interval again — the round-trip the kernel-type probe relies on.
380        let reparsed = crate::comprehension::spec::parse_source(&spec_text).unwrap();
381        assert!(
382            matches!(reparsed, Source::ContinuousInterval { .. }),
383            "reconstructed '{spec_text}' re-parsed to {reparsed:?}, expected ContinuousInterval"
384        );
385    }
386
387    #[test]
388    fn referenced_source_names_grammar_based() {
389        // `eh in eh_values` — a bare source reference parses to
390        // a Generator whose free name is the workload param.
391        let bare = Comprehension::clause(
392            "eh",
393            Source::Generator {
394                expr: "eh_values".into(),
395                cardinality_hint: None,
396            },
397        );
398        let got: Vec<String> = bare.referenced_source_names().into_iter().collect();
399        assert_eq!(got, vec!["eh_values"]);
400
401        // `(nbo) in (concat(nbo_v_values))` — the source is a
402        // function call; the callee `concat` is NOT a reference
403        // but its argument IS.
404        let call = Comprehension::clause(
405            "nbo",
406            Source::Generator {
407                expr: "concat(nbo_v_values)".into(),
408                cardinality_hint: None,
409            },
410        );
411        let got: Vec<String> = call.referenced_source_names().into_iter().collect();
412        assert_eq!(got, vec!["nbo_v_values"]);
413
414        // `{profiles}` — an explicit WorkloadParamList contributes
415        // its name directly.
416        let wpl = Comprehension::clause(
417            "p",
418            Source::WorkloadParamList {
419                name: "profiles".into(),
420                len_hint: None,
421            },
422        );
423        let got: Vec<String> = wpl.referenced_source_names().into_iter().collect();
424        assert_eq!(got, vec!["profiles"]);
425
426        // Literal sources contribute nothing.
427        let lit = lit_int_clause("k", &[1, 2, 3]);
428        assert!(lit.referenced_source_names().is_empty());
429
430        // Cartesian unions the per-clause references.
431        let cart = Comprehension::cartesian(vec![bare, call]);
432        let got: Vec<String> = cart.referenced_source_names().into_iter().collect();
433        assert_eq!(got, vec!["eh_values", "nbo_v_values"]);
434    }
435
436    #[test]
437    fn cartesian_coordinates_in_declaration_order() {
438        let c = Comprehension::cartesian(vec![
439            lit_int_clause("k", &[1, 2]),
440            lit_int_clause("limit", &[10, 20, 30]),
441        ]);
442        assert_eq!(c.coordinate_names(), vec!["k", "limit"]);
443        assert!(c.is_combinator());
444    }
445
446    #[test]
447    fn zip_coordinates() {
448        let c = Comprehension::zip(
449            vec![
450                lit_int_clause("x", &[1, 2, 3]),
451                lit_int_clause("y", &[10, 20, 30]),
452            ],
453            ZipMode::Strict,
454        );
455        assert_eq!(c.coordinate_names(), vec!["x", "y"]);
456    }
457
458    #[test]
459    fn union_takes_first_childs_shape() {
460        let a = Comprehension::cartesian(vec![
461            lit_int_clause("k", &[10]),
462            lit_int_clause("limit", &[10, 20]),
463        ]);
464        let b = Comprehension::cartesian(vec![
465            lit_int_clause("k", &[100]),
466            lit_int_clause("limit", &[100, 200]),
467        ]);
468        let u = Comprehension::union(vec![a, b]);
469        assert_eq!(u.coordinate_names(), vec!["k", "limit"]);
470    }
471
472    #[test]
473    fn filter_and_order_pass_through_coordinates() {
474        let inner = Comprehension::cartesian(vec![
475            lit_int_clause("k", &[1, 2]),
476            lit_int_clause("limit", &[10]),
477        ]);
478        let filtered = Comprehension::filter(inner.clone(), "{k} > 0");
479        assert_eq!(filtered.coordinate_names(), vec!["k", "limit"]);
480        assert!(filtered.is_modifier());
481
482        let ordered = Comprehension::order(inner, StrategyName::Lex, Some(5));
483        assert_eq!(ordered.coordinate_names(), vec!["k", "limit"]);
484        assert!(ordered.is_modifier());
485    }
486
487    #[test]
488    fn node_count_and_depth() {
489        let inner = Comprehension::cartesian(vec![
490            lit_int_clause("k", &[1, 2]),
491            lit_int_clause("limit", &[10]),
492        ]);
493        // inner: 1 (cartesian) + 2 (clauses) = 3 nodes; depth 2
494        assert_eq!(inner.node_count(), 3);
495        assert_eq!(inner.depth(), 2);
496
497        let filtered = Comprehension::filter(inner, "{k} > 0");
498        // filtered: 1 (filter) + 3 = 4 nodes; depth 3
499        assert_eq!(filtered.node_count(), 4);
500        assert_eq!(filtered.depth(), 3);
501    }
502
503    #[test]
504    fn round_trip_serde() {
505        let c = Comprehension::order(
506            Comprehension::filter(
507                Comprehension::cartesian(vec![
508                    lit_int_clause("k", &[1, 2, 3]),
509                    lit_int_clause("limit", &[10, 20]),
510                ]),
511                "{k} * {limit} > 5",
512            ),
513            StrategyName::Halton,
514            Some(10),
515        );
516        let json = serde_json::to_string(&c).unwrap();
517        let back: Comprehension = serde_json::from_str(&json).unwrap();
518        assert_eq!(c, back);
519    }
520
521    /// An order's authored seed rides through serde, and an order
522    /// written without one reads back as unseeded.
523    #[test]
524    fn an_orders_seed_round_trips_and_defaults_to_none() {
525        let c = Comprehension::order_seeded(
526            Comprehension::clause(
527                "k",
528                Source::IntRange {
529                    lo: 1,
530                    hi: 4,
531                    step: 1,
532                },
533            ),
534            StrategyName::Shuffle,
535            Some(2),
536            Some(42),
537        );
538        let json = serde_json::to_string(&c).unwrap();
539        assert!(json.contains("\"seed\":42"), "{json}");
540        let back: Comprehension = serde_json::from_str(&json).unwrap();
541        assert_eq!(back, c);
542        let unseeded = json.replace(",\"seed\":42", "");
543        let back: Comprehension = serde_json::from_str(&unseeded).unwrap();
544        assert!(
545            matches!(back, Comprehension::Order { seed: None, .. }),
546            "{back:?}"
547        );
548    }
549}