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