Files
clawhdf5/crates/clawhdf5-agent/src/knowledge.rs
T
osobhandClaude Opus 5.5 1b3bbb054a perf(agent): cache the knowledge graph's adjacency index
bfs_neighbors and spreading_activation built an adjacency index over the
whole graph on every call (1efd82c), so a 2-hop BFS over 1K entities
paid to index every entity and relation first: 155 us, 6.5x the 24 us
the README quoted. Found by the dated benchmark re-run.

The index is now cached on KnowledgeCache and checked against a
fingerprint of the graph on each use — one pass over entity ids and
relation endpoints, no allocation — so any change, including direct
edits of the public entities/relations Vecs (schema.rs's load path
pushes to them), still triggers a rebuild. A test edits the graph
directly in every way (push, in-place rewire, pop + push at equal
length) between traversals.

tank, 2026-09-24: BFS 1K entities 155.1 -> 23.1 us, 100 entities
17.5 -> 5.23 us, spreading activation 100 22.8 -> 10.1 us.

Co-Authored-By: Claude Opus 5.5 (1M context) <[email protected]>
2026-09-24 23:45:15 -05:00

1376 lines
47 KiB
Rust

//! Knowledge graph data structures and cache.
use std::cmp::Reverse;
use std::collections::{HashMap, HashSet, VecDeque};
// ---------------------------------------------------------------------------
// RelationType
// ---------------------------------------------------------------------------
/// Typed classification for a knowledge graph relation.
#[derive(Debug, Clone, PartialEq)]
pub enum RelationType {
Temporal,
Causal,
Associative,
Hierarchical,
Custom(String),
}
impl RelationType {
/// Convert a string label to a `RelationType`.
pub fn from_label(s: &str) -> Self {
match s {
"temporal" => RelationType::Temporal,
"causal" => RelationType::Causal,
"associative" => RelationType::Associative,
"hierarchical" => RelationType::Hierarchical,
other => RelationType::Custom(other.to_string()),
}
}
/// Convert back to a canonical string label.
pub fn as_str(&self) -> &str {
match self {
RelationType::Temporal => "temporal",
RelationType::Causal => "causal",
RelationType::Associative => "associative",
RelationType::Hierarchical => "hierarchical",
RelationType::Custom(s) => s.as_str(),
}
}
}
// ---------------------------------------------------------------------------
// Entity
// ---------------------------------------------------------------------------
/// A knowledge graph entity.
#[derive(Debug, Clone)]
pub struct Entity {
pub id: u64,
pub name: String,
/// Lowercased `name`, cached at construction time to avoid re-allocating
/// and re-lowercasing on every entity-resolution scan.
pub name_lower: String,
pub entity_type: String,
/// Index into the memory embeddings array, or -1 if none.
pub embedding_idx: i64,
/// Arbitrary key-value properties attached to this entity.
pub properties: HashMap<String, String>,
/// Optional dense embedding vector stored directly on the entity.
pub embedding: Option<Vec<f32>>,
/// Unix timestamp (microseconds) when this entity was first created.
pub created_at: f64,
/// Unix timestamp (microseconds) of the most recent update.
pub updated_at: f64,
}
impl Default for Entity {
fn default() -> Self {
let now = current_ts_us();
Self {
id: 0,
name: String::new(),
name_lower: String::new(),
entity_type: String::new(),
embedding_idx: -1,
properties: HashMap::new(),
embedding: None,
created_at: now,
updated_at: now,
}
}
}
// ---------------------------------------------------------------------------
// Relation
// ---------------------------------------------------------------------------
/// A knowledge graph relation between two entities.
#[derive(Debug, Clone)]
pub struct Relation {
pub src: u64,
pub tgt: u64,
pub relation: String,
pub weight: f32,
pub ts: f64,
/// Arbitrary key-value metadata attached to this relation.
pub metadata: HashMap<String, String>,
}
impl Default for Relation {
fn default() -> Self {
Self {
src: 0,
tgt: 0,
relation: String::new(),
weight: 1.0,
ts: current_ts_us(),
metadata: HashMap::new(),
}
}
}
// ---------------------------------------------------------------------------
// Helpers
// ---------------------------------------------------------------------------
/// Return the current wall-clock time in microseconds since Unix epoch.
/// Falls back to 0.0 if the system clock is unavailable.
fn current_ts_us() -> f64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs_f64()
* 1_000_000.0
}
/// Compute a simple edit-distance between two strings (Levenshtein).
/// Returns the number of single-character edits needed to transform `a` into `b`.
fn levenshtein(a: &str, b: &str) -> usize {
let a: Vec<char> = a.chars().collect();
let b: Vec<char> = b.chars().collect();
let na = a.len();
let nb = b.len();
if na == 0 {
return nb;
}
if nb == 0 {
return na;
}
let mut prev: Vec<usize> = (0..=nb).collect();
let mut curr = vec![0usize; nb + 1];
for i in 1..=na {
curr[0] = i;
for j in 1..=nb {
let cost = if a[i - 1] == b[j - 1] { 0 } else { 1 };
curr[j] = (curr[j - 1] + 1).min(prev[j] + 1).min(prev[j - 1] + cost);
}
std::mem::swap(&mut prev, &mut curr);
}
prev[nb]
}
// ---------------------------------------------------------------------------
// AdjacencyIndex
// ---------------------------------------------------------------------------
/// Adjacency index over a snapshot of `entities`/`relations`: an entity-id ->
/// entities-slice-index map, and an entity-id -> relation-indices map (edges
/// touching that entity as either source or target).
///
/// Cached on `KnowledgeCache` and checked against a fingerprint of the graph
/// on every use ([`graph_fingerprint`]). entities/relations are plain `pub`
/// `Vec`s that get changed directly (e.g. `schema.rs`'s load path bypasses
/// `add_entity`/`add_relation`), so the cache cannot rely on being told about
/// changes; the fingerprint notices any of them. Rebuilding it on every
/// traversal instead made a 2-hop BFS over 1K entities 6.5x slower than the
/// scan it replaced (24 -> 155 µs; `BENCHMARKS.md`, "Knowledge Graph").
struct AdjacencyIndex {
entity_index: HashMap<u64, usize>,
by_entity: HashMap<u64, Vec<usize>>,
}
impl AdjacencyIndex {
fn build(entities: &[Entity], relations: &[Relation]) -> Self {
let mut entity_index = HashMap::with_capacity(entities.len());
for (i, e) in entities.iter().enumerate() {
entity_index.insert(e.id, i);
}
let mut by_entity: HashMap<u64, Vec<usize>> = HashMap::new();
for (i, r) in relations.iter().enumerate() {
by_entity.entry(r.src).or_default().push(i);
if r.tgt != r.src {
by_entity.entry(r.tgt).or_default().push(i);
}
}
Self {
entity_index,
by_entity,
}
}
/// Indices into `relations` of every edge touching `entity_id`.
fn relations_touching(&self, entity_id: u64) -> &[usize] {
self.by_entity
.get(&entity_id)
.map(|v| v.as_slice())
.unwrap_or(&[])
}
}
/// A hash of everything [`AdjacencyIndex`] depends on — each entity's id and
/// position, each relation's endpoints and position. One linear pass, no
/// allocation: far cheaper than building the index, which hashes the same
/// values into two maps.
fn graph_fingerprint(entities: &[Entity], relations: &[Relation]) -> u64 {
// splitmix64-style mixing; order matters, so positions are covered.
fn mix(h: u64, v: u64) -> u64 {
let mut z = (h ^ v).wrapping_add(0x9E37_79B9_7F4A_7C15);
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
let mut h = mix(entities.len() as u64, relations.len() as u64);
for e in entities {
h = mix(h, e.id);
}
for r in relations {
h = mix(mix(h, r.src), r.tgt);
}
h
}
/// The cached [`AdjacencyIndex`] and the fingerprint it was built for.
/// Cloning a `KnowledgeCache` starts the clone with an empty cache.
#[derive(Default)]
struct AdjacencyCache(std::sync::Mutex<Option<(u64, std::sync::Arc<AdjacencyIndex>)>>);
impl Clone for AdjacencyCache {
fn clone(&self) -> Self {
Self::default()
}
}
impl std::fmt::Debug for AdjacencyCache {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("AdjacencyCache")
}
}
// ---------------------------------------------------------------------------
// KnowledgeCache
// ---------------------------------------------------------------------------
/// In-memory cache for the /knowledge_graph group.
#[derive(Debug, Clone)]
pub struct KnowledgeCache {
pub entities: Vec<Entity>,
pub relations: Vec<Relation>,
pub alias_strings: Vec<String>,
pub alias_entity_ids: Vec<i64>,
next_entity_id: u64,
adjacency: AdjacencyCache,
}
impl KnowledgeCache {
pub fn new() -> Self {
Self {
entities: Vec::new(),
relations: Vec::new(),
alias_strings: Vec::new(),
alias_entity_ids: Vec::new(),
next_entity_id: 0,
adjacency: AdjacencyCache::default(),
}
}
pub fn new_with_next_id(next_id: u64) -> Self {
Self {
entities: Vec::new(),
relations: Vec::new(),
alias_strings: Vec::new(),
alias_entity_ids: Vec::new(),
next_entity_id: next_id,
adjacency: AdjacencyCache::default(),
}
}
/// The adjacency index for the graph as it is now: the cached one if the
/// graph's fingerprint still matches, otherwise rebuilt and cached.
fn adjacency_index(&self) -> std::sync::Arc<AdjacencyIndex> {
let fp = graph_fingerprint(&self.entities, &self.relations);
let mut slot = self
.adjacency
.0
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some((cached_fp, idx)) = slot.as_ref()
&& *cached_fp == fp
{
return idx.clone();
}
let idx = std::sync::Arc::new(AdjacencyIndex::build(&self.entities, &self.relations));
*slot = Some((fp, idx.clone()));
idx
}
// -----------------------------------------------------------------------
// Entity management
// -----------------------------------------------------------------------
/// Add an entity, returns its assigned ID.
pub fn add_entity(&mut self, name: &str, entity_type: &str, embedding_idx: i64) -> u64 {
let id = self.next_entity_id;
self.next_entity_id += 1;
let now = current_ts_us();
self.entities.push(Entity {
id,
name: name.to_owned(),
name_lower: name.to_lowercase(),
entity_type: entity_type.to_owned(),
embedding_idx,
properties: HashMap::new(),
embedding: None,
created_at: now,
updated_at: now,
});
id
}
/// Find an entity by ID.
pub fn get_entity(&self, id: u64) -> Option<&Entity> {
self.entities.iter().find(|e| e.id == id)
}
/// Find a mutable entity by ID.
pub fn get_entity_mut(&mut self, id: u64) -> Option<&mut Entity> {
self.entities.iter_mut().find(|e| e.id == id)
}
/// Get entity name by ID.
pub fn get_entity_name(&self, entity_id: i64) -> Option<&str> {
self.entities
.iter()
.find(|e| e.id == entity_id as u64)
.map(|e| e.name.as_str())
}
// -----------------------------------------------------------------------
// Relation management
// -----------------------------------------------------------------------
/// Add a relation between two entities.
pub fn add_relation(&mut self, src: u64, tgt: u64, relation: &str, weight: f32) {
let ts = current_ts_us();
self.relations.push(Relation {
src,
tgt,
relation: relation.to_owned(),
weight,
ts,
metadata: HashMap::new(),
});
}
/// Find all relations where the given entity is the source.
pub fn get_relations_from(&self, src_id: u64) -> Vec<&Relation> {
self.relations.iter().filter(|r| r.src == src_id).collect()
}
/// Find all relations where the given entity is the target.
pub fn get_relations_to(&self, tgt_id: u64) -> Vec<&Relation> {
self.relations.iter().filter(|r| r.tgt == tgt_id).collect()
}
// -----------------------------------------------------------------------
// Alias management
// -----------------------------------------------------------------------
/// Register an alias for an entity. Case-insensitive storage.
pub fn add_alias(&mut self, alias: &str, entity_id: i64) {
self.alias_strings.push(alias.to_lowercase());
self.alias_entity_ids.push(entity_id);
}
/// Get all aliases for a given entity.
pub fn get_aliases(&self, entity_id: i64) -> Vec<&str> {
self.alias_strings
.iter()
.zip(&self.alias_entity_ids)
.filter(|&(_, id)| *id == entity_id)
.map(|(s, _)| s.as_str())
.collect()
}
/// Resolve aliases in free text — greedy longest-match replacement.
pub fn resolve_aliases(&self, query: &str) -> String {
let lower = query.to_lowercase();
let mut pairs: Vec<(&str, String)> = self
.alias_strings
.iter()
.zip(&self.alias_entity_ids)
.filter_map(|(alias, &eid)| {
self.get_entity_name(eid)
.map(|name| (alias.as_str(), name.to_lowercase()))
})
.collect();
pairs.sort_by_key(|b| Reverse(b.0.len()));
let mut result = lower;
for (alias, name) in &pairs {
result = result.replace(alias, name);
}
result
}
// -----------------------------------------------------------------------
// Entity resolution
// -----------------------------------------------------------------------
/// Find an existing entity whose name is within `max_distance` Levenshtein
/// edits of `name` (case-insensitive), or create a new entity if none is
/// found. Returns `(entity_id, was_created)`.
pub fn resolve_or_create(
&mut self,
name: &str,
entity_type: &str,
embedding_idx: i64,
max_distance: usize,
) -> (u64, bool) {
let lower_name = name.to_lowercase();
// Search for the closest existing entity, short-circuiting on an
// exact match since no closer candidate can exist.
let mut best: Option<(u64, usize)> = None;
for e in &self.entities {
let dist = levenshtein(&lower_name, &e.name_lower);
if dist > max_distance {
continue;
}
if dist == 0 {
best = Some((e.id, dist));
break;
}
if best.is_none_or(|(_, best_dist)| dist < best_dist) {
best = Some((e.id, dist));
}
}
if let Some((id, _)) = best {
return (id, false);
}
let id = self.add_entity(name, entity_type, embedding_idx);
(id, true)
}
// -----------------------------------------------------------------------
// Graph traversal: BFS neighbors
// -----------------------------------------------------------------------
/// Return all entities reachable from `entity_id` within `max_depth` hops,
/// together with their discovered depth. The seed entity itself is NOT
/// included. Traversal follows both outgoing and incoming relation edges.
pub fn bfs_neighbors(&self, entity_id: u64, max_depth: usize) -> Vec<(Entity, usize)> {
let idx = self.adjacency_index();
let mut visited: HashSet<u64> = HashSet::new();
let mut queue: VecDeque<(u64, usize)> = VecDeque::new();
let mut results: Vec<(Entity, usize)> = Vec::new();
visited.insert(entity_id);
queue.push_back((entity_id, 0));
while let Some((current_id, depth)) = queue.pop_front() {
if depth >= max_depth {
continue;
}
// Collect neighbour IDs from outgoing and incoming edges touching
// this node only, instead of scanning every relation in the graph.
let neighbours: Vec<u64> = idx
.relations_touching(current_id)
.iter()
.filter_map(|&i| {
let r = &self.relations[i];
if r.src == current_id {
Some(r.tgt)
} else if r.tgt == current_id {
Some(r.src)
} else {
None
}
})
.collect();
for neighbour_id in neighbours {
if visited.insert(neighbour_id)
&& let Some(&entity_idx) = idx.entity_index.get(&neighbour_id)
{
results.push((self.entities[entity_idx].clone(), depth + 1));
queue.push_back((neighbour_id, depth + 1));
}
}
}
results
}
// -----------------------------------------------------------------------
// Graph traversal: subgraph extraction
// -----------------------------------------------------------------------
/// Extract the subgraph reachable from any of `seed_ids` within `max_depth`
/// hops. Returns `(entities, relations)` where `relations` contains only
/// those edges whose both endpoints are in the entity set.
pub fn get_subgraph(&self, seed_ids: &[u64], max_depth: usize) -> (Vec<Entity>, Vec<Relation>) {
let mut entity_ids: HashSet<u64> = HashSet::new();
// Seed all starting nodes.
for &seed in seed_ids {
entity_ids.insert(seed);
}
// BFS from each seed, collecting entity IDs.
for &seed in seed_ids {
for (entity, _depth) in self.bfs_neighbors(seed, max_depth) {
entity_ids.insert(entity.id);
}
}
let entities: Vec<Entity> = self
.entities
.iter()
.filter(|e| entity_ids.contains(&e.id))
.cloned()
.collect();
let relations: Vec<Relation> = self
.relations
.iter()
.filter(|r| entity_ids.contains(&r.src) && entity_ids.contains(&r.tgt))
.cloned()
.collect();
(entities, relations)
}
// -----------------------------------------------------------------------
// Spreading activation
// -----------------------------------------------------------------------
/// Compute spreading activation scores starting from `seed_ids`.
///
/// The algorithm works as follows:
/// 1. Each seed receives an initial activation of 1.0.
/// 2. At each step, every activated node spreads activation to its
/// neighbours proportional to `relation.weight * decay_factor`.
/// 3. A node's total activation is the sum of all received signals.
/// 4. Propagation continues for up to `max_steps` rounds or until no node
/// has activation above `min_activation`.
///
/// Returns `Vec<(entity_id, activation_score)>` sorted descending by score,
/// excluding entities whose final score is below `min_activation`.
pub fn spreading_activation(
&self,
seed_ids: &[u64],
decay_factor: f32,
min_activation: f32,
max_steps: usize,
) -> Vec<(u64, f32)> {
let idx = self.adjacency_index();
let mut activation: HashMap<u64, f32> = HashMap::new();
// Initialise seeds with activation 1.0.
for &seed in seed_ids {
*activation.entry(seed).or_insert(0.0) += 1.0;
}
for _step in 0..max_steps {
// Snapshot current activations.
let current: Vec<(u64, f32)> = activation
.iter()
.filter(|&(_, &score)| score >= min_activation)
.map(|(&id, &score)| (id, score))
.collect();
if current.is_empty() {
break;
}
let mut any_spread = false;
for (source_id, source_score) in current {
// Spread only to edges touching this node, instead of
// scanning every relation in the graph per active node.
for &rel_idx in idx.relations_touching(source_id) {
let rel = &self.relations[rel_idx];
let neighbour_id = if rel.src == source_id {
rel.tgt
} else if rel.tgt == source_id {
rel.src
} else {
continue;
};
let delta = source_score * rel.weight * decay_factor;
if delta >= min_activation {
*activation.entry(neighbour_id).or_insert(0.0) += delta;
any_spread = true;
}
}
}
if !any_spread {
break;
}
}
// Collect results, drop entries below threshold.
let mut result: Vec<(u64, f32)> = activation
.into_iter()
.filter(|&(_, score)| score >= min_activation)
.collect();
result.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
result
}
// -----------------------------------------------------------------------
// Graph-aware search context
// -----------------------------------------------------------------------
/// Return a human-readable context string describing the neighbourhood of
/// `entity_id`. Includes the entity itself, all directly connected entities,
/// and the relations linking them. Suitable for injection into retrieval
/// results.
pub fn get_entity_context(&self, entity_id: u64) -> String {
let entity = match self.get_entity(entity_id) {
Some(e) => e,
None => return format!("[entity {entity_id} not found]"),
};
let mut lines: Vec<String> = Vec::new();
lines.push(format!(
"Entity: {} (id={}, type={})",
entity.name, entity.id, entity.entity_type
));
if !entity.properties.is_empty() {
let mut props: Vec<String> = entity
.properties
.iter()
.map(|(k, v)| format!("{k}={v}"))
.collect();
props.sort();
lines.push(format!(" Properties: {}", props.join(", ")));
}
// Outgoing relations.
let outgoing = self.get_relations_from(entity_id);
for rel in outgoing {
if let Some(tgt) = self.get_entity(rel.tgt) {
lines.push(format!(
" -[{}]-> {} (id={}, weight={:.3})",
rel.relation, tgt.name, tgt.id, rel.weight
));
}
}
// Incoming relations.
let incoming = self.get_relations_to(entity_id);
for rel in incoming {
if let Some(src) = self.get_entity(rel.src) {
lines.push(format!(
" <-[{}]- {} (id={}, weight={:.3})",
rel.relation, src.name, src.id, rel.weight
));
}
}
lines.join("\n")
}
}
impl Default for KnowledgeCache {
fn default() -> Self {
Self::new()
}
}
// ---------------------------------------------------------------------------
// Tests
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cached_adjacency_sees_direct_changes_to_the_graph() {
// The index is cached across traversals, but entities/relations are
// pub Vecs anyone can edit; every kind of edit must be seen.
let mut kg = KnowledgeCache::new();
let a = kg.add_entity("a", "t", -1);
let b = kg.add_entity("b", "t", -1);
let c = kg.add_entity("c", "t", -1);
kg.add_relation(a, b, "r", 1.0);
let ids = |kg: &KnowledgeCache| -> Vec<u64> {
let mut v: Vec<u64> = kg.bfs_neighbors(a, 3).iter().map(|(e, _)| e.id).collect();
v.sort();
v
};
assert_eq!(ids(&kg), vec![b]);
assert_eq!(ids(&kg), vec![b], "cached index reused");
// Pushed directly, bypassing add_relation.
kg.relations.push(Relation {
src: b,
tgt: c,
..Relation::default()
});
assert_eq!(ids(&kg), vec![b, c]);
// Rewired in place: same lengths, different edge.
kg.relations[1].tgt = a;
assert_eq!(ids(&kg), vec![b]);
// Removed and replaced: same lengths again.
kg.relations.pop();
kg.relations.push(Relation {
src: a,
tgt: c,
..Relation::default()
});
assert_eq!(ids(&kg), vec![b, c]);
let act: Vec<u64> = kg
.spreading_activation(&[a], 0.5, 0.0, 2)
.iter()
.map(|(id, _)| *id)
.collect();
assert!(act.contains(&c));
}
// -----------------------------------------------------------------------
// Original tests — must remain passing
// -----------------------------------------------------------------------
#[test]
fn test_add_and_get_aliases() {
let mut cache = KnowledgeCache::new();
let id = cache.add_entity("Henry", "person", -1);
cache.add_alias("my son", id as i64);
cache.add_alias("the kid", id as i64);
let aliases = cache.get_aliases(id as i64);
assert_eq!(aliases.len(), 2);
assert!(aliases.contains(&"my son"));
assert!(aliases.contains(&"the kid"));
}
#[test]
fn test_resolve_single_alias() {
let mut cache = KnowledgeCache::new();
let id = cache.add_entity("Henry", "person", -1);
cache.add_alias("my son", id as i64);
let resolved = cache.resolve_aliases("what does my son do?");
assert!(
resolved.contains("henry"),
"expected 'henry' in '{resolved}'"
);
}
#[test]
fn test_resolve_multiple_aliases() {
let mut cache = KnowledgeCache::new();
let h = cache.add_entity("Henry", "person", -1);
let a = cache.add_entity("Acme Corp", "company", -1);
cache.add_alias("my son", h as i64);
cache.add_alias("our main client", a as i64);
let resolved = cache.resolve_aliases("what does my son do at our main client?");
assert!(
resolved.contains("henry"),
"expected 'henry' in '{resolved}'"
);
assert!(
resolved.contains("acme corp"),
"expected 'acme corp' in '{resolved}'"
);
}
#[test]
fn test_longest_match_wins() {
let mut cache = KnowledgeCache::new();
let id = cache.add_entity("Henry", "person", -1);
cache.add_alias("my son", id as i64);
cache.add_alias("my son henry", id as i64);
let resolved = cache.resolve_aliases("ask my son henry");
// The longer alias "my son henry" should match first, producing "ask henry"
assert_eq!(resolved, "ask henry");
}
#[test]
fn test_unregistered_alias_passthrough() {
let mut cache = KnowledgeCache::new();
cache.add_entity("Henry", "person", -1);
let resolved = cache.resolve_aliases("unknown phrase here");
assert_eq!(resolved, "unknown phrase here");
}
#[test]
fn test_case_insensitive_resolve() {
let mut cache = KnowledgeCache::new();
let id = cache.add_entity("Henry", "person", -1);
cache.add_alias("henry", id as i64);
let resolved = cache.resolve_aliases("Tell HENRY about it");
assert!(
resolved.contains("henry"),
"expected 'henry' in '{resolved}'"
);
}
#[test]
fn test_empty_aliases() {
let cache = KnowledgeCache::new();
let resolved = cache.resolve_aliases("hello");
assert_eq!(resolved, "hello");
}
#[test]
fn test_get_entity_name() {
let mut cache = KnowledgeCache::new();
let id = cache.add_entity("Henry", "person", -1);
assert_eq!(cache.get_entity_name(id as i64), Some("Henry"));
assert_eq!(cache.get_entity_name(999), None);
}
// -----------------------------------------------------------------------
// Entity properties
// -----------------------------------------------------------------------
#[test]
fn test_entity_properties_default_empty() {
let mut cache = KnowledgeCache::new();
let id = cache.add_entity("Alice", "person", -1);
let entity = cache.get_entity(id).unwrap();
assert!(entity.properties.is_empty());
}
#[test]
fn test_entity_properties_set_and_get() {
let mut cache = KnowledgeCache::new();
let id = cache.add_entity("Alice", "person", -1);
{
let entity = cache.get_entity_mut(id).unwrap();
entity
.properties
.insert("role".to_string(), "engineer".to_string());
}
let entity = cache.get_entity(id).unwrap();
assert_eq!(
entity.properties.get("role").map(String::as_str),
Some("engineer")
);
}
// -----------------------------------------------------------------------
// Entity embeddings
// -----------------------------------------------------------------------
#[test]
fn test_entity_embedding_default_none() {
let mut cache = KnowledgeCache::new();
let id = cache.add_entity("Bob", "person", -1);
let entity = cache.get_entity(id).unwrap();
assert!(entity.embedding.is_none());
}
#[test]
fn test_entity_embedding_set_and_get() {
let mut cache = KnowledgeCache::new();
let id = cache.add_entity("Bob", "person", -1);
let vec = vec![0.1_f32, 0.2, 0.3];
{
let entity = cache.get_entity_mut(id).unwrap();
entity.embedding = Some(vec.clone());
}
let entity = cache.get_entity(id).unwrap();
assert_eq!(entity.embedding.as_deref(), Some(vec.as_slice()));
}
// -----------------------------------------------------------------------
// Entity timestamps
// -----------------------------------------------------------------------
#[test]
fn test_entity_timestamps_set_on_creation() {
let mut cache = KnowledgeCache::new();
let id = cache.add_entity("Carol", "person", -1);
let entity = cache.get_entity(id).unwrap();
assert!(entity.created_at > 0.0);
assert!(entity.updated_at > 0.0);
assert_eq!(entity.created_at, entity.updated_at);
}
#[test]
fn test_entity_updated_at_can_be_changed() {
let mut cache = KnowledgeCache::new();
let id = cache.add_entity("Carol", "person", -1);
let original_created = cache.get_entity(id).unwrap().created_at;
{
let entity = cache.get_entity_mut(id).unwrap();
entity.updated_at = original_created + 1000.0;
}
let entity = cache.get_entity(id).unwrap();
assert_eq!(entity.created_at, original_created);
assert!(entity.updated_at > entity.created_at);
}
// -----------------------------------------------------------------------
// Relation metadata
// -----------------------------------------------------------------------
#[test]
fn test_relation_metadata_default_empty() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
let b = cache.add_entity("B", "node", -1);
cache.add_relation(a, b, "link", 1.0);
let rels = cache.get_relations_from(a);
assert_eq!(rels.len(), 1);
assert!(rels[0].metadata.is_empty());
}
// -----------------------------------------------------------------------
// RelationType enum
// -----------------------------------------------------------------------
#[test]
fn test_relation_type_from_str_known_variants() {
assert_eq!(RelationType::from_label("temporal"), RelationType::Temporal);
assert_eq!(RelationType::from_label("causal"), RelationType::Causal);
assert_eq!(
RelationType::from_label("associative"),
RelationType::Associative
);
assert_eq!(
RelationType::from_label("hierarchical"),
RelationType::Hierarchical
);
}
#[test]
fn test_relation_type_custom() {
let rt = RelationType::from_label("something_else");
assert_eq!(rt, RelationType::Custom("something_else".to_string()));
assert_eq!(rt.as_str(), "something_else");
}
#[test]
fn test_relation_type_round_trip() {
let variants = [
RelationType::Temporal,
RelationType::Causal,
RelationType::Associative,
RelationType::Hierarchical,
RelationType::Custom("my_type".to_string()),
];
for v in &variants {
assert_eq!(RelationType::from_label(v.as_str()), *v);
}
}
// -----------------------------------------------------------------------
// Levenshtein helper
// -----------------------------------------------------------------------
#[test]
fn test_levenshtein_identical() {
assert_eq!(levenshtein("hello", "hello"), 0);
}
#[test]
fn test_levenshtein_empty_strings() {
assert_eq!(levenshtein("", ""), 0);
assert_eq!(levenshtein("abc", ""), 3);
assert_eq!(levenshtein("", "abc"), 3);
}
#[test]
fn test_levenshtein_one_edit() {
assert_eq!(levenshtein("cat", "bat"), 1);
}
#[test]
fn test_levenshtein_insertions() {
assert_eq!(levenshtein("ab", "abc"), 1);
}
// -----------------------------------------------------------------------
// resolve_or_create
// -----------------------------------------------------------------------
#[test]
fn test_resolve_or_create_creates_new() {
let mut cache = KnowledgeCache::new();
let (id, created) = cache.resolve_or_create("NewEntity", "thing", -1, 0);
assert!(created);
assert!(cache.get_entity(id).is_some());
}
#[test]
fn test_resolve_or_create_finds_exact_match() {
let mut cache = KnowledgeCache::new();
let orig_id = cache.add_entity("Alice", "person", -1);
let (id, created) = cache.resolve_or_create("Alice", "person", -1, 0);
assert!(!created);
assert_eq!(id, orig_id);
}
#[test]
fn test_resolve_or_create_fuzzy_match_within_threshold() {
let mut cache = KnowledgeCache::new();
let orig_id = cache.add_entity("Alice", "person", -1);
// "Alyce" has edit distance 1 from "Alice"
let (id, created) = cache.resolve_or_create("Alyce", "person", -1, 2);
assert!(!created, "expected fuzzy match, got a new entity");
assert_eq!(id, orig_id);
}
/// An exact match must win even when a near-match with a smaller Levenshtein
/// distance-to-zero gap was scanned first — the early exit on dist == 0
/// must not skip past a later exact match.
#[test]
fn test_resolve_or_create_exact_match_beats_earlier_fuzzy_candidate() {
let mut cache = KnowledgeCache::new();
cache.add_entity("Alyce", "person", -1); // dist 1 from "Alice"
let exact_id = cache.add_entity("Alice", "person", -1); // dist 0
let (id, created) = cache.resolve_or_create("Alice", "person", -1, 2);
assert!(!created);
assert_eq!(id, exact_id);
}
#[test]
fn test_resolve_or_create_no_match_beyond_threshold() {
let mut cache = KnowledgeCache::new();
cache.add_entity("Alice", "person", -1);
// "Zebra" is far from "Alice"
let (_id, created) = cache.resolve_or_create("Zebra", "animal", -1, 2);
assert!(created, "expected a new entity to be created");
assert_eq!(cache.entities.len(), 2);
}
#[test]
fn test_resolve_or_create_case_insensitive() {
let mut cache = KnowledgeCache::new();
let orig_id = cache.add_entity("Alice", "person", -1);
// Exact match after lower-casing
let (id, created) = cache.resolve_or_create("ALICE", "person", -1, 0);
assert!(!created);
assert_eq!(id, orig_id);
}
// -----------------------------------------------------------------------
// bfs_neighbors
// -----------------------------------------------------------------------
#[test]
fn test_bfs_neighbors_empty_graph() {
let mut cache = KnowledgeCache::new();
let id = cache.add_entity("Solo", "node", -1);
let result = cache.bfs_neighbors(id, 3);
assert!(result.is_empty());
}
#[test]
fn test_bfs_neighbors_depth_one() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
let b = cache.add_entity("B", "node", -1);
let c = cache.add_entity("C", "node", -1);
cache.add_relation(a, b, "link", 1.0);
cache.add_relation(a, c, "link", 1.0);
let result = cache.bfs_neighbors(a, 1);
assert_eq!(result.len(), 2);
let depths: Vec<usize> = result.iter().map(|(_, d)| *d).collect();
assert!(depths.iter().all(|&d| d == 1));
}
#[test]
fn test_bfs_neighbors_max_depth_zero_returns_empty() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
let b = cache.add_entity("B", "node", -1);
cache.add_relation(a, b, "link", 1.0);
let result = cache.bfs_neighbors(a, 0);
assert!(result.is_empty());
}
#[test]
fn test_bfs_neighbors_respects_max_depth() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
let b = cache.add_entity("B", "node", -1);
let c = cache.add_entity("C", "node", -1);
cache.add_relation(a, b, "link", 1.0);
cache.add_relation(b, c, "link", 1.0);
// depth=1: only B reachable
let result = cache.bfs_neighbors(a, 1);
assert_eq!(result.len(), 1);
assert_eq!(result[0].0.id, b);
// depth=2: both B and C reachable
let result2 = cache.bfs_neighbors(a, 2);
assert_eq!(result2.len(), 2);
}
#[test]
fn test_bfs_neighbors_follows_incoming_edges() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
let b = cache.add_entity("B", "node", -1);
// Relation goes B -> A, but BFS from A should still find B.
cache.add_relation(b, a, "link", 1.0);
let result = cache.bfs_neighbors(a, 1);
assert_eq!(result.len(), 1);
assert_eq!(result[0].0.id, b);
}
#[test]
fn test_bfs_neighbors_no_duplicate_visits() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
let b = cache.add_entity("B", "node", -1);
// Two edges between A and B.
cache.add_relation(a, b, "link1", 1.0);
cache.add_relation(a, b, "link2", 1.0);
let result = cache.bfs_neighbors(a, 2);
assert_eq!(result.len(), 1, "B should appear exactly once");
}
// -----------------------------------------------------------------------
// get_subgraph
// -----------------------------------------------------------------------
#[test]
fn test_get_subgraph_single_seed_no_neighbours() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
let (entities, relations) = cache.get_subgraph(&[a], 2);
assert_eq!(entities.len(), 1);
assert!(relations.is_empty());
}
#[test]
fn test_get_subgraph_includes_seed_and_neighbours() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
let b = cache.add_entity("B", "node", -1);
let c = cache.add_entity("C", "node", -1);
// D is disconnected.
let _d = cache.add_entity("D", "node", -1);
cache.add_relation(a, b, "link", 1.0);
cache.add_relation(b, c, "link", 1.0);
let (entities, relations) = cache.get_subgraph(&[a], 2);
let ids: HashSet<u64> = entities.iter().map(|e| e.id).collect();
assert!(ids.contains(&a));
assert!(ids.contains(&b));
assert!(ids.contains(&c));
assert!(!ids.contains(&_d));
assert_eq!(relations.len(), 2);
}
#[test]
fn test_get_subgraph_multiple_seeds() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
let b = cache.add_entity("B", "node", -1);
let c = cache.add_entity("C", "node", -1);
// No edges, but both seeds should appear.
let (entities, _) = cache.get_subgraph(&[a, b], 1);
let ids: HashSet<u64> = entities.iter().map(|e| e.id).collect();
assert!(ids.contains(&a));
assert!(ids.contains(&b));
assert!(!ids.contains(&c));
}
// -----------------------------------------------------------------------
// spreading_activation
// -----------------------------------------------------------------------
#[test]
fn test_spreading_activation_single_seed_no_edges() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
let result = cache.spreading_activation(&[a], 0.5, 0.01, 5);
assert_eq!(result.len(), 1);
assert_eq!(result[0].0, a);
assert!((result[0].1 - 1.0).abs() < 1e-6);
}
#[test]
fn test_spreading_activation_propagates_to_neighbour() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
let b = cache.add_entity("B", "node", -1);
cache.add_relation(a, b, "link", 1.0);
let result = cache.spreading_activation(&[a], 0.5, 0.01, 3);
// B should have received activation from A.
let b_score = result.iter().find(|&&(id, _)| id == b).map(|&(_, s)| s);
assert!(b_score.is_some());
assert!(b_score.unwrap() > 0.0);
}
/// A self-loop relation (src == tgt) must be visited exactly once by the
/// adjacency index, matching the pre-index behavior of iterating
/// `self.relations` directly (each relation processed once regardless of
/// how many of its endpoints match the current node).
#[test]
fn test_spreading_activation_self_loop_not_double_counted() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
cache.add_relation(a, a, "self", 1.0);
let result = cache.spreading_activation(&[a], 0.5, 0.0001, 1);
let a_score = result
.iter()
.find(|&&(id, _)| id == a)
.map(|&(_, s)| s)
.unwrap();
// Seed activation (1.0) plus exactly one spread contribution
// (1.0 * weight 1.0 * decay 0.5), not two.
assert!(
(a_score - 1.5).abs() < 1e-5,
"expected 1.5 (one self-loop contribution), got {a_score}"
);
}
#[test]
fn test_spreading_activation_decay_reduces_signal() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
let b = cache.add_entity("B", "node", -1);
let c = cache.add_entity("C", "node", -1);
cache.add_relation(a, b, "link", 1.0);
cache.add_relation(b, c, "link", 1.0);
// Use 1 step to test pure decay without iterative backflow accumulation.
let result = cache.spreading_activation(&[a], 0.5, 0.01, 1);
let score_of = |id: u64| -> f32 {
result
.iter()
.find(|&&(eid, _)| eid == id)
.map(|&(_, s)| s)
.unwrap_or(0.0)
};
// A starts with 1.0; B gets 0.5 after 1 step; C gets nothing (2 hops, only 1 step).
assert!(
score_of(a) >= score_of(b),
"A should have higher or equal activation than B"
);
assert!(
score_of(b) > score_of(c),
"B should have more activation than C (C unreachable in 1 step)"
);
}
#[test]
fn test_spreading_activation_sorted_descending() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
let b = cache.add_entity("B", "node", -1);
let c = cache.add_entity("C", "node", -1);
cache.add_relation(a, b, "link", 1.0);
cache.add_relation(b, c, "link", 1.0);
let result = cache.spreading_activation(&[a], 0.5, 0.01, 5);
let scores: Vec<f32> = result.iter().map(|&(_, s)| s).collect();
for window in scores.windows(2) {
assert!(window[0] >= window[1], "results must be sorted descending");
}
}
#[test]
fn test_spreading_activation_min_activation_filter() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("A", "node", -1);
let b = cache.add_entity("B", "node", -1);
cache.add_relation(a, b, "link", 0.01);
// With a very high min_activation, only the seed should appear.
let result = cache.spreading_activation(&[a], 0.5, 10.0, 5);
assert_eq!(
result.len(),
0,
"all activations below min threshold should be filtered"
);
}
// -----------------------------------------------------------------------
// get_entity_context
// -----------------------------------------------------------------------
#[test]
fn test_get_entity_context_unknown_entity() {
let cache = KnowledgeCache::new();
let ctx = cache.get_entity_context(999);
assert!(ctx.contains("not found"));
}
#[test]
fn test_get_entity_context_basic_info() {
let mut cache = KnowledgeCache::new();
let id = cache.add_entity("Alice", "person", -1);
let ctx = cache.get_entity_context(id);
assert!(ctx.contains("Alice"));
assert!(ctx.contains("person"));
}
#[test]
fn test_get_entity_context_includes_outgoing_relations() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("Alice", "person", -1);
let b = cache.add_entity("Bob", "person", -1);
cache.add_relation(a, b, "knows", 0.9);
let ctx = cache.get_entity_context(a);
assert!(
ctx.contains("knows"),
"context should mention the relation type"
);
assert!(
ctx.contains("Bob"),
"context should mention the target entity"
);
}
#[test]
fn test_get_entity_context_includes_incoming_relations() {
let mut cache = KnowledgeCache::new();
let a = cache.add_entity("Alice", "person", -1);
let b = cache.add_entity("Bob", "person", -1);
cache.add_relation(b, a, "trusts", 0.8);
let ctx = cache.get_entity_context(a);
assert!(ctx.contains("trusts"));
assert!(ctx.contains("Bob"));
}
#[test]
fn test_get_entity_context_includes_properties() {
let mut cache = KnowledgeCache::new();
let id = cache.add_entity("Alice", "person", -1);
{
let entity = cache.get_entity_mut(id).unwrap();
entity
.properties
.insert("occupation".to_string(), "engineer".to_string());
}
let ctx = cache.get_entity_context(id);
assert!(ctx.contains("occupation"));
assert!(ctx.contains("engineer"));
}
}