use crate::term::{BoolTerm, BvTerm, Sort};
use std::collections::HashMap;
pub type Env = HashMap<String, u128>;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum EvalError {
UnboundVar(String),
WidthMismatch { left: u32, right: u32 },
BadExtract { hi: u32, lo: u32, width: u32 },
UnsupportedWidth(u32),
}
fn mask(width: u32) -> u128 {
if width >= 128 {
u128::MAX
} else {
(1u128 << width) - 1
}
}
fn check_width(w: u32) -> Result<u32, EvalError> {
if w == 0 || w > 128 {
Err(EvalError::UnsupportedWidth(w))
} else {
Ok(w)
}
}
fn same_width(a: u32, b: u32) -> Result<u32, EvalError> {
if a == b {
Ok(a)
} else {
Err(EvalError::WidthMismatch { left: a, right: b })
}
}
pub fn bv_sort(term: &BvTerm) -> Result<Sort, EvalError> {
let w = match term {
BvTerm::Const { sort, .. } | BvTerm::Var { sort, .. } => check_width(sort.width)?,
BvTerm::Add(a, b)
| BvTerm::Sub(a, b)
| BvTerm::Mul(a, b)
| BvTerm::Udiv(a, b)
| BvTerm::Urem(a, b)
| BvTerm::And(a, b)
| BvTerm::Or(a, b)
| BvTerm::Xor(a, b)
| BvTerm::Shl(a, b)
| BvTerm::Lshr(a, b)
| BvTerm::Ashr(a, b)
| BvTerm::Rotr(a, b) => same_width(bv_sort(a)?.width, bv_sort(b)?.width)?,
BvTerm::Extract { hi, lo, arg } => {
let w = bv_sort(arg)?.width;
if hi < lo || *hi >= w {
return Err(EvalError::BadExtract {
hi: *hi,
lo: *lo,
width: w,
});
}
hi - lo + 1
}
BvTerm::Concat(a, b) => check_width(bv_sort(a)?.width + bv_sort(b)?.width)?,
BvTerm::ZeroExt { by, arg } | BvTerm::SignExt { by, arg } => {
check_width(bv_sort(arg)?.width + by)?
}
BvTerm::Ite { then_, else_, .. } => {
same_width(bv_sort(then_)?.width, bv_sort(else_)?.width)?
}
};
Ok(Sort::new(w))
}
pub fn eval_bv(term: &BvTerm, env: &Env) -> Result<u128, EvalError> {
let w = bv_sort(term)?.width;
let m = mask(w);
let v = match term {
BvTerm::Const { value, .. } => value & m,
BvTerm::Var { name, .. } => {
*env.get(name)
.ok_or_else(|| EvalError::UnboundVar(name.clone()))?
& m
}
BvTerm::Add(a, b) => eval_bv(a, env)?.wrapping_add(eval_bv(b, env)?) & m,
BvTerm::Sub(a, b) => eval_bv(a, env)?.wrapping_sub(eval_bv(b, env)?) & m,
BvTerm::Mul(a, b) => eval_bv(a, env)?.wrapping_mul(eval_bv(b, env)?) & m,
BvTerm::Udiv(a, b) => {
let (x, y) = (eval_bv(a, env)?, eval_bv(b, env)?);
x.checked_div(y).unwrap_or(m)
}
BvTerm::Urem(a, b) => {
let (x, y) = (eval_bv(a, env)?, eval_bv(b, env)?);
if y == 0 { x } else { x % y }
}
BvTerm::And(a, b) => eval_bv(a, env)? & eval_bv(b, env)?,
BvTerm::Or(a, b) => eval_bv(a, env)? | eval_bv(b, env)?,
BvTerm::Xor(a, b) => eval_bv(a, env)? ^ eval_bv(b, env)?,
BvTerm::Shl(a, b) => {
let (x, sh) = (eval_bv(a, env)?, eval_bv(b, env)?);
if sh >= w as u128 { 0 } else { (x << sh) & m }
}
BvTerm::Lshr(a, b) => {
let (x, sh) = (eval_bv(a, env)?, eval_bv(b, env)?);
if sh >= w as u128 { 0 } else { x >> sh }
}
BvTerm::Ashr(a, b) => {
let (x, sh) = (eval_bv(a, env)?, eval_bv(b, env)?);
let sign = (x >> (w - 1)) & 1 == 1;
if sh >= w as u128 {
if sign { m } else { 0 }
} else if sign {
((x >> sh) | (m & !(m >> sh))) & m
} else {
x >> sh
}
}
BvTerm::Rotr(a, b) => {
let (x, sh) = (eval_bv(a, env)?, eval_bv(b, env)?);
let r = (sh % w as u128) as u32;
if r == 0 {
x
} else {
((x >> r) | (x << (w - r))) & m
}
}
BvTerm::Extract { lo, arg, .. } => (eval_bv(arg, env)? >> lo) & m,
BvTerm::Concat(a, b) => {
let wb = bv_sort(b)?.width;
((eval_bv(a, env)? << wb) | eval_bv(b, env)?) & m
}
BvTerm::Ite { cond, then_, else_ } => {
if eval_bool(cond, env)? {
eval_bv(then_, env)?
} else {
eval_bv(else_, env)?
}
}
BvTerm::ZeroExt { arg, .. } => eval_bv(arg, env)?,
BvTerm::SignExt { arg, .. } => {
let wa = bv_sort(arg)?.width;
let x = eval_bv(arg, env)?;
if (x >> (wa - 1)) & 1 == 1 {
(x | (m ^ mask(wa))) & m
} else {
x
}
}
};
Ok(v)
}
fn to_signed(v: u128, w: u32) -> i128 {
if w < 128 && (v >> (w - 1)) & 1 == 1 {
(v | !mask(w)) as i128
} else {
v as i128
}
}
pub fn eval_bool(term: &BoolTerm, env: &Env) -> Result<bool, EvalError> {
let cmp_w = |a: &BvTerm, b: &BvTerm| -> Result<(u128, u128, u32), EvalError> {
let w = same_width(bv_sort(a)?.width, bv_sort(b)?.width)?;
Ok((eval_bv(a, env)?, eval_bv(b, env)?, w))
};
Ok(match term {
BoolTerm::Eq(a, b) => cmp_w(a, b).map(|(x, y, _)| x == y)?,
BoolTerm::Ne(a, b) => cmp_w(a, b).map(|(x, y, _)| x != y)?,
BoolTerm::Ult(a, b) => cmp_w(a, b).map(|(x, y, _)| x < y)?,
BoolTerm::Ule(a, b) => cmp_w(a, b).map(|(x, y, _)| x <= y)?,
BoolTerm::Ugt(a, b) => cmp_w(a, b).map(|(x, y, _)| x > y)?,
BoolTerm::Uge(a, b) => cmp_w(a, b).map(|(x, y, _)| x >= y)?,
BoolTerm::Slt(a, b) => cmp_w(a, b).map(|(x, y, w)| to_signed(x, w) < to_signed(y, w))?,
BoolTerm::Sle(a, b) => cmp_w(a, b).map(|(x, y, w)| to_signed(x, w) <= to_signed(y, w))?,
BoolTerm::Sgt(a, b) => cmp_w(a, b).map(|(x, y, w)| to_signed(x, w) > to_signed(y, w))?,
BoolTerm::Sge(a, b) => cmp_w(a, b).map(|(x, y, w)| to_signed(x, w) >= to_signed(y, w))?,
BoolTerm::Not(t) => !eval_bool(t, env)?,
BoolTerm::And(a, b) => eval_bool(a, env)? && eval_bool(b, env)?,
BoolTerm::Or(a, b) => eval_bool(a, env)? || eval_bool(b, env)?,
})
}
#[cfg(test)]
mod tests {
use super::*;
fn c(value: u128, w: u32) -> BvTerm {
BvTerm::Const {
value,
sort: Sort::new(w),
}
}
fn ev(t: &BvTerm) -> u128 {
eval_bv(t, &Env::new()).unwrap()
}
fn evb(t: &BoolTerm) -> bool {
eval_bool(t, &Env::new()).unwrap()
}
fn b(t: BvTerm) -> Box<BvTerm> {
Box::new(t)
}
#[test]
fn udiv_by_zero_is_all_ones() {
for w in [8u32, 32, 64] {
let q = BvTerm::Udiv(b(c(42, w)), b(c(0, w)));
assert_eq!(ev(&q), mask(w), "width {w}");
}
assert_eq!(ev(&BvTerm::Udiv(b(c(0, 8)), b(c(0, 8)))), 0xFF);
}
#[test]
fn shifts_out_of_range_saturate() {
assert_eq!(ev(&BvTerm::Shl(b(c(0xAB, 8)), b(c(8, 8)))), 0);
assert_eq!(ev(&BvTerm::Shl(b(c(0xAB, 8)), b(c(200, 8)))), 0);
assert_eq!(ev(&BvTerm::Lshr(b(c(0xAB, 8)), b(c(9, 8)))), 0);
assert_eq!(ev(&BvTerm::Ashr(b(c(0x80, 8)), b(c(8, 8)))), 0xFF);
assert_eq!(ev(&BvTerm::Ashr(b(c(0x7F, 8)), b(c(8, 8)))), 0x00);
}
#[test]
fn ashr_in_range_sign_fills() {
assert_eq!(ev(&BvTerm::Ashr(b(c(0x80, 8)), b(c(1, 8)))), 0xC0);
assert_eq!(ev(&BvTerm::Ashr(b(c(0x80, 8)), b(c(7, 8)))), 0xFF);
assert_eq!(ev(&BvTerm::Ashr(b(c(0x40, 8)), b(c(1, 8)))), 0x20);
let v = 0x8000_0000_0000_0000u128;
assert_eq!(
ev(&BvTerm::Ashr(b(c(v, 64)), b(c(4, 64)))),
0xF800_0000_0000_0000
);
}
#[test]
fn rotr_wraps_modulo_width() {
assert_eq!(ev(&BvTerm::Rotr(b(c(0b0000_0001, 8)), b(c(1, 8)))), 0x80);
assert_eq!(ev(&BvTerm::Rotr(b(c(0xAB, 8)), b(c(8, 8)))), 0xAB);
assert_eq!(ev(&BvTerm::Rotr(b(c(0xAB, 8)), b(c(0, 8)))), 0xAB);
assert_eq!(
ev(&BvTerm::Rotr(b(c(0xAB, 8)), b(c(9, 8)))),
ev(&BvTerm::Rotr(b(c(0xAB, 8)), b(c(1, 8))))
);
}
#[test]
fn structural_ops() {
assert_eq!(
ev(&BvTerm::Extract {
hi: 7,
lo: 4,
arg: b(c(0xAB, 8))
}),
0xA
);
let hi4 = BvTerm::Extract {
hi: 7,
lo: 4,
arg: b(c(0xA0, 8)),
};
let lo4 = BvTerm::Extract {
hi: 3,
lo: 0,
arg: b(c(0x0B, 8)),
};
assert_eq!(ev(&BvTerm::Concat(b(hi4), b(lo4))), 0xAB);
assert_eq!(
ev(&BvTerm::ZeroExt {
by: 8,
arg: b(c(0x80, 8))
}),
0x0080
);
assert_eq!(
ev(&BvTerm::SignExt {
by: 8,
arg: b(c(0x80, 8))
}),
0xFF80
);
assert_eq!(
ev(&BvTerm::SignExt {
by: 8,
arg: b(c(0x7F, 8))
}),
0x007F
);
}
#[test]
fn signed_comparisons_on_boundaries() {
let min = c(0x80, 8);
let neg1 = c(0xFF, 8);
let zero = c(0, 8);
let max = c(0x7F, 8);
assert!(evb(&BoolTerm::Slt(b(min.clone()), b(neg1.clone()))));
assert!(evb(&BoolTerm::Slt(b(neg1.clone()), b(zero.clone()))));
assert!(evb(&BoolTerm::Slt(b(zero.clone()), b(max.clone()))));
assert!(evb(&BoolTerm::Ugt(b(min.clone()), b(max.clone()))));
assert!(evb(&BoolTerm::Sge(b(max), b(min))));
assert!(evb(&BoolTerm::Sle(b(neg1), b(zero))));
}
#[test]
fn modular_arithmetic_wraps() {
assert_eq!(ev(&BvTerm::Add(b(c(0xFF, 8)), b(c(1, 8)))), 0);
assert_eq!(ev(&BvTerm::Sub(b(c(0, 8)), b(c(1, 8)))), 0xFF);
assert_eq!(ev(&BvTerm::Mul(b(c(0x80, 8)), b(c(2, 8)))), 0);
assert_eq!(ev(&BvTerm::Add(b(c(u64::MAX as u128, 64)), b(c(1, 64)))), 0);
}
#[test]
fn width64_reference_cross_check() {
let mut s: u64 = 0x2545F4914F6CDD1D;
let mut next = move || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
for _ in 0..1000 {
let (x, y) = (next(), next());
let (tx, ty) = (c(x as u128, 64), c(y as u128, 64));
assert_eq!(
ev(&BvTerm::Add(b(tx.clone()), b(ty.clone()))),
x.wrapping_add(y) as u128
);
assert_eq!(
ev(&BvTerm::Sub(b(tx.clone()), b(ty.clone()))),
x.wrapping_sub(y) as u128
);
assert_eq!(
ev(&BvTerm::Mul(b(tx.clone()), b(ty.clone()))),
x.wrapping_mul(y) as u128
);
let d = x.checked_div(y).unwrap_or(u64::MAX);
assert_eq!(ev(&BvTerm::Udiv(b(tx.clone()), b(ty.clone()))), d as u128);
let sh = y % 67; let shl = if sh >= 64 { 0 } else { x << sh };
assert_eq!(
ev(&BvTerm::Shl(b(tx.clone()), b(c(sh as u128, 64)))),
shl as u128
);
let ashr = if sh >= 64 {
((x as i64) >> 63) as u64
} else {
((x as i64) >> sh) as u64
};
assert_eq!(
ev(&BvTerm::Ashr(b(tx.clone()), b(c(sh as u128, 64)))),
ashr as u128
);
let rr = x.rotate_right((y % 64) as u32);
assert_eq!(
ev(&BvTerm::Rotr(b(tx.clone()), b(c((y % 64) as u128, 64)))),
rr as u128
);
assert_eq!(
evb(&BoolTerm::Slt(b(tx.clone()), b(ty.clone()))),
(x as i64) < (y as i64)
);
assert_eq!(evb(&BoolTerm::Ult(b(tx), b(ty))), x < y);
}
}
#[test]
fn unbound_var_and_width_mismatch_error() {
let x = BvTerm::Var {
name: "x".into(),
sort: Sort::new(8),
};
assert_eq!(
eval_bv(&x, &Env::new()),
Err(EvalError::UnboundVar("x".into()))
);
let bad = BvTerm::Add(b(c(1, 8)), b(c(1, 32)));
assert_eq!(
eval_bv(&bad, &Env::new()),
Err(EvalError::WidthMismatch { left: 8, right: 32 })
);
let bad_ex = BvTerm::Extract {
hi: 8,
lo: 0,
arg: b(c(0, 8)),
};
assert!(matches!(
eval_bv(&bad_ex, &Env::new()),
Err(EvalError::BadExtract { .. })
));
}
}