use crate::eval;
use crate::solver::{CheckResult, Solver};
use crate::term::{BoolTerm, BvTerm, Sort};
fn bx(t: BvTerm) -> Box<BvTerm> {
Box::new(t)
}
fn bb(t: BoolTerm) -> Box<BoolTerm> {
Box::new(t)
}
fn bool_true() -> BoolTerm {
let z = || {
bx(BvTerm::Const {
value: 0,
sort: Sort::new(8),
})
};
BoolTerm::Eq(z(), z())
}
fn bool_false() -> BoolTerm {
BoolTerm::Not(bb(bool_true()))
}
fn zero_like(t: &BvTerm) -> BvTerm {
let width = eval::bv_sort(t).map(|s| s.width).unwrap_or(8);
BvTerm::Const {
value: 0,
sort: Sort::new(width),
}
}
#[derive(Clone, Debug)]
pub struct DefineOrTrap {
pub value: BvTerm,
pub may_trap: BoolTerm,
}
#[derive(Clone, Copy, Debug)]
pub enum DivOp {
DivU,
DivS,
RemU,
RemS,
}
impl DivOp {
fn is_signed(self) -> bool {
matches!(self, DivOp::DivS | DivOp::RemS)
}
}
pub fn trap_div(op: DivOp, dividend: &BvTerm, divisor: &BvTerm, width: u32) -> BoolTerm {
let sort = Sort::new(width);
let zero = BvTerm::Const { value: 0, sort };
let div_by_zero = BoolTerm::Eq(bx(divisor.clone()), bx(zero));
if !op.is_signed() {
return div_by_zero;
}
let int_min = 1u128 << (width - 1);
let all_ones = if width >= 128 {
u128::MAX
} else {
(1u128 << width) - 1
};
let overflow = BoolTerm::And(
bb(BoolTerm::Eq(
bx(dividend.clone()),
bx(BvTerm::Const {
value: int_min,
sort,
}),
)),
bb(BoolTerm::Eq(
bx(divisor.clone()),
bx(BvTerm::Const {
value: all_ones,
sort,
}),
)),
);
BoolTerm::Or(bb(div_by_zero), bb(overflow))
}
pub fn trap_always() -> BoolTerm {
bool_true()
}
pub fn trap_mem_oob(addr: &BvTerm, size: &BvTerm, mem_bound: &BvTerm) -> BoolTerm {
let ext = |t: &BvTerm| {
bx(BvTerm::ZeroExt {
by: 1,
arg: bx(t.clone()),
})
};
let end = BvTerm::Add(ext(addr), ext(size));
BoolTerm::Ugt(bx(end), ext(mem_bound))
}
pub enum TypeTrap<'a> {
Runtime {
actual_type_id: &'a BvTerm,
expected_id: &'a BvTerm,
},
StaticallyDischarged,
}
pub struct CallIndirect<'a> {
pub index: &'a BvTerm,
pub table_size: &'a BvTerm,
pub slot_ptr: &'a BvTerm,
pub type_trap: TypeTrap<'a>,
}
pub fn trap_call_indirect(ci: &CallIndirect) -> BoolTerm {
let bounds = BoolTerm::Uge(bx(ci.index.clone()), bx(ci.table_size.clone()));
let null_slot = BoolTerm::Eq(bx(ci.slot_ptr.clone()), bx(zero_like(ci.slot_ptr)));
let type_clause = match &ci.type_trap {
TypeTrap::Runtime {
actual_type_id,
expected_id,
} => BoolTerm::Ne(bx((*actual_type_id).clone()), bx((*expected_id).clone())),
TypeTrap::StaticallyDischarged => bool_false(),
};
BoolTerm::Or(bb(BoolTerm::Or(bb(bounds), bb(null_slot))), bb(type_clause))
}
pub fn trap_any(conds: &[BoolTerm]) -> BoolTerm {
match conds.split_first() {
None => bool_false(),
Some((first, rest)) => rest
.iter()
.fold(first.clone(), |acc, c| BoolTerm::Or(bb(acc), bb(c.clone()))),
}
}
fn iff(a: &BoolTerm, b: &BoolTerm) -> BoolTerm {
let imp =
|x: &BoolTerm, y: &BoolTerm| BoolTerm::Or(bb(BoolTerm::Not(bb(x.clone()))), bb(y.clone()));
BoolTerm::And(bb(imp(a, b)), bb(imp(b, a)))
}
pub fn trap_condition_equivalence(orig_may_trap: &BoolTerm, opt_may_trap: &BoolTerm) -> BoolTerm {
iff(orig_may_trap, opt_may_trap)
}
pub fn trap_equivalence_vc(orig: &DefineOrTrap, opt: &DefineOrTrap) -> BoolTerm {
let trap_eq = iff(&orig.may_trap, &opt.may_trap);
let value_eq = BoolTerm::Eq(bx(orig.value.clone()), bx(opt.value.clone()));
let guarded_value = BoolTerm::Or(bb(orig.may_trap.clone()), bb(value_eq));
BoolTerm::And(bb(trap_eq), bb(guarded_value))
}
fn prove_valid(goal: BoolTerm) -> CheckResult {
let mut s = Solver::new();
s.assert(BoolTerm::Not(bb(goal)));
s.check()
}
pub fn prove_trap_equivalence(orig: &DefineOrTrap, opt: &DefineOrTrap) -> CheckResult {
prove_valid(trap_equivalence_vc(orig, opt))
}
pub fn prove_trap_condition_equivalence(
orig_may_trap: &BoolTerm,
opt_may_trap: &BoolTerm,
) -> CheckResult {
prove_valid(trap_condition_equivalence(orig_may_trap, opt_may_trap))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::eval::Env;
fn v(name: &str, w: u32) -> BvTerm {
BvTerm::Var {
name: name.into(),
sort: Sort::new(w),
}
}
fn c(value: u128, w: u32) -> BvTerm {
BvTerm::Const {
value,
sort: Sort::new(w),
}
}
fn env2(a: u128, b: u128) -> Env {
let mut e = Env::new();
e.insert("a".into(), a);
e.insert("b".into(), b);
e
}
#[test]
fn trap_div_matches_wasm_semantics() {
let (a, b) = (v("a", 8), v("b", 8));
for op in [DivOp::DivU, DivOp::DivS, DivOp::RemU, DivOp::RemS] {
let cond = trap_div(op, &a, &b, 8);
for av in 0u128..256 {
for bv in 0u128..256 {
let got = eval::eval_bool(&cond, &env2(av, bv)).unwrap();
let zero = bv == 0;
let overflow = op.is_signed() && av == 0x80 && bv == 0xFF;
assert_eq!(got, zero || overflow, "{op:?} a={av} b={bv}");
}
}
}
}
#[test]
fn trap_always_is_true() {
assert!(eval::eval_bool(&trap_always(), &Env::new()).unwrap());
}
#[test]
fn trap_mem_oob_matches_reference_and_is_wraparound_safe() {
let addr = v("a", 8);
let bound = v("b", 8);
let size = c(4, 8);
let cond = trap_mem_oob(&addr, &size, &bound);
for a in 0u128..256 {
for b in 0u128..256 {
let got = eval::eval_bool(&cond, &env2(a, b)).unwrap();
assert_eq!(got, a + 4 > b, "addr={a} bound={b}");
}
}
assert!(eval::eval_bool(&cond, &env2(254, 255)).unwrap());
}
#[test]
fn trap_call_indirect_covers_bounds_null_and_type() {
let index = v("a", 32);
let table_size = c(10, 32);
let slot = v("b", 32);
let actual = BvTerm::Var {
name: "t".into(),
sort: Sort::new(32),
};
let expected = c(7, 32);
let ci = CallIndirect {
index: &index,
table_size: &table_size,
slot_ptr: &slot,
type_trap: TypeTrap::Runtime {
actual_type_id: &actual,
expected_id: &expected,
},
};
let cond = trap_call_indirect(&ci);
let eval = |idx: u128, slotv: u128, t: u128| {
let mut e = Env::new();
e.insert("a".into(), idx);
e.insert("b".into(), slotv);
e.insert("t".into(), t);
eval::eval_bool(&cond, &e).unwrap()
};
assert!(eval(10, 1, 7), "index == size is out of bounds");
assert!(eval(3, 0, 7), "null slot traps");
assert!(eval(3, 1, 9), "type mismatch traps");
assert!(
!eval(3, 1, 7),
"in-bounds, non-null, matching type: no trap"
);
}
#[test]
fn statically_discharged_type_never_contributes_a_trap() {
let index = v("a", 32);
let table_size = c(10, 32);
let slot = v("b", 32);
let ci = CallIndirect {
index: &index,
table_size: &table_size,
slot_ptr: &slot,
type_trap: TypeTrap::StaticallyDischarged,
};
let cond = trap_call_indirect(&ci);
let mut e = Env::new();
e.insert("a".into(), 3);
e.insert("b".into(), 1);
assert!(!eval::eval_bool(&cond, &e).unwrap());
}
#[test]
fn trap_any_is_the_or_fold() {
assert!(!eval::eval_bool(&trap_any(&[]), &Env::new()).unwrap());
let a_zero = BoolTerm::Eq(Box::new(v("a", 8)), Box::new(c(0, 8)));
let b_zero = BoolTerm::Eq(Box::new(v("b", 8)), Box::new(c(0, 8)));
let any = trap_any(&[a_zero, b_zero]);
assert!(eval::eval_bool(&any, &env2(0, 5)).unwrap());
assert!(eval::eval_bool(&any, &env2(5, 0)).unwrap());
assert!(!eval::eval_bool(&any, &env2(5, 5)).unwrap());
}
#[test]
fn preserved_div_lowering_proves_unsat() {
let (a, b) = (v("a", 8), v("b", 8));
let value = BvTerm::Udiv(Box::new(a.clone()), Box::new(b.clone()));
let d = |val: BvTerm, t: BoolTerm| DefineOrTrap {
value: val,
may_trap: t,
};
let orig = d(value.clone(), trap_div(DivOp::DivU, &a, &b, 8));
let opt = d(value, trap_div(DivOp::DivU, &a, &b, 8));
match prove_trap_equivalence(&orig, &opt) {
CheckResult::Unsat(cert) => cert.recheck().expect("trap-equiv cert re-checks"),
other => panic!("preserved lowering must be Unsat, got {other:?}"),
}
}
#[test]
fn dropped_trap_is_caught_with_counterexample() {
let (a, b) = (v("a", 8), v("b", 8));
let value = BvTerm::Udiv(Box::new(a.clone()), Box::new(b.clone()));
let orig = DefineOrTrap {
value: value.clone(),
may_trap: trap_div(DivOp::DivU, &a, &b, 8),
};
let opt = DefineOrTrap {
value,
may_trap: bool_false(),
};
match prove_trap_equivalence(&orig, &opt) {
CheckResult::Sat(m) => {
let b_val = m
.assignments
.iter()
.find(|(n, _)| n == "b")
.map(|(_, x)| *x);
assert_eq!(b_val, Some(0), "counterexample must set divisor to 0");
}
other => panic!("dropped trap must be Sat, got {other:?}"),
}
}
#[test]
fn dropped_bounds_check_caught_by_conjunct1_gate() {
let addr = v("a", 8);
let bound = v("b", 8);
let size = c(4, 8);
let orig_trap = trap_mem_oob(&addr, &size, &bound);
let opt_trap = bool_false(); match prove_trap_condition_equivalence(&orig_trap, &opt_trap) {
CheckResult::Sat(_) => {}
other => panic!("dropped bounds check must be Sat, got {other:?}"),
}
match prove_trap_condition_equivalence(&orig_trap, &orig_trap) {
CheckResult::Unsat(cert) => cert.recheck().expect("re-check"),
other => panic!("preserved bounds check must be Unsat, got {other:?}"),
}
}
}