use crate::ast::Value;
use crate::dsl::compile::eval_const_expr_for;
use crate::iteration::comprehension::ast::Comprehension;
use crate::iteration::comprehension::runtime::{RuntimeError, RuntimeTuple};
use crate::iteration::comprehension::source::{LiteralValue, Source};
use crate::kernel::interp::{Layered, Lookup, interpolate_with_lookup};
use polydat_grammar::comprehension::predicate::{
Comparison, Predicate, PredicateKind, PredicateLiteral, parse_predicate, predicate_reads,
};
#[derive(Debug, Clone)]
pub struct CompiledPredicate {
text: String,
tree: Predicate,
bare: Vec<std::ops::Range<usize>>,
}
#[derive(Debug, Clone, PartialEq)]
enum Scalar {
Int(i128),
Float(f64),
Str(String),
Bool(bool),
}
impl std::fmt::Display for Scalar {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Scalar::Int(n) => write!(f, "{n}"),
Scalar::Float(x) => write!(f, "{x}"),
Scalar::Str(s) => write!(f, "{s:?}"),
Scalar::Bool(b) => write!(f, "{b}"),
}
}
}
impl CompiledPredicate {
pub fn new(text: &str) -> Self {
let tree = parse_predicate(text).unwrap_or(Predicate {
kind: PredicateKind::Expr,
span: 0..text.len(),
});
let mut bare = Vec::new();
collect_bare(&tree, text, &mut bare);
Self {
text: text.to_string(),
tree,
bare,
}
}
pub fn text(&self) -> &str {
&self.text
}
pub fn keeps(&self, tuple: &RuntimeTuple, scope: &dyn Lookup) -> Result<bool, RuntimeError> {
self.eval(&self.tree, tuple, scope)
.and_then(|v| v.map_or(Ok(false), |v| truth(&v)))
.map_err(|message| RuntimeError::FilterEval {
predicate: self.text.clone(),
message,
})
}
fn eval(
&self,
node: &Predicate,
tuple: &RuntimeTuple,
scope: &dyn Lookup,
) -> Result<Option<Scalar>, String> {
Ok(Some(match &node.kind {
PredicateKind::Or(parts) => {
for part in parts {
let Some(value) = self.eval(part, tuple, scope)? else {
return Ok(None);
};
if truth(&value)? {
return Ok(Some(Scalar::Bool(true)));
}
}
Scalar::Bool(false)
}
PredicateKind::And(parts) => {
for part in parts {
let Some(value) = self.eval(part, tuple, scope)? else {
return Ok(None);
};
if !truth(&value)? {
return Ok(Some(Scalar::Bool(false)));
}
}
Scalar::Bool(true)
}
PredicateKind::Not(inner) => {
let Some(value) = self.eval(inner, tuple, scope)? else {
return Ok(None);
};
Scalar::Bool(!truth(&value)?)
}
PredicateKind::Compare(op, a, b) => {
let Some(a) = self.eval(a, tuple, scope)? else {
return Ok(None);
};
let Some(b) = self.eval(b, tuple, scope)? else {
return Ok(None);
};
Scalar::Bool(compare(*op, &a, &b)?)
}
PredicateKind::In(needle, items) => {
let Some(needle) = self.eval(needle, tuple, scope)? else {
return Ok(None);
};
let mut hit = false;
for item in items {
let Some(item) = self.eval(item, tuple, scope)? else {
return Ok(None);
};
hit |= scalar_eq(&needle, &item);
}
Scalar::Bool(hit)
}
PredicateKind::Element(name) => {
let value = match tuple.iter().find(|(n, _)| n == name) {
Some((_, v)) => Some(v.clone()),
None => scope.lookup(name),
};
match value {
None | Some(Value::None) => return Ok(None),
Some(value) => scalar(&value)
.ok_or_else(|| format!("`{{{name}}}` is {value:?}, not a scalar"))?,
}
}
PredicateKind::Literal(literal) => match literal {
PredicateLiteral::Int(n) => Scalar::Int(*n),
PredicateLiteral::Float(f) => Scalar::Float(*f),
PredicateLiteral::Str(s) => Scalar::Str(s.clone()),
PredicateLiteral::Bool(b) => Scalar::Bool(*b),
},
PredicateKind::Arith(..) | PredicateKind::Expr => {
if self.bare.contains(&node.span) {
return Ok(None);
}
let text = node.text(&self.text);
let layered = Layered {
prefix: tuple,
inner: scope,
};
let none = std::cell::Cell::new(false);
let interpolated =
interpolate_with_lookup(text, |name| match layered.lookup(name) {
None | Some(Value::None) => {
none.set(true);
None
}
Some(value) => Some(value.to_display_string()),
});
if none.get() {
return Ok(None);
}
let value = eval_const_expr_for(&interpolated?, scope.ledger())
.map_err(|e| e.to_string())?;
if matches!(value, Value::None) {
return Ok(None);
}
scalar(&value).ok_or_else(|| format!("`{text}` is {value:?}, not a scalar"))?
}
}))
}
}
fn collect_bare(node: &Predicate, text: &str, out: &mut Vec<std::ops::Range<usize>>) {
match &node.kind {
PredicateKind::Or(parts) | PredicateKind::And(parts) => {
for part in parts {
collect_bare(part, text, out);
}
}
PredicateKind::Not(inner) => collect_bare(inner, text, out),
PredicateKind::Compare(_, a, b) => {
collect_bare(a, text, out);
collect_bare(b, text, out);
}
PredicateKind::In(needle, items) => {
collect_bare(needle, text, out);
for item in items {
collect_bare(item, text, out);
}
}
PredicateKind::Element(_) | PredicateKind::Literal(_) => {}
PredicateKind::Arith(..) | PredicateKind::Expr => {
if !predicate_reads(node.text(text)).bare.is_empty() {
out.push(node.span.clone());
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ValueKind {
Int,
Float,
Str,
Bool,
}
impl CompiledPredicate {
pub fn is_total(&self, kind_of: &dyn Fn(&str) -> Option<ValueKind>) -> bool {
total_kind(&self.tree, kind_of).is_some_and(has_truth)
}
}
fn total_kind(node: &Predicate, kind_of: &dyn Fn(&str) -> Option<ValueKind>) -> Option<ValueKind> {
use polydat_grammar::ast::BinOpKind;
let truthful = |p: &Predicate| total_kind(p, kind_of).is_some_and(has_truth);
match &node.kind {
PredicateKind::Or(parts) | PredicateKind::And(parts) => {
parts.iter().all(truthful).then_some(ValueKind::Bool)
}
PredicateKind::Not(inner) => truthful(inner).then_some(ValueKind::Bool),
PredicateKind::Compare(op, a, b) => {
let (a, b) = (total_kind(a, kind_of)?, total_kind(b, kind_of)?);
let ordered = matches!(
(a, b),
(
ValueKind::Int | ValueKind::Float,
ValueKind::Int | ValueKind::Float
) | (ValueKind::Str, ValueKind::Str)
| (ValueKind::Bool, ValueKind::Bool)
);
(matches!(op, Comparison::Eq | Comparison::Ne) || ordered).then_some(ValueKind::Bool)
}
PredicateKind::In(needle, items) => {
total_kind(needle, kind_of)?;
for item in items {
total_kind(item, kind_of)?;
}
Some(ValueKind::Bool)
}
PredicateKind::Element(name) => kind_of(name),
PredicateKind::Literal(literal) => Some(match literal {
PredicateLiteral::Int(_) => ValueKind::Int,
PredicateLiteral::Float(_) => ValueKind::Float,
PredicateLiteral::Str(_) => ValueKind::Str,
PredicateLiteral::Bool(_) => ValueKind::Bool,
}),
PredicateKind::Arith(op, a, b) => {
let operand = |p: &Predicate| {
let fits = match &p.kind {
PredicateKind::Literal(PredicateLiteral::Int(n)) => u64::try_from(*n).is_ok(),
PredicateKind::Literal(PredicateLiteral::Float(f)) => *f >= 0.0,
_ => true,
};
total_kind(p, kind_of)
.filter(|k| fits && matches!(k, ValueKind::Int | ValueKind::Float))
};
let (ka, kb) = (operand(a)?, operand(b)?);
if matches!(op, BinOpKind::Div | BinOpKind::Mod) {
let non_zero_constant = match &b.kind {
PredicateKind::Literal(PredicateLiteral::Int(n)) => *n != 0,
PredicateKind::Literal(PredicateLiteral::Float(f)) => *f != 0.0,
_ => false,
};
if !non_zero_constant {
return None;
}
}
Some(
if ka == ValueKind::Int && kb == ValueKind::Int && *op != BinOpKind::Pow {
ValueKind::Int
} else {
ValueKind::Float
},
)
}
PredicateKind::Expr => None,
}
}
fn has_truth(kind: ValueKind) -> bool {
kind != ValueKind::Str
}
pub fn element_kind(c: &Comprehension, name: &str) -> Option<ValueKind> {
let mut kinds = Vec::new();
collect_kinds(c, name, &mut kinds);
let first = (*kinds.first()?)?;
kinds.iter().all(|k| *k == Some(first)).then_some(first)
}
fn collect_kinds(c: &Comprehension, name: &str, out: &mut Vec<Option<ValueKind>>) {
match c {
Comprehension::Clause { name: n, source } if n == name => out.push(source_kind(source)),
Comprehension::Clause { .. } => {}
Comprehension::Cartesian { children }
| Comprehension::Zip { children, .. }
| Comprehension::Union { children } => {
for child in children {
collect_kinds(child, name, out);
}
}
Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
collect_kinds(child, name, out);
}
}
}
fn source_kind(source: &Source) -> Option<ValueKind> {
match source {
Source::IntRange { .. } => Some(ValueKind::Int),
Source::ContinuousInterval { .. } | Source::Distribution { .. } => Some(ValueKind::Float),
Source::Literal { values } => {
let kind = |v: &LiteralValue| match v {
LiteralValue::Int(_) | LiteralValue::UInt(_) => Some(ValueKind::Int),
LiteralValue::Float(_) => Some(ValueKind::Float),
LiteralValue::String(_) => Some(ValueKind::Str),
LiteralValue::Bool(_) => Some(ValueKind::Bool),
LiteralValue::Json(_) => None,
};
let first = kind(values.first()?)?;
values
.iter()
.all(|v| kind(v) == Some(first))
.then_some(first)
}
Source::Generator { .. } | Source::WorkloadParamList { .. } => None,
}
}
fn scalar(value: &Value) -> Option<Scalar> {
Some(match value {
Value::U64(n) => Scalar::Int(i128::from(*n)),
Value::I64(n) => Scalar::Int(i128::from(*n)),
Value::F64(f) => Scalar::Float(*f),
Value::Str(s) => Scalar::Str(s.to_string()),
Value::Bool(b) => Scalar::Bool(*b),
Value::Json(j) => match j.as_ref() {
serde_json::Value::Number(n) if n.is_i64() => Scalar::Int(i128::from(n.as_i64()?)),
serde_json::Value::Number(n) if n.is_u64() => Scalar::Int(i128::from(n.as_u64()?)),
serde_json::Value::Number(n) => Scalar::Float(n.as_f64()?),
serde_json::Value::String(s) => Scalar::Str(s.clone()),
serde_json::Value::Bool(b) => Scalar::Bool(*b),
_ => return None,
},
_ => return None,
})
}
fn truth(value: &Scalar) -> Result<bool, String> {
match value {
Scalar::Bool(b) => Ok(*b),
Scalar::Int(n) => Ok(*n != 0),
Scalar::Float(f) => Ok(*f != 0.0),
Scalar::Str(s) => Err(format!("expected bool/u64/f64, got {s:?}")),
}
}
fn scalar_eq(a: &Scalar, b: &Scalar) -> bool {
match (a, b) {
(Scalar::Int(x), Scalar::Float(y)) | (Scalar::Float(y), Scalar::Int(x)) => {
(*x as f64) == *y
}
_ => a == b,
}
}
fn compare(op: Comparison, a: &Scalar, b: &Scalar) -> Result<bool, String> {
use std::cmp::Ordering;
let ordering = match op {
Comparison::Eq => return Ok(scalar_eq(a, b)),
Comparison::Ne => return Ok(!scalar_eq(a, b)),
_ => match (a, b) {
(Scalar::Int(x), Scalar::Int(y)) => Some(x.cmp(y)),
(Scalar::Float(x), Scalar::Float(y)) => x.partial_cmp(y),
(Scalar::Int(x), Scalar::Float(y)) => (*x as f64).partial_cmp(y),
(Scalar::Float(x), Scalar::Int(y)) => x.partial_cmp(&(*y as f64)),
(Scalar::Str(x), Scalar::Str(y)) => Some(x.cmp(y)),
(Scalar::Bool(x), Scalar::Bool(y)) => Some(x.cmp(y)),
_ => return Err(format!("cannot order {a} and {b}")),
},
};
Ok(ordering.is_some_and(|o| match op {
Comparison::Lt => o == Ordering::Less,
Comparison::Le => o != Ordering::Greater,
Comparison::Gt => o == Ordering::Greater,
_ => o != Ordering::Less,
}))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kernel::interp::NoScope;
use std::sync::Arc;
fn tuple(bindings: &[(&str, Value)]) -> RuntimeTuple {
bindings
.iter()
.map(|(n, v)| ((*n).to_string(), v.clone()))
.collect()
}
fn keeps(predicate: &str, t: &RuntimeTuple) -> bool {
CompiledPredicate::new(predicate)
.keeps(t, &NoScope::new())
.unwrap()
}
#[test]
fn not_binds_tighter_than_or() {
for (done, retry, kept) in [
(false, false, true),
(false, true, true),
(true, false, false),
(true, true, true),
] {
let t = tuple(&[("done", Value::Bool(done)), ("retry", Value::Bool(retry))]);
assert_eq!(keeps("!{done} || {retry}", &t), kept, "{done} {retry}");
assert_eq!(keeps("!({done} || {retry})", &t), !done && !retry);
assert_eq!(keeps("!{done} && {retry}", &t), !done && retry);
}
}
#[test]
fn every_operator_pair_evaluates_as_grouped() {
let pairs = [
("{a} || {b} && {c}", "{a} || ({b} && {c})"),
("{a} && {b} || {c}", "({a} && {b}) || {c}"),
("!{a} || {b}", "(!{a}) || {b}"),
("!{a} && {b}", "(!{a}) && {b}"),
("{a} == {b} || {c}", "({a} == {b}) || {c}"),
("{a} || {b} == {c}", "{a} || ({b} == {c})"),
("{a} != {b} && {c}", "({a} != {b}) && {c}"),
(
"{x} < 2 || {y} >= 1 && {a}",
"({x} < 2) || (({y} >= 1) && {a})",
),
("{x} + 1 > {y} + {y}", "({x} + 1) > ({y} + {y})"),
(
"{x} + {x} == {y} + 1 || {c}",
"(({x} + {x}) == ({y} + 1)) || {c}",
),
("{x} < {y} == {a}", "({x} < {y}) == {a}"),
("{x} in [0, 2] || {a}", "({x} in [0, 2]) || {a}"),
("!{a} == {b}", "(!{a}) == {b}"),
];
for bits in 0..8u8 {
for x in 0..3u64 {
for y in 0..3u64 {
let t = tuple(&[
("a", Value::Bool(bits & 1 != 0)),
("b", Value::Bool(bits & 2 != 0)),
("c", Value::Bool(bits & 4 != 0)),
("x", Value::U64(x)),
("y", Value::U64(y)),
]);
for (bare, grouped) in pairs {
assert_eq!(keeps(bare, &t), keeps(grouped, &t), "{bare} at {t:?}");
}
}
}
}
}
#[test]
fn totality_is_read_off_the_tree() {
let kind_of = |name: &str| match name {
"k" | "m" => Some(ValueKind::Int),
"x" => Some(ValueKind::Float),
"w" => Some(ValueKind::Str),
"b" => Some(ValueKind::Bool),
_ => None,
};
let total = [
"{k} > 1",
"{k} < {x}",
"{w} == 2",
"{w} != {k} && {b}",
"{w} >= \"m\" || !{b}",
"{k} in [1, \"a\", true]",
"{k} * 2 + 1 > {m}",
"{k} - {m} >= 0",
"{k} / 2 == 1",
"{x} % 1.5 < 1",
"{k} ** 2 > {x}",
"{k} + 1",
"{x}",
];
let partial = [
"{w} > 2",
"{b} < 1",
"{w}",
"{w} && {b}",
"{k} / {m} == 1",
"{k} % 0 == 1",
"{w} + 1 > 2",
"{k} + -1 > 2",
"{z} > 1",
"u64_add({k}, 1) > 2",
"{x} as u64 > 2",
"{k} & 1 == 1",
];
for p in total {
assert!(CompiledPredicate::new(p).is_total(&kind_of), "{p}");
}
for p in partial {
assert!(!CompiledPredicate::new(p).is_total(&kind_of), "{p}");
}
}
#[test]
fn a_bare_word_is_a_name_that_reads_none() {
let t = tuple(&[("region", Value::Str(Arc::from("us-east")))]);
assert!(keeps("{region} == \"us-east\"", &t));
for predicate in [
"{region} == us-east",
"{region} in [us-west, \"us-east\"]",
"{region} != eu",
"!({region} != eu)",
"u64_add(1, width) > 0",
] {
assert!(!keeps(predicate, &t), "{predicate}");
}
let error = CompiledPredicate::new("nosuch({region}) > 1")
.keeps(&t, &NoScope::new())
.unwrap_err()
.to_string();
assert!(error.contains("nosuch"), "{error}");
}
#[test]
fn mixed_kinds_are_unequal_and_unordered() {
let t = tuple(&[("c", Value::Str(Arc::from("s0")))]);
assert!(keeps("{c} != 2", &t));
assert!(!keeps("{c} == 2", &t));
assert!(keeps("{c} == \"s0\"", &t));
assert!(keeps("{c} == 's0'", &t));
assert!(!keeps("{c} == 2 && {c} > 2", &t));
assert!(keeps("{c} != 2 || {c} > 2", &t));
let error = CompiledPredicate::new("{c} > 2")
.keeps(&t, &NoScope::new())
.unwrap_err();
assert!(
error.to_string().contains("cannot order \"s0\" and 2"),
"{error}"
);
}
#[test]
fn a_call_evaluates_through_the_kernel() {
let t = tuple(&[("a", Value::U64(3)), ("b", Value::U64(5))]);
assert!(keeps("u64_add({a}, {b}) > 7", &t));
assert!(!keeps("u64_add({a}, {b}) > 8", &t));
assert!(keeps("!(u64_add({a}, {b}) > 8)", &t));
}
#[test]
fn none_propagates_and_keeps_no_tuple() {
let unbound = tuple(&[("k", Value::U64(1))]);
let bound_none = tuple(&[("k", Value::U64(1)), ("z", Value::None)]);
for t in [&unbound, &bound_none] {
for predicate in [
"{z} > 1",
"{z} != 1",
"!({z} == 1)",
"{z} in [1, 2]",
"{k} in [{z}, 1]",
"{z} + 1 > 0",
"u64_add({z}, 1) > 0",
"{z} > 1 || {k} == 1",
"{k} == 1 && {z} != 1",
"{k} == 1 && !({z} > 1)",
] {
assert!(!keeps(predicate, t), "{predicate} over {t:?}");
}
assert!(keeps("{k} == 1 || {z} > 1", t));
assert!(!keeps("{k} == 2 && {z} > 1", t));
}
}
}