Skip to main content

polydat_grammar/comprehension/
source.rs

1// Copyright 2024-2026 Jonathan Shook
2// SPDX-License-Identifier: Apache-2.0
3
4//! Clause source values (comprehension_forms.md §3.1).
5//!
6//! A `clause(name, source)` binds a name to the values
7//! produced by its source. Sources split into two families:
8//!
9//! - **Discrete stream producers** — literal lists, integer
10//!   ranges, generator functions, workload-param references.
11//!   Cardinality is `Bounded`, `BoundedAtMost`, or `Unbounded`.
12//! - **Continuous measures** — real intervals with an
13//!   integrable measure (uniform on bounded intervals; named
14//!   probability distributions like Normal / Exponential).
15//!   Cardinality is `Continuous`; V8 requires an enclosing
16//!   sampling `order(_, strategy, Some(n))` before dispense.
17//!
18//! Sources are stream producers — they do not pre-materialize
19//! into `Vec<Value>`. This is the load-bearing model property
20//! of comprehension_forms.md §3.1 and §6.2.
21
22use serde::{Deserialize, Serialize};
23
24use super::cardinality::{CardinalityClass, Interval, MeasureName, ProductMeasure};
25
26/// A clause's source of values.
27///
28/// Discrete variants produce a stream of `Value` via the
29/// runtime evaluator; continuous variants describe a measure
30/// that a downstream sampling strategy will draw from.
31#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
32#[serde(tag = "kind", rename_all = "snake_case")]
33pub enum Source {
34    /// Literal comma list (e.g., `[1, 2, 4, 8]`). Stream
35    /// producer over the list contents.
36    Literal {
37        /// The values, in order.
38        values: Vec<LiteralValue>,
39    },
40
41    /// Integer half-open range `lo..hi` with optional step.
42    /// Default step is 1.
43    IntRange {
44        /// The first value.
45        lo: i64,
46        /// One past the last.
47        hi: i64,
48        /// The step between values, 1 by default.
49        step: i64,
50    },
51
52    /// Generator function call expressed as a Polydat source string.
53    /// Its eval class follows its free names (comprehension_forms.md §10.7.0): an
54    /// expression that references no name is context-free and is
55    /// evaluated at compile, so a clause over it becomes a literal of
56    /// its values; one that references a name (an outer coordinate,
57    /// a parameter, a wire) is evaluated at traversal.
58    Generator {
59        /// The generator call, as Polydat source.
60        expr: String,
61        /// How many values it yields, when known: the count the
62        /// compile established by evaluating a context-free call
63        /// whose values are not literal-representable (a partition
64        /// list, for one). `None` for a call evaluated at traversal,
65        /// whose cardinality is `Unbounded` until then.
66        cardinality_hint: Option<u64>,
67    },
68
69    /// Reference to a workload-level parameter that resolves to
70    /// a list of values. Cardinality is the parameter's
71    /// declared list length.
72    WorkloadParamList {
73        /// The parameter's name.
74        name: String,
75        /// The list's length, when known.
76        len_hint: Option<u64>,
77    },
78
79    /// Real interval (continuous source). Combined with a
80    /// `measure` to form a `Continuous` cardinality.
81    /// Integrability is checked at parse via V8.
82    ContinuousInterval {
83        /// The interval.
84        interval: Interval,
85        /// The measure drawn from.
86        measure: ProductMeasure,
87    },
88
89    /// Named continuous distribution. The distribution carries
90    /// its own support; the `support` field records the
91    /// effective interval for V8's check.
92    Distribution {
93        /// The distribution.
94        distribution: MeasureName,
95        /// Its effective support.
96        support: Interval,
97        /// Its parameters, in the order of
98        /// [`MeasureName::parameter_names`]; empty for the standard
99        /// parameters ([`MeasureName::default_params`]).
100        params: Vec<f64>,
101    },
102}
103
104/// A literal value carried in a `Source::Literal`. Subset of
105/// the polydat `Value` type — the kinds clauses can directly
106/// bind. Extension to richer value types lives in the source
107/// evaluator, not the AST.
108///
109/// Serialized untagged because the variants are primitives;
110/// the JSON/YAML representation is just the bare value
111/// (`1` / `"x"` / `true` / `1.5`).
112#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
113#[serde(untagged)]
114pub enum LiteralValue {
115    /// An integer.
116    Int(i64),
117    /// An integer above `i64::MAX`, which only a `u64` holds. An
118    /// integer that fits `i64` is `Int`, so every integer has one
119    /// form ([`LiteralValue::unsigned`]). After `Int`, so an untagged
120    /// read takes an integer as `Int` when it fits.
121    UInt(u64),
122    /// A float.
123    Float(f64),
124    /// A string.
125    String(String),
126    /// A boolean.
127    Bool(bool),
128    /// A JSON value carrying its own kind: an item of a JSON list a
129    /// generator supplied at run time, bound where the element is
130    /// declared `json`. Last, so an untagged read tries the scalar
131    /// forms first.
132    Json(serde_json::Value),
133}
134
135impl LiteralValue {
136    /// The literal of an unsigned integer: `Int` when it fits `i64`,
137    /// `UInt` above.
138    pub fn unsigned(n: u64) -> Self {
139        match i64::try_from(n) {
140            Ok(n) => LiteralValue::Int(n),
141            Err(_) => LiteralValue::UInt(n),
142        }
143    }
144}
145
146impl Source {
147    /// The names this source references (comprehension_forms.md §10.7.0): a generator
148    /// expression's parsed free identifiers (`concat(foo)`) and its
149    /// `{name}` interpolation placeholders, and a workload parameter
150    /// list's own name. Literals, ranges, and intervals reference
151    /// nothing. A source with no references is context-free.
152    pub fn referenced_names(&self) -> std::collections::BTreeSet<String> {
153        let mut out = std::collections::BTreeSet::new();
154        match self {
155            Source::WorkloadParamList { name, .. } => {
156                out.insert(name.clone());
157            }
158            Source::Generator { expr, .. } => {
159                out.extend(crate::refs::referenced_names(expr));
160                crate::refs::collect_string_interpolation_refs(expr, &mut out);
161            }
162            Source::Literal { .. }
163            | Source::IntRange { .. }
164            | Source::ContinuousInterval { .. }
165            | Source::Distribution { .. } => {}
166        }
167        out
168    }
169
170    /// The names evaluating this source reads from the prior axes and
171    /// the scope it is evaluated in (comprehension_forms.md §5 V3,
172    /// §10.9.1). This differs from [`Self::referenced_names`] in two
173    /// forms. A composed name (`{k_{k}_limits}`) reads its leaves
174    /// (`k`) here; the name they compose to is known only once the
175    /// leaves are bound, and is read when the source is evaluated. The
176    /// cursor form `all(<cursor>)` reads the cursor's extent outputs
177    /// ([`cursor_extent_names`]), not a value named after the cursor.
178    pub fn names_read(&self) -> std::collections::BTreeSet<String> {
179        let mut out = std::collections::BTreeSet::new();
180        match self {
181            Source::WorkloadParamList { name, .. } => {
182                crate::refs::collect_string_interpolation_refs(&format!("{{{name}}}"), &mut out);
183            }
184            Source::Generator { expr, .. } => match all_cursor_argument(expr) {
185                Some(cursor) => out.extend(cursor_extent_names(cursor)),
186                None => out.extend(self.referenced_names()),
187            },
188            Source::Literal { .. }
189            | Source::IntRange { .. }
190            | Source::ContinuousInterval { .. }
191            | Source::Distribution { .. } => {}
192        }
193        out
194    }
195
196    /// Declare this source's cardinality class for use by
197    /// `clause` metadata propagation.
198    pub fn cardinality(&self) -> CardinalityClass {
199        match self {
200            Source::Literal { values } => CardinalityClass::Bounded(values.len() as u64),
201            Source::IntRange { lo, hi, step } => {
202                let step = (*step).max(1).unsigned_abs();
203                if hi <= lo {
204                    CardinalityClass::Bounded(0)
205                } else {
206                    let span = (hi - lo) as u64;
207                    let n = span.div_ceil(step);
208                    CardinalityClass::Bounded(n)
209                }
210            }
211            Source::Generator {
212                cardinality_hint, ..
213            } => match cardinality_hint {
214                Some(n) => CardinalityClass::Bounded(*n),
215                None => CardinalityClass::Unbounded,
216            },
217            Source::WorkloadParamList { len_hint, .. } => match len_hint {
218                Some(n) => CardinalityClass::Bounded(*n),
219                None => CardinalityClass::Unbounded,
220            },
221            Source::ContinuousInterval { interval, measure } => CardinalityClass::Continuous {
222                intervals: vec![interval.clone()],
223                measure: measure.clone(),
224            },
225            Source::Distribution { support, .. } => CardinalityClass::Continuous {
226                intervals: vec![support.clone()],
227                measure: ProductMeasure::Named(*self.distribution_name()),
228            },
229        }
230    }
231
232    /// `true` if this source is continuous (Continuous /
233    /// Distribution variants). Used by V7 (zip must be all
234    /// discrete) and V9 (union must be all discrete) without
235    /// a full cardinality computation.
236    pub fn is_continuous(&self) -> bool {
237        matches!(
238            self,
239            Source::ContinuousInterval { .. } | Source::Distribution { .. }
240        )
241    }
242
243    /// `true` if this source is discrete (every variant except
244    /// the continuous ones).
245    pub fn is_discrete(&self) -> bool {
246        !self.is_continuous()
247    }
248
249    fn distribution_name(&self) -> &MeasureName {
250        match self {
251            Source::Distribution { distribution, .. } => distribution,
252            _ => panic!("distribution_name called on non-Distribution source"),
253        }
254    }
255}
256
257// ── iteration interior and string-comprehension striping
258//    (comprehension_forms.md §3.1.2, §3.1.3) ──
259
260/// The string-comprehension separator rule (comprehension_forms.md
261/// §3.1.3), in one place
262/// so the parse-time (`source_parser`) and runtime (`eval`)
263/// striping can never drift: split on runs of comma / semicolon /
264/// ASCII whitespace, trim, drop empties. Every other character
265/// (`:` `.` `-` `/` …) stays in the token. Returns the raw token
266/// substrings; callers type them (Value or LiteralValue).
267pub fn split_string_comprehension(s: &str) -> Vec<&str> {
268    s.split(|c: char| c == ',' || c == ';' || c.is_ascii_whitespace())
269        .map(str::trim)
270        .filter(|t| !t.is_empty())
271        .collect()
272}
273
274// ── the cursor form `all(<cursor>)` (comprehension_forms.md §10.9.1) ──
275
276/// The cursor `text` enumerates when it is the source form
277/// `all(<cursor>)`, whitespace aside; `None` for any other text.
278pub fn all_cursor_argument(text: &str) -> Option<&str> {
279    let cursor = text.trim().strip_prefix("all(")?.strip_suffix(')')?.trim();
280    let mut chars = cursor.chars();
281    let starts = chars
282        .next()
283        .is_some_and(|c| c.is_ascii_alphabetic() || c == '_');
284    (starts && chars.all(|c| c.is_ascii_alphanumeric() || c == '_')).then_some(cursor)
285}
286
287/// The auxiliary outputs a `cursor <name> = ...` declaration compiles
288/// to, holding its extent: `[start, end]`. `all(<cursor>)` reads these.
289pub fn cursor_extent_names(cursor: &str) -> [String; 2] {
290    [
291        format!("__cursor_extent_{cursor}_start"),
292        format!("__cursor_extent_{cursor}_end"),
293    ]
294}
295
296/// The cursor whose extent the output `name` holds, when `name` is one
297/// of [`cursor_extent_names`].
298pub fn cursor_of_extent_name(name: &str) -> Option<&str> {
299    let rest = name.strip_prefix("__cursor_extent_")?;
300    rest.strip_suffix("_start")
301        .or_else(|| rest.strip_suffix("_end"))
302        .filter(|cursor| !cursor.is_empty())
303}
304
305#[cfg(test)]
306mod tests {
307    use super::*;
308
309    /// A composed name reads its leaves, and the cursor form reads the
310    /// cursor's extent; every other source reads what it references.
311    #[test]
312    fn names_read_are_the_leaves_of_a_composition_and_a_cursor_extent() {
313        let names = |s: Source| s.names_read().into_iter().collect::<Vec<_>>();
314        assert_eq!(
315            names(Source::WorkloadParamList {
316                name: "k_{k}_limits".into(),
317                len_hint: None,
318            }),
319            ["k"]
320        );
321        assert_eq!(
322            names(Source::WorkloadParamList {
323                name: "k_values".into(),
324                len_hint: None,
325            }),
326            ["k_values"]
327        );
328        assert_eq!(
329            names(Source::Generator {
330                expr: " all( row ) ".into(),
331                cardinality_hint: None,
332            }),
333            ["__cursor_extent_row_end", "__cursor_extent_row_start"]
334        );
335        assert_eq!(
336            names(Source::Generator {
337                expr: "pow2({n})".into(),
338                cardinality_hint: None,
339            }),
340            ["n"]
341        );
342        assert_eq!(all_cursor_argument("all(1)"), None);
343        assert_eq!(all_cursor_argument("all(a, b)"), None);
344        for extent in cursor_extent_names("row") {
345            assert_eq!(cursor_of_extent_name(&extent), Some("row"));
346        }
347        assert_eq!(cursor_of_extent_name("row"), None);
348    }
349
350    #[test]
351    fn literal_cardinality_is_list_length() {
352        let s = Source::Literal {
353            values: vec![
354                LiteralValue::Int(1),
355                LiteralValue::Int(2),
356                LiteralValue::Int(3),
357            ],
358        };
359        assert!(matches!(s.cardinality(), CardinalityClass::Bounded(3)));
360    }
361
362    /// An integer has one literal form, `Int` up to `i64::MAX` and
363    /// `UInt` above, and its serialized form reads back as that form.
364    #[test]
365    fn an_unsigned_literal_keeps_its_value_through_serde() {
366        assert_eq!(
367            LiteralValue::unsigned(i64::MAX as u64),
368            LiteralValue::Int(i64::MAX)
369        );
370        assert_eq!(
371            LiteralValue::unsigned(u64::MAX),
372            LiteralValue::UInt(u64::MAX)
373        );
374        for v in [
375            LiteralValue::Int(-3),
376            LiteralValue::Int(i64::MAX),
377            LiteralValue::UInt(1 << 63),
378            LiteralValue::UInt(u64::MAX),
379        ] {
380            let json = serde_json::to_string(&v).unwrap();
381            let back: LiteralValue = serde_json::from_str(&json).unwrap();
382            assert_eq!(back, v, "{json}");
383        }
384    }
385
386    #[test]
387    fn int_range_step_1() {
388        let s = Source::IntRange {
389            lo: 1,
390            hi: 10,
391            step: 1,
392        };
393        assert!(matches!(s.cardinality(), CardinalityClass::Bounded(9)));
394    }
395
396    #[test]
397    fn int_range_with_step() {
398        let s = Source::IntRange {
399            lo: 0,
400            hi: 10,
401            step: 2,
402        };
403        // 0,2,4,6,8 = 5 values
404        assert!(matches!(s.cardinality(), CardinalityClass::Bounded(5)));
405    }
406
407    #[test]
408    fn int_range_empty() {
409        let s = Source::IntRange {
410            lo: 5,
411            hi: 5,
412            step: 1,
413        };
414        assert!(matches!(s.cardinality(), CardinalityClass::Bounded(0)));
415    }
416
417    #[test]
418    fn generator_without_hint_is_unbounded() {
419        let s = Source::Generator {
420            expr: "live_query()".into(),
421            cardinality_hint: None,
422        };
423        assert!(matches!(s.cardinality(), CardinalityClass::Unbounded));
424    }
425
426    #[test]
427    fn generator_with_hint_is_bounded() {
428        let s = Source::Generator {
429            expr: "first_100()".into(),
430            cardinality_hint: Some(100),
431        };
432        assert!(matches!(s.cardinality(), CardinalityClass::Bounded(100)));
433    }
434
435    #[test]
436    fn continuous_interval_produces_continuous_class() {
437        let s = Source::ContinuousInterval {
438            interval: Interval::closed(0.0, 1.0),
439            measure: ProductMeasure::Uniform,
440        };
441        match s.cardinality() {
442            CardinalityClass::Continuous { intervals, measure } => {
443                assert_eq!(intervals.len(), 1);
444                assert!(matches!(measure, ProductMeasure::Uniform));
445            }
446            other => panic!("expected Continuous, got {other:?}"),
447        }
448        assert!(s.is_continuous());
449        assert!(!s.is_discrete());
450    }
451
452    #[test]
453    fn distribution_source_classification() {
454        let s = Source::Distribution {
455            distribution: MeasureName::Normal,
456            support: Interval {
457                lo: f64::NEG_INFINITY,
458                hi: f64::INFINITY,
459                lo_open: true,
460                hi_open: true,
461            },
462            params: vec![0.0, 1.0],
463        };
464        assert!(s.is_continuous());
465        match s.cardinality() {
466            CardinalityClass::Continuous {
467                measure: ProductMeasure::Named(MeasureName::Normal),
468                ..
469            } => {}
470            other => panic!("expected Continuous with Named(Normal), got {other:?}"),
471        }
472    }
473}