diff --git a/tests/representation_audit.rs b/tests/representation_audit.rs index 8043cac..0305325 100644 --- a/tests/representation_audit.rs +++ b/tests/representation_audit.rs @@ -493,8 +493,81 @@ fn rust_files(root: &Path) -> Vec { files } +/// every `mod name;` in one file, as the file it loads and whether it's test-only +struct ModuleFiles<'a> { + known: &'a BTreeSet, + dir: Vec, + test_only: bool, + found: Vec<(String, bool)>, +} + +impl<'ast> Visit<'ast> for ModuleFiles<'_> { + fn visit_item_mod(&mut self, node: &'ast ItemMod) { + let previous_test_only = self.test_only; + self.test_only |= is_test_only(&node.attrs); + self.dir.push(node.ident.to_string()); + match &node.content { + Some((_, items)) => { + for item in items { + self.visit_item(item); + } + } + None => { + let base = self.dir.join("/"); + let file = [format!("{base}.rs"), format!("{base}/mod.rs")] + .into_iter() + .find(|file| self.known.contains(file)); + if let Some(file) = file { + self.found.push((file, self.test_only)); + } + } + } + self.dir.pop(); + self.test_only = previous_test_only; + } +} + +/// files that only compile under `cfg(test)` because their `mod` declaration +/// says so, which the per-file scan can't see from inside the file +fn test_only_files(sources: &BTreeMap) -> BTreeSet { + let known = sources.keys().cloned().collect::>(); + let mut declared = BTreeMap::new(); + for (relative, syntax) in sources { + let path = Path::new(relative); + let stem = path.file_stem().unwrap().to_string_lossy(); + let parent = path.parent().unwrap().to_string_lossy().into_owned(); + let dir = if matches!(stem.as_ref(), "mod" | "lib" | "main") { + vec![parent] + } else { + vec![parent, stem.into_owned()] + }; + let mut modules = ModuleFiles { + known: &known, + dir, + test_only: false, + found: Vec::new(), + }; + modules.visit_file(syntax); + declared.insert(relative.clone(), modules.found); + } + + let mut test_only = BTreeSet::new(); + let mut pending = declared + .values() + .flatten() + .filter(|(_, test_only)| *test_only) + .map(|(file, _)| file.clone()) + .collect::>(); + while let Some(file) = pending.pop() { + if test_only.insert(file.clone()) { + pending.extend(declared[&file].iter().map(|(child, _)| child.clone())); + } + } + test_only +} + fn discover(manifest: &Path) -> Surface { - let mut surface = Surface::default(); + let mut sources = BTreeMap::new(); for path in rust_files(&manifest.join("src")) { let relative = path .strip_prefix(manifest) @@ -505,15 +578,21 @@ fn discover(manifest: &Path) -> Surface { .unwrap_or_else(|error| panic!("failed to read {}: {error}", path.display())); let syntax = syn::parse_file(&source) .unwrap_or_else(|error| panic!("failed to parse {}: {error}", path.display())); + sources.insert(relative, syntax); + } + + let test_only = test_only_files(&sources); + let mut surface = Surface::default(); + for (relative, syntax) in &sources { let mut scanner = Scanner { - relative_path: &relative, + relative_path: relative, modules: Vec::new(), impl_name: None, function: None, - test_only: false, + test_only: test_only.contains(relative), surface: &mut surface, }; - scanner.visit_file(&syntax); + scanner.visit_file(syntax); } surface } @@ -972,6 +1051,47 @@ mod tests { assert_eq!(errors, ["stale: serde-type:src/lib.rs::Removed"]); } +#[test] +fn representation_audit_skips_files_behind_a_test_only_mod() { + let tmp = tempfile::tempdir().unwrap(); + fs::create_dir_all(tmp.path().join("src/stream/helpers")).unwrap(); + fs::create_dir_all(tmp.path().join("tests")).unwrap(); + fs::write( + tmp.path().join("src/lib.rs"), + "mod stream;\npub fn persist(value: &u64) -> Vec { rmp_serde::to_vec(value).unwrap() }\n", + ) + .unwrap(); + fs::write( + tmp.path().join("src/stream.rs"), + "#[cfg(all(test, feature = \"x\"))]\nmod helpers;\n", + ) + .unwrap(); + fs::write( + tmp.path().join("src/stream/helpers/mod.rs"), + concat!( + "mod nested;\n", + "fn put() -> Vec { serde_json::to_vec(&1).unwrap() }\n", + "#[test]\nfn round_trip() {}\n", + ), + ) + .unwrap(); + fs::write( + tmp.path().join("src/stream/helpers/nested.rs"), + "fn get(bytes: &[u8]) -> u64 { rmp_serde::from_slice(bytes).unwrap() }\n", + ) + .unwrap(); + fs::write( + tmp.path().join(INVENTORY_PATH), + concat!( + "codec-call:src/lib.rs::persist::rmp_serde::to_vec", + "\trust:src/stream/helpers/mod.rs::round_trip\n", + ), + ) + .unwrap(); + + assert_eq!(inventory_errors(tmp.path()), Vec::::new()); +} + #[test] #[ignore] fn print_proposed_representation_inventory() {