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}