1use crate::ast::Value;
36use crate::dsl::compile::eval_const_expr_for;
37use crate::iteration::comprehension::ast::Comprehension;
38use crate::iteration::comprehension::runtime::{RuntimeError, RuntimeTuple};
39use crate::iteration::comprehension::source::{LiteralValue, Source};
40use crate::kernel::interp::{Layered, Lookup, interpolate_with_lookup};
41use polydat_grammar::comprehension::predicate::{
42 Comparison, Predicate, PredicateKind, PredicateLiteral, parse_predicate, predicate_reads,
43};
44
45#[derive(Debug, Clone)]
47pub struct CompiledPredicate {
48 text: String,
49 tree: Predicate,
50 bare: Vec<std::ops::Range<usize>>,
53}
54
55#[derive(Debug, Clone, PartialEq)]
57enum Scalar {
58 Int(i128),
59 Float(f64),
60 Str(String),
61 Bool(bool),
62}
63
64impl std::fmt::Display for Scalar {
65 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
66 match self {
67 Scalar::Int(n) => write!(f, "{n}"),
68 Scalar::Float(x) => write!(f, "{x}"),
69 Scalar::Str(s) => write!(f, "{s:?}"),
70 Scalar::Bool(b) => write!(f, "{b}"),
71 }
72 }
73}
74
75impl CompiledPredicate {
76 pub fn new(text: &str) -> Self {
78 let tree = parse_predicate(text).unwrap_or(Predicate {
79 kind: PredicateKind::Expr,
80 span: 0..text.len(),
81 });
82 let mut bare = Vec::new();
83 collect_bare(&tree, text, &mut bare);
84 Self {
85 text: text.to_string(),
86 tree,
87 bare,
88 }
89 }
90
91 pub fn text(&self) -> &str {
93 &self.text
94 }
95
96 pub fn keeps(&self, tuple: &RuntimeTuple, scope: &dyn Lookup) -> Result<bool, RuntimeError> {
99 self.eval(&self.tree, tuple, scope)
100 .and_then(|v| v.map_or(Ok(false), |v| truth(&v)))
101 .map_err(|message| RuntimeError::FilterEval {
102 predicate: self.text.clone(),
103 message,
104 })
105 }
106
107 fn eval(
109 &self,
110 node: &Predicate,
111 tuple: &RuntimeTuple,
112 scope: &dyn Lookup,
113 ) -> Result<Option<Scalar>, String> {
114 Ok(Some(match &node.kind {
115 PredicateKind::Or(parts) => {
116 for part in parts {
117 let Some(value) = self.eval(part, tuple, scope)? else {
118 return Ok(None);
119 };
120 if truth(&value)? {
121 return Ok(Some(Scalar::Bool(true)));
122 }
123 }
124 Scalar::Bool(false)
125 }
126 PredicateKind::And(parts) => {
127 for part in parts {
128 let Some(value) = self.eval(part, tuple, scope)? else {
129 return Ok(None);
130 };
131 if !truth(&value)? {
132 return Ok(Some(Scalar::Bool(false)));
133 }
134 }
135 Scalar::Bool(true)
136 }
137 PredicateKind::Not(inner) => {
138 let Some(value) = self.eval(inner, tuple, scope)? else {
139 return Ok(None);
140 };
141 Scalar::Bool(!truth(&value)?)
142 }
143 PredicateKind::Compare(op, a, b) => {
144 let Some(a) = self.eval(a, tuple, scope)? else {
145 return Ok(None);
146 };
147 let Some(b) = self.eval(b, tuple, scope)? else {
148 return Ok(None);
149 };
150 Scalar::Bool(compare(*op, &a, &b)?)
151 }
152 PredicateKind::In(needle, items) => {
153 let Some(needle) = self.eval(needle, tuple, scope)? else {
154 return Ok(None);
155 };
156 let mut hit = false;
157 for item in items {
158 let Some(item) = self.eval(item, tuple, scope)? else {
159 return Ok(None);
160 };
161 hit |= scalar_eq(&needle, &item);
162 }
163 Scalar::Bool(hit)
164 }
165 PredicateKind::Element(name) => {
166 let value = match tuple.iter().find(|(n, _)| n == name) {
167 Some((_, v)) => Some(v.clone()),
168 None => scope.lookup(name),
169 };
170 match value {
171 None | Some(Value::None) => return Ok(None),
172 Some(value) => scalar(&value)
173 .ok_or_else(|| format!("`{{{name}}}` is {value:?}, not a scalar"))?,
174 }
175 }
176 PredicateKind::Literal(literal) => match literal {
177 PredicateLiteral::Int(n) => Scalar::Int(*n),
178 PredicateLiteral::Float(f) => Scalar::Float(*f),
179 PredicateLiteral::Str(s) => Scalar::Str(s.clone()),
180 PredicateLiteral::Bool(b) => Scalar::Bool(*b),
181 },
182 PredicateKind::Arith(..) | PredicateKind::Expr => {
183 if self.bare.contains(&node.span) {
184 return Ok(None);
185 }
186 let text = node.text(&self.text);
187 let layered = Layered {
188 prefix: tuple,
189 inner: scope,
190 };
191 let none = std::cell::Cell::new(false);
193 let interpolated =
194 interpolate_with_lookup(text, |name| match layered.lookup(name) {
195 None | Some(Value::None) => {
196 none.set(true);
197 None
198 }
199 Some(value) => Some(value.to_display_string()),
200 });
201 if none.get() {
202 return Ok(None);
203 }
204 let value = eval_const_expr_for(&interpolated?, scope.ledger())
205 .map_err(|e| e.to_string())?;
206 if matches!(value, Value::None) {
207 return Ok(None);
208 }
209 scalar(&value).ok_or_else(|| format!("`{text}` is {value:?}, not a scalar"))?
210 }
211 }))
212 }
213}
214
215fn collect_bare(node: &Predicate, text: &str, out: &mut Vec<std::ops::Range<usize>>) {
218 match &node.kind {
219 PredicateKind::Or(parts) | PredicateKind::And(parts) => {
220 for part in parts {
221 collect_bare(part, text, out);
222 }
223 }
224 PredicateKind::Not(inner) => collect_bare(inner, text, out),
225 PredicateKind::Compare(_, a, b) => {
226 collect_bare(a, text, out);
227 collect_bare(b, text, out);
228 }
229 PredicateKind::In(needle, items) => {
230 collect_bare(needle, text, out);
231 for item in items {
232 collect_bare(item, text, out);
233 }
234 }
235 PredicateKind::Element(_) | PredicateKind::Literal(_) => {}
236 PredicateKind::Arith(..) | PredicateKind::Expr => {
237 if !predicate_reads(node.text(text)).bare.is_empty() {
238 out.push(node.span.clone());
239 }
240 }
241 }
242}
243
244#[derive(Debug, Clone, Copy, PartialEq, Eq)]
246pub enum ValueKind {
247 Int,
249 Float,
251 Str,
253 Bool,
255}
256
257impl CompiledPredicate {
258 pub fn is_total(&self, kind_of: &dyn Fn(&str) -> Option<ValueKind>) -> bool {
273 total_kind(&self.tree, kind_of).is_some_and(has_truth)
274 }
275}
276
277fn total_kind(node: &Predicate, kind_of: &dyn Fn(&str) -> Option<ValueKind>) -> Option<ValueKind> {
280 use polydat_grammar::ast::BinOpKind;
281 let truthful = |p: &Predicate| total_kind(p, kind_of).is_some_and(has_truth);
282 match &node.kind {
283 PredicateKind::Or(parts) | PredicateKind::And(parts) => {
284 parts.iter().all(truthful).then_some(ValueKind::Bool)
285 }
286 PredicateKind::Not(inner) => truthful(inner).then_some(ValueKind::Bool),
287 PredicateKind::Compare(op, a, b) => {
288 let (a, b) = (total_kind(a, kind_of)?, total_kind(b, kind_of)?);
289 let ordered = matches!(
290 (a, b),
291 (
292 ValueKind::Int | ValueKind::Float,
293 ValueKind::Int | ValueKind::Float
294 ) | (ValueKind::Str, ValueKind::Str)
295 | (ValueKind::Bool, ValueKind::Bool)
296 );
297 (matches!(op, Comparison::Eq | Comparison::Ne) || ordered).then_some(ValueKind::Bool)
298 }
299 PredicateKind::In(needle, items) => {
300 total_kind(needle, kind_of)?;
301 for item in items {
302 total_kind(item, kind_of)?;
303 }
304 Some(ValueKind::Bool)
305 }
306 PredicateKind::Element(name) => kind_of(name),
307 PredicateKind::Literal(literal) => Some(match literal {
308 PredicateLiteral::Int(_) => ValueKind::Int,
309 PredicateLiteral::Float(_) => ValueKind::Float,
310 PredicateLiteral::Str(_) => ValueKind::Str,
311 PredicateLiteral::Bool(_) => ValueKind::Bool,
312 }),
313 PredicateKind::Arith(op, a, b) => {
314 let operand = |p: &Predicate| {
318 let fits = match &p.kind {
319 PredicateKind::Literal(PredicateLiteral::Int(n)) => u64::try_from(*n).is_ok(),
320 PredicateKind::Literal(PredicateLiteral::Float(f)) => *f >= 0.0,
321 _ => true,
322 };
323 total_kind(p, kind_of)
324 .filter(|k| fits && matches!(k, ValueKind::Int | ValueKind::Float))
325 };
326 let (ka, kb) = (operand(a)?, operand(b)?);
327 if matches!(op, BinOpKind::Div | BinOpKind::Mod) {
328 let non_zero_constant = match &b.kind {
329 PredicateKind::Literal(PredicateLiteral::Int(n)) => *n != 0,
330 PredicateKind::Literal(PredicateLiteral::Float(f)) => *f != 0.0,
331 _ => false,
332 };
333 if !non_zero_constant {
334 return None;
335 }
336 }
337 Some(
338 if ka == ValueKind::Int && kb == ValueKind::Int && *op != BinOpKind::Pow {
339 ValueKind::Int
340 } else {
341 ValueKind::Float
342 },
343 )
344 }
345 PredicateKind::Expr => None,
346 }
347}
348
349fn has_truth(kind: ValueKind) -> bool {
351 kind != ValueKind::Str
352}
353
354pub fn element_kind(c: &Comprehension, name: &str) -> Option<ValueKind> {
359 let mut kinds = Vec::new();
360 collect_kinds(c, name, &mut kinds);
361 let first = (*kinds.first()?)?;
362 kinds.iter().all(|k| *k == Some(first)).then_some(first)
363}
364
365fn collect_kinds(c: &Comprehension, name: &str, out: &mut Vec<Option<ValueKind>>) {
366 match c {
367 Comprehension::Clause { name: n, source } if n == name => out.push(source_kind(source)),
368 Comprehension::Clause { .. } => {}
369 Comprehension::Cartesian { children }
370 | Comprehension::Zip { children, .. }
371 | Comprehension::Union { children } => {
372 for child in children {
373 collect_kinds(child, name, out);
374 }
375 }
376 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
377 collect_kinds(child, name, out);
378 }
379 }
380}
381
382fn source_kind(source: &Source) -> Option<ValueKind> {
384 match source {
385 Source::IntRange { .. } => Some(ValueKind::Int),
386 Source::ContinuousInterval { .. } | Source::Distribution { .. } => Some(ValueKind::Float),
387 Source::Literal { values } => {
388 let kind = |v: &LiteralValue| match v {
389 LiteralValue::Int(_) | LiteralValue::UInt(_) => Some(ValueKind::Int),
390 LiteralValue::Float(_) => Some(ValueKind::Float),
391 LiteralValue::String(_) => Some(ValueKind::Str),
392 LiteralValue::Bool(_) => Some(ValueKind::Bool),
393 LiteralValue::Json(_) => None,
394 };
395 let first = kind(values.first()?)?;
396 values
397 .iter()
398 .all(|v| kind(v) == Some(first))
399 .then_some(first)
400 }
401 Source::Generator { .. } | Source::WorkloadParamList { .. } => None,
402 }
403}
404
405fn scalar(value: &Value) -> Option<Scalar> {
408 Some(match value {
409 Value::U64(n) => Scalar::Int(i128::from(*n)),
410 Value::I64(n) => Scalar::Int(i128::from(*n)),
411 Value::F64(f) => Scalar::Float(*f),
412 Value::Str(s) => Scalar::Str(s.to_string()),
413 Value::Bool(b) => Scalar::Bool(*b),
414 Value::Json(j) => match j.as_ref() {
415 serde_json::Value::Number(n) if n.is_i64() => Scalar::Int(i128::from(n.as_i64()?)),
416 serde_json::Value::Number(n) if n.is_u64() => Scalar::Int(i128::from(n.as_u64()?)),
417 serde_json::Value::Number(n) => Scalar::Float(n.as_f64()?),
418 serde_json::Value::String(s) => Scalar::Str(s.clone()),
419 serde_json::Value::Bool(b) => Scalar::Bool(*b),
420 _ => return None,
421 },
422 _ => return None,
423 })
424}
425
426fn truth(value: &Scalar) -> Result<bool, String> {
428 match value {
429 Scalar::Bool(b) => Ok(*b),
430 Scalar::Int(n) => Ok(*n != 0),
431 Scalar::Float(f) => Ok(*f != 0.0),
432 Scalar::Str(s) => Err(format!("expected bool/u64/f64, got {s:?}")),
433 }
434}
435
436fn scalar_eq(a: &Scalar, b: &Scalar) -> bool {
437 match (a, b) {
438 (Scalar::Int(x), Scalar::Float(y)) | (Scalar::Float(y), Scalar::Int(x)) => {
439 (*x as f64) == *y
440 }
441 _ => a == b,
442 }
443}
444
445fn compare(op: Comparison, a: &Scalar, b: &Scalar) -> Result<bool, String> {
446 use std::cmp::Ordering;
447 let ordering = match op {
448 Comparison::Eq => return Ok(scalar_eq(a, b)),
449 Comparison::Ne => return Ok(!scalar_eq(a, b)),
450 _ => match (a, b) {
451 (Scalar::Int(x), Scalar::Int(y)) => Some(x.cmp(y)),
452 (Scalar::Float(x), Scalar::Float(y)) => x.partial_cmp(y),
453 (Scalar::Int(x), Scalar::Float(y)) => (*x as f64).partial_cmp(y),
454 (Scalar::Float(x), Scalar::Int(y)) => x.partial_cmp(&(*y as f64)),
455 (Scalar::Str(x), Scalar::Str(y)) => Some(x.cmp(y)),
456 (Scalar::Bool(x), Scalar::Bool(y)) => Some(x.cmp(y)),
457 _ => return Err(format!("cannot order {a} and {b}")),
458 },
459 };
460 Ok(ordering.is_some_and(|o| match op {
462 Comparison::Lt => o == Ordering::Less,
463 Comparison::Le => o != Ordering::Greater,
464 Comparison::Gt => o == Ordering::Greater,
465 _ => o != Ordering::Less,
466 }))
467}
468
469#[cfg(test)]
470mod tests {
471 use super::*;
472 use crate::kernel::interp::NoScope;
473 use std::sync::Arc;
474
475 fn tuple(bindings: &[(&str, Value)]) -> RuntimeTuple {
476 bindings
477 .iter()
478 .map(|(n, v)| ((*n).to_string(), v.clone()))
479 .collect()
480 }
481
482 fn keeps(predicate: &str, t: &RuntimeTuple) -> bool {
483 CompiledPredicate::new(predicate)
484 .keeps(t, &NoScope::new())
485 .unwrap()
486 }
487
488 #[test]
491 fn not_binds_tighter_than_or() {
492 for (done, retry, kept) in [
493 (false, false, true),
494 (false, true, true),
495 (true, false, false),
496 (true, true, true),
497 ] {
498 let t = tuple(&[("done", Value::Bool(done)), ("retry", Value::Bool(retry))]);
499 assert_eq!(keeps("!{done} || {retry}", &t), kept, "{done} {retry}");
500 assert_eq!(keeps("!({done} || {retry})", &t), !done && !retry);
501 assert_eq!(keeps("!{done} && {retry}", &t), !done && retry);
502 }
503 }
504
505 #[test]
509 fn every_operator_pair_evaluates_as_grouped() {
510 let pairs = [
511 ("{a} || {b} && {c}", "{a} || ({b} && {c})"),
512 ("{a} && {b} || {c}", "({a} && {b}) || {c}"),
513 ("!{a} || {b}", "(!{a}) || {b}"),
514 ("!{a} && {b}", "(!{a}) && {b}"),
515 ("{a} == {b} || {c}", "({a} == {b}) || {c}"),
516 ("{a} || {b} == {c}", "{a} || ({b} == {c})"),
517 ("{a} != {b} && {c}", "({a} != {b}) && {c}"),
518 (
519 "{x} < 2 || {y} >= 1 && {a}",
520 "({x} < 2) || (({y} >= 1) && {a})",
521 ),
522 ("{x} + 1 > {y} + {y}", "({x} + 1) > ({y} + {y})"),
523 (
524 "{x} + {x} == {y} + 1 || {c}",
525 "(({x} + {x}) == ({y} + 1)) || {c}",
526 ),
527 ("{x} < {y} == {a}", "({x} < {y}) == {a}"),
528 ("{x} in [0, 2] || {a}", "({x} in [0, 2]) || {a}"),
529 ("!{a} == {b}", "(!{a}) == {b}"),
530 ];
531 for bits in 0..8u8 {
532 for x in 0..3u64 {
533 for y in 0..3u64 {
534 let t = tuple(&[
535 ("a", Value::Bool(bits & 1 != 0)),
536 ("b", Value::Bool(bits & 2 != 0)),
537 ("c", Value::Bool(bits & 4 != 0)),
538 ("x", Value::U64(x)),
539 ("y", Value::U64(y)),
540 ]);
541 for (bare, grouped) in pairs {
542 assert_eq!(keeps(bare, &t), keeps(grouped, &t), "{bare} at {t:?}");
543 }
544 }
545 }
546 }
547 }
548
549 #[test]
557 fn totality_is_read_off_the_tree() {
558 let kind_of = |name: &str| match name {
559 "k" | "m" => Some(ValueKind::Int),
560 "x" => Some(ValueKind::Float),
561 "w" => Some(ValueKind::Str),
562 "b" => Some(ValueKind::Bool),
563 _ => None,
564 };
565 let total = [
566 "{k} > 1",
567 "{k} < {x}",
568 "{w} == 2",
569 "{w} != {k} && {b}",
570 "{w} >= \"m\" || !{b}",
571 "{k} in [1, \"a\", true]",
572 "{k} * 2 + 1 > {m}",
573 "{k} - {m} >= 0",
574 "{k} / 2 == 1",
575 "{x} % 1.5 < 1",
576 "{k} ** 2 > {x}",
577 "{k} + 1",
578 "{x}",
579 ];
580 let partial = [
581 "{w} > 2",
582 "{b} < 1",
583 "{w}",
584 "{w} && {b}",
585 "{k} / {m} == 1",
586 "{k} % 0 == 1",
587 "{w} + 1 > 2",
588 "{k} + -1 > 2",
589 "{z} > 1",
590 "u64_add({k}, 1) > 2",
591 "{x} as u64 > 2",
592 "{k} & 1 == 1",
593 ];
594 for p in total {
595 assert!(CompiledPredicate::new(p).is_total(&kind_of), "{p}");
596 }
597 for p in partial {
598 assert!(!CompiledPredicate::new(p).is_total(&kind_of), "{p}");
599 }
600 }
601
602 #[test]
607 fn a_bare_word_is_a_name_that_reads_none() {
608 let t = tuple(&[("region", Value::Str(Arc::from("us-east")))]);
609 assert!(keeps("{region} == \"us-east\"", &t));
610 for predicate in [
611 "{region} == us-east",
612 "{region} in [us-west, \"us-east\"]",
613 "{region} != eu",
614 "!({region} != eu)",
615 "u64_add(1, width) > 0",
616 ] {
617 assert!(!keeps(predicate, &t), "{predicate}");
618 }
619 let error = CompiledPredicate::new("nosuch({region}) > 1")
620 .keeps(&t, &NoScope::new())
621 .unwrap_err()
622 .to_string();
623 assert!(error.contains("nosuch"), "{error}");
624 }
625
626 #[test]
629 fn mixed_kinds_are_unequal_and_unordered() {
630 let t = tuple(&[("c", Value::Str(Arc::from("s0")))]);
631 assert!(keeps("{c} != 2", &t));
632 assert!(!keeps("{c} == 2", &t));
633 assert!(keeps("{c} == \"s0\"", &t));
634 assert!(keeps("{c} == 's0'", &t));
635 assert!(!keeps("{c} == 2 && {c} > 2", &t));
637 assert!(keeps("{c} != 2 || {c} > 2", &t));
638 let error = CompiledPredicate::new("{c} > 2")
639 .keeps(&t, &NoScope::new())
640 .unwrap_err();
641 assert!(
642 error.to_string().contains("cannot order \"s0\" and 2"),
643 "{error}"
644 );
645 }
646
647 #[test]
649 fn a_call_evaluates_through_the_kernel() {
650 let t = tuple(&[("a", Value::U64(3)), ("b", Value::U64(5))]);
651 assert!(keeps("u64_add({a}, {b}) > 7", &t));
652 assert!(!keeps("u64_add({a}, {b}) > 8", &t));
653 assert!(keeps("!(u64_add({a}, {b}) > 8)", &t));
654 }
655
656 #[test]
662 fn none_propagates_and_keeps_no_tuple() {
663 let unbound = tuple(&[("k", Value::U64(1))]);
664 let bound_none = tuple(&[("k", Value::U64(1)), ("z", Value::None)]);
665 for t in [&unbound, &bound_none] {
666 for predicate in [
667 "{z} > 1",
668 "{z} != 1",
669 "!({z} == 1)",
670 "{z} in [1, 2]",
671 "{k} in [{z}, 1]",
672 "{z} + 1 > 0",
673 "u64_add({z}, 1) > 0",
674 "{z} > 1 || {k} == 1",
675 "{k} == 1 && {z} != 1",
676 "{k} == 1 && !({z} > 1)",
677 ] {
678 assert!(!keeps(predicate, t), "{predicate} over {t:?}");
679 }
680 assert!(keeps("{k} == 1 || {z} > 1", t));
681 assert!(!keeps("{k} == 2 && {z} > 1", t));
682 }
683 }
684}