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#[derive(Debug, Clone, GetSize)]
13pub enum BooleanExpr<V>
14where
15 V: Clone,
16{
17 Const(bool),
19 Var(V),
21 Rand(FloatExpr<V>),
23 And(Vec<BooleanExpr<V>>),
28 Or(Vec<BooleanExpr<V>>),
30 Implies(Box<(BooleanExpr<V>, BooleanExpr<V>)>),
32 Not(Box<BooleanExpr<V>>),
34 NatEqual(NaturalExpr<V>, NaturalExpr<V>),
39 IntEqual(IntegerExpr<V>, IntegerExpr<V>),
41 FloatEqual(FloatExpr<V>, FloatExpr<V>),
43 NatGreater(NaturalExpr<V>, NaturalExpr<V>),
45 IntGreater(IntegerExpr<V>, IntegerExpr<V>),
47 FloatGreater(FloatExpr<V>, FloatExpr<V>),
49 NatGreaterEq(NaturalExpr<V>, NaturalExpr<V>),
51 IntGreaterEq(IntegerExpr<V>, IntegerExpr<V>),
53 FloatGreaterEq(FloatExpr<V>, FloatExpr<V>),
55 NatLess(NaturalExpr<V>, NaturalExpr<V>),
57 IntLess(IntegerExpr<V>, IntegerExpr<V>),
59 FloatLess(FloatExpr<V>, FloatExpr<V>),
61 NatLessEq(NaturalExpr<V>, NaturalExpr<V>),
63 IntLessEq(IntegerExpr<V>, IntegerExpr<V>),
65 FloatLessEq(FloatExpr<V>, FloatExpr<V>),
67 Ite(Box<(BooleanExpr<V>, BooleanExpr<V>, BooleanExpr<V>)>),
74}
75
76impl<V> BooleanExpr<V>
77where
78 V: Copy,
79{
80 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 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 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}