use crate::snapshot::{FunctionSnapshot, Snapshot};
use anyhow::{bail, Context, Result};
use linfa::prelude::Fit;
use linfa::Dataset;
use linfa_ensemble::RandomForestParams;
use linfa_trees::DecisionTreeParams;
use ndarray::{Array1, Array2};
use rand::rngs::SmallRng;
use rand::SeedableRng;
use serde::{Deserialize, Serialize};
use std::collections::HashSet;
use std::path::Path;
pub(crate) fn make_rel(path: &str, prefix_can: &str, prefix_raw: &str) -> String {
let p = path.replace('\\', "/");
let rel = if let Some(r) = p.strip_prefix(prefix_can) {
r.to_string()
} else if let Some(r) = p.strip_prefix(prefix_raw) {
r.to_string()
} else if let Some(r) = std::path::Path::new(&p).canonicalize().ok().and_then(|cp| {
cp.to_str()
.and_then(|s| s.strip_prefix(prefix_can))
.map(str::to_string)
}) {
r
} else {
return p;
};
rel.strip_prefix("./").unwrap_or(&rel).to_string()
}
pub(crate) fn repo_prefixes(repo_root: &Path) -> (String, String) {
let canonical = repo_root
.canonicalize()
.unwrap_or_else(|_| repo_root.to_path_buf());
let prefix_can = format!(
"{}/",
canonical.to_str().unwrap_or("").trim_end_matches('/')
);
let prefix_raw = repo_root
.to_str()
.map(|s| format!("{}/", s.trim_end_matches('/')))
.unwrap_or_default();
(prefix_can, prefix_raw)
}
pub const FEATURE_NAMES: [&str; 8] = [
"lrs",
"cc",
"nd",
"loc",
"fo",
"fan_in",
"total_churn",
"authors_90d",
];
pub fn extract_features(func: &FunctionSnapshot) -> [f64; 8] {
let cg = func.callgraph.as_ref();
let total_churn = func
.churn
.as_ref()
.map(|c| (c.lines_added + c.lines_deleted) as f64)
.unwrap_or(0.0);
[
func.lrs,
f64::from(func.metrics.cc),
f64::from(func.metrics.nd),
f64::from(func.metrics.loc),
f64::from(func.metrics.fo),
cg.map(|c| c.fan_in as f64).unwrap_or(0.0),
total_churn,
func.authors_90d.unwrap_or(0) as f64,
]
}
pub fn collect_fix_files(repo_root: &Path, window_days: u32) -> Result<HashSet<String>> {
use std::process::Command;
let after = format!("{}.days.ago", window_days);
let out = Command::new("git")
.args([
"log",
"--after",
&after,
"--name-only",
"--pretty=format:%s",
"--diff-filter=M",
])
.current_dir(repo_root)
.output()
.context("git log failed")?;
if !out.status.success() {
bail!(
"git log exited {}: {}",
out.status,
String::from_utf8_lossy(&out.stderr)
);
}
let text = String::from_utf8_lossy(&out.stdout);
let mut fix_files = HashSet::new();
let mut in_fix_commit = false;
for line in text.lines() {
if line.is_empty() {
in_fix_commit = false;
continue;
}
if !in_fix_commit && is_fix_message(line) {
in_fix_commit = true;
continue;
}
if in_fix_commit && !line.is_empty() {
fix_files.insert(line.replace('\\', "/"));
}
}
Ok(fix_files)
}
fn is_fix_message(msg: &str) -> bool {
let lower = msg.to_lowercase();
lower.contains("fix")
|| lower.contains("bug")
|| lower.contains("patch")
|| lower.contains("regression")
|| lower.contains("defect")
|| lower.contains("hotfix")
}
pub fn collect_fix_functions(
snapshot: &Snapshot,
repo_root: &Path,
window_days: u32,
) -> Result<HashSet<(String, u32)>> {
use std::process::Command;
let (repo_prefix_canonical, repo_prefix_raw) = repo_prefixes(repo_root);
let mut file_index: std::collections::HashMap<String, Vec<(u32, String)>> =
std::collections::HashMap::new();
for func in &snapshot.functions {
let rel = make_rel(&func.file, &repo_prefix_canonical, &repo_prefix_raw);
file_index
.entry(rel)
.or_default()
.push((func.line, func.function_id.clone()));
}
for entries in file_index.values_mut() {
entries.sort_by_key(|(line, _)| *line);
}
let after = format!("{}.days.ago", window_days);
let out = Command::new("git")
.args([
"log",
"--after",
&after,
"--name-only",
"--pretty=format:%H|%s",
"--diff-filter=M",
])
.current_dir(repo_root)
.output()
.context("git log failed")?;
if !out.status.success() {
bail!(
"git log exited {}: {}",
out.status,
String::from_utf8_lossy(&out.stderr)
);
}
let text = String::from_utf8_lossy(&out.stdout);
let mut fix_shas: Vec<String> = Vec::new();
let mut current_sha: Option<String> = None;
for line in text.lines() {
if line.is_empty() {
current_sha = None;
continue;
}
if let Some((sha, subject)) = line.split_once('|') {
if is_fix_message(subject) {
current_sha = Some(sha.to_string());
fix_shas.push(sha.to_string());
}
continue;
}
let _ = current_sha.as_ref();
}
fix_shas.dedup();
let mut labelled: HashSet<(String, u32)> = HashSet::new();
for sha in &fix_shas {
let diff_out = Command::new("git")
.args(["diff-tree", "--no-commit-id", "-r", "--unified=0", sha])
.current_dir(repo_root)
.output();
let diff_out = match diff_out {
Ok(o) if o.status.success() => o,
_ => continue,
};
let diff_text = String::from_utf8_lossy(&diff_out.stdout);
let mut current_file: Option<String> = None;
for dline in diff_text.lines() {
if let Some(rest) = dline.strip_prefix("+++ b/") {
current_file = Some(rest.replace('\\', "/"));
continue;
}
if let Some(rest) = dline.strip_prefix("@@ ") {
if let Some(file) = ¤t_file {
if let Some(old_start) = parse_hunk_old_start(rest) {
if let Some(func_line) = nearest_function_above(
file_index.get(file).map(Vec::as_slice).unwrap_or(&[]),
old_start,
) {
labelled.insert((file.clone(), func_line));
}
}
}
}
}
}
Ok(labelled)
}
pub(crate) fn parse_hunk_old_start(rest: &str) -> Option<u32> {
let after_minus = rest.strip_prefix('-')?;
let num_str = after_minus.split([',', ' ']).next()?;
num_str.parse::<u32>().ok()
}
pub(crate) fn nearest_function_above(entries: &[(u32, String)], line: u32) -> Option<u32> {
if entries.is_empty() {
return None;
}
let pos = entries.partition_point(|(start, _)| *start <= line);
if pos == 0 {
return None;
}
Some(entries[pos - 1].0)
}
#[derive(Debug, Clone)]
pub struct TrainConfig {
pub label_window_days: u32,
pub n_estimators: usize,
pub max_depth: usize,
pub seed: u64,
pub blame_labels: bool,
}
impl Default for TrainConfig {
fn default() -> Self {
Self {
label_window_days: 365,
n_estimators: 200,
max_depth: 6,
seed: 42,
blame_labels: false,
}
}
}
pub fn train(
snapshot: &Snapshot,
repo_root: &Path,
cfg: &TrainConfig,
) -> Result<Option<RankerModel>> {
let mut rows: Vec<([f64; 8], bool)> = Vec::new();
if cfg.blame_labels {
let fix_funcs = collect_fix_functions(snapshot, repo_root, cfg.label_window_days)?;
let (prefix_can, prefix_raw) = repo_prefixes(repo_root);
for func in &snapshot.functions {
let rel = make_rel(&func.file, &prefix_can, &prefix_raw);
let label = fix_funcs.contains(&(rel, func.line));
rows.push((extract_features(func), label));
}
} else {
let fix_files = collect_fix_files(repo_root, cfg.label_window_days)?;
for func in &snapshot.functions {
let file_norm = func.file.replace('\\', "/");
let label = fix_files.contains(&file_norm)
|| fix_files.iter().any(|f| file_norm.ends_with(f.as_str()));
rows.push((extract_features(func), label));
}
}
let n_pos = rows.iter().filter(|(_, l)| *l).count();
let n_neg = rows.len() - n_pos;
if rows.len() < 50 || n_pos < 5 || n_neg < 10 {
return Ok(None);
}
let n = rows.len();
let mut x_data = Vec::with_capacity(n * 8);
let mut y_data: Vec<bool> = Vec::with_capacity(n);
for (feats, label) in &rows {
x_data.extend_from_slice(feats);
y_data.push(*label);
}
let x: Array2<f64> =
Array2::from_shape_vec((n, 8), x_data).context("feature matrix shape error")?;
let y: Array1<bool> = Array1::from_vec(y_data);
let dataset = Dataset::new(x, y).with_feature_names(
FEATURE_NAMES
.iter()
.map(|s| s.to_string())
.collect::<Vec<_>>(),
);
let rng = SmallRng::seed_from_u64(cfg.seed);
let tree_params = DecisionTreeParams::new().max_depth(Some(cfg.max_depth));
let rf = RandomForestParams::new_fixed_rng(tree_params, rng.clone())
.ensemble_size(cfg.n_estimators)
.bootstrap_proportion(0.8)
.fit(&dataset)
.context("RandomForest fit failed")?;
let mut trees: Vec<SerializedTree> = Vec::with_capacity(rf.models.len());
for (tree, feat_indices) in rf.models.iter().zip(rf.model_features.iter()) {
let nodes = serialize_tree(tree);
trees.push(SerializedTree {
nodes,
feature_indices: feat_indices.clone(),
});
}
let meta = TrainMeta {
n_samples: n,
n_pos,
n_neg,
label_window_days: cfg.label_window_days,
n_estimators: cfg.n_estimators,
max_depth: cfg.max_depth,
};
Ok(Some(RankerModel {
model_version: 3,
trees,
meta,
}))
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TreeNodeRecord {
pub feature_idx: usize,
pub threshold: f64,
pub left: usize,
pub right: usize,
pub leaf_value: bool,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SerializedTree {
pub nodes: Vec<TreeNodeRecord>,
pub feature_indices: Vec<usize>,
}
fn serialize_tree<F, L>(tree: &linfa_trees::DecisionTree<F, L>) -> Vec<TreeNodeRecord>
where
F: linfa::Float,
L: linfa::Label + Into<bool> + Copy,
{
let mut nodes: Vec<TreeNodeRecord> = Vec::new();
serialize_node(tree.root_node(), &mut nodes);
nodes
}
fn serialize_node<F, L>(
node: &linfa_trees::TreeNode<F, L>,
nodes: &mut Vec<TreeNodeRecord>,
) -> usize
where
F: linfa::Float,
L: linfa::Label + Into<bool> + Copy,
{
let idx = nodes.len();
nodes.push(TreeNodeRecord {
feature_idx: usize::MAX,
threshold: 0.0,
left: 0,
right: 0,
leaf_value: false,
});
if node.is_leaf() {
let val: bool = node.prediction().map(|p| p.into()).unwrap_or(false);
nodes[idx].leaf_value = val;
} else {
let (feat, thresh, _) = node.split();
nodes[idx].feature_idx = feat;
nodes[idx].threshold = thresh.to_f64().unwrap_or(0.0);
let children = node.children();
let left_child = children[0].as_deref();
let right_child = children[1].as_deref();
let left_idx = if let Some(lc) = left_child {
serialize_node(lc, nodes)
} else {
idx };
let right_idx = if let Some(rc) = right_child {
serialize_node(rc, nodes)
} else {
idx
};
nodes[idx].left = left_idx;
nodes[idx].right = right_idx;
}
idx
}
pub fn score(model: &RankerModel, func: &FunctionSnapshot) -> f64 {
let feats = extract_features(func);
let n = model.trees.len();
if n == 0 {
return 0.0;
}
let votes: usize = model
.trees
.iter()
.map(|tree| vote(tree, &feats) as usize)
.sum();
votes as f64 / n as f64
}
fn vote(tree: &SerializedTree, feats: &[f64; 8]) -> bool {
let nodes = &tree.nodes;
if nodes.is_empty() {
return false;
}
let mut cur = 0usize;
loop {
let node = &nodes[cur];
if node.feature_idx == usize::MAX {
return node.leaf_value;
}
let global_idx = tree
.feature_indices
.get(node.feature_idx)
.copied()
.unwrap_or(node.feature_idx);
let val = if global_idx < 8 {
feats[global_idx]
} else {
0.0
};
if val <= node.threshold {
cur = node.left;
} else {
cur = node.right;
}
if cur >= nodes.len() {
return false;
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TrainMeta {
pub n_samples: usize,
pub n_pos: usize,
pub n_neg: usize,
pub label_window_days: u32,
pub n_estimators: usize,
pub max_depth: usize,
}
fn default_model_version() -> u32 {
1
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RankerModel {
#[serde(default = "default_model_version")]
pub model_version: u32,
pub trees: Vec<SerializedTree>,
pub meta: TrainMeta,
}
impl RankerModel {
pub fn save(&self, path: &Path) -> Result<()> {
let json = serde_json::to_string_pretty(self).context("serialize model")?;
std::fs::write(path, json).context("write model file")?;
Ok(())
}
pub fn load(path: &Path) -> Result<Self> {
let json = std::fs::read_to_string(path).context("read model file")?;
let model: Self = serde_json::from_str(&json).context("deserialize model")?;
if model.model_version < 3 {
bail!(
"{} was trained with an older feature set (model_version={}). \
Run `hotspots train` to retrain with the current feature set.",
path.display(),
model.model_version
);
}
Ok(model)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn feature_names_count_matches_array() {
assert_eq!(FEATURE_NAMES.len(), 8);
}
#[test]
fn no_leaky_windowed_features() {
for name in FEATURE_NAMES {
assert_ne!(name, "touch_count_30d", "leaky feature still present");
assert_ne!(
name, "days_since_last_change",
"leaky feature still present"
);
assert_ne!(name, "activity_risk", "leaky feature still present");
}
assert!(FEATURE_NAMES.contains(&"total_churn"));
}
#[test]
fn parse_hunk_simple() {
assert_eq!(parse_hunk_old_start("-42,5 +50,3 @@"), Some(42));
}
#[test]
fn parse_hunk_no_count() {
assert_eq!(parse_hunk_old_start("-10 +10 @@"), Some(10));
}
#[test]
fn parse_hunk_wrong_prefix() {
assert_eq!(parse_hunk_old_start("+42,5 -50,3 @@"), None);
assert_eq!(parse_hunk_old_start(""), None);
}
#[test]
fn nearest_exact_match() {
let entries = vec![
(1u32, "a".to_string()),
(10u32, "b".to_string()),
(20u32, "c".to_string()),
];
assert_eq!(nearest_function_above(&entries, 10), Some(10));
}
#[test]
fn nearest_between_funcs() {
let entries = vec![
(1u32, "a".to_string()),
(10u32, "b".to_string()),
(20u32, "c".to_string()),
];
assert_eq!(nearest_function_above(&entries, 15), Some(10));
}
#[test]
fn nearest_before_first_func() {
let entries = vec![(10u32, "a".to_string())];
assert_eq!(nearest_function_above(&entries, 5), None);
}
#[test]
fn nearest_empty_entries() {
assert_eq!(nearest_function_above(&[], 42), None);
}
#[test]
fn nearest_after_last_func() {
let entries = vec![
(1u32, "a".to_string()),
(10u32, "b".to_string()),
(50u32, "c".to_string()),
];
assert_eq!(nearest_function_above(&entries, 999), Some(50));
}
#[test]
fn load_v1_model_returns_error() {
let dir = tempfile::tempdir().expect("tempdir");
let path = dir.path().join("ranker.json");
std::fs::write(
&path,
r#"{"trees":[],"meta":{"n_samples":100,"n_pos":50,"n_neg":50,"label_window_days":365,"n_estimators":10,"max_depth":3}}"#,
)
.unwrap();
let err = RankerModel::load(&path).unwrap_err().to_string();
assert!(
err.contains("retrain") || err.contains("model_version"),
"unexpected error: {err}"
);
}
#[test]
fn load_v2_model_returns_error() {
let dir = tempfile::tempdir().expect("tempdir");
let path = dir.path().join("ranker.json");
std::fs::write(
&path,
r#"{"model_version":2,"trees":[],"meta":{"n_samples":100,"n_pos":50,"n_neg":50,"label_window_days":365,"n_estimators":10,"max_depth":3}}"#,
)
.unwrap();
let err = RankerModel::load(&path).unwrap_err().to_string();
assert!(
err.contains("retrain") || err.contains("model_version"),
"unexpected error: {err}"
);
}
#[test]
fn load_v3_model_succeeds() {
let dir = tempfile::tempdir().expect("tempdir");
let path = dir.path().join("ranker.json");
std::fs::write(
&path,
r#"{"model_version":3,"trees":[],"meta":{"n_samples":100,"n_pos":50,"n_neg":50,"label_window_days":365,"n_estimators":10,"max_depth":3}}"#,
)
.unwrap();
let model = RankerModel::load(&path).expect("should load");
assert_eq!(model.model_version, 3);
}
}