use std::collections::{HashMap, HashSet, VecDeque};
use super::{
DatalogAtom, DatalogError, DatalogFact, DatalogProgram, DatalogRule, DatalogTerm, DatalogValue,
FactDatabase, Substitution,
};
pub fn unify(term: &DatalogTerm, value: &DatalogValue, sub: &mut Substitution) -> bool {
match term {
DatalogTerm::Constant(c) => c == value,
DatalogTerm::Variable(v) => {
if let Some(existing) = sub.get(v) {
existing == value
} else {
sub.insert(v.clone(), value.clone());
true
}
}
}
}
fn apply_subst_to_atom(atom: &DatalogAtom, sub: &Substitution) -> DatalogAtom {
DatalogAtom {
predicate: atom.predicate.clone(),
terms: atom
.terms
.iter()
.map(|t| match t {
DatalogTerm::Variable(v) => sub
.get(v)
.map(|val| DatalogTerm::Constant(val.clone()))
.unwrap_or_else(|| DatalogTerm::Variable(v.clone())),
DatalogTerm::Constant(_) => t.clone(),
})
.collect(),
}
}
fn extend_substitutions(
partial_subs: Vec<Substitution>,
atom: &DatalogAtom,
db: &FactDatabase,
) -> Vec<Substitution> {
let mut result = Vec::new();
for sub in &partial_subs {
let atom_inst = apply_subst_to_atom(atom, sub);
for tuple in db.tuples_for(&atom_inst.predicate) {
if atom_inst.terms.len() != tuple.len() {
continue;
}
let mut new_sub = sub.clone();
let mut ok = true;
for (term, value) in atom_inst.terms.iter().zip(tuple.iter()) {
if !unify(term, value, &mut new_sub) {
ok = false;
break;
}
}
if ok {
result.push(new_sub);
}
}
}
result
}
pub fn apply_rule(
rule: &DatalogRule,
edb: &FactDatabase,
idb: &FactDatabase,
) -> HashSet<DatalogFact> {
let mut merged = edb.clone();
merged.merge(idb);
apply_rule_against(&merged, rule)
}
fn apply_rule_against(db: &FactDatabase, rule: &DatalogRule) -> HashSet<DatalogFact> {
let mut derived = HashSet::new();
if rule.body.is_empty() {
if let Some(fact) = ground_atom(&rule.head, &HashMap::new()) {
derived.insert(fact);
}
return derived;
}
let mut subs: Vec<Substitution> = vec![HashMap::new()];
for atom in &rule.body {
subs = extend_substitutions(subs, atom, db);
if subs.is_empty() {
return derived; }
}
for sub in subs {
if let Some(fact) = ground_atom(&rule.head, &sub) {
derived.insert(fact);
}
}
derived
}
fn ground_atom(atom: &DatalogAtom, sub: &Substitution) -> Option<DatalogFact> {
let mut args = Vec::with_capacity(atom.terms.len());
for term in &atom.terms {
match term {
DatalogTerm::Constant(c) => args.push(c.clone()),
DatalogTerm::Variable(v) => {
args.push(sub.get(v)?.clone());
}
}
}
Some(DatalogFact {
predicate: atom.predicate.clone(),
args,
})
}
struct DepEdge {
from: String,
to: String,
negated: bool,
}
fn build_dependency_graph(rules: &[DatalogRule]) -> Vec<DepEdge> {
let mut edges = Vec::new();
for rule in rules {
let head_pred = &rule.head.predicate;
for body_atom in &rule.body {
edges.push(DepEdge {
from: head_pred.clone(),
to: body_atom.predicate.clone(),
negated: false, });
}
}
edges
}
fn stratify(rules: &[DatalogRule]) -> Result<Vec<String>, DatalogError> {
let edges = build_dependency_graph(rules);
let mut predicates: HashSet<String> = HashSet::new();
for rule in rules {
predicates.insert(rule.head.predicate.clone());
for atom in &rule.body {
predicates.insert(atom.predicate.clone());
}
}
for edge in &edges {
if edge.negated {
if has_path(&edges, &edge.to, &edge.from) {
return Err(DatalogError::StratificationError(format!(
"cyclic negation between '{}' and '{}'",
edge.from, edge.to
)));
}
}
}
let mut in_degree: HashMap<String, usize> = HashMap::new();
let mut reverse_adj: HashMap<String, Vec<String>> = HashMap::new();
for pred in &predicates {
in_degree.entry(pred.clone()).or_insert(0);
reverse_adj.entry(pred.clone()).or_default();
}
for edge in &edges {
*in_degree.entry(edge.from.clone()).or_insert(0) += 1;
reverse_adj
.entry(edge.to.clone())
.or_default()
.push(edge.from.clone());
}
let mut queue: VecDeque<String> = in_degree
.iter()
.filter(|(_, °)| deg == 0)
.map(|(p, _)| p.clone())
.collect();
let mut order = Vec::new();
while let Some(pred) = queue.pop_front() {
order.push(pred.clone());
if let Some(dependents) = reverse_adj.get(&pred) {
for dep in dependents {
let deg = in_degree.entry(dep.clone()).or_insert(0);
if *deg > 0 {
*deg -= 1;
}
if *deg == 0 {
queue.push_back(dep.clone());
}
}
}
}
if order.len() < predicates.len() {
for pred in &predicates {
if !order.contains(pred) {
order.push(pred.clone());
}
}
}
Ok(order)
}
fn has_path(edges: &[DepEdge], start: &str, goal: &str) -> bool {
let mut visited = HashSet::new();
let mut stack = vec![start.to_string()];
while let Some(current) = stack.pop() {
if current == goal {
return true;
}
if visited.contains(¤t) {
continue;
}
visited.insert(current.clone());
for edge in edges {
if edge.from == current {
stack.push(edge.to.clone());
}
}
}
false
}
pub struct SemiNaiveEvaluator {}
impl SemiNaiveEvaluator {
pub fn new() -> Self {
Self {}
}
pub fn evaluate(&self, program: &DatalogProgram) -> Result<FactDatabase, DatalogError> {
let _strata_order = stratify(&program.rules)?;
let idb_preds: HashSet<String> = program.idb_predicates();
let mut edb = FactDatabase::new();
for fact in &program.edb {
edb.insert_fact(fact);
}
let mut idb = FactDatabase::new();
let mut delta: HashMap<String, HashSet<Vec<DatalogValue>>> = HashMap::new();
for pred in &idb_preds {
delta.insert(pred.clone(), HashSet::new());
}
for rule in &program.rules {
let new_facts = apply_rule(rule, &edb, &idb);
for fact in new_facts {
let args = fact.args.clone();
let pred = fact.predicate.clone();
if idb.insert_fact(&fact) {
delta.entry(pred).or_default().insert(args);
}
}
}
loop {
let prev_idb_size = idb.len();
let mut new_delta: HashMap<String, HashSet<Vec<DatalogValue>>> = HashMap::new();
for pred in &idb_preds {
new_delta.insert(pred.clone(), HashSet::new());
}
for rule in &program.rules {
let head_pred = &rule.head.predicate;
if rule.body.is_empty() {
continue;
}
let body_len = rule.body.len();
for i in 0..body_len {
let delta_pred = &rule.body[i].predicate;
let delta_tuples = match delta.get(delta_pred) {
Some(d) if !d.is_empty() => d.clone(),
_ => continue, };
let mut delta_db = FactDatabase::new();
for tuple in &delta_tuples {
delta_db.insert(delta_pred, tuple.clone());
}
let mut subs: Vec<Substitution> = vec![HashMap::new()];
for (j, body_atom) in rule.body.iter().enumerate() {
let db_to_use = if j == i { &delta_db } else { &idb };
if idb_preds.contains(&body_atom.predicate) || j == i {
subs = extend_substitutions(subs, body_atom, db_to_use);
} else {
subs = extend_substitutions(subs, body_atom, &edb);
}
if subs.is_empty() {
break;
}
}
for sub in subs {
if let Some(fact) = ground_atom(&rule.head, &sub) {
let args = fact.args.clone();
let is_new = idb.insert_fact(&fact);
if is_new {
new_delta.entry(head_pred.clone()).or_default().insert(args);
}
}
}
}
}
delta = new_delta;
if idb.len() == prev_idb_size {
break;
}
}
let mut result = edb;
result.merge(&idb);
Ok(result)
}
}
impl Default for SemiNaiveEvaluator {
fn default() -> Self {
Self::new()
}
}