use std::collections::HashMap;
use std::path::Path;
use serde::{Deserialize, Serialize};
use crate::classify::rules::Rule;
use crate::classify::tiers::ClassificationResult;
pub const CATCH_ALL_RULE_ID: &str = "catch-all";
pub const BUILTIN_SOURCE: &str = "builtin";
pub const UNKNOWN_SOURCE: &str = "ruleset";
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum TraceTier {
Manual,
Exact,
IssueType,
JiraProject,
Regex,
CatchAll,
WeightedSum,
Fuzzy,
ExternalSource,
Llm,
RepoCategory,
RepoMap,
Unclassified,
}
impl TraceTier {
pub fn as_str(self) -> &'static str {
match self {
Self::Manual => "manual",
Self::Exact => "exact",
Self::IssueType => "issue_type",
Self::JiraProject => "jira_project",
Self::Regex => "regex",
Self::CatchAll => "catch_all",
Self::WeightedSum => "weighted_sum",
Self::Fuzzy => "fuzzy",
Self::ExternalSource => "external_source",
Self::Llm => "llm",
Self::RepoCategory => "repo_category",
Self::RepoMap => "repo_map",
Self::Unclassified => "unclassified",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct RuleTrace {
pub tier: TraceTier,
pub rule_id: String,
}
impl RuleTrace {
pub fn new(tier: TraceTier, rule_id: impl Into<String>) -> Self {
Self {
tier,
rule_id: rule_id.into(),
}
}
pub fn for_rule(tier: TraceTier, rule: &Rule, sources: &RuleSources) -> Self {
let source = sources.source_of(&rule.id);
if rule.id == CATCH_ALL_RULE_ID && tier == TraceTier::Regex {
let rule_id = if source == BUILTIN_SOURCE {
"catch_all".to_string()
} else {
format!("{source}#{}", rule.id)
};
return Self::new(TraceTier::CatchAll, rule_id);
}
Self::new(tier, format!("{source}#{}", rule.id))
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct TracedVerdict {
pub verdict: ClassificationResult,
pub trace: RuleTrace,
}
impl TracedVerdict {
pub fn unclassified() -> Self {
Self {
verdict: ClassificationResult::unclassified(),
trace: RuleTrace::new(TraceTier::Unclassified, "unclassified"),
}
}
}
#[derive(Debug, Clone)]
pub struct RuleSources {
by_id: HashMap<String, String>,
fallback: &'static str,
}
impl Default for RuleSources {
fn default() -> Self {
Self {
by_id: HashMap::new(),
fallback: UNKNOWN_SOURCE,
}
}
}
impl RuleSources {
pub fn builtin() -> Self {
Self {
by_id: HashMap::new(),
fallback: BUILTIN_SOURCE,
}
}
pub fn record_file(&mut self, path: &Path, rules: &[Rule]) {
let label = path.display().to_string();
for rule in rules {
self.by_id.insert(rule.id.clone(), label.clone());
}
}
pub fn source_of(&self, rule_id: &str) -> &str {
self.by_id
.get(rule_id)
.map(String::as_str)
.unwrap_or(self.fallback)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn rule(id: &str) -> Rule {
Rule {
id: id.to_string(),
category: "x".to_string(),
subcategory: None,
keywords: vec![],
patterns: vec![],
priority: 1,
confidence: 0.3,
}
}
#[test]
fn rule_trace_qualifies_by_source() {
let mut sources = RuleSources::builtin();
sources.record_file(Path::new("team/rules.yaml"), &[rule("deploy")]);
let file = RuleTrace::for_rule(TraceTier::Exact, &rule("deploy"), &sources);
assert_eq!(file.rule_id, "team/rules.yaml#deploy");
assert_eq!(file.tier, TraceTier::Exact);
let builtin = RuleTrace::for_rule(TraceTier::Regex, &rule("cc-fix"), &sources);
assert_eq!(builtin.rule_id, "builtin#cc-fix");
let unknown = RuleTrace::for_rule(TraceTier::Regex, &rule("r"), &RuleSources::default());
assert_eq!(unknown.rule_id, "ruleset#r");
}
#[test]
fn rule_trace_names_the_builtin_catch_all() {
let builtin = RuleTrace::for_rule(
TraceTier::Regex,
&rule(CATCH_ALL_RULE_ID),
&RuleSources::builtin(),
);
assert_eq!(builtin, RuleTrace::new(TraceTier::CatchAll, "catch_all"));
let mut sources = RuleSources::builtin();
sources.record_file(Path::new("r.yaml"), &[rule(CATCH_ALL_RULE_ID)]);
let custom = RuleTrace::for_rule(TraceTier::Regex, &rule(CATCH_ALL_RULE_ID), &sources);
assert_eq!(custom.tier, TraceTier::CatchAll);
assert_eq!(custom.rule_id, "r.yaml#catch-all");
}
}