use std::sync::Arc;
use klieo_core::error::Error;
use klieo_core::llm::LlmClient;
use klieo_core::memory::{LongTermMemory, Scope};
use klieo_embed_common::Embedder;
use klieo_memory_graph::{
EntityExtractor, FilterableLongTermMemory, KnowledgeGraph, RecallMetrics,
};
use klieo_memory_graph_rag::{
BuiltinExtractor, FallbackExtractor, GraphAwareLongTerm, InMemoryFilterableLongTerm,
LlmEntityExtractor, ProvenanceKnowledgeGraph,
};
pub use klieo_memory_neo4j::Neo4jConfig;
const DEFAULT_MIN_GRAPH_HITS: usize = 1;
const RESERVED_PROJECTION_SCOPE: &str = "__klieo_graph_projection";
const PROVENANCE_ACTOR: &str = "klieo-app-graph-rag";
const DEFAULT_PROVENANCE_DB_PATH: &str = "klieo-provenance.sqlite";
pub(crate) const FASTEMBED_ID: &str = "fastembed:all-minilm-l6-v2";
#[non_exhaustive]
pub enum ExtractorMode {
BuiltinThenLlm,
BuiltinOnly,
Custom(Arc<dyn EntityExtractor>),
}
pub(crate) enum Backend {
InMemory,
Remote {
qdrant_url: String,
neo4j: klieo_memory_neo4j::Neo4jConfig,
},
}
#[non_exhaustive]
pub struct GraphRagConfig {
pub(crate) backend: Backend,
pub(crate) embedder: Option<(Arc<dyn Embedder>, String)>,
pub(crate) provenance: bool,
pub(crate) projector: bool,
pub(crate) extractor_mode: ExtractorMode,
pub(crate) projection_scope: Scope,
pub(crate) min_graph_hits: usize,
pub(crate) provenance_actor: String,
}
impl GraphRagConfig {
fn defaults(backend: Backend) -> Self {
Self {
backend,
embedder: None,
provenance: true,
projector: true,
extractor_mode: ExtractorMode::BuiltinThenLlm,
projection_scope: Scope::Workspace(RESERVED_PROJECTION_SCOPE.to_string()),
min_graph_hits: DEFAULT_MIN_GRAPH_HITS,
provenance_actor: PROVENANCE_ACTOR.to_string(),
}
}
pub fn in_memory() -> Self {
Self::defaults(Backend::InMemory)
}
pub fn remote(qdrant_url: impl Into<String>, neo4j: klieo_memory_neo4j::Neo4jConfig) -> Self {
Self::defaults(Backend::Remote {
qdrant_url: qdrant_url.into(),
neo4j,
})
}
pub fn embedder(mut self, embedder: Arc<dyn Embedder>, id: impl Into<String>) -> Self {
self.embedder = Some((embedder, id.into()));
self
}
pub fn without_provenance(mut self) -> Self {
self.provenance = false;
self
}
pub fn without_projector(mut self) -> Self {
self.projector = false;
self
}
pub fn builtin_extractor_only(mut self) -> Self {
self.extractor_mode = ExtractorMode::BuiltinOnly;
self
}
pub fn custom_extractor(mut self, extractor: Arc<dyn EntityExtractor>) -> Self {
self.extractor_mode = ExtractorMode::Custom(extractor);
self
}
pub fn projection_scope(mut self, scope: Scope) -> Self {
self.projection_scope = scope;
self
}
pub fn provenance_actor(mut self, actor: impl Into<String>) -> Self {
self.provenance_actor = actor.into();
self
}
pub fn min_graph_hits(mut self, n: usize) -> Self {
self.min_graph_hits = n;
self
}
}
pub(crate) struct WiredGraphRag {
pub long_term: Arc<dyn LongTermMemory>,
pub graph_aware: Arc<GraphAwareLongTerm>,
pub metrics: Arc<RecallMetrics>,
pub graph: Arc<dyn KnowledgeGraph>,
pub extractor: Arc<dyn EntityExtractor>,
pub projection_scope: Scope,
pub projector_enabled: bool,
}
struct VectorAndGraph {
vector: Arc<dyn FilterableLongTermMemory>,
raw_graph: Arc<dyn KnowledgeGraph>,
}
struct ProvenanceWiring {
enabled: bool,
is_remote: bool,
actor: String,
sqlite_path: Option<std::path::PathBuf>,
}
impl GraphRagConfig {
pub(crate) async fn wire(
self,
llm: Arc<dyn LlmClient>,
sqlite_path: Option<std::path::PathBuf>,
) -> Result<WiredGraphRag, Error> {
let (embedder, embedder_id) = self.resolve_embedder()?;
let provenance = ProvenanceWiring {
enabled: self.provenance,
is_remote: matches!(self.backend, Backend::Remote { .. }),
actor: self.provenance_actor,
sqlite_path,
};
let parts = build_vector_and_graph(self.backend, embedder, embedder_id).await?;
let graph = wrap_provenance(parts.raw_graph, provenance)?;
let extractor = build_extractor(self.extractor_mode, llm);
let metrics = Arc::new(RecallMetrics::default());
let graph_aware = Arc::new(
GraphAwareLongTerm::builder()
.extractor(extractor.clone())
.metrics(metrics.clone())
.min_graph_hits(self.min_graph_hits)
.build(parts.vector, graph.clone()),
);
let long_term: Arc<dyn LongTermMemory> = graph_aware.clone();
Ok(WiredGraphRag {
long_term,
graph_aware,
metrics,
graph,
extractor,
projection_scope: self.projection_scope,
projector_enabled: self.projector,
})
}
fn resolve_embedder(&self) -> Result<(Arc<dyn Embedder>, String), Error> {
if let Some((embedder, id)) = &self.embedder {
return Ok((embedder.clone(), id.clone()));
}
let fastembed = klieo_embed_common::FastEmbedEmbedder::new()?;
Ok((Arc::new(fastembed), FASTEMBED_ID.to_string()))
}
}
async fn build_vector_and_graph(
backend: Backend,
embedder: Arc<dyn Embedder>,
embedder_id: String,
) -> Result<VectorAndGraph, Error> {
match backend {
Backend::InMemory => Ok(VectorAndGraph {
vector: Arc::new(InMemoryFilterableLongTerm::new(embedder, embedder_id)),
raw_graph: Arc::new(klieo_memory_graph::InMemoryGraph::default()),
}),
Backend::Remote { qdrant_url, neo4j } => {
let cfg =
klieo_memory_qdrant::QdrantConfig::new(qdrant_url).with_embedder_id(embedder_id);
let qdrant = klieo_memory_qdrant::MemoryQdrant::new(cfg, embedder).await?;
let neo = klieo_memory_neo4j::MemoryNeo4j::new(neo4j).await?;
let graph = klieo_memory_graph_neo4j::Neo4jKnowledgeGraph::new(neo.neo4j_handle());
Ok(VectorAndGraph {
vector: qdrant.qdrant_long_term,
raw_graph: Arc::new(graph),
})
}
}
}
fn wrap_provenance(
raw_graph: Arc<dyn KnowledgeGraph>,
wiring: ProvenanceWiring,
) -> Result<Arc<dyn KnowledgeGraph>, Error> {
if !wiring.enabled {
return Ok(raw_graph);
}
let repo = if wiring.is_remote {
let path = wiring
.sqlite_path
.unwrap_or_else(|| std::path::PathBuf::from(DEFAULT_PROVENANCE_DB_PATH));
klieo_provenance::SqliteProvenanceRepository::open(path)
} else {
klieo_provenance::SqliteProvenanceRepository::open_in_memory()
}
.map_err(|e| Error::Downstream(Box::new(e)))?;
Ok(Arc::new(ProvenanceKnowledgeGraph::new(
raw_graph,
Arc::new(repo),
wiring.actor,
)))
}
fn build_extractor(mode: ExtractorMode, llm: Arc<dyn LlmClient>) -> Arc<dyn EntityExtractor> {
match mode {
ExtractorMode::Custom(extractor) => extractor,
ExtractorMode::BuiltinOnly => Arc::new(BuiltinExtractor::default()),
ExtractorMode::BuiltinThenLlm => Arc::new(FallbackExtractor::new(
BuiltinExtractor::default(),
LlmEntityExtractor::new(llm),
)),
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicBool, Ordering};
#[test]
fn in_memory_defaults_all_on() {
let c = GraphRagConfig::in_memory();
assert!(c.provenance);
assert!(c.projector);
assert!(matches!(c.extractor_mode, ExtractorMode::BuiltinThenLlm));
assert_eq!(c.min_graph_hits, 1);
}
#[test]
fn opt_outs_flip_flags() {
let c = GraphRagConfig::in_memory()
.without_provenance()
.without_projector()
.builtin_extractor_only();
assert!(!c.provenance);
assert!(!c.projector);
assert!(matches!(c.extractor_mode, ExtractorMode::BuiltinOnly));
}
#[test]
fn remote_stores_qdrant_url() {
let neo4j = Neo4jConfig::new("bolt://localhost:7687", "neo4j", "password".to_string());
let c = GraphRagConfig::remote("http://localhost:6333", neo4j);
match c.backend {
Backend::Remote { qdrant_url, .. } => {
assert_eq!(qdrant_url, "http://localhost:6333");
}
Backend::InMemory => panic!("remote() must produce Backend::Remote"),
}
}
const TEST_EMBEDDER_DIM: usize = 384;
fn fake_llm() -> Arc<dyn LlmClient> {
Arc::new(klieo_core::test_utils::FakeLlmClient::new("fake"))
}
fn test_embedder() -> Arc<klieo_embed_common::FakeEmbedder> {
Arc::new(klieo_embed_common::FakeEmbedder::new(TEST_EMBEDDER_DIM))
}
#[test]
fn provenance_actor_defaults_and_overrides() {
let default = GraphRagConfig::in_memory();
assert_eq!(default.provenance_actor, PROVENANCE_ACTOR);
let custom = GraphRagConfig::in_memory().provenance_actor("tenant-a");
assert_eq!(custom.provenance_actor, "tenant-a");
}
#[tokio::test]
async fn in_memory_wire_produces_working_long_term() {
let cfg = GraphRagConfig::in_memory().embedder(test_embedder(), "test-fake");
let wired = cfg
.wire(fake_llm(), None)
.await
.expect("in-memory wire should not touch any external service");
assert!(wired.projector_enabled, "projector defaults on");
let scope = Scope::Workspace("wire-test".to_string());
let target = "the wired port round-trips facts";
wired
.long_term
.remember(scope.clone(), klieo_core::memory::Fact::new(target))
.await
.expect("remember through the wired long-term port");
wired
.long_term
.remember(
scope.clone(),
klieo_core::memory::Fact::new("unrelated weather report"),
)
.await
.expect("remember of a distractor fact");
let hits = wired
.long_term
.recall(scope, "the wired port round-trips facts", 5)
.await
.expect("recall through the wired long-term port");
assert_eq!(
hits.first().map(|f| f.text.as_str()),
Some(target),
"the query-matching fact must rank first under real cosine ranking; got {hits:?}"
);
}
#[tokio::test]
async fn wire_without_provenance_skips_recorder() {
let cfg = GraphRagConfig::in_memory()
.embedder(test_embedder(), "test-fake")
.without_provenance();
let wired = cfg
.wire(fake_llm(), None)
.await
.expect("wire should succeed with provenance disabled");
let scope = Scope::Workspace("wire-no-prov".to_string());
wired
.long_term
.remember(
scope.clone(),
klieo_core::memory::Fact::new("no provenance here"),
)
.await
.expect("remember should still work without provenance");
let hits = wired
.long_term
.recall(scope, "no provenance here", 5)
.await
.expect("recall should still work without provenance");
assert!(hits.iter().any(|f| f.text.contains("no provenance")));
}
struct SentinelExtractor(Arc<AtomicBool>);
#[async_trait::async_trait]
impl klieo_memory_graph::EntityExtractor for SentinelExtractor {
async fn extract(
&self,
_text: &str,
_hints: &[klieo_memory_graph::EntityRef],
) -> Result<Vec<klieo_memory_graph::EntityRef>, klieo_core::error::MemoryError> {
self.0.store(true, Ordering::Release);
Ok(vec![])
}
}
#[tokio::test]
async fn custom_extractor_is_called_on_remember() {
let called = Arc::new(AtomicBool::new(false));
let sentinel = Arc::new(SentinelExtractor(called.clone()));
let cfg = GraphRagConfig::in_memory()
.embedder(test_embedder(), "test-fake")
.custom_extractor(sentinel);
let wired = cfg
.wire(fake_llm(), None)
.await
.expect("in-memory wire with custom extractor should succeed");
let scope = Scope::Workspace("custom-extractor-test".to_string());
wired
.long_term
.remember(scope, klieo_core::memory::Fact::new("sentinel text"))
.await
.expect("remember should invoke the extractor");
assert!(
called.load(Ordering::Acquire),
"custom_extractor must be called during remember"
);
}
}