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),
UInt(u64),
Float(f64),
String(String),
Bool(bool),
Json(serde_json::Value),
}
impl LiteralValue {
pub fn unsigned(n: u64) -> Self {
match i64::try_from(n) {
Ok(n) => LiteralValue::Int(n),
Err(_) => LiteralValue::UInt(n),
}
}
}
impl Source {
pub fn referenced_names(&self) -> std::collections::BTreeSet<String> {
let mut out = std::collections::BTreeSet::new();
match self {
Source::WorkloadParamList { name, .. } => {
out.insert(name.clone());
}
Source::Generator { expr, .. } => {
out.extend(crate::refs::referenced_names(expr));
crate::refs::collect_string_interpolation_refs(expr, &mut out);
}
Source::Literal { .. }
| Source::IntRange { .. }
| Source::ContinuousInterval { .. }
| Source::Distribution { .. } => {}
}
out
}
pub fn names_read(&self) -> std::collections::BTreeSet<String> {
let mut out = std::collections::BTreeSet::new();
match self {
Source::WorkloadParamList { name, .. } => {
crate::refs::collect_string_interpolation_refs(&format!("{{{name}}}"), &mut out);
}
Source::Generator { expr, .. } => match all_cursor_argument(expr) {
Some(cursor) => out.extend(cursor_extent_names(cursor)),
None => out.extend(self.referenced_names()),
},
Source::Literal { .. }
| Source::IntRange { .. }
| Source::ContinuousInterval { .. }
| Source::Distribution { .. } => {}
}
out
}
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()
}
pub fn all_cursor_argument(text: &str) -> Option<&str> {
let cursor = text.trim().strip_prefix("all(")?.strip_suffix(')')?.trim();
let mut chars = cursor.chars();
let starts = chars
.next()
.is_some_and(|c| c.is_ascii_alphabetic() || c == '_');
(starts && chars.all(|c| c.is_ascii_alphanumeric() || c == '_')).then_some(cursor)
}
pub fn cursor_extent_names(cursor: &str) -> [String; 2] {
[
format!("__cursor_extent_{cursor}_start"),
format!("__cursor_extent_{cursor}_end"),
]
}
pub fn cursor_of_extent_name(name: &str) -> Option<&str> {
let rest = name.strip_prefix("__cursor_extent_")?;
rest.strip_suffix("_start")
.or_else(|| rest.strip_suffix("_end"))
.filter(|cursor| !cursor.is_empty())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn names_read_are_the_leaves_of_a_composition_and_a_cursor_extent() {
let names = |s: Source| s.names_read().into_iter().collect::<Vec<_>>();
assert_eq!(
names(Source::WorkloadParamList {
name: "k_{k}_limits".into(),
len_hint: None,
}),
["k"]
);
assert_eq!(
names(Source::WorkloadParamList {
name: "k_values".into(),
len_hint: None,
}),
["k_values"]
);
assert_eq!(
names(Source::Generator {
expr: " all( row ) ".into(),
cardinality_hint: None,
}),
["__cursor_extent_row_end", "__cursor_extent_row_start"]
);
assert_eq!(
names(Source::Generator {
expr: "pow2({n})".into(),
cardinality_hint: None,
}),
["n"]
);
assert_eq!(all_cursor_argument("all(1)"), None);
assert_eq!(all_cursor_argument("all(a, b)"), None);
for extent in cursor_extent_names("row") {
assert_eq!(cursor_of_extent_name(&extent), Some("row"));
}
assert_eq!(cursor_of_extent_name("row"), None);
}
#[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 an_unsigned_literal_keeps_its_value_through_serde() {
assert_eq!(
LiteralValue::unsigned(i64::MAX as u64),
LiteralValue::Int(i64::MAX)
);
assert_eq!(
LiteralValue::unsigned(u64::MAX),
LiteralValue::UInt(u64::MAX)
);
for v in [
LiteralValue::Int(-3),
LiteralValue::Int(i64::MAX),
LiteralValue::UInt(1 << 63),
LiteralValue::UInt(u64::MAX),
] {
let json = serde_json::to_string(&v).unwrap();
let back: LiteralValue = serde_json::from_str(&json).unwrap();
assert_eq!(back, v, "{json}");
}
}
#[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:?}"),
}
}
}