use std::sync::Arc;
use exocortex_kernel::{Memory, MemoryId, Provenance, Relationship, RelationshipId};
use exocortex_storage::{Storage, StorageError};
use tokio::sync::mpsc;
use tracing::{instrument, warn};
use crate::rules::{self, Edge, EntityFact, TagFact};
const MAX_REASONING_NODES: usize = 512;
const MAX_REASONING_EDGES: usize = 4096;
const MAX_DERIVED_RELATIONSHIPS: usize = 10_000;
pub enum ReasoningWork {
KHopOver {
seed: MemoryId,
k: u8,
},
SessionWrapup {
memories: Vec<MemoryId>,
},
#[doc(hidden)]
DurableSessionWrapup {
memories: Vec<MemoryId>,
operation_key: smol_str::SmolStr,
completion: tokio::sync::oneshot::Sender<Result<(), StorageError>>,
},
#[cfg(debug_assertions)]
#[doc(hidden)]
StopWorker,
}
pub struct ReasoningEngine<S: Storage> {
storage: Arc<S>,
tx_work: mpsc::Sender<ReasoningWork>,
rx_work: tokio::sync::Mutex<mpsc::Receiver<ReasoningWork>>,
k_hop: u8,
}
impl<S: Storage + 'static> ReasoningEngine<S> {
pub fn new(storage: Arc<S>, queue_depth: usize, k_hop: u8) -> Self {
rules::prime(&storage_ontology(&storage));
let (tx, rx) = mpsc::channel(queue_depth);
Self {
storage,
tx_work: tx,
rx_work: tokio::sync::Mutex::new(rx),
k_hop,
}
}
pub async fn enqueue(&self, w: ReasoningWork) {
if self.tx_work.try_send(w).is_err() {
metrics::counter!("exocortex_reasoning_dropped_total").increment(1);
warn!("reasoning queue full; dropping work");
}
}
pub async fn run(self: Arc<Self>) {
loop {
let Some(w) = self.rx_work.lock().await.recv().await else {
return;
};
match w {
ReasoningWork::KHopOver { seed, k } => self.k_hop_reason(seed, k).await,
ReasoningWork::SessionWrapup { memories } => {
if let Err(error) = self.process_session_wrapup(&memories, None).await {
warn!(?error, "session reasoning failed");
}
}
ReasoningWork::DurableSessionWrapup {
memories,
operation_key,
completion,
} => {
let _ = completion.send(
self.process_session_wrapup(&memories, Some(operation_key.as_str()))
.await,
);
}
#[cfg(debug_assertions)]
ReasoningWork::StopWorker => return,
}
}
}
pub async fn process_durable_session_wrapup(
&self,
operation_key: smol_str::SmolStr,
memories: Vec<MemoryId>,
) -> Result<(), StorageError> {
let (completion, completed) = tokio::sync::oneshot::channel();
self.tx_work
.send(ReasoningWork::DurableSessionWrapup {
memories,
operation_key,
completion,
})
.await
.map_err(|_| StorageError::Backend("reasoning worker is unavailable".into()))?;
completed.await.map_err(|_| {
StorageError::Backend("reasoning worker stopped before completion".into())
})?
}
#[cfg(debug_assertions)]
#[doc(hidden)]
pub async fn stop_worker_for_testing(&self) {
let _ = self.tx_work.send(ReasoningWork::StopWorker).await;
}
#[instrument(skip(self))]
pub async fn k_hop_reason(&self, seed: MemoryId, k: u8) {
if let Err(error) = self.try_k_hop_reason(&[seed], k, None).await {
warn!(?error, "bounded reasoning pass failed");
}
}
async fn try_k_hop_reason(
&self,
seeds: &[MemoryId],
k: u8,
operation_key: Option<&str>,
) -> Result<(), StorageError> {
let k = k.clamp(1, self.k_hop.max(1));
let mut edges: Vec<Edge> = Vec::new();
let entities: Vec<EntityFact>;
let tags: Vec<TagFact>;
let mut seed_ids: Vec<_> = seeds.to_vec();
seed_ids.sort();
seed_ids.dedup();
if seed_ids.len() > MAX_REASONING_NODES {
return Err(StorageError::Backend(format!(
"reasoning seed cohort exceeds {MAX_REASONING_NODES} memories"
)));
}
let mut neighborhood: std::collections::HashSet<MemoryId> =
seed_ids.iter().copied().collect();
let mut seen_edges: std::collections::HashSet<RelationshipId> =
std::collections::HashSet::new();
let mut relationship_rows = Vec::new();
let mut frontier = seed_ids;
for _hop in 0..k {
if frontier.is_empty() {
break;
}
let mut next = Vec::new();
let rows = match self
.storage
.relationships_touching(&frontier, MAX_REASONING_EDGES as u32)
.await
{
Ok(rows) => rows,
Err(error) => return Err(error),
};
for row in rows {
if !seen_edges.insert(row.id) {
continue;
}
for other in [row.from, row.to] {
if neighborhood.len() < MAX_REASONING_NODES && neighborhood.insert(other) {
next.push(other);
}
}
relationship_rows.push(row);
}
if next.is_empty() || relationship_rows.len() >= MAX_REASONING_EDGES {
break;
}
next.sort();
next.dedup();
frontier = next;
}
for relationship in &relationship_rows {
if neighborhood.contains(&relationship.from) && neighborhood.contains(&relationship.to)
{
edges.push(Edge(relationship.from, relationship.to, relationship.kind));
}
}
const MAX_POSTING_LIST: usize = 256;
let mut neighborhood_ids: Vec<_> = neighborhood.iter().copied().collect();
neighborhood_ids.sort();
let mut memory_rows = match self.storage.get_memories(&neighborhood_ids).await {
Ok(rows) => rows,
Err(error) => return Err(error),
};
let attribute_tags: std::collections::HashSet<_> = memory_rows
.iter()
.flat_map(|memory| memory.tags.iter().cloned())
.collect();
let attribute_entities: std::collections::HashSet<_> = memory_rows
.iter()
.flat_map(|memory| memory.context.entities.iter().copied())
.collect();
if !attribute_tags.is_empty() || !attribute_entities.is_empty() {
let mut expansion = match self
.storage
.memories_sharing_attributes(
&attribute_tags.into_iter().collect::<Vec<_>>(),
&attribute_entities.into_iter().collect::<Vec<_>>(),
4096,
)
.await
{
Ok(rows) => rows,
Err(error) => return Err(error),
};
let mut seen: std::collections::HashSet<_> =
memory_rows.iter().map(|memory| memory.id).collect();
expansion.retain(|memory| seen.insert(memory.id));
memory_rows.extend(expansion);
}
let mut memories: Vec<rules::MemoryFact> = Vec::new();
let mut raw_tags: Vec<(MemoryId, u32)> = Vec::new();
let mut raw_entities: Vec<(MemoryId, exocortex_kernel::EntityId)> = Vec::new();
for memory in &memory_rows {
memories.push(rules::MemoryFact(memory.id, memory.memory_type));
for tag in &memory.tags {
raw_tags.push((memory.id, fxhash_tag(tag.as_str())));
}
for entity in &memory.context.entities {
raw_entities.push((memory.id, *entity));
}
}
{
use std::collections::HashMap;
let mut tag_counts: HashMap<u32, usize> = HashMap::new();
for (_, t) in &raw_tags {
*tag_counts.entry(*t).or_insert(0) += 1;
}
let dropped: usize = tag_counts
.values()
.filter(|c| **c > MAX_POSTING_LIST)
.count();
if dropped > 0 {
metrics::counter!("exocortex_reasoning_high_frequency_attributes_dropped_total")
.increment(dropped as u64);
}
tags = raw_tags
.into_iter()
.filter(|(_, t)| tag_counts.get(t).is_some_and(|c| *c <= MAX_POSTING_LIST))
.map(|(m, t)| TagFact(m, t))
.collect();
let mut ent_counts: HashMap<exocortex_kernel::EntityId, usize> = HashMap::new();
for (_, e) in &raw_entities {
*ent_counts.entry(*e).or_insert(0) += 1;
}
let dropped: usize = ent_counts
.values()
.filter(|c| **c > MAX_POSTING_LIST)
.count();
if dropped > 0 {
metrics::counter!("exocortex_reasoning_high_frequency_attributes_dropped_total")
.increment(dropped as u64);
}
entities = raw_entities
.into_iter()
.filter(|(_, e)| ent_counts.get(e).is_some_and(|c| *c <= MAX_POSTING_LIST))
.map(|(m, e)| EntityFact(m, e))
.collect();
}
let mut derived = rules::evaluate(edges, entities, tags, memories);
let before = derived.co_occurrence_affinity.len() + derived.similar_tags_affinity.len();
if before > MAX_DERIVED_RELATIONSHIPS {
let scale = MAX_DERIVED_RELATIONSHIPS as f64 / before as f64;
let keep = |v: &mut Vec<(MemoryId, MemoryId)>| {
let n = ((v.len() as f64) * scale).ceil() as usize;
v.truncate(n.min(v.len()));
};
keep(&mut derived.co_occurrence_affinity);
keep(&mut derived.similar_tags_affinity);
metrics::counter!("exocortex_reasoning_derived_pairs_capped_total")
.increment((before - MAX_DERIVED_RELATIONSHIPS) as u64);
}
self.write_back(derived, &memory_rows, &relationship_rows, operation_key)
.await
}
async fn process_session_wrapup(
&self,
ms: &[MemoryId],
operation_key: Option<&str>,
) -> Result<(), StorageError> {
let cohort_key = operation_key.map(|key| durable_cohort_key(key, ms));
self.try_k_hop_reason(ms, 3, cohort_key.as_deref()).await
}
async fn write_back(
&self,
mut derived: rules::Derived,
memory_rows: &[Memory],
relationship_rows: &[Relationship],
operation_key: Option<&str>,
) -> Result<(), StorageError> {
let ontology = self.storage_ontology();
let mut new_rels: Vec<Relationship> = Vec::new();
let now = chrono::Utc::now();
let mut memory_visibility = std::collections::HashMap::new();
for memory in memory_rows {
memory_visibility.insert(memory.id, memory.visibility);
}
let mut adj: std::collections::HashMap<
MemoryId,
Vec<(MemoryId, exocortex_kernel::RelKindId, RelationshipId)>,
> = std::collections::HashMap::new();
let mut evidence_visibility = std::collections::HashMap::new();
for relationship in relationship_rows {
evidence_visibility.insert(relationship.id, relationship.visibility);
adj.entry(relationship.from).or_default().push((
relationship.to,
relationship.kind,
relationship.id,
));
}
let support = |from: MemoryId,
to: MemoryId,
k1: exocortex_kernel::RelKindId,
k2: exocortex_kernel::RelKindId|
-> Vec<RelationshipId> {
let empty = Vec::new();
let outs = adj.get(&from).unwrap_or(&empty);
for (mid, ek1, id1) in outs {
if *ek1 != k1 {
continue;
}
if let Some(mids) = adj.get(mid) {
for (t, ek2, id2) in mids {
if *t == to && *ek2 == k2 {
return vec![*id1, *id2];
}
}
}
}
vec![]
};
let kind_of = |name: &str| ontology.kind_id(name).expect("kind");
let mut push = |from: MemoryId,
to: MemoryId,
rule_id: &str,
strength: f32,
shared_count: u32,
evidence: Vec<RelationshipId>| {
let Some(from_visibility) = memory_visibility.get(&from).copied() else {
return;
};
let Some(to_visibility) = memory_visibility.get(&to).copied() else {
return;
};
let visibility = exocortex_kernel::narrowest_visibility(
[from_visibility, to_visibility].into_iter().chain(
evidence
.iter()
.filter_map(|id| evidence_visibility.get(id).copied()),
),
)
.expect("two endpoint visibilities");
let kind = derived_kind(&ontology, rule_id);
let id = RelationshipId::derive(from, kind, to, Some(rule_id));
new_rels.push(Relationship {
id,
kind,
from,
to,
visibility,
provenance: Provenance::Derived {
rule_id: rule_id.into(),
evidence,
},
properties: exocortex_kernel::RelationshipProperties {
strength,
confidence: derived_confidence(rule_id, shared_count),
context: None,
evidence_count: 1,
success_rate: None,
validation_count: 0,
counter_evidence_count: 0,
last_validated: now,
},
description: None,
bidirectional: false,
valid_from: now,
valid_until: None,
recorded_at: now,
invalidated_by: None,
lsn: exocortex_kernel::LSN::new_local(0),
});
};
let dep = kind_of("DependsOn");
let req = kind_of("Requires");
let builds = kind_of("BuildsOn");
let blocks = kind_of("Blocks");
let contradicts = kind_of("Contradicts");
let confirms = kind_of("Confirms");
for (a, c) in derived.transitive_depends_on {
let ev = support(a, c, dep, dep);
push(a, c, "R4", 0.5, 0, ev);
}
for (a, c) in derived.transitive_requires {
let ev = support(a, c, req, req);
push(a, c, "R5", 0.5, 0, ev);
}
let r7_counts = rules::pair_counts(std::mem::take(&mut derived.co_occurrence_affinity));
for (a, b, shared) in r7_counts {
push(a, b, "R7", 0.3, shared, vec![]);
}
for (a, b) in derived.problem_solution_bridge {
push(a, b, "R8", 0.3, 0, vec![]);
}
let r9_counts = rules::pair_counts(std::mem::take(&mut derived.similar_tags_affinity));
for (a, b, shared) in r9_counts {
push(a, b, "R9", 0.3, shared, vec![]);
}
for (a, b) in derived.implied_solves {
push(a, b, "D1", 0.8, 0, vec![]);
}
for (a, c) in derived.transitive_builds_on {
let ev = support(a, c, builds, builds);
push(a, c, "D2", 0.5, 0, ev);
}
for (a, c) in derived.indirect_blocker {
let ev = support(a, c, blocks, req);
push(a, c, "D3", 0.5, 0, ev);
}
for (a, c) in derived.contradiction_propagates {
let ev = support(a, c, contradicts, confirms);
push(a, c, "D4", 0.5, 0, ev);
}
for (a, b) in derived.shared_target {
push(a, b, "D5", 0.4, 0, vec![]);
}
{
let mut by_session: std::collections::HashMap<MemoryId, Vec<MemoryId>> =
std::collections::HashMap::new();
for (m, sess) in &derived.session_cohort {
by_session.entry(*sess).or_default().push(*m);
}
let mut members: Vec<MemoryId> = by_session.keys().copied().collect();
members.sort();
for sess in members {
let group = by_session
.get_mut(&sess)
.expect("session key collected above");
group.sort();
group.dedup();
for (i, m1) in group.iter().enumerate() {
for m2 in group.iter().skip(i + 1) {
if m1 != m2 {
push(*m1, *m2, "D6", 0.6, 0, vec![]);
}
}
}
}
}
new_rels.sort_by_key(|relationship| relationship.id);
new_rels.dedup_by_key(|relationship| relationship.id);
if new_rels.len() > MAX_DERIVED_RELATIONSHIPS {
let dropped = new_rels.len() - MAX_DERIVED_RELATIONSHIPS;
new_rels.truncate(MAX_DERIVED_RELATIONSHIPS);
metrics::counter!("exocortex_reasoning_derived_pairs_capped_total")
.increment(dropped as u64);
}
if new_rels.is_empty() && operation_key.is_none() {
return Ok(());
}
let candidate_ids = new_rels
.iter()
.map(|relationship| relationship.id)
.collect::<Vec<_>>();
let existing = self
.storage
.get_relationships(&candidate_ids)
.await?
.into_iter()
.map(|relationship| relationship.id)
.collect::<std::collections::HashSet<_>>();
let fresh = new_rels
.into_iter()
.filter(|relationship| !existing.contains(&relationship.id))
.collect::<Vec<_>>();
if let Some(operation_key) = operation_key {
let committed = self
.storage
.upsert_batch_once(operation_key, &[], &fresh)
.await?;
if committed && !fresh.is_empty() {
metrics::counter!("exocortex_rules_executed_total", "engine" => "crepe")
.increment(fresh.len() as u64);
}
} else if !fresh.is_empty() {
self.storage.upsert_batch(&[], &fresh).await?;
metrics::counter!("exocortex_rules_executed_total", "engine" => "crepe")
.increment(fresh.len() as u64);
}
Ok(())
}
pub async fn inferred_type(&self, id: MemoryId) -> Option<u8> {
let mut edges = Vec::new();
let rels = self
.storage
.relationships_touching(&[id], 4096)
.await
.ok()?;
for r in rels {
if r.from == id || r.to == id {
edges.push(Edge(r.from, r.to, r.kind));
}
}
let derived = rules::evaluate(edges, vec![], vec![], vec![]);
derived
.type_from_solves
.iter()
.chain(derived.type_from_fixes.iter())
.chain(derived.type_from_causes.iter())
.find(|(m, _)| *m == id)
.map(|(_, t)| *t)
}
fn storage_ontology(&self) -> exocortex_kernel::Ontology {
storage_ontology(&self.storage)
}
}
fn durable_cohort_key(effect_key: &str, seeds: &[MemoryId]) -> String {
use std::fmt::Write as _;
let mut seeds = seeds.to_vec();
seeds.sort();
seeds.dedup();
let mut key = format!(
"reasoning:v2:{}:{effect_key}:{}:",
effect_key.len(),
seeds.len()
);
for seed in seeds {
for byte in seed.0 {
let _ = write!(key, "{byte:02x}");
}
}
key
}
fn storage_ontology<S: Storage>(_: &Arc<S>) -> exocortex_kernel::Ontology {
exocortex_kernel::Ontology::from_packs(vec![exocortex_pack_dev_v1::pack_def()])
.expect("linked pack assembles")
}
fn derived_kind(onto: &exocortex_kernel::Ontology, rule_id: &str) -> exocortex_kernel::RelKindId {
match rule_id {
"R4" => onto.kind_id("DependsOn").expect("kind"),
"R5" => onto.kind_id("Requires").expect("kind"),
"D1" => exocortex_kernel::kinds::SOLVES,
"D2" => onto.kind_id("BuildsOn").expect("kind"),
"D3" => onto.kind_id("Blocks").expect("kind"),
"D4" => onto.kind_id("Contradicts").expect("kind"),
"D6" => onto.kind_id("RelatedTo").expect("kind"),
_ => onto.kind_id("RelatedTo").expect("kind"),
}
}
fn derived_confidence(rule_id: &str, n: u32) -> f32 {
match rule_id {
"R4" | "R5" => 0.5, "R7" | "R9" => (n as f32 / 5.0).min(1.0),
"R8" | "D1" => 0.8,
_ => 0.5,
}
}
fn fxhash_tag(tag: &str) -> u32 {
let mut h: u32 = 0x811c_9dc5;
for b in tag.as_bytes() {
h ^= *b as u32;
h = h.wrapping_mul(0x0100_0193);
}
h
}
#[cfg(test)]
mod durable_key_tests {
use super::*;
#[test]
fn durable_cohort_identity_is_versioned_length_framed_sorted_raw_hex() {
let first = MemoryId([
0x01, 0x23, 0x45, 0x67, 0x89, 0xab, 0xcd, 0xef, 0, 0, 0, 0, 0, 0, 0, 0,
]);
let second = MemoryId([0xff; 16]);
assert_eq!(
durable_cohort_key("effect", &[second, first, first]),
"reasoning:v2:6:effect:2:0123456789abcdef0000000000000000ffffffffffffffffffffffffffffffff"
);
}
}