Skip to main content

forge_lint/sol/low/
missing_events_arithmetic.rs

1use super::MissingEventsArithmetic;
2use crate::{
3    linter::{LateLintPass, LintContext},
4    sol::{
5        Severity, SolLint,
6        analysis::{
7            dispatched_function, is_protected, lhs_local_var, loop_stmts, state_lhs_vars,
8            underlying_var,
9        },
10    },
11};
12use solar::{
13    ast::{ContractKind, StateMutability},
14    data_structures::map::FxIndexSet,
15    interface::Span,
16    sema::{
17        Gcx,
18        builtins::Builtin,
19        hir::{
20            self, BinOpKind, ContractId, ElementaryType, Expr, ExprKind, FunctionId, StmtKind,
21            TypeKind, VariableId, Visit,
22        },
23    },
24};
25use std::{
26    collections::{HashMap, HashSet},
27    ops::ControlFlow,
28};
29
30declare_forge_lint!(
31    MISSING_EVENTS_ARITHMETIC,
32    Severity::Low,
33    "missing-events-arithmetic",
34    "critical arithmetic state changes without an event"
35);
36
37impl<'gcx> LateLintPass<'gcx> for MissingEventsArithmetic {
38    fn check_nested_contract(
39        &mut self,
40        ctx: &LintContext,
41        gcx: Gcx<'gcx>,
42        contract_id: ContractId,
43    ) {
44        let contract = gcx.hir.contract(contract_id);
45        if contract.kind != ContractKind::Contract || contract.linearization_failed() {
46            return;
47        }
48
49        // State variables and their entry points are commonly split across a base/derived pair,
50        // so candidates come from the whole inheritance chain.
51        let candidates: HashSet<_> = contract
52            .linearized_bases
53            .iter()
54            .flat_map(|&cid| gcx.hir.contract(cid).variables())
55            .filter(|&id| {
56                let var = gcx.hir.variable(id);
57                var.kind.is_state()
58                    && !var.is_constant()
59                    && !var.is_immutable()
60                    && matches!(
61                        var.ty.kind,
62                        TypeKind::Elementary(ElementaryType::Int(_) | ElementaryType::UInt(_))
63                    )
64            })
65            .collect();
66        if candidates.is_empty() {
67            return;
68        }
69
70        // The externally reachable functions, with overridden base implementations resolved to
71        // the most derived one.
72        let (protected, unprotected): (Vec<_>, Vec<_>) = gcx
73            .interface_functions(contract_id)
74            .all()
75            .iter()
76            .map(|func| func.id)
77            .partition(|&id| is_protected(gcx, id));
78        let entry_points: Vec<_> = protected
79            .into_iter()
80            .filter(|&id| {
81                !matches!(
82                    gcx.hir.function(id).state_mutability,
83                    StateMutability::Pure | StateMutability::View
84                )
85            })
86            .collect();
87        if entry_points.is_empty() {
88            return;
89        }
90
91        // Candidates that flow into arithmetic reachable from an unprotected function.
92        let mut uses = UseAnalyzer {
93            gcx,
94            contract_id,
95            targets: &candidates,
96            mode: Mode::Uses,
97            taint: HashMap::new(),
98            used: HashSet::new(),
99            returned: HashSet::new(),
100            call_stack: Vec::new(),
101        };
102        for func_id in unprotected {
103            uses.taint.clear();
104            uses.analyze_function(func_id);
105        }
106        if uses.used.is_empty() {
107            return;
108        }
109
110        for func_id in entry_points {
111            let mut analyzer =
112                WriteAnalyzer { gcx, contract_id, targets: &uses.used, call_stack: Vec::new() };
113            let mut emitted = HashSet::new();
114            for write in analyzer.analyze_entry_point(func_id) {
115                if !emitted.insert(write.var_id) {
116                    continue;
117                }
118                let name = gcx
119                    .hir
120                    .variable(write.var_id)
121                    .name
122                    .map_or_else(|| "state variable".to_string(), |name| name.to_string());
123                ctx.emit_with_msg(
124                    &MISSING_EVENTS_ARITHMETIC,
125                    write.span,
126                    format!("`{name}` is changed without an event but is used in arithmetic"),
127                );
128            }
129        }
130    }
131}
132
133const fn is_arithmetic_op(kind: BinOpKind) -> bool {
134    matches!(
135        kind,
136        BinOpKind::Add
137            | BinOpKind::Sub
138            | BinOpKind::Mul
139            | BinOpKind::Div
140            | BinOpKind::Rem
141            | BinOpKind::Pow
142    )
143}
144
145// --- Arithmetic uses --------------------------------------------------------------------------
146
147#[derive(Clone, Copy, PartialEq, Eq)]
148enum Mode {
149    /// Record target variables that reach arithmetic operators.
150    Uses,
151    /// Record target variables that reach `return` statements.
152    Returns,
153}
154
155/// Finds target state variables that flow into arithmetic, following locals and internal calls.
156struct UseAnalyzer<'a, 'gcx> {
157    gcx: Gcx<'gcx>,
158    contract_id: ContractId,
159    targets: &'a HashSet<VariableId>,
160    mode: Mode,
161    /// Target state variables each local may currently hold.
162    taint: HashMap<VariableId, HashSet<VariableId>>,
163    used: HashSet<VariableId>,
164    returned: HashSet<VariableId>,
165    call_stack: Vec<FunctionId>,
166}
167
168impl<'gcx> UseAnalyzer<'_, 'gcx> {
169    fn analyze_function(&mut self, func_id: FunctionId) {
170        if self.call_stack.contains(&func_id) {
171            return;
172        }
173        let Some(body) = self.gcx.hir.function(func_id).body else { return };
174        self.call_stack.push(func_id);
175        for stmt in body.stmts {
176            let _ = self.visit_stmt(stmt);
177        }
178        self.call_stack.pop();
179    }
180
181    /// Analyzes `callee_id` with its parameters bound to the call's argument sources, restoring the
182    /// caller's taint afterwards.
183    fn analyze_call(&mut self, callee_id: FunctionId, call: &Expr<'gcx>) {
184        if self.call_stack.contains(&callee_id) {
185            return;
186        }
187        let params = self
188            .gcx
189            .hir
190            .function(callee_id)
191            .parameters
192            .iter()
193            .enumerate()
194            .filter_map(|(index, &param)| {
195                let sources = self.sources(self.gcx.call_arg(call, index)?);
196                (!sources.is_empty()).then_some((param, sources))
197            })
198            .collect();
199        let saved = std::mem::replace(&mut self.taint, params);
200        self.analyze_function(callee_id);
201        self.taint = saved;
202    }
203
204    /// Target state variables `expr` may evaluate to, including through helper return values.
205    fn sources(&mut self, expr: &Expr<'gcx>) -> HashSet<VariableId> {
206        let mut out = HashSet::new();
207        let _ = expr.visit(&mut |e| {
208            if let Some(var_id) = underlying_var(self.gcx, e) {
209                if self.targets.contains(&var_id) {
210                    out.insert(var_id);
211                }
212                if let Some(sources) = self.taint.get(&var_id) {
213                    out.extend(sources);
214                }
215            }
216            if let ExprKind::Call(callee, ..) = &e.kind
217                && let Some(callee_id) = dispatched_function(self.gcx, self.contract_id, callee)
218            {
219                out.extend(self.return_sources(callee_id, e));
220            }
221            ControlFlow::<()>::Continue(())
222        });
223        out
224    }
225
226    fn return_sources(&mut self, callee_id: FunctionId, call: &Expr<'gcx>) -> HashSet<VariableId> {
227        let outer_mode = std::mem::replace(&mut self.mode, Mode::Returns);
228        let outer_returned = std::mem::take(&mut self.returned);
229        self.analyze_call(callee_id, call);
230        self.mode = outer_mode;
231        std::mem::replace(&mut self.returned, outer_returned)
232    }
233
234    fn set_taint(&mut self, var_id: VariableId, sources: HashSet<VariableId>) {
235        if sources.is_empty() {
236            self.taint.remove(&var_id);
237        } else {
238            self.taint.insert(var_id, sources);
239        }
240    }
241}
242
243impl<'gcx> Visit<'gcx> for UseAnalyzer<'_, 'gcx> {
244    type BreakValue = solar::interface::data_structures::Never;
245
246    fn hir(&self) -> &'gcx hir::Hir<'gcx> {
247        &self.gcx.hir
248    }
249
250    fn visit_stmt(&mut self, stmt: &'gcx hir::Stmt<'gcx>) -> ControlFlow<Self::BreakValue> {
251        match stmt.kind {
252            StmtKind::DeclSingle(var_id) => {
253                if let Some(init) = self.gcx.hir.variable(var_id).initializer {
254                    let sources = self.sources(init);
255                    self.set_taint(var_id, sources);
256                }
257            }
258            StmtKind::DeclMulti(vars, expr) => {
259                let sources = self.sources(expr);
260                for var_id in vars.iter().flatten() {
261                    self.set_taint(*var_id, sources.clone());
262                }
263            }
264            StmtKind::Return(Some(expr)) if self.mode == Mode::Returns => {
265                let sources = self.sources(expr);
266                self.returned.extend(sources);
267            }
268            _ => {}
269        }
270        self.walk_stmt(stmt)
271    }
272
273    fn visit_expr(&mut self, expr: &'gcx Expr<'gcx>) -> ControlFlow<Self::BreakValue> {
274        match &expr.kind {
275            ExprKind::Assign(lhs, _, rhs) => {
276                if let Some(local) = lhs_local_var(self.gcx, lhs) {
277                    let sources = self.sources(rhs);
278                    self.set_taint(local, sources);
279                }
280            }
281            ExprKind::Binary(lhs, op, rhs)
282                if self.mode == Mode::Uses && is_arithmetic_op(op.kind) =>
283            {
284                let sources = self.sources(lhs);
285                self.used.extend(sources);
286                let sources = self.sources(rhs);
287                self.used.extend(sources);
288            }
289            ExprKind::Call(callee, ..) if self.mode == Mode::Uses => {
290                self.walk_expr(expr)?;
291                if let Some(callee_id) = dispatched_function(self.gcx, self.contract_id, callee) {
292                    self.analyze_call(callee_id, expr);
293                }
294                return ControlFlow::Continue(());
295            }
296            _ => {}
297        }
298        self.walk_expr(expr)
299    }
300}
301
302// --- Writes without events --------------------------------------------------------------------
303
304#[derive(Clone, Copy, PartialEq, Eq, Hash)]
305struct StateWrite {
306    var_id: VariableId,
307    span: Span,
308}
309
310/// Analysis state along one control-flow path.
311#[derive(Clone, Default)]
312struct WriteState {
313    /// Locals holding a value that is not a compile-time constant.
314    dynamic: HashSet<VariableId>,
315    /// Unique target writes not yet followed by an `emit`, in first-seen order.
316    /// Deduplication prevents branch joins from multiplying identical pending writes.
317    writes: FxIndexSet<StateWrite>,
318}
319
320fn merge(lhs: Option<WriteState>, rhs: Option<WriteState>) -> Option<WriteState> {
321    match (lhs, rhs) {
322        (Some(mut lhs), Some(rhs)) => {
323            lhs.dynamic.extend(rhs.dynamic);
324            lhs.writes.extend(rhs.writes);
325            Some(lhs)
326        }
327        (lhs, rhs) => lhs.or(rhs),
328    }
329}
330
331/// Paths leaving a statement: those continuing to the next statement and those that `return`ed
332/// (which skip the rest of the body but still run the modifiers' trailing code).
333#[derive(Default)]
334struct Flow {
335    fallthrough: Option<WriteState>,
336    returned: Option<WriteState>,
337}
338
339impl Flow {
340    const fn fallthrough(state: WriteState) -> Self {
341        Self { fallthrough: Some(state), returned: None }
342    }
343
344    fn merge(self, other: Self) -> Self {
345        Self {
346            fallthrough: merge(self.fallthrough, other.fallthrough),
347            returned: merge(self.returned, other.returned),
348        }
349    }
350
351    fn merged(self) -> Option<WriteState> {
352        merge(self.fallthrough, self.returned)
353    }
354}
355
356/// Collects writes to target variables that no later `emit` on the same path covers.
357struct WriteAnalyzer<'a, 'gcx> {
358    gcx: Gcx<'gcx>,
359    contract_id: ContractId,
360    targets: &'a HashSet<VariableId>,
361    call_stack: Vec<FunctionId>,
362}
363
364impl<'gcx> WriteAnalyzer<'_, 'gcx> {
365    fn analyze_entry_point(&mut self, func_id: FunctionId) -> FxIndexSet<StateWrite> {
366        let func = self.gcx.hir.function(func_id);
367        let state = WriteState {
368            dynamic: func.parameters.iter().copied().collect(),
369            writes: FxIndexSet::default(),
370        };
371        let mut state = self.analyze_function(func_id, state).merged();
372        // Modifier code after `_` runs once the body finished, innermost modifier first, and may
373        // still emit for the body's writes.
374        for modifier in func.modifiers.iter().rev() {
375            let Some(body) =
376                modifier.id.as_function().and_then(|id| self.gcx.hir.function(id).body)
377            else {
378                continue;
379            };
380            let Some(pos) = body.stmts.iter().position(|s| matches!(s.kind, StmtKind::Placeholder))
381            else {
382                continue;
383            };
384            let suffix = &body.stmts[pos + 1..];
385            state = state.and_then(|state| self.analyze_stmts(suffix, state).merged());
386        }
387        state.map(|state| state.writes).unwrap_or_default()
388    }
389
390    fn analyze_function(&mut self, func_id: FunctionId, state: WriteState) -> Flow {
391        if self.call_stack.contains(&func_id) {
392            return Flow::fallthrough(state);
393        }
394        let Some(body) = self.gcx.hir.function(func_id).body else {
395            return Flow::fallthrough(state);
396        };
397        self.call_stack.push(func_id);
398        let flow = self.analyze_stmts(body.stmts, state);
399        self.call_stack.pop();
400        flow
401    }
402
403    fn analyze_stmts(
404        &mut self,
405        stmts: impl IntoIterator<Item = &'gcx hir::Stmt<'gcx>>,
406        state: WriteState,
407    ) -> Flow {
408        let mut flow = Flow::fallthrough(state);
409        for stmt in stmts {
410            let Some(state) = flow.fallthrough.take() else { break };
411            let next = self.analyze_stmt(stmt, state);
412            flow.fallthrough = next.fallthrough;
413            flow.returned = merge(flow.returned, next.returned);
414        }
415        flow
416    }
417
418    fn analyze_stmt(&mut self, stmt: &'gcx hir::Stmt<'gcx>, mut state: WriteState) -> Flow {
419        match stmt.kind {
420            StmtKind::DeclSingle(var_id) => {
421                if let Some(init) = self.gcx.hir.variable(var_id).initializer {
422                    self.analyze_expr(init, &mut state);
423                    self.set_dynamic(&mut state, var_id, init);
424                }
425                Flow::fallthrough(state)
426            }
427            StmtKind::DeclMulti(vars, expr) => {
428                self.analyze_expr(expr, &mut state);
429                for var_id in vars.iter().flatten() {
430                    self.set_dynamic(&mut state, *var_id, expr);
431                }
432                Flow::fallthrough(state)
433            }
434            StmtKind::Block(block) | StmtKind::UncheckedBlock(block) => {
435                self.analyze_stmts(block.stmts, state)
436            }
437            StmtKind::Loop(block, source) => self.analyze_stmts(loop_stmts(block, source), state),
438            StmtKind::If(cond, then_stmt, else_stmt) => {
439                self.analyze_expr(cond, &mut state);
440                let then_flow = self.analyze_stmt(then_stmt, state.clone());
441                let else_flow = match else_stmt {
442                    Some(else_stmt) => self.analyze_stmt(else_stmt, state),
443                    None => Flow::fallthrough(state),
444                };
445                then_flow.merge(else_flow)
446            }
447            StmtKind::Try(try_stmt) => {
448                self.analyze_expr(&try_stmt.expr, &mut state);
449                try_stmt.clauses.iter().fold(Flow::default(), |flow, clause| {
450                    flow.merge(self.analyze_stmts(clause.block.stmts, state.clone()))
451                })
452            }
453            StmtKind::Expr(expr) => {
454                self.analyze_expr(expr, &mut state);
455                Flow::fallthrough(state)
456            }
457            StmtKind::Revert(expr) => {
458                self.analyze_expr(expr, &mut state);
459                Flow::default()
460            }
461            StmtKind::Emit(expr) => {
462                self.analyze_expr(expr, &mut state);
463                state.writes.clear();
464                Flow::fallthrough(state)
465            }
466            StmtKind::Return(expr) => {
467                if let Some(expr) = expr {
468                    self.analyze_expr(expr, &mut state);
469                }
470                Flow { fallthrough: None, returned: Some(state) }
471            }
472            _ => Flow::fallthrough(state),
473        }
474    }
475
476    fn analyze_expr(&mut self, expr: &'gcx Expr<'gcx>, state: &mut WriteState) {
477        let _ = expr.visit(&mut |e| {
478            match &e.kind {
479                ExprKind::Assign(lhs, op, rhs) => {
480                    let dynamic = self.is_dynamic(state, rhs);
481                    if dynamic || op.is_some_and(|op| is_arithmetic_op(op.kind)) {
482                        self.record_writes(state, lhs);
483                    }
484                    if let Some(local) = lhs_local_var(self.gcx, lhs) {
485                        self.set_dynamic(state, local, rhs);
486                    }
487                }
488                ExprKind::Unary(op, inner) if op.kind.has_side_effects() => {
489                    self.record_writes(state, inner);
490                }
491                ExprKind::Call(callee, ..) => {
492                    if let Some(callee_id) = dispatched_function(self.gcx, self.contract_id, callee)
493                    {
494                        self.analyze_call(callee_id, e, state);
495                    }
496                }
497                _ => {}
498            }
499            ControlFlow::<()>::Continue(())
500        });
501    }
502
503    /// Inlines `callee_id` with its parameters marked dynamic when the matching argument is; the
504    /// callee's pending writes (and any `emit` clearing them) flow back into the caller.
505    fn analyze_call(&mut self, callee_id: FunctionId, call: &Expr<'gcx>, state: &mut WriteState) {
506        let callee_state = WriteState {
507            dynamic: self
508                .gcx
509                .hir
510                .function(callee_id)
511                .parameters
512                .iter()
513                .enumerate()
514                .filter(|(index, _)| {
515                    self.gcx.call_arg(call, *index).is_some_and(|arg| self.is_dynamic(state, arg))
516                })
517                .map(|(_, &param)| param)
518                .collect(),
519            writes: state.writes.clone(),
520        };
521        if let Some(merged) = self.analyze_function(callee_id, callee_state).merged() {
522            state.writes = merged.writes;
523        }
524    }
525
526    fn record_writes(&self, state: &mut WriteState, lhs: &Expr<'_>) {
527        for var_id in state_lhs_vars(self.gcx, lhs) {
528            if self.targets.contains(&var_id) {
529                state.writes.insert(StateWrite { var_id, span: lhs.span });
530            }
531        }
532    }
533
534    fn set_dynamic(&self, state: &mut WriteState, var_id: VariableId, value: &Expr<'_>) {
535        if self.is_dynamic(state, value) {
536            state.dynamic.insert(var_id);
537        } else {
538            state.dynamic.remove(&var_id);
539        }
540    }
541
542    /// True unless `expr` is a compile-time constant: reads of mutable state, dynamic locals,
543    /// calls and `block`/`msg`/`tx` members are all dynamic.
544    fn is_dynamic(&self, state: &WriteState, expr: &Expr<'_>) -> bool {
545        expr.visit(&mut |e| {
546            let dynamic = match &e.kind {
547                ExprKind::Call(..) => true,
548                ExprKind::Member(base, _) => {
549                    matches!(
550                        self.gcx.resolved_builtin(base),
551                        Some(Builtin::Block | Builtin::Msg | Builtin::Tx)
552                    )
553                }
554                _ => underlying_var(self.gcx, e).is_some_and(|var_id| {
555                    let var = self.gcx.hir.variable(var_id);
556                    state.dynamic.contains(&var_id)
557                        || (var.kind.is_state() && !var.is_constant() && !var.is_immutable())
558                }),
559            };
560            if dynamic { ControlFlow::Break(()) } else { ControlFlow::Continue(()) }
561        })
562        .is_break()
563    }
564}
565
566#[cfg(test)]
567mod tests {
568    use super::*;
569    use solar::interface::BytePos;
570
571    #[test]
572    fn merge_deduplicates_writes_in_order() {
573        let first = StateWrite { var_id: VariableId::new(0), span: Span::DUMMY };
574        let another_span =
575            StateWrite { var_id: first.var_id, span: Span::new(BytePos(1), BytePos(2)) };
576        let another_var = StateWrite { var_id: VariableId::new(1), span: first.span };
577        let lhs = WriteState { writes: [first].into_iter().collect(), ..Default::default() };
578        let rhs = WriteState {
579            writes: [first, another_span, another_var].into_iter().collect(),
580            ..Default::default()
581        };
582        let mut state = merge(Some(lhs), Some(rhs)).unwrap();
583        let expected = [first, another_span, another_var].map(|write| (write.var_id, write.span));
584
585        // Each join must retain one copy per write site, even after repeated branching.
586        for _ in 0..32 {
587            assert_eq!(
588                state.writes.iter().map(|write| (write.var_id, write.span)).collect::<Vec<_>>(),
589                expected,
590            );
591            state = merge(Some(state.clone()), Some(state)).unwrap();
592        }
593    }
594}