Skip to main content

miden_objtool/
decorators.rs

1use std::{fs, path::PathBuf};
2
3use anyhow::{Context, Result};
4use clap::Args;
5use miden_core::{
6    mast::MastForest,
7    serde::{Deserializable, Serializable},
8};
9use miden_mast_package::{Package, TargetType};
10
11#[derive(Debug, Clone, Args)]
12#[command(arg_required_else_help = true)]
13pub struct DecoratorsCommand {
14    /// Path to the input .masp file
15    #[arg(required = true)]
16    pub path: PathBuf,
17}
18
19#[derive(Debug, Clone, Copy)]
20enum ArtifactKind {
21    Program,
22    Library,
23}
24
25impl ArtifactKind {
26    fn from_package(package: &Package) -> Self {
27        if package.is_program() {
28            ArtifactKind::Program
29        } else {
30            ArtifactKind::Library
31        }
32    }
33}
34
35impl std::fmt::Display for ArtifactKind {
36    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
37        match self {
38            Self::Program => write!(f, "program"),
39            Self::Library => write!(f, "library"),
40        }
41    }
42}
43
44pub fn run(command: &DecoratorsCommand) -> Result<()> {
45    let input_bytes = fs::read(&command.path)
46        .with_context(|| format!("failed to read input file '{}'", command.path.display()))?;
47    let masp_size = input_bytes.len();
48
49    let package = Package::read_from_bytes(&input_bytes)
50        .with_context(|| format!("failed to decode package '{}'", command.path.display()))?;
51
52    let original_forest = package.mast.mast_forest().as_ref().clone();
53    let original_forest_size = forest_size(&original_forest);
54
55    let mut stripped_forest = original_forest.clone();
56    stripped_forest.clear_debug_info();
57
58    let (compacted_forest, _) = stripped_forest.clone().compact();
59
60    let report = Report {
61        input: command.path.display().to_string(),
62        package_kind: package.kind,
63        artifact_kind: ArtifactKind::from_package(&package),
64        metric_points: vec![
65            MetricPoint::reference("original masp", masp_size),
66            MetricPoint::baseline("original forest", original_forest_size),
67            MetricPoint::delta(
68                "without decorators",
69                forest_size(&stripped_forest),
70                original_forest_size,
71            ),
72            MetricPoint::delta(
73                "compacted forest",
74                forest_size(&compacted_forest),
75                original_forest_size,
76            ),
77        ],
78    };
79
80    println!("{report}");
81
82    Ok(())
83}
84
85fn forest_size(forest: &MastForest) -> usize {
86    forest.to_bytes().len()
87}
88
89fn bytes_to_kb(bytes: usize) -> f64 {
90    bytes as f64 / 1024.0
91}
92
93#[derive(Debug, Clone)]
94struct Report {
95    input: String,
96    package_kind: TargetType,
97    artifact_kind: ArtifactKind,
98    metric_points: Vec<MetricPoint>,
99}
100
101#[derive(Debug, Clone)]
102struct MetricPoint {
103    label: &'static str,
104    bytes: usize,
105    delta: Option<i64>,
106    delta_percent: Option<f64>,
107}
108
109impl MetricPoint {
110    fn reference(label: &'static str, bytes: usize) -> Self {
111        Self {
112            label,
113            bytes,
114            delta: None,
115            delta_percent: None,
116        }
117    }
118
119    fn baseline(label: &'static str, bytes: usize) -> Self {
120        Self {
121            label,
122            bytes,
123            delta: Some(0),
124            delta_percent: Some(0.0),
125        }
126    }
127
128    fn delta(label: &'static str, bytes: usize, baseline: usize) -> Self {
129        let delta = i64::try_from(bytes).unwrap() - i64::try_from(baseline).unwrap();
130        let delta_percent = if baseline == 0 {
131            0.0
132        } else {
133            (delta as f64 / baseline as f64) * 100.0
134        };
135
136        Self {
137            label,
138            bytes,
139            delta: Some(delta),
140            delta_percent: Some(delta_percent),
141        }
142    }
143}
144
145fn format_delta(row: &MetricPoint) -> String {
146    match row.delta {
147        Some(delta) => {
148            let delta_kb = delta as f64 / 1024.0;
149            if delta_kb.abs() < 0.01 {
150                "0.00".to_string()
151            } else if delta_kb > 0.0 {
152                format!("+{delta_kb:.2}")
153            } else {
154                format!("{delta_kb:.2}")
155            }
156        }
157        None => "-".to_string(),
158    }
159}
160
161fn format_delta_percent(row: &MetricPoint) -> String {
162    match row.delta_percent {
163        Some(percent) => format!("{percent:+.2}%"),
164        None => "-".to_string(),
165    }
166}
167
168impl std::fmt::Display for Report {
169    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
170        writeln!(f, "Input: {}", self.input)?;
171        writeln!(f, "Package kind: {}", self.package_kind)?;
172        writeln!(f, "Artifact: {}", self.artifact_kind)?;
173        writeln!(f)?;
174
175        let kb_strings: Vec<String> = self
176            .metric_points
177            .iter()
178            .map(|row| format!("{:.2}", bytes_to_kb(row.bytes)))
179            .collect();
180        let delta_strings: Vec<String> = self.metric_points.iter().map(format_delta).collect();
181        let delta_percent_strings: Vec<String> =
182            self.metric_points.iter().map(format_delta_percent).collect();
183
184        let metric_width = self
185            .metric_points
186            .iter()
187            .map(|row| row.label.len())
188            .chain(std::iter::once("Metric".len()))
189            .max()
190            .unwrap_or("Metric".len());
191        let kb_width = kb_strings
192            .iter()
193            .map(String::len)
194            .chain(std::iter::once("KB".len()))
195            .max()
196            .unwrap_or("KB".len());
197        let delta_width = delta_strings
198            .iter()
199            .map(String::len)
200            .chain(std::iter::once("Delta".len()))
201            .max()
202            .unwrap_or("Delta".len());
203        let delta_percent_width = delta_percent_strings
204            .iter()
205            .map(String::len)
206            .chain(std::iter::once("Delta %".len()))
207            .max()
208            .unwrap_or("Delta %".len());
209
210        writeln!(
211            f,
212            "{:<metric_width$}  {:>kb_width$}  {:>delta_width$}  {:>delta_percent_width$}",
213            "Metric", "KB", "Delta", "Delta %",
214        )?;
215
216        for ((row, kb), (delta, delta_percent)) in self
217            .metric_points
218            .iter()
219            .zip(kb_strings.iter())
220            .zip(delta_strings.iter().zip(delta_percent_strings.iter()))
221        {
222            writeln!(
223                f,
224                "{:<metric_width$}  {:>kb_width$}  {:>delta_width$}  {:>delta_percent_width$}",
225                row.label, kb, delta, delta_percent,
226            )?;
227        }
228
229        Ok(())
230    }
231}