Skip to main content

candle_graph/
model_baseline.rs

1//! Deterministic ModelIr fingerprints for crate-wide CI checks.
2//!
3//! Unlike the legacy structure baseline, this captures stable identities from the unified
4//! `candle-graph/model/1` IR: components, parameters, entrypoints, and finding codes.
5
6use std::fmt;
7use std::fs;
8use std::io::{self, Write};
9use std::path::{Path, PathBuf};
10
11use crate::model_ir::{FindingSeverity, ModelIr};
12
13/// Schema line written at the top of every model baseline file.
14pub const SCHEMA: &str = "candle-graph/model-baseline/1";
15
16const HEADER: &str = "# candle-graph/model-baseline/1";
17
18/// Parsed / rendered model baseline document.
19#[derive(Debug, Clone, Default, PartialEq, Eq)]
20pub struct ModelBaseline {
21    pub components: Vec<String>,
22    pub parameters: Vec<String>,
23    pub entrypoints: Vec<String>,
24    pub findings: Vec<String>,
25}
26
27/// Line-oriented comparison of two model baselines.
28#[derive(Debug, Clone, Default, PartialEq, Eq)]
29pub struct ModelBaselineDiff {
30    pub added: Vec<String>,
31    pub removed: Vec<String>,
32}
33
34impl ModelBaselineDiff {
35    pub fn is_empty(&self) -> bool {
36        self.added.is_empty() && self.removed.is_empty()
37    }
38}
39
40impl fmt::Display for ModelBaselineDiff {
41    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
42        for line in &self.removed {
43            writeln!(f, "- {line}")?;
44        }
45        for line in &self.added {
46            writeln!(f, "+ {line}")?;
47        }
48        Ok(())
49    }
50}
51
52#[derive(Debug)]
53pub enum ModelBaselineError {
54    Io { path: PathBuf, source: io::Error },
55    Parse { path: PathBuf, message: String },
56    Mismatch(ModelBaselineDiff),
57}
58
59impl fmt::Display for ModelBaselineError {
60    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
61        match self {
62            Self::Io { path, source } => write!(f, "{}: {source}", path.display()),
63            Self::Parse { path, message } => write!(f, "{}: {message}", path.display()),
64            Self::Mismatch(diff) => write!(f, "model baseline mismatch:\n{diff}"),
65        }
66    }
67}
68
69impl std::error::Error for ModelBaselineError {}
70
71impl ModelBaseline {
72    pub fn from_model(model: &ModelIr) -> Self {
73        let mut components: Vec<String> = model
74            .components
75            .iter()
76            .map(|component| {
77                format!(
78                    "component\t{}\t{}",
79                    escape_field(&component.qualified_name),
80                    escape_field(&component.id.0)
81                )
82            })
83            .collect();
84        let mut parameters: Vec<String> = model
85            .parameters
86            .iter()
87            .map(|parameter| {
88                format!(
89                    "parameter\t{}\t{}\t{}",
90                    escape_field(&parameter.builder_root),
91                    escape_field(&parameter.key),
92                    escape_field(&parameter.kind)
93                )
94            })
95            .collect();
96        let mut entrypoints: Vec<String> = model
97            .functions
98            .iter()
99            .filter(|function| function.is_entrypoint)
100            .map(|function| format!("entrypoint\t{}", escape_field(&function.qualified_name)))
101            .collect();
102        let mut findings: Vec<String> = model
103            .findings
104            .iter()
105            .map(|finding| {
106                format!(
107                    "finding\t{}\t{}",
108                    escape_field(&finding.rule),
109                    severity_name(&finding.severity)
110                )
111            })
112            .collect();
113
114        components.sort();
115        components.dedup();
116        parameters.sort();
117        parameters.dedup();
118        entrypoints.sort();
119        entrypoints.dedup();
120        findings.sort();
121        findings.dedup();
122
123        Self {
124            components,
125            parameters,
126            entrypoints,
127            findings,
128        }
129    }
130
131    pub fn render(&self) -> String {
132        let mut lines = vec![HEADER.to_string()];
133        lines.extend(self.components.iter().cloned());
134        lines.extend(self.parameters.iter().cloned());
135        lines.extend(self.entrypoints.iter().cloned());
136        lines.extend(self.findings.iter().cloned());
137        lines.push(String::new());
138        lines.join("\n")
139    }
140
141    pub fn parse(text: &str) -> Result<Self, String> {
142        let mut baseline = Self::default();
143        let mut saw_header = false;
144        for (index, raw) in text.lines().enumerate() {
145            let line = raw.trim_end();
146            if line.is_empty() {
147                continue;
148            }
149            if line.starts_with('#') {
150                if line == HEADER {
151                    saw_header = true;
152                    continue;
153                }
154                return Err(format!("line {}: unsupported header `{line}`", index + 1));
155            }
156            let mut parts = line.split('\t');
157            let kind = parts.next().unwrap_or_default();
158            match kind {
159                "component" => {
160                    let qualified = unescape_field(parts.next().unwrap_or_default());
161                    let id = unescape_field(parts.next().unwrap_or_default());
162                    baseline.components.push(format!(
163                        "component\t{}\t{}",
164                        escape_field(&qualified),
165                        escape_field(&id)
166                    ));
167                }
168                "parameter" => {
169                    let root = unescape_field(parts.next().unwrap_or_default());
170                    let key = unescape_field(parts.next().unwrap_or_default());
171                    let param_kind = unescape_field(parts.next().unwrap_or_default());
172                    baseline.parameters.push(format!(
173                        "parameter\t{}\t{}\t{}",
174                        escape_field(&root),
175                        escape_field(&key),
176                        escape_field(&param_kind)
177                    ));
178                }
179                "entrypoint" => {
180                    let name = unescape_field(parts.next().unwrap_or_default());
181                    baseline
182                        .entrypoints
183                        .push(format!("entrypoint\t{}", escape_field(&name)));
184                }
185                "finding" => {
186                    let rule = unescape_field(parts.next().unwrap_or_default());
187                    let severity = parts.next().unwrap_or_default();
188                    baseline
189                        .findings
190                        .push(format!("finding\t{}\t{severity}", escape_field(&rule)));
191                }
192                other => {
193                    return Err(format!("line {}: unknown record kind `{other}`", index + 1));
194                }
195            }
196        }
197        if !saw_header {
198            return Err(format!("missing `{HEADER}` header"));
199        }
200        baseline.components.sort();
201        baseline.components.dedup();
202        baseline.parameters.sort();
203        baseline.parameters.dedup();
204        baseline.entrypoints.sort();
205        baseline.entrypoints.dedup();
206        baseline.findings.sort();
207        baseline.findings.dedup();
208        Ok(baseline)
209    }
210}
211
212/// Compare two baselines as sorted identity lines.
213pub fn compare(actual: &ModelBaseline, expected: &ModelBaseline) -> ModelBaselineDiff {
214    let actual_lines = all_lines(actual);
215    let expected_lines = all_lines(expected);
216    let mut added = Vec::new();
217    let mut removed = Vec::new();
218    let mut i = 0;
219    let mut j = 0;
220    while i < actual_lines.len() && j < expected_lines.len() {
221        match actual_lines[i].cmp(&expected_lines[j]) {
222            std::cmp::Ordering::Equal => {
223                i += 1;
224                j += 1;
225            }
226            std::cmp::Ordering::Less => {
227                added.push(actual_lines[i].clone());
228                i += 1;
229            }
230            std::cmp::Ordering::Greater => {
231                removed.push(expected_lines[j].clone());
232                j += 1;
233            }
234        }
235    }
236    added.extend(actual_lines[i..].iter().cloned());
237    removed.extend(expected_lines[j..].iter().cloned());
238    ModelBaselineDiff { added, removed }
239}
240
241pub fn load(path: impl AsRef<Path>) -> Result<ModelBaseline, ModelBaselineError> {
242    let path = path.as_ref();
243    let text = fs::read_to_string(path).map_err(|source| ModelBaselineError::Io {
244        path: path.to_path_buf(),
245        source,
246    })?;
247    ModelBaseline::parse(&text).map_err(|message| ModelBaselineError::Parse {
248        path: path.to_path_buf(),
249        message,
250    })
251}
252
253pub fn check(model: &ModelIr, path: impl AsRef<Path>) -> Result<(), ModelBaselineError> {
254    let expected = load(path)?;
255    let actual = ModelBaseline::from_model(model);
256    let diff = compare(&actual, &expected);
257    if diff.is_empty() {
258        Ok(())
259    } else {
260        Err(ModelBaselineError::Mismatch(diff))
261    }
262}
263
264pub fn update(model: &ModelIr, path: impl AsRef<Path>) -> Result<(), ModelBaselineError> {
265    let path = path.as_ref();
266    let text = ModelBaseline::from_model(model).render();
267    atomic_write(path, text.as_bytes())
268}
269
270pub fn atomic_write(path: &Path, bytes: &[u8]) -> Result<(), ModelBaselineError> {
271    let parent = path.parent().unwrap_or_else(|| Path::new("."));
272    fs::create_dir_all(parent).map_err(|source| ModelBaselineError::Io {
273        path: parent.to_path_buf(),
274        source,
275    })?;
276
277    let mut tmp_name = path
278        .file_name()
279        .map(|s| s.to_os_string())
280        .unwrap_or_else(|| "model-baseline".into());
281    tmp_name.push(".tmp");
282    let tmp_path = parent.join(tmp_name);
283
284    let write_tmp = || -> io::Result<()> {
285        let mut file = fs::File::create(&tmp_path)?;
286        file.write_all(bytes)?;
287        file.sync_all()?;
288        Ok(())
289    };
290    if let Err(source) = write_tmp() {
291        let _ = fs::remove_file(&tmp_path);
292        return Err(ModelBaselineError::Io {
293            path: tmp_path,
294            source,
295        });
296    }
297    fs::rename(&tmp_path, path).map_err(|source| ModelBaselineError::Io {
298        path: path.to_path_buf(),
299        source,
300    })
301}
302
303fn all_lines(baseline: &ModelBaseline) -> Vec<String> {
304    let mut lines = Vec::new();
305    lines.extend(baseline.components.iter().cloned());
306    lines.extend(baseline.parameters.iter().cloned());
307    lines.extend(baseline.entrypoints.iter().cloned());
308    lines.extend(baseline.findings.iter().cloned());
309    lines
310}
311
312fn severity_name(severity: &FindingSeverity) -> &'static str {
313    match severity {
314        FindingSeverity::Error => "error",
315        FindingSeverity::Warning => "warning",
316        FindingSeverity::Information => "information",
317    }
318}
319
320fn escape_field(value: &str) -> String {
321    value
322        .replace('\\', "\\\\")
323        .replace('\t', "\\t")
324        .replace('\n', "\\n")
325}
326
327fn unescape_field(value: &str) -> String {
328    let mut out = String::with_capacity(value.len());
329    let mut chars = value.chars().peekable();
330    while let Some(ch) = chars.next() {
331        if ch == '\\' {
332            match chars.next() {
333                Some('t') => out.push('\t'),
334                Some('n') => out.push('\n'),
335                Some('\\') => out.push('\\'),
336                Some(other) => {
337                    out.push('\\');
338                    out.push(other);
339                }
340                None => out.push('\\'),
341            }
342        } else {
343            out.push(ch);
344        }
345    }
346    out
347}
348
349#[cfg(test)]
350mod tests {
351    use super::*;
352    use crate::model_ir::{
353        Confidence, Finding, FindingSeverity, Function, ModelIr, StableId, Visibility,
354    };
355
356    fn sample_model() -> ModelIr {
357        let mut model = ModelIr::empty(StableId::new("analysis", ["test"]));
358        model.functions.push(Function {
359            id: StableId::new("fn", ["Root::forward"]),
360            name: "forward".into(),
361            qualified_name: "Root::forward".into(),
362            owner_type: Some("Root".into()),
363            visibility: Visibility::Public,
364            parameters: Vec::new(),
365            return_type: None,
366            cfg_predicates: Vec::new(),
367            cfg_active: Some(true),
368            source: "model.rs".into(),
369            calls: Vec::new(),
370            tensor_inputs: Vec::new(),
371            tensor_outputs: Vec::new(),
372            is_entrypoint: true,
373            is_loss: false,
374        });
375        model.findings.push(Finding {
376            id: StableId::new("finding", ["compiler-semantic-evidence"]),
377            rule: "compiler-semantic-evidence".into(),
378            severity: FindingSeverity::Information,
379            confidence: Confidence::Unknown,
380            message: "pending compiler frontend".into(),
381            source: None,
382            related: Vec::new(),
383            evidence: Vec::new(),
384        });
385        model
386    }
387
388    #[test]
389    fn round_trip_render_parse() {
390        let baseline = ModelBaseline::from_model(&sample_model());
391        let text = baseline.render();
392        assert!(text.starts_with(HEADER));
393        let parsed = ModelBaseline::parse(&text).unwrap();
394        assert_eq!(parsed, baseline);
395    }
396}