use std::collections::BTreeSet;
use crate::ast::Value;
use crate::iteration::comprehension::ast::Comprehension;
use crate::iteration::comprehension::eval_source::{EvalContext, SourceEval};
use crate::iteration::comprehension::source::{LiteralValue, Source};
use crate::kernel::interp::Lookup;
pub fn flatten_static_sources(ast: &Comprehension, scope: &dyn Lookup) -> Comprehension {
let bound: BTreeSet<String> = ast.coordinate_names().into_iter().collect();
flatten_node(ast, scope, &bound)
}
fn flatten_node(c: &Comprehension, scope: &dyn Lookup, bound: &BTreeSet<String>) -> Comprehension {
match c {
Comprehension::Clause { name, source } => Comprehension::Clause {
name: name.clone(),
source: flatten_source(name, source, scope, bound),
},
Comprehension::Cartesian { children } => Comprehension::Cartesian {
children: flatten_all(children, scope, bound),
},
Comprehension::Zip { children, mode } => Comprehension::Zip {
children: flatten_all(children, scope, bound),
mode: *mode,
},
Comprehension::Union { children } => Comprehension::Union {
children: flatten_all(children, scope, bound),
},
Comprehension::Filter { child, predicate } => Comprehension::Filter {
child: Box::new(flatten_node(child, scope, bound)),
predicate: predicate.clone(),
},
Comprehension::Order {
child,
strategy,
truncation,
seed,
} => Comprehension::Order {
child: Box::new(flatten_node(child, scope, bound)),
strategy: *strategy,
truncation: *truncation,
seed: *seed,
},
}
}
fn flatten_all(
children: &[Comprehension],
scope: &dyn Lookup,
bound: &BTreeSet<String>,
) -> Vec<Comprehension> {
children
.iter()
.map(|c| flatten_node(c, scope, bound))
.collect()
}
fn flatten_source(
name: &str,
source: &Source,
scope: &dyn Lookup,
bound: &BTreeSet<String>,
) -> Source {
let Source::Generator { expr, .. } = source else {
return source.clone();
};
if context_required(source, scope, bound) {
return source.clone();
}
let ctx = EvalContext {
var_name: name,
scope,
prefix: &[],
};
let Ok(evaluated) = source.evaluate(Some(&ctx)) else {
return source.clone();
};
let count = evaluated.cardinality;
if count == 0 {
return Source::Generator {
expr: expr.clone(),
cardinality_hint: Some(0),
};
}
let literals: Option<Vec<LiteralValue>> = evaluated.values.iter().map(literal_of).collect();
match literals {
Some(values) => Source::Literal { values },
None => Source::Generator {
expr: expr.clone(),
cardinality_hint: Some(count),
},
}
}
fn context_required(source: &Source, scope: &dyn Lookup, bound: &BTreeSet<String>) -> bool {
source
.referenced_names()
.iter()
.any(|n| bound.contains(n) || scope.lookup(n).is_none())
}
pub fn first_refused_generator(
ast: &Comprehension,
scope: &dyn Lookup,
) -> Option<(String, String)> {
let bound: BTreeSet<String> = ast.coordinate_names().into_iter().collect();
refused_in(ast, scope, &bound)
}
fn refused_in(
c: &Comprehension,
scope: &dyn Lookup,
bound: &BTreeSet<String>,
) -> Option<(String, String)> {
match c {
Comprehension::Clause { name, source } => match source {
Source::Generator { expr, .. } if !context_required(source, scope, bound) => {
crate::iteration::comprehension::eval::refused_generator_call(expr, scope)
.map(|message| (name.clone(), message))
}
_ => None,
},
Comprehension::Cartesian { children }
| Comprehension::Zip { children, .. }
| Comprehension::Union { children } => {
children.iter().find_map(|c| refused_in(c, scope, bound))
}
Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
refused_in(child, scope, bound)
}
}
}
fn literal_of(v: &Value) -> Option<LiteralValue> {
Some(match v {
Value::U64(n) => LiteralValue::unsigned(*n),
Value::I64(n) => LiteralValue::Int(*n),
Value::F64(x) => LiteralValue::Float(*x),
Value::Bool(b) => LiteralValue::Bool(*b),
Value::Str(s) => LiteralValue::String(s.to_string()),
Value::Json(j) => LiteralValue::Json(j.as_ref().clone()),
_ => return None,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::iteration::comprehension::source::Source;
use crate::kernel::interp::NoScope;
fn generator(name: &str, expr: &str) -> Comprehension {
Comprehension::clause(
name,
Source::Generator {
expr: expr.into(),
cardinality_hint: None,
},
)
}
#[test]
fn a_context_free_generator_becomes_a_literal_of_its_values() {
let out = flatten_static_sources(&generator("k", "fib(6)"), &NoScope::new());
let Comprehension::Clause {
source: Source::Literal { values },
..
} = out
else {
panic!("expected a literal clause, got {out:?}");
};
assert_eq!(
values,
vec![
LiteralValue::Int(1),
LiteralValue::Int(1),
LiteralValue::Int(2),
LiteralValue::Int(3),
LiteralValue::Int(5),
LiteralValue::Int(8),
]
);
}
#[test]
fn a_generator_over_a_bound_coordinate_or_an_unresolved_name_is_kept() {
let ast =
Comprehension::cartesian(vec![generator("k", "fib(3)"), generator("j", "pow2({k})")]);
let out = flatten_static_sources(&ast, &NoScope::new());
let Comprehension::Cartesian { children } = &out else {
panic!("{out:?}");
};
assert!(matches!(
&children[0],
Comprehension::Clause {
source: Source::Literal { .. },
..
}
));
assert!(matches!(
&children[1],
Comprehension::Clause {
source: Source::Generator { expr, cardinality_hint: None },
..
} if expr == "pow2({k})"
));
let out = flatten_static_sources(&generator("k", "pow2(n)"), &NoScope::new());
assert!(matches!(
out,
Comprehension::Clause {
source: Source::Generator {
cardinality_hint: None,
..
},
..
}
));
}
#[test]
fn values_without_a_literal_form_keep_the_call_and_gain_the_count() {
let out =
flatten_static_sources(&generator("p", "partitions(\"*/4\", 100)"), &NoScope::new());
assert!(
matches!(
out,
Comprehension::Clause {
source: Source::Generator {
cardinality_hint: Some(4),
..
},
..
}
),
"{out:?}"
);
}
#[test]
fn a_call_the_compile_cannot_evaluate_is_left_to_the_traversal() {
let out = flatten_static_sources(&generator("k", "fib(-1)"), &NoScope::new());
assert_eq!(out, generator("k", "fib(-1)"));
}
#[test]
fn a_value_above_i64_max_flattens_to_an_unsigned_literal() {
let out = flatten_static_sources(&generator("k", "pow2(64)"), &NoScope::new());
let Comprehension::Clause {
source: Source::Literal { values },
..
} = out
else {
panic!("expected a literal clause, got {out:?}");
};
assert_eq!(values.len(), 64);
assert_eq!(values[62], LiteralValue::Int(1 << 62));
assert_eq!(values[63], LiteralValue::UInt(1 << 63));
}
#[test]
fn a_refused_generator_is_found_when_its_arguments_resolve() {
let scope = NoScope::new();
let ast = Comprehension::cartesian(vec![
generator("a", "fib(3)"),
generator("b", "binomial(70)"),
]);
let flat = flatten_static_sources(&ast, &scope);
let (name, message) = first_refused_generator(&flat, &scope).unwrap();
assert_eq!(name, "b");
assert!(
message.starts_with("binomial(70): term C(70, "),
"{message}"
);
let dependent =
Comprehension::cartesian(vec![generator("k", "fib(3)"), generator("j", "fib({k})")]);
assert_eq!(first_refused_generator(&dependent, &scope), None);
assert_eq!(
first_refused_generator(&generator("k", "fib(n)"), &scope),
None
);
}
}