Skip to main content

scan_core/grammar/
boolean.rs

1use std::ops::{BitAnd, BitOr, Not};
2
3use get_size2::GetSize;
4use rand::{Rng, RngExt};
5
6use crate::{
7    Expression, Type, TypeError, Val,
8    grammar::{FloatExpr, IntegerExpr, NaturalExpr},
9};
10
11/// Boolean expressions.
12#[derive(Debug, Clone, GetSize)]
13pub enum BooleanExpr<V>
14where
15    V: Clone,
16{
17    /// A Boolean constant (`true` or `false`).
18    Const(bool),
19    /// A Boolean variable.
20    Var(V),
21    /// A Bernoulli distribution with the given probability.
22    Rand(FloatExpr<V>),
23    // -----------------
24    // Logical operators
25    // -----------------
26    /// n-ary logical conjunction.
27    And(Vec<BooleanExpr<V>>),
28    /// n-ary logical disjunction.
29    Or(Vec<BooleanExpr<V>>),
30    /// Logical implication.
31    Implies(Box<(BooleanExpr<V>, BooleanExpr<V>)>),
32    /// Logical negation.
33    Not(Box<BooleanExpr<V>>),
34    // ------------
35    // (In)Equality
36    // ------------
37    /// Equality of Natural expressions.
38    NatEqual(NaturalExpr<V>, NaturalExpr<V>),
39    /// Equality of Integer expressions.
40    IntEqual(IntegerExpr<V>, IntegerExpr<V>),
41    /// Equality of Float expressions.
42    FloatEqual(FloatExpr<V>, FloatExpr<V>),
43    /// Inequality of Natural expressions: LHS greater than RHS.
44    NatGreater(NaturalExpr<V>, NaturalExpr<V>),
45    /// Inequality of Integer expressions: LHS greater than RHS.
46    IntGreater(IntegerExpr<V>, IntegerExpr<V>),
47    /// Inequality of Float expressions: LHS greater than RHS.
48    FloatGreater(FloatExpr<V>, FloatExpr<V>),
49    /// Inequality of Natural expressions: LHS greater than, or equal to,  RHS.
50    NatGreaterEq(NaturalExpr<V>, NaturalExpr<V>),
51    /// Inequality of Integer expressions: LHS greater than, or equal to,  RHS.
52    IntGreaterEq(IntegerExpr<V>, IntegerExpr<V>),
53    /// Inequality of Float expressions: LHS greater than, or equal to,  RHS.
54    FloatGreaterEq(FloatExpr<V>, FloatExpr<V>),
55    /// Inequality of Natural expressions: LHS less than RHS.
56    NatLess(NaturalExpr<V>, NaturalExpr<V>),
57    /// Inequality of Integer expressions: LHS less than RHS.
58    IntLess(IntegerExpr<V>, IntegerExpr<V>),
59    /// Inequality of Float expressions: LHS less than RHS.
60    FloatLess(FloatExpr<V>, FloatExpr<V>),
61    /// Inequality of Natural expressions: LHS less than, or equal to, RHS.
62    NatLessEq(NaturalExpr<V>, NaturalExpr<V>),
63    /// Inequality of Integer expressions: LHS less than, or equal to, RHS.
64    IntLessEq(IntegerExpr<V>, IntegerExpr<V>),
65    /// Inequality of Float expressions: LHS less than, or equal to, RHS.
66    FloatLessEq(FloatExpr<V>, FloatExpr<V>),
67    // -----
68    // Flow
69    // -----
70    /// If-Then-Else construct, where If must be a boolean expression,
71    /// Then and Else must have the same type,
72    /// and this is also the type of the whole expression.
73    Ite(Box<(BooleanExpr<V>, BooleanExpr<V>, BooleanExpr<V>)>),
74}
75
76impl<V> BooleanExpr<V>
77where
78    V: Copy,
79{
80    /// Returns `true` if the expression is constant, i.e., it contains no variables, and `false` otherwise.
81    pub fn is_constant(&self) -> bool {
82        match self {
83            BooleanExpr::Const(_) => true,
84            BooleanExpr::Var(_) => false,
85            BooleanExpr::Rand(_float_expr) => false,
86            BooleanExpr::And(boolean_exprs) | BooleanExpr::Or(boolean_exprs) => {
87                boolean_exprs.iter().all(Self::is_constant)
88            }
89            BooleanExpr::Implies(args) => {
90                let (lhs, rhs) = args.as_ref();
91                lhs.is_constant() && rhs.is_constant()
92            }
93            BooleanExpr::Not(boolean_expr) => boolean_expr.is_constant(),
94            BooleanExpr::NatEqual(natural_expr_lhs, natural_expr_rhs)
95            | BooleanExpr::NatGreater(natural_expr_lhs, natural_expr_rhs)
96            | BooleanExpr::NatGreaterEq(natural_expr_lhs, natural_expr_rhs)
97            | BooleanExpr::NatLess(natural_expr_lhs, natural_expr_rhs)
98            | BooleanExpr::NatLessEq(natural_expr_lhs, natural_expr_rhs) => {
99                natural_expr_lhs.is_constant() && natural_expr_rhs.is_constant()
100            }
101            BooleanExpr::IntEqual(integer_expr, integer_expr1)
102            | BooleanExpr::IntGreater(integer_expr, integer_expr1)
103            | BooleanExpr::IntGreaterEq(integer_expr, integer_expr1)
104            | BooleanExpr::IntLess(integer_expr, integer_expr1)
105            | BooleanExpr::IntLessEq(integer_expr, integer_expr1) => {
106                integer_expr.is_constant() && integer_expr1.is_constant()
107            }
108            BooleanExpr::FloatEqual(float_expr, float_expr1)
109            | BooleanExpr::FloatLess(float_expr, float_expr1)
110            | BooleanExpr::FloatLessEq(float_expr, float_expr1)
111            | BooleanExpr::FloatGreater(float_expr, float_expr1)
112            | BooleanExpr::FloatGreaterEq(float_expr, float_expr1) => {
113                float_expr.is_constant() && float_expr1.is_constant()
114            }
115            BooleanExpr::Ite(args) => {
116                let (ite, lhs, rhs) = args.as_ref();
117                ite.is_constant() && lhs.is_constant() && rhs.is_constant()
118            }
119        }
120    }
121
122    /// Returns the Boolean value computed from the expression,
123    /// given the variable evaluation.
124    /// It panics if the evaluation is not possible, including:
125    ///
126    /// - If a variable is not included in the evaluation;
127    /// - If a variable included in the evaluation is not of Boolean type.
128    pub fn eval<R: Rng>(&self, vars: &dyn Fn(V) -> Val, mut rng: Option<&mut R>) -> bool {
129        match self {
130            BooleanExpr::Const(b) => *b,
131            BooleanExpr::Var(var) => {
132                if let Val::Boolean(b) = vars(*var) {
133                    b
134                } else {
135                    panic!("type mismatch: expected boolean variable")
136                }
137            }
138            BooleanExpr::Rand(float_expr) => {
139                let bernoulli = float_expr.eval(vars, rng.as_deref_mut());
140                rng.as_mut().expect("rng").random_bool(bernoulli)
141            }
142            BooleanExpr::And(boolean_exprs) => boolean_exprs
143                .iter()
144                .all(|boolean_expr| boolean_expr.eval(vars, rng.as_deref_mut())),
145            BooleanExpr::Or(boolean_exprs) => boolean_exprs
146                .iter()
147                .any(|boolean_expr| boolean_expr.eval(vars, rng.as_deref_mut())),
148            BooleanExpr::Implies(boolean_exprs) => {
149                let (lhs, rhs) = boolean_exprs.as_ref();
150                rhs.eval(vars, rng.as_deref_mut()) || !lhs.eval(vars, rng)
151            }
152            BooleanExpr::Not(boolean_expr) => !&boolean_expr.eval(vars, rng),
153            BooleanExpr::NatEqual(natural_expr_lhs, natural_expr_rhs) => {
154                natural_expr_lhs.eval(vars, rng.as_deref_mut()) == natural_expr_rhs.eval(vars, rng)
155            }
156            BooleanExpr::IntEqual(integer_expr_lhs, integer_expr_rhs) => {
157                integer_expr_lhs.eval(vars, rng.as_deref_mut()) == integer_expr_rhs.eval(vars, rng)
158            }
159            BooleanExpr::FloatEqual(float_expr_lhs, float_expr_rhs) => {
160                float_expr_lhs.eval(vars, rng.as_deref_mut()) == float_expr_rhs.eval(vars, rng)
161            }
162            BooleanExpr::NatGreater(natural_expr_lhs, natural_expr_rhs) => {
163                natural_expr_lhs.eval(vars, rng.as_deref_mut()) > natural_expr_rhs.eval(vars, rng)
164            }
165            BooleanExpr::IntGreater(integer_expr_lhs, integer_expr_rhs) => {
166                integer_expr_lhs.eval(vars, rng.as_deref_mut()) > integer_expr_rhs.eval(vars, rng)
167            }
168            BooleanExpr::FloatGreater(float_expr_lhs, float_expr_rhs) => {
169                float_expr_lhs.eval(vars, rng.as_deref_mut()) > float_expr_rhs.eval(vars, rng)
170            }
171            BooleanExpr::NatGreaterEq(natural_expr_lhs, natural_expr_rhs) => {
172                natural_expr_lhs.eval(vars, rng.as_deref_mut()) >= natural_expr_rhs.eval(vars, rng)
173            }
174            BooleanExpr::IntGreaterEq(integer_expr_lhs, integer_expr_rhs) => {
175                integer_expr_lhs.eval(vars, rng.as_deref_mut()) >= integer_expr_rhs.eval(vars, rng)
176            }
177            BooleanExpr::FloatGreaterEq(float_expr_lhs, float_expr_rhs) => {
178                float_expr_lhs.eval(vars, rng.as_deref_mut()) >= float_expr_rhs.eval(vars, rng)
179            }
180            BooleanExpr::NatLess(natural_expr_lhs, natural_expr_rhs) => {
181                natural_expr_lhs.eval(vars, rng.as_deref_mut()) < natural_expr_rhs.eval(vars, rng)
182            }
183            BooleanExpr::IntLess(integer_expr_lhs, integer_expr_rhs) => {
184                integer_expr_lhs.eval(vars, rng.as_deref_mut()) < integer_expr_rhs.eval(vars, rng)
185            }
186            BooleanExpr::FloatLess(float_expr_lhs, float_expr_rhs) => {
187                float_expr_lhs.eval(vars, rng.as_deref_mut()) < float_expr_rhs.eval(vars, rng)
188            }
189            BooleanExpr::NatLessEq(natural_expr_lhs, natural_expr_rhs) => {
190                natural_expr_lhs.eval(vars, rng.as_deref_mut()) <= natural_expr_rhs.eval(vars, rng)
191            }
192            BooleanExpr::IntLessEq(integer_expr_lhs, integer_expr_rhs) => {
193                integer_expr_lhs.eval(vars, rng.as_deref_mut()) <= integer_expr_rhs.eval(vars, rng)
194            }
195            BooleanExpr::FloatLessEq(float_expr_lhs, float_expr_rhs) => {
196                float_expr_lhs.eval(vars, rng.as_deref_mut()) <= float_expr_rhs.eval(vars, rng)
197            }
198            BooleanExpr::Ite(args) => {
199                let (ite, lhs, rhs) = args.as_ref();
200                if ite.eval(vars, rng.as_deref_mut()) {
201                    lhs.eval(vars, rng)
202                } else {
203                    rhs.eval(vars, rng)
204                }
205            }
206        }
207    }
208
209    pub(crate) fn map<W: Clone>(self, map: &dyn Fn(V) -> W) -> BooleanExpr<W> {
210        match self {
211            BooleanExpr::Const(b) => BooleanExpr::Const(b),
212            BooleanExpr::Var(var) => BooleanExpr::Var(map(var)),
213            BooleanExpr::Rand(float_expr) => BooleanExpr::Rand(float_expr.map(map)),
214            BooleanExpr::And(boolean_exprs) => BooleanExpr::And(
215                boolean_exprs
216                    .into_iter()
217                    .map(|expr| expr.map(map))
218                    .collect(),
219            ),
220            BooleanExpr::Or(boolean_exprs) => BooleanExpr::Or(
221                boolean_exprs
222                    .into_iter()
223                    .map(|expr| expr.map(map))
224                    .collect(),
225            ),
226            BooleanExpr::Implies(args) => {
227                let (lhs, rhs) = *args;
228                BooleanExpr::Implies(Box::new((lhs.map(map), rhs.map(map))))
229            }
230            BooleanExpr::Not(boolean_expr) => BooleanExpr::Not(Box::new(boolean_expr.map(map))),
231            BooleanExpr::NatEqual(natural_expr_lhs, natural_expr_rhs) => {
232                BooleanExpr::NatEqual(natural_expr_lhs.map(map), natural_expr_rhs.map(map))
233            }
234            BooleanExpr::IntEqual(integer_expr_lhs, integer_expr_rhs) => {
235                BooleanExpr::IntEqual(integer_expr_lhs.map(map), integer_expr_rhs.map(map))
236            }
237            BooleanExpr::FloatEqual(float_expr_lhs, float_expr_rhs) => {
238                BooleanExpr::FloatEqual(float_expr_lhs.map(map), float_expr_rhs.map(map))
239            }
240            BooleanExpr::NatGreater(natural_expr_lhs, natural_expr_rhs) => {
241                BooleanExpr::NatGreater(natural_expr_lhs.map(map), natural_expr_rhs.map(map))
242            }
243            BooleanExpr::IntGreater(integer_expr_lhs, integer_expr_rhs) => {
244                BooleanExpr::IntGreater(integer_expr_lhs.map(map), integer_expr_rhs.map(map))
245            }
246            BooleanExpr::FloatGreater(float_expr_lhs, float_expr_rhs) => {
247                BooleanExpr::FloatGreater(float_expr_lhs.map(map), float_expr_rhs.map(map))
248            }
249            BooleanExpr::NatGreaterEq(natural_expr_lhs, natural_expr_rhs) => {
250                BooleanExpr::NatGreaterEq(natural_expr_lhs.map(map), natural_expr_rhs.map(map))
251            }
252            BooleanExpr::IntGreaterEq(integer_expr_lhs, integer_expr_rhs) => {
253                BooleanExpr::IntGreaterEq(integer_expr_lhs.map(map), integer_expr_rhs.map(map))
254            }
255            BooleanExpr::FloatGreaterEq(float_expr_lhs, float_expr_rhs) => {
256                BooleanExpr::FloatGreaterEq(float_expr_lhs.map(map), float_expr_rhs.map(map))
257            }
258            BooleanExpr::NatLess(natural_expr_lhs, natural_expr_rhs) => {
259                BooleanExpr::NatLess(natural_expr_lhs.map(map), natural_expr_rhs.map(map))
260            }
261            BooleanExpr::IntLess(integer_expr_lhs, integer_expr_rhs) => {
262                BooleanExpr::IntLess(integer_expr_lhs.map(map), integer_expr_rhs.map(map))
263            }
264            BooleanExpr::FloatLess(float_expr_lhs, float_expr_rhs) => {
265                BooleanExpr::FloatLess(float_expr_lhs.map(map), float_expr_rhs.map(map))
266            }
267            BooleanExpr::NatLessEq(natural_expr_lhs, natural_expr_rhs) => {
268                BooleanExpr::NatLessEq(natural_expr_lhs.map(map), natural_expr_rhs.map(map))
269            }
270            BooleanExpr::IntLessEq(integer_expr_lhs, integer_expr_rhs) => {
271                BooleanExpr::IntLessEq(integer_expr_lhs.map(map), integer_expr_rhs.map(map))
272            }
273            BooleanExpr::FloatLessEq(float_expr_lhs, float_expr_rhs) => {
274                BooleanExpr::FloatLessEq(float_expr_lhs.map(map), float_expr_rhs.map(map))
275            }
276            BooleanExpr::Ite(args) => {
277                let (r#if, then, r#else) = *args;
278                BooleanExpr::Ite(Box::new((r#if.map(map), then.map(map), r#else.map(map))))
279            }
280        }
281    }
282
283    pub(crate) fn context(&self, vars: &dyn Fn(V) -> Option<Type>) -> Result<(), TypeError> {
284        match self {
285            BooleanExpr::Const(_) => Ok(()),
286            BooleanExpr::Var(v) => matches!(vars(*v), Some(Type::Boolean))
287                .then_some(())
288                .ok_or(TypeError::TypeMismatch),
289            BooleanExpr::Rand(float_expr) => float_expr.context(vars),
290            BooleanExpr::And(boolean_exprs) | BooleanExpr::Or(boolean_exprs) => {
291                boolean_exprs.iter().try_for_each(|expr| expr.context(vars))
292            }
293            BooleanExpr::Implies(exprs) => {
294                exprs.0.context(vars).and_then(|()| exprs.1.context(vars))
295            }
296            BooleanExpr::Not(boolean_expr) => boolean_expr.context(vars),
297            BooleanExpr::NatEqual(natural_expr_lhs, natural_expr_rhs)
298            | BooleanExpr::NatGreater(natural_expr_lhs, natural_expr_rhs)
299            | BooleanExpr::NatGreaterEq(natural_expr_lhs, natural_expr_rhs)
300            | BooleanExpr::NatLess(natural_expr_lhs, natural_expr_rhs)
301            | BooleanExpr::NatLessEq(natural_expr_lhs, natural_expr_rhs) => natural_expr_lhs
302                .context(vars)
303                .and_then(|()| natural_expr_rhs.context(vars)),
304            BooleanExpr::IntEqual(integer_expr_lhs, integer_expr_rhs)
305            | BooleanExpr::IntGreater(integer_expr_lhs, integer_expr_rhs)
306            | BooleanExpr::IntGreaterEq(integer_expr_lhs, integer_expr_rhs)
307            | BooleanExpr::IntLess(integer_expr_lhs, integer_expr_rhs)
308            | BooleanExpr::IntLessEq(integer_expr_lhs, integer_expr_rhs) => integer_expr_lhs
309                .context(vars)
310                .and_then(|()| integer_expr_rhs.context(vars)),
311            BooleanExpr::FloatGreater(float_expr_lhs, float_expr_rhs)
312            | BooleanExpr::FloatLess(float_expr_lhs, float_expr_rhs)
313            | BooleanExpr::FloatEqual(float_expr_lhs, float_expr_rhs)
314            | BooleanExpr::FloatGreaterEq(float_expr_lhs, float_expr_rhs)
315            | BooleanExpr::FloatLessEq(float_expr_lhs, float_expr_rhs) => float_expr_lhs
316                .context(vars)
317                .and_then(|()| float_expr_rhs.context(vars)),
318            BooleanExpr::Ite(exprs) => exprs
319                .0
320                .context(vars)
321                .and_then(|()| exprs.1.context(vars))
322                .and_then(|()| exprs.2.context(vars)),
323        }
324    }
325
326    /// Creates an if-then-else expression
327    pub fn ite(
328        self,
329        then: Expression<V>,
330        r#else: Expression<V>,
331    ) -> Result<Expression<V>, TypeError> {
332        match then {
333            Expression::Boolean(if_boolean_expr) => {
334                if let Expression::Boolean(else_boolean_expr) = r#else {
335                    Ok(Expression::Boolean(BooleanExpr::Ite(Box::new((
336                        self,
337                        if_boolean_expr,
338                        else_boolean_expr,
339                    )))))
340                } else {
341                    Err(TypeError::TypeMismatch)
342                }
343            }
344            Expression::Natural(if_natural_expr) => match r#else {
345                Expression::Boolean(_) => Err(TypeError::TypeMismatch),
346                Expression::Natural(else_natural_expr) => Ok(Expression::Natural(
347                    NaturalExpr::Ite(Box::new((self, if_natural_expr, else_natural_expr))),
348                )),
349                Expression::Integer(else_integer_expr) => {
350                    Ok(Expression::Integer(IntegerExpr::Ite(Box::new((
351                        self,
352                        IntegerExpr::from(if_natural_expr),
353                        else_integer_expr,
354                    )))))
355                }
356                Expression::Float(else_float_expr) => Ok(Expression::Float(FloatExpr::Ite(
357                    Box::new((self, FloatExpr::from(if_natural_expr), else_float_expr)),
358                ))),
359            },
360            Expression::Integer(if_integer_expr) => match r#else {
361                Expression::Boolean(_) => Err(TypeError::TypeMismatch),
362                Expression::Natural(else_natural_expr) => {
363                    Ok(Expression::Integer(IntegerExpr::Ite(Box::new((
364                        self,
365                        if_integer_expr,
366                        IntegerExpr::from(else_natural_expr),
367                    )))))
368                }
369                Expression::Integer(else_integer_expr) => Ok(Expression::Integer(
370                    IntegerExpr::Ite(Box::new((self, if_integer_expr, else_integer_expr))),
371                )),
372                Expression::Float(else_float_expr) => Ok(Expression::Float(FloatExpr::Ite(
373                    Box::new((self, FloatExpr::from(if_integer_expr), else_float_expr)),
374                ))),
375            },
376            Expression::Float(if_float_expr) => match r#else {
377                Expression::Boolean(_) => Err(TypeError::TypeMismatch),
378                Expression::Natural(else_natural_expr) => Ok(Expression::Float(FloatExpr::Ite(
379                    Box::new((self, if_float_expr, FloatExpr::from(else_natural_expr))),
380                ))),
381                Expression::Integer(else_integer_expr) => Ok(Expression::Float(FloatExpr::Ite(
382                    Box::new((self, if_float_expr, FloatExpr::from(else_integer_expr))),
383                ))),
384                Expression::Float(else_float_expr) => Ok(Expression::Float(FloatExpr::Ite(
385                    Box::new((self, if_float_expr, else_float_expr)),
386                ))),
387            },
388        }
389    }
390}
391
392impl<V> From<bool> for BooleanExpr<V>
393where
394    V: Clone,
395{
396    fn from(value: bool) -> Self {
397        Self::Const(value)
398    }
399}
400
401impl<V> TryFrom<Expression<V>> for BooleanExpr<V>
402where
403    V: Clone,
404{
405    type Error = TypeError;
406
407    fn try_from(value: Expression<V>) -> Result<Self, Self::Error> {
408        if let Expression::Boolean(bool_expr) = value {
409            Ok(bool_expr)
410        } else {
411            Err(TypeError::TypeMismatch)
412        }
413    }
414}
415
416impl<V> Not for BooleanExpr<V>
417where
418    V: Clone,
419{
420    type Output = Self;
421
422    fn not(self) -> Self::Output {
423        if let Self::Not(expr) = self {
424            *expr
425        } else {
426            Self::Not(Box::new(self))
427        }
428    }
429}
430
431impl<V> BitAnd for BooleanExpr<V>
432where
433    V: Clone,
434{
435    type Output = Self;
436
437    fn bitand(mut self, mut rhs: Self) -> Self::Output {
438        if let BooleanExpr::And(ref mut exprs) = self {
439            if let BooleanExpr::And(rhs_exprs) = rhs {
440                exprs.extend(rhs_exprs);
441            } else {
442                exprs.push(rhs);
443            }
444            self
445        } else if let BooleanExpr::And(ref mut rhs_exprs) = rhs {
446            rhs_exprs.push(self);
447            rhs
448        } else {
449            BooleanExpr::And(vec![self, rhs])
450        }
451    }
452}
453
454impl<V> BitOr for BooleanExpr<V>
455where
456    V: Clone,
457{
458    type Output = Self;
459
460    fn bitor(mut self, mut rhs: Self) -> Self::Output {
461        if let BooleanExpr::And(ref mut exprs) = self {
462            if let BooleanExpr::Or(rhs_exprs) = rhs {
463                exprs.extend(rhs_exprs);
464            } else {
465                exprs.push(rhs);
466            }
467            self
468        } else if let BooleanExpr::Or(ref mut rhs_exprs) = rhs {
469            rhs_exprs.push(self);
470            rhs
471        } else {
472            BooleanExpr::Or(vec![self, rhs])
473        }
474    }
475}