use crate::{
analysis::CfgInfo,
mir::{
BlockId, Function, InstKind, Terminator, Value, ValueId, utils::repair_reachability_phis,
},
pass::FunctionPass,
};
use alloy_primitives::U256;
use solar_data_structures::map::{FxHashMap, FxHashSet};
const MAX_DEPTH: usize = 12;
#[derive(Debug, Default, Clone)]
pub struct CheckElimStats {
pub branches_folded: usize,
}
pub struct CheckElimPass;
impl FunctionPass for CheckElimPass {
fn name(&self) -> &str {
"check-elim"
}
fn run_on_function(&mut self, func: &mut Function) -> bool {
CheckEliminator::new().run(func) != 0
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct Range {
lo: U256,
hi: U256,
}
impl Range {
const FULL: Self = Self { lo: U256::ZERO, hi: U256::MAX };
const fn new(lo: U256, hi: U256) -> Self {
Self { lo, hi }
}
fn singleton(value: U256) -> Self {
Self { lo: value, hi: value }
}
fn is_singleton(self) -> bool {
self.lo == self.hi
}
fn intersect(self, other: Self) -> Option<Self> {
let lo = self.lo.max(other.lo);
let hi = self.hi.min(other.hi);
(lo <= hi).then_some(Self { lo, hi })
}
fn union(self, other: Self) -> Self {
Self { lo: self.lo.min(other.lo), hi: self.hi.max(other.hi) }
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
enum Relation {
Lt(ValueId, ValueId),
Le(ValueId, ValueId),
Eq(ValueId, ValueId),
Ne(ValueId, ValueId),
}
fn ordered(a: ValueId, b: ValueId) -> (ValueId, ValueId) {
if a.index() <= b.index() { (a, b) } else { (b, a) }
}
#[derive(Default)]
pub struct CheckEliminator {
pub stats: CheckElimStats,
ranges: FxHashMap<ValueId, Range>,
relations: FxHashSet<Relation>,
range_undo: Vec<(ValueId, Option<Range>)>,
relation_undo: Vec<Relation>,
}
impl CheckEliminator {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn run(&mut self, func: &mut Function) -> usize {
self.stats = CheckElimStats::default();
if func.blocks.is_empty() {
return 0;
}
let cfg = CfgInfo::new(func);
let mut preds: Vec<Vec<BlockId>> = vec![Vec::new(); func.blocks.len()];
for &block in cfg.rpo() {
for &succ in cfg.successors(block) {
preds[succ.index()].push(block);
}
}
let folds = self.collect_folds(func, &cfg, &preds);
self.ranges.clear();
self.relations.clear();
self.range_undo.clear();
self.relation_undo.clear();
if folds.is_empty() {
return 0;
}
for &(block, keep) in &folds {
func.blocks[block].terminator = Some(Terminator::Jump(keep));
}
repair_reachability_phis(func);
self.stats.branches_folded = folds.len();
folds.len()
}
fn collect_folds(
&mut self,
func: &Function,
cfg: &CfgInfo,
preds: &[Vec<BlockId>],
) -> Vec<(BlockId, BlockId)> {
enum Walk {
Enter(BlockId),
Exit { range_mark: usize, relation_mark: usize },
}
let mut folds = Vec::new();
let mut stack = vec![Walk::Enter(func.entry_block)];
while let Some(item) = stack.pop() {
match item {
Walk::Exit { range_mark, relation_mark } => {
while self.range_undo.len() > range_mark {
let (value, old) = self.range_undo.pop().expect("checked len");
match old {
Some(range) => self.ranges.insert(value, range),
None => self.ranges.remove(&value),
};
}
while self.relation_undo.len() > relation_mark {
let relation = self.relation_undo.pop().expect("checked len");
self.relations.remove(&relation);
}
}
Walk::Enter(block) => {
stack.push(Walk::Exit {
range_mark: self.range_undo.len(),
relation_mark: self.relation_undo.len(),
});
if let Some((condition, is_true)) = dominating_edge_fact(func, preds, block) {
self.assume(func, condition, is_true, MAX_DEPTH);
}
if let Some(Terminator::Branch { condition, then_block, else_block }) =
func.blocks[block].terminator.as_ref()
&& then_block != else_block
&& let Some(truth) = self.eval_truth(func, *condition, MAX_DEPTH)
{
folds.push((block, if truth { *then_block } else { *else_block }));
}
for &child in cfg.dominators().children(block) {
stack.push(Walk::Enter(child));
}
}
}
}
folds
}
fn assume(&mut self, func: &Function, value: ValueId, truth: bool, depth: usize) {
if truth {
self.narrow(value, Range::new(U256::from(1), U256::MAX));
} else {
self.narrow(value, Range::singleton(U256::ZERO));
}
let Some(depth) = depth.checked_sub(1) else { return };
let Some(kind) = inst_kind(func, value) else { return };
match *kind {
InstKind::IsZero(a) => self.assume(func, a, !truth, depth),
InstKind::Lt(a, b) => self.assume_lt(func, a, b, truth, depth),
InstKind::Gt(a, b) => self.assume_lt(func, b, a, truth, depth),
InstKind::Eq(a, b) => self.assume_eq(func, a, b, truth, depth),
InstKind::Sub(a, b) | InstKind::Xor(a, b) => self.assume_eq(func, a, b, !truth, depth),
InstKind::And(a, b) if truth => {
self.assume(func, a, true, depth);
self.assume(func, b, true, depth);
}
InstKind::Or(a, b) if !truth => {
self.assume(func, a, false, depth);
self.assume(func, b, false, depth);
}
_ => {}
}
}
fn assume_lt(&mut self, func: &Function, a: ValueId, b: ValueId, truth: bool, depth: usize) {
if truth {
self.add_relation(Relation::Lt(a, b));
let hi_b = self.range_of(func, b, depth).hi;
if hi_b > U256::ZERO {
self.narrow(a, Range::new(U256::ZERO, hi_b - U256::from(1)));
}
let lo_a = self.range_of(func, a, depth).lo;
if lo_a < U256::MAX {
self.narrow(b, Range::new(lo_a + U256::from(1), U256::MAX));
}
} else {
self.add_relation(Relation::Le(b, a));
let lo_b = self.range_of(func, b, depth).lo;
self.narrow(a, Range::new(lo_b, U256::MAX));
let hi_a = self.range_of(func, a, depth).hi;
self.narrow(b, Range::new(U256::ZERO, hi_a));
}
}
fn assume_eq(&mut self, func: &Function, a: ValueId, b: ValueId, truth: bool, depth: usize) {
let (x, y) = ordered(a, b);
if truth {
self.add_relation(Relation::Eq(x, y));
let range = self.range_of(func, a, depth);
self.narrow(b, range);
let range = self.range_of(func, b, depth);
self.narrow(a, range);
} else {
self.add_relation(Relation::Ne(x, y));
self.exclude_boundary(func, a, b, depth);
self.exclude_boundary(func, b, a, depth);
}
}
fn exclude_boundary(&mut self, func: &Function, a: ValueId, b: ValueId, depth: usize) {
let rb = self.range_of(func, b, depth);
if !rb.is_singleton() {
return;
}
let ra = self.range_of(func, a, depth);
if ra.is_singleton() {
return;
}
if ra.lo == rb.lo {
self.narrow(a, Range::new(ra.lo + U256::from(1), ra.hi));
} else if ra.hi == rb.hi {
self.narrow(a, Range::new(ra.lo, ra.hi - U256::from(1)));
}
}
fn narrow(&mut self, value: ValueId, range: Range) {
let old = self.ranges.get(&value).copied();
let Some(new) = old.unwrap_or(Range::FULL).intersect(range) else { return };
if Some(new) == old {
return;
}
self.range_undo.push((value, old));
self.ranges.insert(value, new);
}
fn add_relation(&mut self, relation: Relation) {
if self.relations.insert(relation) {
self.relation_undo.push(relation);
}
}
fn has_relation(&self, relation: Relation) -> bool {
self.relations.contains(&relation)
}
fn range_of(&mut self, func: &Function, value: ValueId, depth: usize) -> Range {
if let Some(constant) = const_of(func, value) {
return Range::singleton(constant);
}
let mut range = self.ranges.get(&value).copied().unwrap_or(Range::FULL);
let Some(depth) = depth.checked_sub(1) else { return range };
let Some(kind) = inst_kind(func, value) else { return range };
let derived = match *kind {
InstKind::Add(a, b) => {
let ra = self.range_of(func, a, depth);
let rb = self.range_of(func, b, depth);
match ra.hi.checked_add(rb.hi) {
Some(hi) => Range::new(ra.lo.wrapping_add(rb.lo), hi),
None => Range::FULL,
}
}
InstKind::Sub(a, b) => {
let ra = self.range_of(func, a, depth);
let rb = self.range_of(func, b, depth);
if ra.lo >= rb.hi { Range::new(ra.lo - rb.hi, ra.hi - rb.lo) } else { Range::FULL }
}
InstKind::Mul(a, b) => {
let ra = self.range_of(func, a, depth);
let rb = self.range_of(func, b, depth);
match ra.hi.checked_mul(rb.hi) {
Some(hi) => Range::new(ra.lo.wrapping_mul(rb.lo), hi),
None => Range::FULL,
}
}
InstKind::Div(a, b) => {
let ra = self.range_of(func, a, depth);
let rb = self.range_of(func, b, depth);
let lo = if rb.lo > U256::ZERO { ra.lo / rb.hi } else { U256::ZERO };
Range::new(lo, ra.hi)
}
InstKind::Mod(a, b) => {
let ra = self.range_of(func, a, depth);
let rb = self.range_of(func, b, depth);
let bound = if rb.hi > U256::ZERO { rb.hi - U256::from(1) } else { U256::ZERO };
Range::new(U256::ZERO, bound.min(ra.hi))
}
InstKind::And(a, b) => {
let ra = self.range_of(func, a, depth);
let rb = self.range_of(func, b, depth);
Range::new(U256::ZERO, ra.hi.min(rb.hi))
}
InstKind::Lt(..)
| InstKind::Gt(..)
| InstKind::SLt(..)
| InstKind::SGt(..)
| InstKind::Eq(..)
| InstKind::IsZero(..) => match self.eval_truth(func, value, depth) {
Some(true) => Range::singleton(U256::from(1)),
Some(false) => Range::singleton(U256::ZERO),
None => Range::new(U256::ZERO, U256::from(1)),
},
InstKind::Select(condition, then_value, else_value) => {
match self.eval_truth(func, condition, depth) {
Some(true) => self.range_of(func, then_value, depth),
Some(false) => self.range_of(func, else_value, depth),
None => self
.range_of(func, then_value, depth)
.union(self.range_of(func, else_value, depth)),
}
}
_ => Range::FULL,
};
if let Some(intersection) = range.intersect(derived) {
range = intersection;
}
range
}
fn eval_truth(&mut self, func: &Function, value: ValueId, depth: usize) -> Option<bool> {
if let Some(constant) = const_of(func, value) {
return Some(!constant.is_zero());
}
if let Some(range) = self.ranges.get(&value) {
if range.lo > U256::ZERO {
return Some(true);
}
if range.hi.is_zero() {
return Some(false);
}
}
let depth = depth.checked_sub(1)?;
let kind = inst_kind(func, value)?;
match *kind {
InstKind::Lt(a, b) => self.eval_lt(func, a, b, depth),
InstKind::Gt(a, b) => self.eval_lt(func, b, a, depth),
InstKind::Eq(a, b) => self.eval_eq(func, a, b, depth),
InstKind::IsZero(a) => self.eval_truth(func, a, depth).map(|truth| !truth),
InstKind::Sub(a, b) | InstKind::Xor(a, b) => {
self.eval_eq(func, a, b, depth).map(|eq| !eq)
}
InstKind::And(a, b) => {
let ta = self.eval_truth(func, a, depth);
let tb = self.eval_truth(func, b, depth);
if ta == Some(false) || tb == Some(false) {
return Some(false);
}
let one = Range::singleton(U256::from(1));
if self.range_of(func, a, depth) == one && self.range_of(func, b, depth) == one {
return Some(true);
}
None
}
InstKind::Or(a, b) => {
let ta = self.eval_truth(func, a, depth);
let tb = self.eval_truth(func, b, depth);
if ta == Some(true) || tb == Some(true) {
return Some(true);
}
if ta == Some(false) && tb == Some(false) {
return Some(false);
}
None
}
_ => {
let range = self.range_of(func, value, depth);
if range.lo > U256::ZERO {
return Some(true);
}
if range.hi.is_zero() {
return Some(false);
}
None
}
}
}
fn eval_lt(&mut self, func: &Function, a: ValueId, b: ValueId, depth: usize) -> Option<bool> {
if a == b {
return Some(false);
}
let (x, y) = ordered(a, b);
if self.has_relation(Relation::Lt(a, b)) {
return Some(true);
}
if self.has_relation(Relation::Lt(b, a))
|| self.has_relation(Relation::Le(b, a))
|| self.has_relation(Relation::Eq(x, y))
{
return Some(false);
}
let ra = self.range_of(func, a, depth);
let rb = self.range_of(func, b, depth);
if ra.hi < rb.lo {
return Some(true);
}
if ra.lo >= rb.hi {
return Some(false);
}
if let Some(&InstKind::Add(x, y)) = inst_kind(func, a)
&& (b == x || b == y)
{
let rx = self.range_of(func, x, depth);
let ry = self.range_of(func, y, depth);
if rx.hi.checked_add(ry.hi).is_some() {
return Some(false);
}
if rx.lo.checked_add(ry.lo).is_none() {
return Some(true);
}
}
if let Some(&InstKind::Sub(x, y)) = inst_kind(func, b)
&& a == x
&& let Some(reduced_depth) = depth.checked_sub(1)
{
return self.eval_lt(func, x, y, reduced_depth);
}
None
}
fn eval_eq(&mut self, func: &Function, a: ValueId, b: ValueId, depth: usize) -> Option<bool> {
if a == b {
return Some(true);
}
let (x, y) = ordered(a, b);
if self.has_relation(Relation::Eq(x, y)) {
return Some(true);
}
if self.has_relation(Relation::Ne(x, y))
|| self.has_relation(Relation::Lt(a, b))
|| self.has_relation(Relation::Lt(b, a))
{
return Some(false);
}
let ra = self.range_of(func, a, depth);
let rb = self.range_of(func, b, depth);
if ra.hi < rb.lo || rb.hi < ra.lo {
return Some(false);
}
if ra.is_singleton() && ra == rb {
return Some(true);
}
if let Some(truth) = self.eval_muldiv_roundtrip(func, a, b, depth) {
return Some(truth);
}
if let Some(truth) = self.eval_muldiv_roundtrip(func, b, a, depth) {
return Some(truth);
}
None
}
fn eval_muldiv_roundtrip(
&mut self,
func: &Function,
div_value: ValueId,
expected: ValueId,
depth: usize,
) -> Option<bool> {
let InstKind::Div(mul_value, divisor) = *inst_kind(func, div_value)? else { return None };
let InstKind::Mul(p, q) = *inst_kind(func, mul_value)? else { return None };
for (x, y) in [(p, q), (q, p)] {
if x != expected || !values_equal(func, divisor, y) {
continue;
}
let ry = self.range_of(func, y, depth);
if ry.lo.is_zero() {
continue;
}
let rx = self.range_of(func, x, depth);
if rx.hi.checked_mul(ry.hi).is_some() {
return Some(true);
}
}
None
}
}
fn dominating_edge_fact(
func: &Function,
preds: &[Vec<BlockId>],
block: BlockId,
) -> Option<(ValueId, bool)> {
if block == func.entry_block {
return None;
}
let preds = &preds[block.index()];
let (&first, rest) = preds.split_first()?;
if rest.iter().any(|&pred| pred != first) {
return None;
}
let Terminator::Branch { condition, then_block, else_block } =
func.blocks[first].terminator.as_ref()?
else {
return None;
};
if then_block == else_block {
return None;
}
if *then_block == block {
Some((*condition, true))
} else if *else_block == block {
Some((*condition, false))
} else {
None
}
}
fn const_of(func: &Function, value: ValueId) -> Option<U256> {
match func.value(value) {
Value::Immediate(imm) => imm.as_u256(),
_ => None,
}
}
fn inst_kind(func: &Function, value: ValueId) -> Option<&InstKind> {
match func.value(value) {
Value::Inst(inst_id) => Some(&func.instructions[*inst_id].kind),
_ => None,
}
}
fn values_equal(func: &Function, a: ValueId, b: ValueId) -> bool {
if a == b {
return true;
}
match (const_of(func, a), const_of(func, b)) {
(Some(a), Some(b)) => a == b,
_ => false,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn range_intersection_and_union() {
let a = Range::new(U256::from(0), U256::from(10));
let b = Range::new(U256::from(5), U256::from(20));
assert_eq!(a.intersect(b), Some(Range::new(U256::from(5), U256::from(10))));
assert_eq!(a.union(b), Range::new(U256::from(0), U256::from(20)));
let disjoint = Range::new(U256::from(11), U256::from(12));
assert_eq!(a.intersect(disjoint), None);
}
}