diff --git a/Cargo.lock b/Cargo.lock index 6c308e6..9376670 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2840,6 +2840,7 @@ dependencies = [ "futures", "lancedb", "openraft", + "petgraph 0.6.5", "prometheus", "prost", "rand 0.9.4", @@ -2965,6 +2966,12 @@ version = "0.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" +[[package]] +name = "fixedbitset" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ce7134b9999ecaf8bcd65542e436736ef32ddca1b3e06094cb6ec5755203b80" + [[package]] name = "fixedbitset" version = "0.5.7" @@ -5649,13 +5656,23 @@ version = "0.4.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "df202b0b0f5b8e389955afd5f27b007b00fb948162953f1db9c70d2c7e3157d7" +[[package]] +name = "petgraph" +version = "0.6.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b4c5cc86750666a3ed20bdaf5ca2a0344f9c67674cae0515bec2da16fbaa47db" +dependencies = [ + "fixedbitset 0.4.2", + "indexmap 2.14.0", +] + [[package]] name = "petgraph" version = "0.7.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3672b37090dbd86368a4145bc067582552b29c27377cad4e0a306c97f9bd7772" dependencies = [ - "fixedbitset", + "fixedbitset 0.5.7", "indexmap 2.14.0", ] @@ -5665,7 +5682,7 @@ version = "0.8.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8701b58ea97060d5e5b155d383a69952a60943f0e6dfe30b04c287beb0b27455" dependencies = [ - "fixedbitset", + "fixedbitset 0.5.7", "hashbrown 0.15.5", "indexmap 2.14.0", "serde", diff --git a/Cargo.toml b/Cargo.toml index 289c28d..63bb947 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -29,6 +29,7 @@ utoipa-swagger-ui = { version = "9", features = ["axum", "vendored"] } uuid = { version = "1", features = ["v4", "serde"] } anyhow = "1" openraft = { version = "0.9", features = ["serde", "storage-v2"] } +petgraph = "0.6" tonic = "0.12" prost = "0.13" tokio-stream = "0.1" diff --git a/benches/e2e_throughput.rs b/benches/e2e_throughput.rs index 30d8981..92ee803 100644 --- a/benches/e2e_throughput.rs +++ b/benches/e2e_throughput.rs @@ -175,6 +175,8 @@ impl BenchmarkHarness { )); let (embedding_job_sender, receiver) = embedding_job_channel(config.channel_size); + let (knowledge_job_sender, mut krx) = tokio::sync::mpsc::channel::(16); + tokio::spawn(async move { while krx.recv().await.is_some() {} }); let state = Arc::new(AppState { short_term_memory: short_term_memory.clone(), vector_store: vector_store.clone(), @@ -191,6 +193,10 @@ impl BenchmarkHarness { raft_addr: None, raft_advertise_addr: None, cluster_peers: vec![], + knowledge_graph: Arc::new(tokio::sync::RwLock::new( + engram::knowledge::graph::KnowledgeGraph::new(), + )), + knowledge_job_sender, }); let _worker_handles = spawn_embedding_workers( diff --git a/docker-compose.cluster.yml b/docker-compose.cluster.yml index cb11744..9d7d731 100644 --- a/docker-compose.cluster.yml +++ b/docker-compose.cluster.yml @@ -25,7 +25,8 @@ services: REDIS_URL: "redis://redis-1:6379" LANCE_DB_PATH: "/data/lancedb" ENGRAM_BIND_ADDR: "0.0.0.0:3000" - OPENAI_API_KEY: "${OPENAI_API_KEY}" + OPENAI_API_KEY: "${OPENAI_API_KEY:-}" + KNOWLEDGE_EXTRACTOR: "${KNOWLEDGE_EXTRACTOR:-mock}" RUST_LOG: "info,openraft=debug" ports: - "3000:3000" @@ -45,7 +46,8 @@ services: REDIS_URL: "redis://redis-2:6379" LANCE_DB_PATH: "/data/lancedb" ENGRAM_BIND_ADDR: "0.0.0.0:3000" - OPENAI_API_KEY: "${OPENAI_API_KEY}" + OPENAI_API_KEY: "${OPENAI_API_KEY:-}" + KNOWLEDGE_EXTRACTOR: "${KNOWLEDGE_EXTRACTOR:-mock}" RUST_LOG: "info,openraft=debug" ports: - "3001:3000" @@ -65,7 +67,8 @@ services: REDIS_URL: "redis://redis-3:6379" LANCE_DB_PATH: "/data/lancedb" ENGRAM_BIND_ADDR: "0.0.0.0:3000" - OPENAI_API_KEY: "${OPENAI_API_KEY}" + OPENAI_API_KEY: "${OPENAI_API_KEY:-}" + KNOWLEDGE_EXTRACTOR: "${KNOWLEDGE_EXTRACTOR:-mock}" RUST_LOG: "info,openraft=debug" ports: - "3002:3000" diff --git a/scripts/cluster-verify.sh b/scripts/cluster-verify.sh index 58ba9cf..e7ec1ec 100755 --- a/scripts/cluster-verify.sh +++ b/scripts/cluster-verify.sh @@ -96,5 +96,97 @@ echo "$METRICS" | grep -q "engram_raft_term" && pass "engram_raft_term p echo "$METRICS" | grep -q "engram_raft_commit_index" && pass "engram_raft_commit_index present" || fail "engram_raft_commit_index missing" echo "$METRICS" | grep -q "engram_raft_is_leader" && pass "engram_raft_is_leader present" || fail "engram_raft_is_leader missing" + +echo "[6] Knowledge replication (deterministic mock extraction)" +# Find current leader — may have changed after failover in check [4] +WRITE_LEADER="$N1" +for port in 3000 3001 3002; do + ROLE=$(curl -sf "http://localhost:$port/cluster" | \ + python3 -c "import sys,json; print(json.load(sys.stdin)['role'])" 2>/dev/null || echo "") + if [ "$ROLE" = "Leader" ]; then + WRITE_LEADER="http://localhost:$port" + break + fi +done + +SESSION_K=$(curl -sf -X POST "$WRITE_LEADER/sessions" | \ + python3 -c "import sys,json; print(json.load(sys.stdin)['session_id'])") + +# Three separate messages: one pattern per message so the mock extractor handles each cleanly +curl -sf -X POST "$WRITE_LEADER/sessions/$SESSION_K/messages" \ + -H "Content-Type: application/json" \ + -d '{"role":"user","content":"Alice works at OpenAI"}' > /dev/null +curl -sf -X POST "$WRITE_LEADER/sessions/$SESSION_K/messages" \ + -H "Content-Type: application/json" \ + -d '{"role":"user","content":"Bob works at OpenAI"}' > /dev/null +curl -sf -X POST "$WRITE_LEADER/sessions/$SESSION_K/messages" \ + -H "Content-Type: application/json" \ + -d '{"role":"user","content":"Alice knows Bob"}' > /dev/null + +echo " Waiting 3 seconds for extraction and Raft replication..." +sleep 3 + +LEADER_PORT="" +for port in 3000 3001 3002; do + ROLE=$(curl -sf "http://localhost:$port/cluster" | \ + python3 -c "import sys,json; print(json.load(sys.stdin)['role'])" 2>/dev/null || echo "") + if [ "$ROLE" = "Leader" ]; then + LEADER_PORT="$port" + break + fi +done + +[ -z "$LEADER_PORT" ] && fail "no leader found for knowledge check" + +LEADER_ENTITIES=$(curl -sf "http://localhost:$LEADER_PORT/sessions/$SESSION_K/knowledge" | \ + python3 -c "import sys,json; print(len(json.load(sys.stdin)['entities']))" 2>/dev/null || echo "-1") + +[ "$LEADER_ENTITIES" -ge 3 ] \ + && pass "leader (:$LEADER_PORT) has $LEADER_ENTITIES entities" \ + || fail "leader (:$LEADER_PORT) has $LEADER_ENTITIES entities (expected >= 3)" + +LEADER_EDGES=$(curl -sf "http://localhost:$LEADER_PORT/sessions/$SESSION_K/knowledge" | \ + python3 -c "import sys,json; print(len(json.load(sys.stdin)['edges']))" 2>/dev/null || echo "-1") + +[ "$LEADER_EDGES" -ge 3 ] \ + && pass "leader (:$LEADER_PORT) has $LEADER_EDGES relationships" \ + || fail "leader (:$LEADER_PORT) has $LEADER_EDGES relationships (expected >= 3)" + +for port in 3000 3001 3002; do + [ "$port" -eq "$LEADER_PORT" ] && continue + FOLLOWER_ENTITIES=$(curl -sf "http://localhost:$port/sessions/$SESSION_K/knowledge" | \ + python3 -c "import sys,json; print(len(json.load(sys.stdin)['entities']))" 2>/dev/null || echo "-1") + [ "$FOLLOWER_ENTITIES" -eq "$LEADER_ENTITIES" ] \ + && pass "follower :$port converged to $FOLLOWER_ENTITIES entities (matches leader)" \ + || fail "follower :$port has $FOLLOWER_ENTITIES entities (leader has $LEADER_ENTITIES)" +done + +echo "[6b] Capability criterion: graph answers questions without LLM or vector search" +RELATED=$(curl -sf "http://localhost:$LEADER_PORT/sessions/$SESSION_K/knowledge/entities/OpenAI" | \ + python3 -c "import sys,json; d=json.load(sys.stdin); print([r['name'] for r in d['related']])" 2>/dev/null || echo "[]") +echo "$RELATED" | grep -q "Alice" \ + && pass "OpenAI is related to Alice (works_at)" \ + || fail "OpenAI not related to Alice" +echo "$RELATED" | grep -q "Bob" \ + && pass "OpenAI is related to Bob (works_at)" \ + || fail "OpenAI not related to Bob" + +PATH_RESP=$(curl -sf "http://localhost:$LEADER_PORT/sessions/$SESSION_K/knowledge/path?from=Alice&to=Bob" | \ + python3 -c "import sys,json; d=json.load(sys.stdin); print(d.get('path'))" 2>/dev/null || echo "None") +[ "$PATH_RESP" != "None" ] && [ "$PATH_RESP" != "null" ] \ + && pass "shortest path Alice→Bob found via graph traversal" \ + || fail "no path found from Alice to Bob" + +echo "[6c] Delete-session removes knowledge graph state from all nodes" +curl -sf -X DELETE "$WRITE_LEADER/sessions/$SESSION_K" > /dev/null +sleep 1 +for port in 3000 3001 3002; do + ENTITIES_AFTER=$(curl -sf "http://localhost:$port/sessions/$SESSION_K/knowledge" | \ + python3 -c "import sys,json; print(len(json.load(sys.stdin)['entities']))" 2>/dev/null || echo "-1") + [ "$ENTITIES_AFTER" -eq 0 ] \ + && pass "node :$port knowledge graph empty after delete" \ + || fail "node :$port still has $ENTITIES_AFTER entities after delete" +done + echo "" -echo "=== All Stage 1 criteria PASSED ===" +echo "=== All Stage 1 + Stage 2 criteria PASSED ===" diff --git a/src/app.rs b/src/app.rs index 2240691..c2cc033 100644 --- a/src/app.rs +++ b/src/app.rs @@ -11,6 +11,10 @@ use crate::core::{ StoreError, TokenCounter, VectorStore, }; use crate::embedding::OpenAIEmbedder; +use crate::config::KnowledgeExtractorType; +use crate::knowledge::extractor::{MockKnowledgeExtractor, OpenAIKnowledgeExtractor}; +use crate::knowledge::graph::KnowledgeGraph; +use crate::knowledge::worker::{knowledge_job_channel, spawn_knowledge_workers}; use crate::metrics::AppMetrics; use crate::server::AppState; use crate::stores::{LanceDBStore, RedisCoreMemoryStore, RedisShortTermMemory}; @@ -36,6 +40,8 @@ pub async fn build_raft_node( core_memory: Arc, vector_store: Arc, embedding_tx: tokio::sync::mpsc::Sender, + knowledge_graph: Arc>, + knowledge_tx: tokio::sync::mpsc::Sender, ) -> anyhow::Result> { use crate::raft::{ log_store::EngRaftLogStore, network::EngRaftNetwork, @@ -61,7 +67,7 @@ pub async fn build_raft_node( raft_config, EngRaftNetwork, EngRaftLogStore::default(), - EngStateMachineStore::new(short_term, core_memory, vector_store, embedding_tx), + EngStateMachineStore::new(short_term, core_memory, vector_store, embedding_tx, knowledge_graph, knowledge_tx), ) .await?; @@ -96,12 +102,16 @@ pub fn spawn_raft_metrics_watcher( } pub async fn build_real_app_state(config: &Config) -> Result, AppBuildError> { - let embedding_provider: Arc = match &config.openai_base_url { - Some(base_url) => Arc::new(OpenAIEmbedder::new_with_base_url( - config.openai_api_key.clone(), - base_url.clone(), - )?), - None => Arc::new(OpenAIEmbedder::new_with_api_key(config.openai_api_key.clone())?), + let embedding_provider: Arc = if config.openai_api_key.is_empty() { + Arc::new(crate::core::RandomEmbeddingProvider) + } else { + match &config.openai_base_url { + Some(base_url) => Arc::new(OpenAIEmbedder::new_with_base_url( + config.openai_api_key.clone(), + base_url.clone(), + )?), + None => Arc::new(OpenAIEmbedder::new_with_api_key(config.openai_api_key.clone())?), + } }; build_app_state_with_embedding_provider(config, embedding_provider).await @@ -136,6 +146,21 @@ pub async fn build_app_state_with_embedding_provider( config.embedding_max_concurrency, ); + let knowledge_graph = Arc::new(tokio::sync::RwLock::new(KnowledgeGraph::new())); + let (knowledge_job_sender, knowledge_receiver) = knowledge_job_channel(config.knowledge_channel_size); + + let knowledge_extractor: Arc = + match config.knowledge_extractor { + KnowledgeExtractorType::Mock => Arc::new(MockKnowledgeExtractor), + KnowledgeExtractorType::OpenAI => match &config.openai_base_url { + Some(base_url) => Arc::new(OpenAIKnowledgeExtractor::new_with_base_url( + config.openai_api_key.clone(), + base_url.clone(), + )), + None => Arc::new(OpenAIKnowledgeExtractor::new(config.openai_api_key.clone())), + }, + }; + let (raft, node_id, peer_http_addrs, raft_addr, raft_advertise_addr, cluster_peers) = if config.node_id.is_some() { let raft = build_raft_node( config, @@ -143,6 +168,8 @@ pub async fn build_app_state_with_embedding_provider( core_memory_store.clone(), vector_store.clone(), embedding_job_sender.clone(), + knowledge_graph.clone(), + knowledge_job_sender.clone(), ) .await .map_err(|e| AppBuildError::Other(e.into()))?; @@ -160,6 +187,16 @@ pub async fn build_app_state_with_embedding_provider( spawn_raft_metrics_watcher(raft_handle.clone(), metrics.clone()); } + let _knowledge_worker_handles = spawn_knowledge_workers( + knowledge_extractor, + raft.clone(), + config.node_id.unwrap_or(0), + knowledge_graph.clone(), + metrics.clone(), + knowledge_receiver, + config.knowledge_max_workers, + ); + Ok(Arc::new(AppState { short_term_memory, vector_store, @@ -176,5 +213,7 @@ pub async fn build_app_state_with_embedding_provider( raft_addr, raft_advertise_addr, cluster_peers, + knowledge_graph, + knowledge_job_sender, })) } diff --git a/src/cluster.rs b/src/cluster.rs index 21421cd..ed1fdf9 100644 --- a/src/cluster.rs +++ b/src/cluster.rs @@ -202,12 +202,17 @@ mod tests { async fn build_test_app_with_single_node_raft() -> TestServer { let c = build_test_components(); let config = Config { node_id: Some(1), ..Config::default() }; + let knowledge_graph = Arc::new(tokio::sync::RwLock::new(crate::knowledge::graph::KnowledgeGraph::new())); + let (knowledge_tx, mut knowledge_rx) = tokio::sync::mpsc::channel::(500); + tokio::spawn(async move { while knowledge_rx.recv().await.is_some() {} }); let raft = build_raft_node( &config, c.short_term.clone(), c.core_memory.clone(), c.vector_store.clone(), c.embedding_job_sender.clone(), + knowledge_graph.clone(), + knowledge_tx.clone(), ) .await .unwrap(); @@ -234,12 +239,16 @@ mod tests { raft_addr: Some("127.0.0.1:0".to_string()), raft_advertise_addr: None, cluster_peers: vec![], + knowledge_graph, + knowledge_job_sender: knowledge_tx, }); TestServer::new(build_router(state)).unwrap() } fn build_test_app_standalone() -> TestServer { let c = build_test_components(); + let (knowledge_job_sender, mut krx) = tokio::sync::mpsc::channel::(16); + tokio::spawn(async move { while krx.recv().await.is_some() {} }); let state = Arc::new(AppState { short_term_memory: c.short_term, vector_store: c.vector_store, @@ -256,6 +265,10 @@ mod tests { raft_addr: None, raft_advertise_addr: None, cluster_peers: vec![], + knowledge_graph: Arc::new(tokio::sync::RwLock::new( + crate::knowledge::graph::KnowledgeGraph::new(), + )), + knowledge_job_sender, }); TestServer::new(build_router(state)).unwrap() } diff --git a/src/config.rs b/src/config.rs index 811aaf2..289d883 100644 --- a/src/config.rs +++ b/src/config.rs @@ -15,6 +15,12 @@ pub struct PeerConfig { pub addr: String, // "host:grpc_port" } +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum KnowledgeExtractorType { + OpenAI, + Mock, +} + const DEFAULT_REDIS_URL: &str = "redis://localhost:6379"; const DEFAULT_LANCE_DB_PATH: &str = "./data/lancedb"; const DEFAULT_EMBEDDING_DIMENSION: usize = 1536; @@ -48,6 +54,9 @@ pub struct Config { /// HTTP addresses of peer nodes keyed by node ID, parsed from CLUSTER_HTTP_PEERS. /// Format: "id:host:http_port,id:host:http_port" pub cluster_http_peers: HashMap, + pub knowledge_max_workers: usize, + pub knowledge_channel_size: usize, + pub knowledge_extractor: KnowledgeExtractorType, } #[derive(Debug, Error)] @@ -76,18 +85,29 @@ impl Default for Config { raft_advertise_addr: None, cluster_peers: vec![], cluster_http_peers: HashMap::new(), + knowledge_max_workers: 4, + knowledge_channel_size: 500, + knowledge_extractor: KnowledgeExtractorType::OpenAI, } } } impl Config { pub fn from_env() -> Result { + let knowledge_extractor = match env::var("KNOWLEDGE_EXTRACTOR").as_deref() { + Ok("mock") => KnowledgeExtractorType::Mock, + _ => KnowledgeExtractorType::OpenAI, + }; Ok(Self { redis_url: env::var("REDIS_URL") .ok() .filter(|value| !value.trim().is_empty()) .unwrap_or_else(|| DEFAULT_REDIS_URL.to_string()), - openai_api_key: required_env("OPENAI_API_KEY")?, + openai_api_key: if knowledge_extractor == KnowledgeExtractorType::Mock { + env::var("OPENAI_API_KEY").unwrap_or_default() + } else { + required_env("OPENAI_API_KEY")? + }, openai_base_url: optional_env("OPENAI_BASE_URL")?, lance_db_path: PathBuf::from(optional_lance_db_path()?), embedding_dimension: positive_usize_env( @@ -115,6 +135,9 @@ impl Config { cluster_http_peers: Self::parse_http_peers( &env::var("CLUSTER_HTTP_PEERS").unwrap_or_default(), ), + knowledge_max_workers: positive_usize_env("KNOWLEDGE_MAX_WORKERS", 4)?, + knowledge_channel_size: positive_usize_env("KNOWLEDGE_CHANNEL_SIZE", 500)?, + knowledge_extractor, }) } @@ -214,7 +237,7 @@ mod tests { use std::env; use std::sync::{Mutex, OnceLock}; - use super::{Config, ConfigError}; + use super::{Config, ConfigError, KnowledgeExtractorType}; fn env_lock() -> &'static Mutex<()> { static ENV_LOCK: OnceLock> = OnceLock::new(); @@ -342,6 +365,51 @@ mod tests { )); } + #[test] + fn knowledge_extractor_defaults_to_openai() { + let _guard = env_lock().lock().unwrap(); + let old_key = env::var("OPENAI_API_KEY").ok(); + let old_extractor = env::var("KNOWLEDGE_EXTRACTOR").ok(); + unsafe { + env::set_var("OPENAI_API_KEY", "test-key"); + env::remove_var("KNOWLEDGE_EXTRACTOR"); + } + let config = Config::from_env().unwrap(); + assert_eq!(config.knowledge_extractor, KnowledgeExtractorType::OpenAI); + restore_env("OPENAI_API_KEY", old_key); + restore_env("KNOWLEDGE_EXTRACTOR", old_extractor); + } + + #[test] + fn knowledge_extractor_mock_parsed_from_env() { + let _guard = env_lock().lock().unwrap(); + let old_key = env::var("OPENAI_API_KEY").ok(); + let old_extractor = env::var("KNOWLEDGE_EXTRACTOR").ok(); + unsafe { + env::remove_var("OPENAI_API_KEY"); + env::set_var("KNOWLEDGE_EXTRACTOR", "mock"); + } + let config = Config::from_env().unwrap(); + assert_eq!(config.knowledge_extractor, KnowledgeExtractorType::Mock); + restore_env("OPENAI_API_KEY", old_key); + restore_env("KNOWLEDGE_EXTRACTOR", old_extractor); + } + + #[test] + fn mock_mode_does_not_require_openai_api_key() { + let _guard = env_lock().lock().unwrap(); + let old_key = env::var("OPENAI_API_KEY").ok(); + let old_extractor = env::var("KNOWLEDGE_EXTRACTOR").ok(); + unsafe { + env::remove_var("OPENAI_API_KEY"); + env::set_var("KNOWLEDGE_EXTRACTOR", "mock"); + } + let result = Config::from_env(); + assert!(result.is_ok(), "mock mode should not require OPENAI_API_KEY"); + restore_env("OPENAI_API_KEY", old_key); + restore_env("KNOWLEDGE_EXTRACTOR", old_extractor); + } + fn restore_env(name: &str, value: Option) { match value { Some(value) => unsafe { env::set_var(name, value) }, diff --git a/src/knowledge/export.rs b/src/knowledge/export.rs new file mode 100644 index 0000000..6bbbb92 --- /dev/null +++ b/src/knowledge/export.rs @@ -0,0 +1,104 @@ +use serde::{Deserialize, Serialize}; +use crate::knowledge::types::{Entity, Relationship}; + +#[derive(Debug, Serialize, Deserialize)] +pub struct GraphExport { + pub session_id: String, + pub entities: Vec, + pub edges: Vec, +} + +impl GraphExport { + pub fn new(session_id: impl Into, entities: Vec, edges: Vec) -> Self { + Self { session_id: session_id.into(), entities, edges } + } +} + +pub fn to_dot(export: &GraphExport) -> String { + let mut lines = vec!["digraph knowledge {".to_string()]; + for entity in &export.entities { + lines.push(format!( + " \"{}\" [label=\"{}\\n({})\"];", + escape_dot(&entity.name), + escape_dot(&entity.name), + entity.entity_type, + )); + } + for edge in &export.edges { + lines.push(format!( + " \"{}\" -> \"{}\" [label=\"{}\"];", + escape_dot(&edge.from), + escape_dot(&edge.to), + edge.relationship_type, + )); + } + lines.push("}".to_string()); + lines.join("\n") +} + +fn escape_dot(s: &str) -> String { + s.replace('\\', "\\\\").replace('"', "\\\"") +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::knowledge::types::{Entity, Relationship}; + use std::collections::HashMap; + + fn make_export() -> GraphExport { + GraphExport { + session_id: "s1".to_string(), + entities: vec![ + Entity { name: "Alice".into(), entity_type: "Person".into(), attributes: HashMap::new() }, + Entity { name: "OpenAI".into(), entity_type: "Organization".into(), attributes: HashMap::new() }, + ], + edges: vec![ + Relationship { from: "Alice".into(), to: "OpenAI".into(), relationship_type: "works_at".into() }, + ], + } + } + + #[test] + fn empty_graph_produces_minimal_dot() { + let empty = GraphExport { session_id: "s1".into(), entities: vec![], edges: vec![] }; + let dot = to_dot(&empty); + assert_eq!(dot.trim(), "digraph knowledge {\n}"); + } + + #[test] + fn dot_contains_entity_nodes_with_type() { + let dot = to_dot(&make_export()); + assert!(dot.contains("\"Alice\"")); + assert!(dot.contains("\"OpenAI\"")); + assert!(dot.contains("Person")); + } + + #[test] + fn dot_contains_edge_with_label() { + let dot = to_dot(&make_export()); + assert!(dot.contains("\"Alice\" -> \"OpenAI\"")); + assert!(dot.contains("works_at")); + } + + #[test] + fn dot_escapes_double_quotes_in_names() { + let export = GraphExport { + session_id: "s1".into(), + entities: vec![Entity { name: "Alice \"A\"".into(), entity_type: "Person".into(), attributes: HashMap::new() }], + edges: vec![], + }; + let dot = to_dot(&export); + assert!(!dot.contains("Alice \"A\"")); + } + + #[test] + fn graph_export_json_round_trips() { + let export = make_export(); + let json = serde_json::to_string(&export).unwrap(); + let back: GraphExport = serde_json::from_str(&json).unwrap(); + assert_eq!(back.entities.len(), 2); + assert_eq!(back.edges.len(), 1); + assert_eq!(back.session_id, "s1"); + } +} diff --git a/src/knowledge/extractor.rs b/src/knowledge/extractor.rs new file mode 100644 index 0000000..54efdf0 --- /dev/null +++ b/src/knowledge/extractor.rs @@ -0,0 +1,409 @@ +use std::collections::HashMap; +use async_trait::async_trait; +use reqwest::Client; +use serde::{Deserialize, Serialize}; +use thiserror::Error; + +use crate::knowledge::types::{Entity, ExtractionResult, Relationship}; + +#[derive(Debug, Error)] +pub enum ExtractError { + #[error("extraction API error: {0}")] + Api(String), + #[error("extraction parse error: {0}")] + Parse(String), + #[error("rate limit exceeded after {retries} retries")] + RateLimitExceeded { retries: u32 }, +} + +#[async_trait] +pub trait KnowledgeExtractor: Send + Sync { + async fn extract(&self, text: &str) -> Result; +} + +const SYSTEM_PROMPT: &str = r#"You are a knowledge extraction system. Extract named entities and relationships from the given text. + +Respond with valid JSON in exactly this format: +{"entities": [{"name": "string", "type": "Person|Organization|Place|Concept|Event|Other"}], "relationships": [{"from": "entity_name", "to": "entity_name", "type": "relationship_type"}]} + +Rules: +- Keep entity names as they appear in the text. +- Relationship types must be snake_case (e.g. works_at, knows, located_in, created_by, part_of). +- Only include relationships between entities you extracted. +- If no entities or relationships are found, return empty arrays. +- Respond with only the JSON object, no surrounding text."#; + +#[derive(Serialize)] +struct ChatRequest<'a> { + model: &'a str, + messages: Vec>, + response_format: ResponseFormat, + temperature: f32, +} + +#[derive(Serialize)] +struct ChatMessage<'a> { + role: &'a str, + content: &'a str, +} + +#[derive(Serialize)] +struct ResponseFormat { + #[serde(rename = "type")] + format_type: &'static str, +} + +#[derive(Deserialize)] +struct ChatResponse { + choices: Vec, +} + +#[derive(Deserialize)] +struct Choice { + message: AssistantMessage, +} + +#[derive(Deserialize)] +struct AssistantMessage { + content: String, +} + +#[derive(Deserialize)] +struct RawExtractionResult { + entities: Vec, + #[serde(default)] + relationships: Vec, +} + +#[derive(Deserialize)] +struct RawEntity { + name: String, + #[serde(rename = "type")] + entity_type: String, +} + +#[derive(Deserialize)] +struct RawRelationship { + from: String, + to: String, + #[serde(rename = "type")] + relationship_type: String, +} + +pub struct MockKnowledgeExtractor; + +fn push_unique_entity(entities: &mut Vec, name: &str, entity_type: &str) { + if !entities.iter().any(|e| e.name == name) { + entities.push(Entity { + name: name.to_string(), + entity_type: entity_type.to_string(), + attributes: HashMap::new(), + }); + } +} + +#[async_trait] +impl KnowledgeExtractor for MockKnowledgeExtractor { + async fn extract(&self, text: &str) -> Result { + let mut entities: Vec = Vec::new(); + let mut relationships: Vec = Vec::new(); + + for sentence in text.split(['.', '!', '?']) { + let s = sentence.trim(); + if s.is_empty() { + continue; + } + if let Some((left, right)) = s.split_once(" works at ") { + let (p, o) = (left.trim().to_string(), right.trim().to_string()); + if !p.is_empty() && !o.is_empty() { + push_unique_entity(&mut entities, &p, "Person"); + push_unique_entity(&mut entities, &o, "Organization"); + relationships.push(Relationship { from: p, to: o, relationship_type: "works_at".into() }); + } + } else if let Some((left, right)) = s.split_once(" knows ") { + let (p1, p2) = (left.trim().to_string(), right.trim().to_string()); + if !p1.is_empty() && !p2.is_empty() { + push_unique_entity(&mut entities, &p1, "Person"); + push_unique_entity(&mut entities, &p2, "Person"); + relationships.push(Relationship { from: p1, to: p2, relationship_type: "knows".into() }); + } + } else if let Some((left, right)) = s.split_once(" likes ") { + let (p, o) = (left.trim().to_string(), right.trim().to_string()); + if !p.is_empty() && !o.is_empty() { + push_unique_entity(&mut entities, &p, "Person"); + push_unique_entity(&mut entities, &o, "Thing"); + relationships.push(Relationship { from: p, to: o, relationship_type: "likes".into() }); + } + } else if let Some((left, right)) = s.split_once(" lives in ") { + let (p, o) = (left.trim().to_string(), right.trim().to_string()); + if !p.is_empty() && !o.is_empty() { + push_unique_entity(&mut entities, &p, "Person"); + push_unique_entity(&mut entities, &o, "Place"); + relationships.push(Relationship { from: p, to: o, relationship_type: "lives_in".into() }); + } + } + } + + Ok(ExtractionResult { entities, relationships }) + } +} + +pub struct OpenAIKnowledgeExtractor { + client: Client, + api_key: String, + base_url: String, + model: String, + max_retries: u32, +} + +impl OpenAIKnowledgeExtractor { + pub fn new(api_key: String) -> Self { + Self::new_with_base_url(api_key, "https://api.openai.com".to_string()) + } + + pub fn new_with_base_url(api_key: String, base_url: String) -> Self { + Self { + client: Client::new(), + api_key, + base_url, + model: "gpt-4o-mini".to_string(), + max_retries: 3, + } + } +} + +#[async_trait] +impl KnowledgeExtractor for OpenAIKnowledgeExtractor { + async fn extract(&self, text: &str) -> Result { + let url = format!("{}/v1/chat/completions", self.base_url); + let req = ChatRequest { + model: &self.model, + messages: vec![ + ChatMessage { role: "system", content: SYSTEM_PROMPT }, + ChatMessage { role: "user", content: text }, + ], + response_format: ResponseFormat { format_type: "json_object" }, + temperature: 0.0, + }; + + let mut attempt = 0u32; + loop { + let resp = self + .client + .post(&url) + .bearer_auth(&self.api_key) + .json(&req) + .send() + .await + .map_err(|e| ExtractError::Api(e.to_string()))?; + + match resp.status().as_u16() { + 200..=299 => { + let chat: ChatResponse = resp + .json() + .await + .map_err(|e| ExtractError::Parse(e.to_string()))?; + let content = chat + .choices + .into_iter() + .next() + .ok_or_else(|| ExtractError::Parse("empty choices array".to_string()))? + .message + .content; + let raw: RawExtractionResult = serde_json::from_str(&content) + .map_err(|e| ExtractError::Parse(format!("{e}: {content}")))?; + let entities = raw + .entities + .into_iter() + .map(|e| Entity { + name: e.name, + entity_type: e.entity_type, + attributes: HashMap::new(), + }) + .collect(); + let relationships = raw + .relationships + .into_iter() + .map(|r| Relationship { + from: r.from, + to: r.to, + relationship_type: r.relationship_type, + }) + .collect(); + return Ok(ExtractionResult { entities, relationships }); + } + 429 => { + attempt += 1; + if attempt > self.max_retries { + return Err(ExtractError::RateLimitExceeded { retries: self.max_retries }); + } + let backoff_ms = std::cmp::min(1000u64 << attempt.saturating_sub(1), 30_000); + tokio::time::sleep(tokio::time::Duration::from_millis(backoff_ms)).await; + } + status => { + let body = resp.text().await.unwrap_or_default(); + return Err(ExtractError::Api(format!("HTTP {status}: {body}"))); + } + } + } + } +} + +#[cfg(test)] +mod mock_tests { + use super::*; + + #[tokio::test] + async fn mock_extracts_works_at_relationship() { + let result = MockKnowledgeExtractor.extract("Alice works at OpenAI").await.unwrap(); + assert_eq!(result.entities.len(), 2); + let alice = result.entities.iter().find(|e| e.name == "Alice").unwrap(); + assert_eq!(alice.entity_type, "Person"); + let openai = result.entities.iter().find(|e| e.name == "OpenAI").unwrap(); + assert_eq!(openai.entity_type, "Organization"); + assert_eq!(result.relationships.len(), 1); + assert_eq!(result.relationships[0].from, "Alice"); + assert_eq!(result.relationships[0].to, "OpenAI"); + assert_eq!(result.relationships[0].relationship_type, "works_at"); + } + + #[tokio::test] + async fn mock_extracts_knows_relationship() { + let result = MockKnowledgeExtractor.extract("Bob knows Alice").await.unwrap(); + assert_eq!(result.entities.len(), 2); + let bob = result.entities.iter().find(|e| e.name == "Bob").unwrap(); + assert_eq!(bob.entity_type, "Person"); + let alice = result.entities.iter().find(|e| e.name == "Alice").unwrap(); + assert_eq!(alice.entity_type, "Person"); + assert_eq!(result.relationships[0].relationship_type, "knows"); + } + + #[tokio::test] + async fn mock_extracts_likes_relationship() { + let result = MockKnowledgeExtractor.extract("Alice likes Rust").await.unwrap(); + assert_eq!(result.relationships[0].relationship_type, "likes"); + let thing = result.entities.iter().find(|e| e.name == "Rust").unwrap(); + assert_eq!(thing.entity_type, "Thing"); + } + + #[tokio::test] + async fn mock_extracts_lives_in_relationship() { + let result = MockKnowledgeExtractor.extract("Bob lives in Paris").await.unwrap(); + assert_eq!(result.relationships[0].relationship_type, "lives_in"); + let place = result.entities.iter().find(|e| e.name == "Paris").unwrap(); + assert_eq!(place.entity_type, "Place"); + } + + #[tokio::test] + async fn mock_returns_empty_for_unknown_input() { + let result = MockKnowledgeExtractor.extract("The sky is blue").await.unwrap(); + assert!(result.entities.is_empty()); + assert!(result.relationships.is_empty()); + } + + #[tokio::test] + async fn mock_handles_multi_sentence_input() { + let result = MockKnowledgeExtractor + .extract("Alice works at OpenAI. Bob knows Alice.") + .await + .unwrap(); + assert_eq!(result.entities.len(), 3); // Alice, OpenAI, Bob + assert_eq!(result.relationships.len(), 2); + } + + #[tokio::test] + async fn mock_deduplicates_entities_across_sentences() { + let result = MockKnowledgeExtractor + .extract("Alice works at OpenAI. Bob knows Alice.") + .await + .unwrap(); + let alice_count = result.entities.iter().filter(|e| e.name == "Alice").count(); + assert_eq!(alice_count, 1); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + fn chat_response(content: &str) -> serde_json::Value { + json!({ "choices": [{ "message": { "content": content } }] }) + } + + #[tokio::test] + async fn extracts_entities_and_relationships() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/chat/completions")) + .respond_with(ResponseTemplate::new(200).set_body_json(chat_response( + r#"{"entities":[{"name":"Alice","type":"Person"},{"name":"OpenAI","type":"Organization"}],"relationships":[{"from":"Alice","to":"OpenAI","type":"works_at"}]}"#, + ))) + .mount(&server) + .await; + + let extractor = OpenAIKnowledgeExtractor::new_with_base_url("sk-test".into(), server.uri()); + let result = extractor.extract("Alice works at OpenAI").await.unwrap(); + + assert_eq!(result.entities.len(), 2); + assert_eq!(result.entities[0].name, "Alice"); + assert_eq!(result.entities[0].entity_type, "Person"); + assert_eq!(result.relationships.len(), 1); + assert_eq!(result.relationships[0].relationship_type, "works_at"); + } + + #[tokio::test] + async fn returns_empty_when_no_entities_found() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/chat/completions")) + .respond_with(ResponseTemplate::new(200).set_body_json(chat_response( + r#"{"entities":[],"relationships":[]}"#, + ))) + .mount(&server) + .await; + + let extractor = OpenAIKnowledgeExtractor::new_with_base_url("sk-test".into(), server.uri()); + let result = extractor.extract("the sky is blue").await.unwrap(); + assert!(result.entities.is_empty()); + assert!(result.relationships.is_empty()); + } + + #[tokio::test] + async fn retries_on_rate_limit_and_succeeds() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/chat/completions")) + .respond_with(ResponseTemplate::new(429)) + .up_to_n_times(2) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/v1/chat/completions")) + .respond_with(ResponseTemplate::new(200).set_body_json(chat_response( + r#"{"entities":[],"relationships":[]}"#, + ))) + .mount(&server) + .await; + + let extractor = OpenAIKnowledgeExtractor::new_with_base_url("sk-test".into(), server.uri()); + let result = extractor.extract("test").await.unwrap(); + assert!(result.entities.is_empty()); + } + + #[tokio::test] + async fn exhausted_retries_returns_error() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/chat/completions")) + .respond_with(ResponseTemplate::new(429)) + .mount(&server) + .await; + + let extractor = OpenAIKnowledgeExtractor::new_with_base_url("sk-test".into(), server.uri()); + let err = extractor.extract("test").await.unwrap_err(); + assert!(matches!(err, ExtractError::RateLimitExceeded { .. })); + } +} diff --git a/src/knowledge/graph.rs b/src/knowledge/graph.rs new file mode 100644 index 0000000..31e95c2 --- /dev/null +++ b/src/knowledge/graph.rs @@ -0,0 +1,330 @@ +use std::collections::{HashMap, HashSet, VecDeque}; +use petgraph::graph::{DiGraph, NodeIndex}; +use petgraph::visit::EdgeRef; +use petgraph::Direction; +use serde::{Deserialize, Serialize}; + +use crate::knowledge::types::{Entity, Relationship}; + +struct EntityNode { + name: String, + entity_type: String, + attributes: HashMap, +} + +struct RelEdge { + relationship_type: String, +} + +struct SessionGraph { + graph: DiGraph, + name_to_idx: HashMap, +} + +impl SessionGraph { + fn new() -> Self { + Self { graph: DiGraph::new(), name_to_idx: HashMap::new() } + } + + fn ensure_entity(&mut self, name: &str, entity_type: &str, attributes: HashMap) -> NodeIndex { + if let Some(&idx) = self.name_to_idx.get(name) { + return idx; + } + let idx = self.graph.add_node(EntityNode { + name: name.to_string(), + entity_type: entity_type.to_string(), + attributes, + }); + self.name_to_idx.insert(name.to_string(), idx); + idx + } +} + +#[derive(Debug, Serialize, Deserialize)] +pub enum RelationshipDirection { + Incoming, + Outgoing, +} + +#[derive(Debug, Serialize, Deserialize)] +pub struct RelatedEntity { + pub name: String, + pub entity_type: String, + pub relationship_type: String, + pub direction: RelationshipDirection, +} + +#[derive(Debug, Serialize, Deserialize)] +pub struct PathEdge { + pub from: String, + pub relationship_type: String, + pub to: String, +} + +/// Per-session in-memory knowledge graph. +/// Wrap in `Arc>` for shared access. +pub struct KnowledgeGraph { + sessions: HashMap, + /// Dedup set using "session_id\x00message_id" keys. + processed: HashSet, +} + +impl KnowledgeGraph { + pub fn new() -> Self { + Self { sessions: HashMap::new(), processed: HashSet::new() } + } + + fn dedup_key(session_id: &str, message_id: &str) -> String { + format!("{}\x00{}", session_id, message_id) + } + + pub fn is_processed(&self, session_id: &str, message_id: &str) -> bool { + self.processed.contains(&Self::dedup_key(session_id, message_id)) + } + + /// Returns `false` if this (session_id, message_id) was already processed. + pub fn apply_extraction( + &mut self, + session_id: &str, + message_id: &str, + entities: Vec, + relationships: Vec, + ) -> bool { + let key = Self::dedup_key(session_id, message_id); + if self.processed.contains(&key) { + return false; + } + self.processed.insert(key); + + let session = self.sessions.entry(session_id.to_string()).or_insert_with(SessionGraph::new); + + for entity in &entities { + session.ensure_entity(&entity.name, &entity.entity_type, entity.attributes.clone()); + } + for rel in &relationships { + let from_idx = session.ensure_entity(&rel.from, "Other", HashMap::new()); + let to_idx = session.ensure_entity(&rel.to, "Other", HashMap::new()); + session.graph.add_edge(from_idx, to_idx, RelEdge { relationship_type: rel.relationship_type.clone() }); + } + true + } + + pub fn get_related(&self, session_id: &str, entity_name: &str) -> Vec { + let Some(session) = self.sessions.get(session_id) else { return vec![] }; + let Some(&idx) = session.name_to_idx.get(entity_name) else { return vec![] }; + + let mut related = Vec::new(); + for edge in session.graph.edges_directed(idx, Direction::Outgoing) { + let node = &session.graph[edge.target()]; + related.push(RelatedEntity { + name: node.name.clone(), + entity_type: node.entity_type.clone(), + relationship_type: edge.weight().relationship_type.clone(), + direction: RelationshipDirection::Outgoing, + }); + } + for edge in session.graph.edges_directed(idx, Direction::Incoming) { + let node = &session.graph[edge.source()]; + related.push(RelatedEntity { + name: node.name.clone(), + entity_type: node.entity_type.clone(), + relationship_type: edge.weight().relationship_type.clone(), + direction: RelationshipDirection::Incoming, + }); + } + related + } + + /// BFS shortest path following outgoing edges. Returns `None` if no path exists. + pub fn find_path(&self, session_id: &str, from: &str, to: &str) -> Option> { + let session = self.sessions.get(session_id)?; + let &from_idx = session.name_to_idx.get(from)?; + let &to_idx = session.name_to_idx.get(to)?; + + if from_idx == to_idx { return Some(vec![]); } + + let mut parent: HashMap = HashMap::new(); + let mut queue = VecDeque::new(); + queue.push_back(from_idx); + + 'bfs: while let Some(current) = queue.pop_front() { + for edge in session.graph.edges_directed(current, Direction::Outgoing) { + let next = edge.target(); + if parent.contains_key(&next) { continue; } + parent.insert(next, (current, edge.weight().relationship_type.clone())); + if next == to_idx { break 'bfs; } + queue.push_back(next); + } + } + + if !parent.contains_key(&to_idx) { return None; } + + let mut path = Vec::new(); + let mut node = to_idx; + while node != from_idx { + let (prev, rel) = parent.remove(&node).unwrap(); + path.push(PathEdge { + from: session.graph[prev].name.clone(), + relationship_type: rel, + to: session.graph[node].name.clone(), + }); + node = prev; + } + path.reverse(); + Some(path) + } + + pub fn all_entities(&self, session_id: &str) -> Vec { + let Some(session) = self.sessions.get(session_id) else { return vec![] }; + session.graph.node_indices() + .map(|idx| { + let n = &session.graph[idx]; + Entity { name: n.name.clone(), entity_type: n.entity_type.clone(), attributes: n.attributes.clone() } + }) + .collect() + } + + pub fn all_relationships(&self, session_id: &str) -> Vec { + let Some(session) = self.sessions.get(session_id) else { return vec![] }; + session.graph.edge_indices() + .map(|eidx| { + let (src, tgt) = session.graph.edge_endpoints(eidx).unwrap(); + Relationship { + from: session.graph[src].name.clone(), + to: session.graph[tgt].name.clone(), + relationship_type: session.graph[eidx].relationship_type.clone(), + } + }) + .collect() + } + + pub fn delete_session(&mut self, session_id: &str) { + self.sessions.remove(session_id); + let prefix = format!("{}\x00", session_id); + self.processed.retain(|k| !k.starts_with(&prefix)); + } +} + +impl Default for KnowledgeGraph { + fn default() -> Self { Self::new() } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::knowledge::types::{Entity, Relationship}; + use std::collections::HashMap; + + fn entity(name: &str, t: &str) -> Entity { + Entity { name: name.into(), entity_type: t.into(), attributes: HashMap::new() } + } + fn rel(from: &str, to: &str, t: &str) -> Relationship { + Relationship { from: from.into(), to: to.into(), relationship_type: t.into() } + } + + #[test] + fn who_works_at_openai() { + let mut kg = KnowledgeGraph::new(); + kg.apply_extraction("s1", "m1", + vec![entity("Alice","Person"), entity("OpenAI","Organization")], + vec![rel("Alice","OpenAI","works_at")]); + kg.apply_extraction("s1", "m2", + vec![entity("Bob","Person"), entity("OpenAI","Organization")], + vec![rel("Bob","OpenAI","works_at")]); + + let related = kg.get_related("s1", "OpenAI"); + let workers: Vec<&str> = related.iter() + .filter(|r| r.relationship_type == "works_at" && matches!(r.direction, RelationshipDirection::Incoming)) + .map(|r| r.name.as_str()) + .collect(); + + assert!(workers.contains(&"Alice"), "Alice should work at OpenAI"); + assert!(workers.contains(&"Bob"), "Bob should work at OpenAI"); + } + + #[test] + fn who_does_alice_know() { + let mut kg = KnowledgeGraph::new(); + kg.apply_extraction("s1", "m1", + vec![entity("Alice","Person"), entity("Bob","Person")], + vec![rel("Alice","Bob","knows")]); + + let related = kg.get_related("s1", "Alice"); + let known: Vec<&str> = related.iter() + .filter(|r| r.relationship_type == "knows" && matches!(r.direction, RelationshipDirection::Outgoing)) + .map(|r| r.name.as_str()) + .collect(); + + assert!(known.contains(&"Bob"), "Alice should know Bob"); + } + + #[test] + fn path_alice_to_bob_direct() { + let mut kg = KnowledgeGraph::new(); + kg.apply_extraction("s1", "m1", + vec![entity("Alice","Person"), entity("Bob","Person")], + vec![rel("Alice","Bob","knows")]); + + let path = kg.find_path("s1", "Alice", "Bob").unwrap(); + assert_eq!(path.len(), 1); + assert_eq!(path[0].relationship_type, "knows"); + assert_eq!(path[0].from, "Alice"); + assert_eq!(path[0].to, "Bob"); + } + + #[test] + fn path_alice_to_bob_via_openai() { + let mut kg = KnowledgeGraph::new(); + kg.apply_extraction("s1", "m1", + vec![entity("Alice","Person"), entity("OpenAI","Organization")], + vec![rel("Alice","OpenAI","works_at")]); + kg.apply_extraction("s1", "m2", + vec![entity("Bob","Person"), entity("OpenAI","Organization")], + vec![rel("OpenAI","Bob","employs")]); + + let path = kg.find_path("s1", "Alice", "Bob").unwrap(); + assert_eq!(path.len(), 2); + assert_eq!(path[0].from, "Alice"); + assert_eq!(path[1].to, "Bob"); + } + + #[test] + fn apply_extraction_is_idempotent() { + let mut kg = KnowledgeGraph::new(); + assert!(kg.apply_extraction("s1", "m1", vec![entity("Alice","Person")], vec![])); + assert!(!kg.apply_extraction("s1", "m1", vec![entity("Alice","Person")], vec![])); + assert_eq!(kg.all_entities("s1").len(), 1); + } + + #[test] + fn delete_session_removes_all_data_and_resets_dedup() { + let mut kg = KnowledgeGraph::new(); + kg.apply_extraction("s1", "m1", + vec![entity("Alice","Person"), entity("Bob","Person")], + vec![rel("Alice","Bob","knows")]); + kg.delete_session("s1"); + assert!(kg.all_entities("s1").is_empty()); + assert!(kg.apply_extraction("s1", "m1", vec![entity("Alice","Person")], vec![])); + } + + #[test] + fn get_related_empty_for_unknown_entity() { + let kg = KnowledgeGraph::new(); + assert!(kg.get_related("s1", "nobody").is_empty()); + } + + #[test] + fn find_path_returns_none_when_no_path_exists() { + let mut kg = KnowledgeGraph::new(); + kg.apply_extraction("s1", "m1", vec![entity("Alice","Person")], vec![]); + kg.apply_extraction("s1", "m2", vec![entity("Bob","Person")], vec![]); + assert!(kg.find_path("s1", "Alice", "Bob").is_none()); + } + + #[test] + fn session_isolation_prevents_cross_session_queries() { + let mut kg = KnowledgeGraph::new(); + kg.apply_extraction("s1", "m1", vec![entity("Alice","Person")], vec![]); + assert!(kg.all_entities("s2").is_empty()); + } +} diff --git a/src/knowledge/handler.rs b/src/knowledge/handler.rs new file mode 100644 index 0000000..2ab5382 --- /dev/null +++ b/src/knowledge/handler.rs @@ -0,0 +1,246 @@ +use std::sync::Arc; +use axum::{ + extract::{Path, Query, State}, + http::StatusCode, + Json, +}; +use serde::{Deserialize, Serialize}; + +use crate::knowledge::export::{to_dot, GraphExport}; +use crate::knowledge::graph::RelatedEntity; +use crate::server::AppState; + +#[derive(Serialize)] +pub struct KnowledgeResponse { + session_id: String, + entities: Vec, + edges: Vec, +} + +pub async fn get_knowledge( + State(state): State>, + Path(session_id): Path, +) -> Json { + let kg = state.knowledge_graph.read().await; + Json(KnowledgeResponse { + session_id: session_id.clone(), + entities: kg.all_entities(&session_id), + edges: kg.all_relationships(&session_id), + }) +} + +#[derive(Serialize)] +pub struct RelatedResponse { + entity_name: String, + related: Vec, +} + +pub async fn get_related( + State(state): State>, + Path((session_id, entity_name)): Path<(String, String)>, +) -> Result, StatusCode> { + let kg = state.knowledge_graph.read().await; + let exists = kg.all_entities(&session_id).iter().any(|e| e.name == entity_name); + if !exists { + return Err(StatusCode::NOT_FOUND); + } + let related = kg.get_related(&session_id, &entity_name); + Ok(Json(RelatedResponse { entity_name, related })) +} + +#[derive(Deserialize)] +pub struct PathQuery { + from: String, + to: String, +} + +#[derive(Serialize)] +pub struct PathResponse { + from: String, + to: String, + path: Option>, +} + +pub async fn find_path( + State(state): State>, + Path(session_id): Path, + Query(params): Query, +) -> Json { + let kg = state.knowledge_graph.read().await; + let path = kg.find_path(&session_id, ¶ms.from, ¶ms.to); + Json(PathResponse { from: params.from, to: params.to, path }) +} + +#[derive(Deserialize)] +pub struct ExportQuery { + #[serde(default = "default_format")] + format: String, +} + +fn default_format() -> String { "json".to_string() } + +pub async fn export_knowledge( + State(state): State>, + Path(session_id): Path, + Query(params): Query, +) -> (StatusCode, [(axum::http::header::HeaderName, &'static str); 1], String) { + let kg = state.knowledge_graph.read().await; + let export = GraphExport::new( + session_id.clone(), + kg.all_entities(&session_id), + kg.all_relationships(&session_id), + ); + match params.format.as_str() { + "dot" => ( + StatusCode::OK, + [(axum::http::header::CONTENT_TYPE, "text/vnd.graphviz")], + to_dot(&export), + ), + _ => ( + StatusCode::OK, + [(axum::http::header::CONTENT_TYPE, "application/json")], + serde_json::to_string(&export).unwrap_or_default(), + ), + } +} + +#[cfg(test)] +mod tests { + use axum_test::TestServer; + use std::collections::HashMap; + use std::sync::Arc; + use tokio::sync::RwLock; + + use crate::knowledge::graph::KnowledgeGraph; + use crate::knowledge::types::{Entity, KnowledgeJob, Relationship}; + use crate::server::{AppState, build_router}; + + fn make_state_with_graph(kg: Arc>) -> Arc { + use crate::assembler::ContextAssembler; + use crate::core::{InMemoryCoreMemoryStore, InMemoryStore, InMemoryVectorStore, + OpenAITokenCounter, RandomEmbeddingProvider}; + use crate::metrics::AppMetrics; + use crate::worker::embedding_job_channel; + + let short_term_memory = Arc::new(InMemoryStore::default()); + let vector_store = Arc::new(InMemoryVectorStore::default()); + let embedding_provider = Arc::new(RandomEmbeddingProvider); + let token_counter = Arc::new(OpenAITokenCounter::new().unwrap()); + let core_memory_store = Arc::new(InMemoryCoreMemoryStore::default()); + let metrics = Arc::new(AppMetrics::new().unwrap()); + let context_assembler = Arc::new(ContextAssembler::new( + short_term_memory.clone(), vector_store.clone(), + embedding_provider.clone(), token_counter.clone(), core_memory_store.clone(), + )); + let (embedding_job_sender, mut rx) = embedding_job_channel(16); + tokio::spawn(async move { while rx.recv().await.is_some() {} }); + let (knowledge_job_sender, mut krx) = tokio::sync::mpsc::channel::(16); + tokio::spawn(async move { while krx.recv().await.is_some() {} }); + + Arc::new(AppState { + short_term_memory, + vector_store, + embedding_provider, + token_counter, + core_memory_store, + context_assembler, + metrics, + embedding_job_sender, + short_term_count: 20, + raft: None, + node_id: 0, + peer_http_addrs: std::collections::HashMap::new(), + raft_addr: None, + raft_advertise_addr: None, + cluster_peers: vec![], + knowledge_graph: kg, + knowledge_job_sender, + }) + } + + #[tokio::test] + async fn get_knowledge_returns_empty_for_new_session() { + let kg = Arc::new(RwLock::new(KnowledgeGraph::new())); + let server = TestServer::new(build_router(make_state_with_graph(kg))).unwrap(); + + let resp = server.get("/sessions/s1/knowledge").await; + resp.assert_status_ok(); + let body: serde_json::Value = resp.json(); + assert!(body["entities"].as_array().unwrap().is_empty()); + assert!(body["edges"].as_array().unwrap().is_empty()); + } + + #[tokio::test] + async fn get_knowledge_returns_graph_state() { + let kg = Arc::new(RwLock::new(KnowledgeGraph::new())); + { + let mut graph = kg.write().await; + graph.apply_extraction("s1", "m1", + vec![ + Entity { name: "Alice".into(), entity_type: "Person".into(), attributes: HashMap::new() }, + Entity { name: "OpenAI".into(), entity_type: "Organization".into(), attributes: HashMap::new() }, + ], + vec![Relationship { from: "Alice".into(), to: "OpenAI".into(), relationship_type: "works_at".into() }], + ); + } + let server = TestServer::new(build_router(make_state_with_graph(kg))).unwrap(); + + let resp = server.get("/sessions/s1/knowledge").await; + resp.assert_status_ok(); + let body: serde_json::Value = resp.json(); + assert_eq!(body["entities"].as_array().unwrap().len(), 2); + assert_eq!(body["edges"].as_array().unwrap().len(), 1); + } + + #[tokio::test] + async fn get_related_returns_connections() { + let kg = Arc::new(RwLock::new(KnowledgeGraph::new())); + { + let mut graph = kg.write().await; + graph.apply_extraction("s1", "m1", + vec![ + Entity { name: "Alice".into(), entity_type: "Person".into(), attributes: HashMap::new() }, + Entity { name: "OpenAI".into(), entity_type: "Organization".into(), attributes: HashMap::new() }, + ], + vec![Relationship { from: "Alice".into(), to: "OpenAI".into(), relationship_type: "works_at".into() }], + ); + } + let server = TestServer::new(build_router(make_state_with_graph(kg))).unwrap(); + + let resp = server.get("/sessions/s1/knowledge/entities/Alice").await; + resp.assert_status_ok(); + let body: serde_json::Value = resp.json(); + let related = body["related"].as_array().unwrap(); + assert!(!related.is_empty()); + assert_eq!(related[0]["name"], "OpenAI"); + assert_eq!(related[0]["relationship_type"], "works_at"); + } + + #[tokio::test] + async fn get_related_returns_404_for_unknown_entity() { + let kg = Arc::new(RwLock::new(KnowledgeGraph::new())); + let server = TestServer::new(build_router(make_state_with_graph(kg))).unwrap(); + + let resp = server.get("/sessions/s1/knowledge/entities/nobody").await; + resp.assert_status(axum::http::StatusCode::NOT_FOUND); + } + + #[tokio::test] + async fn export_dot_returns_dot_string() { + let kg = Arc::new(RwLock::new(KnowledgeGraph::new())); + { + let mut graph = kg.write().await; + graph.apply_extraction("s1", "m1", + vec![Entity { name: "Alice".into(), entity_type: "Person".into(), attributes: HashMap::new() }], + vec![], + ); + } + let server = TestServer::new(build_router(make_state_with_graph(kg))).unwrap(); + + let resp = server.get("/sessions/s1/knowledge/export?format=dot").await; + resp.assert_status_ok(); + let body = resp.text(); + assert!(body.contains("digraph knowledge")); + assert!(body.contains("Alice")); + } +} diff --git a/src/knowledge/mod.rs b/src/knowledge/mod.rs new file mode 100644 index 0000000..5ec1622 --- /dev/null +++ b/src/knowledge/mod.rs @@ -0,0 +1,8 @@ +pub mod extractor; +pub mod export; +pub mod graph; +pub mod handler; +pub mod types; +pub mod worker; + +pub use types::{Entity, ExtractionResult, KnowledgeJob, Relationship}; diff --git a/src/knowledge/types.rs b/src/knowledge/types.rs new file mode 100644 index 0000000..aa88b48 --- /dev/null +++ b/src/knowledge/types.rs @@ -0,0 +1,89 @@ +use std::collections::HashMap; +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +pub struct Entity { + pub name: String, + pub entity_type: String, + #[serde(default)] + pub attributes: HashMap, +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +pub struct Relationship { + pub from: String, + pub to: String, + pub relationship_type: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ExtractionResult { + pub entities: Vec, + pub relationships: Vec, +} + +#[derive(Debug, Clone)] +pub struct KnowledgeJob { + pub session_id: String, + pub message_id: String, + pub text: String, +} + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::HashMap; + + #[test] + fn entity_round_trips_with_attributes() { + let entity = Entity { + name: "Alice".to_string(), + entity_type: "Person".to_string(), + attributes: [("role".to_string(), "engineer".to_string())].into(), + }; + let json = serde_json::to_string(&entity).unwrap(); + let back: Entity = serde_json::from_str(&json).unwrap(); + assert_eq!(back, entity); + } + + #[test] + fn entity_empty_attributes_round_trips() { + let entity = Entity { + name: "OpenAI".to_string(), + entity_type: "Organization".to_string(), + attributes: HashMap::new(), + }; + let json = serde_json::to_string(&entity).unwrap(); + let back: Entity = serde_json::from_str(&json).unwrap(); + assert!(back.attributes.is_empty()); + } + + #[test] + fn relationship_round_trips() { + let rel = Relationship { + from: "Alice".to_string(), + to: "OpenAI".to_string(), + relationship_type: "works_at".to_string(), + }; + let json = serde_json::to_string(&rel).unwrap(); + let back: Relationship = serde_json::from_str(&json).unwrap(); + assert_eq!(back, rel); + } + + #[test] + fn extraction_result_round_trips() { + let result = ExtractionResult { + entities: vec![ + Entity { name: "Alice".into(), entity_type: "Person".into(), attributes: HashMap::new() }, + Entity { name: "OpenAI".into(), entity_type: "Organization".into(), attributes: HashMap::new() }, + ], + relationships: vec![ + Relationship { from: "Alice".into(), to: "OpenAI".into(), relationship_type: "works_at".into() }, + ], + }; + let json = serde_json::to_string(&result).unwrap(); + let back: ExtractionResult = serde_json::from_str(&json).unwrap(); + assert_eq!(back.entities.len(), 2); + assert_eq!(back.relationships[0].relationship_type, "works_at"); + } +} diff --git a/src/knowledge/worker.rs b/src/knowledge/worker.rs new file mode 100644 index 0000000..2041321 --- /dev/null +++ b/src/knowledge/worker.rs @@ -0,0 +1,251 @@ +use std::sync::Arc; +use tokio::sync::{Mutex, RwLock, mpsc}; +use tokio::task::JoinHandle; + +use crate::knowledge::extractor::KnowledgeExtractor; +use crate::knowledge::graph::KnowledgeGraph; +use crate::knowledge::types::KnowledgeJob; +use crate::metrics::AppMetrics; +use crate::raft::types::RaftHandle; + +pub fn knowledge_job_channel(capacity: usize) -> (mpsc::Sender, mpsc::Receiver) { + mpsc::channel(capacity.max(1)) +} + +pub fn spawn_knowledge_workers( + extractor: Arc, + raft: Option>, + node_id: u64, + knowledge_graph: Arc>, + metrics: Arc, + receiver: mpsc::Receiver, + worker_count: usize, +) -> Vec> { + let shared_receiver = Arc::new(Mutex::new(receiver)); + + (0..worker_count.max(1)) + .map(|_| { + let extractor = extractor.clone(); + let raft = raft.clone(); + let knowledge_graph = knowledge_graph.clone(); + let metrics = metrics.clone(); + let shared_receiver = shared_receiver.clone(); + + tokio::spawn(async move { + worker_loop(extractor, raft, node_id, knowledge_graph, metrics, shared_receiver).await; + }) + }) + .collect() +} + +async fn worker_loop( + extractor: Arc, + raft: Option>, + node_id: u64, + knowledge_graph: Arc>, + metrics: Arc, + receiver: Arc>>, +) { + loop { + let (job, queue_size) = { + let mut rx = receiver.lock().await; + let job = rx.recv().await; + let queue_size = rx.len(); + (job, queue_size) + }; + + metrics.set_knowledge_queue_size(queue_size); + + let Some(job) = job else { break }; + + process_knowledge_job(job, &extractor, &raft, node_id, &knowledge_graph, &metrics).await; + } +} + +#[tracing::instrument( + skip_all, + fields(session_id = %job.session_id, message_id = %job.message_id) +)] +async fn process_knowledge_job( + job: KnowledgeJob, + extractor: &Arc, + raft: &Option>, + node_id: u64, + knowledge_graph: &Arc>, + metrics: &AppMetrics, +) { + // Leader only extraction: only the current Raft leader calls the extractor. + // Followers skip because they receive AddKnowledge via Raft replication instead. + // In standalone mode (raft is None), always extract and apply directly. + // + // Leader change safety: if this node loses leadership while the HTTP request + // is in-flight, client_write() will be rejected by Raft. The result is + // discarded. This check only avoids spending tokens on a write that will + // almost certainly be rejected anyway. + if let Some(raft) = raft { + let current_leader = raft.metrics().borrow().current_leader; + if current_leader != Some(node_id) { + tracing::debug!("skipping knowledge extraction: not leader"); + return; + } + } + + // Dedup: skip if already processed (guards against replayed jobs). + if knowledge_graph.read().await.is_processed(&job.session_id, &job.message_id) { + tracing::debug!("skipping knowledge extraction: already processed"); + return; + } + + let timer = metrics.start_knowledge_extraction_timer(); + + let result = match extractor.extract(&job.text).await { + Ok(r) => r, + Err(e) => { + drop(timer); + tracing::error!(error = %e, "knowledge extraction failed"); + return; + } + }; + + let entity_count = result.entities.len() as u64; + let relationship_count = result.relationships.len() as u64; + drop(timer); + + metrics.increment_knowledge_entities(entity_count); + metrics.increment_knowledge_relationships(relationship_count); + + tracing::info!(entities = entity_count, relationships = relationship_count, "extraction complete"); + + let cmd = crate::raft::types::MemoryCommand::AddKnowledge { + session_id: job.session_id.clone(), + message_id: job.message_id.clone(), + entities: result.entities, + relationships: result.relationships, + }; + + match raft { + Some(raft) => { + if let Err(e) = raft.client_write(cmd).await { + tracing::error!(error = %e, "failed to submit AddKnowledge via Raft"); + } + } + None => { + if let crate::raft::types::MemoryCommand::AddKnowledge { + session_id, message_id, entities, relationships, + } = cmd + { + knowledge_graph + .write() + .await + .apply_extraction(&session_id, &message_id, entities, relationships); + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::knowledge::extractor::{ExtractError, KnowledgeExtractor}; + use crate::knowledge::graph::KnowledgeGraph; + use crate::knowledge::types::{Entity, ExtractionResult, KnowledgeJob, Relationship}; + use crate::metrics::AppMetrics; + use async_trait::async_trait; + use std::collections::HashMap; + use std::sync::Arc; + use tokio::sync::{RwLock, mpsc}; + + struct MockExtractor { + result: ExtractionResult, + call_count: Arc>, + } + + #[async_trait] + impl KnowledgeExtractor for MockExtractor { + async fn extract(&self, _text: &str) -> Result { + *self.call_count.lock().await += 1; + Ok(self.result.clone()) + } + } + + struct FailingExtractor; + + #[async_trait] + impl KnowledgeExtractor for FailingExtractor { + async fn extract(&self, _text: &str) -> Result { + Err(ExtractError::Api("service unavailable".to_string())) + } + } + + fn make_result() -> ExtractionResult { + ExtractionResult { + entities: vec![ + Entity { name: "Alice".into(), entity_type: "Person".into(), attributes: HashMap::new() }, + Entity { name: "OpenAI".into(), entity_type: "Organization".into(), attributes: HashMap::new() }, + ], + relationships: vec![ + Relationship { from: "Alice".into(), to: "OpenAI".into(), relationship_type: "works_at".into() }, + ], + } + } + + #[tokio::test] + async fn standalone_mode_applies_directly_to_graph() { + let call_count = Arc::new(tokio::sync::Mutex::new(0u32)); + let extractor: Arc = Arc::new(MockExtractor { + result: make_result(), + call_count: call_count.clone(), + }); + let kg = Arc::new(RwLock::new(KnowledgeGraph::new())); + let metrics = Arc::new(AppMetrics::new().unwrap()); + let (tx, rx) = mpsc::channel(10); + + spawn_knowledge_workers(extractor, None, 0, kg.clone(), metrics, rx, 1); + + tx.send(KnowledgeJob { session_id: "s1".into(), message_id: "m1".into(), text: "Alice works at OpenAI".into() }) + .await.unwrap(); + tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; + + assert_eq!(*call_count.lock().await, 1); + let graph = kg.read().await; + assert_eq!(graph.all_entities("s1").len(), 2); + } + + #[tokio::test] + async fn dedup_skips_extractor_if_already_processed() { + let call_count = Arc::new(tokio::sync::Mutex::new(0u32)); + let extractor: Arc = Arc::new(MockExtractor { + result: make_result(), + call_count: call_count.clone(), + }); + let kg = Arc::new(RwLock::new(KnowledgeGraph::new())); + let metrics = Arc::new(AppMetrics::new().unwrap()); + let (tx, rx) = mpsc::channel(10); + + spawn_knowledge_workers(extractor, None, 0, kg.clone(), metrics, rx, 1); + + let job = KnowledgeJob { session_id: "s1".into(), message_id: "m1".into(), text: "test".into() }; + tx.send(job.clone()).await.unwrap(); + tokio::time::sleep(tokio::time::Duration::from_millis(50)).await; + tx.send(job).await.unwrap(); + tokio::time::sleep(tokio::time::Duration::from_millis(50)).await; + + assert_eq!(*call_count.lock().await, 1, "extractor should only be called once"); + } + + #[tokio::test] + async fn extractor_failure_does_not_panic() { + let extractor: Arc = Arc::new(FailingExtractor); + let kg = Arc::new(RwLock::new(KnowledgeGraph::new())); + let metrics = Arc::new(AppMetrics::new().unwrap()); + let (tx, rx) = mpsc::channel(10); + + spawn_knowledge_workers(extractor, None, 0, kg.clone(), metrics, rx, 1); + + tx.send(KnowledgeJob { session_id: "s1".into(), message_id: "m1".into(), text: "test".into() }) + .await.unwrap(); + tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; + + assert!(kg.read().await.all_entities("s1").is_empty()); + } +} diff --git a/src/lib.rs b/src/lib.rs index 21ec638..1fe24b4 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -4,6 +4,7 @@ pub mod cluster; pub mod config; pub mod core; pub mod embedding; +pub mod knowledge; pub mod logging; pub mod metrics; pub mod models; diff --git a/src/metrics.rs b/src/metrics.rs index 9435b57..5e9e323 100644 --- a/src/metrics.rs +++ b/src/metrics.rs @@ -18,6 +18,10 @@ pub struct AppMetrics { pub raft_commit_index: IntGauge, pub raft_is_leader: IntGauge, pub raft_leader_changes_total: IntCounter, + knowledge_extraction_duration_seconds: HistogramVec, + knowledge_entities_extracted_total: IntCounter, + knowledge_relationships_extracted_total: IntCounter, + knowledge_queue_size: IntGauge, } impl AppMetrics { @@ -96,6 +100,33 @@ impl AppMetrics { ))?; registry.register(Box::new(raft_leader_changes_total.clone()))?; + let knowledge_extraction_duration_seconds = HistogramVec::new( + HistogramOpts::new( + "knowledge_extraction_duration_seconds", + "Duration of knowledge extraction calls in seconds.", + ), + &["model"], + )?; + registry.register(Box::new(knowledge_extraction_duration_seconds.clone()))?; + + let knowledge_entities_extracted_total = IntCounter::with_opts(Opts::new( + "knowledge_entities_extracted_total", + "Total entities extracted from messages.", + ))?; + registry.register(Box::new(knowledge_entities_extracted_total.clone()))?; + + let knowledge_relationships_extracted_total = IntCounter::with_opts(Opts::new( + "knowledge_relationships_extracted_total", + "Total relationships extracted from messages.", + ))?; + registry.register(Box::new(knowledge_relationships_extracted_total.clone()))?; + + let knowledge_queue_size = IntGauge::with_opts(Opts::new( + "knowledge_queue_size", + "Current number of pending knowledge extraction jobs.", + ))?; + registry.register(Box::new(knowledge_queue_size.clone()))?; + Ok(Self { registry, messages_added_total, @@ -108,6 +139,10 @@ impl AppMetrics { raft_commit_index, raft_is_leader, raft_leader_changes_total, + knowledge_extraction_duration_seconds, + knowledge_entities_extracted_total, + knowledge_relationships_extracted_total, + knowledge_queue_size, }) } @@ -157,6 +192,22 @@ impl AppMetrics { self.embedding_queue_size.get() } + pub fn start_knowledge_extraction_timer(&self) -> HistogramTimer { + self.knowledge_extraction_duration_seconds.with_label_values(&["gpt-4o-mini"]).start_timer() + } + + pub fn increment_knowledge_entities(&self, count: u64) { + self.knowledge_entities_extracted_total.inc_by(count); + } + + pub fn increment_knowledge_relationships(&self, count: u64) { + self.knowledge_relationships_extracted_total.inc_by(count); + } + + pub fn set_knowledge_queue_size(&self, size: usize) { + self.knowledge_queue_size.set(size as i64); + } + pub fn render(&self) -> Result { let mut buffer = Vec::new(); let encoder = TextEncoder::new(); diff --git a/src/raft/state_machine.rs b/src/raft/state_machine.rs index edeaba1..6c5fcd8 100644 --- a/src/raft/state_machine.rs +++ b/src/raft/state_machine.rs @@ -1,7 +1,7 @@ use std::io; use std::io::Cursor; use std::sync::Arc; -use tokio::sync::{mpsc, Mutex}; +use tokio::sync::{mpsc, Mutex, RwLock}; use openraft::{ BasicNode, Entry, EntryPayload, ErrorSubject, ErrorVerb, LogId, Snapshot, SnapshotMeta, StorageError, StoredMembership, RaftSnapshotBuilder, @@ -9,6 +9,8 @@ use openraft::{ }; use crate::core::{CoreMemoryStore, ShortTermMemory}; +use crate::knowledge::graph::KnowledgeGraph; +use crate::knowledge::types::KnowledgeJob; use crate::models::{EmbeddingStatus, Message}; use crate::raft::types::{CommandResponse, MemoryCommand, TypeConfig}; use crate::worker::EmbeddingJob; @@ -23,6 +25,8 @@ struct SmInner { short_term: Arc, core_memory: Arc, embedding_tx: mpsc::Sender, + knowledge_graph: Arc>, + knowledge_tx: mpsc::Sender, } impl EngStateMachineStore { @@ -33,6 +37,8 @@ impl EngStateMachineStore { // the LanceDB handle and handles deletes via the EmbeddingJob channel. _vector_store: Arc, embedding_tx: mpsc::Sender, + knowledge_graph: Arc>, + knowledge_tx: mpsc::Sender, ) -> Self { Self { inner: Arc::new(Mutex::new(SmInner { @@ -41,6 +47,8 @@ impl EngStateMachineStore { short_term, core_memory, embedding_tx, + knowledge_graph, + knowledge_tx, })), } } @@ -64,9 +72,15 @@ impl RaftStateMachine for EngStateMachineStore { I::IntoIter: Send, { // Clone Arcs once so the lock is not held across async apply_cmd calls. - let (short_term, core_memory, embedding_tx) = { + let (short_term, core_memory, embedding_tx, knowledge_graph, knowledge_tx) = { let inner = self.inner.lock().await; - (inner.short_term.clone(), inner.core_memory.clone(), inner.embedding_tx.clone()) + ( + inner.short_term.clone(), + inner.core_memory.clone(), + inner.embedding_tx.clone(), + inner.knowledge_graph.clone(), + inner.knowledge_tx.clone(), + ) }; let mut responses = Vec::new(); @@ -79,7 +93,7 @@ impl RaftStateMachine for EngStateMachineStore { last_membership = Some(StoredMembership::new(Some(entry.log_id.clone()), mem.clone())); } if let EntryPayload::Normal(cmd) = entry.payload { - apply_cmd(cmd, &short_term, &core_memory, &embedding_tx).await; + apply_cmd(cmd, &short_term, &core_memory, &embedding_tx, &knowledge_graph, &knowledge_tx).await; } responses.push(CommandResponse::default()); } @@ -149,6 +163,8 @@ async fn apply_cmd( short_term: &Arc, core_memory: &Arc, embedding_tx: &mpsc::Sender, + knowledge_graph: &Arc>, + knowledge_tx: &mpsc::Sender, ) { match cmd { MemoryCommand::AddMessage { session_id, message } => { @@ -165,7 +181,14 @@ async fn apply_cmd( tracing::error!(error = %e, session_id = %session_id, "failed to add message to short-term store"); } // Drop the job if the channel is full. embedding is eventually consistent. - let _ = embedding_tx.try_send(EmbeddingJob::new(session_id, message.id, message.content)); + let _ = embedding_tx.try_send(EmbeddingJob::new(session_id.clone(), message.id.clone(), message.content.clone())); + // Enqueue knowledge extraction. The worker checks whether this node is the + // leader before calling the extractor, so only the leader extracts. + let _ = knowledge_tx.try_send(KnowledgeJob { + session_id, + message_id: message.id, + text: message.content, + }); } MemoryCommand::AddFact { session_id, fact } => { if let Err(e) = core_memory.add_fact(&session_id, &fact).await { @@ -180,7 +203,11 @@ async fn apply_cmd( tracing::error!(error = %e, session_id = %session_id, "failed to delete session from core memory"); } // Signal embedding worker to delete from local LanceDB. - let _ = embedding_tx.try_send(EmbeddingJob::DeleteSession { session_id }); + let _ = embedding_tx.try_send(EmbeddingJob::DeleteSession { session_id: session_id.clone() }); + knowledge_graph.write().await.delete_session(&session_id); + } + MemoryCommand::AddKnowledge { session_id, message_id, entities, relationships } => { + knowledge_graph.write().await.apply_extraction(&session_id, &message_id, entities, relationships); } MemoryCommand::NoOp => {} } @@ -190,17 +217,35 @@ async fn apply_cmd( mod tests { use super::*; use crate::core::{InMemoryCoreMemoryStore, InMemoryStore, InMemoryVectorStore}; + use crate::knowledge::graph::KnowledgeGraph; + use crate::knowledge::types::{Entity, KnowledgeJob, Relationship}; use crate::raft::types::MessagePayload; + use std::collections::HashMap; use std::sync::Arc; - use tokio::sync::mpsc; + use tokio::sync::{RwLock, mpsc}; - fn make_sm() -> (EngStateMachineStore, Arc, mpsc::Receiver) { + fn make_sm() -> ( + EngStateMachineStore, + Arc, + mpsc::Receiver, + mpsc::Receiver, + Arc>, + ) { let short_term = Arc::new(InMemoryStore::default()); let core_memory = Arc::new(InMemoryCoreMemoryStore::default()); let vector_store = Arc::new(InMemoryVectorStore::default()); - let (tx, rx) = mpsc::channel(10); - let sm = EngStateMachineStore::new(short_term.clone(), core_memory, vector_store, tx); - (sm, short_term, rx) + let (embed_tx, embed_rx) = mpsc::channel(10); + let (know_tx, know_rx) = mpsc::channel(10); + let kg = Arc::new(RwLock::new(KnowledgeGraph::new())); + let sm = EngStateMachineStore::new( + short_term.clone(), + core_memory, + vector_store as Arc, + embed_tx, + kg.clone(), + know_tx, + ); + (sm, short_term, embed_rx, know_rx, kg) } fn make_entry(index: u64, cmd: MemoryCommand) -> openraft::Entry { @@ -212,7 +257,7 @@ mod tests { #[tokio::test] async fn add_message_writes_to_short_term() { - let (mut sm, short_term, _rx) = make_sm(); + let (mut sm, short_term, _embed, _know, _kg) = make_sm(); sm.apply(vec![make_entry( 0, MemoryCommand::AddMessage { @@ -234,7 +279,7 @@ mod tests { #[tokio::test] async fn add_message_enqueues_embedding_job() { - let (mut sm, _st, mut rx) = make_sm(); + let (mut sm, _st, mut embed_rx, _know, _kg) = make_sm(); sm.apply(vec![make_entry( 0, MemoryCommand::AddMessage { @@ -249,13 +294,13 @@ mod tests { )]) .await .unwrap(); - let job = rx.try_recv().expect("embedding job should be enqueued"); + let job = embed_rx.try_recv().expect("embedding job should be enqueued"); assert!(matches!(job, EmbeddingJob::Embed { text, .. } if text == "embed me")); } #[tokio::test] async fn delete_session_clears_redis_and_enqueues_lancedb_delete() { - let (mut sm, short_term, mut rx) = make_sm(); + let (mut sm, short_term, mut embed_rx, _know, _kg) = make_sm(); sm.apply(vec![ make_entry( 0, @@ -275,16 +320,68 @@ mod tests { .unwrap(); let msgs = short_term.get_recent("s2", 10).await.unwrap(); assert_eq!(msgs.len(), 0); - let _ = rx.try_recv(); // drain the Embed job from AddMessage - let del_job = rx.try_recv().expect("delete job should be enqueued"); + let _ = embed_rx.try_recv(); // drain the Embed job from AddMessage + let del_job = embed_rx.try_recv().expect("delete job should be enqueued"); assert!(matches!(del_job, EmbeddingJob::DeleteSession { session_id } if session_id == "s2")); } #[tokio::test] async fn noop_command_is_ignored() { - let (mut sm, short_term, _rx) = make_sm(); + let (mut sm, short_term, _embed, _know, _kg) = make_sm(); sm.apply(vec![make_entry(0, MemoryCommand::NoOp)]).await.unwrap(); let msgs = short_term.get_recent("any", 10).await.unwrap(); assert_eq!(msgs.len(), 0); } + + #[tokio::test] + async fn add_message_enqueues_knowledge_job() { + let (mut sm, _st, _embed, mut know_rx, _kg) = make_sm(); + sm.apply(vec![make_entry(0, MemoryCommand::AddMessage { + session_id: "s1".into(), + message: MessagePayload { + id: "m1".into(), role: "user".into(), + content: "Alice works at OpenAI".into(), + timestamp: chrono::Utc::now(), + }, + })]).await.unwrap(); + let job = know_rx.try_recv().expect("knowledge job should be enqueued"); + assert_eq!(job.session_id, "s1"); + assert_eq!(job.message_id, "m1"); + assert_eq!(job.text, "Alice works at OpenAI"); + } + + #[tokio::test] + async fn add_knowledge_updates_graph() { + let (mut sm, _st, _embed, _know, kg) = make_sm(); + sm.apply(vec![make_entry(0, MemoryCommand::AddKnowledge { + session_id: "s1".into(), + message_id: "m1".into(), + entities: vec![ + Entity { name: "Alice".into(), entity_type: "Person".into(), attributes: HashMap::new() }, + Entity { name: "OpenAI".into(), entity_type: "Organization".into(), attributes: HashMap::new() }, + ], + relationships: vec![ + Relationship { from: "Alice".into(), to: "OpenAI".into(), relationship_type: "works_at".into() }, + ], + })]).await.unwrap(); + + let kg = kg.read().await; + assert_eq!(kg.all_entities("s1").len(), 2); + let related = kg.get_related("s1", "OpenAI"); + assert!(related.iter().any(|r| r.name == "Alice")); + } + + #[tokio::test] + async fn delete_session_clears_knowledge_graph() { + let (mut sm, _st, _embed, _know, kg) = make_sm(); + sm.apply(vec![make_entry(0, MemoryCommand::AddKnowledge { + session_id: "s1".into(), message_id: "m1".into(), + entities: vec![Entity { name: "Alice".into(), entity_type: "Person".into(), attributes: HashMap::new() }], + relationships: vec![], + })]).await.unwrap(); + sm.apply(vec![make_entry(1, MemoryCommand::DeleteSession { session_id: "s1".into() })]).await.unwrap(); + + let kg = kg.read().await; + assert!(kg.all_entities("s1").is_empty()); + } } diff --git a/src/raft/types.rs b/src/raft/types.rs index e8731c4..5982b44 100644 --- a/src/raft/types.rs +++ b/src/raft/types.rs @@ -34,6 +34,15 @@ pub enum MemoryCommand { /// Deletes all Redis short-term + core memory for a session on all nodes. /// Also signals each node's embedding worker to delete from local LanceDB. DeleteSession { session_id: String }, + /// Extract and store knowledge from a message. Idempotent by (session_id, message_id). + /// Only submitted by the leader's knowledge worker. All nodes receive this command + /// via Raft replication and apply it to their local KnowledgeGraph. + AddKnowledge { + session_id: String, + message_id: String, + entities: Vec, + relationships: Vec, + }, /// No-op placeholder. Applied by the state machine without side effects. /// Reserved for future cluster operations (e.g., leadership probes). NoOp, @@ -60,6 +69,34 @@ pub type RaftHandle = openraft::Raft; mod tests { use super::*; + #[test] + fn add_knowledge_command_round_trips() { + use crate::knowledge::types::{Entity, Relationship}; + use std::collections::HashMap; + + let cmd = MemoryCommand::AddKnowledge { + session_id: "s1".to_string(), + message_id: "m1".to_string(), + entities: vec![ + Entity { name: "Alice".into(), entity_type: "Person".into(), attributes: HashMap::new() }, + ], + relationships: vec![ + Relationship { from: "Alice".into(), to: "OpenAI".into(), relationship_type: "works_at".into() }, + ], + }; + let json = serde_json::to_string(&cmd).unwrap(); + let back: MemoryCommand = serde_json::from_str(&json).unwrap(); + match back { + MemoryCommand::AddKnowledge { session_id, message_id, entities, relationships } => { + assert_eq!(session_id, "s1"); + assert_eq!(message_id, "m1"); + assert_eq!(entities[0].name, "Alice"); + assert_eq!(relationships[0].relationship_type, "works_at"); + } + _ => panic!("wrong variant"), + } + } + #[test] fn memory_command_serializes_round_trip() { let cmd = MemoryCommand::AddMessage { diff --git a/src/server.rs b/src/server.rs index c6da93f..ff9a009 100644 --- a/src/server.rs +++ b/src/server.rs @@ -10,6 +10,7 @@ use axum::{ response::IntoResponse, routing::{delete, get, post, put}, }; +use crate::knowledge::handler::{export_knowledge, find_path, get_knowledge, get_related}; use axum_prometheus::{PrometheusMetricLayer, PrometheusMetricLayerBuilder}; use axum_prometheus::metrics_exporter_prometheus::PrometheusHandle; use chrono::Utc; @@ -123,6 +124,10 @@ pub struct AppState { pub raft_advertise_addr: Option, /// gRPC addresses of peer nodes, used to build the initial cluster membership. pub cluster_peers: Vec, + /// Per-session in-memory knowledge graph, shared with the Raft state machine. + pub knowledge_graph: Arc>, + /// Channel for sending knowledge extraction jobs to the worker pool. + pub knowledge_job_sender: tokio::sync::mpsc::Sender, } @@ -215,6 +220,10 @@ pub fn build_router(state: Arc) -> Router { .route("/sessions/{session_id}/context", get(get_context)) .route("/sessions/{session_id}/search", post(search_session)) .route("/sessions/{session_id}/core-memory", put(put_core_memory)) + .route("/sessions/{session_id}/knowledge", get(get_knowledge)) + .route("/sessions/{session_id}/knowledge/entities/{entity_name}", get(get_related)) + .route("/sessions/{session_id}/knowledge/path", get(find_path)) + .route("/sessions/{session_id}/knowledge/export", get(export_knowledge)) .route("/cluster", get(crate::cluster::get_cluster_status)) .route("/cluster/init", post(crate::cluster::init_cluster)) .route("/cluster/add-learner", post(crate::cluster::add_learner)) @@ -760,6 +769,14 @@ mod tests { raft_addr: None, raft_advertise_addr: None, cluster_peers: vec![], + knowledge_graph: Arc::new(tokio::sync::RwLock::new( + crate::knowledge::graph::KnowledgeGraph::new(), + )), + knowledge_job_sender: { + let (tx, mut rx) = tokio::sync::mpsc::channel::(16); + tokio::spawn(async move { while rx.recv().await.is_some() {} }); + tx + }, }) } @@ -770,6 +787,14 @@ mod tests { let _ = s.node_id; } + #[tokio::test] + async fn appstate_has_knowledge_fields() { + // If AppState is missing the new fields, this won't compile. + let state = build_test_state(); + let _ = state.knowledge_graph.read().await; + let _ = state.knowledge_job_sender.capacity(); + } + #[tokio::test] async fn health_route_returns_ok() { let server = TestServer::new(build_router(build_test_state())).unwrap(); @@ -1210,6 +1235,46 @@ mod tests { response.assert_status(StatusCode::BAD_REQUEST); } + #[tokio::test] + async fn knowledge_routes_are_registered() { + let server = TestServer::new(build_router(build_test_state())).unwrap(); + server.get("/sessions/test-session/knowledge").await.assert_status_ok(); + server + .get("/sessions/test-session/knowledge/export?format=json") + .await + .assert_status_ok(); + } + + #[tokio::test] + async fn knowledge_metrics_appear_in_prometheus_scrape() { + let state = build_test_state(); + // Observe each metric so the HistogramVec emits output (Vec types only appear + // in Prometheus text format once at least one label set has been recorded). + let timer = state.metrics.start_knowledge_extraction_timer(); + drop(timer); + state.metrics.increment_knowledge_entities(1); + state.metrics.increment_knowledge_relationships(1); + state.metrics.set_knowledge_queue_size(0); + let server = TestServer::new(build_router(state)).unwrap(); + let body = server.get("/metrics").await.text(); + assert!( + body.contains("engram_knowledge_extraction_duration_seconds"), + "missing knowledge_extraction_duration_seconds" + ); + assert!( + body.contains("engram_knowledge_entities_extracted_total"), + "missing knowledge_entities_extracted_total" + ); + assert!( + body.contains("engram_knowledge_relationships_extracted_total"), + "missing knowledge_relationships_extracted_total" + ); + assert!( + body.contains("engram_knowledge_queue_size"), + "missing knowledge_queue_size" + ); + } + #[tokio::test] #[traced_test] async fn handler_spans_are_logged_without_content_fields() { diff --git a/tests/e2e_test.rs b/tests/e2e_test.rs index fb67b7f..fdfd11c 100644 --- a/tests/e2e_test.rs +++ b/tests/e2e_test.rs @@ -128,6 +128,9 @@ async fn e2e_flow_uses_real_stores_and_background_worker() { raft_advertise_addr: None, cluster_peers: vec![], cluster_http_peers: std::collections::HashMap::new(), + knowledge_max_workers: 4, + knowledge_channel_size: 500, + knowledge_extractor: engram::config::KnowledgeExtractorType::OpenAI, }; let embedding_provider: Arc = Arc::new( OpenAIEmbedder::new_with_base_url("test-key", mock_server.uri()).unwrap(), diff --git a/tests/integration_test.rs b/tests/integration_test.rs index 3a781ba..7f014f9 100644 --- a/tests/integration_test.rs +++ b/tests/integration_test.rs @@ -91,6 +91,9 @@ async fn setup_test_app() -> TestApp { raft_advertise_addr: None, cluster_peers: vec![], cluster_http_peers: std::collections::HashMap::new(), + knowledge_max_workers: 4, + knowledge_channel_size: 500, + knowledge_extractor: engram::config::KnowledgeExtractorType::OpenAI, }; let embedding_provider: Arc = Arc::new( diff --git a/tests/raft_write_test.rs b/tests/raft_write_test.rs index e7ff4de..5854473 100644 --- a/tests/raft_write_test.rs +++ b/tests/raft_write_test.rs @@ -18,12 +18,16 @@ async fn single_node_raft_write_commits_to_state_machine() { node_id: Some(1), ..Config::default() }; + let knowledge_graph = Arc::new(tokio::sync::RwLock::new(engram::knowledge::graph::KnowledgeGraph::new())); + let (knowledge_tx, _knowledge_rx) = mpsc::channel(500); let raft = build_raft_node( &config, short_term.clone(), Arc::new(InMemoryCoreMemoryStore::default()), Arc::new(InMemoryVectorStore::default()), tx, + knowledge_graph, + knowledge_tx, ) .await .unwrap(); diff --git a/tests/token_efficiency.rs b/tests/token_efficiency.rs index 157e0bc..1f734c4 100644 --- a/tests/token_efficiency.rs +++ b/tests/token_efficiency.rs @@ -58,6 +58,8 @@ fn build_test_state() -> Arc { core_memory_store.clone(), )); + let (knowledge_job_sender, mut krx) = tokio::sync::mpsc::channel::(16); + tokio::spawn(async move { while krx.recv().await.is_some() {} }); Arc::new(AppState { short_term_memory, vector_store, @@ -74,6 +76,10 @@ fn build_test_state() -> Arc { raft_addr: None, raft_advertise_addr: None, cluster_peers: vec![], + knowledge_graph: Arc::new(tokio::sync::RwLock::new( + engram::knowledge::graph::KnowledgeGraph::new(), + )), + knowledge_job_sender, }) }