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 || visit.binding.is_some() || !Self::bool_can_bind(expr) {
208            return;
209        }
210        let idx = self.bindings.len();
211        visit.binding = Some(idx);
212        self.bindings.push(SmtBinding::Bool(expr.clone()));
213    }
214
215    fn expr_binding(&self, expr: &SymExpr) -> Option<usize> {
216        self.expr_visits.get(expr).and_then(|visit| visit.binding)
217    }
218
219    fn bool_binding(&self, expr: &SymBoolExpr) -> Option<usize> {
220        self.bool_visits.get(expr).and_then(|visit| visit.binding)
221    }
222
223    fn expr_can_bind(expr: &SymExpr) -> bool {
224        !matches!(
225            expr.kind(),
226            SymExprKind::Const(_)
227                | SymExprKind::Var(_)
228                | SymExprKind::GasLeft(_)
229                | SymExprKind::Keccak { .. }
230                | SymExprKind::Hash { .. }
231        )
232    }
233
234    fn bool_can_bind(expr: &SymBoolExpr) -> bool {
235        !matches!(expr.kind(), SymBoolExprKind::Const(_))
236    }
237}
238
239enum SmtBinding {
240    Expr(SymExpr),
241    Bool(SymBoolExpr),
242}
243
244impl SmtBinding {
245    fn write_definition_header(&self, out: &mut String, idx: usize) {
246        match self {
247            Self::Expr(_) => {
248                Self::write_expr_name(out, idx);
249                out.push_str(" () (_ BitVec 256) ");
250            }
251            Self::Bool(_) => {
252                Self::write_bool_name(out, idx);
253                out.push_str(" () Bool ");
254            }
255        }
256    }
257
258    fn write_expr_name(out: &mut String, idx: usize) {
259        let _ = write!(out, "__sym_expr_{idx}");
260    }
261
262    fn write_bool_name(out: &mut String, idx: usize) {
263        let _ = write!(out, "__sym_bool_{idx}");
264    }
265}
266
267struct SmtCseWriter<'a> {
268    cx: &'a SymCx,
269    plan: &'a SmtCsePlan,
270}
271
272impl SmtCseWriter<'_> {
273    fn write_expr(
274        &self,
275        out: &mut String,
276        expr: &SymExpr,
277        skip_expr: Option<usize>,
278        skip_bool: Option<usize>,
279    ) {
280        if let Some(idx) = self.plan.expr_binding(expr)
281            && Some(idx) != skip_expr
282        {
283            SmtBinding::write_expr_name(out, idx);
284            return;
285        }
286
287        match expr.kind() {
288            SymExprKind::Const(value) => {
289                let _ = write!(out, "(_ bv{value} 256)");
290            }
291            SymExprKind::Var(symbol)
292            | SymExprKind::GasLeft(symbol)
293            | SymExprKind::Keccak { name: symbol, .. }
294            | SymExprKind::Hash { name: symbol, .. } => out.push_str(self.cx.symbol_name(*symbol)),
295            SymExprKind::Not(value) => {
296                out.push_str("(bvnot ");
297                self.write_expr(out, value, skip_expr, skip_bool);
298                out.push(')');
299            }
300            SymExprKind::BinOp(op, left, right) => {
301                let _ = write!(out, "({} ", op.smt());
302                self.write_expr(out, left, skip_expr, skip_bool);
303                out.push(' ');
304                self.write_expr(out, right, skip_expr, skip_bool);
305                out.push(')');
306            }
307            SymExprKind::TernOp(op, left, right, modulus) => {
308                self.write_wide_modular_arithmetic(out, op.smt(), left, right, modulus);
309            }
310            SymExprKind::Ite(cond, left, right) => {
311                out.push_str("(ite ");
312                self.write_bool(out, cond, skip_expr, skip_bool);
313                out.push(' ');
314                self.write_expr(out, left, skip_expr, skip_bool);
315                out.push(' ');
316                self.write_expr(out, right, skip_expr, skip_bool);
317                out.push(')');
318            }
319        }
320    }
321
322    fn write_wide_modular_arithmetic(
323        &self,
324        out: &mut String,
325        op: &'static str,
326        left: &SymExpr,
327        right: &SymExpr,
328        modulus: &SymExpr,
329    ) {
330        // if modulus == 0:
331        //   0
332        // else:
333        //   low_256((zext(left) op zext(right)) urem zext(modulus))
334        out.push_str("(ite (= ");
335        self.write_expr(out, modulus, None, None);
336        out.push_str(" (_ bv0 256)) (_ bv0 256) ((_ extract 255 0) (bvurem (");
337        out.push_str(op);
338        out.push_str(" ((_ zero_extend 256) ");
339        self.write_expr(out, left, None, None);
340        out.push_str(") ((_ zero_extend 256) ");
341        self.write_expr(out, right, None, None);
342        out.push_str(")) ((_ zero_extend 256) ");
343        self.write_expr(out, modulus, None, None);
344        out.push_str("))))");
345    }
346
347    fn write_bool(
348        &self,
349        out: &mut String,
350        expr: &SymBoolExpr,
351        skip_expr: Option<usize>,
352        skip_bool: Option<usize>,
353    ) {
354        if let Some(idx) = self.plan.bool_binding(expr)
355            && Some(idx) != skip_bool
356        {
357            SmtBinding::write_bool_name(out, idx);
358            return;
359        }
360
361        match expr.kind() {
362            SymBoolExprKind::Const(value) => out.push_str(if *value { "true" } else { "false" }),
363            SymBoolExprKind::Not(value) => {
364                out.push_str("(not ");
365                self.write_bool(out, value, skip_expr, skip_bool);
366                out.push(')');
367            }
368            SymBoolExprKind::And(values) => {
369                out.push_str("(and");
370                for value in values.iter() {
371                    out.push(' ');
372                    self.write_bool(out, value, skip_expr, skip_bool);
373                }
374                out.push(')');
375            }
376            SymBoolExprKind::Cmp(op, left, right) => {
377                let _ = write!(out, "({} ", op.smt());
378                self.write_expr(out, left, skip_expr, skip_bool);
379                out.push(' ');
380                self.write_expr(out, right, skip_expr, skip_bool);
381                out.push(')');
382            }
383        }
384    }
385}