use crate::{Rule, RuleAtom};
use anyhow::Result;
use std::cmp::Ordering;
use std::collections::HashMap;
use tracing::{debug, info, warn};
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Default)]
pub enum Priority {
VeryLow = 0,
Low = 1,
#[default]
Normal = 2,
High = 3,
VeryHigh = 4,
}
#[derive(Debug, Clone, PartialEq)]
pub enum ResolutionStrategy {
Priority,
Specificity,
Combined,
KeepAll,
RejectAll,
Confidence,
}
#[derive(Debug, Clone)]
pub struct Conflict {
pub rules: Vec<String>,
pub facts: Vec<RuleAtom>,
pub description: String,
pub severity: ConflictSeverity,
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
pub enum ConflictSeverity {
Low,
Medium,
High,
Critical,
}
pub struct ConflictResolver {
priorities: HashMap<String, Priority>,
specificity: HashMap<String, f64>,
strategy: ResolutionStrategy,
conflicts: Vec<Conflict>,
confidence_scores: HashMap<RuleAtom, f64>,
}
impl Default for ConflictResolver {
fn default() -> Self {
Self::new()
}
}
impl ConflictResolver {
pub fn new() -> Self {
Self {
priorities: HashMap::new(),
specificity: HashMap::new(),
strategy: ResolutionStrategy::Combined,
conflicts: Vec::new(),
confidence_scores: HashMap::new(),
}
}
pub fn set_priority(&mut self, rule_name: &str, priority: Priority) {
self.priorities.insert(rule_name.to_string(), priority);
debug!("Set priority for rule '{}': {:?}", rule_name, priority);
}
pub fn get_priority(&self, rule_name: &str) -> Priority {
self.priorities.get(rule_name).copied().unwrap_or_default()
}
pub fn set_strategy(&mut self, strategy: ResolutionStrategy) {
info!("Set resolution strategy: {:?}", strategy);
self.strategy = strategy;
}
pub fn calculate_specificity(&mut self, rule: &Rule) -> f64 {
let body_size = rule.body.len() as f64;
let mut variable_count = 0;
let mut constant_count = 0;
for atom in &rule.body {
self.count_terms_in_atom(atom, &mut variable_count, &mut constant_count);
}
let specificity = body_size + (constant_count as f64 * 2.0) - (variable_count as f64 * 0.5);
self.specificity.insert(rule.name.clone(), specificity);
specificity
}
fn count_terms_in_atom(&self, atom: &RuleAtom, variables: &mut usize, constants: &mut usize) {
match atom {
RuleAtom::Triple {
subject,
predicate,
object,
} => {
Self::count_term(subject, variables, constants);
Self::count_term(predicate, variables, constants);
Self::count_term(object, variables, constants);
}
RuleAtom::Builtin { args, .. } => {
for arg in args {
Self::count_term(arg, variables, constants);
}
}
RuleAtom::NotEqual { left, right }
| RuleAtom::GreaterThan { left, right }
| RuleAtom::LessThan { left, right } => {
Self::count_term(left, variables, constants);
Self::count_term(right, variables, constants);
}
}
}
fn count_term(term: &crate::Term, variables: &mut usize, constants: &mut usize) {
match term {
crate::Term::Variable(_) => *variables += 1,
crate::Term::Constant(_) | crate::Term::Literal(_) => *constants += 1,
crate::Term::Function { args, .. } => {
for arg in args {
Self::count_term(arg, variables, constants);
}
}
}
}
pub fn resolve_conflicts(
&mut self,
rules: &[Rule],
facts: &[RuleAtom],
) -> Result<Vec<RuleAtom>> {
info!("Resolving conflicts using strategy: {:?}", self.strategy);
self.detect_conflicts(rules, facts)?;
let resolved = match self.strategy {
ResolutionStrategy::Priority => self.resolve_by_priority(rules, facts)?,
ResolutionStrategy::Specificity => self.resolve_by_specificity(rules, facts)?,
ResolutionStrategy::Combined => self.resolve_combined(rules, facts)?,
ResolutionStrategy::KeepAll => facts.to_vec(),
ResolutionStrategy::RejectAll => {
if !self.conflicts.is_empty() {
warn!("Rejecting all facts due to conflicts");
vec![]
} else {
facts.to_vec()
}
}
ResolutionStrategy::Confidence => self.resolve_by_confidence(facts)?,
};
info!(
"Conflict resolution complete: {} facts retained",
resolved.len()
);
Ok(resolved)
}
fn detect_conflicts(&mut self, rules: &[Rule], facts: &[RuleAtom]) -> Result<()> {
self.conflicts.clear();
let mut groups: HashMap<(String, String), Vec<(RuleAtom, String)>> = HashMap::new();
for fact in facts {
if let RuleAtom::Triple {
subject,
predicate,
object,
} = fact
{
let key = (Self::term_key(subject), Self::term_key(predicate));
groups
.entry(key)
.or_default()
.push((fact.clone(), Self::term_key(object)));
}
}
for members in groups.values() {
let distinct_objects: std::collections::HashSet<&String> =
members.iter().map(|(_, o)| o).collect();
if distinct_objects.len() > 1 {
let mut involved: Vec<String> = Vec::new();
for (fact, _) in members {
for name in self.deriving_rule_names(rules, fact) {
if !involved.contains(&name) {
involved.push(name);
}
}
}
self.conflicts.push(Conflict {
rules: involved,
facts: members.iter().map(|(f, _)| f.clone()).collect(),
description: "Competing values for the same subject/predicate".to_string(),
severity: ConflictSeverity::High,
});
}
}
for fact in facts {
let names = self.deriving_rule_names(rules, fact);
if names.len() > 1 {
self.conflicts.push(Conflict {
rules: names,
facts: vec![fact.clone()],
description: "Multiple derivation paths".to_string(),
severity: ConflictSeverity::Low,
});
}
}
debug!("Detected {} conflicts", self.conflicts.len());
Ok(())
}
fn deriving_rule_names(&self, rules: &[Rule], fact: &RuleAtom) -> Vec<String> {
let mut names = Vec::new();
for rule in rules {
if rule.head.iter().any(|h| Self::head_atom_matches(h, fact))
&& !names.contains(&rule.name)
{
names.push(rule.name.clone());
}
}
names
}
fn head_atom_matches(head: &RuleAtom, fact: &RuleAtom) -> bool {
match (head, fact) {
(
RuleAtom::Triple {
subject: hs,
predicate: hp,
object: ho,
},
RuleAtom::Triple {
subject: fs,
predicate: fp,
object: fo,
},
) => {
Self::term_matches(hs, fs)
&& Self::term_matches(hp, fp)
&& Self::term_matches(ho, fo)
}
(
RuleAtom::NotEqual {
left: hl,
right: hr,
},
RuleAtom::NotEqual {
left: fl,
right: fr,
},
)
| (
RuleAtom::GreaterThan {
left: hl,
right: hr,
},
RuleAtom::GreaterThan {
left: fl,
right: fr,
},
)
| (
RuleAtom::LessThan {
left: hl,
right: hr,
},
RuleAtom::LessThan {
left: fl,
right: fr,
},
) => Self::term_matches(hl, fl) && Self::term_matches(hr, fr),
_ => false,
}
}
fn term_matches(pattern: &crate::Term, value: &crate::Term) -> bool {
use crate::Term;
match (pattern, value) {
(Term::Variable(_), _) => true,
(Term::Constant(a), Term::Constant(b)) => a == b,
(Term::Literal(a), Term::Literal(b)) => a == b,
(Term::Constant(a), Term::Literal(b)) | (Term::Literal(a), Term::Constant(b)) => a == b,
_ => false,
}
}
fn term_key(term: &crate::Term) -> String {
use crate::Term;
match term {
Term::Variable(v) => format!("?{v}"),
Term::Constant(c) => c.clone(),
Term::Literal(l) => format!("\"{l}\""),
Term::Function { name, args } => {
let inner: Vec<String> = args.iter().map(Self::term_key).collect();
format!("{name}({})", inner.join(","))
}
}
}
fn rule_specificity_score(rule: &Rule) -> f64 {
let body_size = rule.body.len() as f64;
let mut variables = 0usize;
let mut constants = 0usize;
for atom in &rule.body {
match atom {
RuleAtom::Triple {
subject,
predicate,
object,
} => {
Self::count_term(subject, &mut variables, &mut constants);
Self::count_term(predicate, &mut variables, &mut constants);
Self::count_term(object, &mut variables, &mut constants);
}
RuleAtom::Builtin { args, .. } => {
for arg in args {
Self::count_term(arg, &mut variables, &mut constants);
}
}
RuleAtom::NotEqual { left, right }
| RuleAtom::GreaterThan { left, right }
| RuleAtom::LessThan { left, right } => {
Self::count_term(left, &mut variables, &mut constants);
Self::count_term(right, &mut variables, &mut constants);
}
}
}
body_size + (constants as f64 * 2.0) - (variables as f64 * 0.5)
}
fn fact_score<F: Fn(&Rule) -> f64>(&self, rules: &[Rule], fact: &RuleAtom, metric: &F) -> f64 {
let mut best = f64::NEG_INFINITY;
let mut any = false;
for rule in rules {
if rule.head.iter().any(|h| Self::head_atom_matches(h, fact)) {
any = true;
best = best.max(metric(rule));
}
}
if any {
best
} else {
f64::INFINITY
}
}
fn resolve_by_metric<F: Fn(&Rule) -> f64>(
&self,
rules: &[Rule],
facts: &[RuleAtom],
metric: F,
) -> Vec<RuleAtom> {
let scores: Vec<f64> = facts
.iter()
.map(|f| self.fact_score(rules, f, &metric))
.collect();
let mut group_max: HashMap<(String, String), f64> = HashMap::new();
for (idx, fact) in facts.iter().enumerate() {
if let RuleAtom::Triple {
subject, predicate, ..
} = fact
{
let key = (Self::term_key(subject), Self::term_key(predicate));
let entry = group_max.entry(key).or_insert(f64::NEG_INFINITY);
*entry = entry.max(scores[idx]);
}
}
let mut resolved = Vec::new();
for (idx, fact) in facts.iter().enumerate() {
let keep = match fact {
RuleAtom::Triple {
subject, predicate, ..
} => {
let key = (Self::term_key(subject), Self::term_key(predicate));
let max = group_max.get(&key).copied().unwrap_or(scores[idx]);
scores[idx] == max || (scores[idx] - max).abs() < f64::EPSILON
}
_ => true,
};
if keep {
resolved.push(fact.clone());
}
}
resolved
}
fn resolve_by_priority(&self, rules: &[Rule], facts: &[RuleAtom]) -> Result<Vec<RuleAtom>> {
Ok(self.resolve_by_metric(rules, facts, |rule| {
self.get_priority(&rule.name) as u8 as f64
}))
}
fn resolve_by_specificity(&self, rules: &[Rule], facts: &[RuleAtom]) -> Result<Vec<RuleAtom>> {
Ok(self.resolve_by_metric(rules, facts, Self::rule_specificity_score))
}
fn resolve_combined(&self, rules: &[Rule], facts: &[RuleAtom]) -> Result<Vec<RuleAtom>> {
Ok(self.resolve_by_metric(rules, facts, |rule| {
(self.get_priority(&rule.name) as u8 as f64) * 1_000_000.0
+ Self::rule_specificity_score(rule)
}))
}
fn resolve_by_confidence(&self, facts: &[RuleAtom]) -> Result<Vec<RuleAtom>> {
let mut scored_facts: Vec<(RuleAtom, f64)> = facts
.iter()
.map(|f| {
let confidence = self.confidence_scores.get(f).copied().unwrap_or(1.0);
(f.clone(), confidence)
})
.collect();
scored_facts.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(Ordering::Equal));
let threshold = 0.5;
let resolved: Vec<RuleAtom> = scored_facts
.into_iter()
.filter(|(_, conf)| *conf >= threshold)
.map(|(fact, _)| fact)
.collect();
Ok(resolved)
}
pub fn set_confidence(&mut self, fact: RuleAtom, confidence: f64) {
self.confidence_scores
.insert(fact, confidence.clamp(0.0, 1.0));
}
pub fn get_conflicts(&self) -> &[Conflict] {
&self.conflicts
}
pub fn has_conflicts(&self) -> bool {
!self.conflicts.is_empty()
}
pub fn get_stats(&self) -> ConflictStats {
let critical_conflicts = self
.conflicts
.iter()
.filter(|c| c.severity == ConflictSeverity::Critical)
.count();
let high_conflicts = self
.conflicts
.iter()
.filter(|c| c.severity == ConflictSeverity::High)
.count();
ConflictStats {
total_conflicts: self.conflicts.len(),
critical_conflicts,
high_conflicts,
rules_with_priority: self.priorities.len(),
active_strategy: self.strategy.clone(),
}
}
}
#[derive(Debug, Clone)]
pub struct ConflictStats {
pub total_conflicts: usize,
pub critical_conflicts: usize,
pub high_conflicts: usize,
pub rules_with_priority: usize,
pub active_strategy: ResolutionStrategy,
}
impl std::fmt::Display for ConflictStats {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"Conflicts: {} (critical: {}, high: {}), Rules with priority: {}, Strategy: {:?}",
self.total_conflicts,
self.critical_conflicts,
self.high_conflicts,
self.rules_with_priority,
self.active_strategy
)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::Term;
#[test]
fn test_priority_setting() {
let mut resolver = ConflictResolver::new();
resolver.set_priority("rule1", Priority::High);
assert_eq!(resolver.get_priority("rule1"), Priority::High);
assert_eq!(resolver.get_priority("rule2"), Priority::Normal);
}
#[test]
fn test_specificity_calculation() {
let mut resolver = ConflictResolver::new();
let rule = Rule {
name: "test".to_string(),
body: vec![RuleAtom::Triple {
subject: Term::Variable("X".to_string()),
predicate: Term::Constant("type".to_string()),
object: Term::Constant("Person".to_string()),
}],
head: vec![],
};
let specificity = resolver.calculate_specificity(&rule);
assert!(specificity > 0.0);
}
#[test]
fn test_conflict_detection() -> Result<(), Box<dyn std::error::Error>> {
let mut resolver = ConflictResolver::new();
let rules = vec![];
let facts = vec![RuleAtom::Triple {
subject: Term::Constant("john".to_string()),
predicate: Term::Constant("age".to_string()),
object: Term::Literal("30".to_string()),
}];
resolver.detect_conflicts(&rules, &facts)?;
assert_eq!(resolver.conflicts.len(), 0);
Ok(())
}
#[test]
fn test_confidence_scoring() -> Result<(), Box<dyn std::error::Error>> {
let mut resolver = ConflictResolver::new();
let fact = RuleAtom::Triple {
subject: Term::Constant("john".to_string()),
predicate: Term::Constant("type".to_string()),
object: Term::Constant("Person".to_string()),
};
resolver.set_confidence(fact.clone(), 0.9);
let facts = vec![fact];
resolver.set_strategy(ResolutionStrategy::Confidence);
let resolved = resolver.resolve_conflicts(&[], &facts)?;
assert_eq!(resolved.len(), 1);
Ok(())
}
#[test]
fn test_stats() {
let resolver = ConflictResolver::new();
let stats = resolver.get_stats();
assert_eq!(stats.total_conflicts, 0);
assert_eq!(stats.rules_with_priority, 0);
}
#[test]
fn regression_priority_resolution_drops_competing_fact(
) -> Result<(), Box<dyn std::error::Error>> {
let mut resolver = ConflictResolver::new();
let high_fact = RuleAtom::Triple {
subject: Term::Constant("john".to_string()),
predicate: Term::Constant("age".to_string()),
object: Term::Literal("30".to_string()),
};
let low_fact = RuleAtom::Triple {
subject: Term::Constant("john".to_string()),
predicate: Term::Constant("age".to_string()),
object: Term::Literal("25".to_string()),
};
let rules = vec![
Rule {
name: "specific".to_string(),
body: vec![],
head: vec![high_fact.clone()],
},
Rule {
name: "general".to_string(),
body: vec![],
head: vec![low_fact.clone()],
},
];
resolver.set_priority("specific", Priority::VeryHigh);
resolver.set_priority("general", Priority::Low);
resolver.set_strategy(ResolutionStrategy::Priority);
let resolved =
resolver.resolve_conflicts(&rules, &[high_fact.clone(), low_fact.clone()])?;
assert!(
resolved.contains(&high_fact),
"high-priority fact must survive: {resolved:?}"
);
assert!(
!resolved.contains(&low_fact),
"low-priority competing fact must be dropped: {resolved:?}"
);
assert!(resolver.has_conflicts());
let named: Vec<String> = resolver
.get_conflicts()
.iter()
.flat_map(|c| c.rules.clone())
.collect();
assert!(named.contains(&"specific".to_string()));
assert!(named.contains(&"general".to_string()));
assert!(!named.contains(&"source".to_string()));
Ok(())
}
}