use serde::{Deserialize, Serialize};
use super::cardinality::{CardinalityClass, Interval, MeasureName, ProductMeasure};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum Source {
Literal {
values: Vec<LiteralValue>,
},
IntRange {
lo: i64,
hi: i64,
step: i64,
},
Generator {
expr: String,
cardinality_hint: Option<u64>,
},
WorkloadParamList {
name: String,
len_hint: Option<u64>,
},
ContinuousInterval {
interval: Interval,
measure: ProductMeasure,
},
Distribution {
distribution: MeasureName,
support: Interval,
params: Vec<f64>,
},
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum LiteralValue {
Int(i64),
Float(f64),
String(String),
Bool(bool),
Json(serde_json::Value),
}
impl Source {
pub fn cardinality(&self) -> CardinalityClass {
match self {
Source::Literal { values } => CardinalityClass::Bounded(values.len() as u64),
Source::IntRange { lo, hi, step } => {
let step = (*step).max(1).unsigned_abs();
if hi <= lo {
CardinalityClass::Bounded(0)
} else {
let span = (hi - lo) as u64;
let n = span.div_ceil(step);
CardinalityClass::Bounded(n)
}
}
Source::Generator {
cardinality_hint, ..
} => match cardinality_hint {
Some(n) => CardinalityClass::Bounded(*n),
None => CardinalityClass::Unbounded,
},
Source::WorkloadParamList { len_hint, .. } => match len_hint {
Some(n) => CardinalityClass::Bounded(*n),
None => CardinalityClass::Unbounded,
},
Source::ContinuousInterval { interval, measure } => CardinalityClass::Continuous {
intervals: vec![interval.clone()],
measure: measure.clone(),
},
Source::Distribution { support, .. } => CardinalityClass::Continuous {
intervals: vec![support.clone()],
measure: ProductMeasure::Named(*self.distribution_name()),
},
}
}
pub fn is_continuous(&self) -> bool {
matches!(
self,
Source::ContinuousInterval { .. } | Source::Distribution { .. }
)
}
pub fn is_discrete(&self) -> bool {
!self.is_continuous()
}
fn distribution_name(&self) -> &MeasureName {
match self {
Source::Distribution { distribution, .. } => distribution,
_ => panic!("distribution_name called on non-Distribution source"),
}
}
}
pub fn split_string_comprehension(s: &str) -> Vec<&str> {
s.split(|c: char| c == ',' || c == ';' || c.is_ascii_whitespace())
.map(str::trim)
.filter(|t| !t.is_empty())
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn literal_cardinality_is_list_length() {
let s = Source::Literal {
values: vec![
LiteralValue::Int(1),
LiteralValue::Int(2),
LiteralValue::Int(3),
],
};
assert!(matches!(s.cardinality(), CardinalityClass::Bounded(3)));
}
#[test]
fn int_range_step_1() {
let s = Source::IntRange {
lo: 1,
hi: 10,
step: 1,
};
assert!(matches!(s.cardinality(), CardinalityClass::Bounded(9)));
}
#[test]
fn int_range_with_step() {
let s = Source::IntRange {
lo: 0,
hi: 10,
step: 2,
};
assert!(matches!(s.cardinality(), CardinalityClass::Bounded(5)));
}
#[test]
fn int_range_empty() {
let s = Source::IntRange {
lo: 5,
hi: 5,
step: 1,
};
assert!(matches!(s.cardinality(), CardinalityClass::Bounded(0)));
}
#[test]
fn generator_without_hint_is_unbounded() {
let s = Source::Generator {
expr: "live_query()".into(),
cardinality_hint: None,
};
assert!(matches!(s.cardinality(), CardinalityClass::Unbounded));
}
#[test]
fn generator_with_hint_is_bounded() {
let s = Source::Generator {
expr: "first_100()".into(),
cardinality_hint: Some(100),
};
assert!(matches!(s.cardinality(), CardinalityClass::Bounded(100)));
}
#[test]
fn continuous_interval_produces_continuous_class() {
let s = Source::ContinuousInterval {
interval: Interval::closed(0.0, 1.0),
measure: ProductMeasure::Uniform,
};
match s.cardinality() {
CardinalityClass::Continuous { intervals, measure } => {
assert_eq!(intervals.len(), 1);
assert!(matches!(measure, ProductMeasure::Uniform));
}
other => panic!("expected Continuous, got {other:?}"),
}
assert!(s.is_continuous());
assert!(!s.is_discrete());
}
#[test]
fn distribution_source_classification() {
let s = Source::Distribution {
distribution: MeasureName::Normal,
support: Interval {
lo: f64::NEG_INFINITY,
hi: f64::INFINITY,
lo_open: true,
hi_open: true,
},
params: vec![0.0, 1.0],
};
assert!(s.is_continuous());
match s.cardinality() {
CardinalityClass::Continuous {
measure: ProductMeasure::Named(MeasureName::Normal),
..
} => {}
other => panic!("expected Continuous with Named(Normal), got {other:?}"),
}
}
}