Skip to main content

foundry_evm_coverage/
analysis.rs

1use super::{CoverageItem, CoverageItemKind, SourceLocation};
2use alloy_primitives::map::HashMap;
3use foundry_common::TestFunctionExt;
4use foundry_compilers::ProjectCompileOutput;
5use rayon::prelude::*;
6use solar::{
7    ast::{self, ExprKind, ItemKind, StmtKind, yul},
8    data_structures::{Never, map::FxHashSet},
9    interface::{BytePos, Span},
10    sema::{Gcx, hir},
11};
12use std::{
13    ops::{ControlFlow, Range},
14    path::PathBuf,
15    sync::Arc,
16};
17
18/// A visitor that walks the AST of a single contract and finds coverage items.
19#[derive(Clone)]
20struct SourceVisitor<'gcx> {
21    /// The source ID of the contract.
22    source_id: u32,
23    /// The solar session for span resolution.
24    gcx: Gcx<'gcx>,
25
26    /// The name of the contract being walked.
27    contract_name: Arc<str>,
28
29    /// The current branch ID
30    branch_id: u32,
31
32    /// Coverage items
33    items: Vec<CoverageItem>,
34
35    all_lines: Vec<u32>,
36    function_calls: Vec<Span>,
37    function_calls_set: FxHashSet<Span>,
38}
39
40struct SourceVisitorCheckpoint {
41    items: usize,
42    all_lines: usize,
43    function_calls: usize,
44}
45
46impl<'gcx> SourceVisitor<'gcx> {
47    fn new(source_id: u32, gcx: Gcx<'gcx>) -> Self {
48        Self {
49            source_id,
50            gcx,
51            contract_name: Arc::default(),
52            branch_id: 0,
53            all_lines: Default::default(),
54            function_calls: Default::default(),
55            function_calls_set: Default::default(),
56            items: Default::default(),
57        }
58    }
59
60    const fn checkpoint(&self) -> SourceVisitorCheckpoint {
61        SourceVisitorCheckpoint {
62            items: self.items.len(),
63            all_lines: self.all_lines.len(),
64            function_calls: self.function_calls.len(),
65        }
66    }
67
68    fn restore_checkpoint(&mut self, checkpoint: SourceVisitorCheckpoint) {
69        let SourceVisitorCheckpoint { items, all_lines, function_calls } = checkpoint;
70        self.items.truncate(items);
71        self.all_lines.truncate(all_lines);
72        self.function_calls.truncate(function_calls);
73    }
74
75    fn visit_contract<'ast>(&mut self, contract: &'ast ast::ItemContract<'ast>) {
76        let _ = ast::Visit::visit_item_contract(self, contract);
77    }
78
79    /// Returns `true` if the contract has any test functions.
80    fn has_tests(&self, checkpoint: &SourceVisitorCheckpoint) -> bool {
81        self.items[checkpoint.items..].iter().any(|item| {
82            if let CoverageItemKind::Function { name } = &item.kind {
83                name.is_any_test()
84            } else {
85                false
86            }
87        })
88    }
89
90    /// Disambiguate functions with the same name in the same contract.
91    fn disambiguate_functions(&mut self) {
92        let mut dups = HashMap::<_, Vec<usize>>::default();
93        for (i, item) in self.items.iter().enumerate() {
94            if let CoverageItemKind::Function { name } = &item.kind {
95                dups.entry(name.clone()).or_default().push(i);
96            }
97        }
98        for dups in dups.values() {
99            if dups.len() > 1 {
100                for (i, &dup) in dups.iter().enumerate() {
101                    let item = &mut self.items[dup];
102                    if let CoverageItemKind::Function { name } = &item.kind {
103                        item.kind =
104                            CoverageItemKind::Function { name: format!("{name}.{i}").into() };
105                    }
106                }
107            }
108        }
109    }
110
111    fn resolve_function_calls(&mut self, hir_source_id: hir::SourceId) {
112        self.function_calls_set = self.function_calls.iter().copied().collect();
113        let _ = hir::Visit::visit_nested_source(self, hir_source_id);
114    }
115
116    fn sort(&mut self) {
117        self.items.sort();
118    }
119
120    fn push_lines(&mut self) {
121        self.all_lines.sort_unstable();
122        self.all_lines.dedup();
123        let mut lines = Vec::new();
124        for &line in &self.all_lines {
125            if let Some(reference_item) =
126                self.items.iter().find(|item| item.loc.lines.start == line)
127            {
128                lines.push(CoverageItem {
129                    kind: CoverageItemKind::Line,
130                    loc: reference_item.loc.clone(),
131                    anchor_loc: None,
132                    hits: 0,
133                });
134            }
135        }
136        self.items.extend(lines);
137    }
138
139    fn push_stmt(&mut self, span: Span) {
140        self.push_item_kind(CoverageItemKind::Statement, span);
141    }
142
143    /// Creates a coverage item for a given kind and source location. Pushes item to the internal
144    /// collection (plus additional coverage line if item is a statement).
145    fn push_item_kind(&mut self, kind: CoverageItemKind, span: Span) -> &mut CoverageItem {
146        let item =
147            CoverageItem { kind, loc: self.source_location_for(span), anchor_loc: None, hits: 0 };
148
149        debug_assert!(!matches!(item.kind, CoverageItemKind::Line));
150        self.all_lines.push(item.loc.lines.start);
151
152        self.items.push(item);
153        self.items.last_mut().unwrap()
154    }
155
156    fn source_location_for(&self, mut span: Span) -> SourceLocation {
157        // Statements' ranges in the solc source map do not include the semicolon.
158        if let Ok(snippet) = self.gcx.sess.source_map().span_to_snippet(span)
159            && let Some(stripped) = snippet.strip_suffix(';')
160        {
161            let stripped = stripped.trim_end();
162            let skipped = snippet.len() - stripped.len();
163            span = span.with_hi(span.hi() - BytePos::from_usize(skipped));
164        }
165
166        SourceLocation {
167            source_id: self.source_id as usize,
168            contract_name: self.contract_name.clone(),
169            bytes: self.byte_range(span),
170            lines: self.line_range(span),
171        }
172    }
173
174    fn byte_range(&self, span: Span) -> Range<u32> {
175        let bytes_usize = self.gcx.sess.source_map().span_to_source(span).unwrap().data;
176        bytes_usize.start as u32..bytes_usize.end as u32
177    }
178
179    fn line_range(&self, span: Span) -> Range<u32> {
180        let lines = self.gcx.sess.source_map().span_to_lines(span).unwrap().data;
181        assert!(!lines.is_empty());
182        let first = lines.first().unwrap();
183        let last = lines.last().unwrap();
184        first.line_index as u32 + 1..last.line_index as u32 + 2
185    }
186
187    const fn next_branch_id(&mut self) -> u32 {
188        let id = self.branch_id;
189        self.branch_id = id + 1;
190        id
191    }
192}
193
194impl<'ast> ast::Visit<'ast> for SourceVisitor<'_> {
195    type BreakValue = Never;
196
197    fn visit_item_contract(
198        &mut self,
199        contract: &'ast ast::ItemContract<'ast>,
200    ) -> ControlFlow<Self::BreakValue> {
201        self.contract_name = contract.name.as_str().into();
202        self.walk_item_contract(contract)
203    }
204
205    #[expect(clippy::single_match)]
206    fn visit_item(&mut self, item: &'ast ast::Item<'ast>) -> ControlFlow<Self::BreakValue> {
207        match &item.kind {
208            ItemKind::Function(func) => {
209                // TODO: We currently can only detect empty bodies in normal functions, not any of
210                // the other kinds: https://github.com/foundry-rs/foundry/issues/9458
211                if func.kind != ast::FunctionKind::Function && !has_statements(func.body.as_ref()) {
212                    return ControlFlow::Continue(());
213                }
214
215                let name = func.header.name.as_ref().map(|n| n.as_str()).unwrap_or_else(|| {
216                    match func.kind {
217                        ast::FunctionKind::Constructor => "constructor",
218                        ast::FunctionKind::Receive => "receive",
219                        ast::FunctionKind::Fallback => "fallback",
220                        ast::FunctionKind::Function | ast::FunctionKind::Modifier => unreachable!(),
221                    }
222                });
223
224                // Exclude function from coverage report if it is virtual without implementation.
225                let exclude_func = func.header.virtual_() && !func.is_implemented();
226                if !exclude_func {
227                    self.push_item_kind(
228                        CoverageItemKind::Function { name: name.into() },
229                        item.span,
230                    );
231                }
232
233                self.walk_item(item)?;
234            }
235            _ => {}
236        }
237        // Only walk functions.
238        ControlFlow::Continue(())
239    }
240
241    fn visit_stmt(&mut self, stmt: &'ast ast::Stmt<'ast>) -> ControlFlow<Self::BreakValue> {
242        match &stmt.kind {
243            StmtKind::Break | StmtKind::Continue | StmtKind::Emit(..) | StmtKind::Revert(..) => {
244                self.push_stmt(stmt.span);
245                // TODO(dani): these probably shouldn't be excluded.
246                return ControlFlow::Continue(());
247            }
248            StmtKind::Return(_) | StmtKind::DeclSingle(_) | StmtKind::DeclMulti(..) => {
249                self.push_stmt(stmt.span);
250            }
251
252            StmtKind::If(_cond, then_stmt, else_stmt) => {
253                let branch_id = self.next_branch_id();
254
255                // Add branch coverage items only if one of true/branch bodies contains statements.
256                if stmt_has_statements(then_stmt)
257                    || else_stmt.as_ref().is_some_and(|s| stmt_has_statements(s))
258                {
259                    // The branch instruction is mapped to the first opcode within the true
260                    // body source range.
261                    self.push_item_kind(
262                        CoverageItemKind::Branch { branch_id, path_id: 0, is_first_opcode: true },
263                        then_stmt.span,
264                    );
265                    if let Some(else_stmt) = else_stmt {
266                        let is_first_opcode = stmt_has_statements(else_stmt);
267                        let anchor_loc =
268                            is_first_opcode.then(|| self.source_location_for(else_stmt.span));
269                        self.push_item_kind(
270                            CoverageItemKind::Branch { branch_id, path_id: 1, is_first_opcode },
271                            stmt.span,
272                        )
273                        .anchor_loc = anchor_loc;
274                    }
275                }
276            }
277
278            StmtKind::Try(ast::StmtTry { expr: _, clauses }) => {
279                let branch_id = self.next_branch_id();
280
281                let mut path_id = 0;
282                for catch in clauses.iter() {
283                    let ast::TryCatchClause { span, name: _, args, block } = catch;
284                    let span = if path_id == 0 { stmt.span.to(*span) } else { *span };
285                    if path_id == 0 || has_statements(Some(block)) {
286                        self.push_item_kind(
287                            CoverageItemKind::Branch { branch_id, path_id, is_first_opcode: true },
288                            span,
289                        );
290                        path_id += 1;
291                    } else if !args.is_empty() {
292                        // Add coverage for clause with parameters and empty statements.
293                        // (`catch (bytes memory reason) {}`).
294                        // Catch all clause without statements is ignored (`catch {}`).
295                        self.push_stmt(span);
296                    }
297                }
298            }
299
300            // Skip placeholder statements as they are never referenced in source maps.
301            StmtKind::Assembly(_)
302            | StmtKind::Block(_)
303            | StmtKind::UncheckedBlock(_)
304            | StmtKind::Placeholder
305            | StmtKind::Expr(_)
306            | StmtKind::While(..)
307            | StmtKind::DoWhile(..)
308            | StmtKind::For { .. } => {}
309        }
310        self.walk_stmt(stmt)
311    }
312
313    fn visit_expr(&mut self, expr: &'ast ast::Expr<'ast>) -> ControlFlow<Self::BreakValue> {
314        match &expr.kind {
315            ExprKind::Assign(..)
316            | ExprKind::Unary(..)
317            | ExprKind::Binary(..)
318            | ExprKind::Ternary(..) => {
319                self.push_stmt(expr.span);
320                if matches!(expr.kind, ExprKind::Binary(..)) {
321                    return self.walk_expr(expr);
322                }
323            }
324            ExprKind::Call(callee, _args) => {
325                // Resolve later.
326                self.function_calls.push(expr.span);
327
328                if let ExprKind::Ident(ident) = &callee.kind {
329                    // Might be a require call, add branch coverage.
330                    // Asserts should not be considered branches: <https://github.com/foundry-rs/foundry/issues/9460>.
331                    if ident.as_str() == "require" {
332                        let branch_id = self.next_branch_id();
333                        self.push_item_kind(
334                            CoverageItemKind::Branch {
335                                branch_id,
336                                path_id: 0,
337                                is_first_opcode: false,
338                            },
339                            expr.span,
340                        );
341                        self.push_item_kind(
342                            CoverageItemKind::Branch {
343                                branch_id,
344                                path_id: 1,
345                                is_first_opcode: false,
346                            },
347                            expr.span,
348                        );
349                    }
350                }
351            }
352            _ => {}
353        }
354        // Intentionally do not walk all expressions.
355        ControlFlow::Continue(())
356    }
357
358    fn visit_yul_stmt(&mut self, stmt: &'ast yul::Stmt<'ast>) -> ControlFlow<Self::BreakValue> {
359        use yul::StmtKind;
360        match &stmt.kind {
361            StmtKind::VarDecl(..)
362            | StmtKind::AssignSingle(..)
363            | StmtKind::AssignMulti(..)
364            | StmtKind::Leave
365            | StmtKind::Break
366            | StmtKind::Continue => {
367                self.push_stmt(stmt.span);
368                // Don't walk assignments.
369                return ControlFlow::Continue(());
370            }
371            StmtKind::If(..) => {
372                let branch_id = self.next_branch_id();
373                self.push_item_kind(
374                    CoverageItemKind::Branch { branch_id, path_id: 0, is_first_opcode: false },
375                    stmt.span,
376                );
377            }
378            StmtKind::For(yul::StmtFor { body, .. }) => {
379                self.push_stmt(body.span);
380            }
381            StmtKind::Switch(switch) => {
382                for case in switch.cases.iter() {
383                    self.push_stmt(case.span);
384                    self.push_stmt(case.body.span);
385                }
386            }
387            StmtKind::FunctionDef(func) => {
388                let name = func.name.as_str();
389                self.push_item_kind(CoverageItemKind::Function { name: name.into() }, stmt.span);
390            }
391            // TODO(dani): merge with Block below on next solar release: https://github.com/paradigmxyz/solar/pull/496
392            StmtKind::Expr(_) => {
393                self.push_stmt(stmt.span);
394                return ControlFlow::Continue(());
395            }
396            StmtKind::Block(_) => {}
397        }
398        self.walk_yul_stmt(stmt)
399    }
400
401    fn visit_yul_expr(&mut self, expr: &'ast yul::Expr<'ast>) -> ControlFlow<Self::BreakValue> {
402        use yul::ExprKind;
403        match &expr.kind {
404            ExprKind::Path(_) | ExprKind::Lit(_) => {}
405            ExprKind::Call(_) => self.push_stmt(expr.span),
406        }
407        // Intentionally do not walk all expressions.
408        ControlFlow::Continue(())
409    }
410}
411
412impl<'gcx> hir::Visit<'gcx> for SourceVisitor<'gcx> {
413    type BreakValue = Never;
414
415    fn hir(&self) -> &'gcx hir::Hir<'gcx> {
416        &self.gcx.hir
417    }
418
419    fn visit_expr(&mut self, expr: &'gcx hir::Expr<'gcx>) -> ControlFlow<Self::BreakValue> {
420        if let hir::ExprKind::Call(lhs, ..) = &expr.kind
421            && self.function_calls_set.contains(&expr.span)
422            && is_regular_call(lhs)
423        {
424            self.push_stmt(expr.span);
425        }
426        self.walk_expr(expr)
427    }
428}
429
430// https://github.com/argotorg/solidity/blob/965166317bbc2b02067eb87f222a2dce9d24e289/libsolidity/ast/ASTAnnotations.h#L336-L341
431// https://github.com/argotorg/solidity/blob/965166317bbc2b02067eb87f222a2dce9d24e289/libsolidity/analysis/TypeChecker.cpp#L2720
432fn is_regular_call(lhs: &hir::Expr<'_>) -> bool {
433    match lhs.peel_parens().kind {
434        // StructConstructorCall
435        hir::ExprKind::Ident([hir::Res::Item(hir::ItemId::Struct(_))]) => false,
436        // TypeConversion
437        hir::ExprKind::Type(_) => false,
438        _ => true,
439    }
440}
441
442fn has_statements(block: Option<&ast::Block<'_>>) -> bool {
443    block.is_some_and(|block| !block.is_empty())
444}
445
446fn stmt_has_statements(stmt: &ast::Stmt<'_>) -> bool {
447    match &stmt.kind {
448        StmtKind::Assembly(a) => !a.block.is_empty(),
449        StmtKind::Block(b) | StmtKind::UncheckedBlock(b) => has_statements(Some(b)),
450        _ => true,
451    }
452}
453
454/// Coverage source analysis.
455#[derive(Clone, Debug, Default)]
456pub struct SourceAnalysis {
457    /// All the coverage items.
458    all_items: Vec<CoverageItem>,
459    /// Source ID to `(offset, len)` into `all_items`.
460    map: Vec<(u32, u32)>,
461}
462
463impl SourceAnalysis {
464    /// Analyzes contracts in the sources held by the source analyzer.
465    ///
466    /// Coverage items are found by:
467    /// - Walking the AST of each contract (except interfaces)
468    /// - Recording the items of each contract
469    ///
470    /// Each coverage item contains relevant information to find opcodes corresponding to them: the
471    /// source ID the item is in, the source code range of the item, and the contract name the item
472    /// is in.
473    ///
474    /// Note: Source IDs are only unique per compilation job; that is, a code base compiled with
475    /// two different solc versions will produce overlapping source IDs if the compiler version is
476    /// not taken into account.
477    #[instrument(name = "SourceAnalysis::new", skip_all)]
478    pub fn new(data: &SourceFiles, output: &ProjectCompileOutput) -> eyre::Result<Self> {
479        let mut sourced_items = output.parser().solc().compiler().enter(|compiler| {
480            data.sources
481                .par_iter()
482                .map(|(&source_id, path)| {
483                    let _guard = debug_span!("SourceAnalysis::new::visit", ?path).entered();
484
485                    let (_, source) = compiler.gcx().get_ast_source(path).unwrap();
486                    let ast = source.ast.as_ref().unwrap();
487                    let (hir_source_id, _) = compiler.gcx().get_hir_source(path).unwrap();
488
489                    let mut visitor = SourceVisitor::new(source_id, compiler.gcx());
490                    for item in ast.items.iter() {
491                        // Visit only top-level contracts.
492                        let ItemKind::Contract(contract) = &item.kind else { continue };
493
494                        // Skip interfaces which have no function implementations.
495                        if contract.kind.is_interface() {
496                            continue;
497                        }
498
499                        let checkpoint = visitor.checkpoint();
500                        visitor.visit_contract(contract);
501                        if visitor.has_tests(&checkpoint) {
502                            visitor.restore_checkpoint(checkpoint);
503                        }
504                    }
505
506                    if !visitor.function_calls.is_empty() {
507                        visitor.resolve_function_calls(hir_source_id);
508                    }
509
510                    if !visitor.items.is_empty() {
511                        visitor.disambiguate_functions();
512                        visitor.sort();
513                        visitor.push_lines();
514                        visitor.sort();
515                    }
516                    (source_id, visitor.items)
517                })
518                .collect::<Vec<(u32, Vec<CoverageItem>)>>()
519        });
520
521        // Create mapping and merge items.
522        sourced_items.sort_by_key(|(id, items)| (*id, items.first().map(|i| i.loc.bytes.start)));
523        let Some(&(max_idx, _)) = sourced_items.last() else { return Ok(Self::default()) };
524        let len = max_idx + 1;
525        let mut all_items = Vec::new();
526        let mut map = vec![(u32::MAX, 0); len as usize];
527        for (idx, items) in sourced_items {
528            // Assumes that all `idx` items are consecutive, guaranteed by the sort above.
529            let idx = idx as usize;
530            if map[idx].0 == u32::MAX {
531                map[idx].0 = all_items.len() as u32;
532            }
533            map[idx].1 += items.len() as u32;
534            all_items.extend(items);
535        }
536
537        Ok(Self { all_items, map })
538    }
539
540    /// Returns all the coverage items.
541    pub fn all_items(&self) -> &[CoverageItem] {
542        &self.all_items
543    }
544
545    /// Returns all the mutable coverage items.
546    pub const fn all_items_mut(&mut self) -> &mut Vec<CoverageItem> {
547        &mut self.all_items
548    }
549
550    /// Returns an iterator over the coverage items and their IDs for the given source.
551    pub fn items_for_source_enumerated(
552        &self,
553        source_id: u32,
554    ) -> impl Iterator<Item = (u32, &CoverageItem)> {
555        let (base_id, items) = self.items_for_source(source_id);
556        items.iter().enumerate().map(move |(idx, item)| (base_id + idx as u32, item))
557    }
558
559    /// Returns the base item ID and all the coverage items for the given source.
560    pub fn items_for_source(&self, source_id: u32) -> (u32, &[CoverageItem]) {
561        let (mut offset, len) = self.map.get(source_id as usize).copied().unwrap_or_default();
562        if offset == u32::MAX {
563            offset = 0;
564        }
565        (offset, &self.all_items[offset as usize..][..len as usize])
566    }
567
568    /// Returns the coverage item for the given item ID.
569    #[inline]
570    pub fn get(&self, item_id: u32) -> Option<&CoverageItem> {
571        self.all_items.get(item_id as usize)
572    }
573}
574
575/// A list of versioned sources and their ASTs.
576#[derive(Default)]
577pub struct SourceFiles {
578    /// The versioned sources.
579    pub sources: HashMap<u32, PathBuf>,
580}