Skip to main content

foundry_evm_symbolic/runtime/solver/smt/
mod.rs

1//! Query-level SMT-LIB assertion emission with common-subexpression bindings.
2
3use super::*;
4
5pub(super) fn write_smt_assertions(
6    cx: &SymCx,
7    out: &mut String,
8    constraints: &[SymBoolExpr],
9) -> Result<(), SymbolicError> {
10    if constraints.is_empty() {
11        return Ok(());
12    }
13    if constraints.iter().any(SymBoolExpr::contains_gasleft) {
14        return Err(SymbolicError::Unsupported("GAS/gasleft() not modeled"));
15    }
16
17    let plan = SmtCsePlan::new(constraints);
18    if plan.bindings.is_empty() {
19        for constraint in constraints {
20            let _ = writeln!(out, "(assert {})", constraint.smt(cx));
21        }
22        return Ok(());
23    }
24
25    let writer = SmtCseWriter { cx, plan: &plan };
26    // define binding_0 = term_0
27    // ...
28    // define binding_n = term_n
29    // assert constraint_0
30    // ...
31    // assert constraint_n
32    for (idx, binding) in plan.bindings.iter().enumerate() {
33        out.push_str("(define-fun ");
34        binding.write_definition_header(out, idx);
35        match binding {
36            SmtBinding::Expr(expr) => writer.write_expr(out, expr, Some(idx), None),
37            SmtBinding::Bool(expr) => writer.write_bool(out, expr, None, Some(idx)),
38        }
39        out.push_str(")\n");
40    }
41    for constraint in constraints {
42        out.push_str("(assert ");
43        writer.write_bool(out, constraint, None, None);
44        out.push_str(")\n");
45    }
46    Ok(())
47}
48
49#[derive(Default)]
50struct SmtCseVisit {
51    count: usize,
52    binding: Option<usize>,
53    collected: bool,
54}
55
56struct SmtCsePlan {
57    expr_visits: HashMap<SymExpr, SmtCseVisit>,
58    bool_visits: HashMap<SymBoolExpr, SmtCseVisit>,
59    bindings: Vec<SmtBinding>,
60}
61
62impl SmtCsePlan {
63    fn new(constraints: &[SymBoolExpr]) -> Self {
64        let mut plan = Self {
65            expr_visits: HashMap::default(),
66            bool_visits: HashMap::default(),
67            bindings: Vec::new(),
68        };
69        for constraint in constraints {
70            plan.count_bool(constraint);
71        }
72        for constraint in constraints {
73            plan.collect_bool_binding(constraint);
74        }
75        plan
76    }
77
78    fn count_expr(&mut self, expr: &SymExpr) {
79        let visit = self.expr_visits.entry(expr.clone()).or_default();
80        visit.count += 1;
81        if visit.count != 1 {
82            return;
83        }
84        match expr.kind() {
85            SymExprKind::Const(_)
86            | SymExprKind::Var(_)
87            | SymExprKind::GasLeft(_)
88            | SymExprKind::Keccak { .. }
89            | SymExprKind::Hash { .. } => {}
90            SymExprKind::Not(value) => self.count_expr(value),
91            SymExprKind::BinOp(_, left, right) => {
92                self.count_expr(left);
93                self.count_expr(right);
94            }
95            SymExprKind::TernOp(_, left, right, modulus) => {
96                self.count_expr(modulus);
97                self.count_expr(left);
98                self.count_expr(right);
99                self.count_expr(modulus);
100            }
101            SymExprKind::Ite(cond, left, right) => {
102                self.count_bool(cond);
103                self.count_expr(left);
104                self.count_expr(right);
105            }
106        }
107    }
108
109    fn count_bool(&mut self, expr: &SymBoolExpr) {
110        let visit = self.bool_visits.entry(expr.clone()).or_default();
111        visit.count += 1;
112        if visit.count != 1 {
113            return;
114        }
115        match expr.kind() {
116            SymBoolExprKind::Const(_) => {}
117            SymBoolExprKind::Not(value) => self.count_bool(value),
118            SymBoolExprKind::And(values) => {
119                for value in values.iter() {
120                    self.count_bool(value);
121                }
122            }
123            SymBoolExprKind::Cmp(_, left, right) => {
124                self.count_expr(left);
125                self.count_expr(right);
126            }
127        }
128    }
129
130    fn collect_expr_binding(&mut self, expr: &SymExpr) {
131        {
132            let Some(visit) = self.expr_visits.get_mut(expr) else {
133                return;
134            };
135            if visit.collected {
136                return;
137            }
138            visit.collected = true;
139        }
140        match expr.kind() {
141            SymExprKind::Const(_)
142            | SymExprKind::Var(_)
143            | SymExprKind::GasLeft(_)
144            | SymExprKind::Keccak { .. }
145            | SymExprKind::Hash { .. } => {}
146            SymExprKind::Not(value) => self.collect_expr_binding(value),
147            SymExprKind::BinOp(_, left, right) => {
148                self.collect_expr_binding(left);
149                self.collect_expr_binding(right);
150            }
151            SymExprKind::TernOp(_, left, right, modulus) => {
152                self.collect_expr_binding(modulus);
153                self.collect_expr_binding(left);
154                self.collect_expr_binding(right);
155            }
156            SymExprKind::Ite(cond, left, right) => {
157                self.collect_bool_binding(cond);
158                self.collect_expr_binding(left);
159                self.collect_expr_binding(right);
160            }
161        }
162        self.bind_expr(expr);
163    }
164
165    fn collect_bool_binding(&mut self, expr: &SymBoolExpr) {
166        {
167            let Some(visit) = self.bool_visits.get_mut(expr) else {
168                return;
169            };
170            if visit.collected {
171                return;
172            }
173            visit.collected = true;
174        }
175        match expr.kind() {
176            SymBoolExprKind::Const(_) => {}
177            SymBoolExprKind::Not(value) => self.collect_bool_binding(value),
178            SymBoolExprKind::And(values) => {
179                for value in values.iter() {
180                    self.collect_bool_binding(value);
181                }
182            }
183            SymBoolExprKind::Cmp(_, left, right) => {
184                self.collect_expr_binding(left);
185                self.collect_expr_binding(right);
186            }
187        }
188        self.bind_bool(expr);
189    }
190
191    fn bind_expr(&mut self, expr: &SymExpr) {
192        let Some(visit) = self.expr_visits.get_mut(expr) else {
193            return;
194        };
195        if visit.count <= 1 || visit.binding.is_some() || !Self::expr_can_bind(expr) {
196            return;
197        }
198        let idx = self.bindings.len();
199        visit.binding = Some(idx);
200        self.bindings.push(SmtBinding::Expr(expr.clone()));
201    }
202
203    fn bind_bool(&mut self, expr: &SymBoolExpr) {
204        let Some(visit) = self.bool_visits.get_mut(expr) else {
205            return;
206        };
207        if visit.count <= 1
208            || visit.binding.is_some()
209            || matches!(expr.kind(), SymBoolExprKind::Const(_))
210        {
211            return;
212        }
213        let idx = self.bindings.len();
214        visit.binding = Some(idx);
215        self.bindings.push(SmtBinding::Bool(expr.clone()));
216    }
217
218    fn expr_can_bind(expr: &SymExpr) -> bool {
219        !matches!(
220            expr.kind(),
221            SymExprKind::Const(_)
222                | SymExprKind::Var(_)
223                | SymExprKind::GasLeft(_)
224                | SymExprKind::Keccak { .. }
225                | SymExprKind::Hash { .. }
226        )
227    }
228}
229
230enum SmtBinding {
231    Expr(SymExpr),
232    Bool(SymBoolExpr),
233}
234
235impl SmtBinding {
236    fn write_definition_header(&self, out: &mut String, idx: usize) {
237        match self {
238            Self::Expr(_) => {
239                let _ = write!(out, "__sym_expr_{idx}");
240                out.push_str(" () (_ BitVec 256) ");
241            }
242            Self::Bool(_) => {
243                let _ = write!(out, "__sym_bool_{idx}");
244                out.push_str(" () Bool ");
245            }
246        }
247    }
248}
249
250struct SmtCseWriter<'a> {
251    cx: &'a SymCx,
252    plan: &'a SmtCsePlan,
253}
254
255impl SmtCseWriter<'_> {
256    fn write_expr(
257        &self,
258        out: &mut String,
259        expr: &SymExpr,
260        skip_expr: Option<usize>,
261        skip_bool: Option<usize>,
262    ) {
263        if let Some(idx) = self.plan.expr_visits.get(expr).and_then(|visit| visit.binding)
264            && Some(idx) != skip_expr
265        {
266            let _ = write!(out, "__sym_expr_{idx}");
267            return;
268        }
269
270        match expr.kind() {
271            SymExprKind::Const(value) => {
272                let _ = write!(out, "(_ bv{value} 256)");
273            }
274            SymExprKind::Var(symbol)
275            | SymExprKind::GasLeft(symbol)
276            | SymExprKind::Keccak { name: symbol, .. }
277            | SymExprKind::Hash { name: symbol, .. } => out.push_str(self.cx.symbol_name(*symbol)),
278            SymExprKind::Not(value) => {
279                out.push_str("(bvnot ");
280                self.write_expr(out, value, skip_expr, skip_bool);
281                out.push(')');
282            }
283            SymExprKind::BinOp(op, left, right) => {
284                let _ = write!(out, "({} ", op.smt());
285                self.write_expr(out, left, skip_expr, skip_bool);
286                out.push(' ');
287                self.write_expr(out, right, skip_expr, skip_bool);
288                out.push(')');
289            }
290            SymExprKind::TernOp(op, left, right, modulus) => {
291                self.write_wide_modular_arithmetic(out, op.smt(), left, right, modulus);
292            }
293            SymExprKind::Ite(cond, left, right) => {
294                out.push_str("(ite ");
295                self.write_bool(out, cond, skip_expr, skip_bool);
296                out.push(' ');
297                self.write_expr(out, left, skip_expr, skip_bool);
298                out.push(' ');
299                self.write_expr(out, right, skip_expr, skip_bool);
300                out.push(')');
301            }
302        }
303    }
304
305    fn write_wide_modular_arithmetic(
306        &self,
307        out: &mut String,
308        op: &'static str,
309        left: &SymExpr,
310        right: &SymExpr,
311        modulus: &SymExpr,
312    ) {
313        // if modulus == 0:
314        //   0
315        // else:
316        //   low_256((zext(left) op zext(right)) urem zext(modulus))
317        out.push_str("(ite (= ");
318        self.write_expr(out, modulus, None, None);
319        out.push_str(" (_ bv0 256)) (_ bv0 256) ((_ extract 255 0) (bvurem (");
320        out.push_str(op);
321        out.push_str(" ((_ zero_extend 256) ");
322        self.write_expr(out, left, None, None);
323        out.push_str(") ((_ zero_extend 256) ");
324        self.write_expr(out, right, None, None);
325        out.push_str(")) ((_ zero_extend 256) ");
326        self.write_expr(out, modulus, None, None);
327        out.push_str("))))");
328    }
329
330    fn write_bool(
331        &self,
332        out: &mut String,
333        expr: &SymBoolExpr,
334        skip_expr: Option<usize>,
335        skip_bool: Option<usize>,
336    ) {
337        if let Some(idx) = self.plan.bool_visits.get(expr).and_then(|visit| visit.binding)
338            && Some(idx) != skip_bool
339        {
340            let _ = write!(out, "__sym_bool_{idx}");
341            return;
342        }
343
344        match expr.kind() {
345            SymBoolExprKind::Const(value) => out.push_str(if *value { "true" } else { "false" }),
346            SymBoolExprKind::Not(value) => {
347                out.push_str("(not ");
348                self.write_bool(out, value, skip_expr, skip_bool);
349                out.push(')');
350            }
351            SymBoolExprKind::And(values) => {
352                out.push_str("(and");
353                for value in values.iter() {
354                    out.push(' ');
355                    self.write_bool(out, value, skip_expr, skip_bool);
356                }
357                out.push(')');
358            }
359            SymBoolExprKind::Cmp(op, left, right) => {
360                let _ = write!(out, "({} ", op.smt());
361                self.write_expr(out, left, skip_expr, skip_bool);
362                out.push(' ');
363                self.write_expr(out, right, skip_expr, skip_bool);
364                out.push(')');
365            }
366        }
367    }
368}