use serde::{Deserialize, Serialize};
use super::ast::Comprehension;
use super::cardinality::{CardinalityClass, Hybrid, Interval, ProductMeasure};
use super::source::Source;
use super::strategy::{StrategyName, ZipMode};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Metadata {
pub cardinality: CardinalityClass,
pub index_addressable: Option<IndexFn>,
pub natural_order: NaturalOrder,
pub materialization: Materialization,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum IndexFn {
Lattice {
axis_sizes: Vec<u64>,
},
Lockstep {
length: u64,
},
Modular {
axis_sizes: Vec<u64>,
},
Concatenation {
segment_sizes: Vec<u64>,
},
Continuous {
intervals: Vec<Interval>,
measure: ProductMeasure,
},
Hybrid {
discrete_axes: Vec<u64>,
continuous_axes: Vec<Interval>,
measure: ProductMeasure,
},
}
impl IndexFn {
pub fn has_continuous_axis(&self) -> bool {
matches!(self, IndexFn::Continuous { .. } | IndexFn::Hybrid { .. })
}
pub fn is_multi_axis_lattice(&self) -> bool {
matches!(self, IndexFn::Lattice { axis_sizes } if axis_sizes.len() >= 2)
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum NaturalOrder {
Lex,
Lockstep,
Sequential,
Strategy(StrategyName),
PendingSampling,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum Materialization {
Streaming,
BoundedBarrier {
working_set_size: u64,
},
UnboundedBarrier,
}
impl Comprehension {
pub fn metadata(&self) -> Metadata {
match self {
Comprehension::Clause { source, .. } => clause_metadata(source),
Comprehension::Cartesian { children } => cartesian_metadata(children),
Comprehension::Zip { children, mode } => zip_metadata(children, *mode),
Comprehension::Union { children } => union_metadata(children),
Comprehension::Filter { child, .. } => filter_metadata(child),
Comprehension::Order {
child,
strategy,
truncation,
..
} => order_metadata(child, *strategy, *truncation),
}
}
}
fn clause_metadata(source: &Source) -> Metadata {
let cardinality = source.cardinality();
let (index_addressable, natural_order) = match &cardinality {
CardinalityClass::Bounded(n) => (
Some(IndexFn::Lattice {
axis_sizes: vec![*n],
}),
NaturalOrder::Lex,
),
CardinalityClass::Continuous { intervals, measure } => (
Some(IndexFn::Continuous {
intervals: intervals.clone(),
measure: measure.clone(),
}),
NaturalOrder::PendingSampling,
),
_ => (None, NaturalOrder::Lex),
};
Metadata {
cardinality,
index_addressable,
natural_order,
materialization: Materialization::Streaming,
}
}
fn cartesian_metadata(children: &[Comprehension]) -> Metadata {
let dependent = detect_dependent_sources(children);
let child_meta: Vec<Metadata> = children.iter().map(|c| c.metadata()).collect();
let cardinality = combine_cartesian_cardinality(&child_meta);
let index_addressable = if dependent {
None
} else {
combine_cartesian_index_fn(&child_meta)
};
let natural_order = if matches!(
cardinality,
CardinalityClass::Continuous { .. } | CardinalityClass::Hybrid(_)
) {
NaturalOrder::PendingSampling
} else {
NaturalOrder::Lex
};
Metadata {
cardinality,
index_addressable,
natural_order,
materialization: Materialization::Streaming,
}
}
fn zip_metadata(children: &[Comprehension], mode: ZipMode) -> Metadata {
let child_meta: Vec<Metadata> = children.iter().map(|c| c.metadata()).collect();
let cardinality = combine_zip_cardinality(&child_meta, mode);
let index_addressable = combine_zip_index_fn(&child_meta, mode);
let materialization = match mode {
ZipMode::Strict | ZipMode::Truncate => Materialization::Streaming,
ZipMode::Cycle => cycle_materialization(&cycle_operands(&child_meta)),
};
Metadata {
cardinality,
index_addressable,
natural_order: NaturalOrder::Lockstep,
materialization,
}
}
pub fn cycle_length(counts: &[u64]) -> u64 {
if counts.contains(&0) {
0
} else {
counts.iter().copied().max().unwrap_or(0)
}
}
fn known_empty(m: &Metadata) -> bool {
matches!(
m.cardinality,
CardinalityClass::Bounded(0) | CardinalityClass::BoundedAtMost(0)
)
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum CycleOperand {
Indexed,
Streamed,
Buffered {
bound: Option<u64>,
},
}
pub fn cycle_operands(children: &[Metadata]) -> Vec<CycleOperand> {
let indexed = |m: &Metadata| {
!known_empty(m)
&& m.index_addressable
.as_ref()
.is_some_and(|idx| !idx.has_continuous_axis())
};
let bound = |m: &Metadata| match &m.cardinality {
CardinalityClass::Bounded(n) | CardinalityClass::BoundedAtMost(n) => Some(*n),
_ => None,
};
let rest: Vec<usize> = (0..children.len())
.filter(|&i| !indexed(&children[i]) && !known_empty(&children[i]))
.collect();
let streamed = rest
.iter()
.copied()
.find(|&i| bound(&children[i]).is_none())
.or_else(|| {
rest.iter()
.copied()
.rev()
.max_by_key(|&i| bound(&children[i]))
});
children
.iter()
.enumerate()
.map(|(i, m)| {
if indexed(m) {
CycleOperand::Indexed
} else if Some(i) == streamed {
CycleOperand::Streamed
} else {
CycleOperand::Buffered { bound: bound(m) }
}
})
.collect()
}
pub fn cycle_plan_is_empty(plan: &[CycleOperand]) -> bool {
plan.contains(&CycleOperand::Buffered { bound: Some(0) })
}
pub fn cycle_materialization(plan: &[CycleOperand]) -> Materialization {
if cycle_plan_is_empty(plan) {
return Materialization::Streaming;
}
let mut total: u64 = 0;
let mut buffered = false;
for operand in plan {
if let CycleOperand::Buffered { bound } = operand {
buffered = true;
match bound {
Some(n) => total = total.saturating_add(*n),
None => return Materialization::UnboundedBarrier,
}
}
}
if buffered {
Materialization::BoundedBarrier {
working_set_size: total,
}
} else {
Materialization::Streaming
}
}
fn union_metadata(children: &[Comprehension]) -> Metadata {
let child_meta: Vec<Metadata> = children.iter().map(|c| c.metadata()).collect();
let cardinality = combine_union_cardinality(&child_meta);
let index_addressable = combine_union_index_fn(&child_meta);
Metadata {
cardinality,
index_addressable,
natural_order: NaturalOrder::Sequential,
materialization: Materialization::Streaming,
}
}
fn filter_metadata(child: &Comprehension) -> Metadata {
let child_meta = child.metadata();
let cardinality = match &child_meta.cardinality {
CardinalityClass::Bounded(0) | CardinalityClass::BoundedAtMost(0) => {
CardinalityClass::Bounded(0)
}
CardinalityClass::Bounded(n) | CardinalityClass::BoundedAtMost(n) => {
CardinalityClass::BoundedAtMost(*n)
}
CardinalityClass::Unbounded => CardinalityClass::Unbounded,
CardinalityClass::Continuous { intervals, measure }
| CardinalityClass::ContinuousAtMost {
intervals,
measure_at_most: measure,
} => CardinalityClass::ContinuousAtMost {
intervals: intervals.clone(),
measure_at_most: measure.clone(),
},
CardinalityClass::Hybrid(h) => CardinalityClass::Hybrid(h.clone()),
};
Metadata {
cardinality,
index_addressable: None, natural_order: child_meta.natural_order,
materialization: child_meta.materialization,
}
}
fn order_metadata(
child: &Comprehension,
strategy: StrategyName,
truncation: Option<u64>,
) -> Metadata {
let child_meta = child.metadata();
let cardinality = order_cardinality(child, &child_meta.cardinality, strategy, truncation);
let selected = || match &cardinality {
CardinalityClass::Bounded(n) | CardinalityClass::BoundedAtMost(n) => {
Some(IndexFn::Lattice {
axis_sizes: vec![*n],
})
}
_ => None,
};
let (index_addressable, natural_order, materialization) = match strategy {
StrategyName::Lex => (
match truncation {
None => child_meta.index_addressable,
Some(_) => child_meta.index_addressable.and_then(|_| selected()),
},
NaturalOrder::Lex,
child_meta.materialization, ),
non_lex => {
let materialization = match &child_meta.index_addressable {
Some(_) => Materialization::BoundedBarrier {
working_set_size: strategy_working_set(
non_lex,
&child_meta.index_addressable,
truncation,
),
},
None => match &child_meta.cardinality {
CardinalityClass::Bounded(n) | CardinalityClass::BoundedAtMost(n) => {
Materialization::BoundedBarrier {
working_set_size: *n,
}
}
_ => Materialization::UnboundedBarrier,
},
};
(selected(), NaturalOrder::Strategy(non_lex), materialization)
}
};
Metadata {
cardinality,
index_addressable,
natural_order,
materialization,
}
}
fn combine_cartesian_cardinality(children: &[Metadata]) -> CardinalityClass {
let mut has_continuous = false;
let mut has_discrete = false;
let mut counts: Vec<Count> = Vec::new();
let mut discrete_axes: Vec<u64> = Vec::new();
let mut continuous_intervals: Vec<Interval> = Vec::new();
let mut continuous_measures: Vec<ProductMeasure> = Vec::new();
for m in children {
match &m.cardinality {
CardinalityClass::Bounded(n) | CardinalityClass::BoundedAtMost(n) => {
has_discrete = true;
discrete_axes.push(*n); counts.extend(Count::of(&m.cardinality));
}
CardinalityClass::Unbounded => {
has_discrete = true;
discrete_axes.push(0);
counts.push(Count::Unknown);
}
CardinalityClass::Continuous { intervals, measure }
| CardinalityClass::ContinuousAtMost {
intervals,
measure_at_most: measure,
} => {
has_continuous = true;
continuous_intervals.extend(intervals.iter().cloned());
continuous_measures.push(measure.clone());
}
CardinalityClass::Hybrid(h) => {
has_continuous = true;
has_discrete = true;
discrete_axes.extend(h.discrete_axes.iter().copied());
continuous_intervals.extend(h.continuous_axes.iter().cloned());
continuous_measures.push(h.measure.clone());
}
}
}
if has_continuous && has_discrete {
CardinalityClass::Hybrid(Hybrid {
discrete_axes,
continuous_axes: continuous_intervals,
measure: simplify_measures(continuous_measures),
})
} else if has_continuous {
CardinalityClass::Continuous {
intervals: continuous_intervals,
measure: simplify_measures(continuous_measures),
}
} else {
cartesian_count(&counts).class()
}
}
fn combine_cartesian_index_fn(children: &[Metadata]) -> Option<IndexFn> {
let all_addressable = children.iter().all(|m| m.index_addressable.is_some());
if !all_addressable {
return None;
}
let mut all_discrete = true;
let mut all_continuous = true;
let mut discrete_axes: Vec<u64> = Vec::new();
let mut continuous_intervals: Vec<Interval> = Vec::new();
let mut continuous_measures: Vec<ProductMeasure> = Vec::new();
for m in children {
match m.index_addressable.as_ref().unwrap() {
IndexFn::Lattice { axis_sizes } => {
all_continuous = false;
discrete_axes.extend(axis_sizes.iter().copied());
}
IndexFn::Continuous { intervals, measure } => {
all_discrete = false;
continuous_intervals.extend(intervals.iter().cloned());
continuous_measures.push(measure.clone());
}
IndexFn::Hybrid {
discrete_axes: d,
continuous_axes: c,
measure,
} => {
all_discrete = false;
all_continuous = false;
discrete_axes.extend(d.iter().copied());
continuous_intervals.extend(c.iter().cloned());
continuous_measures.push(measure.clone());
}
IndexFn::Lockstep { .. } | IndexFn::Modular { .. } | IndexFn::Concatenation { .. } => {
return None;
}
}
}
if all_discrete {
Some(IndexFn::Lattice {
axis_sizes: discrete_axes,
})
} else if all_continuous {
Some(IndexFn::Continuous {
intervals: continuous_intervals,
measure: simplify_measures(continuous_measures),
})
} else {
Some(IndexFn::Hybrid {
discrete_axes,
continuous_axes: continuous_intervals,
measure: simplify_measures(continuous_measures),
})
}
}
fn combine_zip_cardinality(children: &[Metadata], mode: ZipMode) -> CardinalityClass {
let counts: Vec<Count> = children
.iter()
.map(|m| Count::of(&m.cardinality).unwrap_or(Count::Unknown))
.collect();
match mode {
ZipMode::Strict => strict_zip_count(&counts),
ZipMode::Truncate => truncate_zip_count(&counts),
ZipMode::Cycle => cycle_zip_count(&counts),
}
.class()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Count {
Exact(u64),
AtMost(u64),
Unknown,
}
impl Count {
fn of(class: &CardinalityClass) -> Option<Self> {
Some(match class {
CardinalityClass::Bounded(n) | CardinalityClass::BoundedAtMost(n @ 0) => {
Count::Exact(*n)
}
CardinalityClass::BoundedAtMost(n) => Count::AtMost(*n),
CardinalityClass::Unbounded => Count::Unknown,
_ => return None,
})
}
fn bound(self) -> Option<u64> {
match self {
Count::Exact(n) | Count::AtMost(n) => Some(n),
Count::Unknown => None,
}
}
fn combined(n: u64, counts: &[Count]) -> Self {
if counts.iter().all(|c| matches!(c, Count::Exact(_))) {
Count::Exact(n)
} else {
Count::AtMost(n)
}
}
fn class(self) -> CardinalityClass {
match self {
Count::Exact(n) | Count::AtMost(n @ 0) => CardinalityClass::Bounded(n),
Count::AtMost(n) => CardinalityClass::BoundedAtMost(n),
Count::Unknown => CardinalityClass::Unbounded,
}
}
}
fn cartesian_count(counts: &[Count]) -> Count {
if counts.contains(&Count::Exact(0)) {
return Count::Exact(0);
}
let Some(bounds) = counts.iter().map(|c| c.bound()).collect::<Option<Vec<_>>>() else {
return Count::Unknown;
};
Count::combined(bounds.into_iter().fold(1, u64::saturating_mul), counts)
}
fn truncate_zip_count(counts: &[Count]) -> Count {
if counts.contains(&Count::Exact(0)) {
return Count::Exact(0);
}
match counts.iter().filter_map(|c| c.bound()).min() {
Some(m) => Count::combined(m, counts),
None => Count::Unknown,
}
}
fn strict_zip_count(counts: &[Count]) -> Count {
if let Some(exact) = counts.iter().find(|c| matches!(c, Count::Exact(_))) {
return *exact;
}
match counts.iter().filter_map(|c| c.bound()).min() {
Some(m) => Count::AtMost(m),
None => Count::Unknown,
}
}
fn cycle_zip_count(counts: &[Count]) -> Count {
if counts.contains(&Count::Exact(0)) {
return Count::Exact(0);
}
let Some(bounds) = counts.iter().map(|c| c.bound()).collect::<Option<Vec<_>>>() else {
return Count::Unknown;
};
Count::combined(bounds.into_iter().max().unwrap_or(0), counts)
}
fn union_count(counts: &[Count]) -> Count {
let Some(bounds) = counts.iter().map(|c| c.bound()).collect::<Option<Vec<_>>>() else {
return Count::Unknown;
};
Count::combined(bounds.into_iter().fold(0, u64::saturating_add), counts)
}
fn order_cardinality(
child: &Comprehension,
child_class: &CardinalityClass,
strategy: StrategyName,
truncation: Option<u64>,
) -> CardinalityClass {
let strata = matches!(strategy, StrategyName::Extrema | StrategyName::Shells);
let Some(count) = Count::of(child_class) else {
let Some(n) = truncation.filter(|_| !matches!(strategy, StrategyName::Lex)) else {
return child_class.clone();
};
let mut space = SampledSpace::default();
space.collect(child);
if space.discrete.contains(&Count::Exact(0)) {
return CardinalityClass::Bounded(0);
}
return if matches!(strategy, StrategyName::Extrema) {
let mut axes = space.discrete;
axes.extend(std::iter::repeat_n(Count::Exact(2), space.continuous));
match cartesian_count(&axes) {
Count::Exact(m) | Count::AtMost(m) => Count::AtMost(m),
Count::Unknown => Count::Unknown,
}
} else if !space.filtered && space.discrete.iter().all(|c| matches!(c, Count::Exact(_))) {
Count::Exact(n)
} else {
Count::AtMost(n)
}
.class();
};
match (count, truncation) {
(Count::Exact(0), _) | (_, None) => count,
(Count::Exact(c) | Count::AtMost(c), Some(_)) if strata => Count::AtMost(c),
(Count::Exact(c), Some(n)) => Count::Exact(c.min(n)),
(Count::AtMost(c), Some(n)) => Count::AtMost(c.min(n)),
(Count::Unknown, Some(_)) if strata => Count::Unknown,
(Count::Unknown, Some(n)) => Count::AtMost(n),
}
.class()
}
#[derive(Default)]
struct SampledSpace {
discrete: Vec<Count>,
continuous: usize,
filtered: bool,
}
impl SampledSpace {
fn collect(&mut self, c: &Comprehension) {
match c {
Comprehension::Clause { source, .. } => match source.cardinality() {
CardinalityClass::Continuous { .. } => self.continuous += 1,
class => self
.discrete
.push(Count::of(&class).unwrap_or(Count::Unknown)),
},
Comprehension::Cartesian { children } => {
children.iter().for_each(|child| self.collect(child));
}
Comprehension::Filter { child, .. } => {
self.filtered = true;
self.collect(child);
}
other => self
.discrete
.push(Count::of(&other.metadata().cardinality).unwrap_or(Count::Unknown)),
}
}
}
fn combine_zip_index_fn(children: &[Metadata], mode: ZipMode) -> Option<IndexFn> {
let all_addressable = children.iter().all(|m| m.index_addressable.is_some());
if !all_addressable {
return None;
}
let counts: Vec<u64> = children
.iter()
.filter_map(|m| match &m.cardinality {
CardinalityClass::Bounded(n) | CardinalityClass::BoundedAtMost(n) => Some(*n),
_ => None,
})
.collect();
if counts.len() != children.len() {
return None;
}
match mode {
ZipMode::Strict | ZipMode::Truncate => {
let length = match mode {
ZipMode::Strict => counts[0],
ZipMode::Truncate => *counts.iter().min().unwrap(),
ZipMode::Cycle => unreachable!(),
};
Some(IndexFn::Lockstep { length })
}
ZipMode::Cycle => Some(IndexFn::Modular { axis_sizes: counts }),
}
}
fn combine_union_cardinality(children: &[Metadata]) -> CardinalityClass {
let counts: Vec<Count> = children
.iter()
.map(|m| Count::of(&m.cardinality).unwrap_or(Count::Unknown))
.collect();
union_count(&counts).class()
}
fn combine_union_index_fn(children: &[Metadata]) -> Option<IndexFn> {
let all_addressable = children.iter().all(|m| m.index_addressable.is_some());
if !all_addressable {
return None;
}
let segment_sizes: Vec<u64> = children
.iter()
.filter_map(|m| match &m.cardinality {
CardinalityClass::Bounded(n) | CardinalityClass::BoundedAtMost(n) => Some(*n),
_ => None,
})
.collect();
if segment_sizes.len() != children.len() {
return None;
}
Some(IndexFn::Concatenation { segment_sizes })
}
fn simplify_measures(measures: Vec<ProductMeasure>) -> ProductMeasure {
match measures.len() {
0 => ProductMeasure::Uniform,
1 => measures.into_iter().next().unwrap(),
_ => ProductMeasure::Product(measures),
}
}
fn strategy_working_set(
strategy: StrategyName,
input: &Option<IndexFn>,
truncation: Option<u64>,
) -> u64 {
match (strategy, input, truncation) {
(StrategyName::Halton, Some(_), Some(n))
| (StrategyName::Sobol, Some(_), Some(n))
| (StrategyName::Shuffle, Some(_), Some(n))
| (StrategyName::ReverseLex, Some(_), Some(n)) => n,
(StrategyName::Lhs, Some(idx), Some(n)) => {
let dim = lattice_dim(idx).max(1);
n.saturating_mul(dim as u64)
}
(StrategyName::Extrema, Some(idx), Some(_))
| (StrategyName::Shells, Some(idx), Some(_)) => index_fn_cardinality(idx),
(StrategyName::Diagonal, Some(_), Some(n))
| (StrategyName::Antidiagonal, Some(_), Some(n)) => n,
(_, Some(idx), None) => index_fn_cardinality(idx),
(_, None, Some(n)) => n,
(_, None, None) => 0,
(StrategyName::Lex, Some(_), Some(n)) => n,
}
}
fn lattice_dim(idx: &IndexFn) -> usize {
match idx {
IndexFn::Lattice { axis_sizes } => axis_sizes.len(),
IndexFn::Continuous { intervals, .. } => intervals.len(),
IndexFn::Hybrid {
discrete_axes,
continuous_axes,
..
} => discrete_axes.len() + continuous_axes.len(),
IndexFn::Lockstep { .. } | IndexFn::Modular { .. } | IndexFn::Concatenation { .. } => 1,
}
}
fn index_fn_cardinality(idx: &IndexFn) -> u64 {
match idx {
IndexFn::Lattice { axis_sizes } => axis_sizes
.iter()
.copied()
.fold(1u64, |a, b| a.saturating_mul(b)),
IndexFn::Lockstep { length } => *length,
IndexFn::Modular { axis_sizes } => cycle_length(axis_sizes),
IndexFn::Concatenation { segment_sizes } => segment_sizes
.iter()
.copied()
.fold(0u64, |a, b| a.saturating_add(b)),
IndexFn::Continuous { .. } | IndexFn::Hybrid { .. } => 0,
}
}
fn detect_dependent_sources(children: &[Comprehension]) -> bool {
let mut prior_names: Vec<String> = Vec::new();
for child in children {
for name in collect_source_name_references(child) {
if prior_names.contains(&name) {
return true;
}
}
for n in child.coordinate_names() {
if !prior_names.contains(&n) {
prior_names.push(n);
}
}
}
false
}
fn collect_source_name_references(c: &Comprehension) -> std::collections::BTreeSet<String> {
c.source_names_read()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::comprehension::source::{LiteralValue, Source};
fn clause(name: &str, vs: &[i64]) -> Comprehension {
Comprehension::clause(
name,
Source::Literal {
values: vs.iter().map(|n| LiteralValue::Int(*n)).collect(),
},
)
}
fn continuous_clause(name: &str) -> Comprehension {
Comprehension::clause(
name,
Source::ContinuousInterval {
interval: Interval::closed(0.0, 1.0),
measure: ProductMeasure::Uniform,
},
)
}
#[test]
fn clause_metadata_for_bounded_source() {
let m = clause("k", &[1, 2, 3]).metadata();
assert_eq!(m.cardinality, CardinalityClass::Bounded(3));
assert_eq!(
m.index_addressable,
Some(IndexFn::Lattice {
axis_sizes: vec![3]
})
);
assert_eq!(m.natural_order, NaturalOrder::Lex);
assert_eq!(m.materialization, Materialization::Streaming);
}
#[test]
fn clause_metadata_for_continuous_source() {
let m = continuous_clause("alpha").metadata();
assert!(matches!(m.cardinality, CardinalityClass::Continuous { .. }));
assert!(matches!(
m.index_addressable,
Some(IndexFn::Continuous { .. })
));
assert_eq!(m.natural_order, NaturalOrder::PendingSampling);
assert_eq!(m.materialization, Materialization::Streaming);
}
#[test]
fn cartesian_metadata_combines_lattice_axes() {
let c =
Comprehension::cartesian(vec![clause("k", &[1, 2]), clause("limit", &[10, 20, 30])]);
let m = c.metadata();
assert_eq!(m.cardinality, CardinalityClass::Bounded(6));
assert_eq!(
m.index_addressable,
Some(IndexFn::Lattice {
axis_sizes: vec![2, 3]
})
);
assert_eq!(m.natural_order, NaturalOrder::Lex);
}
#[test]
fn cartesian_metadata_for_hybrid() {
let c =
Comprehension::cartesian(vec![clause("k", &[1, 2, 3, 4]), continuous_clause("theta")]);
let m = c.metadata();
match m.cardinality {
CardinalityClass::Hybrid(h) => {
assert_eq!(h.discrete_axes, vec![4]);
assert_eq!(h.continuous_axes.len(), 1);
}
other => panic!("expected Hybrid, got {other:?}"),
}
assert!(matches!(m.index_addressable, Some(IndexFn::Hybrid { .. })));
assert_eq!(m.natural_order, NaturalOrder::PendingSampling);
}
#[test]
fn dependent_cartesian_produces_none_addressable() {
let dependent = Comprehension::cartesian(vec![
clause("k", &[1, 2, 3]),
Comprehension::clause(
"replicas",
Source::Generator {
expr: "range(0, 2 * {k})".into(),
cardinality_hint: Some(6),
},
),
]);
let m = dependent.metadata();
assert!(m.index_addressable.is_none());
}
#[test]
fn zip_strict_produces_lockstep_index_fn() {
let c = Comprehension::zip(
vec![clause("x", &[1, 2, 3]), clause("y", &[10, 20, 30])],
ZipMode::Strict,
);
let m = c.metadata();
assert_eq!(m.index_addressable, Some(IndexFn::Lockstep { length: 3 }));
assert_eq!(m.natural_order, NaturalOrder::Lockstep);
assert_eq!(m.materialization, Materialization::Streaming);
}
#[test]
fn zip_cycle_produces_modular_index_fn_and_barrier() {
let c = Comprehension::zip(
vec![clause("k", &[1, 2, 3, 4, 5]), clause("color", &[1, 2, 3])],
ZipMode::Cycle,
);
let m = c.metadata();
match m.index_addressable {
Some(IndexFn::Modular { axis_sizes }) => {
assert_eq!(axis_sizes, vec![5, 3]);
}
other => panic!("expected Modular, got {other:?}"),
}
assert_eq!(m.materialization, Materialization::Streaming);
}
fn unknown_count(name: &str) -> Comprehension {
Comprehension::clause(
name,
Source::Generator {
expr: "values({n})".into(),
cardinality_hint: None,
},
)
}
#[test]
fn zip_cycle_with_an_unknown_count_buffers_every_finite_unaddressable_operand() {
let c = Comprehension::zip(
vec![
unknown_count("tick"),
Comprehension::filter(clause("a", &(0..1000).collect::<Vec<_>>()), "{a} > 1"),
Comprehension::filter(clause("b", &[1, 2, 3]), "{b} > 0"),
clause("color", &[1, 2, 3]),
],
ZipMode::Cycle,
);
let plan = cycle_operands(
&match &c {
Comprehension::Zip { children, .. } => children,
_ => unreachable!(),
}
.iter()
.map(Comprehension::metadata)
.collect::<Vec<_>>(),
);
assert_eq!(
plan,
vec![
CycleOperand::Streamed,
CycleOperand::Buffered { bound: Some(1000) },
CycleOperand::Buffered { bound: Some(3) },
CycleOperand::Indexed,
]
);
assert_eq!(
c.metadata().materialization,
Materialization::BoundedBarrier {
working_set_size: 1003
}
);
}
#[test]
fn zip_cycle_streams_the_largest_unaddressable_operand() {
let c = Comprehension::zip(
vec![
Comprehension::filter(clause("a", &[1, 2]), "{a} > 0"),
Comprehension::filter(clause("b", &[1, 2, 3, 4]), "{b} > 0"),
clause("k", &(0..100).collect::<Vec<_>>()),
],
ZipMode::Cycle,
);
assert_eq!(
c.metadata().materialization,
Materialization::BoundedBarrier {
working_set_size: 2
}
);
}
#[test]
fn zip_cycle_with_an_operand_known_empty_is_empty() {
let operands = || {
vec![
unknown_count("tick"),
Comprehension::filter(clause("a", &[1, 2, 3]), "{a} > 1"),
clause("color", &[1, 2, 3]),
]
};
let empties = [
clause("e", &[]),
Comprehension::filter(clause("e", &[]), "{e} > 1"),
];
for empty in &empties {
for at in 0..=3 {
let mut children = operands();
children.insert(at, empty.clone());
let plan = cycle_operands(
&children
.iter()
.map(Comprehension::metadata)
.collect::<Vec<_>>(),
);
assert_eq!(plan[at], CycleOperand::Buffered { bound: Some(0) });
assert!(cycle_plan_is_empty(&plan));
let m = Comprehension::zip(children, ZipMode::Cycle).metadata();
assert_eq!(m.cardinality, CardinalityClass::Bounded(0));
assert_eq!(m.materialization, Materialization::Streaming);
}
}
let addressable = Comprehension::zip(
vec![clause("k", &[1, 2, 3]), clause("e", &[])],
ZipMode::Cycle,
);
let m = addressable.metadata();
assert_eq!(
m.index_addressable,
Some(IndexFn::Modular {
axis_sizes: vec![3, 0]
})
);
assert_eq!(
index_fn_cardinality(m.index_addressable.as_ref().unwrap()),
0
);
assert_eq!(cycle_length(&[3, 0]), 0);
assert_eq!(cycle_length(&[3, 5]), 5);
}
#[test]
fn zip_cycle_over_an_operand_at_most_counts_at_most() {
let c = Comprehension::zip(
vec![
clause("k", &[1, 2, 3, 4, 5]),
Comprehension::filter(clause("a", &[1, 2]), "{a} > 1"),
],
ZipMode::Cycle,
);
assert_eq!(c.metadata().cardinality, CardinalityClass::BoundedAtMost(5));
}
fn witnesses(c: Count) -> Vec<u64> {
match c {
Count::Exact(n) => vec![n],
Count::AtMost(n) => (0..=n).collect(),
Count::Unknown => (0..=7).chain([1000]).collect(),
}
}
fn assert_describes(claim: Count, yields: &[u64], what: &str) {
let Some(&max) = yields.iter().max() else {
return; };
match claim {
Count::Exact(n) => assert!(yields.iter().all(|&y| y == n), "{what}: {yields:?}"),
Count::AtMost(n) => {
assert!(max <= n, "{what}: {yields:?} exceed {n}");
assert_eq!(max, n, "{what}: the bound is not tight");
assert!(n > 0, "{what}: at most zero is exactly zero");
assert!(
yields.iter().any(|&y| y != n),
"{what}: always {n}, so exact"
);
}
Count::Unknown => assert!(max >= 1000, "{what}: bounded by {max}"),
}
}
const KINDS: [Count; 5] = [
Count::Exact(0),
Count::Exact(3),
Count::Exact(5),
Count::AtMost(4),
Count::Unknown,
];
fn combinations(counts: &[Count]) -> Vec<Vec<u64>> {
counts.iter().fold(vec![Vec::new()], |acc, c| {
acc.iter()
.flat_map(|prefix| {
witnesses(*c).into_iter().map(move |w| {
let mut next = prefix.clone();
next.push(w);
next
})
})
.collect()
})
}
#[test]
fn every_kind_combination_counts_what_the_operator_yields() {
let mut shapes: Vec<Vec<Count>> = Vec::new();
for a in KINDS {
for b in KINDS {
shapes.push(vec![a, b]);
for c in KINDS {
shapes.push(vec![a, b, c]);
}
}
}
for counts in &shapes {
let combos = combinations(counts);
let product: Vec<u64> = combos.iter().map(|c| c.iter().product()).collect();
assert_describes(
cartesian_count(counts),
&product,
&format!("cartesian {counts:?}"),
);
let sum: Vec<u64> = combos.iter().map(|c| c.iter().sum()).collect();
assert_describes(union_count(counts), &sum, &format!("union {counts:?}"));
let shortest: Vec<u64> = combos.iter().map(|c| *c.iter().min().unwrap()).collect();
assert_describes(
truncate_zip_count(counts),
&shortest,
&format!("truncate {counts:?}"),
);
let cycled: Vec<u64> = combos.iter().map(|c| cycle_length(c)).collect();
assert_describes(
cycle_zip_count(counts),
&cycled,
&format!("cycle {counts:?}"),
);
let strict: Vec<u64> = combos
.iter()
.filter(|c| c.iter().all(|&n| n == c[0]))
.map(|c| c[0])
.collect();
assert_describes(
strict_zip_count(counts),
&strict,
&format!("strict {counts:?}"),
);
}
}
fn filtered(c: Comprehension) -> Comprehension {
Comprehension::filter(c, "true")
}
#[test]
fn an_operand_at_most_makes_a_combination_at_most() {
let at_most = || filtered(clause("k", &[1, 2, 3, 4, 5, 6, 7, 8, 9]));
let colors = || clause("c", &[1, 2]);
let product = Comprehension::cartesian(vec![at_most(), colors()]);
assert_eq!(
product.metadata().cardinality,
CardinalityClass::BoundedAtMost(18)
);
let zip = Comprehension::zip(vec![at_most(), colors()], ZipMode::Truncate);
assert_eq!(
zip.metadata().cardinality,
CardinalityClass::BoundedAtMost(2)
);
let zip = Comprehension::zip(vec![unknown_count("u"), colors()], ZipMode::Truncate);
assert_eq!(
zip.metadata().cardinality,
CardinalityClass::BoundedAtMost(2)
);
let product = Comprehension::cartesian(vec![unknown_count("u"), clause("e", &[])]);
assert_eq!(product.metadata().cardinality, CardinalityClass::Bounded(0));
let empty = filtered(clause("e", &[]));
assert_eq!(empty.metadata().cardinality, CardinalityClass::Bounded(0));
}
#[test]
fn an_order_counts_by_its_strategy() {
let order = |c, s, t| Comprehension::order(c, s, t).metadata().cardinality;
let ks = || clause("k", &[1, 2, 3, 4, 5, 6]);
for s in [
StrategyName::Lex,
StrategyName::ReverseLex,
StrategyName::Diagonal,
StrategyName::Halton,
StrategyName::Sobol,
StrategyName::Lhs,
StrategyName::Shuffle,
] {
assert_eq!(
order(ks(), s, Some(4)),
CardinalityClass::Bounded(4),
"{s:?}"
);
assert_eq!(
order(ks(), s, Some(9)),
CardinalityClass::Bounded(6),
"{s:?}"
);
assert_eq!(order(ks(), s, None), CardinalityClass::Bounded(6), "{s:?}");
assert_eq!(
order(filtered(ks()), s, Some(4)),
CardinalityClass::BoundedAtMost(4),
"{s:?}"
);
assert_eq!(
order(unknown_count("u"), s, Some(4)),
CardinalityClass::BoundedAtMost(4),
"{s:?}"
);
}
for s in [StrategyName::Extrema, StrategyName::Shells] {
assert_eq!(
order(ks(), s, Some(1)),
CardinalityClass::BoundedAtMost(6),
"{s:?}"
);
assert_eq!(order(ks(), s, None), CardinalityClass::Bounded(6), "{s:?}");
assert_eq!(
order(clause("e", &[]), s, Some(1)),
CardinalityClass::Bounded(0)
);
}
let space = || Comprehension::cartesian(vec![clause("k", &[1, 2]), continuous_clause("u")]);
assert_eq!(
order(space(), StrategyName::Halton, Some(5)),
CardinalityClass::Bounded(5)
);
assert_eq!(
order(filtered(space()), StrategyName::Halton, Some(5)),
CardinalityClass::BoundedAtMost(5)
);
assert_eq!(
order(space(), StrategyName::Extrema, Some(1)),
CardinalityClass::BoundedAtMost(4)
);
let empty_axis = Comprehension::cartesian(vec![clause("k", &[]), continuous_clause("u")]);
assert_eq!(
order(empty_axis, StrategyName::Sobol, Some(5)),
CardinalityClass::Bounded(0)
);
}
#[test]
fn zip_cycle_with_two_unknown_counts_is_unbounded() {
let c = Comprehension::zip(vec![unknown_count("x"), unknown_count("y")], ZipMode::Cycle);
assert_eq!(
c.metadata().materialization,
Materialization::UnboundedBarrier
);
}
#[test]
fn union_produces_concatenation_index_fn() {
let a = Comprehension::cartesian(vec![clause("k", &[1, 2]), clause("limit", &[10])]);
let b = Comprehension::cartesian(vec![clause("k", &[3, 4]), clause("limit", &[20])]);
let u = Comprehension::union(vec![a, b]);
let m = u.metadata();
assert_eq!(m.cardinality, CardinalityClass::Bounded(4));
assert_eq!(
m.index_addressable,
Some(IndexFn::Concatenation {
segment_sizes: vec![2, 2]
})
);
assert_eq!(m.natural_order, NaturalOrder::Sequential);
}
#[test]
fn filter_destroys_addressability() {
let inner =
Comprehension::cartesian(vec![clause("k", &[1, 2]), clause("limit", &[10, 20])]);
let filtered = Comprehension::filter(inner, "{k} > 0");
let m = filtered.metadata();
assert_eq!(m.cardinality, CardinalityClass::BoundedAtMost(4));
assert_eq!(m.index_addressable, None);
}
#[test]
fn lex_order_inherits_addressability_untruncated() {
let inner =
Comprehension::cartesian(vec![clause("k", &[1, 2]), clause("limit", &[10, 20])]);
let whole = Comprehension::order(inner.clone(), StrategyName::Lex, None).metadata();
assert_eq!(whole.cardinality, CardinalityClass::Bounded(4));
assert_eq!(
whole.index_addressable,
Some(IndexFn::Lattice {
axis_sizes: vec![2, 2]
})
);
assert_eq!(whole.natural_order, NaturalOrder::Lex);
let prefix = Comprehension::order(inner.clone(), StrategyName::Lex, Some(3)).metadata();
assert_eq!(prefix.cardinality, CardinalityClass::Bounded(3));
assert_eq!(
prefix.index_addressable,
Some(IndexFn::Lattice {
axis_sizes: vec![3]
})
);
assert_eq!(prefix.natural_order, NaturalOrder::Lex);
let streamed = Comprehension::order(
Comprehension::filter(inner, "{k} > 1"),
StrategyName::Lex,
Some(3),
)
.metadata();
assert_eq!(streamed.index_addressable, None);
}
#[test]
fn non_lex_order_addresses_its_selection() {
let inner =
Comprehension::cartesian(vec![clause("k", &[1, 2]), clause("limit", &[10, 20])]);
let ordered = Comprehension::order(inner, StrategyName::Halton, Some(2));
let m = ordered.metadata();
assert_eq!(
m.index_addressable,
Some(IndexFn::Lattice {
axis_sizes: vec![2]
})
);
let reordered = Comprehension::order(ordered.clone(), StrategyName::Shuffle, None);
let r = reordered.metadata();
assert_eq!(r.cardinality, CardinalityClass::Bounded(2));
assert_eq!(
r.index_addressable,
Some(IndexFn::Lattice {
axis_sizes: vec![2]
})
);
assert_eq!(
r.materialization,
Materialization::BoundedBarrier {
working_set_size: 2
}
);
match m.natural_order {
NaturalOrder::Strategy(StrategyName::Halton) => {}
other => panic!("expected Strategy(Halton), got {other:?}"),
}
assert_eq!(
m.materialization,
Materialization::BoundedBarrier {
working_set_size: 2
}
);
}
#[test]
fn continuous_sampling_yields_bounded_cardinality() {
let inner =
Comprehension::cartesian(vec![continuous_clause("alpha"), continuous_clause("beta")]);
let ordered = Comprehension::order(inner, StrategyName::Halton, Some(100));
let m = ordered.metadata();
assert_eq!(m.cardinality, CardinalityClass::Bounded(100));
assert_eq!(
m.materialization,
Materialization::BoundedBarrier {
working_set_size: 100
}
);
}
#[test]
fn metadata_propagation_is_idempotent() {
let c = Comprehension::order(
Comprehension::filter(
Comprehension::cartesian(vec![clause("k", &[1, 2, 3]), clause("limit", &[10, 20])]),
"{k} * {limit} > 5",
),
StrategyName::Halton,
Some(5),
);
let m1 = c.metadata();
let m2 = c.metadata();
assert_eq!(m1, m2);
}
#[test]
fn has_continuous_axis_classifier() {
let lat = IndexFn::Lattice {
axis_sizes: vec![3, 4],
};
assert!(!lat.has_continuous_axis());
let cont = IndexFn::Continuous {
intervals: vec![Interval::closed(0.0, 1.0)],
measure: ProductMeasure::Uniform,
};
assert!(cont.has_continuous_axis());
}
#[test]
fn multi_axis_lattice_classifier() {
assert!(
IndexFn::Lattice {
axis_sizes: vec![3, 4]
}
.is_multi_axis_lattice()
);
assert!(
!IndexFn::Lattice {
axis_sizes: vec![3]
}
.is_multi_axis_lattice()
);
assert!(
!IndexFn::Continuous {
intervals: vec![Interval::closed(0.0, 1.0), Interval::closed(0.0, 1.0)],
measure: ProductMeasure::Uniform,
}
.is_multi_axis_lattice()
);
}
}