Skip to main content

candle_graph/
cli.rs

1//! Shared CLI engine for `candle-graph` and `cargo-candle-graph`.
2
3use anyhow::{Context, Result};
4use std::{
5    collections::HashMap,
6    io::{self, Write},
7    path::{Path, PathBuf},
8};
9
10use crate::{
11    analysis_cache,
12    cargo_context::CargoOptions,
13    diagnostics::{self, MessageFormat},
14    discover::{self, ScanOptions},
15    model_baseline,
16    model_ir::ModelIr,
17    query::{self, QueryKind},
18    verify,
19};
20
21/// Output format for model IR / query responses.
22#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum OutputFormat {
24    Json,
25    Tree,
26    #[cfg(feature = "visualizer")]
27    Html,
28}
29
30/// How to render analyzed model facts.
31#[derive(Debug, Clone)]
32pub enum ReportMode {
33    FullIr {
34        format: OutputFormat,
35    },
36    Query {
37        kind: String,
38        selector: Option<String>,
39        to: Option<String>,
40        limit: usize,
41        offset: usize,
42        format: OutputFormat,
43    },
44}
45
46/// Inputs for a unified-model scan.
47#[derive(Debug, Clone)]
48pub struct AnalyzeRequest {
49    pub path: PathBuf,
50    pub cargo: CargoOptions,
51    #[cfg(feature = "runtime")]
52    pub runtime_trace: Option<PathBuf>,
53    pub component_root: Option<String>,
54    pub dataflow: bool,
55    pub heuristic_architecture: bool,
56    pub use_cache: bool,
57}
58
59/// Bundle written by `--audit-dir` / `cargo candle-graph audit`.
60#[derive(Debug, Clone)]
61pub struct AuditBundle {
62    pub output_dir: PathBuf,
63    pub checkpoint: Option<PathBuf>,
64    pub verify_root: Option<String>,
65    pub deny_rules: Vec<String>,
66    pub strict: bool,
67}
68
69/// One invocation of the shared model-mode engine.
70#[derive(Debug, Clone)]
71pub struct ModelRun {
72    pub analyze: AnalyzeRequest,
73    pub report: ReportMode,
74    pub output: Option<PathBuf>,
75    /// When set, finding diagnostics are written to stderr in this format.
76    pub diagnostics: Option<MessageFormat>,
77    /// Exit non-zero when findings include proven Error defects.
78    pub fail_on_warning_error: bool,
79    /// Exit non-zero when proven findings match any of these rule names.
80    pub deny_rules: Vec<String>,
81    pub check_baseline: Option<PathBuf>,
82    pub update_baseline: Option<PathBuf>,
83    pub audit: Option<AuditBundle>,
84    /// Optional safetensors checkpoint merged before report/audit (header verify only).
85    pub checkpoint: Option<PathBuf>,
86    /// VarBuilder root for safetensors checkpoint verification (default: first builder or `vb`).
87    pub verify_root: Option<String>,
88}
89
90/// Analyze a crate into the unified model IR, optionally reusing a disk cache.
91pub fn analyze_model(request: &AnalyzeRequest) -> Result<ModelIr> {
92    if analysis_cache::cache_enabled(request.use_cache) {
93        let canonical = request
94            .path
95            .canonicalize()
96            .unwrap_or_else(|_| request.path.clone());
97        let key = format!("{canonical:?}|{:?}", request.cargo);
98        let cache_path = analysis_cache::cache_path(&key);
99        if let Some(cached) = analysis_cache::load(&cache_path)? {
100            return Ok(cached);
101        }
102        let model = analyze_model_uncached(request)?;
103        analysis_cache::save(&cache_path, &model)?;
104        return Ok(model);
105    }
106    analyze_model_uncached(request)
107}
108
109fn analyze_model_uncached(request: &AnalyzeRequest) -> Result<ModelIr> {
110    let options = ScanOptions {
111        cargo: request.cargo.clone(),
112        #[cfg(feature = "runtime")]
113        runtime_trace: request.runtime_trace.clone(),
114        component_root: request.component_root.clone(),
115        dataflow: request.dataflow,
116        heuristic_architecture: request.heuristic_architecture,
117    };
118    discover::analyze(&request.path, &options)
119}
120
121/// Render a report from an already-analyzed model.
122pub fn render_report(model: &ModelIr, report: &ReportMode) -> Result<String> {
123    match report {
124        ReportMode::FullIr { format } => match format {
125            OutputFormat::Json => Ok(serde_json::to_string_pretty(model)? + "\n"),
126            OutputFormat::Tree => {
127                let request = query::QueryRequest::new(query::QueryKind::Summary);
128                Ok(query::render_text(&query::execute(model, &request)?))
129            }
130            #[cfg(feature = "visualizer")]
131            OutputFormat::Html => {
132                let payload = crate::viewer_projection::project(model);
133                Ok(crate::viewer::render_html(&payload))
134            }
135        },
136        ReportMode::Query {
137            kind,
138            selector,
139            to,
140            limit,
141            offset,
142            format,
143        } => {
144            let mut request = query::QueryRequest::new(kind.parse()?);
145            request.selector = selector.clone();
146            request.to = to.clone();
147            request.limit = *limit;
148            request.offset = *offset;
149            let response = query::execute(model, &request)?;
150            match format {
151                OutputFormat::Json => Ok(serde_json::to_string_pretty(&response)? + "\n"),
152                OutputFormat::Tree => Ok(query::render_text(&response)),
153                #[cfg(feature = "visualizer")]
154                OutputFormat::Html => {
155                    let payload = crate::viewer_projection::project(model);
156                    Ok(crate::viewer::render_html(&payload))
157                }
158            }
159        }
160    }
161}
162
163fn write_query(model: &ModelIr, kind: QueryKind, path: &Path) -> Result<()> {
164    let request = query::QueryRequest::new(kind);
165    let response = query::execute(model, &request)?;
166    let rendered = serde_json::to_string_pretty(&response)? + "\n";
167    write_output(Some(path), rendered.as_bytes())
168}
169
170fn write_checkpoint_audit(
171    model: &mut ModelIr,
172    checkpoint: &Path,
173    verify_root: Option<&str>,
174    path: &Path,
175) -> Result<()> {
176    let header = verify::read_header(checkpoint)?;
177    let root = verify_root
178        .map(str::to_string)
179        .or_else(|| {
180            model
181                .components
182                .first()
183                .and_then(|c| c.builders.first().map(|b| b.name.clone()))
184        })
185        .unwrap_or_else(|| "vb".to_string());
186    let report = verify::verify_model(model, &header, &root);
187    write_output(
188        Some(path),
189        (serde_json::to_string_pretty(&report)? + "\n").as_bytes(),
190    )
191}
192
193fn run_audit_bundle(
194    model: &mut ModelIr,
195    _request: &AnalyzeRequest,
196    audit: &AuditBundle,
197) -> Result<()> {
198    std::fs::create_dir_all(&audit.output_dir)
199        .with_context(|| format!("creating {}", audit.output_dir.display()))?;
200    write_query(
201        model,
202        QueryKind::Summary,
203        &audit.output_dir.join("summary.json"),
204    )?;
205    write_query(
206        model,
207        QueryKind::Doctor,
208        &audit.output_dir.join("doctor.json"),
209    )?;
210    write_query(
211        model,
212        QueryKind::ModelImprovement,
213        &audit.output_dir.join("model-improvement.json"),
214    )?;
215    write_query(
216        model,
217        QueryKind::Findings,
218        &audit.output_dir.join("findings.json"),
219    )?;
220    let rendered = render_report(
221        model,
222        &ReportMode::FullIr {
223            format: OutputFormat::Json,
224        },
225    )?;
226    write_output(
227        Some(&audit.output_dir.join("model-ir.json")),
228        rendered.as_bytes(),
229    )?;
230    #[cfg(feature = "runtime")]
231    if let Some(runtime) = &_request.runtime_trace {
232        write_query(
233            model,
234            QueryKind::Runtime,
235            &audit.output_dir.join("runtime.json"),
236        )?;
237        let _ = runtime;
238    }
239    #[cfg(feature = "visualizer")]
240    {
241        let payload = crate::viewer_projection::project(model);
242        write_output(
243            Some(&audit.output_dir.join("model.html")),
244            crate::viewer::render_html(&payload).as_bytes(),
245        )?;
246    }
247    if let Some(checkpoint) = &audit.checkpoint {
248        write_checkpoint_audit(
249            model,
250            checkpoint,
251            audit.verify_root.as_deref(),
252            &audit.output_dir.join("checkpoint.json"),
253        )?;
254    }
255    Ok(())
256}
257
258fn enforce_exit_policy(model: &ModelIr, run: &ModelRun) -> Result<()> {
259    if run.fail_on_warning_error && diagnostics::has_proven_defect_findings(model) {
260        anyhow::bail!(
261            "strict: unified model analysis contains proven error findings \
262             (coverage gaps and non-proven warnings do not fail)"
263        );
264    }
265    let denied = diagnostics::denied_findings(model, &run.deny_rules);
266    if !denied.is_empty() {
267        let rules: Vec<_> = denied.iter().map(|f| f.rule.as_str()).collect();
268        anyhow::bail!(
269            "deny: proven findings matched blocked rules: {}",
270            rules.join(", ")
271        );
272    }
273    if run.audit.as_ref().is_some_and(|a| a.strict)
274        && (diagnostics::has_proven_defect_findings(model)
275            || diagnostics::has_denied_findings(model, &run.deny_rules))
276    {
277        anyhow::bail!("audit strict gate failed");
278    }
279    Ok(())
280}
281
282fn apply_checkpoint(
283    model: &mut ModelIr,
284    checkpoint: &Path,
285    verify_root: Option<&str>,
286) -> Result<()> {
287    let header = verify::read_header(checkpoint)?;
288    let root = verify_root
289        .map(str::to_string)
290        .or_else(|| {
291            model
292                .components
293                .first()
294                .and_then(|c| c.builders.first().map(|b| b.name.clone()))
295        })
296        .unwrap_or_else(|| "vb".to_string());
297    verify::verify_model(model, &header, &root);
298    crate::dtype_propagate::propagate_tensor_dtypes(model, &HashMap::new());
299    model.normalize();
300    Ok(())
301}
302
303/// Run a full model-mode invocation: analyze, optional baseline, report, diagnostics, exit policy.
304pub fn run_model(run: &ModelRun) -> Result<()> {
305    let mut model = analyze_model(&run.analyze)?;
306
307    if let Some(path) = &run.update_baseline {
308        model_baseline::update(&model, path)
309            .with_context(|| format!("updating model baseline {}", path.display()))?;
310        eprintln!("updated model baseline {}", path.display());
311    }
312    if let Some(path) = &run.check_baseline {
313        model_baseline::check(&model, path)
314            .with_context(|| format!("checking model baseline {}", path.display()))?;
315    }
316
317    if run.audit.is_none() {
318        if let Some(checkpoint) = &run.checkpoint {
319            apply_checkpoint(&mut model, checkpoint, run.verify_root.as_deref())?;
320        }
321    }
322
323    if let Some(audit) = &run.audit {
324        run_audit_bundle(&mut model, &run.analyze, audit)?;
325    }
326
327    let rendered = render_report(&model, &run.report)?;
328    if run.audit.is_none() || run.output.is_some() {
329        write_output(run.output.as_deref(), rendered.as_bytes())?;
330    }
331
332    if let Some(format) = run.diagnostics {
333        let diagnostics = diagnostics::from_model(&model);
334        let text = diagnostics::render(&diagnostics, format);
335        if !text.is_empty() {
336            eprint!("{text}");
337        }
338    }
339
340    enforce_exit_policy(&model, run)?;
341    Ok(())
342}
343
344/// Resolve a package path from an optional directory / manifest path.
345pub fn resolve_package_path(path: Option<&Path>, manifest_path: Option<&Path>) -> Result<PathBuf> {
346    if let Some(manifest) = manifest_path {
347        if manifest
348            .file_name()
349            .is_some_and(|name| name == "Cargo.toml")
350        {
351            return Ok(manifest
352                .parent()
353                .map(Path::to_path_buf)
354                .unwrap_or_else(|| PathBuf::from(".")));
355        }
356        return Ok(manifest.to_path_buf());
357    }
358    Ok(path
359        .map(Path::to_path_buf)
360        .unwrap_or_else(|| PathBuf::from(".")))
361}
362
363pub fn write_output(path: Option<&Path>, bytes: &[u8]) -> Result<()> {
364    match path {
365        Some(path) => {
366            if let Some(parent) = path
367                .parent()
368                .filter(|parent| !parent.as_os_str().is_empty())
369            {
370                std::fs::create_dir_all(parent)
371                    .with_context(|| format!("creating {}", parent.display()))?;
372            }
373            std::fs::write(path, bytes).with_context(|| format!("writing {}", path.display()))?;
374        }
375        None => io::stdout().lock().write_all(bytes)?,
376    }
377    Ok(())
378}