Skip to main content

polydat_grammar/comprehension/
cardinality.rs

1// Copyright 2024-2026 Jonathan Shook
2// SPDX-License-Identifier: Apache-2.0
3
4//! Cardinality classes — spec §6.1.
5//!
6//! Six classes describe every comprehension's dispense count.
7//! Three are discrete (`Bounded`, `BoundedAtMost`, `Unbounded`);
8//! two are continuous-domain (`Continuous`, `ContinuousAtMost`);
9//! one is hybrid (`Hybrid`). The class propagates through every
10//! constructor per spec §6.1's table.
11
12use serde::{Deserialize, Serialize};
13
14/// Cardinality of a comprehension's dispense stream.
15///
16/// Six variants per spec §6.1:
17///
18/// - **Discrete classes** enumerate distinct tuples; the count
19///   may be known exactly (`Bounded`), bounded above
20///   (`BoundedAtMost`), or unknown (`Unbounded`).
21/// - **Continuous classes** describe a measure-theoretic value
22///   space; they cannot enumerate and must be sampled via an
23///   enclosing `order(_, strategy, Some(n))` per V8.
24/// - **Hybrid** is a cartesian whose children mix discrete and
25///   continuous axes.
26#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
27pub enum CardinalityClass {
28    /// Discrete, exactly `n` tuples.
29    Bounded(u64),
30
31    /// Discrete, between 0 and `n` tuples (post-filter).
32    BoundedAtMost(u64),
33
34    /// Discrete, no known upper bound (generator, live stream).
35    Unbounded,
36
37    /// Continuous source — bounded or unbounded real intervals
38    /// with an integrable product measure. Sampled rather than
39    /// enumerated; V8 requires an enclosing
40    /// `order(_, strategy, Some(n))` before reaching a
41    /// `PolyStreamer`.
42    Continuous {
43        /// The interval of each axis.
44        intervals: Vec<Interval>,
45        /// The measure sampled.
46        measure: ProductMeasure,
47    },
48
49    /// Filtered continuous source. Measure reduced by the
50    /// predicate; still requires sampling.
51    ContinuousAtMost {
52        /// The interval of each axis.
53        intervals: Vec<Interval>,
54        /// The measure before the predicate reduces it.
55        measure_at_most: ProductMeasure,
56    },
57
58    /// Mixed discrete × continuous cartesian. The discrete part
59    /// is enumerable; the continuous part needs sampling. V8
60    /// applies to the continuous component.
61    Hybrid(Hybrid),
62}
63
64/// Mixed discrete × continuous cartesian shape.
65///
66/// Each `discrete_axes` entry is the axis size in tuples; each
67/// `continuous_axes` entry is the interval the axis spans.
68/// `measure` covers the continuous part.
69#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
70pub struct Hybrid {
71    /// Per-axis cardinality for the discrete axes, in
72    /// declaration order.
73    pub discrete_axes: Vec<u64>,
74    /// Per-axis intervals for the continuous axes, in
75    /// declaration order.
76    pub continuous_axes: Vec<Interval>,
77    /// Product measure over the continuous axes.
78    pub measure: ProductMeasure,
79}
80
81/// Real interval `[lo, hi]` (or open variants) for continuous
82/// sources. Unbounded sides use `f64::NEG_INFINITY` /
83/// `f64::INFINITY`; the V8 integrability check determines
84/// whether such intervals are valid given the measure.
85#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
86pub struct Interval {
87    /// The lower end.
88    pub lo: f64,
89    /// The upper end.
90    pub hi: f64,
91    /// Whether the lower end is excluded.
92    pub lo_open: bool,
93    /// Whether the upper end is excluded.
94    pub hi_open: bool,
95}
96
97impl Interval {
98    /// Closed-closed interval `[lo, hi]`.
99    pub fn closed(lo: f64, hi: f64) -> Self {
100        Self {
101            lo,
102            hi,
103            lo_open: false,
104            hi_open: false,
105        }
106    }
107
108    /// Half-open `[lo, hi)`.
109    pub fn half_open(lo: f64, hi: f64) -> Self {
110        Self {
111            lo,
112            hi,
113            lo_open: false,
114            hi_open: true,
115        }
116    }
117
118    /// Open interval `(lo, hi)`.
119    pub fn open(lo: f64, hi: f64) -> Self {
120        Self {
121            lo,
122            hi,
123            lo_open: true,
124            hi_open: true,
125        }
126    }
127
128    /// `true` if the interval has finite Lebesgue measure
129    /// (both endpoints finite). Used by V8's integrability
130    /// check together with the measure variant.
131    pub fn is_bounded(&self) -> bool {
132        self.lo.is_finite() && self.hi.is_finite()
133    }
134}
135
136/// Product measure over one or more continuous axes.
137///
138/// `Uniform` is the Lebesgue measure scaled by interval width
139/// (requires bounded intervals; V8 rejects unbounded + Uniform).
140/// `Named(D)` is a probability distribution with proper density
141/// over its declared support — Normal, Exponential, Pareto,
142/// Beta, etc. `Product(_)` carries a per-axis product of measures
143/// for K-D continuous cartesians.
144#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
145pub enum ProductMeasure {
146    /// Lebesgue measure scaled to the interval; needs bounded intervals.
147    Uniform,
148    /// A named probability distribution over its support.
149    Named(MeasureName),
150    /// One measure per axis.
151    Product(Vec<ProductMeasure>),
152}
153
154impl ProductMeasure {
155    /// `true` if this measure has finite total mass given the
156    /// supplied intervals. Used by V8.
157    ///
158    /// - `Uniform` is integrable iff every interval is bounded.
159    /// - `Named(D)` is integrable per its distribution: proper
160    ///   probability distributions are always integrable
161    ///   (they have unit total mass by definition).
162    /// - `Product(children)` is integrable iff every child is.
163    pub fn is_integrable(&self, intervals: &[Interval]) -> bool {
164        match self {
165            ProductMeasure::Uniform => intervals.iter().all(Interval::is_bounded),
166            ProductMeasure::Named(name) => name.is_proper_probability_measure(),
167            ProductMeasure::Product(children) => {
168                if children.len() != intervals.len() {
169                    return false;
170                }
171                children
172                    .iter()
173                    .zip(intervals.iter())
174                    .all(|(m, i)| m.is_integrable(std::slice::from_ref(i)))
175            }
176        }
177    }
178}
179
180/// Named continuous distribution. Closed enum per spec
181/// §10.7.5's "User-defined extensions" non-goal — new
182/// distributions land as coordinated additions.
183#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
184pub enum MeasureName {
185    /// The normal distribution.
186    Normal,
187    /// The exponential distribution.
188    Exponential,
189    /// The Pareto distribution.
190    Pareto,
191    /// The beta distribution.
192    Beta,
193    /// The log-normal distribution.
194    LogNormal,
195    /// The gamma distribution.
196    Gamma,
197    /// The uniform distribution on `[0, 1]`.
198    Uniform01,
199}
200
201impl MeasureName {
202    /// All currently-named distributions are proper probability
203    /// measures (unit total mass). V8 accepts them on any
204    /// interval that matches the distribution's support.
205    pub fn is_proper_probability_measure(self) -> bool {
206        true
207    }
208
209    /// The distribution's parameters, in the order a
210    /// `Source::Distribution`'s `params` lists them:
211    ///
212    /// | Measure | Parameters | Support |
213    /// |---|---|---|
214    /// | `Normal` | mean, stddev | (-∞, ∞) |
215    /// | `Exponential` | rate | [0, ∞) |
216    /// | `Pareto` | scale, shape | [scale, ∞) |
217    /// | `Beta` | alpha, beta | [0, 1] |
218    /// | `LogNormal` | mean, stddev (of the log) | (0, ∞) |
219    /// | `Gamma` | shape, scale | [0, ∞) |
220    /// | `Uniform01` | none | [0, 1] |
221    pub fn parameter_names(self) -> &'static [&'static str] {
222        match self {
223            MeasureName::Normal | MeasureName::LogNormal => &["mean", "stddev"],
224            MeasureName::Exponential => &["rate"],
225            MeasureName::Pareto => &["scale", "shape"],
226            MeasureName::Beta => &["alpha", "beta"],
227            MeasureName::Gamma => &["shape", "scale"],
228            MeasureName::Uniform01 => &[],
229        }
230    }
231
232    /// The standard parameters, used when a source names the
233    /// measure without parameters: `Normal(0, 1)`, `Exponential(1)`,
234    /// `Pareto(1, 1)`, `Beta(1, 1)`, `LogNormal(0, 1)`, `Gamma(1, 1)`.
235    pub fn default_params(self) -> &'static [f64] {
236        match self {
237            MeasureName::Normal | MeasureName::LogNormal => &[0.0, 1.0],
238            MeasureName::Exponential => &[1.0],
239            MeasureName::Pareto | MeasureName::Beta | MeasureName::Gamma => &[1.0, 1.0],
240            MeasureName::Uniform01 => &[],
241        }
242    }
243
244    /// The parameters as the sampler takes them: `params` when it
245    /// has the measure's count, the standard parameters when it is
246    /// empty, and an error naming the expected parameters otherwise.
247    pub fn resolve_params(self, params: &[f64]) -> Result<Vec<f64>, String> {
248        let names = self.parameter_names();
249        if params.is_empty() {
250            return Ok(self.default_params().to_vec());
251        }
252        if params.len() != names.len() {
253            return Err(format!(
254                "{self:?} takes {} parameter(s) ({}); found {}",
255                names.len(),
256                names.join(", "),
257                params.len()
258            ));
259        }
260        if let Some(bad) = params.iter().find(|p| !p.is_finite()) {
261            return Err(format!("{self:?}: parameter {bad} is not finite"));
262        }
263        Ok(params.to_vec())
264    }
265
266    /// The measure's name in comprehension source text, the spelling
267    /// a clause writes: `normal(0, 1)`, `log_normal(0, 0.5)`.
268    pub fn text(self) -> &'static str {
269        match self {
270            MeasureName::Normal => "normal",
271            MeasureName::Exponential => "exponential",
272            MeasureName::Pareto => "pareto",
273            MeasureName::Beta => "beta",
274            MeasureName::LogNormal => "log_normal",
275            MeasureName::Gamma => "gamma",
276            MeasureName::Uniform01 => "uniform01",
277        }
278    }
279
280    /// The measure a source-text name denotes, or `None` when the
281    /// name is not one of the closed set (§10.7.5): the caller reads
282    /// the text as something else, a generator call for one.
283    pub fn from_text(name: &str) -> Option<Self> {
284        [
285            MeasureName::Normal,
286            MeasureName::Exponential,
287            MeasureName::Pareto,
288            MeasureName::Beta,
289            MeasureName::LogNormal,
290            MeasureName::Gamma,
291            MeasureName::Uniform01,
292        ]
293        .into_iter()
294        .find(|m| m.text() == name)
295    }
296
297    /// The measure's own support under `params`, the interval a
298    /// source that names no narrower one draws from. A Pareto's
299    /// support starts at its scale; every other measure's is fixed.
300    pub fn support(self, params: &[f64]) -> Interval {
301        let p = self
302            .resolve_params(params)
303            .unwrap_or_else(|_| self.default_params().to_vec());
304        let inf = f64::INFINITY;
305        match self {
306            MeasureName::Normal => Interval::open(-inf, inf),
307            MeasureName::Exponential | MeasureName::Gamma => Interval {
308                lo: 0.0,
309                hi: inf,
310                lo_open: false,
311                hi_open: true,
312            },
313            MeasureName::Pareto => Interval {
314                lo: p[0],
315                hi: inf,
316                lo_open: false,
317                hi_open: true,
318            },
319            MeasureName::LogNormal => Interval {
320                lo: 0.0,
321                hi: inf,
322                lo_open: true,
323                hi_open: true,
324            },
325            MeasureName::Beta | MeasureName::Uniform01 => Interval::closed(0.0, 1.0),
326        }
327    }
328}
329
330#[cfg(test)]
331mod tests {
332    use super::*;
333
334    #[test]
335    fn bounded_interval_is_bounded() {
336        assert!(Interval::closed(0.0, 1.0).is_bounded());
337        assert!(Interval::open(-1.0, 1.0).is_bounded());
338    }
339
340    #[test]
341    fn unbounded_interval_is_not_bounded() {
342        let i = Interval {
343            lo: 0.0,
344            hi: f64::INFINITY,
345            lo_open: false,
346            hi_open: true,
347        };
348        assert!(!i.is_bounded());
349    }
350
351    #[test]
352    fn uniform_integrable_on_bounded_interval() {
353        let m = ProductMeasure::Uniform;
354        assert!(m.is_integrable(&[Interval::closed(0.0, 1.0)]));
355    }
356
357    #[test]
358    fn uniform_not_integrable_on_unbounded_interval() {
359        let m = ProductMeasure::Uniform;
360        let unbounded = Interval {
361            lo: 0.0,
362            hi: f64::INFINITY,
363            lo_open: false,
364            hi_open: true,
365        };
366        assert!(!m.is_integrable(&[unbounded]));
367    }
368
369    #[test]
370    fn named_measure_always_integrable() {
371        let m = ProductMeasure::Named(MeasureName::Normal);
372        let unbounded = Interval {
373            lo: f64::NEG_INFINITY,
374            hi: f64::INFINITY,
375            lo_open: true,
376            hi_open: true,
377        };
378        assert!(m.is_integrable(&[unbounded]));
379    }
380
381    #[test]
382    fn product_measure_requires_matching_arity() {
383        let m = ProductMeasure::Product(vec![ProductMeasure::Uniform, ProductMeasure::Uniform]);
384        assert!(m.is_integrable(&[Interval::closed(0.0, 1.0), Interval::closed(0.0, 1.0)]));
385        assert!(!m.is_integrable(&[Interval::closed(0.0, 1.0)]));
386    }
387
388    #[test]
389    fn cardinality_class_round_trip_serde() {
390        let c = CardinalityClass::Continuous {
391            intervals: vec![Interval::closed(0.0, 1.0)],
392            measure: ProductMeasure::Uniform,
393        };
394        let json = serde_json::to_string(&c).unwrap();
395        let back: CardinalityClass = serde_json::from_str(&json).unwrap();
396        assert_eq!(c, back);
397    }
398}