use crate::extractor::BuiltinExtractor;
use crate::metadata::{KEY_ENTITIES, KEY_VALID_FROM};
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use klieo_core::error::MemoryError;
use klieo_core::ids::FactId;
use klieo_core::memory::{Fact, LongTermMemory, Scope};
use klieo_memory_graph::{
EntityExtractor, EntityRef, FilterableLongTermMemory, KnowledgeGraph, RecallMetrics,
RetrievalPath,
};
use std::sync::Arc;
const DEFAULT_MIN_GRAPH_HITS: usize = 1;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct RecallTrace {
pub extracted_entities: Vec<EntityRef>,
pub paths: Vec<RetrievalPath>,
pub ranked_facts: Vec<Fact>,
pub fell_back_to_vector: bool,
}
struct RecallCore {
entities: Vec<EntityRef>,
facts: Vec<Fact>,
fell_back_to_vector: bool,
graph_candidates: usize,
}
pub struct GraphAwareLongTerm {
vector: Arc<dyn FilterableLongTermMemory>,
graph: Arc<dyn KnowledgeGraph>,
extractor: Arc<dyn EntityExtractor>,
metrics: Arc<RecallMetrics>,
min_graph_hits: usize,
}
impl GraphAwareLongTerm {
pub fn new(
vector: Arc<dyn FilterableLongTermMemory>,
graph: Arc<dyn KnowledgeGraph>,
extractor: Arc<dyn EntityExtractor>,
metrics: Arc<RecallMetrics>,
) -> Self {
Self {
vector,
graph,
extractor,
metrics,
min_graph_hits: DEFAULT_MIN_GRAPH_HITS,
}
}
pub fn with_min_graph_hits(mut self, n: usize) -> Self {
self.min_graph_hits = n;
self
}
pub fn builder() -> GraphAwareLongTermBuilder {
GraphAwareLongTermBuilder::default()
}
async fn recall_core(
&self,
scope: &Scope,
query: &str,
k: usize,
) -> Result<RecallCore, MemoryError> {
let entities = self
.extractor
.extract(query, &[])
.await
.unwrap_or_else(|e| {
tracing::warn!(error = %e, "extractor failed on recall; falling back to pure vector");
Vec::new()
});
if !entities.is_empty() {
match self.graph.neighbors(scope, &entities).await {
Ok(candidate_ids) if candidate_ids.len() >= self.min_graph_hits => {
let facts = self
.vector
.recall_filtered(scope.clone(), query, k, &candidate_ids)
.await?;
return Ok(RecallCore {
entities,
facts,
fell_back_to_vector: false,
graph_candidates: candidate_ids.len(),
});
}
Ok(_) => { }
Err(e) => tracing::warn!(
error = %e,
"graph neighbors failed on recall; falling back to pure vector"
),
}
}
let facts = self.vector.recall(scope.clone(), query, k).await?;
Ok(RecallCore {
entities,
facts,
fell_back_to_vector: true,
graph_candidates: 0,
})
}
pub async fn recall_traced(
&self,
scope: Scope,
query: &str,
k: usize,
) -> Result<RecallTrace, MemoryError> {
let core = self.recall_core(&scope, query, k).await?;
let paths = match self.graph.recall_paths(&scope, &core.entities).await {
Ok(paths) => paths,
Err(e) => {
tracing::warn!(
?e,
"recall_paths failed in recall_traced; returning empty paths"
);
Vec::new()
}
};
Ok(RecallTrace {
extracted_entities: core.entities,
paths,
ranked_facts: core.facts,
fell_back_to_vector: core.fell_back_to_vector,
})
}
}
#[derive(Default)]
pub struct GraphAwareLongTermBuilder {
extractor: Option<Arc<dyn EntityExtractor>>,
metrics: Option<Arc<RecallMetrics>>,
min_graph_hits: Option<usize>,
}
impl GraphAwareLongTermBuilder {
pub fn extractor(mut self, extractor: Arc<dyn EntityExtractor>) -> Self {
self.extractor = Some(extractor);
self
}
pub fn metrics(mut self, metrics: Arc<RecallMetrics>) -> Self {
self.metrics = Some(metrics);
self
}
pub fn min_graph_hits(mut self, n: usize) -> Self {
self.min_graph_hits = Some(n);
self
}
pub fn build(
self,
vector: Arc<dyn FilterableLongTermMemory>,
graph: Arc<dyn KnowledgeGraph>,
) -> GraphAwareLongTerm {
GraphAwareLongTerm {
vector,
graph,
extractor: self
.extractor
.unwrap_or_else(|| Arc::new(BuiltinExtractor::default())),
metrics: self
.metrics
.unwrap_or_else(|| Arc::new(RecallMetrics::default())),
min_graph_hits: self.min_graph_hits.unwrap_or(DEFAULT_MIN_GRAPH_HITS),
}
}
}
fn extract_hints_from_metadata(metadata: &serde_json::Value) -> Vec<EntityRef> {
metadata
.get(KEY_ENTITIES)
.and_then(|v| serde_json::from_value::<Vec<EntityRef>>(v.clone()).ok())
.unwrap_or_default()
}
fn parse_valid_from(metadata: &serde_json::Value) -> Option<DateTime<Utc>> {
metadata
.get(KEY_VALID_FROM)
.and_then(|v| v.as_str())
.and_then(|s| DateTime::parse_from_rfc3339(s).ok())
.map(|dt| dt.with_timezone(&Utc))
}
#[async_trait]
impl LongTermMemory for GraphAwareLongTerm {
async fn remember(&self, scope: Scope, fact: Fact) -> Result<FactId, MemoryError> {
let fact_id = self.vector.remember(scope.clone(), fact.clone()).await?;
let hints = extract_hints_from_metadata(&fact.metadata);
let valid_from = parse_valid_from(&fact.metadata);
let entities = self
.extractor
.extract(&fact.text, &hints)
.await
.unwrap_or_else(|e| {
tracing::warn!(error = %e, "extractor failed; indexing graph with caller hints only");
hints.clone()
});
if entities.is_empty() {
return Ok(fact_id);
}
if let Err(e) = self
.graph
.index(scope, &fact_id, &entities, &fact.text, valid_from)
.await
{
tracing::warn!(
fact_id = %fact_id,
error = %e,
"graph index failed; vector remains authoritative"
);
}
Ok(fact_id)
}
async fn recall(&self, scope: Scope, query: &str, k: usize) -> Result<Vec<Fact>, MemoryError> {
let core = self.recall_core(&scope, query, k).await?;
if core.fell_back_to_vector {
self.metrics.record_vector_fallback();
} else {
self.metrics
.record_graph_hit_with_candidates(core.graph_candidates as u64);
}
Ok(core.facts)
}
async fn forget(&self, id: FactId) -> Result<(), MemoryError> {
self.vector.forget(id).await
}
}