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,
}
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()?,
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()),
})
}
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 t = policy
.repo_map
.traced(&c.repo, c.is_merge, &c.message, engine.taxonomy())
.unwrap_or(t);
let mut superseded = false;
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) => {
return Resolved {
tier,
rule_id: rule.to_string(),
category: cat.clone(),
confidence: *conf,
carried: true,
superseded: false,
};
}
Some(_) => superseded = true,
None if cat != &t.verdict.category => drifted += 1,
None => {}
}
}
Resolved {
tier: t.trace.tier,
rule_id: t.trace.rule_id,
category: t.verdict.category,
confidence: t.verdict.confidence,
carried: false,
superseded,
}
})
.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")
);
}
}