use std::sync::Arc;
use crate::iteration::comprehension::cardinality::ProductMeasure;
use crate::iteration::comprehension::metadata::IndexFn;
use crate::iteration::comprehension::source::{LiteralValue, Source};
use crate::kernel::PolydatKernel;
use crate::ast::Value;
#[derive(Debug, Clone)]
pub struct EvaluatedSource {
pub values: Vec<Value>,
pub cardinality: u64,
pub index_fn: IndexFn,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum EvalClass {
Static,
ContextRequired,
Distribution,
}
#[derive(Debug, Clone)]
pub enum EvalError {
NeedsContext,
EvalFailed {
var: String,
source: String,
message: String,
},
}
impl std::fmt::Display for EvalError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
EvalError::NeedsContext => f.write_str("source evaluation needs a kernel context"),
EvalError::EvalFailed { var, source, message } => {
write!(f, "source '{var} in {source}': {message}")
}
}
}
}
impl std::error::Error for EvalError {}
pub struct EvalContext<'a> {
pub var_name: &'a str,
pub parent: &'a Arc<PolydatKernel>,
pub canonical: &'a Arc<PolydatKernel>,
pub prefix: &'a [(String, Value)],
}
pub trait SourceEval {
fn eval_class(&self) -> EvalClass;
fn evaluate(&self, ctx: Option<&EvalContext<'_>>) -> Result<EvaluatedSource, EvalError>;
}
impl SourceEval for Source {
fn eval_class(&self) -> EvalClass {
match self {
Source::Literal { .. } | Source::IntRange { .. } => EvalClass::Static,
Source::ContinuousInterval { .. } | Source::Distribution { .. } => {
EvalClass::Distribution
}
Source::Generator { .. } => EvalClass::ContextRequired,
Source::WorkloadParamList { .. } => EvalClass::ContextRequired,
}
}
fn evaluate(&self, ctx: Option<&EvalContext<'_>>) -> Result<EvaluatedSource, EvalError> {
match self {
Source::Literal { values } => {
let vals: Vec<Value> = values.iter().map(literal_to_value).collect();
let n = vals.len() as u64;
Ok(EvaluatedSource {
values: vals,
cardinality: n,
index_fn: IndexFn::Lattice { axis_sizes: vec![n] },
})
}
Source::IntRange { lo, hi, step } => {
let step = (*step).max(1);
let mut vals = Vec::new();
let mut cur = *lo;
while cur < *hi {
vals.push(Value::U64(cur as u64));
cur += step;
}
let n = vals.len() as u64;
Ok(EvaluatedSource {
values: vals,
cardinality: n,
index_fn: IndexFn::Lattice { axis_sizes: vec![n] },
})
}
Source::Generator { .. } | Source::WorkloadParamList { .. } => {
let ctx = ctx.ok_or(EvalError::NeedsContext)?;
let spec_text = match self {
Source::Generator { expr, .. } => expr.clone(),
Source::WorkloadParamList { name, .. } => format!("{{{name}}}"),
_ => unreachable!(),
};
let kernel = ctx
.parent
.materialize_subscope(ctx.canonical.program().clone(), ctx.prefix);
let vals = crate::iteration::comprehension::eval::evaluate_spec(&spec_text, &kernel)
.map_err(|e| EvalError::EvalFailed {
var: ctx.var_name.to_string(),
source: spec_text,
message: e.to_string(),
})?;
let n = vals.len() as u64;
let index_fn = classify_observed_values(&vals);
Ok(EvaluatedSource {
values: vals,
cardinality: n,
index_fn,
})
}
Source::ContinuousInterval { interval, measure } => Ok(EvaluatedSource {
values: Vec::new(),
cardinality: 0,
index_fn: IndexFn::Continuous {
intervals: vec![interval.clone()],
measure: measure.clone(),
},
}),
Source::Distribution { support, .. } => Ok(EvaluatedSource {
values: Vec::new(),
cardinality: 0,
index_fn: IndexFn::Continuous {
intervals: vec![support.clone()],
measure: ProductMeasure::Uniform,
},
}),
}
}
}
fn classify_observed_values(vals: &[Value]) -> IndexFn {
let n = vals.len() as u64;
IndexFn::Lattice { axis_sizes: vec![n] }
}
fn literal_to_value(lv: &LiteralValue) -> Value {
match lv {
LiteralValue::Int(n) => Value::U64(*n as u64),
LiteralValue::Float(f) => Value::F64(*f),
LiteralValue::String(s) => Value::Str(Arc::from(s.as_str())),
LiteralValue::Bool(b) => Value::Bool(*b),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::iteration::comprehension::cardinality::{Interval, MeasureName, ProductMeasure};
use crate::iteration::comprehension::source::LiteralValue;
#[test]
fn literal_evaluates_without_context() {
let s = Source::Literal {
values: vec![LiteralValue::Int(1), LiteralValue::Int(2), LiteralValue::Int(3)],
};
assert_eq!(s.eval_class(), EvalClass::Static);
let ev = s.evaluate(None).unwrap();
assert_eq!(ev.cardinality, 3);
assert_eq!(ev.values.len(), 3);
assert!(matches!(ev.index_fn, IndexFn::Lattice { axis_sizes: ref a } if a == &vec![3]));
}
#[test]
fn int_range_evaluates_without_context() {
let s = Source::IntRange { lo: 0, hi: 10, step: 2 };
assert_eq!(s.eval_class(), EvalClass::Static);
let ev = s.evaluate(None).unwrap();
assert_eq!(ev.cardinality, 5);
assert!(matches!(ev.index_fn, IndexFn::Lattice { axis_sizes: ref a } if a == &vec![5]));
}
#[test]
fn generator_without_context_errors() {
let s = Source::Generator {
expr: "range(0, 10)".into(),
cardinality_hint: Some(10),
};
assert_eq!(s.eval_class(), EvalClass::ContextRequired);
match s.evaluate(None) {
Err(EvalError::NeedsContext) => {}
other => panic!("expected NeedsContext, got {other:?}"),
}
}
#[test]
fn workload_param_list_without_context_errors() {
let s = Source::WorkloadParamList {
name: "k_values".into(),
len_hint: Some(5),
};
assert_eq!(s.eval_class(), EvalClass::ContextRequired);
assert!(matches!(s.evaluate(None), Err(EvalError::NeedsContext)));
}
#[test]
fn continuous_interval_yields_continuous_index_fn() {
let s = Source::ContinuousInterval {
interval: Interval::closed(0.0, 1.0),
measure: ProductMeasure::Uniform,
};
assert_eq!(s.eval_class(), EvalClass::Distribution);
let ev = s.evaluate(None).unwrap();
assert_eq!(ev.cardinality, 0);
assert!(ev.values.is_empty());
match ev.index_fn {
IndexFn::Continuous { intervals, .. } => assert_eq!(intervals.len(), 1),
other => panic!("expected Continuous, got {other:?}"),
}
}
#[test]
fn distribution_yields_continuous_index_fn() {
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_eq!(s.eval_class(), EvalClass::Distribution);
let ev = s.evaluate(None).unwrap();
assert_eq!(ev.cardinality, 0);
assert!(matches!(ev.index_fn, IndexFn::Continuous { .. }));
}
#[test]
fn generator_with_context_evaluates_to_lattice() {
let parent = Arc::new(crate::dsl::compile_polydat("\n").unwrap());
let canonical = Arc::new(crate::dsl::compile_polydat("\n").unwrap());
let s = Source::Generator {
expr: "1, 2, 3, 4, 5".into(),
cardinality_hint: Some(5),
};
let ctx = EvalContext {
var_name: "k",
parent: &parent,
canonical: &canonical,
prefix: &[],
};
let ev = s.evaluate(Some(&ctx)).unwrap();
assert_eq!(ev.cardinality, 5);
assert!(matches!(ev.index_fn, IndexFn::Lattice { .. }));
}
}