Skip to main content

candle_graph/
model_baseline.rs

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