use crate::path::{PathHop, RetrievalPath};
use crate::traits::KnowledgeGraph;
use crate::types::{EdgeKind, EntityRef, GraphEdge, GraphNode, GraphView};
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use klieo_core::error::MemoryError;
use klieo_core::ids::FactId;
use klieo_core::memory::Scope;
use petgraph::stable_graph::{NodeIndex, StableGraph};
use petgraph::Undirected;
use std::collections::{HashMap, HashSet};
use std::sync::Mutex;
#[derive(Debug, Clone)]
enum Node {
Entity(EntityRef),
FactRef(FactId),
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct ScopeKey {
kind: &'static str,
value: String,
}
impl From<&Scope> for ScopeKey {
fn from(scope: &Scope) -> Self {
match scope {
Scope::Workspace(s) => Self {
kind: "workspace",
value: s.clone(),
},
Scope::Agent(s) => Self {
kind: "agent",
value: s.clone(),
},
Scope::Global => Self {
kind: "global",
value: String::new(),
},
}
}
}
#[derive(Debug, Clone)]
enum Edge {
MentionedIn,
CoOccurs { count: u32 },
}
struct Inner {
graph: StableGraph<Node, Edge, Undirected>,
entity_idx: HashMap<(String, String, ScopeKey), NodeIndex>,
factref_idx: HashMap<(String, ScopeKey), NodeIndex>,
fact_texts: HashMap<FactId, String>,
}
pub struct InMemoryGraph {
inner: Mutex<Inner>,
}
impl Default for InMemoryGraph {
fn default() -> Self {
Self {
inner: Mutex::new(Inner {
graph: StableGraph::default(),
entity_idx: HashMap::new(),
factref_idx: HashMap::new(),
fact_texts: HashMap::new(),
}),
}
}
}
#[async_trait]
impl KnowledgeGraph for InMemoryGraph {
async fn index(
&self,
scope: Scope,
fact_id: &FactId,
entities: &[EntityRef],
text: &str,
_valid_from: Option<DateTime<Utc>>,
) -> Result<(), MemoryError> {
if entities.is_empty() {
tracing::debug!(%fact_id, "index called with empty entities; no-op");
return Ok(());
}
let sk = ScopeKey::from(&scope);
let mut guard = self.inner.lock().map_err(|_| {
tracing::error!(operation = "InMemoryGraph::index", "mutex poisoned");
MemoryError::Store("InMemoryGraph mutex poisoned".into())
})?;
let Inner {
graph,
entity_idx,
factref_idx,
fact_texts,
} = &mut *guard;
if !text.is_empty() {
fact_texts.insert(fact_id.clone(), text.to_string());
}
let fact_node = *factref_idx
.entry((fact_id.to_string(), sk.clone()))
.or_insert_with(|| graph.add_node(Node::FactRef(fact_id.clone())));
let mut entity_nodes: Vec<NodeIndex> = Vec::with_capacity(entities.len());
for entity in entities {
let key = (
entity.entity_type.as_str().to_owned(),
entity.name.clone(),
sk.clone(),
);
let ent_node = *entity_idx
.entry(key)
.or_insert_with(|| graph.add_node(Node::Entity(entity.clone())));
entity_nodes.push(ent_node);
if !graph.contains_edge(ent_node, fact_node) {
graph.add_edge(ent_node, fact_node, Edge::MentionedIn);
}
}
for i in 0..entity_nodes.len() {
for j in (i + 1)..entity_nodes.len() {
let (a, b) = (entity_nodes[i], entity_nodes[j]);
if let Some(edge) = graph.find_edge(a, b) {
if let Some(Edge::CoOccurs { count }) = graph.edge_weight_mut(edge) {
*count += 1;
}
} else {
graph.add_edge(a, b, Edge::CoOccurs { count: 1 });
}
}
}
Ok(())
}
async fn neighbors(
&self,
scope: &Scope,
entities: &[EntityRef],
) -> Result<Vec<FactId>, MemoryError> {
if entities.is_empty() {
tracing::debug!("neighbors called with empty entities; no-op");
return Ok(Vec::new());
}
let sk = ScopeKey::from(scope);
let guard = self.inner.lock().map_err(|_| {
tracing::error!(operation = "InMemoryGraph::neighbors", "mutex poisoned");
MemoryError::Store("InMemoryGraph mutex poisoned".into())
})?;
let Inner {
graph, entity_idx, ..
} = &*guard;
let mut seen: HashSet<FactId> = HashSet::new();
for entity in entities {
let key = (
entity.entity_type.as_str().to_owned(),
entity.name.clone(),
sk.clone(),
);
let Some(&ent_node) = entity_idx.get(&key) else {
continue;
};
collect_reachable_fact_ids(graph, ent_node, &mut seen);
}
Ok(seen.into_iter().collect())
}
async fn recall_paths(
&self,
scope: &Scope,
entities: &[EntityRef],
) -> Result<Vec<RetrievalPath>, MemoryError> {
if entities.is_empty() {
return Ok(Vec::new());
}
let sk = ScopeKey::from(scope);
let guard = self.inner.lock().map_err(|_| {
tracing::error!(operation = "InMemoryGraph::recall_paths", "mutex poisoned");
MemoryError::Store("InMemoryGraph mutex poisoned".into())
})?;
let Inner {
graph, entity_idx, ..
} = &*guard;
let mut paths: Vec<RetrievalPath> = Vec::new();
let mut seen: HashSet<(String, String, FactId)> = HashSet::new();
for entity in entities {
let key = (
entity.entity_type.as_str().to_owned(),
entity.name.clone(),
sk.clone(),
);
let Some(&ent_node) = entity_idx.get(&key) else {
continue;
};
let mut fact_ids: HashSet<FactId> = HashSet::new();
collect_reachable_fact_ids(graph, ent_node, &mut fact_ids);
for fid in fact_ids {
let dedupe_key = (
entity.entity_type.as_str().to_owned(),
entity.name.clone(),
fid.clone(),
);
if !seen.insert(dedupe_key) {
continue;
}
paths.push(RetrievalPath {
hops: vec![PathHop::new(entity.clone(), fid)],
});
}
}
Ok(paths)
}
async fn subgraph(&self, scope: &Scope, limit: usize) -> Result<GraphView, MemoryError> {
let sk = ScopeKey::from(scope);
let guard = self.inner.lock().map_err(|_| {
tracing::error!(operation = "InMemoryGraph::subgraph", "mutex poisoned");
MemoryError::Store("InMemoryGraph mutex poisoned".into())
})?;
let Inner {
graph,
entity_idx,
factref_idx,
fact_texts,
} = &*guard;
let mut scope_nodes: Vec<NodeIndex> = entity_idx
.iter()
.filter(|(key, _)| key.2 == sk)
.map(|(_, &idx)| idx)
.chain(
factref_idx
.iter()
.filter(|(key, _)| key.1 == sk)
.map(|(_, &idx)| idx),
)
.collect();
scope_nodes.sort_unstable();
let mut nodes = Vec::new();
let mut keep: HashSet<NodeIndex> = HashSet::new();
let mut truncated = false;
for &idx in &scope_nodes {
if nodes.len() >= limit {
truncated = true;
break;
}
nodes.push(node_dto(graph, idx));
keep.insert(idx);
}
let edges = graph
.edge_indices()
.filter_map(|e| {
let (a, b) = graph.edge_endpoints(e)?;
if !keep.contains(&a) || !keep.contains(&b) {
return None;
}
let kind = match graph[e] {
Edge::MentionedIn => EdgeKind::MentionedIn,
Edge::CoOccurs { .. } => EdgeKind::CoOccurs,
};
Some(GraphEdge::new(node_dto(graph, a), node_dto(graph, b), kind))
})
.collect();
let view_fact_texts = nodes
.iter()
.filter_map(|node| match node {
GraphNode::Fact(fid) => fact_texts.get(fid).map(|t| (fid.clone(), t.clone())),
_ => None,
})
.collect();
let view = if truncated {
GraphView::truncated(nodes, edges)
} else {
GraphView::complete(nodes, edges)
};
Ok(view.with_fact_texts(view_fact_texts))
}
async fn forget(&self, scope: &Scope, fact_id: &FactId) -> Result<(), MemoryError> {
let sk = ScopeKey::from(scope);
let mut guard = self.inner.lock().map_err(|_| {
tracing::error!(operation = "InMemoryGraph::forget", "mutex poisoned");
MemoryError::Store("InMemoryGraph mutex poisoned".into())
})?;
let Inner {
graph,
factref_idx,
fact_texts,
..
} = &mut *guard;
fact_texts.remove(fact_id);
let Some(fact_node) = factref_idx.remove(&(fact_id.to_string(), sk)) else {
return Ok(());
};
let fact_edges: Vec<_> = graph
.edges(fact_node)
.map(|e| petgraph::visit::EdgeRef::id(&e))
.collect();
for edge in fact_edges {
graph.remove_edge(edge);
}
graph.remove_node(fact_node);
Ok(())
}
}
fn node_dto(graph: &StableGraph<Node, Edge, Undirected>, idx: NodeIndex) -> GraphNode {
match &graph[idx] {
Node::Entity(entity) => GraphNode::Entity(entity.clone()),
Node::FactRef(fid) => GraphNode::Fact(fid.clone()),
}
}
fn collect_reachable_fact_ids(
graph: &StableGraph<Node, Edge, Undirected>,
entry_entity: NodeIndex,
seen: &mut HashSet<FactId>,
) {
for neighbor in graph.neighbors(entry_entity) {
match &graph[neighbor] {
Node::FactRef(fid) => {
seen.insert(fid.clone());
}
Node::Entity(_) => {
for deeper in graph.neighbors(neighbor) {
if let Node::FactRef(fid) = &graph[deeper] {
seen.insert(fid.clone());
}
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::EntityType;
#[tokio::test]
async fn subgraph_returns_indexed_nodes_and_edges() {
let g = InMemoryGraph::default();
let scope = Scope::Workspace("w".into());
let fact = FactId::new("f1");
let alice = EntityRef::new(EntityType::Member, "alice");
let acme = EntityRef::new(EntityType::Member, "acme");
g.index(
scope.clone(),
&fact,
&[alice.clone(), acme.clone()],
"alice at acme",
None,
)
.await
.unwrap();
let view = g.subgraph(&scope, 100).await.unwrap();
assert!(!view.truncated);
assert!(view.nodes.contains(&GraphNode::Fact(fact.clone())));
assert!(view.nodes.contains(&GraphNode::Entity(alice.clone())));
assert!(view.edges.iter().any(|e| e.kind == EdgeKind::MentionedIn
&& e.from == GraphNode::Entity(alice.clone())
&& e.to == GraphNode::Fact(fact.clone())));
assert!(view.edges.iter().any(|e| e.kind == EdgeKind::CoOccurs));
}
#[tokio::test]
async fn subgraph_surfaces_fact_text_and_forget_drops_it() {
let g = InMemoryGraph::default();
let scope = Scope::Workspace("w".into());
let fact = FactId::new("f1");
let alice = EntityRef::new(EntityType::Member, "alice");
g.index(scope.clone(), &fact, &[alice], "alice filed a claim", None)
.await
.unwrap();
let view = g.subgraph(&scope, 100).await.unwrap();
assert_eq!(
view.fact_texts.get(&fact).map(String::as_str),
Some("alice filed a claim"),
);
g.forget(&scope, &fact).await.unwrap();
let after = g.subgraph(&scope, 100).await.unwrap();
assert!(
after.fact_texts.is_empty(),
"forget must drop the fact's stored text"
);
}
#[tokio::test]
async fn subgraph_empty_scope_is_empty_view() {
let g = InMemoryGraph::default();
let view = g.subgraph(&Scope::Global, 100).await.unwrap();
assert!(view.nodes.is_empty() && view.edges.is_empty() && !view.truncated);
}
#[tokio::test]
async fn subgraph_truncates_at_limit() {
let g = InMemoryGraph::default();
let scope = Scope::Workspace("w".into());
for i in 0..5 {
let f = FactId::new(format!("f{i}"));
g.index(
scope.clone(),
&f,
&[EntityRef::new(EntityType::Member, format!("p{i}"))],
"x",
None,
)
.await
.unwrap();
}
let view = g.subgraph(&scope, 3).await.unwrap();
assert!(view.truncated);
assert_eq!(view.nodes.len(), 3);
}
}