Skip to main content

forge_lint/sol/med/
div_mul.rs

1use super::DivideBeforeMultiply;
2use crate::{
3    linter::{LateLintPass, LintContext},
4    sol::{
5        Severity, SolLint,
6        analysis::{is_revert_call, loop_update, tuple_elems},
7    },
8};
9use solar::sema::{
10    Gcx,
11    builtins::Builtin,
12    hir::{BinOpKind, Block, Expr, ExprKind, Function, Stmt, StmtKind, VariableId},
13};
14use std::collections::HashSet;
15
16declare_forge_lint!(
17    DIVIDE_BEFORE_MULTIPLY,
18    Severity::Med,
19    "divide-before-multiply",
20    "division before multiplication may lose precision"
21);
22
23/// Locals whose current value is the result of a division.
24type Tainted = HashSet<VariableId>;
25
26impl<'gcx> LateLintPass<'gcx> for DivideBeforeMultiply {
27    fn check_function(&mut self, ctx: &LintContext, gcx: Gcx<'gcx>, func: &'gcx Function<'gcx>) {
28        if let Some(body) = func.body {
29            check_block(ctx, gcx, body, &mut Tainted::new());
30        }
31    }
32}
33
34/// Checks `block`, returning `false` once control cannot continue past a statement.
35fn check_block<'gcx>(
36    ctx: &LintContext,
37    gcx: Gcx<'gcx>,
38    block: Block<'gcx>,
39    tainted: &mut Tainted,
40) -> bool {
41    block.stmts.iter().all(|stmt| check_stmt(ctx, gcx, stmt, tainted))
42}
43
44/// Checks the bodies of mutually exclusive branches and keeps the taint of every branch that
45/// falls through, on top of the taint before the branch.
46fn check_branches<'gcx>(
47    ctx: &LintContext,
48    gcx: Gcx<'gcx>,
49    blocks: impl Iterator<Item = Block<'gcx>>,
50    tainted: &mut Tainted,
51) {
52    let mut merged = Tainted::new();
53    for block in blocks {
54        let mut branch_tainted = tainted.clone();
55        if check_block(ctx, gcx, block, &mut branch_tainted) {
56            merged.extend(branch_tainted);
57        }
58    }
59    tainted.extend(merged);
60}
61
62fn check_stmt<'gcx>(
63    ctx: &LintContext,
64    gcx: Gcx<'gcx>,
65    stmt: &'gcx Stmt<'gcx>,
66    tainted: &mut Tainted,
67) -> bool {
68    match &stmt.kind {
69        StmtKind::DeclSingle(var_id) => {
70            if let Some(init) = gcx.hir.variable(*var_id).initializer {
71                check_expr(ctx, gcx, init, tainted);
72                set_taint(gcx, *var_id, is_division_or_tainted(gcx, init, tainted), tainted);
73            }
74            true
75        }
76        StmtKind::DeclMulti(vars, expr) => {
77            check_expr(ctx, gcx, expr, tainted);
78            for (var_id, is_tainted) in vars.iter().zip(rhs_taints(gcx, expr, vars.len(), tainted))
79            {
80                if let Some(var_id) = var_id {
81                    set_taint(gcx, *var_id, is_tainted, tainted);
82                }
83            }
84            true
85        }
86        StmtKind::Expr(expr) => {
87            check_expr(ctx, gcx, expr, tainted);
88            !is_revert_call(gcx, expr)
89        }
90        StmtKind::Emit(expr) => {
91            check_expr(ctx, gcx, expr, tainted);
92            true
93        }
94        StmtKind::Revert(expr) | StmtKind::Return(Some(expr)) => {
95            check_expr(ctx, gcx, expr, tainted);
96            false
97        }
98        StmtKind::Return(None) => false,
99        StmtKind::If(cond, then_stmt, else_stmt) => {
100            check_expr(ctx, gcx, cond, tainted);
101            let mut merged = Tainted::new();
102            let mut falls_through = false;
103            for branch in [Some(*then_stmt), *else_stmt] {
104                let mut branch_tainted = tainted.clone();
105                if branch.is_none_or(|stmt| check_stmt(ctx, gcx, stmt, &mut branch_tainted)) {
106                    merged.extend(branch_tainted);
107                    falls_through = true;
108                }
109            }
110            if falls_through {
111                *tainted = merged;
112            }
113            falls_through
114        }
115        StmtKind::Loop(block, source) => {
116            let mut branch = tainted.clone();
117            if check_block(ctx, gcx, *block, &mut branch)
118                && loop_update(*source)
119                    .is_none_or(|update| check_stmt(ctx, gcx, update, &mut branch))
120            {
121                tainted.extend(branch);
122            }
123            true
124        }
125        StmtKind::Try(try_stmt) => {
126            check_expr(ctx, gcx, &try_stmt.expr, tainted);
127            check_branches(ctx, gcx, try_stmt.clauses.iter().map(|c| c.block), tainted);
128            true
129        }
130        StmtKind::Switch(switch) => {
131            check_expr(ctx, gcx, switch.selector, tainted);
132            check_branches(ctx, gcx, switch.cases.iter().map(|c| c.body), tainted);
133            true
134        }
135        StmtKind::Block(block)
136        | StmtKind::UncheckedBlock(block)
137        | StmtKind::AssemblyBlock(block) => check_block(ctx, gcx, *block, tainted),
138        StmtKind::Break | StmtKind::Continue | StmtKind::Placeholder | StmtKind::Err(_) => true,
139    }
140}
141
142fn check_expr<'gcx>(
143    ctx: &LintContext,
144    gcx: Gcx<'gcx>,
145    expr: &'gcx Expr<'gcx>,
146    tainted: &mut Tainted,
147) {
148    match &expr.peel_parens().kind {
149        ExprKind::Assign(lhs, op, rhs) => {
150            check_expr(ctx, gcx, rhs, tainted);
151            check_expr(ctx, gcx, lhs, tainted);
152            match op.map(|op| op.kind) {
153                None => match tuple_elems(lhs) {
154                    Some(elems) => {
155                        for (lhs, is_tainted) in
156                            elems.iter().zip(rhs_taints(gcx, rhs, elems.len(), tainted))
157                        {
158                            if let Some(lhs) = lhs {
159                                set_lhs_taint(gcx, lhs, is_tainted, tainted);
160                            }
161                        }
162                    }
163                    None => {
164                        set_lhs_taint(gcx, lhs, is_division_or_tainted(gcx, rhs, tainted), tainted)
165                    }
166                },
167                Some(BinOpKind::Mul) => {
168                    let is_tainted = is_division_or_tainted(gcx, lhs, tainted)
169                        || is_division_or_tainted(gcx, rhs, tainted);
170                    if is_tainted {
171                        ctx.emit(&DIVIDE_BEFORE_MULTIPLY, expr.span);
172                    }
173                    set_lhs_taint(gcx, lhs, is_tainted, tainted);
174                }
175                Some(op) => set_lhs_taint(gcx, lhs, op == BinOpKind::Div, tainted),
176            }
177        }
178        ExprKind::Binary(left, op, right) => {
179            check_expr(ctx, gcx, left, tainted);
180            check_expr(ctx, gcx, right, tainted);
181            if op.kind == BinOpKind::Mul
182                && (is_division_or_tainted(gcx, left, tainted)
183                    || is_division_or_tainted(gcx, right, tainted))
184            {
185                ctx.emit(&DIVIDE_BEFORE_MULTIPLY, expr.span);
186            }
187        }
188        ExprKind::Call(callee, args) => {
189            let (callee, named_args) = callee.split_call_options();
190            check_expr(ctx, gcx, callee, tainted);
191            for arg in args.exprs() {
192                check_expr(ctx, gcx, arg, tainted);
193            }
194            for arg in named_args.iter().flat_map(|opts| opts.args) {
195                check_expr(ctx, gcx, &arg.value, tainted);
196            }
197            if is_yul_call(gcx, expr, &[Builtin::YulMul])
198                && args.exprs().any(|arg| is_division_or_tainted(gcx, arg, tainted))
199            {
200                ctx.emit(&DIVIDE_BEFORE_MULTIPLY, expr.span);
201            }
202        }
203        ExprKind::CallOptions(callee, options) => {
204            check_expr(ctx, gcx, callee, tainted);
205            for option in options.args {
206                check_expr(ctx, gcx, &option.value, tainted);
207            }
208        }
209        ExprKind::Ternary(cond, then_expr, else_expr) => {
210            check_expr(ctx, gcx, cond, tainted);
211            let mut then_tainted = tainted.clone();
212            check_expr(ctx, gcx, then_expr, &mut then_tainted);
213            check_expr(ctx, gcx, else_expr, tainted);
214            tainted.extend(then_tainted);
215        }
216        ExprKind::Unary(op, inner) => {
217            check_expr(ctx, gcx, inner, tainted);
218            if op.kind.has_side_effects() {
219                set_lhs_taint(gcx, inner, false, tainted);
220            }
221        }
222        ExprKind::Array(exprs) => exprs.iter().for_each(|e| check_expr(ctx, gcx, e, tainted)),
223        ExprKind::Tuple(exprs) => {
224            exprs.iter().flatten().for_each(|e| check_expr(ctx, gcx, e, tainted))
225        }
226        ExprKind::Index(base, index) => {
227            check_expr(ctx, gcx, base, tainted);
228            if let Some(index) = index {
229                check_expr(ctx, gcx, index, tainted);
230            }
231        }
232        ExprKind::Slice(base, start, end) => {
233            check_expr(ctx, gcx, base, tainted);
234            start.iter().chain(end).for_each(|e| check_expr(ctx, gcx, e, tainted));
235        }
236        ExprKind::Delete(inner)
237        | ExprKind::Member(inner, _)
238        | ExprKind::YulMember(inner, _)
239        | ExprKind::Payable(inner) => check_expr(ctx, gcx, inner, tainted),
240        ExprKind::Ident(_)
241        | ExprKind::Lit(_)
242        | ExprKind::New(_)
243        | ExprKind::TypeCall(_)
244        | ExprKind::Type(_)
245        | ExprKind::Err(_) => {}
246    }
247}
248
249/// Taint of each of the `n` slots assigned from `rhs`: elementwise for a tuple of matching arity,
250/// otherwise the taint of the whole expression.
251fn rhs_taints(gcx: Gcx<'_>, rhs: &Expr<'_>, n: usize, tainted: &Tainted) -> Vec<bool> {
252    match tuple_elems(rhs) {
253        Some(elems) if elems.len() == n => elems
254            .iter()
255            .map(|e| e.is_some_and(|e| is_division_or_tainted(gcx, e, tainted)))
256            .collect(),
257        _ => vec![is_division_or_tainted(gcx, rhs, tainted); n],
258    }
259}
260
261fn set_lhs_taint(gcx: Gcx<'_>, lhs: &Expr<'_>, is_tainted: bool, tainted: &mut Tainted) {
262    match &lhs.peel_parens().kind {
263        ExprKind::Ident(_) => {
264            if let Some(var_id) = gcx.resolved_variable(lhs) {
265                set_taint(gcx, var_id, is_tainted, tainted);
266            }
267        }
268        ExprKind::Tuple(exprs) => {
269            exprs.iter().flatten().for_each(|e| set_lhs_taint(gcx, e, is_tainted, tainted))
270        }
271        _ => {}
272    }
273}
274
275fn set_taint(gcx: Gcx<'_>, var_id: VariableId, is_tainted: bool, tainted: &mut Tainted) {
276    if gcx.hir.variable(var_id).is_local_or_return() {
277        if is_tainted {
278            tainted.insert(var_id);
279        } else {
280            tainted.remove(&var_id);
281        }
282    }
283}
284
285/// The value of `expr` is a division result, directly or through a tainted local.
286fn is_division_or_tainted(gcx: Gcx<'_>, expr: &Expr<'_>, tainted: &Tainted) -> bool {
287    match &expr.peel_parens().kind {
288        ExprKind::Binary(_, op, _) => op.kind == BinOpKind::Div,
289        ExprKind::Ident(_) => gcx.resolved_variable(expr).is_some_and(|v| tainted.contains(&v)),
290        ExprKind::Call(..) => is_yul_call(gcx, expr, &[Builtin::YulDiv, Builtin::YulSdiv]),
291        ExprKind::YulMember(inner, _) => is_division_or_tainted(gcx, inner, tainted),
292        _ => false,
293    }
294}
295
296/// A two-argument call to one of the given Yul builtins.
297fn is_yul_call(gcx: Gcx<'_>, expr: &Expr<'_>, candidates: &[Builtin]) -> bool {
298    matches!(&expr.peel_parens().kind, ExprKind::Call(callee, args)
299        if args.len() == 2 && gcx.resolved_builtin(callee).is_some_and(|b| candidates.contains(&b)))
300}