Skip to main content

foundry_evm_symbolic/runtime/solver/
hard_arith_fallback.rs

1use super::*;
2
3impl SymBoolExpr {
4    pub(crate) fn contains_hard_arith(&self) -> bool {
5        self.visit_bool(is_hard_arith_node)
6    }
7
8    fn contains_symbolic_hash(&self) -> bool {
9        self.visit_bool(|expr| matches!(expr.kind(), SymExprKind::Hash { .. }))
10    }
11}
12
13impl SymExpr {
14    #[cfg(test)]
15    pub(crate) fn contains_hard_arith(&self) -> bool {
16        self.visit_bool(is_hard_arith_node)
17    }
18
19    fn contains_var(&self) -> bool {
20        self.visit_bool(|expr| {
21            matches!(
22                expr.kind(),
23                SymExprKind::Var(_) | SymExprKind::Keccak { .. } | SymExprKind::Hash { .. }
24            )
25        })
26    }
27}
28
29fn is_hard_arith_node(expr: &SymExpr) -> bool {
30    match expr.kind() {
31        SymExprKind::BinOp(SymBinOp::Mul, left, right) => {
32            left.contains_var() && right.contains_var()
33        }
34        SymExprKind::BinOp(
35            SymBinOp::UDiv | SymBinOp::URem | SymBinOp::SDiv | SymBinOp::SRem,
36            left,
37            right,
38        ) => left.contains_var() || right.contains_var(),
39        SymExprKind::TernOp(_, left, right, modulus) => {
40            left.contains_var() || right.contains_var() || modulus.contains_var()
41        }
42        _ => false,
43    }
44}
45
46/// Returns whether local hard-arithmetic search should run before asking the solver.
47pub(crate) fn constraints_prefer_hard_arith_fallback_first(
48    cx: &SymCx,
49    constraints: &[SymBoolExpr],
50) -> bool {
51    if !constraints.iter().any(SymBoolExpr::contains_hard_arith)
52        || constraints.iter().any(SymBoolExpr::contains_symbolic_hash)
53    {
54        return false;
55    }
56
57    let mut vars = SymbolicVars::default();
58    for constraint in constraints {
59        collect_bool_fallback_vars(constraint, &mut vars);
60    }
61    let vars = fallback_search_vars(cx, vars, constraints);
62    !vars.is_empty() && vars.len() <= HARD_ARITH_FALLBACK_MAX_VARS
63}
64
65pub(crate) fn hard_arith_fallback_model(
66    cx: &SymCx,
67    constraints: &[SymBoolExpr],
68) -> Option<SymbolicModel> {
69    if !constraints.iter().any(SymBoolExpr::contains_hard_arith)
70        || constraints.iter().any(SymBoolExpr::contains_symbolic_hash)
71    {
72        return None;
73    }
74
75    let mut vars = SymbolicVars::default();
76    let mut constants = HashSet::<U256>::default();
77    for constraint in constraints {
78        collect_bool_fallback_vars(constraint, &mut vars);
79        collect_bool_constants(constraint, &mut constants);
80    }
81    let mut constants = constants.into_iter().collect::<Vec<_>>();
82    constants.sort_unstable();
83    let vars = fallback_search_vars(cx, vars, constraints);
84    if vars.is_empty() || vars.len() > HARD_ARITH_FALLBACK_MAX_VARS {
85        return None;
86    }
87
88    let candidates = vars
89        .iter()
90        .map(|var| fallback_candidates_for_var(var, constraints, &constants))
91        .collect::<Option<Vec<_>>>()?;
92    let searched_vars = vars.iter().copied().collect::<SymbolicVars>();
93    let constraint_vars = constraints
94        .iter()
95        .map(|constraint| {
96            let mut vars = SymbolicVars::default();
97            constraint.collect_vars(&mut vars);
98            vars
99        })
100        .collect::<Vec<_>>();
101    let mut model = SymbolicModel::default();
102    let mut assignments = 0usize;
103    let search = FallbackSearch {
104        constraints,
105        constraint_vars: &constraint_vars,
106        searched_vars: &searched_vars,
107        vars: &vars,
108        candidates: &candidates,
109    };
110    search.model(0, &mut model, &mut assignments)
111}
112
113fn fallback_search_vars(
114    cx: &SymCx,
115    vars: SymbolicVars,
116    constraints: &[SymBoolExpr],
117) -> Vec<Symbol> {
118    if vars.len() <= HARD_ARITH_FALLBACK_MAX_VARS {
119        return vars.into_iter().collect();
120    }
121
122    let hard_arith_vars = hard_arith_fallback_vars(constraints);
123    if !hard_arith_vars.is_empty() && hard_arith_vars.len() <= HARD_ARITH_FALLBACK_MAX_VARS {
124        let mut vars = hard_arith_vars;
125        add_zero_invalid_support_vars(&mut vars, constraints);
126        return vars.into_iter().collect();
127    }
128
129    vars.into_iter()
130        .filter(|var| {
131            let var = cx.symbol_name(*var);
132            var.starts_with("calldata")
133                || var.starts_with("sequence")
134                || var.starts_with("create_address")
135                || var.starts_with("create2_address")
136                || !var.contains('_')
137        })
138        .collect()
139}
140
141fn hard_arith_fallback_vars(constraints: &[SymBoolExpr]) -> SymbolicVars {
142    let mut vars = SymbolicVars::default();
143    for constraint in constraints {
144        collect_bool_hard_arith_vars(constraint, &mut vars);
145    }
146    vars
147}
148
149fn add_zero_invalid_support_vars(vars: &mut SymbolicVars, constraints: &[SymBoolExpr]) {
150    let zero_model = SymbolicModel::default();
151    for constraint in constraints {
152        if constraint.eval_model(&zero_model).unwrap_or(false) {
153            continue;
154        }
155
156        let mut constraint_vars = SymbolicVars::default();
157        constraint.collect_vars(&mut constraint_vars);
158        let missing =
159            constraint_vars.iter().filter(|var| !vars.contains(*var)).copied().collect::<Vec<_>>();
160        if vars.len() + missing.len() > HARD_ARITH_FALLBACK_MAX_VARS {
161            continue;
162        }
163        vars.extend(missing);
164    }
165}
166
167fn fallback_candidates_for_var(
168    var: &Symbol,
169    constraints: &[SymBoolExpr],
170    constants: &[U256],
171) -> Option<Vec<U256>> {
172    let hints = MaskHints::for_var(var, constraints);
173    if (hints.one & hints.zero) != U256::ZERO {
174        return None;
175    }
176
177    let mut candidates = HashSet::<U256>::default();
178    for candidate in [
179        U256::ZERO,
180        U256::from(1),
181        U256::from(2),
182        U256::from(3),
183        U256::MAX,
184        U256::MAX - U256::from(1),
185        U256::MAX - U256::from(2),
186    ] {
187        push_fallback_candidate(&mut candidates, candidate, hints);
188    }
189
190    for constant in constants.iter().copied() {
191        push_fallback_candidate(&mut candidates, constant, hints);
192        push_fallback_candidate(&mut candidates, constant.wrapping_add(U256::from(1)), hints);
193        push_fallback_candidate(&mut candidates, constant.wrapping_sub(U256::from(1)), hints);
194        if candidates.len() >= HARD_ARITH_FALLBACK_MAX_CANDIDATES_PER_VAR {
195            break;
196        }
197    }
198
199    for bit in 0..256 {
200        let power = U256::from(1) << bit;
201        push_fallback_candidate(&mut candidates, power, hints);
202        if candidates.len() >= HARD_ARITH_FALLBACK_MAX_CANDIDATES_PER_VAR {
203            break;
204        }
205    }
206
207    let mut candidates = candidates.into_iter().collect::<Vec<_>>();
208    candidates.sort_unstable();
209    candidates.truncate(HARD_ARITH_FALLBACK_MAX_CANDIDATES_PER_VAR);
210    Some(candidates)
211}
212
213struct FallbackSearch<'a> {
214    constraints: &'a [SymBoolExpr],
215    constraint_vars: &'a [SymbolicVars],
216    searched_vars: &'a SymbolicVars,
217    vars: &'a [Symbol],
218    candidates: &'a [Vec<U256>],
219}
220
221impl FallbackSearch<'_> {
222    fn model(
223        &self,
224        index: usize,
225        model: &mut SymbolicModel,
226        assignments: &mut usize,
227    ) -> Option<SymbolicModel> {
228        if index == self.vars.len() {
229            *assignments += 1;
230            if *assignments > HARD_ARITH_FALLBACK_MAX_ASSIGNMENTS {
231                return None;
232            }
233            let mut completed = model.clone();
234            return complete_fallback_support_model(self.constraints, &mut completed)
235                .then_some(completed);
236        }
237
238        for candidate in &self.candidates[index] {
239            model.insert(self.vars[index], *candidate);
240            if fallback_partial_model_satisfies_known_constraints(
241                self.constraints,
242                self.constraint_vars,
243                self.searched_vars,
244                model,
245            ) && let Some(model) = self.model(index + 1, model, assignments)
246            {
247                return Some(model);
248            }
249            if *assignments > HARD_ARITH_FALLBACK_MAX_ASSIGNMENTS {
250                return None;
251            }
252        }
253        model.remove(&self.vars[index]);
254        None
255    }
256}
257
258fn fallback_model_satisfies_all_constraints(
259    constraints: &[SymBoolExpr],
260    model: &(impl SymbolicModelLookup + ?Sized),
261) -> bool {
262    constraints.iter().all(|constraint| constraint.eval_model(model).unwrap_or(false))
263}
264
265fn complete_fallback_support_model(constraints: &[SymBoolExpr], model: &mut SymbolicModel) -> bool {
266    for _ in 0..constraints.len() {
267        let mut changed = false;
268        for constraint in constraints {
269            match constraint.eval_model_if_complete(model) {
270                Ok(Some(true)) => {}
271                Ok(Some(false)) | Err(_) => return false,
272                Ok(None) => {
273                    changed |= complete_support_constraint(constraint, model);
274                }
275            }
276        }
277        if changed {
278            continue;
279        }
280        // Default checked-add bases to zero only after exact/lower-bound completions had a chance
281        // to assign a stronger value required by another constraint.
282        for constraint in constraints {
283            match constraint.eval_model_if_complete(model) {
284                Ok(Some(true)) => {}
285                Ok(Some(false)) | Err(_) => return false,
286                Ok(None) => {
287                    changed |= complete_default_support_constraint(constraint, model);
288                }
289            }
290        }
291        if !changed {
292            break;
293        }
294    }
295    fallback_model_satisfies_all_constraints(constraints, model)
296}
297
298fn complete_support_constraint(constraint: &SymBoolExpr, model: &mut SymbolicModel) -> bool {
299    complete_support_bool(constraint, model, false, false)
300}
301
302fn complete_default_support_constraint(
303    constraint: &SymBoolExpr,
304    model: &mut SymbolicModel,
305) -> bool {
306    complete_support_bool(constraint, model, false, true)
307}
308
309fn complete_support_bool(
310    constraint: &SymBoolExpr,
311    model: &mut SymbolicModel,
312    inverted: bool,
313    defaults_only: bool,
314) -> bool {
315    match constraint.kind() {
316        SymBoolExprKind::Const(_) => false,
317        SymBoolExprKind::Not(value) => {
318            complete_support_bool(value, model, !inverted, defaults_only)
319        }
320        SymBoolExprKind::And(values) if !inverted => {
321            let mut changed = false;
322            for value in values.iter() {
323                changed |= complete_support_bool(value, model, false, defaults_only);
324            }
325            changed
326        }
327        SymBoolExprKind::Cmp(op, left, right) => {
328            let Some(op) = support_cmp_op(*op, inverted) else {
329                return false;
330            };
331            if defaults_only {
332                complete_default_support_comparison(op, left, right, model)
333            } else {
334                complete_support_comparison(op, left, right, model)
335            }
336        }
337        SymBoolExprKind::And(_) => false,
338    }
339}
340
341const fn support_cmp_op(op: SymCmpOp, inverted: bool) -> Option<SymCmpOp> {
342    if !inverted {
343        return Some(op);
344    }
345
346    match op {
347        SymCmpOp::Ult => Some(SymCmpOp::Uge),
348        SymCmpOp::Ugt => Some(SymCmpOp::Ule),
349        SymCmpOp::Ule => Some(SymCmpOp::Ugt),
350        SymCmpOp::Uge => Some(SymCmpOp::Ult),
351        SymCmpOp::Eq | SymCmpOp::Slt | SymCmpOp::Sgt => None,
352    }
353}
354
355fn complete_support_comparison(
356    op: SymCmpOp,
357    left: &SymExpr,
358    right: &SymExpr,
359    model: &mut SymbolicModel,
360) -> bool {
361    if complete_checked_sub_guard(op, left, right, model) {
362        return true;
363    }
364    if let Ok(Some(value)) = left.eval_model_if_complete(model)
365        && let Some(target) = support_target_for_known_left(op, value)
366    {
367        return right.assign_model_value(model, target);
368    }
369    if let Ok(Some(value)) = right.eval_model_if_complete(model)
370        && let Some(target) = support_target_for_known_right(op, value)
371    {
372        return left.assign_model_value(model, target);
373    }
374    false
375}
376
377fn complete_default_support_comparison(
378    op: SymCmpOp,
379    left: &SymExpr,
380    right: &SymExpr,
381    model: &mut SymbolicModel,
382) -> bool {
383    complete_checked_add_guard(op, left, right, model)
384}
385
386fn complete_checked_sub_guard(
387    op: SymCmpOp,
388    left: &SymExpr,
389    right: &SymExpr,
390    model: &mut SymbolicModel,
391) -> bool {
392    match op {
393        SymCmpOp::Uge => assign_checked_sub_minuend(left, right, model),
394        SymCmpOp::Ule => assign_checked_sub_minuend(right, left, model),
395        _ => false,
396    }
397}
398
399fn assign_checked_sub_minuend(
400    minuend: &SymExpr,
401    sub_expr: &SymExpr,
402    model: &mut SymbolicModel,
403) -> bool {
404    let SymExprKind::BinOp(SymBinOp::Sub, sub_minuend, amount) = sub_expr.kind() else {
405        return false;
406    };
407    if sub_minuend != minuend {
408        return false;
409    }
410    let Ok(Some(amount)) = amount.eval_model_if_complete(model) else {
411        return false;
412    };
413    minuend.assign_model_value(model, amount)
414}
415
416fn complete_checked_add_guard(
417    op: SymCmpOp,
418    left: &SymExpr,
419    right: &SymExpr,
420    model: &mut SymbolicModel,
421) -> bool {
422    match op {
423        SymCmpOp::Uge => assign_checked_add_base(left, right, model),
424        SymCmpOp::Ule => assign_checked_add_base(right, left, model),
425        _ => false,
426    }
427}
428
429fn assign_checked_add_base(sum: &SymExpr, base: &SymExpr, model: &mut SymbolicModel) -> bool {
430    let SymExprKind::BinOp(SymBinOp::Add, left, right) = sum.kind() else {
431        return false;
432    };
433    if left == base && right.eval_model_if_complete(model).ok().flatten().is_some() {
434        return base.assign_model_value(model, U256::ZERO);
435    }
436    if right == base && left.eval_model_if_complete(model).ok().flatten().is_some() {
437        return base.assign_model_value(model, U256::ZERO);
438    }
439    false
440}
441
442fn support_target_for_known_left(op: SymCmpOp, value: U256) -> Option<U256> {
443    match op {
444        SymCmpOp::Eq | SymCmpOp::Ule | SymCmpOp::Uge => Some(value),
445        SymCmpOp::Ult => value.checked_add(U256::from(1)),
446        SymCmpOp::Ugt => value.checked_sub(U256::from(1)),
447        SymCmpOp::Slt | SymCmpOp::Sgt => None,
448    }
449}
450
451fn support_target_for_known_right(op: SymCmpOp, value: U256) -> Option<U256> {
452    match op {
453        SymCmpOp::Eq | SymCmpOp::Ule | SymCmpOp::Uge => Some(value),
454        SymCmpOp::Ult => value.checked_sub(U256::from(1)),
455        SymCmpOp::Ugt => value.checked_add(U256::from(1)),
456        SymCmpOp::Slt | SymCmpOp::Sgt => None,
457    }
458}
459
460fn fallback_partial_model_satisfies_known_constraints(
461    constraints: &[SymBoolExpr],
462    constraint_vars: &[SymbolicVars],
463    searched_vars: &SymbolicVars,
464    model: &SymbolicModel,
465) -> bool {
466    constraints.iter().zip(constraint_vars).all(|(constraint, vars)| {
467        !vars.is_subset(searched_vars)
468            || !vars.iter().all(|var| model.contains_name(*var))
469            || constraint.eval_model(model).unwrap_or(false)
470    })
471}
472
473fn collect_bool_fallback_vars(expr: &SymBoolExpr, vars: &mut SymbolicVars) {
474    let _ = expr.visit_exprs(&mut |expr| {
475        if let Some(var) = expr.kind().get_eval_var() {
476            vars.insert(var);
477        }
478        ControlFlow::<()>::Continue(())
479    });
480}
481
482fn collect_bool_hard_arith_vars(expr: &SymBoolExpr, vars: &mut SymbolicVars) {
483    let _ = expr.visit_exprs(&mut |expr| {
484        if is_hard_arith_node(expr) {
485            expr.collect_eval_vars(vars);
486        }
487        ControlFlow::<()>::Continue(())
488    });
489}
490
491pub(crate) fn fallback_single_var_model(constraints: &[SymBoolExpr]) -> Option<SymbolicModel> {
492    let mut vars = SymbolicVars::default();
493    let mut constants = HashSet::<U256>::default();
494    for constraint in constraints {
495        constraint.collect_vars(&mut vars);
496        collect_bool_constants(constraint, &mut constants);
497    }
498    let mut constants = constants.into_iter().collect::<Vec<_>>();
499    constants.sort_unstable();
500
501    let var = if vars.len() == 1 { *vars.iter().next()? } else { return None };
502    let hints = MaskHints::for_var(&var, constraints);
503    if (hints.one & hints.zero) != U256::ZERO {
504        return None;
505    }
506
507    let mut candidates = HashSet::<U256>::default();
508    for candidate in [
509        U256::ZERO,
510        U256::from(1),
511        U256::from(2),
512        U256::MAX,
513        U256::MAX - U256::from(1),
514        U256::MAX - U256::from(2),
515    ] {
516        push_fallback_candidate(&mut candidates, candidate, hints);
517    }
518
519    for constant in constants.iter().copied() {
520        push_fallback_candidate(&mut candidates, constant, hints);
521        push_fallback_candidate(&mut candidates, constant.wrapping_add(U256::from(1)), hints);
522        push_fallback_candidate(&mut candidates, constant.wrapping_sub(U256::from(1)), hints);
523    }
524
525    for bit in 0..256 {
526        let power = U256::from(1) << bit;
527        push_fallback_candidate(&mut candidates, power, hints);
528        for constant in constants.iter().copied().take(64) {
529            push_fallback_candidate(&mut candidates, power | constant, hints);
530            push_fallback_candidate(&mut candidates, power.wrapping_add(constant), hints);
531        }
532    }
533
534    let mut candidates = candidates.into_iter().collect::<Vec<_>>();
535    candidates.sort_unstable();
536    for candidate in candidates {
537        let mut model = SymbolicModel::default();
538        model.insert(var, candidate);
539        if constraints.iter().all(|constraint| constraint.eval_model(&model).unwrap_or(false)) {
540            return Some(model);
541        }
542    }
543
544    None
545}
546
547pub(crate) fn fallback_two_var_model(constraints: &[SymBoolExpr]) -> Option<SymbolicModel> {
548    if constraints.iter().any(SymBoolExpr::contains_hard_arith) {
549        return None;
550    }
551
552    let mut vars = SymbolicVars::default();
553    for constraint in constraints {
554        collect_bool_fallback_vars(constraint, &mut vars);
555        if vars.len() > 2 {
556            return None;
557        }
558    }
559    if vars.len() != 2 {
560        return None;
561    }
562    if constraints.iter().any(SymBoolExpr::contains_symbolic_hash)
563        || constraints.iter().any(SymBoolExpr::contains_gasleft)
564    {
565        return None;
566    }
567    if !constraints_have_two_var_relation(constraints, &vars)
568        || !constraints_bind_each_search_var(constraints, &vars)
569    {
570        return None;
571    }
572
573    let mut constants = HashSet::<U256>::default();
574    for constraint in constraints {
575        collect_bool_constants(constraint, &mut constants);
576    }
577    let mut constants = constants.into_iter().collect::<Vec<_>>();
578    constants.sort_unstable();
579    let vars = vars.into_iter().collect::<Vec<_>>();
580    let candidates = vars
581        .iter()
582        .map(|var| fallback_candidates_for_var(var, constraints, &constants))
583        .collect::<Option<Vec<_>>>()?;
584    let searched_vars = vars.iter().copied().collect::<SymbolicVars>();
585    let constraint_vars = constraints
586        .iter()
587        .map(|constraint| {
588            let mut vars = SymbolicVars::default();
589            constraint.collect_vars(&mut vars);
590            vars
591        })
592        .collect::<Vec<_>>();
593    let search = FallbackSearch {
594        constraints,
595        constraint_vars: &constraint_vars,
596        searched_vars: &searched_vars,
597        vars: &vars,
598        candidates: &candidates,
599    };
600    let mut model = SymbolicModel::default();
601    let mut assignments = 0usize;
602    search.model(0, &mut model, &mut assignments)
603}
604
605fn constraints_have_two_var_relation(
606    constraints: &[SymBoolExpr],
607    searched_vars: &SymbolicVars,
608) -> bool {
609    constraints
610        .iter()
611        .any(|constraint| bool_expr_has_two_var_relation(constraint, searched_vars, false))
612}
613
614fn bool_expr_has_two_var_relation(
615    expr: &SymBoolExpr,
616    searched_vars: &SymbolicVars,
617    inverted: bool,
618) -> bool {
619    match expr.kind() {
620        SymBoolExprKind::Const(_) => false,
621        SymBoolExprKind::Not(expr) => {
622            bool_expr_has_two_var_relation(expr, searched_vars, !inverted)
623        }
624        SymBoolExprKind::And(exprs) if !inverted => {
625            exprs.iter().any(|expr| bool_expr_has_two_var_relation(expr, searched_vars, false))
626        }
627        SymBoolExprKind::And(_) => false,
628        SymBoolExprKind::Cmp(_, left, right) => {
629            let mut vars = SymbolicVars::default();
630            collect_expr_fallback_vars(left, &mut vars);
631            collect_expr_fallback_vars(right, &mut vars);
632            vars.len() == 2 && vars.is_subset(searched_vars)
633        }
634    }
635}
636
637fn constraints_bind_each_search_var(
638    constraints: &[SymBoolExpr],
639    searched_vars: &SymbolicVars,
640) -> bool {
641    searched_vars.iter().all(|var| {
642        constraints.iter().any(|constraint| bool_expr_binds_single_var(constraint, *var, false))
643    })
644}
645
646fn bool_expr_binds_single_var(expr: &SymBoolExpr, bound_var: Symbol, inverted: bool) -> bool {
647    match expr.kind() {
648        SymBoolExprKind::Const(_) => false,
649        SymBoolExprKind::Not(expr) => bool_expr_binds_single_var(expr, bound_var, !inverted),
650        SymBoolExprKind::And(exprs) if !inverted => {
651            exprs.iter().any(|expr| bool_expr_binds_single_var(expr, bound_var, false))
652        }
653        SymBoolExprKind::And(_) => false,
654        SymBoolExprKind::Cmp(_, left, right) => {
655            let mut vars = SymbolicVars::default();
656            collect_expr_fallback_vars(left, &mut vars);
657            collect_expr_fallback_vars(right, &mut vars);
658            vars.len() == 1
659                && vars.contains(&bound_var)
660                && (expr_contains_const(left) || expr_contains_const(right))
661        }
662    }
663}
664
665fn collect_expr_fallback_vars(expr: &SymExpr, vars: &mut SymbolicVars) {
666    let _ = expr.visit(&mut |expr| {
667        if let Some(var) = expr.kind().get_eval_var() {
668            vars.insert(var);
669        }
670        ControlFlow::<()>::Continue(())
671    });
672}
673
674fn expr_contains_const(expr: &SymExpr) -> bool {
675    expr.visit_bool(|expr| matches!(expr.kind(), SymExprKind::Const(_)))
676}
677
678fn push_fallback_candidate(candidates: &mut HashSet<U256>, candidate: U256, hints: MaskHints) {
679    candidates.insert((candidate | hints.one) & !hints.zero);
680}
681
682fn collect_bool_constants(expr: &SymBoolExpr, constants: &mut HashSet<U256>) {
683    let _ = expr.visit_exprs(&mut |expr| {
684        if let SymExprKind::Const(value) = expr.kind() {
685            constants.insert(*value);
686        }
687        ControlFlow::<()>::Continue(())
688    });
689}
690
691#[derive(Clone, Copy, Debug, Default)]
692struct MaskHints {
693    one: U256,
694    zero: U256,
695}
696
697impl MaskHints {
698    fn for_var(var: &Symbol, constraints: &[SymBoolExpr]) -> Self {
699        let mut hints = Self::default();
700        for constraint in constraints {
701            hints.apply_bool(var, constraint, false);
702        }
703        hints
704    }
705
706    fn apply_bool(&mut self, var: &Symbol, expr: &SymBoolExpr, inverted: bool) {
707        match expr.kind() {
708            SymBoolExprKind::Const(_) => {}
709            SymBoolExprKind::Not(value) => self.apply_bool(var, value, !inverted),
710            SymBoolExprKind::And(values) if !inverted => {
711                for value in values.iter() {
712                    self.apply_bool(var, value, false);
713                }
714            }
715            SymBoolExprKind::Cmp(SymCmpOp::Eq, left, right) => {
716                self.apply_equality(var, left, right, inverted)
717            }
718            SymBoolExprKind::Cmp(_, _, _) | SymBoolExprKind::And(_) => {}
719        }
720    }
721
722    fn apply_equality(&mut self, var: &Symbol, left: &SymExpr, right: &SymExpr, inverted: bool) {
723        if let Some(mask) =
724            zero_mask_equality(var, left, right).or_else(|| zero_mask_equality(var, right, left))
725        {
726            if inverted {
727                if is_single_bit(mask) {
728                    self.one |= mask;
729                }
730            } else {
731                self.zero |= mask;
732            }
733        }
734    }
735}
736
737fn is_single_bit(value: U256) -> bool {
738    !value.is_zero() && (value & (value - U256::from(1))).is_zero()
739}
740
741fn zero_mask_equality(var: &Symbol, masked: &SymExpr, zero: &SymExpr) -> Option<U256> {
742    if !zero.as_const().is_some_and(|value| value.is_zero()) {
743        return None;
744    }
745    match masked.kind() {
746        SymExprKind::BinOp(SymBinOp::And, left, right)
747            if left.kind().get_var().is_some_and(|name| &name == var) =>
748        {
749            right.as_const()
750        }
751        _ => None,
752    }
753}
754
755#[cfg(test)]
756mod tests {
757    use super::*;
758
759    #[test]
760    fn hard_arith_fallback_ignores_unrelated_abi_vars() {
761        let mut cx = SymCx::new();
762        let amount = SymExpr::var(&mut cx, "sequence_0_0_0_1");
763        let zero = SymExpr::zero(&mut cx);
764        let scale = SymExpr::constant(&mut cx, U256::from(1_000_000));
765        let product = SymExpr::binop(&mut cx, SymBinOp::Mul, scale.clone(), amount.clone());
766        let div = SymExpr::binop(&mut cx, SymBinOp::UDiv, product, amount.clone());
767        let amount_is_zero = SymBoolExpr::eq(&mut cx, amount, zero);
768        let guarded_zero = SymExpr::zero(&mut cx);
769        let guarded_div = SymExpr::ite(&mut cx, amount_is_zero.clone(), guarded_zero, div);
770        let overflow_branch = SymBoolExpr::eq(&mut cx, guarded_div, scale).not(&mut cx);
771
772        let address_bound = U256::from(1) << 160;
773        let mut constraints = vec![amount_is_zero.not(&mut cx), overflow_branch];
774        for idx in 0..6 {
775            let abi_word = SymExpr::var(&mut cx, &format!("sequence_0_0_0_addr_{idx}"));
776            constraints.push(SymBoolExpr::cmp_word_const(
777                &mut cx,
778                SymCmpOp::Ult,
779                &abi_word,
780                address_bound,
781            ));
782        }
783
784        assert!(constraints_prefer_hard_arith_fallback_first(&cx, &constraints));
785        let model = hard_arith_fallback_model(&cx, &constraints).expect("fallback model");
786        assert!(model.contains_name(cx.symbol("sequence_0_0_0_1")));
787        assert!(constraints.iter().all(|constraint| constraint.eval_model(&model).unwrap()));
788    }
789
790    #[test]
791    fn hard_arith_fallback_keeps_prior_path_vars_needed_by_zero_model() {
792        let mut cx = SymCx::new();
793        let setup_amount = SymExpr::var(&mut cx, "sequence_0_0_0_1");
794        let borrow_amount = SymExpr::var(&mut cx, "sequence_2_2_0_1");
795        let zero = SymExpr::zero(&mut cx);
796        let scale = SymExpr::constant(&mut cx, U256::from(1_000_000));
797        let product = SymExpr::binop(&mut cx, SymBinOp::Mul, scale.clone(), borrow_amount.clone());
798        let quotient = SymExpr::binop(&mut cx, SymBinOp::UDiv, product, borrow_amount.clone());
799
800        let constraints = vec![
801            SymBoolExpr::eq(&mut cx, setup_amount, zero.clone()).not(&mut cx),
802            SymBoolExpr::eq(&mut cx, borrow_amount, zero).not(&mut cx),
803            SymBoolExpr::eq(&mut cx, quotient, scale),
804        ];
805
806        assert!(constraints_prefer_hard_arith_fallback_first(&cx, &constraints));
807        let model = hard_arith_fallback_model(&cx, &constraints).expect("fallback model");
808        assert!(model.contains_name(cx.symbol("sequence_0_0_0_1")));
809        assert!(model.contains_name(cx.symbol("sequence_2_2_0_1")));
810        assert!(constraints.iter().all(|constraint| constraint.eval_model(&model).unwrap()));
811    }
812
813    #[test]
814    fn hard_arith_fallback_completes_checked_storage_guards() {
815        let mut cx = SymCx::new();
816        let amount = SymExpr::var(&mut cx, "sequence_0_0_0_1");
817        let from_balance = SymExpr::var(&mut cx, "storage_from_balance");
818        let to_balance = SymExpr::var(&mut cx, "storage_to_balance");
819        let zero = SymExpr::zero(&mut cx);
820        let scale = SymExpr::constant(&mut cx, U256::from(1_000_000));
821        let product = SymExpr::binop(&mut cx, SymBinOp::Mul, scale.clone(), amount.clone());
822        let quotient = SymExpr::binop(&mut cx, SymBinOp::UDiv, product, amount.clone());
823
824        let debited = SymExpr::binop(&mut cx, SymBinOp::Sub, from_balance.clone(), amount.clone());
825        let credited = SymExpr::binop(&mut cx, SymBinOp::Add, to_balance.clone(), amount.clone());
826        let mut constraints = vec![
827            SymBoolExpr::eq(&mut cx, amount, zero).not(&mut cx),
828            SymBoolExpr::eq(&mut cx, quotient, scale),
829            SymBoolExpr::cmp(&mut cx, SymCmpOp::Ult, from_balance, debited).not(&mut cx),
830            SymBoolExpr::cmp(&mut cx, SymCmpOp::Ult, credited, to_balance).not(&mut cx),
831        ];
832
833        let address_bound = U256::from(1) << 160;
834        for idx in 0..6 {
835            let abi_word = SymExpr::var(&mut cx, &format!("sequence_0_0_0_addr_{idx}"));
836            constraints.push(SymBoolExpr::cmp_word_const(
837                &mut cx,
838                SymCmpOp::Ult,
839                &abi_word,
840                address_bound,
841            ));
842        }
843
844        assert!(constraints_prefer_hard_arith_fallback_first(&cx, &constraints));
845        let model = hard_arith_fallback_model(&cx, &constraints).expect("fallback model");
846        assert!(model.contains_name(cx.symbol("sequence_0_0_0_1")));
847        assert!(model.contains_name(cx.symbol("storage_from_balance")));
848        assert!(model.contains_name(cx.symbol("storage_to_balance")));
849        assert!(constraints.iter().all(|constraint| constraint.eval_model(&model).unwrap()));
850    }
851}