Skip to main content

aprender_contracts_cli/commands/
coverage.rs

1use std::path::Path;
2
3use provable_contracts::binding::parse_binding;
4use provable_contracts::coverage::{coverage_report, overall_percentage, CoverageReport};
5use provable_contracts::reverse_coverage::reverse_coverage;
6use provable_contracts::schema::Contract;
7
8use crate::contract_walk::collect_corpus;
9
10pub fn run(
11    contract_dir: &Path,
12    binding_path: Option<&Path>,
13    _show_fuzz: bool,
14    reverse_crate: Option<&Path>,
15    enforcement_crate: Option<&Path>,
16) -> Result<(), Box<dyn std::error::Error>> {
17    if let Some(crate_dir) = reverse_crate {
18        return run_reverse_coverage(crate_dir, binding_path);
19    }
20
21    let binding = match binding_path {
22        Some(bp) => Some(parse_binding(bp)?),
23        None => None,
24    };
25
26    let contracts = collect_corpus(contract_dir)?;
27    // PVL-1 (PMAT-1099): an empty corpus is refused (exit 2), never reported as 0/0.
28    let refs: Vec<(String, &Contract)> = contracts.iter().map(|(s, c)| (s.clone(), c)).collect();
29    let report = coverage_report(&refs, binding.as_ref());
30    let pct = overall_percentage(&report);
31
32    print_coverage_report(&report, pct, binding_path.is_some());
33
34    if let Some(crate_dir) = enforcement_crate {
35        let bp = binding_path.ok_or("--enforcement requires --binding <path>")?;
36        let binding = provable_contracts::binding::parse_binding(bp)?;
37        print_enforcement_report(crate_dir, &binding);
38    }
39
40    Ok(())
41}
42
43/// Run `--reverse` mode: list public functions not covered by `binding.yaml`.
44fn run_reverse_coverage(
45    crate_dir: &Path,
46    binding_path: Option<&Path>,
47) -> Result<(), Box<dyn std::error::Error>> {
48    let bp = binding_path.ok_or("--reverse requires --binding <path>")?;
49    let report = reverse_coverage(crate_dir, bp);
50    println!("Reverse Coverage Report");
51    println!("=======================");
52    println!("  Public functions: {}", report.total_pub_fns);
53    println!("  Bound (in binding.yaml): {}", report.bound_fns);
54    println!("  Annotated (#[contract]): {}", report.annotated_fns);
55    println!("  Auto-exempt (trivial): {}", report.exempt_fns);
56    println!("  Unbound: {}", report.unbound.len());
57    println!(
58        "  Coverage: {:.1}% (bound + exempt / total)",
59        report.coverage_pct
60    );
61
62    let gated_count = report
63        .unbound
64        .iter()
65        .filter(|f| f.feature_gate.is_some())
66        .count();
67    if gated_count > 0 {
68        println!("  Feature-gated:   {gated_count} (require --features to test)");
69    }
70
71    print_unbound_functions(&report.unbound);
72    Ok(())
73}
74
75fn print_unbound_functions(unbound: &[provable_contracts::reverse_coverage::PubFn]) {
76    if unbound.is_empty() {
77        return;
78    }
79    println!("\nUnbound functions:");
80    for f in unbound.iter().take(20) {
81        let gate_note = match &f.feature_gate {
82            Some(feat) => format!(" [requires --features {feat}]"),
83            None => String::new(),
84        };
85        println!("  {} ({}:{}){gate_note}", f.path, f.file, f.line);
86    }
87    if unbound.len() > 20 {
88        println!("  ... and {} more", unbound.len() - 20);
89    }
90}
91
92fn print_coverage_report(report: &CoverageReport, pct: f64, has_binding: bool) {
93    println!("Obligation Coverage Report");
94    println!("==========================");
95    println!();
96
97    for cc in &report.contracts {
98        println!(
99            "  {:<35} eq={} ob={} ft={} kani={} impl={}/{}",
100            cc.stem,
101            cc.equations,
102            cc.obligations,
103            cc.falsification_covered,
104            cc.kani_covered,
105            cc.binding_implemented,
106            cc.equations,
107        );
108    }
109
110    println!();
111    println!("Totals:");
112    println!("  Contracts:            {}", report.totals.contracts);
113    println!("  Equations:            {}", report.totals.equations);
114    println!("  Obligations:          {}", report.totals.obligations);
115    println!(
116        "  Falsification tests:  {}",
117        report.totals.falsification_tests
118    );
119    println!("  Kani harnesses:       {}", report.totals.kani_harnesses);
120    if has_binding {
121        println!(
122            "  Binding implemented:  {}",
123            report.totals.binding_implemented
124        );
125        println!("  Binding partial:      {}", report.totals.binding_partial);
126        println!("  Binding missing:      {}", report.totals.binding_missing);
127    }
128    println!();
129    println!("Overall obligation coverage: {pct:.1}%");
130}
131
132/// Scan crate source for contract call sites and classify enforcement quality.
133fn print_enforcement_report(
134    crate_dir: &Path,
135    binding: &provable_contracts::binding::BindingRegistry,
136) {
137    let src_dir = crate_dir.join("src");
138    if !src_dir.exists() {
139        eprintln!("warning: {}/src not found", crate_dir.display());
140        return;
141    }
142
143    // Find all contract_pre_* call sites in .rs files
144    let mut call_sites: Vec<CallSite> = Vec::new();
145    scan_call_sites(&src_dir, &mut call_sites);
146
147    // Read generated_contracts.rs to classify macro quality
148    let gen_path = src_dir.join("generated_contracts.rs");
149    let gen_content = std::fs::read_to_string(&gen_path).unwrap_or_default();
150
151    // Classify each call site
152    for site in &mut call_sites {
153        site.level = classify_macro(&site.macro_name, &gen_content);
154    }
155
156    let total_bindings = binding.bindings.len();
157    let total_sites = call_sites.len();
158    let e0_count = call_sites
159        .iter()
160        .filter(|s| s.level == EnforcementLevel::E0)
161        .count();
162    let e1_count = call_sites
163        .iter()
164        .filter(|s| s.level == EnforcementLevel::E1)
165        .count();
166    let e2_count = call_sites
167        .iter()
168        .filter(|s| s.level == EnforcementLevel::E2)
169        .count();
170
171    #[allow(clippy::cast_precision_loss)]
172    let penetration = if total_bindings > 0 {
173        total_sites as f64 / total_bindings as f64
174    } else {
175        0.0
176    };
177
178    #[allow(clippy::cast_precision_loss)]
179    let quality_score = if total_sites > 0 {
180        (e0_count as f64 * 0.1 + e1_count as f64 * 0.5 + e2_count as f64 * 1.0) / total_sites as f64
181    } else {
182        0.0
183    };
184
185    let enforcement_score = penetration * quality_score;
186
187    println!();
188    println!("Enforcement Quality Report");
189    println!("==========================");
190    println!();
191    println!("  Bindings declared:    {total_bindings}");
192    println!("  Call sites found:     {total_sites}");
193    println!(
194        "  Penetration:          {:.1}% ({total_sites}/{total_bindings})",
195        penetration * 100.0
196    );
197    println!();
198    println!("  E0 (generic !is_empty):  {e0_count}");
199    println!("  E1 (domain pre-checks):  {e1_count}");
200    println!("  E2 (pre + post checks):  {e2_count}");
201    println!("  Quality score:           {quality_score:.2} (E0=0.1, E1=0.5, E2=1.0)");
202    println!();
203    println!("  Enforcement score:       {enforcement_score:.4} (penetration × quality)");
204    println!();
205
206    if !call_sites.is_empty() {
207        println!("  Call sites:");
208        for site in &call_sites {
209            let level_str = match site.level {
210                EnforcementLevel::E0 => "E0",
211                EnforcementLevel::E1 => "E1",
212                EnforcementLevel::E2 => "E2",
213            };
214            println!(
215                "    [{level_str}] {}:{} — {}",
216                site.file, site.line, site.macro_name
217            );
218        }
219    }
220}
221
222#[derive(Debug, Clone, Copy, PartialEq, Eq)]
223enum EnforcementLevel {
224    E0,
225    E1,
226    E2,
227}
228
229struct CallSite {
230    file: String,
231    line: usize,
232    macro_name: String,
233    level: EnforcementLevel,
234}
235
236/// Recursively scan `.rs` files for `contract_pre_*` and `contract_post_*` invocations.
237fn scan_call_sites(dir: &Path, sites: &mut Vec<CallSite>) {
238    let Ok(entries) = std::fs::read_dir(dir) else {
239        return;
240    };
241    for entry in entries.flatten() {
242        let path = entry.path();
243        if path.is_dir() {
244            scan_call_sites(&path, sites);
245        } else if path.extension().and_then(|e| e.to_str()) == Some("rs")
246            && path.file_name().and_then(|n| n.to_str()) != Some("generated_contracts.rs")
247        {
248            let Ok(content) = std::fs::read_to_string(&path) else {
249                continue;
250            };
251            for (i, line) in content.lines().enumerate() {
252                if let Some(pos) = line.find("contract_pre_") {
253                    let rest = &line[pos..];
254                    let end = rest.find('!').unwrap_or(rest.len());
255                    let macro_name = rest[..end].to_string();
256                    sites.push(CallSite {
257                        file: path.display().to_string(),
258                        line: i + 1,
259                        macro_name,
260                        level: EnforcementLevel::E0,
261                    });
262                }
263            }
264        }
265    }
266}
267
268/// Classify a macro's enforcement level by inspecting its body in `generated_contracts.rs`.
269fn classify_macro(macro_name: &str, gen_content: &str) -> EnforcementLevel {
270    // Find the macro definition
271    let pattern = format!("macro_rules! {macro_name} {{");
272    let Some(start) = gen_content.find(&pattern) else {
273        return EnforcementLevel::E0;
274    };
275    // Extract the macro body (next ~20 lines)
276    let body: String = gen_content[start..]
277        .lines()
278        .take(20)
279        .collect::<Vec<_>>()
280        .join("\n");
281
282    let has_domain_pre = body.contains("is_finite")
283        || body.contains("len() >")
284        || body.contains("len() %")
285        || body.contains("len() ==")
286        || body.contains("is_empty()");
287
288    // Check if there's a corresponding post macro
289    let post_name = macro_name.replace("contract_pre_", "contract_post_");
290    let has_post = gen_content.contains(&format!("macro_rules! {post_name} {{"));
291
292    if has_domain_pre && has_post {
293        EnforcementLevel::E2
294    } else if has_domain_pre {
295        EnforcementLevel::E1
296    } else {
297        EnforcementLevel::E0
298    }
299}