use std::collections::HashMap;
use std::path::Path;
use crate::classify::classifier::{ClassificationEngine, ClassificationEngineConfig};
use crate::classify::rules::{default_rules, CategoryDef, Rule, RuleSet};
use crate::classify::tiers::weighted_sum::WeightedSumConfig;
use crate::classify::trace::{RuleSources, TraceTier};
fn deploy_rule() -> Rule {
Rule {
id: "deploy".to_string(),
category: "deployment".to_string(),
subcategory: None,
keywords: vec!["deploy:".to_string()],
patterns: vec![],
priority: 110,
confidence: 0.9,
}
}
fn custom_engine(extend_defaults: bool, weighted_sum: bool) -> ClassificationEngine {
custom_engine_with(extend_defaults, weighted_sum, &[])
}
fn custom_engine_with(
extend_defaults: bool,
weighted_sum: bool,
categories: &[&str],
) -> ClassificationEngine {
let ruleset = RuleSet {
version: None,
extend_defaults,
rules: vec![deploy_rule()],
categories: categories
.iter()
.map(|name| CategoryDef {
name: (*name).to_string(),
description: None,
})
.collect(),
buckets: None,
};
let cfg = ClassificationEngineConfig {
weighted_sum: WeightedSumConfig {
enabled: weighted_sum,
..Default::default()
},
..ClassificationEngineConfig::default()
};
ClassificationEngine::new(ruleset, cfg).expect("engine")
}
fn jira_engine() -> ClassificationEngine {
let mut mappings = HashMap::new();
mappings.insert("TQL".to_string(), "bugfix".to_string());
ClassificationEngine::with_taxonomy_and_mappings(
default_rules(),
ClassificationEngineConfig::default(),
Vec::new(),
mappings,
None,
)
.expect("engine")
}
#[test]
fn traced_cascade_matches_untraced_verdicts() {
let engines = [
ClassificationEngine::new(default_rules(), ClassificationEngineConfig::default())
.expect("engine"),
jira_engine(),
custom_engine(false, true),
custom_engine_with(false, true, &["bugfix", "docs", "feature"]),
custom_engine(true, false),
];
let messages: [(&str, bool); 10] = [
("fix: handle null user", false),
("deploy: prod", false),
("TQL-9 login fails", false),
("update the readme wording for install steps", false),
("some unstructured prose about nothing in particular", false),
("Merge branch 'main'", true),
("Revert \"add cache\"", false),
("PROJ-12 adjust thing", false),
("tidy", false),
("", false),
];
for engine in &engines {
let untraced = engine.classify_batch(&messages);
let traced = engine.classify_batch_traced(&messages);
for ((u, t), (msg, _)) in untraced.iter().zip(&traced).zip(&messages) {
assert_eq!(u, &t.verdict, "verdict drifted for {msg:?}");
}
}
}
#[test]
fn traced_cascade_names_rule_sources() {
let trace = |engine: &ClassificationEngine, msg: &str, merge: bool| {
engine
.classify_sync_traced(msg, merge, None, None, None)
.map(|t| (t.trace.tier, t.trace.rule_id))
};
let builtin = jira_engine().with_rule_sources(RuleSources::builtin());
let (tier, id) = trace(&builtin, "fix: handle null user", false).expect("exact");
assert_eq!(tier, TraceTier::Exact);
assert!(id.starts_with("builtin#"), "{id}");
assert_eq!(
trace(&builtin, "TQL-9 login fails", false),
Some((TraceTier::JiraProject, "jira_project:TQL".to_string()))
);
assert_eq!(
trace(&builtin, "zqx vbn wrt plk", false),
Some((TraceTier::CatchAll, "catch_all".to_string()))
);
let mut sources = RuleSources::builtin();
sources.record_file(Path::new("team-rules.yaml"), &[deploy_rule()]);
let custom = custom_engine(false, true).with_rule_sources(sources);
assert_eq!(
trace(&custom, "deploy: prod", false),
Some((TraceTier::Exact, "team-rules.yaml#deploy".to_string()))
);
assert_eq!(
trace(&custom, "fix null pointer regression hotfix", false),
None
);
let declared = custom_engine_with(false, true, &["bugfix"]);
let (tier, id) =
trace(&declared, "fix null pointer regression hotfix", false).expect("weighted");
assert_eq!(tier, TraceTier::WeightedSum);
assert_eq!(id, "weighted_sum:bugfix/keyword");
let fuzzy = custom_engine(true, false);
assert_eq!(
trace(&fuzzy, "Merge branch 'main'", true),
Some((TraceTier::Fuzzy, "fuzzy:merge".to_string()))
);
assert_eq!(
fuzzy.classify_batch_traced(&[("", false)])[0].trace.tier,
TraceTier::Unclassified
);
}