use super::*;
pub(crate) type EvalFn = fn(&str, &[LogicalTerm]) -> Result<bool, String>;
pub(crate) type BatchEvalFn = fn(&[ComputeRequest]) -> Vec<Result<bool, String>>;
pub(super) fn extract_num_value(
term: &LogicalTerm,
subs: &HashMap<String, GroundTerm>,
) -> Option<f64> {
match term {
LogicalTerm::Number(n) => Some(*n),
LogicalTerm::Variable(v) => {
let gt = subs.get(v.as_str())?;
gt.as_f64()
}
_ => None,
}
}
pub(super) fn try_numeric_comparison(
rel: &str,
args: &[LogicalTerm],
subs: &HashMap<String, GroundTerm>,
) -> Option<QueryResult> {
let a = extract_num_value(args.get(0)?, subs)?;
let b = extract_num_value(args.get(1)?, subs)?;
let holds = match rel {
"greater" => a > b,
"less" => a < b,
"num_equal" => a == b,
_ => return None,
};
if !a.is_finite() || !b.is_finite() {
return Some(QueryResult::Unknown(UnknownReason::NonFinite));
}
Some(if holds {
QueryResult::True
} else {
QueryResult::False
})
}
pub(super) fn try_arithmetic_evaluation(
rel: &str,
args: &[LogicalTerm],
subs: &HashMap<String, GroundTerm>,
) -> Option<bool> {
let x1 = extract_num_value(args.get(0)?, subs)?;
let x2 = extract_num_value(args.get(1)?, subs)?;
let x3 = extract_num_value(args.get(2)?, subs)?;
nibli_types::eval_arithmetic(rel, &[x1, x2, x3])
}
pub(super) fn ground_term_to_logical_term(gt: &GroundTerm) -> LogicalTerm {
match gt {
GroundTerm::Constant(c) => LogicalTerm::Constant(c.clone()),
GroundTerm::Number(bits) => LogicalTerm::Number(f64::from_bits(*bits)),
GroundTerm::Description(d) => LogicalTerm::Description(d.clone()),
GroundTerm::Unspecified => LogicalTerm::Unspecified,
GroundTerm::PatternVar(v) => LogicalTerm::Variable(v.clone()),
GroundTerm::SkolemFn(name, _) => LogicalTerm::Constant(name.clone()),
GroundTerm::DepPair(_, _) => LogicalTerm::Unspecified,
}
}
pub(super) fn witness_term_to_logical_term(gt: &GroundTerm) -> LogicalTerm {
match gt {
GroundTerm::SkolemFn(..) | GroundTerm::DepPair(..) => {
LogicalTerm::Constant(gt.to_display_string())
}
other => ground_term_to_logical_term(other),
}
}
pub(super) struct NumericGroupVerdict {
pub relation: String,
pub method: &'static str,
pub verdict: QueryResult,
}
pub(super) fn try_evaluate_numeric_group(
inner: &KnowledgeBaseInner,
buffer: &LogicBuffer,
exists_var: &str,
body_id: u32,
subs: &HashMap<String, GroundTerm>,
) -> Option<NumericGroupVerdict> {
let mut conjuncts: Vec<u32> = Vec::new();
let mut stack = vec![body_id];
while let Some(id) = stack.pop() {
match get_node(buffer, id).ok()? {
LogicNode::AndNode((l, r)) => {
stack.push(*l);
stack.push(*r);
}
LogicNode::Predicate(_) | LogicNode::ComputeNode(_) => conjuncts.push(id),
_ => return None,
}
}
let is_head_var =
|args: &[LogicalTerm]| matches!(args, [LogicalTerm::Variable(v)] if v == exists_var);
let mut head: Option<(&str, bool)> = None; for &id in &conjuncts {
match get_node(buffer, id).ok()? {
LogicNode::ComputeNode((rel, args)) if is_head_var(args) => {
if head.is_some() {
return None; }
head = Some((rel.as_str(), true));
}
LogicNode::Predicate((rel, args))
if is_head_var(args)
&& nibli_types::relations::is_numeric_comparison(rel.as_str()) =>
{
if head.is_some() {
return None;
}
head = Some((rel.as_str(), false));
}
_ => {}
}
}
let (rel, head_is_compute) = head?;
let role_prefix = format!("{rel}_x");
let mut roles: Vec<Option<&LogicalTerm>> = Vec::new();
for &id in &conjuncts {
match get_node(buffer, id).ok()? {
LogicNode::ComputeNode((r, args)) if is_head_var(args) && r.as_str() == rel => {}
LogicNode::Predicate((r, args)) if is_head_var(args) && r.as_str() == rel => {}
LogicNode::Predicate((r, args)) if r.starts_with(&role_prefix) => {
let n: usize = r[role_prefix.len()..].parse().ok()?;
if n == 0 {
return None;
}
match args.as_slice() {
[LogicalTerm::Variable(v), arg] if v == exists_var => {
if roles.len() < n {
roles.resize(n, None);
}
if roles[n - 1].is_some() {
return None; }
roles[n - 1] = Some(arg);
}
_ => return None,
}
}
_ => return None, }
}
if roles.is_empty() || roles.iter().any(|r| r.is_none()) {
return None;
}
let collected: Vec<LogicalTerm> = roles.into_iter().map(|r| r.unwrap().clone()).collect();
{
let operands: Vec<f64> = collected
.iter()
.filter_map(|t| extract_num_value(t, subs))
.collect();
let non_finite = match rel {
r if nibli_types::relations::is_builtin_arithmetic(r) => {
operands.len() == 3 && nibli_types::eval_arithmetic(rel, &operands).is_none()
}
r if nibli_types::relations::is_numeric_comparison(r) => {
operands.len() >= 2 && operands.iter().take(2).any(|n| !n.is_finite())
}
_ => false,
};
if non_finite {
return Some(NumericGroupVerdict {
relation: rel.to_string(),
method: "non_finite",
verdict: QueryResult::Unknown(UnknownReason::NonFinite),
});
}
}
if let Some(verdict) = try_numeric_comparison(rel, &collected, subs) {
return Some(NumericGroupVerdict {
relation: rel.to_string(),
method: if matches!(verdict, QueryResult::Unknown(_)) {
"non_finite"
} else {
"numeric"
},
verdict,
});
}
if let Some(holds) = try_arithmetic_evaluation(rel, &collected, subs) {
return Some(NumericGroupVerdict {
relation: rel.to_string(),
method: "arithmetic",
verdict: bool_verdict(holds),
});
}
if head_is_compute {
let resolved = resolve_args_for_dispatch(&collected, subs);
let dispatchable = resolved
.iter()
.all(|t| matches!(t, LogicalTerm::Number(_) | LogicalTerm::Unspecified));
if dispatchable {
return Some(match dispatch_to_backend(inner, rel, &resolved) {
Ok(holds) => NumericGroupVerdict {
relation: rel.to_string(),
method: "backend",
verdict: bool_verdict(holds),
},
Err(_) => NumericGroupVerdict {
relation: rel.to_string(),
method: "backend_unavailable",
verdict: QueryResult::Unknown(UnknownReason::BackendUnavailable),
},
});
}
}
None
}
fn bool_verdict(holds: bool) -> QueryResult {
if holds {
QueryResult::True
} else {
QueryResult::False
}
}
pub(super) fn resolve_args_for_dispatch(
args: &[LogicalTerm],
subs: &HashMap<String, GroundTerm>,
) -> Vec<LogicalTerm> {
args.iter()
.map(|a| match a {
LogicalTerm::Variable(v) => {
if let Some(gt) = subs.get(v.as_str()) {
ground_term_to_logical_term(gt)
} else {
a.clone()
}
}
_ => a.clone(),
})
.collect()
}
pub(super) fn dispatch_to_backend(
inner: &KnowledgeBaseInner,
rel: &str,
args: &[LogicalTerm],
) -> Result<bool, String> {
match inner.compute_eval {
Some(eval) => eval(rel, args),
None => Err("Compute backend not registered".to_string()),
}
}
pub struct ComputeRequest {
pub relation: String,
pub args: Vec<LogicalTerm>,
}
fn dispatch_batch_to_backend(
inner: &KnowledgeBaseInner,
requests: &[ComputeRequest],
) -> Vec<Result<bool, String>> {
match inner.compute_batch_eval {
Some(batch_eval) => batch_eval(requests),
None => requests
.iter()
.map(|_| Err("Compute backend not registered".to_string()))
.collect(),
}
}
pub(super) fn build_ground_fact_from_resolved(
rel: &str,
resolved_args: &[LogicalTerm],
) -> Option<StoredFact> {
for arg in resolved_args {
if matches!(arg, LogicalTerm::Variable(_)) {
return None;
}
}
let args: Vec<GroundTerm> = resolved_args
.iter()
.map(|arg| match arg {
LogicalTerm::Number(n) => GroundTerm::from_f64(*n),
LogicalTerm::Constant(c) => GroundTerm::Constant(c.clone()),
LogicalTerm::Description(d) => GroundTerm::Description(d.clone()),
LogicalTerm::Unspecified => GroundTerm::Unspecified,
LogicalTerm::Variable(v) => {
unreachable!("Variable '{}' in compute result — should be ground", v)
}
})
.collect();
Some(StoredFact::Bare(GroundFact::new(rel, args)))
}
pub(super) struct BatchComputeResult {
pub results: Vec<bool>,
pub deferred_facts: Vec<StoredFact>,
}
pub(super) fn batch_evaluate_compute_for_members(
inner: &KnowledgeBaseInner,
rel: &str,
args: &[LogicalTerm],
var: &str,
members: &[GroundTerm],
subs: &HashMap<String, GroundTerm>,
) -> Option<BatchComputeResult> {
let mut results = vec![false; members.len()];
let mut deferred_facts = Vec::new();
let mut pending: Vec<(usize, Vec<LogicalTerm>)> = Vec::new();
for (i, member) in members.iter().enumerate() {
let mut s = subs.clone();
s.insert(var.to_string(), member.clone());
if let Some(r) = try_arithmetic_evaluation(rel, args, &s) {
results[i] = r;
if r {
let resolved = resolve_args_for_dispatch(args, &s);
if let Some(fact) = build_ground_fact_from_resolved(rel, &resolved) {
deferred_facts.push(fact);
}
}
} else {
let resolved = resolve_args_for_dispatch(args, &s);
pending.push((i, resolved));
}
}
if pending.is_empty() {
return Some(BatchComputeResult {
results,
deferred_facts,
});
}
let requests: Vec<ComputeRequest> = pending
.iter()
.map(|(_, resolved)| ComputeRequest {
relation: rel.to_string(),
args: resolved.clone(),
})
.collect();
let batch_results = dispatch_batch_to_backend(inner, &requests);
for (batch_idx, result) in batch_results.into_iter().enumerate() {
let member_idx = pending[batch_idx].0;
match result {
Ok(r) => {
results[member_idx] = r;
if r {
if let Some(fact) = build_ground_fact_from_resolved(rel, &pending[batch_idx].1)
{
deferred_facts.push(fact);
}
}
}
Err(_) => return None,
}
}
Some(BatchComputeResult {
results,
deferred_facts,
})
}