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),
}
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 iteration_interior(v: &crate::ast::Value) -> Option<Vec<crate::ast::Value>> {
use crate::ast::Value;
match v {
Value::VecF32(s) => Some(s.as_slice().iter().map(|x| Value::F64(*x as f64)).collect()),
Value::VecF64(s) => Some(s.as_slice().iter().map(|x| Value::F64(*x)).collect()),
Value::VecF16(s) => Some(s.as_slice().iter().map(|x| Value::F64(x.to_f64())).collect()),
Value::VecI32(s) => Some(s.as_slice().iter().map(|x| Value::I64(*x as i64)).collect()),
Value::VecI64(s) => Some(s.as_slice().iter().map(|x| Value::I64(*x)).collect()),
Value::VecI16(s) => Some(s.as_slice().iter().map(|x| Value::I64(*x as i64)).collect()),
Value::VecI8(s) => Some(s.as_slice().iter().map(|x| Value::I64(*x as i64)).collect()),
Value::Json(j) => j.as_array().map(|arr| {
arr.iter().map(|e| Value::Json(std::sync::Arc::new(e.clone()))).collect()
}),
Value::Ext(_) => v.as_partition_list().map(|list| {
list.as_slice().iter().map(|p| Value::from_partition(*p)).collect()
}),
Value::Reg128(b, view) => {
use crate::ast::RegLanes;
match view {
RegLanes::Raw => None,
RegLanes::I8x16 => Some(b.lanes_i8().iter().map(|x| Value::I64(*x as i64)).collect()),
RegLanes::I16x8 => Some(b.lanes_i16().iter().map(|x| Value::I64(*x as i64)).collect()),
RegLanes::I32x4 => Some(b.lanes_i32().iter().map(|x| Value::I64(*x as i64)).collect()),
RegLanes::I64x2 => Some(b.lanes_i64().iter().map(|x| Value::I64(*x)).collect()),
RegLanes::F16x8 => Some(b.lanes_f16().iter().map(|x| Value::F64(x.to_f64())).collect()),
RegLanes::F32x4 => Some(b.lanes_f32().iter().map(|x| Value::F64(*x as f64)).collect()),
RegLanes::F64x2 => Some(b.lanes_f64().iter().map(|x| Value::F64(*x)).collect()),
}
}
Value::Str(s) => Some(strip_string_tokens(s)),
Value::U64(_) | Value::I64(_) | Value::U128(_) | Value::I128(_)
| Value::F64(_) | Value::Bool(_)
| Value::Bytes(_) | Value::Handle(_) | Value::None => None,
}
}
pub fn strip_string_tokens(s: &str) -> Vec<crate::ast::Value> {
use crate::ast::Value;
split_string_comprehension(s)
.into_iter()
.map(|t| {
if let Ok(n) = t.parse::<u64>() {
Value::U64(n)
} else if let Ok(f) = t.parse::<f64>() {
Value::F64(f)
} else if t == "true" {
Value::Bool(true)
} else if t == "false" {
Value::Bool(false)
} else {
Value::Str(t.to_string().into())
}
})
.collect()
}
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::*;
use crate::ast::Value;
#[test]
fn string_strips_on_comma_semicolon_whitespace_retaining_colons() {
let got = strip_string_tokens("a:1, b:2; c:3 d:4");
assert_eq!(got, vec![
Value::Str("a:1".into()), Value::Str("b:2".into()),
Value::Str("c:3".into()), Value::Str("d:4".into()),
]);
}
#[test]
fn string_tokens_are_typed_like_literals() {
assert_eq!(strip_string_tokens("1, 2, 3"),
vec![Value::U64(1), Value::U64(2), Value::U64(3)]);
assert_eq!(strip_string_tokens("1.5, 2.5"),
vec![Value::F64(1.5), Value::F64(2.5)]);
}
#[test]
fn single_token_string_degenerates_to_singleton() {
assert_eq!(strip_string_tokens("OTHER"), vec![Value::Str("OTHER".into())]);
}
#[test]
fn iteration_interior_string_is_its_tokens() {
let v = Value::Str("x, y, z".into());
assert_eq!(iteration_interior(&v),
Some(vec![Value::Str("x".into()), Value::Str("y".into()), Value::Str("z".into())]));
}
#[test]
fn iteration_interior_vector_peels_to_elements() {
let v = Value::VecI32(crate::ast::SliceArc::from_vec(vec![10, 20, 30]));
assert_eq!(iteration_interior(&v),
Some(vec![Value::I64(10), Value::I64(20), Value::I64(30)]));
}
#[test]
fn iteration_interior_signed_vector_peels_signed() {
let v = Value::VecI64(crate::ast::SliceArc::from_vec(vec![-5_i64, 7]));
assert_eq!(iteration_interior(&v),
Some(vec![Value::I64(-5), Value::I64(7)]));
}
#[test]
fn iteration_interior_scalars_are_none() {
assert_eq!(iteration_interior(&Value::U64(5)), None);
assert_eq!(iteration_interior(&Value::F64(1.5)), None);
assert_eq!(iteration_interior(&Value::Bool(true)), 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 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:?}"),
}
}
}