use std::path::Path;
use anyhow::{Context, Result};
use serde::{Deserialize, Serialize};
use crate::analysis::walker::Language;
pub mod comments;
pub mod rust;
pub mod c;
pub mod cpp;
pub mod elixir;
pub mod go;
pub mod haskell;
pub mod java;
pub mod javascript;
pub mod python;
pub mod ruby;
pub mod scala;
pub mod typescript;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum SignalTier {
High,
Medium,
Low,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum SignalKind {
Panic,
Assert,
WarnComment,
LinterDisable,
UnwrapLike,
Guard,
RawApi,
}
impl SignalKind {
pub fn default_tier(self) -> SignalTier {
match self {
SignalKind::WarnComment => SignalTier::High,
SignalKind::Panic => SignalTier::High,
SignalKind::Assert => SignalTier::High,
SignalKind::LinterDisable => SignalTier::Medium,
SignalKind::Guard => SignalTier::Medium,
SignalKind::UnwrapLike => SignalTier::Medium,
SignalKind::RawApi => SignalTier::Low,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct Signal {
pub file_line: u32,
pub tier: SignalTier,
pub kind: SignalKind,
pub evidence: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SignalReport {
pub file: String,
pub language: String,
pub signal_count: usize,
pub signals: Vec<Signal>,
}
impl SignalReport {
pub fn truncate(&mut self, limit: usize) {
if limit > 0 && self.signals.len() > limit {
self.signals.truncate(limit);
self.signal_count = self.signals.len();
}
}
}
pub fn language_label(lang: Language) -> &'static str {
match lang {
Language::Rust => "rust",
Language::TypeScript => "typescript",
Language::JavaScript => "javascript",
Language::Python => "python",
Language::Go => "go",
Language::Java => "java",
Language::C => "c",
Language::Cpp => "cpp",
Language::Ruby => "ruby",
Language::Scala => "scala",
Language::Elixir => "elixir",
Language::Haskell => "haskell",
Language::Unknown => "unknown",
}
}
pub fn extract_signals(path: &Path, language: Language) -> Result<SignalReport> {
let source = std::fs::read_to_string(path)
.with_context(|| format!("failed to read {}", path.display()))?;
let parent_gated = language == Language::Rust && rust::parent_declares_test_module(path);
if is_test_path(path) || parent_gated {
return Ok(SignalReport {
file: path.display().to_string(),
language: language_label(language).to_string(),
signal_count: 0,
signals: Vec::new(),
});
}
let mut signals = match language {
Language::Rust => rust::extract(&source)?,
Language::Python => python::extract(&source)?,
Language::TypeScript => typescript::extract(&source)?,
Language::JavaScript => javascript::extract(&source)?,
Language::Go => go::extract(&source)?,
Language::Java => java::extract(&source)?,
Language::C => c::extract(&source)?,
Language::Cpp => cpp::extract(&source)?,
Language::Ruby => ruby::extract(&source)?,
Language::Scala => scala::extract(&source)?,
Language::Elixir => elixir::extract(&source)?,
Language::Haskell => haskell::extract(&source)?,
Language::Unknown => comments::scan_unknown(&source, language),
};
sort_canonical(&mut signals);
Ok(SignalReport {
file: path.display().to_string(),
language: language_label(language).to_string(),
signal_count: signals.len(),
signals,
})
}
pub fn is_test_path(path: &Path) -> bool {
if path
.components()
.any(|c| matches!(c.as_os_str().to_str(), Some("tests") | Some("__tests__")))
{
return true;
}
let Some(name) = path.file_name().and_then(|n| n.to_str()) else {
return false;
};
if name.ends_with("_test.go") {
return true;
}
if name.ends_with(".py") && (name.starts_with("test_") || name.ends_with("_test.py")) {
return true;
}
const JS_EXTS: [&str; 4] = ["js", "ts", "jsx", "tsx"];
JS_EXTS.iter().any(|ext| {
name.ends_with(&format!(".test.{ext}")) || name.ends_with(&format!(".spec.{ext}"))
})
}
pub(crate) fn node_text(source: &[u8], node: tree_sitter::Node) -> String {
let start = node.start_byte();
let end = node.end_byte().min(source.len());
if start >= end {
return String::new();
}
String::from_utf8_lossy(&source[start..end]).into_owned()
}
pub(crate) fn macro_body_signals(
parser: &mut tree_sitter::Parser,
query: &tree_sitter::Query,
body: tree_sitter::Node,
source: &[u8],
panic_idx: u32,
assert_idx: u32,
) -> Vec<Signal> {
let fragment = format!("void _(){{ {} ;}}", node_text(source, body));
let Some(tree) = parser.parse(fragment.as_bytes(), None) else {
return Vec::new();
};
let base = body.start_position().row as u32;
let bytes = fragment.as_bytes();
let mut out = Vec::new();
let mut cursor = tree_sitter::QueryCursor::new();
for m in cursor.matches(query, tree.root_node(), bytes) {
for c in m.captures {
let kind = if c.index == panic_idx {
SignalKind::Panic
} else if c.index == assert_idx {
SignalKind::Assert
} else {
continue;
};
out.push(Signal {
file_line: base + c.node.start_position().row as u32 + 1,
tier: SignalTier::High,
kind,
evidence: trim_evidence(&node_text(bytes, c.node)),
});
}
}
out
}
pub(crate) fn collect_test_ranges<F>(
root: tree_sitter::Node,
is_test: F,
) -> Vec<std::ops::Range<usize>>
where
F: Fn(tree_sitter::Node) -> bool,
{
let mut ranges = Vec::new();
let mut stack = vec![root];
while let Some(node) = stack.pop() {
if node.id() != root.id() && is_test(node) {
ranges.push(node.start_byte()..node.end_byte());
continue;
}
let mut cursor = node.walk();
stack.extend(node.named_children(&mut cursor));
}
ranges
}
pub(crate) fn in_test_range(ranges: &[std::ops::Range<usize>], node: tree_sitter::Node) -> bool {
ranges.iter().any(|r| r.contains(&node.start_byte()))
}
pub(crate) fn is_js_test_call(node: tree_sitter::Node, source: &[u8]) -> bool {
if node.kind() != "call_expression" {
return false;
}
let Some(callee) = node.child_by_field_name("function") else {
return false;
};
let text = node_text(source, callee);
let base = text.split('.').next().unwrap_or("").trim();
if !matches!(
base,
"describe" | "it" | "test" | "suite" | "context" | "beforeEach" | "afterEach"
) {
return false;
}
let Some(args) = node.child_by_field_name("arguments") else {
return false;
};
let mut cursor = args.walk();
for arg in args.named_children(&mut cursor) {
if matches!(
arg.kind(),
"arrow_function" | "function_expression" | "function"
) {
return true;
}
}
false
}
pub(crate) fn named_field_matches<F>(
node: tree_sitter::Node,
source: &[u8],
field: &str,
pred: F,
) -> bool
where
F: Fn(&str) -> bool,
{
node.child_by_field_name(field)
.map(|n| pred(node_text(source, n).trim()))
.unwrap_or(false)
}
pub(crate) fn trim_evidence(text: &str) -> String {
let one_line = text.replace('\n', " ");
if one_line.chars().count() <= 200 {
one_line.trim().to_string()
} else {
let truncated: String = one_line.chars().take(200).collect();
format!("{}…", truncated.trim_end())
}
}
pub fn sort_canonical(signals: &mut [Signal]) {
signals.sort_by(|a, b| {
let tier_rank = |t: SignalTier| match t {
SignalTier::High => 2,
SignalTier::Medium => 1,
SignalTier::Low => 0,
};
tier_rank(b.tier)
.cmp(&tier_rank(a.tier))
.then(a.file_line.cmp(&b.file_line))
});
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn signal_tier_default_mapping() {
assert_eq!(SignalKind::Panic.default_tier(), SignalTier::High);
assert_eq!(SignalKind::WarnComment.default_tier(), SignalTier::High);
assert_eq!(SignalKind::Assert.default_tier(), SignalTier::High);
assert_eq!(SignalKind::LinterDisable.default_tier(), SignalTier::Medium);
assert_eq!(SignalKind::Guard.default_tier(), SignalTier::Medium);
assert_eq!(SignalKind::UnwrapLike.default_tier(), SignalTier::Medium);
assert_eq!(SignalKind::RawApi.default_tier(), SignalTier::Low);
}
#[test]
fn sort_canonical_orders_by_tier_then_line() {
let mut signals = vec![
Signal {
file_line: 5,
tier: SignalTier::Low,
kind: SignalKind::RawApi,
evidence: "a".into(),
},
Signal {
file_line: 2,
tier: SignalTier::High,
kind: SignalKind::Panic,
evidence: "b".into(),
},
Signal {
file_line: 10,
tier: SignalTier::High,
kind: SignalKind::WarnComment,
evidence: "c".into(),
},
Signal {
file_line: 1,
tier: SignalTier::Medium,
kind: SignalKind::Guard,
evidence: "d".into(),
},
];
sort_canonical(&mut signals);
assert_eq!(signals[0].file_line, 2); assert_eq!(signals[1].file_line, 10); assert_eq!(signals[2].file_line, 1); assert_eq!(signals[3].file_line, 5); }
#[test]
fn language_label_is_stable_snake_case() {
assert_eq!(language_label(Language::Rust), "rust");
assert_eq!(language_label(Language::TypeScript), "typescript");
assert_eq!(language_label(Language::Cpp), "cpp");
assert_eq!(language_label(Language::Haskell), "haskell");
assert_eq!(language_label(Language::Unknown), "unknown");
}
#[test]
fn test_paths_are_recognised() {
for p in [
"tests/integration.rs",
"crates/foo/tests/smoke.rs",
"src/__tests__/helper.js",
"pkg/server_test.go",
"app/test_views.py",
"app/views_test.py",
"src/Button.test.tsx",
"src/Button.spec.js",
] {
assert!(is_test_path(Path::new(p)), "{p} should be a test path");
}
}
#[test]
fn production_paths_are_not_test_paths() {
for p in [
"src/store/durability.rs",
"src/latest/contest.go",
"src/protest.py",
"src/testing.py",
"src/attest.ts",
"src/harness/testutil.go",
] {
assert!(!is_test_path(Path::new(p)), "{p} should not be a test path");
}
}
#[test]
fn skipped_file_returns_empty_report_not_error() {
let dir = std::env::temp_dir().join("mati_enrich_signals_skip");
std::fs::create_dir_all(&dir).unwrap();
let file = dir.join("thing_test.go");
std::fs::write(&file, "func TestX(t *testing.T) { panic(\"x\") }").unwrap();
let report = extract_signals(&file, Language::Go).unwrap();
assert_eq!(report.signal_count, 0);
assert!(report.signals.is_empty());
assert_eq!(report.language, "go");
std::fs::remove_dir_all(&dir).ok();
}
fn write_gated_module(name: &str, parent: &str, decl: &str) -> std::path::PathBuf {
let root = std::env::temp_dir().join(name);
std::fs::remove_dir_all(&root).ok();
let src = root.join("src");
std::fs::create_dir_all(src.join("hooks")).unwrap();
std::fs::write(src.join(parent), decl).unwrap();
let child = src.join("hooks").join("compliance.rs");
std::fs::write(&child, "fn go() { panic!(\"boom\"); assert!(true); }").unwrap();
child
}
#[test]
fn parent_mod_rs_test_gate_yields_empty_report() {
let child = write_gated_module(
"mati_enrich_parent_modrs",
"hooks/mod.rs",
"pub mod other;\n\n#[cfg(test)]\nmod compliance;\n",
);
let report = extract_signals(&child, Language::Rust).unwrap();
assert_eq!(report.signal_count, 0);
assert!(report.signals.is_empty());
assert_eq!(report.language, "rust");
std::fs::remove_dir_all(child.parent().unwrap().parent().unwrap().parent().unwrap()).ok();
}
#[test]
fn parent_sibling_form_test_gate_yields_empty_report() {
let child = write_gated_module(
"mati_enrich_parent_sibling",
"hooks.rs",
"#[cfg(all(test, unix))]\nmod compliance;\n",
);
let report = extract_signals(&child, Language::Rust).unwrap();
assert_eq!(report.signal_count, 0);
std::fs::remove_dir_all(child.parent().unwrap().parent().unwrap().parent().unwrap()).ok();
}
#[test]
fn ungated_parent_declaration_extracts_normally() {
let child = write_gated_module(
"mati_enrich_parent_ungated",
"hooks/mod.rs",
"#[cfg(not(test))]\nmod compliance;\n",
);
let report = extract_signals(&child, Language::Rust).unwrap();
assert!(report.signal_count > 0, "cfg(not(test)) is production code");
std::fs::remove_dir_all(child.parent().unwrap().parent().unwrap().parent().unwrap()).ok();
}
#[test]
fn missing_parent_file_falls_through_to_extraction() {
let root = std::env::temp_dir().join("mati_enrich_parent_missing");
std::fs::remove_dir_all(&root).ok();
std::fs::create_dir_all(root.join("src")).unwrap();
let file = root.join("src").join("orphan.rs");
std::fs::write(&file, "fn go() { panic!(\"boom\"); }").unwrap();
let report = extract_signals(&file, Language::Rust).unwrap();
assert!(report.signal_count > 0);
std::fs::remove_dir_all(&root).ok();
}
#[test]
fn truncate_respects_limit_zero_means_unlimited() {
let mut report = SignalReport {
file: "x".into(),
language: "rust".into(),
signal_count: 3,
signals: vec![
Signal {
file_line: 1,
tier: SignalTier::High,
kind: SignalKind::Panic,
evidence: "a".into(),
};
3
],
};
report.truncate(0); assert_eq!(report.signal_count, 3);
report.truncate(2);
assert_eq!(report.signal_count, 2);
}
}