Skip to main content

forge/
gas_report.rs

1//! Gas reports.
2
3use crate::{
4    constants::{CHEATCODE_ADDRESS, HARDHAT_CONSOLE_ADDRESS},
5    traces::{CallTraceArena, CallTraceDecoder, CallTraceNode},
6};
7use alloy_primitives::{Address, map::HashSet};
8use comfy_table::{
9    Cell, CellAlignment, Color, Table,
10    presets::{ASCII_FULL, ASCII_MARKDOWN},
11};
12use foundry_common::{TestFunctionExt, calc, get_contract_name, shell};
13use foundry_evm::traces::CallKind;
14use serde::{Deserialize, Serialize};
15use serde_json::json;
16use std::{collections::BTreeMap, fmt::Display};
17
18/// Represents the gas report for a set of contracts.
19#[derive(Clone, Debug, Default, Serialize, Deserialize)]
20pub struct GasReport {
21    /// Whether to report any contracts.
22    report_any: bool,
23    /// Contracts to generate the report for.
24    report_for: HashSet<String>,
25    /// Contracts to ignore when generating the report.
26    ignore: HashSet<String>,
27    /// Whether to include gas reports for tests.
28    include_tests: bool,
29    /// Additional network-specific cheatcode addresses omitted from reports.
30    #[serde(skip)]
31    extra_cheatcode_addresses: HashSet<Address>,
32    /// All contracts that were analyzed grouped by their identifier
33    /// ``test/Counter.t.sol:CounterTest
34    pub contracts: BTreeMap<String, ContractInfo>,
35}
36
37impl GasReport {
38    pub fn new(
39        report_for: impl IntoIterator<Item = String>,
40        ignore: impl IntoIterator<Item = String>,
41        include_tests: bool,
42        extra_cheatcode_addresses: impl IntoIterator<Item = Address>,
43    ) -> Self {
44        let report_for = report_for.into_iter().collect::<HashSet<_>>();
45        let report_any = report_for.is_empty() || report_for.contains("*");
46        Self {
47            report_any,
48            report_for,
49            ignore: ignore.into_iter().collect(),
50            include_tests,
51            extra_cheatcode_addresses: extra_cheatcode_addresses.into_iter().collect(),
52            contracts: BTreeMap::new(),
53        }
54    }
55
56    /// Whether the given contract should be reported.
57    #[instrument(level = "trace", skip(self), ret)]
58    fn should_report(&self, contract_name: &str) -> bool {
59        if self.ignore.contains(contract_name) {
60            let contains_anyway = self.report_for.contains(contract_name);
61            if contains_anyway {
62                // If the user listed the contract in 'gas_reports' (the foundry.toml field) a
63                // report for the contract is generated even if it's listed in the ignore
64                // list. This is addressed this way because getting a report you don't expect is
65                // preferable than not getting one you expect. A warning is printed to stderr
66                // indicating the "double listing".
67                let _ = sh_warn!(
68                    "{contract_name} is listed in both 'gas_reports' and 'gas_reports_ignore'."
69                );
70            }
71            return contains_anyway;
72        }
73        self.report_any || self.report_for.contains(contract_name)
74    }
75
76    fn is_internal_address(&self, address: Address) -> bool {
77        address == CHEATCODE_ADDRESS
78            || address == HARDHAT_CONSOLE_ADDRESS
79            || self.extra_cheatcode_addresses.contains(&address)
80    }
81
82    /// Analyzes the given traces and generates a gas report.
83    pub async fn analyze(
84        &mut self,
85        arenas: impl IntoIterator<Item = &CallTraceArena>,
86        decoder: &CallTraceDecoder,
87    ) {
88        for node in arenas.into_iter().flat_map(|arena| arena.nodes()) {
89            self.analyze_node(node, decoder).await;
90        }
91    }
92
93    async fn analyze_node(&mut self, node: &CallTraceNode, decoder: &CallTraceDecoder) {
94        let trace = &node.trace;
95        if self.is_internal_address(trace.address) {
96            return;
97        }
98        let Some(name) = decoder.contracts.get(&trace.address) else { return };
99        let contract_name = get_contract_name(name);
100        if !self.should_report(contract_name) {
101            return;
102        }
103
104        let contract_info = self.contracts.entry(name.clone()).or_default();
105        let is_create_call = trace.kind.is_any_create();
106        if is_create_call {
107            trace!(contract_name, "adding create size info");
108            contract_info.size = trace.data.len();
109        }
110
111        // Only include top-level calls which account for calldata and base (21.000) cost.
112        // Only include Calls and Creates as only these calls are isolated in inspector.
113        if trace.depth > 1 && (trace.kind == CallKind::Call || is_create_call) {
114            return;
115        }
116
117        if is_create_call {
118            trace!(contract_name, "adding create gas info");
119            contract_info.gas = trace.gas_used;
120        } else if let Some(signature) = decoder.decode_function_signature(trace).await {
121            let name = signature.split('(').next().unwrap();
122            // Ignore any test/setup functions.
123            if self.include_tests || !name.test_function_kind().is_known() {
124                trace!(contract_name, signature, "adding gas info");
125                contract_info
126                    .functions
127                    .entry(name.to_string())
128                    .or_default()
129                    .entry(signature)
130                    .or_default()
131                    .frames
132                    .push(trace.gas_used);
133            }
134        }
135    }
136
137    /// Finalizes the gas report by calculating the min, max, mean, and median for each function.
138    #[must_use]
139    pub fn finalize(mut self) -> Self {
140        trace!("finalizing gas report");
141        for func in self
142            .contracts
143            .values_mut()
144            .flat_map(|c| c.functions.values_mut().flat_map(|s| s.values_mut()))
145        {
146            func.frames.sort_unstable();
147            func.min = func.frames.first().copied().unwrap_or_default();
148            func.max = func.frames.last().copied().unwrap_or_default();
149            func.mean = calc::mean(&func.frames);
150            func.median = calc::median_sorted(&func.frames);
151            func.calls = func.frames.len() as u64;
152        }
153        self
154    }
155
156    /// Contracts with at least one recorded function call.
157    fn reported_contracts(&self) -> impl Iterator<Item = (&String, &ContractInfo)> {
158        self.contracts.iter().filter(|(name, contract)| {
159            if contract.functions.is_empty() {
160                trace!(name, "gas report contract without functions");
161            }
162            !contract.functions.is_empty()
163        })
164    }
165}
166
167impl Display for GasReport {
168    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> Result<(), std::fmt::Error> {
169        if shell::is_json() {
170            return writeln!(f, "{}", self.format_json_output());
171        }
172        for (name, contract) in self.reported_contracts() {
173            writeln!(f, "\n{}", format_table_output(contract, name))?;
174        }
175        Ok(())
176    }
177}
178
179impl GasReport {
180    fn format_json_output(&self) -> String {
181        let contracts = self
182            .reported_contracts()
183            .map(|(name, contract)| {
184                let functions = contract
185                    .functions
186                    .values()
187                    .flat_map(|sigs| sigs.iter().map(|(sig, info)| (sig.replace(':', ""), info)))
188                    .collect::<BTreeMap<_, _>>();
189                json!({
190                    "contract": name,
191                    "deployment": { "gas": contract.gas, "size": contract.size },
192                    "functions": functions,
193                })
194            })
195            .collect::<Vec<_>>();
196        serde_json::to_string(&contracts).unwrap()
197    }
198}
199
200fn format_table_output(contract: &ContractInfo, name: &str) -> Table {
201    let num = |value: &dyn Display, color: Option<Color>| {
202        let cell = Cell::new(value.to_string()).set_alignment(CellAlignment::Right);
203        match color {
204            Some(color) => cell.fg(color),
205            None => cell,
206        }
207    };
208
209    let mut table = Table::new();
210    if shell::is_markdown() {
211        table.load_style(ASCII_MARKDOWN);
212    } else {
213        table.load_style(ASCII_FULL.with_rounded_corners());
214    }
215    table.set_header(vec![Cell::new(format!("{name} Contract")).fg(Color::Magenta)]);
216    table.add_row(vec![
217        Cell::new("Deployment Cost").fg(Color::Cyan),
218        Cell::new("Deployment Size").fg(Color::Cyan),
219    ]);
220    table.add_row(vec![num(&contract.gas, None), num(&contract.size, None)]);
221    // Add a blank row to separate deployment info from function info.
222    table.add_row(vec![Cell::new("")]);
223    table.add_row(vec![
224        Cell::new("Function Name"),
225        Cell::new("Min").fg(Color::Green),
226        Cell::new("Avg").fg(Color::Yellow),
227        Cell::new("Median").fg(Color::Yellow),
228        Cell::new("Max").fg(Color::Red),
229        Cell::new("# Calls").fg(Color::Cyan),
230    ]);
231    for (fname, sigs) in &contract.functions {
232        for (sig, gas_info) in sigs {
233            // Show function signature if overloaded else display function name.
234            let display_name = if sigs.len() == 1 { fname.clone() } else { sig.replace(':', "") };
235            table.add_row(vec![
236                Cell::new(display_name),
237                num(&gas_info.min, Some(Color::Green)),
238                num(&gas_info.mean, Some(Color::Yellow)),
239                num(&gas_info.median, Some(Color::Yellow)),
240                num(&gas_info.max, Some(Color::Red)),
241                num(&gas_info.calls, None),
242            ]);
243        }
244    }
245    table
246}
247
248#[derive(Clone, Debug, Default, Serialize, Deserialize)]
249pub struct ContractInfo {
250    pub gas: u64,
251    pub size: usize,
252    /// Function name -> Function signature -> GasInfo
253    pub functions: BTreeMap<String, BTreeMap<String, GasInfo>>,
254}
255
256#[derive(Clone, Debug, Default, Serialize, Deserialize)]
257pub struct GasInfo {
258    pub calls: u64,
259    pub min: u64,
260    pub mean: u64,
261    pub median: u64,
262    pub max: u64,
263
264    #[serde(skip)]
265    pub frames: Vec<u64>,
266}
267
268#[cfg(test)]
269mod tests {
270    use super::*;
271    use foundry_evm::constants::MONAD_CHEATCODE_ADDRESS;
272
273    #[test]
274    fn network_cheatcode_addresses_are_opt_in() {
275        let ethereum = GasReport::new([], [], false, []);
276        assert!(!ethereum.is_internal_address(MONAD_CHEATCODE_ADDRESS));
277
278        let monad = GasReport::new([], [], false, [MONAD_CHEATCODE_ADDRESS]);
279        assert!(monad.is_internal_address(MONAD_CHEATCODE_ADDRESS));
280    }
281}