use super::population::CommitRow;
use crate::classify::{ClassificationEngine, TraceTier, TracedVerdict};
use crate::core::config::Config;
pub(crate) struct Resolved {
pub tier: TraceTier,
pub rule_id: String,
pub category: String,
pub confidence: f64,
pub carried: bool,
pub superseded: bool,
}
pub(crate) struct CarryPolicy {
use_llm: bool,
llm_threshold: f64,
llm_scope: crate::core::config::LlmFallbackScope,
llm_categories: Option<Vec<String>>,
external: bool,
repo_map: crate::classify::pipeline_repo_map::RepoCategoryMap,
paths: std::collections::HashMap<i64, Vec<String>>,
}
impl CarryPolicy {
pub(crate) fn from_config(config: &Config) -> crate::classify::Result<Self> {
let c = config.classification.as_ref();
let pipeline = crate::classify::ClassificationPipeline::new(config.clone());
let llm_categories = pipeline
.llm_categories()?
.map(|cats| cats.into_iter().map(|c| c.name).collect());
Ok(Self {
repo_map: pipeline.repo_category_map()?,
paths: std::collections::HashMap::new(),
llm_categories,
use_llm: config.llm.is_some() || c.is_some_and(|c| c.use_llm),
llm_threshold: c.map_or(0.65, |c| c.llm_fallback_threshold),
llm_scope: c.map(|c| c.llm_fallback_scope).unwrap_or_default(),
external: c.is_some_and(|c| !c.no_external && !c.sources.is_empty()),
})
}
pub(crate) fn prepare(
&mut self,
conn: &rusqlite::Connection,
commits: &[&CommitRow],
) -> crate::classify::Result<()> {
self.repo_map.warn_unmatched_keys(conn)?;
self.paths = self
.repo_map
.load_paths(conn, commits.iter().map(|c| (c.id, c.repo.as_str())))?;
Ok(())
}
fn reaches(
&self,
stored: TraceTier,
stored_category: &str,
t: &TracedVerdict,
is_merge: bool,
) -> bool {
if t.trace.tier == TraceTier::RepoMap {
return false;
}
match stored {
TraceTier::Manual => true,
TraceTier::Llm if is_merge => false,
TraceTier::Llm => {
let in_set = self
.llm_categories
.as_ref()
.is_none_or(|cats| cats.iter().any(|n| n == stored_category));
self.use_llm
&& in_set
&& crate::classify::pipeline_llm::llm_eligible(
self.llm_scope,
&t.verdict,
self.llm_threshold,
)
}
TraceTier::RepoCategory => false,
TraceTier::ExternalSource => self.external,
_ => false,
}
}
}
fn stored_override(method: &str, traced: TraceTier) -> Option<(TraceTier, &'static str)> {
match method {
"manual" => Some((TraceTier::Manual, "manual_override")),
"llm_fallback" => Some((TraceTier::Llm, "llm")),
"repo_category_fallback" => Some((TraceTier::RepoCategory, "repo_category")),
"external_source" if !matches!(traced, TraceTier::JiraProject | TraceTier::IssueType) => {
Some((TraceTier::ExternalSource, "external_source"))
}
_ => None,
}
}
pub(crate) fn resolve_verdicts(
engine: &ClassificationEngine,
policy: &CarryPolicy,
commits: &[&CommitRow],
) -> (Vec<Resolved>, u64) {
let pairs: Vec<(&str, bool)> = commits
.iter()
.map(|c| (c.message.as_str(), c.is_merge))
.collect();
let traced = engine.classify_batch_traced(&pairs);
let mut drifted = 0u64;
let resolved = commits
.iter()
.zip(traced)
.map(|(c, t)| {
let paths = policy.paths.get(&c.id).map_or(&[][..], Vec::as_slice);
let mapped =
policy
.repo_map
.traced(&c.repo, c.is_merge, paths, &c.message, engine.taxonomy());
let (t, floor) = match mapped {
Some(m) if !policy.repo_map.is_floor() => (m, None),
m => (t, m),
};
let mut superseded = false;
let mut carried = None;
let mut stored_rule_category = None;
if let Some((cat, conf, method)) = &c.stored {
match stored_override(method, t.trace.tier) {
Some((tier, rule)) if policy.reaches(tier, cat, &t, c.is_merge) => {
carried = Some(Resolved {
tier,
rule_id: rule.to_string(),
category: cat.clone(),
confidence: *conf,
carried: true,
superseded: false,
});
}
Some(_) => superseded = true,
None => stored_rule_category = Some(cat),
}
}
let mut r = carried.unwrap_or(Resolved {
tier: t.trace.tier,
rule_id: t.trace.rule_id,
category: t.verdict.category,
confidence: t.verdict.confidence,
carried: false,
superseded,
});
if let Some(m) = floor.filter(|_| !policy.repo_map.keeps(&r.category, r.confidence)) {
r = Resolved {
tier: m.trace.tier,
rule_id: m.trace.rule_id,
category: m.verdict.category,
confidence: m.verdict.confidence,
carried: false,
superseded: r.carried || r.superseded,
};
}
if stored_rule_category.is_some_and(|cat| cat != &r.category) {
drifted += 1;
}
r
})
.collect();
(resolved, drifted)
}
#[cfg(test)]
mod tests {
use std::io::Write;
use super::*;
use crate::classify::ClassificationPipeline;
use crate::core::config::ClassificationConfig;
fn commit(repo: &str, stored: Option<(&str, &str)>) -> CommitRow {
CommitRow {
id: 1,
sha: "sha-a".into(),
repo: repo.into(),
author_email: "a@x".into(),
timestamp: "2024-01-01T00:00:00Z".into(),
ts: None,
message: "fix: close the security hole".into(),
is_merge: false,
files: 1,
insertions: 1,
deletions: 0,
ticket_id: None,
stored: stored.map(|(c, m)| (c.to_string(), 0.9, m.to_string())),
}
}
#[test]
fn a_mapped_repo_resolves_to_the_repo_map_tier() {
let mut rules = tempfile::Builder::new()
.suffix(".yaml")
.tempfile()
.expect("tempfile");
rules
.write_all(
b"extend_defaults: false\nrules:\n - id: sec\n category: security\n \
keywords: [\"security\"]\ncategories:\n - name: qa\n",
)
.expect("write");
let config = Config {
classification: Some(ClassificationConfig {
rules_files: vec![rules.path().to_path_buf()],
repo_categories: [("e2e".to_string(), "qa".to_string())].into(),
..ClassificationConfig::default()
}),
..Config::default()
};
let engine = ClassificationPipeline::new(config.clone())
.build_rule_engine()
.expect("engine");
let policy = CarryPolicy::from_config(&config).expect("policy");
let rows = [
commit("e2e", None),
commit("e2e", Some(("security", "manual"))),
commit("api", None),
];
let refs: Vec<&CommitRow> = rows.iter().collect();
let (resolved, _) = resolve_verdicts(&engine, &policy, &refs);
for r in &resolved[..2] {
assert_eq!((r.tier.as_str(), r.category.as_str()), ("repo_map", "qa"));
assert_eq!(r.rule_id, "repo_map:e2e");
assert!(!r.carried);
}
assert_eq!(
(resolved[2].tier.as_str(), resolved[2].category.as_str()),
("exact", "security")
);
}
#[test]
fn floor_mode_keeps_an_exception_and_floors_the_rest() {
let mut rules = tempfile::Builder::new()
.suffix(".yaml")
.tempfile()
.expect("tempfile");
rules
.write_all(
b"extend_defaults: false\nrules:\n - id: sec\n category: security\n \
keywords: [\"security\"]\n confidence: 0.9\ncategories:\n - name: qa\n \
- name: new_feature\n",
)
.expect("write");
let config = Config {
classification: Some(ClassificationConfig {
rules_files: vec![rules.path().to_path_buf()],
repo_categories: [("e2e".to_string(), "qa".to_string())].into(),
repo_map: crate::core::config::RepoMapConfig {
mode: crate::core::config::RepoMapMode::Floor,
..Default::default()
},
..ClassificationConfig::default()
}),
..Config::default()
};
let engine = ClassificationPipeline::new(config.clone())
.build_rule_engine()
.expect("engine");
let policy = CarryPolicy::from_config(&config).expect("policy");
let mut quiet = commit("e2e", None);
quiet.message = "zzz qqq vvv".into();
let mut manual = commit("e2e", Some(("new_feature", "manual")));
manual.message = "zzz qqq vvv".into();
let rows = [commit("e2e", None), quiet, manual];
let refs: Vec<&CommitRow> = rows.iter().collect();
let (resolved, _) = resolve_verdicts(&engine, &policy, &refs);
assert_eq!(
(resolved[0].tier.as_str(), resolved[0].category.as_str()),
("exact", "security")
);
for r in &resolved[1..] {
assert_eq!((r.tier.as_str(), r.category.as_str()), ("repo_map", "qa"));
assert_eq!(r.rule_id, "repo_map:e2e");
assert!(!r.carried);
}
assert!(resolved[2].superseded, "the manual verdict was floored");
}
#[test]
fn prepare_reads_the_paths_a_prefix_key_needs() {
let mut rules = tempfile::Builder::new()
.suffix(".yaml")
.tempfile()
.expect("tempfile");
rules
.write_all(
b"extend_defaults: false\nrules:\n - id: sec\n category: security\n \
keywords: [\"security\"]\ncategories:\n - name: qa\n - name: new_feature\n",
)
.expect("write");
let config = Config {
classification: Some(ClassificationConfig {
rules_files: vec![rules.path().to_path_buf()],
repo_categories: [
("mono".to_string(), "qa".to_string()),
("mono:services".to_string(), "new_feature".to_string()),
]
.into(),
..ClassificationConfig::default()
}),
..Config::default()
};
let db = crate::core::db::Database::open_in_memory().expect("db");
let conn = db.connection();
let mut rows = Vec::new();
for (sha, path) in [("sha-svc", "services/a.rs"), ("sha-docs", "docs/x.md")] {
conn.execute(
"INSERT INTO commits (sha, author_name, author_email, timestamp, message, \
repository, is_merge) VALUES (?1, 'a', 'a@x', '2024-01-01T00:00:00Z', \
'zzz qqq vvv', 'mono', 0)",
[sha],
)
.expect("commit");
let id = conn.last_insert_rowid();
conn.execute(
"INSERT INTO files (commit_id, path, change_type) VALUES (?1, ?2, 'M')",
rusqlite::params![id, path],
)
.expect("file");
let mut row = commit("mono", None);
row.id = id;
row.sha = sha.into();
row.message = "zzz qqq vvv".into();
rows.push(row);
}
let engine = ClassificationPipeline::new(config.clone())
.build_rule_engine()
.expect("engine");
let mut policy = CarryPolicy::from_config(&config).expect("policy");
let refs: Vec<&CommitRow> = rows.iter().collect();
policy.prepare(conn, &refs).expect("prepare");
let (resolved, _) = resolve_verdicts(&engine, &policy, &refs);
assert_eq!(
(resolved[0].rule_id.as_str(), resolved[0].category.as_str()),
("repo_map:mono:services", "new_feature")
);
assert_eq!(
(resolved[1].rule_id.as_str(), resolved[1].category.as_str()),
("repo_map:mono", "qa")
);
}
}