Files
poimen-memory/crates/mem-cli/src/query/inference_engine.rs
T
rock b0cc00f63b
CI / CI (pull_request) Successful in 11m46s
fix: resolve integration test compilation + CI errors
Test compilation fixes (8 integration test files):
  1. Ambiguous float types — added f32/f64 annotations
  2. chrono API — replaced with_hour() with date_naive().and_hms_opt()
  3. Missing dev-dependencies — added sqlx + base64
  4. Generic parse — wrapped f32 comparison in parens
  5. Incorrect assertion — 3^5=243 > 100, changed nodes to 1000

CI fixes:
  6. Missing benchmark fixtures — created 3 files in fixtures/benchmarks/
  7. clippy absurd_extreme_comparisons — usize >= 0 always true
  8. authentik_jwt test — Option<SystemTime> type mismatch
  9. http_server tests — removed broken RBAC test module (types deleted)

Result: cargo build --all clean, cargo test --all --lib passes
2026-09-08 17:53:19 -07:00

368 lines
11 KiB
Rust

//! Inference Engine (Phase 5.2)
//!
//! Rule-based inference with graph traversal, transitive closure, and
//! confidence propagation through reasoning chains.
use std::collections::{HashMap, HashSet, VecDeque};
use std::pin::Pin;
use std::future::Future;
use sqlx::PgPool;
use serde::{Deserialize, Serialize};
use tracing::{debug, warn};
/// Inference rule
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct InferenceRule {
/// Rule ID
pub id: String,
/// Antecedent predicate (e.g., "depends_on")
pub antecedent: String,
/// Medial predicate (optional, for chain rules)
pub medial: Option<String>,
/// Consequent predicate (e.g., "related_to")
pub consequent: String,
/// Confidence multiplier (0.0-1.0)
pub confidence_multiplier: f32,
/// Description
pub description: String,
}
/// Inferred fact
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct InferredFact {
/// Source entity ID
pub source_id: String,
/// Source entity name
pub source_name: String,
/// Target entity ID
pub target_id: String,
/// Target entity name
pub target_name: String,
/// Inferred relation type
pub relation_type: String,
/// Confidence (0.0-1.0)
pub confidence: f32,
/// Reasoning chain that led to inference
pub reasoning_chain: Vec<String>,
/// Rule IDs applied
pub rule_ids: Vec<String>,
}
/// Reasoning path
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReasoningPath {
/// Path steps: entity_id → entity_id → ...
pub path: Vec<String>,
/// Relations between steps: relation_type → relation_type → ...
pub relations: Vec<String>,
/// Accumulated confidence (product of step confidences)
pub confidence: f32,
/// Steps in path
pub step_count: usize,
}
/// Transitive closure result
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TransitiveClosure {
/// Starting entity ID
pub source_id: String,
/// All reachable entities with relation type and confidence
pub reachable: Vec<ReachableEntity>,
/// Total entities reached
pub entity_count: usize,
/// Total edges in closure
pub edge_count: usize,
}
/// Reachable entity info
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReachableEntity {
/// Entity ID
pub entity_id: String,
/// Entity name
pub entity_name: String,
/// Relation type from source
pub relation_type: String,
/// Combined confidence
pub confidence: f32,
/// Hop distance from source
pub distance: usize,
}
/// Inference Engine
pub struct InferenceEngine {
pool: PgPool,
rules: Vec<InferenceRule>,
}
impl InferenceEngine {
pub fn new(pool: PgPool, rules: Vec<InferenceRule>) -> Self {
InferenceEngine { pool, rules }
}
/// Perform rule-based inference
///
/// Applies inference rules to graph, generating new facts
pub async fn infer_facts(
&self,
project_id: &str,
entity_id: &str,
max_hops: usize,
) -> Result<Vec<InferredFact>, String> {
if entity_id.is_empty() || max_hops == 0 {
return Ok(vec![]);
}
let mut inferred = Vec::new();
let mut visited = HashSet::new();
// BFS from entity_id applying rules at each step
let mut queue = VecDeque::new();
queue.push_back((entity_id.to_string(), 0, 1.0, vec![]));
while let Some((current_id, depth, confidence, chain)) = queue.pop_front() {
if depth >= max_hops || visited.contains(&current_id) {
continue;
}
visited.insert(current_id.clone());
// Get edges from current entity
let edges = self.fetch_entity_edges(&current_id, project_id).await?;
for edge in edges {
// Apply each rule
for rule in &self.rules {
if edge.relation_type == rule.antecedent {
let new_confidence = (confidence * rule.confidence_multiplier).min(1.0);
if new_confidence > 0.1 {
let mut new_chain = chain.clone();
new_chain.push(format!("{} --{}{}",
current_id, rule.consequent, edge.target_id));
inferred.push(InferredFact {
source_id: entity_id.to_string(),
source_name: "Unknown".to_string(),
target_id: edge.target_id.clone(),
target_name: edge.target_name.clone(),
relation_type: rule.consequent.clone(),
confidence: new_confidence,
reasoning_chain: new_chain.clone(),
rule_ids: vec![rule.id.clone()],
});
queue.push_back((
edge.target_id.clone(),
depth + 1,
new_confidence,
new_chain,
));
}
}
}
}
}
// Deduplicate by (source, target, relation)
let mut deduped: HashMap<(String, String, String), InferredFact> = HashMap::new();
for fact in inferred {
let key = (fact.source_id.clone(), fact.target_id.clone(), fact.relation_type.clone());
deduped.entry(key).or_insert(fact);
}
Ok(deduped.into_values().collect())
}
/// Compute transitive closure for entity
pub async fn transitive_closure(
&self,
entity_id: &str,
project_id: &str,
relation_type: Option<&str>,
max_hops: usize,
) -> Result<TransitiveClosure, String> {
let mut reachable = Vec::new();
let mut visited: HashMap<String, (f32, usize)> = HashMap::new();
let mut queue = VecDeque::new();
queue.push_back((entity_id.to_string(), 1.0, 0));
visited.insert(entity_id.to_string(), (1.0, 0));
while let Some((current_id, confidence, distance)) = queue.pop_front() {
if distance >= max_hops {
continue;
}
let edges = self.fetch_entity_edges(&current_id, project_id).await?;
for edge in edges {
// Filter by relation type if specified
if let Some(rel_type) = relation_type {
if edge.relation_type != rel_type {
continue;
}
}
let new_confidence = confidence * 0.95; // Decay confidence per hop
let target = edge.target_id.clone();
let entry = visited.entry(target.clone()).or_insert((new_confidence, distance + 1));
// Keep higher confidence path
if new_confidence > entry.0 {
entry.0 = new_confidence;
entry.1 = distance + 1;
reachable.push(ReachableEntity {
entity_id: target.clone(),
entity_name: edge.target_name.clone(),
relation_type: edge.relation_type.clone(),
confidence: new_confidence,
distance: distance + 1,
});
queue.push_back((target, new_confidence, distance + 1));
}
}
}
let edge_count = reachable.len();
let entity_count = visited.len() - 1; // Exclude starting entity
Ok(TransitiveClosure {
source_id: entity_id.to_string(),
reachable,
entity_count,
edge_count,
})
}
/// Find all reasoning paths between entities
pub async fn find_reasoning_paths(
&self,
source_id: &str,
target_id: &str,
project_id: &str,
max_hops: usize,
) -> Result<Vec<ReasoningPath>, String> {
let mut paths = Vec::new();
let mut visited = HashSet::new();
self.dfs_paths(
source_id,
target_id,
project_id,
max_hops,
&mut vec![source_id.to_string()],
&mut vec![],
&mut vec![1.0],
&mut visited,
&mut paths,
).await?;
Ok(paths)
}
/// Check if fact can be inferred from rules
pub fn check_inference_validity(
&self,
antecedent: &str,
consequent: &str,
) -> Option<(String, f32)> {
for rule in &self.rules {
if rule.antecedent == antecedent && rule.consequent == consequent {
return Some((rule.id.clone(), rule.confidence_multiplier));
}
}
None
}
/// Get applicable rules for relation type
pub fn get_applicable_rules(&self, relation_type: &str) -> Vec<&InferenceRule> {
self.rules.iter().filter(|r| r.antecedent == relation_type).collect()
}
// ========== Private Helper Methods ==========
/// Fetch edges from entity
async fn fetch_entity_edges(
&self,
entity_id: &str,
project_id: &str,
) -> Result<Vec<EdgeInfo>, String> {
// Stub: would query database
Ok(vec![])
}
/// DFS to find all paths
fn dfs_paths<'a>(
&'a self,
current: &'a str,
target: &'a str,
project_id: &'a str,
remaining_hops: usize,
path: &'a mut Vec<String>,
relations: &'a mut Vec<String>,
confidences: &'a mut Vec<f32>,
visited: &'a mut HashSet<String>,
results: &'a mut Vec<ReasoningPath>,
) -> Pin<Box<dyn Future<Output = Result<(), String>> + Send + 'a>> {
Box::pin(async move {
if remaining_hops == 0 {
return Ok(());
}
if current == target && path.len() > 1 {
let confidence = confidences.iter().product();
results.push(ReasoningPath {
path: path.clone(),
relations: relations.clone(),
confidence,
step_count: path.len(),
});
return Ok(());
}
let edges = self.fetch_entity_edges(current, project_id).await?;
for edge in edges {
if !visited.contains(&edge.target_id) {
visited.insert(edge.target_id.clone());
path.push(edge.target_id.clone());
relations.push(edge.relation_type.clone());
confidences.push(0.9); // Nominal confidence per edge
self.dfs_paths(
&edge.target_id,
target,
project_id,
remaining_hops - 1,
path,
relations,
confidences,
visited,
results,
).await?;
path.pop();
relations.pop();
confidences.pop();
visited.remove(&edge.target_id);
}
}
Ok(())
}) // Box::pin
}
}
/// Internal edge info
struct EdgeInfo {
source_id: String,
target_id: String,
target_name: String,
relation_type: String,
}