1use std::fmt;
4use std::fs;
5use std::io::{self, Write};
6use std::path::{Path, PathBuf};
7
8use crate::model_ir::{FindingSeverity, ModelIr};
9
10pub const SCHEMA: &str = "candle-graph/model-baseline/1";
12
13const HEADER: &str = "# candle-graph/model-baseline/1";
14
15#[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#[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(¶meter.builder_root),
88 escape_field(¶meter.key),
89 escape_field(¶meter.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(¶m_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
209pub 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}