use std::sync::Arc;
use crate::ast::Value;
use crate::iteration::comprehension::cardinality::ProductMeasure;
use crate::iteration::comprehension::eval::{SpecError, spec_error};
use crate::iteration::comprehension::metadata::IndexFn;
use crate::iteration::comprehension::source::{LiteralValue, Source};
use crate::kernel::interp::{Layered, Lookup};
#[derive(Debug, Clone)]
pub struct EvaluatedSource {
pub values: Vec<Value>,
pub cardinality: u64,
pub index_fn: IndexFn,
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum NoneRead {
Unbound(String),
BoundNone(String),
}
impl NoneRead {
pub fn name(&self) -> &str {
match self {
NoneRead::Unbound(name) | NoneRead::BoundNone(name) => name,
}
}
}
impl std::fmt::Display for NoneRead {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
NoneRead::Unbound(name) => write!(f, "`{name}` is not bound"),
NoneRead::BoundNone(name) => write!(f, "`{name}` is None"),
}
}
}
#[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 scope: &'a dyn Lookup,
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 { .. } if self.names_read().is_empty() => EvalClass::Static,
Source::Generator { .. } => EvalClass::ContextRequired,
Source::WorkloadParamList { .. } => EvalClass::ContextRequired,
}
}
fn evaluate(&self, ctx: Option<&EvalContext<'_>>) -> Result<EvaluatedSource, EvalError> {
evaluate_reading(self, ctx).map(|(evaluated, _)| evaluated)
}
}
pub(crate) fn evaluate_reading(
source: &Source,
ctx: Option<&EvalContext<'_>>,
) -> Result<(EvaluatedSource, Vec<NoneRead>), EvalError> {
let evaluated = match source {
Source::Generator { .. } | Source::WorkloadParamList { .. } => {
return evaluate_spec_source(source, ctx);
}
other => evaluate_static(other),
};
Ok((evaluated, Vec::new()))
}
fn evaluate_spec_source(
source: &Source,
ctx: Option<&EvalContext<'_>>,
) -> Result<(EvaluatedSource, Vec<NoneRead>), EvalError> {
let spec_text = match source {
Source::Generator { expr, .. } => expr.clone(),
Source::WorkloadParamList { name, .. } => format!("{{{name}}}"),
_ => unreachable!("only spec-text sources"),
};
let empty = crate::kernel::interp::NoScope::new();
let (var_name, scope): (&str, Layered<'_>) = match ctx {
Some(ctx) => (
ctx.var_name,
Layered {
prefix: ctx.prefix,
inner: ctx.scope,
},
),
None if source.eval_class() == EvalClass::Static => (
"<context-free>",
Layered {
prefix: &[],
inner: &empty,
},
),
None => return Err(EvalError::NeedsContext),
};
match crate::iteration::comprehension::eval::evaluate_spec_internal(&spec_text, &scope) {
Ok(vals) => {
let n = vals.len() as u64;
let index_fn = classify_observed_values(&vals);
Ok((
EvaluatedSource {
values: vals,
cardinality: n,
index_fn,
},
Vec::new(),
))
}
Err(SpecError::ReadsNone { reads, .. }) => Ok((
EvaluatedSource {
values: Vec::new(),
cardinality: 0,
index_fn: IndexFn::Lattice {
axis_sizes: vec![0],
},
},
reads,
)),
Err(SpecError::Failed(message)) => Err(EvalError::EvalFailed {
var: var_name.to_string(),
message: spec_error(&spec_text, message).to_string(),
source: spec_text,
}),
}
}
fn evaluate_static(source: &Source) -> EvaluatedSource {
match source {
Source::Literal { values } => {
let vals: Vec<Value> = values.iter().map(literal_to_value).collect();
let n = vals.len() as u64;
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;
EvaluatedSource {
values: vals,
cardinality: n,
index_fn: IndexFn::Lattice {
axis_sizes: vec![n],
},
}
}
Source::ContinuousInterval { interval, measure } => EvaluatedSource {
values: Vec::new(),
cardinality: 0,
index_fn: IndexFn::Continuous {
intervals: vec![interval.clone()],
measure: measure.clone(),
},
},
Source::Distribution {
distribution,
support,
..
} => EvaluatedSource {
values: Vec::new(),
cardinality: 0,
index_fn: IndexFn::Continuous {
intervals: vec![support.clone()],
measure: ProductMeasure::Named(*distribution),
},
},
Source::Generator { .. } | Source::WorkloadParamList { .. } => {
unreachable!("a spec-text source reads names")
}
}
}
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::UInt(n) => Value::U64(*n),
LiteralValue::Float(f) => Value::F64(*f),
LiteralValue::String(s) => Value::Str(Arc::from(s.as_str())),
LiteralValue::Bool(b) => Value::Bool(*b),
LiteralValue::Json(j) => Value::Json(Arc::new(j.clone())),
}
}
#[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 a_context_free_generator_evaluates_without_context() {
let s = Source::Generator {
expr: "fib(6)".into(),
cardinality_hint: None,
};
assert_eq!(s.eval_class(), EvalClass::Static);
let ev = s.evaluate(None).unwrap();
assert_eq!(ev.cardinality, 6);
}
#[test]
fn generator_without_context_errors() {
let s = Source::Generator {
expr: "range(0, {n})".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 {
measure: ProductMeasure::Named(MeasureName::Normal),
..
}
));
}
#[test]
fn generator_with_context_evaluates_to_lattice() {
let canonical = Arc::new(crate::dsl::compile_polydat_interpreter("\n").unwrap());
let s = Source::Generator {
expr: "1, 2, 3, 4, 5".into(),
cardinality_hint: Some(5),
};
let ctx = EvalContext {
var_name: "k",
scope: &*canonical,
prefix: &[],
};
let ev = s.evaluate(Some(&ctx)).unwrap();
assert_eq!(ev.cardinality, 5);
assert!(matches!(ev.index_fn, IndexFn::Lattice { .. }));
}
}