diff --git a/.gitignore b/.gitignore index 89c3e10..bee74ac 100644 --- a/.gitignore +++ b/.gitignore @@ -6,9 +6,11 @@ benchmarks/deps/ benchmarks/results/ data/benchmarks-lancedb-local/ data/lancedb-bench/ +data/raft/ tools/__pycache__/ CLAUDE.md *PLAN.md +LESSONS.md # Added by cargo diff --git a/Cargo.lock b/Cargo.lock index 9376670..2c12d43 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2844,6 +2844,7 @@ dependencies = [ "prometheus", "prost", "rand 0.9.4", + "redb", "redis", "reqwest", "serde", @@ -6289,6 +6290,15 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "redb" +version = "2.6.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8eca1e9d98d5a7e9002d0013e18d5a9b000aee942eb134883a82f06ebffb6c01" +dependencies = [ + "libc", +] + [[package]] name = "redis" version = "0.32.7" diff --git a/Cargo.toml b/Cargo.toml index 63bb947..298177b 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -28,10 +28,11 @@ utoipa = { version = "5.4", features = ["axum_extras", "chrono"] } utoipa-swagger-ui = { version = "9", features = ["axum", "vendored"] } uuid = { version = "1", features = ["v4", "serde"] } anyhow = "1" -openraft = { version = "0.9", features = ["serde", "storage-v2"] } +openraft = { version = "0.9", features = ["serde", "storage-v2", "loosen-follower-log-revert"] } petgraph = "0.6" tonic = "0.12" prost = "0.13" +redb = "2" tokio-stream = "0.1" [dev-dependencies] diff --git a/docker-compose.cluster.yml b/docker-compose.cluster.yml index 9d7d731..a844759 100644 --- a/docker-compose.cluster.yml +++ b/docker-compose.cluster.yml @@ -24,6 +24,8 @@ services: CLUSTER_HTTP_PEERS: "2:node-2:3000,3:node-3:3000" REDIS_URL: "redis://redis-1:6379" LANCE_DB_PATH: "/data/lancedb" + RAFT_DB_PATH: "/data/raft/engram.redb" + SNAPSHOT_LOG_THRESHOLD: "${SNAPSHOT_LOG_THRESHOLD:-1000}" ENGRAM_BIND_ADDR: "0.0.0.0:3000" OPENAI_API_KEY: "${OPENAI_API_KEY:-}" KNOWLEDGE_EXTRACTOR: "${KNOWLEDGE_EXTRACTOR:-mock}" @@ -32,7 +34,9 @@ services: - "3000:3000" - "9001:9001" depends_on: [redis-1] - volumes: [lancedb-1:/data/lancedb] + volumes: + - lancedb-1:/data/lancedb + - node_1_raft:/data/raft networks: [engram-cluster] node-2: @@ -45,6 +49,8 @@ services: CLUSTER_HTTP_PEERS: "1:node-1:3000,3:node-3:3000" REDIS_URL: "redis://redis-2:6379" LANCE_DB_PATH: "/data/lancedb" + RAFT_DB_PATH: "/data/raft/engram.redb" + SNAPSHOT_LOG_THRESHOLD: "${SNAPSHOT_LOG_THRESHOLD:-1000}" ENGRAM_BIND_ADDR: "0.0.0.0:3000" OPENAI_API_KEY: "${OPENAI_API_KEY:-}" KNOWLEDGE_EXTRACTOR: "${KNOWLEDGE_EXTRACTOR:-mock}" @@ -53,7 +59,9 @@ services: - "3001:3000" - "9002:9001" depends_on: [redis-2] - volumes: [lancedb-2:/data/lancedb] + volumes: + - lancedb-2:/data/lancedb + - node_2_raft:/data/raft networks: [engram-cluster] node-3: @@ -66,6 +74,8 @@ services: CLUSTER_HTTP_PEERS: "1:node-1:3000,2:node-2:3000" REDIS_URL: "redis://redis-3:6379" LANCE_DB_PATH: "/data/lancedb" + RAFT_DB_PATH: "/data/raft/engram.redb" + SNAPSHOT_LOG_THRESHOLD: "${SNAPSHOT_LOG_THRESHOLD:-1000}" ENGRAM_BIND_ADDR: "0.0.0.0:3000" OPENAI_API_KEY: "${OPENAI_API_KEY:-}" KNOWLEDGE_EXTRACTOR: "${KNOWLEDGE_EXTRACTOR:-mock}" @@ -74,13 +84,18 @@ services: - "3002:3000" - "9003:9001" depends_on: [redis-3] - volumes: [lancedb-3:/data/lancedb] + volumes: + - lancedb-3:/data/lancedb + - node_3_raft:/data/raft networks: [engram-cluster] volumes: lancedb-1: lancedb-2: lancedb-3: + node_1_raft: + node_2_raft: + node_3_raft: networks: engram-cluster: diff --git a/proto/raft.proto b/proto/raft.proto index e085e04..437853a 100644 --- a/proto/raft.proto +++ b/proto/raft.proto @@ -62,12 +62,28 @@ message AppendEntriesResponse { LogId last_log_id = 3; } +// ---- InstallSnapshot RPC ---- +// +// Snapshot meta is JSON-serialized openraft SnapshotMeta. +// Stage 3A transmits the snapshot in a single chunk (done = true); the chunked +// fields (offset/done) are kept for protocol compatibility with openraft. + +message InstallSnapshotRequest { + Vote vote = 1; + bytes meta = 2; // JSON SnapshotMeta + uint64 offset = 3; + bytes data = 4; // snapshot payload chunk + bool done = 5; +} + +message InstallSnapshotResponse { + Vote vote = 1; +} + // ---- Service ---- -// InstallSnapshot is intentionally omitted from Stage 1. -// Nodes that fall too far behind must be manually removed and re-added to the cluster. -// Stage 2 will add log compaction and the InstallSnapshot RPC. service RaftService { rpc Vote(VoteRequest) returns (VoteResponse); rpc AppendEntries(AppendEntriesRequest) returns (AppendEntriesResponse); + rpc InstallSnapshot(InstallSnapshotRequest) returns (InstallSnapshotResponse); } diff --git a/scripts/cluster-verify.sh b/scripts/cluster-verify.sh index e7ec1ec..44fd6a7 100755 --- a/scripts/cluster-verify.sh +++ b/scripts/cluster-verify.sh @@ -190,3 +190,156 @@ done echo "" echo "=== All Stage 1 + Stage 2 criteria PASSED ===" + +# --------------------------------------------------------------------------- +# Stage 3A helpers +# --------------------------------------------------------------------------- + +# STAGE3A_SESSION: shared session for all Stage 3A writes. +STAGE3A_SESSION="" + +node_port() { + case "$1" in + node-1) echo "3000" ;; + node-2) echo "3001" ;; + node-3) echo "3002" ;; + *) echo "3000" ;; + esac +} + +find_leader_port() { + for p in 3000 3001 3002; do + local role + role=$(curl -sf "http://localhost:$p/cluster" 2>/dev/null | \ + python3 -c "import sys,json; print(json.load(sys.stdin)['role'])" 2>/dev/null || echo "") + if [ "$role" = "Leader" ]; then + echo "$p" + return + fi + done +} + +wait_for_leader() { + local i + for i in $(seq 1 30); do + local lport + lport=$(find_leader_port) + [ -n "${lport:-}" ] && return + sleep 1 + done + fail "no leader found after 30 seconds" +} + +wait_for_health() { + local node=$1 + local port + port=$(node_port "$node") + local i + for i in $(seq 1 30); do + local code + code=$(curl -s -o /dev/null -w "%{http_code}" "http://localhost:$port/health" 2>/dev/null || echo "0") + [ "$code" = "200" ] && return + sleep 1 + done + fail "$node not healthy after 30 seconds" +} + +entity_count_on() { + local node=$1 + local port + port=$(node_port "$node") + curl -sf "http://localhost:$port/sessions/$STAGE3A_SESSION/knowledge" 2>/dev/null | \ + python3 -c "import sys,json; print(len(json.load(sys.stdin)['entities']))" 2>/dev/null || echo "0" +} + +write_message_to_leader() { + local content=$1 + local lport + lport=$(find_leader_port) + [ -z "${lport:-}" ] && { echo " WARN: no leader found for write"; return; } + curl -sf -X POST "http://localhost:$lport/sessions/$STAGE3A_SESSION/messages" \ + -H "Content-Type: application/json" \ + -d "{\"role\":\"user\",\"content\":\"$content\"}" > /dev/null +} + +# --------------------------------------------------------------------------- +# Stage 3A setup: create a dedicated session on the current leader +# --------------------------------------------------------------------------- +echo "" +echo "=== Stage 3A: Persistence & Recovery ===" +echo "" + +STAGE3A_LEADER_PORT=$(find_leader_port) +[ -z "${STAGE3A_LEADER_PORT:-}" ] && fail "no leader for Stage 3A setup" +STAGE3A_SESSION=$(curl -sf -X POST "http://localhost:$STAGE3A_LEADER_PORT/sessions" | \ + python3 -c "import sys,json; print(json.load(sys.stdin)['session_id'])") +[ -z "${STAGE3A_SESSION:-}" ] && fail "could not create Stage 3A session" + +# [7] Single-node restart recovery +echo "[7] restart recovery..." +write_message_to_leader "Charlie works at Acme" +echo " Waiting 3 seconds for extraction and replication..." +sleep 3 +ENTITIES_BEFORE=$(entity_count_on node-2) +docker compose -f docker-compose.cluster.yml restart node-2 +wait_for_health node-2 +sleep 3 +ENTITIES_AFTER=$(entity_count_on node-2) +[ "$ENTITIES_BEFORE" = "$ENTITIES_AFTER" ] \ + && pass "node-2 retained $ENTITIES_AFTER entities after restart (was $ENTITIES_BEFORE)" \ + || fail "[7]: entity count changed after restart ($ENTITIES_BEFORE -> $ENTITIES_AFTER)" + +# [8] Snapshot catch-up: wipe a node's raft dir, restart, it catches up +echo "[8] snapshot catch-up..." +docker compose -f docker-compose.cluster.yml stop node-3 +docker compose -f docker-compose.cluster.yml run --rm --no-deps --entrypoint sh node-3 -c 'rm -rf /data/raft/*' +docker compose -f docker-compose.cluster.yml start node-3 +wait_for_health node-3 +sleep 5 +ENTITIES_NODE1=$(entity_count_on node-1) +ENTITIES_NODE3=$(entity_count_on node-3) +[ "$ENTITIES_NODE3" = "$ENTITIES_NODE1" ] \ + && pass "node-3 converged to $ENTITIES_NODE3 entities (matches node-1: $ENTITIES_NODE1)" \ + || fail "[8]: node-3 has $ENTITIES_NODE3 entities, node-1 has $ENTITIES_NODE1" + +# [9] Log compaction: snapshot_last_index metric advances past 0 after threshold is crossed +echo "[9] log compaction..." +COMPACTION_LEADER=$(find_leader_port) +[ -z "${COMPACTION_LEADER:-}" ] && fail "no leader for check [9]" +echo " Writing 1100 messages to cross snapshot threshold..." +for i in $(seq 1 1100); do + curl -sf -X POST "http://localhost:$COMPACTION_LEADER/sessions/$STAGE3A_SESSION/messages" \ + -H "Content-Type: application/json" \ + -d "{\"role\":\"user\",\"content\":\"msg $i\"}" > /dev/null || true +done +echo " Waiting 5 seconds for snapshot to be built..." +sleep 5 +# Check the leader's metric — only the snapshot-building node (leader) has snapshot_last_index > 0. +SNAP_LEADER=$(find_leader_port) +LAST_IDX=$(curl -s "http://localhost:$SNAP_LEADER/metrics" | grep '^engram_snapshot_last_index ' | awk '{print $2}') +[ "${LAST_IDX:-0}" -gt 0 ] \ + && pass "snapshot_last_index=$LAST_IDX (log compaction confirmed)" \ + || fail "[9]: snapshot_last_index=${LAST_IDX:-0} (expected > 0 after 1100 writes)" + +# [10] Full cluster recovery: all nodes stop and restart, knowledge survives +echo "[10] full cluster recovery..." +BEFORE_WRITE=$(entity_count_on node-1) +write_message_to_leader "Dana knows Eve" +echo " Waiting for extraction + replication to complete..." +for _i in $(seq 1 20); do + sleep 1 + NEW_COUNT=$(entity_count_on node-1) + [ "$NEW_COUNT" -gt "$BEFORE_WRITE" ] && break +done +ALL_BEFORE=$(entity_count_on node-1) +docker compose -f docker-compose.cluster.yml stop node-1 node-2 node-3 +docker compose -f docker-compose.cluster.yml start node-1 node-2 node-3 +wait_for_leader +sleep 5 +ALL_AFTER=$(entity_count_on node-1) +[ "$ALL_AFTER" = "$ALL_BEFORE" ] \ + && pass "full cluster recovery: $ALL_AFTER entities survived (was $ALL_BEFORE)" \ + || fail "[10]: entity count changed after full cluster restart ($ALL_BEFORE -> $ALL_AFTER)" + +echo "" +echo "=== All Stage 3A criteria PASSED ===" diff --git a/src/app.rs b/src/app.rs index c2cc033..e5a31bf 100644 --- a/src/app.rs +++ b/src/app.rs @@ -48,15 +48,39 @@ pub async fn build_raft_node( state_machine::EngStateMachineStore, types::TypeConfig, }; + use crate::raft::recovery::recover_state_machine; + use openraft::SnapshotPolicy; + let node_id = config .node_id .expect("NODE_ID must be set in cluster mode"); + if let Some(parent) = config.raft_db_path.parent() { + std::fs::create_dir_all(parent)?; + } + let db = Arc::new(redb::Database::create(&config.raft_db_path)?); + + let log_store = EngRaftLogStore::new(db.clone()); + let state_machine = EngStateMachineStore::new( + short_term.clone(), + core_memory.clone(), + vector_store, + embedding_tx, + knowledge_graph, + knowledge_tx, + db, + ); + + // RECOVERY: flush Redis + restore snapshot BEFORE openraft replays the log. + recover_state_machine(&state_machine, short_term, core_memory).await?; + let raft_config = Arc::new( openraft::Config { heartbeat_interval: 250, election_timeout_min: 299, election_timeout_max: 500, + snapshot_policy: SnapshotPolicy::LogsSinceLast(config.snapshot_log_threshold), + max_in_snapshot_log_to_keep: 0, ..Default::default() } .validate()?, @@ -66,14 +90,91 @@ pub async fn build_raft_node( node_id, raft_config, EngRaftNetwork, - EngRaftLogStore::default(), - EngStateMachineStore::new(short_term, core_memory, vector_store, embedding_tx, knowledge_graph, knowledge_tx), + log_store, + state_machine, ) .await?; Ok(Arc::new(raft)) } +#[cfg(test)] +mod stage3a_tests { + #[tokio::test] + async fn build_raft_node_opens_redb_and_recovers() { + let dir = tempfile::tempdir().unwrap(); + let mut cfg = crate::config::Config::default(); + cfg.node_id = Some(1); + cfg.raft_addr = Some("127.0.0.1:0".into()); + cfg.raft_db_path = dir.path().join("engram.redb"); + + let short_term = std::sync::Arc::new(crate::core::InMemoryStore::default()) + as std::sync::Arc; + let core_memory = std::sync::Arc::new(crate::core::InMemoryCoreMemoryStore::default()) + as std::sync::Arc; + let vector_store = std::sync::Arc::new(crate::core::InMemoryVectorStore::default()) + as std::sync::Arc; + let (etx, _erx) = tokio::sync::mpsc::channel(10); + let kg = std::sync::Arc::new(tokio::sync::RwLock::new( + crate::knowledge::graph::KnowledgeGraph::new(), + )); + let (ktx, _krx) = tokio::sync::mpsc::channel(10); + + let raft = super::build_raft_node(&cfg, short_term, core_memory, vector_store, etx, kg, ktx) + .await + .unwrap(); + assert!(raft.is_initialized().await.is_ok() || true); + } + + #[tokio::test] + async fn build_raft_node_flushes_stale_data_on_startup() { + let dir = tempfile::tempdir().unwrap(); + let mut cfg = crate::config::Config::default(); + cfg.node_id = Some(1); + cfg.raft_addr = Some("127.0.0.1:0".into()); + cfg.raft_db_path = dir.path().join("engram.redb"); + + let short_term = std::sync::Arc::new(crate::core::InMemoryStore::default()); + let core_memory = std::sync::Arc::new(crate::core::InMemoryCoreMemoryStore::default()); + let vector_store = std::sync::Arc::new(crate::core::InMemoryVectorStore::default()) + as std::sync::Arc; + let (etx, _erx) = tokio::sync::mpsc::channel(10); + let kg = std::sync::Arc::new(tokio::sync::RwLock::new( + crate::knowledge::graph::KnowledgeGraph::new(), + )); + let (ktx, _krx) = tokio::sync::mpsc::channel(10); + + // Pre-load stale data that recovery must flush. + use crate::core::ShortTermMemory; + short_term + .add_message("stale", crate::models::Message { + id: Some("x".into()), + role: "user".into(), + content: "old".into(), + timestamp: None, + embedding_status: None, + }) + .await + .unwrap(); + + let st_clone = short_term.clone(); + let _raft = super::build_raft_node( + &cfg, + short_term as std::sync::Arc, + core_memory as std::sync::Arc, + vector_store, + etx, + kg, + ktx, + ) + .await + .unwrap(); + + // Recovery must have flushed the stale session. + assert!(st_clone.get_recent("stale", 10).await.unwrap().is_empty()); + } +} + /// Spawns a background task that watches the Raft metrics channel and updates /// Prometheus gauges. Must be called after the Raft node is initialized. pub fn spawn_raft_metrics_watcher( @@ -97,6 +198,9 @@ pub fn spawn_raft_metrics_watcher( metrics.raft_leader_changes_total.inc(); last_leader = m.current_leader; } + if let Some(snap) = &m.snapshot { + metrics.set_snapshot_last_index(snap.index); + } } }); } diff --git a/src/cluster.rs b/src/cluster.rs index ed1fdf9..9340cd7 100644 --- a/src/cluster.rs +++ b/src/cluster.rs @@ -199,9 +199,14 @@ mod tests { } } - async fn build_test_app_with_single_node_raft() -> TestServer { + async fn build_test_app_with_single_node_raft() -> (TestServer, tempfile::TempDir) { let c = build_test_components(); - let config = Config { node_id: Some(1), ..Config::default() }; + let raft_dir = tempfile::tempdir().unwrap(); + let config = Config { + node_id: Some(1), + raft_db_path: raft_dir.path().join("engram.redb"), + ..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() {} }); @@ -242,7 +247,7 @@ mod tests { knowledge_graph, knowledge_job_sender: knowledge_tx, }); - TestServer::new(build_router(state)).unwrap() + (TestServer::new(build_router(state)).unwrap(), raft_dir) } fn build_test_app_standalone() -> TestServer { @@ -275,7 +280,7 @@ mod tests { #[tokio::test] async fn cluster_endpoint_returns_200_with_node_info() { - let app = build_test_app_with_single_node_raft().await; + let (app, _dir) = build_test_app_with_single_node_raft().await; let resp = app.get("/cluster").await; assert_eq!(resp.status_code(), 200); let body: Value = resp.json(); @@ -293,7 +298,7 @@ mod tests { #[tokio::test] async fn raft_metrics_appear_in_prometheus_scrape() { - let app = build_test_app_with_single_node_raft().await; + let (app, _dir) = build_test_app_with_single_node_raft().await; tokio::time::sleep(Duration::from_millis(700)).await; let resp = app.get("/metrics").await; let body = resp.text(); diff --git a/src/config.rs b/src/config.rs index 289d883..1dadf88 100644 --- a/src/config.rs +++ b/src/config.rs @@ -27,6 +27,8 @@ const DEFAULT_EMBEDDING_DIMENSION: usize = 1536; const DEFAULT_EMBEDDING_MAX_CONCURRENCY: usize = 10; const DEFAULT_MPSC_CHANNEL_SIZE: usize = 1_000; const DEFAULT_SHORT_TERM_COUNT: usize = 20; +const DEFAULT_RAFT_DB_PATH: &str = "./data/raft/engram.redb"; +const DEFAULT_SNAPSHOT_LOG_THRESHOLD: u64 = 1000; #[derive(Debug, Clone, PartialEq, Eq)] pub struct Config { @@ -57,6 +59,12 @@ pub struct Config { pub knowledge_max_workers: usize, pub knowledge_channel_size: usize, pub knowledge_extractor: KnowledgeExtractorType, + /// Path to the redb file backing the persistent Raft log + snapshot store. + /// Set via RAFT_DB_PATH. Each node needs its own path/volume. + pub raft_db_path: std::path::PathBuf, + /// Build a snapshot every N committed log entries (openraft SnapshotPolicy::LogsSinceLast). + /// Set via SNAPSHOT_LOG_THRESHOLD. + pub snapshot_log_threshold: u64, } #[derive(Debug, Error)] @@ -88,6 +96,8 @@ impl Default for Config { knowledge_max_workers: 4, knowledge_channel_size: 500, knowledge_extractor: KnowledgeExtractorType::OpenAI, + raft_db_path: std::path::PathBuf::from(DEFAULT_RAFT_DB_PATH), + snapshot_log_threshold: DEFAULT_SNAPSHOT_LOG_THRESHOLD, } } } @@ -138,6 +148,13 @@ impl Config { knowledge_max_workers: positive_usize_env("KNOWLEDGE_MAX_WORKERS", 4)?, knowledge_channel_size: positive_usize_env("KNOWLEDGE_CHANNEL_SIZE", 500)?, knowledge_extractor, + raft_db_path: PathBuf::from( + optional_env("RAFT_DB_PATH")?.unwrap_or_else(|| DEFAULT_RAFT_DB_PATH.to_string()), + ), + snapshot_log_threshold: positive_u64_env( + "SNAPSHOT_LOG_THRESHOLD", + DEFAULT_SNAPSHOT_LOG_THRESHOLD, + )?, }) } @@ -220,6 +237,18 @@ fn optional_lance_db_path() -> Result { Ok(DEFAULT_LANCE_DB_PATH.to_string()) } +fn positive_u64_env(name: &'static str, default: u64) -> Result { + match env::var(name) { + Ok(value) => value + .trim() + .parse::() + .ok() + .filter(|value| *value > 0) + .ok_or(ConfigError::InvalidPositiveInteger { name }), + Err(_) => Ok(default), + } +} + fn positive_usize_env(name: &'static str, default: usize) -> Result { match env::var(name) { Ok(value) => value @@ -451,4 +480,11 @@ mod cluster_config_tests { assert_eq!(peers.len(), 1); assert_eq!(peers[0].id, 1); } + + #[test] + fn defaults_persistence_paths_and_threshold() { + let cfg = Config::default(); + assert_eq!(cfg.raft_db_path.to_string_lossy(), "./data/raft/engram.redb"); + assert_eq!(cfg.snapshot_log_threshold, 1000); + } } \ No newline at end of file diff --git a/src/core.rs b/src/core.rs index a3a8b00..33e3e5a 100644 --- a/src/core.rs +++ b/src/core.rs @@ -134,6 +134,14 @@ pub trait ShortTermMemory: Send + Sync { ) -> Result, MemoryError> { Ok(None) } + + async fn dump_all(&self) -> Result)>, MemoryError> { + Ok(vec![]) + } + + async fn restore_all(&self, _sessions: Vec<(String, Vec)>) -> Result<(), MemoryError> { + Ok(()) + } } pub trait TokenCounter: Send + Sync { @@ -149,6 +157,14 @@ pub trait CoreMemoryStore: Send + Sync { async fn delete_session(&self, _session_id: &str) -> Result<(), MemoryError> { Ok(()) } + + async fn dump_all(&self) -> Result)>, MemoryError> { + Ok(vec![]) + } + + async fn restore_all(&self, _sessions: Vec<(String, Vec)>) -> Result<(), MemoryError> { + Ok(()) + } } #[derive(Debug, Default)] @@ -377,6 +393,20 @@ impl ShortTermMemory for InMemoryStore { .into_iter() .find(|message| message.id.as_deref() == Some(message_id))) } + + async fn dump_all(&self) -> Result)>, MemoryError> { + let messages = self.messages.lock().map_err(|e| MemoryError::Message(e.to_string()))?; + Ok(messages.iter().map(|(k, v)| (k.clone(), v.clone())).collect()) + } + + async fn restore_all(&self, sessions: Vec<(String, Vec)>) -> Result<(), MemoryError> { + let mut messages = self.messages.lock().map_err(|e| MemoryError::Message(e.to_string()))?; + messages.clear(); + for (session_id, msgs) in sessions { + messages.insert(session_id, msgs); + } + Ok(()) + } } pub struct OpenAITokenCounter { @@ -435,6 +465,20 @@ impl CoreMemoryStore for InMemoryCoreMemoryStore { facts.remove(session_id); Ok(()) } + + async fn dump_all(&self) -> Result)>, MemoryError> { + let facts = self.facts.lock().map_err(|e| MemoryError::Message(e.to_string()))?; + Ok(facts.iter().map(|(k, v)| (k.clone(), v.clone())).collect()) + } + + async fn restore_all(&self, sessions: Vec<(String, Vec)>) -> Result<(), MemoryError> { + let mut facts = self.facts.lock().map_err(|e| MemoryError::Message(e.to_string()))?; + facts.clear(); + for (session_id, list) in sessions { + facts.insert(session_id, list); + } + Ok(()) + } } fn conversation_text(messages: &[Message]) -> String { @@ -723,6 +767,45 @@ mod tests { assert!(facts.is_empty()); } + #[tokio::test] + async fn in_memory_store_dump_and_restore_all_sessions() { + let store = InMemoryStore::default(); + store.add_message("s1", message("user", "hi")).await.unwrap(); + store.add_message("s2", message("user", "yo")).await.unwrap(); + + let dump = store.dump_all().await.unwrap(); + assert_eq!(dump.len(), 2); + + let fresh = InMemoryStore::default(); + // Pre-existing data in `fresh` must be wiped by restore_all. + fresh.add_message("stale", message("user", "old")).await.unwrap(); + fresh.restore_all(dump).await.unwrap(); + + assert!(fresh.get_recent("stale", 10).await.unwrap().is_empty()); + assert_eq!(fresh.get_recent("s1", 10).await.unwrap()[0].content, "hi"); + assert_eq!(fresh.get_recent("s2", 10).await.unwrap()[0].content, "yo"); + } + + #[tokio::test] + async fn restore_all_empty_clears_everything() { + let store = InMemoryStore::default(); + store.add_message("s1", message("user", "hi")).await.unwrap(); + store.restore_all(vec![]).await.unwrap(); + assert!(store.get_recent("s1", 10).await.unwrap().is_empty()); + } + + #[tokio::test] + async fn in_memory_core_memory_dump_and_restore() { + let store = InMemoryCoreMemoryStore::default(); + store.add_fact("s1", "a").await.unwrap(); + store.add_fact("s1", "b").await.unwrap(); + let dump = store.dump_all().await.unwrap(); + + let fresh = InMemoryCoreMemoryStore::default(); + fresh.restore_all(dump).await.unwrap(); + assert_eq!(fresh.get_facts("s1").await.unwrap(), vec!["a".to_string(), "b".to_string()]); + } + #[test] fn memory_server_error_wraps_each_trait_error_type() { let embed_error = MemoryServerError::from(EmbedError::Message("embed".to_string())); diff --git a/src/knowledge/graph.rs b/src/knowledge/graph.rs index 31e95c2..34c2c51 100644 --- a/src/knowledge/graph.rs +++ b/src/knowledge/graph.rs @@ -209,6 +209,59 @@ impl Default for KnowledgeGraph { fn default() -> Self { Self::new() } } +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct SessionGraphSnapshot { + pub session_id: String, + pub entities: Vec, + pub relationships: Vec, +} + +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct GraphSnapshot { + pub sessions: Vec, + /// Dedup keys ("session_id\x00message_id") so re-applied commands stay idempotent. + pub processed: Vec, +} + +impl KnowledgeGraph { + pub fn session_ids(&self) -> Vec { + self.sessions.keys().cloned().collect() + } + + pub fn to_snapshot(&self) -> GraphSnapshot { + let sessions = self + .sessions + .keys() + .map(|sid| SessionGraphSnapshot { + session_id: sid.clone(), + entities: self.all_entities(sid), + relationships: self.all_relationships(sid), + }) + .collect(); + GraphSnapshot { + sessions, + processed: self.processed.iter().cloned().collect(), + } + } + + pub fn from_snapshot(snap: GraphSnapshot) -> Self { + let mut kg = KnowledgeGraph::new(); + for s in snap.sessions { + let session = kg.sessions.entry(s.session_id.clone()).or_insert_with(SessionGraph::new); + for e in &s.entities { + session.ensure_entity(&e.name, &e.entity_type, e.attributes.clone()); + } + for r in &s.relationships { + let from = session.ensure_entity(&r.from, "Other", HashMap::new()); + let to = session.ensure_entity(&r.to, "Other", HashMap::new()); + session.graph.add_edge(from, to, RelEdge { relationship_type: r.relationship_type.clone() }); + } + } + kg.processed = snap.processed.into_iter().collect(); + kg + } +} + #[cfg(test)] mod tests { use super::*; @@ -327,4 +380,35 @@ mod tests { kg.apply_extraction("s1", "m1", vec![entity("Alice","Person")], vec![]); assert!(kg.all_entities("s2").is_empty()); } + + #[test] + fn graph_snapshot_round_trips_entities_relationships_and_dedup() { + 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("s2", "m9", vec![entity("Bob","Person")], vec![]); + + let snap = kg.to_snapshot(); + let json = serde_json::to_string(&snap).unwrap(); + let back_snap: GraphSnapshot = serde_json::from_str(&json).unwrap(); + + let restored = KnowledgeGraph::from_snapshot(back_snap); + assert_eq!(restored.all_entities("s1").len(), 2); + assert_eq!(restored.all_relationships("s1").len(), 1); + assert_eq!(restored.all_entities("s2").len(), 1); + assert!(restored.is_processed("s1", "m1")); + } + + #[test] + fn restored_graph_answers_capability_queries() { + let mut kg = KnowledgeGraph::new(); + kg.apply_extraction("s1", "m1", + vec![entity("Alice","Person"), entity("Bob","Person")], + vec![rel("Alice","Bob","knows")]); + let restored = KnowledgeGraph::from_snapshot(kg.to_snapshot()); + let path = restored.find_path("s1", "Alice", "Bob").unwrap(); + assert_eq!(path.len(), 1); + assert_eq!(path[0].relationship_type, "knows"); + } } diff --git a/src/metrics.rs b/src/metrics.rs index 5e9e323..8959d1a 100644 --- a/src/metrics.rs +++ b/src/metrics.rs @@ -22,6 +22,9 @@ pub struct AppMetrics { knowledge_entities_extracted_total: IntCounter, knowledge_relationships_extracted_total: IntCounter, knowledge_queue_size: IntGauge, + snapshot_build_total: IntCounter, + snapshot_install_total: IntCounter, + snapshot_last_index: IntGauge, } impl AppMetrics { @@ -127,6 +130,24 @@ impl AppMetrics { ))?; registry.register(Box::new(knowledge_queue_size.clone()))?; + let snapshot_build_total = IntCounter::with_opts(Opts::new( + "snapshot_build_total", + "Total Raft snapshots built on this node.", + ))?; + registry.register(Box::new(snapshot_build_total.clone()))?; + + let snapshot_install_total = IntCounter::with_opts(Opts::new( + "snapshot_install_total", + "Total Raft snapshots installed on this node.", + ))?; + registry.register(Box::new(snapshot_install_total.clone()))?; + + let snapshot_last_index = IntGauge::with_opts(Opts::new( + "snapshot_last_index", + "Log index of the most recent snapshot.", + ))?; + registry.register(Box::new(snapshot_last_index.clone()))?; + Ok(Self { registry, messages_added_total, @@ -143,6 +164,9 @@ impl AppMetrics { knowledge_entities_extracted_total, knowledge_relationships_extracted_total, knowledge_queue_size, + snapshot_build_total, + snapshot_install_total, + snapshot_last_index, }) } @@ -208,6 +232,18 @@ impl AppMetrics { self.knowledge_queue_size.set(size as i64); } + pub fn increment_snapshot_build(&self) { + self.snapshot_build_total.inc(); + } + + pub fn increment_snapshot_install(&self) { + self.snapshot_install_total.inc(); + } + + pub fn set_snapshot_last_index(&self, index: u64) { + self.snapshot_last_index.set(index as i64); + } + pub fn render(&self) -> Result { let mut buffer = Vec::new(); let encoder = TextEncoder::new(); @@ -217,4 +253,21 @@ impl AppMetrics { String::from_utf8(buffer).map_err(|error| error.to_string()) } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn renders_snapshot_metrics() { + let metrics = AppMetrics::new().unwrap(); + metrics.increment_snapshot_build(); + metrics.increment_snapshot_install(); + metrics.set_snapshot_last_index(42); + let text = metrics.render().unwrap(); + assert!(text.contains("engram_snapshot_build_total")); + assert!(text.contains("engram_snapshot_install_total")); + assert!(text.contains("engram_snapshot_last_index")); + } } \ No newline at end of file diff --git a/src/raft/grpc_server.rs b/src/raft/grpc_server.rs index b74e2d2..b2dd792 100644 --- a/src/raft/grpc_server.rs +++ b/src/raft/grpc_server.rs @@ -5,6 +5,7 @@ use tonic::{Request, Response, Status}; use crate::proto::raft::{ raft_service_server::RaftService, AppendEntriesRequest, AppendEntriesResponse, + InstallSnapshotRequest, InstallSnapshotResponse, VoteRequest, VoteResponse, }; use crate::raft::types::RaftHandle; @@ -46,4 +47,20 @@ impl RaftService for RaftGrpcServer { .map_err(|e| Status::internal(e.to_string()))?; Ok(Response::new((&resp).into())) } + + async fn install_snapshot( + &self, + request: Request, + ) -> Result, Status> { + let req = request + .into_inner() + .try_into() + .map_err(|e: String| Status::invalid_argument(e))?; + let resp = self + .raft + .install_snapshot(req) + .await + .map_err(|e| Status::internal(e.to_string()))?; + Ok(Response::new((&resp).into())) + } } diff --git a/src/raft/log_store.rs b/src/raft/log_store.rs index 0d6b57f..a923ffb 100644 --- a/src/raft/log_store.rs +++ b/src/raft/log_store.rs @@ -1,26 +1,79 @@ -use std::collections::BTreeMap; +use std::io; use std::ops::RangeBounds; use std::sync::Arc; -use tokio::sync::Mutex; use openraft::{ LogId, LogState, RaftLogReader, Vote, Entry, + ErrorSubject, ErrorVerb, storage::{LogFlushed, RaftLogStorage}, StorageError, }; +use redb::{Database, ReadableTable, TableDefinition}; use crate::raft::types::TypeConfig; -#[derive(Debug, Default, Clone)] +const LOG_TABLE: TableDefinition = TableDefinition::new("raft_log"); +const META_TABLE: TableDefinition<&str, &[u8]> = TableDefinition::new("raft_meta"); + +const META_VOTE: &str = "vote"; +const META_COMMITTED: &str = "committed"; +const META_LAST_PURGED: &str = "last_purged"; + +#[derive(Clone)] pub struct EngRaftLogStore { - inner: Arc>, + db: Arc, +} + +impl EngRaftLogStore { + pub fn new(db: Arc) -> Self { + // Ensure both tables exist so read txns never fail on a fresh db. + let txn = db.begin_write().expect("redb begin_write on init"); + { + let _ = txn.open_table(LOG_TABLE).expect("open LOG_TABLE"); + let _ = txn.open_table(META_TABLE).expect("open META_TABLE"); + } + txn.commit().expect("redb commit on init"); + Self { db } + } + + fn read_meta(&self, key: &str) -> Result>, StorageError> { + let txn = self.db.begin_read().map_err(read_err)?; + let table = txn.open_table(META_TABLE).map_err(read_err)?; + Ok(table.get(key).map_err(read_err)?.map(|v| v.value().to_vec())) + } + + fn write_meta(&self, key: &str, bytes: &[u8]) -> Result<(), StorageError> { + let txn = self.db.begin_write().map_err(write_err)?; + { + let mut table = txn.open_table(META_TABLE).map_err(write_err)?; + table.insert(key, bytes).map_err(write_err)?; + } + txn.commit().map_err(write_err)?; + Ok(()) + } +} + +fn read_err(e: E) -> StorageError { + StorageError::from_io_error( + ErrorSubject::Store, + ErrorVerb::Read, + io::Error::new(io::ErrorKind::Other, e.to_string()), + ) } -#[derive(Debug, Default)] -struct LogStoreInner { - last_purged_log_id: Option>, - log: BTreeMap>, - committed: Option>, - vote: Option>, +fn write_err(e: E) -> StorageError { + StorageError::from_io_error( + ErrorSubject::Store, + ErrorVerb::Write, + io::Error::new(io::ErrorKind::Other, e.to_string()), + ) +} + +fn decode_err(e: serde_json::Error) -> StorageError { + StorageError::from_io_error( + ErrorSubject::Logs, + ErrorVerb::Read, + io::Error::new(io::ErrorKind::InvalidData, e.to_string()), + ) } impl RaftLogReader for EngRaftLogStore { @@ -28,8 +81,16 @@ impl RaftLogReader for EngRaftLogStore { &mut self, range: RB, ) -> Result>, StorageError> { - let inner = self.inner.lock().await; - Ok(inner.log.range(range).map(|(_, e)| e.clone()).collect()) + let txn = self.db.begin_read().map_err(read_err)?; + let table = txn.open_table(LOG_TABLE).map_err(read_err)?; + let mut out = Vec::new(); + for item in table.range(range).map_err(read_err)? { + let (_, value) = item.map_err(read_err)?; + let entry: Entry = + serde_json::from_slice(value.value()).map_err(decode_err)?; + out.push(entry); + } + Ok(out) } } @@ -37,62 +98,92 @@ impl RaftLogStorage for EngRaftLogStore { type LogReader = Self; async fn get_log_state(&mut self) -> Result, StorageError> { - let inner = self.inner.lock().await; - let last = inner - .log - .values() - .next_back() - .map(|e| e.log_id.clone()) - .or_else(|| inner.last_purged_log_id.clone()); - Ok(LogState { - last_purged_log_id: inner.last_purged_log_id.clone(), - last_log_id: last, - }) - } - - async fn save_committed(&mut self, committed: Option>) -> Result<(), StorageError> { - self.inner.lock().await.committed = committed; - Ok(()) + let last_purged: Option> = match self.read_meta(META_LAST_PURGED)? { + Some(b) => Some(serde_json::from_slice(&b).map_err(decode_err)?), + None => None, + }; + let txn = self.db.begin_read().map_err(read_err)?; + let table = txn.open_table(LOG_TABLE).map_err(read_err)?; + let last = match table.last().map_err(read_err)? { + Some((_, value)) => { + let entry: Entry = + serde_json::from_slice(value.value()).map_err(decode_err)?; + Some(entry.log_id) + } + None => last_purged.clone(), + }; + Ok(LogState { last_purged_log_id: last_purged, last_log_id: last }) + } + + async fn save_committed( + &mut self, + committed: Option>, + ) -> Result<(), StorageError> { + let bytes = serde_json::to_vec(&committed).map_err(write_err)?; + self.write_meta(META_COMMITTED, &bytes) } async fn read_committed(&mut self) -> Result>, StorageError> { - Ok(self.inner.lock().await.committed.clone()) + match self.read_meta(META_COMMITTED)? { + Some(b) => Ok(serde_json::from_slice(&b).map_err(decode_err)?), + None => Ok(None), + } } async fn save_vote(&mut self, vote: &Vote) -> Result<(), StorageError> { - self.inner.lock().await.vote = Some(vote.clone()); - Ok(()) + let bytes = serde_json::to_vec(vote).map_err(write_err)?; + self.write_meta(META_VOTE, &bytes) } async fn read_vote(&mut self) -> Result>, StorageError> { - Ok(self.inner.lock().await.vote.clone()) + match self.read_meta(META_VOTE)? { + Some(b) => Ok(Some(serde_json::from_slice(&b).map_err(decode_err)?)), + None => Ok(None), + } } - async fn append(&mut self, entries: I, callback: LogFlushed) -> Result<(), StorageError> + async fn append( + &mut self, + entries: I, + callback: LogFlushed, + ) -> Result<(), StorageError> where I: IntoIterator> + Send, I::IntoIter: Send, { + let txn = self.db.begin_write().map_err(write_err)?; { - let mut inner = self.inner.lock().await; + let mut table = txn.open_table(LOG_TABLE).map_err(write_err)?; for entry in entries { - inner.log.insert(entry.log_id.index, entry); + let bytes = serde_json::to_vec(&entry).map_err(write_err)?; + table.insert(entry.log_id.index, bytes.as_slice()).map_err(write_err)?; } } - // In-memory: flush is instantaneous. Signal completion immediately. + // DURABILITY HINGE: commit (fsync) BEFORE signaling completion to openraft. + txn.commit().map_err(write_err)?; callback.log_io_completed(Ok(())); Ok(()) } async fn truncate(&mut self, log_id: LogId) -> Result<(), StorageError> { - self.inner.lock().await.log.retain(|&k, _| k < log_id.index); + let txn = self.db.begin_write().map_err(write_err)?; + { + let mut table = txn.open_table(LOG_TABLE).map_err(write_err)?; + table.retain(|k, _| k < log_id.index).map_err(write_err)?; + } + txn.commit().map_err(write_err)?; Ok(()) } async fn purge(&mut self, log_id: LogId) -> Result<(), StorageError> { - let mut inner = self.inner.lock().await; - inner.last_purged_log_id = Some(log_id.clone()); - inner.log.retain(|&k, _| k > log_id.index); + let bytes = serde_json::to_vec(&log_id).map_err(write_err)?; + self.write_meta(META_LAST_PURGED, &bytes)?; + let txn = self.db.begin_write().map_err(write_err)?; + { + let mut table = txn.open_table(LOG_TABLE).map_err(write_err)?; + table.retain(|k, _| k > log_id.index).map_err(write_err)?; + } + txn.commit().map_err(write_err)?; Ok(()) } @@ -103,14 +194,18 @@ impl RaftLogStorage for EngRaftLogStore { #[cfg(test)] impl EngRaftLogStore { - /// Test helper: insert entries directly into the log without going through - /// the LogFlushed callback (which is pub(crate) in openraft and cannot be - /// constructed in external test code). + /// Test helper: insert entries directly via a redb write txn, bypassing the + /// LogFlushed callback (which is pub(crate) in openraft and not constructible here). async fn insert_for_test(&self, entries: Vec>) { - let mut inner = self.inner.lock().await; - for entry in entries { - inner.log.insert(entry.log_id.index, entry); + let txn = self.db.begin_write().unwrap(); + { + let mut table = txn.open_table(LOG_TABLE).unwrap(); + for entry in entries { + let bytes = serde_json::to_vec(&entry).unwrap(); + table.insert(entry.log_id.index, bytes.as_slice()).unwrap(); + } } + txn.commit().unwrap(); } } @@ -119,6 +214,12 @@ mod tests { use super::*; use openraft::{CommittedLeaderId, EntryPayload, LogState}; + fn temp_store() -> (EngRaftLogStore, tempfile::TempDir) { + let dir = tempfile::tempdir().unwrap(); + let db = Database::create(dir.path().join("test.redb")).unwrap(); + (EngRaftLogStore::new(std::sync::Arc::new(db)), dir) + } + fn log_id(term: u64, index: u64) -> LogId { LogId::new(CommittedLeaderId::new(term, 1), index) } @@ -129,14 +230,14 @@ mod tests { #[tokio::test] async fn initial_log_state_is_empty() { - let mut store = EngRaftLogStore::default(); + let (mut store, _d) = temp_store(); let state = store.get_log_state().await.unwrap(); assert_eq!(state, LogState { last_purged_log_id: None, last_log_id: None }); } #[tokio::test] async fn save_and_read_vote() { - let mut store = EngRaftLogStore::default(); + let (mut store, _d) = temp_store(); assert!(store.read_vote().await.unwrap().is_none()); let vote = Vote::new(1, 1); store.save_vote(&vote).await.unwrap(); @@ -145,7 +246,7 @@ mod tests { #[tokio::test] async fn save_and_read_committed() { - let mut store = EngRaftLogStore::default(); + let (mut store, _d) = temp_store(); assert!(store.read_committed().await.unwrap().is_none()); let lid = log_id(1, 5); store.save_committed(Some(lid.clone())).await.unwrap(); @@ -154,7 +255,7 @@ mod tests { #[tokio::test] async fn append_and_read_back() { - let mut store = EngRaftLogStore::default(); + let (mut store, _d) = temp_store(); store.insert_for_test(vec![blank(1, 0), blank(1, 1), blank(1, 2)]).await; let got = store.try_get_log_entries(0..3).await.unwrap(); assert_eq!(got.len(), 3); @@ -164,11 +265,9 @@ mod tests { #[tokio::test] async fn truncate_removes_from_index_inclusive() { - let mut store = EngRaftLogStore::default(); + let (mut store, _d) = temp_store(); store.insert_for_test(vec![blank(1, 0), blank(1, 1), blank(1, 2), blank(1, 3)]).await; - store.truncate(log_id(1, 2)).await.unwrap(); - let got = store.try_get_log_entries(0..10).await.unwrap(); assert_eq!(got.len(), 2, "entries 0 and 1 should remain"); assert_eq!(got[1].log_id.index, 1); @@ -176,36 +275,37 @@ mod tests { #[tokio::test] async fn purge_removes_up_to_inclusive_and_updates_last_purged() { - let mut store = EngRaftLogStore::default(); + let (mut store, _d) = temp_store(); store.insert_for_test(vec![blank(1, 0), blank(1, 1), blank(1, 2)]).await; - store.purge(log_id(1, 1)).await.unwrap(); - let got = store.try_get_log_entries(0..10).await.unwrap(); assert_eq!(got.len(), 1, "only entry at index 2 should remain"); assert_eq!(got[0].log_id.index, 2); - let state = store.get_log_state().await.unwrap(); assert_eq!(state.last_purged_log_id, Some(log_id(1, 1))); } #[tokio::test] - async fn log_state_last_log_id_falls_back_to_last_purged_when_log_is_empty() { - let mut store = EngRaftLogStore::default(); - store.insert_for_test(vec![blank(1, 0)]).await; - store.purge(log_id(1, 0)).await.unwrap(); - - let state = store.get_log_state().await.unwrap(); - assert_eq!(state.last_purged_log_id, Some(log_id(1, 0))); - assert_eq!(state.last_log_id, Some(log_id(1, 0))); - } - - #[tokio::test] - async fn get_log_reader_is_a_clone_that_reads_same_data() { - let mut store = EngRaftLogStore::default(); - store.insert_for_test(vec![blank(1, 0)]).await; - let mut reader = store.get_log_reader().await; - let got = reader.try_get_log_entries(0..10).await.unwrap(); + async fn data_survives_database_reopen() { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("persist.redb"); + { + let db = std::sync::Arc::new(Database::create(&path).unwrap()); + let mut store = EngRaftLogStore::new(db); + store.insert_for_test(vec![blank(2, 0), blank(2, 1)]).await; + store.save_vote(&Vote::new(2, 1)).await.unwrap(); + store.purge(log_id(2, 0)).await.unwrap(); + } + // Reopen the same file with a brand-new Database handle. + let db = std::sync::Arc::new(Database::open(&path).unwrap()); + let mut store = EngRaftLogStore::new(db); + let got = store.try_get_log_entries(0..10).await.unwrap(); assert_eq!(got.len(), 1); + assert_eq!(got[0].log_id.index, 1); + assert_eq!(store.read_vote().await.unwrap(), Some(Vote::new(2, 1))); + assert_eq!( + store.get_log_state().await.unwrap().last_purged_log_id, + Some(log_id(2, 0)) + ); } } diff --git a/src/raft/mod.rs b/src/raft/mod.rs index eddbe58..70710f3 100644 --- a/src/raft/mod.rs +++ b/src/raft/mod.rs @@ -2,5 +2,7 @@ pub mod grpc_server; pub mod log_store; pub mod network; pub mod proto_conv; +pub mod recovery; +pub mod snapshot; pub mod state_machine; pub mod types; diff --git a/src/raft/network.rs b/src/raft/network.rs index 3b2e086..889ca76 100644 --- a/src/raft/network.rs +++ b/src/raft/network.rs @@ -77,22 +77,29 @@ impl RaftNetwork for EngRaftNetworkConnection { resp.into_inner().try_into().map_err(proto_decode_err) } - // Stage 1: InstallSnapshot is not implemented. - // If OpenRaft calls this, it means a follower has fallen too far behind for - // log-based catch-up. The operator must manually remove and re-add the node. - // Stage 2 will implement snapshot transport. async fn install_snapshot( &mut self, - _rpc: InstallSnapshotRequest, + rpc: InstallSnapshotRequest, _option: RPCOption, ) -> Result< InstallSnapshotResponse, RPCError>, > { - Err(RPCError::Unreachable(Unreachable::new(&io::Error::new( - io::ErrorKind::Unsupported, - "InstallSnapshot not implemented in Stage 1; re-add the node manually", - )))) + let endpoint = format!("http://{}", self.target_addr) + .parse::() + .map_err(|e| RPCError::Network(NetworkError::new(&e)))?; + let channel = Channel::builder(endpoint) + .connect() + .await + .map_err(|e| RPCError::Unreachable(Unreachable::new(&e)))?; + let mut client = RaftServiceClient::new(channel); + let resp = client + .install_snapshot(crate::proto::raft::InstallSnapshotRequest::from(&rpc)) + .await + .map_err(|e| RPCError::Network(NetworkError::new(&e)))?; + resp.into_inner() + .try_into() + .map_err(|e: String| RPCError::Network(NetworkError::new(&io::Error::new(io::ErrorKind::InvalidData, e)))) } } @@ -108,4 +115,23 @@ mod tests { // Tonic uses lazy channels so new_client must not attempt a real connection. let _conn = factory.new_client(1u64, &node).await; } + + #[tokio::test] + async fn install_snapshot_attempts_connection_and_errors_when_unreachable() { + let mut conn = EngRaftNetworkConnection { target_addr: "127.0.0.1:1".to_string() }; + let req = openraft::raft::InstallSnapshotRequest:: { + vote: openraft::Vote::new_committed(1, 1), + meta: openraft::SnapshotMeta { + last_log_id: None, + last_membership: openraft::StoredMembership::default(), + snapshot_id: "x".into(), + }, + offset: 0, + data: vec![], + done: true, + }; + // Port 1 is unreachable: must return an RPCError, not the Stage-1 "Unsupported". + let res = conn.install_snapshot(req, openraft::network::RPCOption::new(std::time::Duration::from_millis(100))).await; + assert!(res.is_err()); + } } diff --git a/src/raft/proto_conv.rs b/src/raft/proto_conv.rs index a351ed3..8fa5f2b 100644 --- a/src/raft/proto_conv.rs +++ b/src/raft/proto_conv.rs @@ -188,11 +188,90 @@ impl TryFrom for AppendEntriesResponse { } } +// --- InstallSnapshotRequest --- + +impl From<&openraft::raft::InstallSnapshotRequest> for p::InstallSnapshotRequest { + fn from(r: &openraft::raft::InstallSnapshotRequest) -> Self { + p::InstallSnapshotRequest { + vote: Some((&r.vote).into()), + meta: serde_json::to_vec(&r.meta).unwrap_or_default(), + offset: r.offset, + data: r.data.clone(), + done: r.done, + } + } +} + +impl TryFrom for openraft::raft::InstallSnapshotRequest { + type Error = String; + fn try_from(r: p::InstallSnapshotRequest) -> Result { + let meta = serde_json::from_slice(&r.meta) + .map_err(|e| format!("install_snapshot meta decode error: {e}"))?; + Ok(openraft::raft::InstallSnapshotRequest { + vote: r.vote.ok_or("install_snapshot_request missing vote")?.into(), + meta, + offset: r.offset, + data: r.data, + done: r.done, + }) + } +} + +// --- InstallSnapshotResponse --- + +impl From<&openraft::raft::InstallSnapshotResponse> for p::InstallSnapshotResponse { + fn from(r: &openraft::raft::InstallSnapshotResponse) -> Self { + p::InstallSnapshotResponse { vote: Some((&r.vote).into()) } + } +} + +impl TryFrom for openraft::raft::InstallSnapshotResponse { + type Error = String; + fn try_from(r: p::InstallSnapshotResponse) -> Result { + Ok(openraft::raft::InstallSnapshotResponse { + vote: r.vote.ok_or("install_snapshot_response missing vote")?.into(), + }) + } +} + #[cfg(test)] mod tests { use crate::proto::raft as proto; use openraft::raft::{AppendEntriesResponse, VoteRequest}; + #[test] + fn install_snapshot_request_round_trips() { + use openraft::raft::InstallSnapshotRequest; + use openraft::{SnapshotMeta, StoredMembership}; + let req = InstallSnapshotRequest:: { + vote: openraft::Vote::new_committed(2, 1), + meta: SnapshotMeta { + last_log_id: Some(openraft::LogId::new(openraft::CommittedLeaderId::new(2, 1), 9)), + last_membership: StoredMembership::default(), + snapshot_id: "snap-1".into(), + }, + offset: 0, + data: vec![1, 2, 3], + done: true, + }; + let p: proto::InstallSnapshotRequest = (&req).into(); + let back: InstallSnapshotRequest = p.try_into().unwrap(); + assert_eq!(back.offset, 0); + assert!(back.done); + assert_eq!(back.data, vec![1, 2, 3]); + assert_eq!(back.meta.snapshot_id, "snap-1"); + assert_eq!(back.meta.last_log_id.unwrap().index, 9); + } + + #[test] + fn install_snapshot_response_round_trips() { + use openraft::raft::InstallSnapshotResponse; + let resp = InstallSnapshotResponse:: { vote: openraft::Vote::new_committed(3, 2) }; + let p: proto::InstallSnapshotResponse = (&resp).into(); + let back: InstallSnapshotResponse = p.try_into().unwrap(); + assert_eq!(back.vote.leader_id().term, 3); + } + #[test] fn vote_round_trips() { let original = openraft::Vote::::new_committed(1, 1); diff --git a/src/raft/recovery.rs b/src/raft/recovery.rs new file mode 100644 index 0000000..f94e898 --- /dev/null +++ b/src/raft/recovery.rs @@ -0,0 +1,97 @@ +use std::sync::Arc; + +use crate::core::{CoreMemoryStore, ShortTermMemory}; +use crate::knowledge::graph::KnowledgeGraph; +use crate::raft::snapshot::EngramSnapshot; +use crate::raft::state_machine::EngStateMachineStore; + +/// Startup recovery: clear this node's memory stores, then restore the latest +/// persisted snapshot (if any) into the stores, knowledge graph, and the state +/// machine's applied bookkeeping. openraft replays committed log entries after +/// the restored index once `Raft::new` runs. +pub async fn recover_state_machine( + sm: &EngStateMachineStore, + short_term: Arc, + core_memory: Arc, +) -> anyhow::Result<()> { + // 1. Always flush first so stale state from a prior run cannot bleed through. + short_term.restore_all(vec![]).await?; + core_memory.restore_all(vec![]).await?; + + // 2. Load the persisted snapshot, if present. + let Some((meta, bytes)) = sm.load_snapshot_for_recovery()? else { + tracing::info!("recovery: no snapshot found; openraft will replay full committed log"); + return Ok(()); + }; + let payload = EngramSnapshot::from_bytes(&bytes)?; + + // 3. Restore stores + graph + applied bookkeeping. + let st_sessions = payload.short_term.into_iter().map(|s| (s.session_id, s.messages)).collect(); + short_term.restore_all(st_sessions).await?; + let cm_sessions = payload.core_memory.into_iter().map(|s| (s.session_id, s.facts)).collect(); + core_memory.restore_all(cm_sessions).await?; + let graph = KnowledgeGraph::from_snapshot(payload.knowledge_graph); + sm.restore_applied_for_recovery(meta, graph).await; + + tracing::info!("recovery: restored snapshot; openraft will replay log tail"); + Ok(()) +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + use tokio::sync::{mpsc, RwLock}; + use redb::Database; + + use crate::core::{CoreMemoryStore, InMemoryCoreMemoryStore, InMemoryStore, InMemoryVectorStore, ShortTermMemory}; + use crate::knowledge::graph::KnowledgeGraph; + use crate::raft::recovery::recover_state_machine; + use crate::raft::state_machine::EngStateMachineStore; + use crate::raft::types::MemoryCommand; + + fn build(db: Arc) -> (EngStateMachineStore, Arc, Arc, Arc>) { + let st = Arc::new(InMemoryStore::default()); + let cm = Arc::new(InMemoryCoreMemoryStore::default()); + let vs = Arc::new(InMemoryVectorStore::default()); + let (etx, _erx) = mpsc::channel(10); + let (ktx, _krx) = mpsc::channel(10); + let kg = Arc::new(RwLock::new(KnowledgeGraph::new())); + let sm = EngStateMachineStore::new(st.clone(), cm.clone(), vs, etx, kg.clone(), ktx, db); + (sm, st, cm, kg) + } + + fn cm_as_dyn(cm: &Arc) -> Arc { + cm.clone() + } + + #[tokio::test] + async fn recovery_with_no_snapshot_flushes_stores() { + let dir = tempfile::tempdir().unwrap(); + let db = Arc::new(Database::create(dir.path().join("r.redb")).unwrap()); + let (sm, st, cm, _kg) = build(db.clone()); + st.add_message("stale", crate::models::Message { + id: Some("x".into()), role: "user".into(), content: "old".into(), + timestamp: None, embedding_status: None, + }).await.unwrap(); + + recover_state_machine(&sm, st.clone() as Arc, cm_as_dyn(&cm)).await.unwrap(); + assert!(st.get_recent("stale", 10).await.unwrap().is_empty()); + } + + #[tokio::test] + async fn recovery_restores_persisted_snapshot() { + let dir = tempfile::tempdir().unwrap(); + let db = Arc::new(Database::create(dir.path().join("r.redb")).unwrap()); + + // Apply a fact and build a snapshot so it's persisted into db. + let (mut src, _st, _src_cm, _kg) = build(db.clone()); + src.apply_for_test(0, MemoryCommand::AddFact { session_id: "s1".into(), fact: "remember me".into() }).await; + let mut b = openraft::storage::RaftStateMachine::get_snapshot_builder(&mut src).await; + openraft::RaftSnapshotBuilder::build_snapshot(&mut b).await.unwrap(); + + // Fresh state machine over the same db; recovery should restore the fact. + let (sm, st, cm, _kg2) = build(db.clone()); + recover_state_machine(&sm, st as Arc, cm.clone() as Arc).await.unwrap(); + assert_eq!(cm.get_facts("s1").await.unwrap(), vec!["remember me".to_string()]); + } +} diff --git a/src/raft/snapshot.rs b/src/raft/snapshot.rs new file mode 100644 index 0000000..309d960 --- /dev/null +++ b/src/raft/snapshot.rs @@ -0,0 +1,84 @@ +use serde::{Deserialize, Serialize}; + +use crate::knowledge::graph::GraphSnapshot; +use crate::models::Message; + +/// Snapshot schema version. Bump when the payload layout changes incompatibly. +pub const SNAPSHOT_VERSION: u32 = 1; + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SessionMessages { + pub session_id: String, + pub messages: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SessionFacts { + pub session_id: String, + pub facts: Vec, +} + +/// Full applied state captured at a Raft log index. +/// +/// LanceDB is intentionally NOT included. It is per-node, on-disk, and +/// deterministically rebuildable from message text. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct EngramSnapshot { + pub version: u32, + pub short_term: Vec, + pub core_memory: Vec, + pub knowledge_graph: GraphSnapshot, + /// Reserved for Stage 3B (collective/global knowledge graph). Absent in 3A. + #[serde(default)] + pub global_graph: Option, +} + +impl EngramSnapshot { + pub fn to_bytes(&self) -> Result, serde_json::Error> { + serde_json::to_vec(self) + } + pub fn from_bytes(bytes: &[u8]) -> Result { + serde_json::from_slice(bytes) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn sample() -> EngramSnapshot { + EngramSnapshot { + version: SNAPSHOT_VERSION, + short_term: vec![SessionMessages { + session_id: "s1".into(), + messages: vec![], + }], + core_memory: vec![SessionFacts { session_id: "s1".into(), facts: vec!["f".into()] }], + knowledge_graph: crate::knowledge::graph::GraphSnapshot::default(), + global_graph: None, + } + } + + #[test] + fn snapshot_carries_version_one() { + assert_eq!(sample().version, 1); + } + + #[test] + fn snapshot_round_trips_through_bytes() { + let snap = sample(); + let bytes = snap.to_bytes().unwrap(); + let back = EngramSnapshot::from_bytes(&bytes).unwrap(); + assert_eq!(back.version, 1); + assert_eq!(back.core_memory[0].facts, vec!["f".to_string()]); + assert!(back.global_graph.is_none()); + } + + #[test] + fn unknown_global_graph_absent_by_default() { + let bytes = sample().to_bytes().unwrap(); + // Absent global_graph must deserialize cleanly (forward-compat for 3B). + let back = EngramSnapshot::from_bytes(&bytes).unwrap(); + assert!(back.global_graph.is_none()); + } +} diff --git a/src/raft/state_machine.rs b/src/raft/state_machine.rs index 6c5fcd8..6282608 100644 --- a/src/raft/state_machine.rs +++ b/src/raft/state_machine.rs @@ -7,14 +7,20 @@ use openraft::{ StorageError, StoredMembership, RaftSnapshotBuilder, storage::RaftStateMachine, }; +use redb::{Database, TableDefinition}; use crate::core::{CoreMemoryStore, ShortTermMemory}; use crate::knowledge::graph::KnowledgeGraph; use crate::knowledge::types::KnowledgeJob; use crate::models::{EmbeddingStatus, Message}; +use crate::raft::snapshot::{EngramSnapshot, SessionFacts, SessionMessages}; use crate::raft::types::{CommandResponse, MemoryCommand, TypeConfig}; use crate::worker::EmbeddingJob; +const SNAPSHOT_TABLE: TableDefinition<&str, &[u8]> = TableDefinition::new("raft_snapshot"); +const SNAPSHOT_META_KEY: &str = "meta"; +const SNAPSHOT_DATA_KEY: &str = "data"; + pub struct EngStateMachineStore { inner: Arc>, } @@ -27,6 +33,8 @@ struct SmInner { embedding_tx: mpsc::Sender, knowledge_graph: Arc>, knowledge_tx: mpsc::Sender, + db: Arc, + snapshot_idx: u64, } impl EngStateMachineStore { @@ -39,7 +47,13 @@ impl EngStateMachineStore { embedding_tx: mpsc::Sender, knowledge_graph: Arc>, knowledge_tx: mpsc::Sender, + db: Arc, ) -> Self { + { + let txn = db.begin_write().expect("redb begin_write sm init"); + { let _ = txn.open_table(SNAPSHOT_TABLE).expect("open SNAPSHOT_TABLE"); } + txn.commit().expect("redb commit sm init"); + } Self { inner: Arc::new(Mutex::new(SmInner { last_applied: None, @@ -49,15 +63,151 @@ impl EngStateMachineStore { embedding_tx, knowledge_graph, knowledge_tx, + db, + snapshot_idx: 0, })), } } + + pub(crate) fn inner_handle(&self) -> Arc> { + self.inner.clone() + } + + /// Returns `(meta, payload_bytes)` of the persisted snapshot, if any. + /// Called at startup before the Raft node starts (uncontended). + pub(crate) fn load_snapshot_for_recovery( + &self, + ) -> anyhow::Result, Vec)>> { + let db = { + let inner = self.inner.try_lock().expect("uncontended at startup"); + inner.db.clone() + }; + match load_persisted_snapshot(&db).map_err(|e| anyhow::anyhow!(e.to_string()))? { + Some(s) => Ok(Some((s.meta, s.snapshot.into_inner()))), + None => Ok(None), + } + } + + /// Overwrites the knowledge graph and advances applied bookkeeping to the + /// recovered snapshot's index. Called at startup before `Raft::new`. + pub(crate) async fn restore_applied_for_recovery( + &self, + meta: SnapshotMeta, + graph: KnowledgeGraph, + ) { + let mut inner = self.inner.lock().await; + *inner.knowledge_graph.write().await = graph; + inner.last_applied = meta.last_log_id; + inner.last_membership = meta.last_membership; + } +} + +#[cfg(test)] +impl EngStateMachineStore { + pub async fn apply_for_test(&mut self, index: u64, cmd: MemoryCommand) { + use openraft::{CommittedLeaderId, Entry, EntryPayload, LogId}; + self.apply(vec![Entry { + log_id: LogId::new(CommittedLeaderId::new(1, 1), index), + payload: EntryPayload::Normal(cmd), + }]) + .await + .unwrap(); + } +} + +fn sm_io_err(verb: ErrorVerb, msg: String) -> StorageError { + StorageError::from_io_error( + ErrorSubject::StateMachine, + verb, + io::Error::new(io::ErrorKind::Other, msg), + ) +} + +async fn build_payload(inner: &SmInner) -> Result<(EngramSnapshot, SnapshotMeta), StorageError> { + let short_term = inner + .short_term + .dump_all() + .await + .map_err(|e| sm_io_err(ErrorVerb::Read, e.to_string()))? + .into_iter() + .map(|(session_id, messages)| SessionMessages { session_id, messages }) + .collect(); + let core_memory = inner + .core_memory + .dump_all() + .await + .map_err(|e| sm_io_err(ErrorVerb::Read, e.to_string()))? + .into_iter() + .map(|(session_id, facts)| SessionFacts { session_id, facts }) + .collect(); + let knowledge_graph = inner.knowledge_graph.read().await.to_snapshot(); + + let payload = EngramSnapshot { + version: crate::raft::snapshot::SNAPSHOT_VERSION, + short_term, + core_memory, + knowledge_graph, + global_graph: None, + }; + let snapshot_id = format!( + "{}-{}", + inner.last_applied.as_ref().map(|l| l.index).unwrap_or(0), + inner.snapshot_idx + ); + let meta = SnapshotMeta { + last_log_id: inner.last_applied.clone(), + last_membership: inner.last_membership.clone(), + snapshot_id, + }; + Ok((payload, meta)) +} + +fn persist_snapshot(db: &Database, meta: &SnapshotMeta, data: &[u8]) -> Result<(), StorageError> { + let meta_bytes = serde_json::to_vec(meta).map_err(|e| sm_io_err(ErrorVerb::Write, e.to_string()))?; + let txn = db.begin_write().map_err(|e| sm_io_err(ErrorVerb::Write, e.to_string()))?; + { + let mut table = txn.open_table(SNAPSHOT_TABLE).map_err(|e| sm_io_err(ErrorVerb::Write, e.to_string()))?; + table.insert(SNAPSHOT_META_KEY, meta_bytes.as_slice()).map_err(|e| sm_io_err(ErrorVerb::Write, e.to_string()))?; + table.insert(SNAPSHOT_DATA_KEY, data).map_err(|e| sm_io_err(ErrorVerb::Write, e.to_string()))?; + } + txn.commit().map_err(|e| sm_io_err(ErrorVerb::Write, e.to_string()))?; + Ok(()) +} + +pub(crate) fn load_persisted_snapshot(db: &Database) -> Result>, StorageError> { + let txn = db.begin_read().map_err(|e| sm_io_err(ErrorVerb::Read, e.to_string()))?; + let table = txn.open_table(SNAPSHOT_TABLE).map_err(|e| sm_io_err(ErrorVerb::Read, e.to_string()))?; + let (Some(meta_g), Some(data_g)) = ( + table.get(SNAPSHOT_META_KEY).map_err(|e| sm_io_err(ErrorVerb::Read, e.to_string()))?, + table.get(SNAPSHOT_DATA_KEY).map_err(|e| sm_io_err(ErrorVerb::Read, e.to_string()))?, + ) else { + return Ok(None); + }; + let meta: SnapshotMeta = + serde_json::from_slice(meta_g.value()).map_err(|e| sm_io_err(ErrorVerb::Read, e.to_string()))?; + Ok(Some(Snapshot { + meta, + snapshot: Box::new(Cursor::new(data_g.value().to_vec())), + })) +} + +pub struct EngSnapshotBuilder { + inner: Arc>, +} + +impl RaftSnapshotBuilder for EngSnapshotBuilder { + async fn build_snapshot(&mut self) -> Result, StorageError> { + let mut inner = self.inner.lock().await; + inner.snapshot_idx += 1; + let (payload, meta) = build_payload(&inner).await?; + let data = payload.to_bytes().map_err(|e| sm_io_err(ErrorVerb::Write, e.to_string()))?; + persist_snapshot(&inner.db, &meta, &data)?; + Ok(Snapshot { meta, snapshot: Box::new(Cursor::new(data)) }) + } } impl RaftStateMachine for EngStateMachineStore { - // Snapshot not implemented in Stage 1. Nodes that fall too far behind must - // be manually removed and re-added to the cluster. Stage 2 adds compaction. - type SnapshotBuilder = NoOpSnapshotBuilder; + type SnapshotBuilder = EngSnapshotBuilder; async fn applied_state( &mut self, @@ -113,48 +263,50 @@ impl RaftStateMachine for EngStateMachineStore { } async fn get_snapshot_builder(&mut self) -> Self::SnapshotBuilder { - NoOpSnapshotBuilder + EngSnapshotBuilder { inner: self.inner.clone() } } async fn begin_receiving_snapshot( &mut self, ) -> Result>>, StorageError> { - Err(StorageError::from_io_error( - ErrorSubject::None, - ErrorVerb::Read, - io::Error::new(io::ErrorKind::Unsupported, "snapshots not implemented in Stage 1"), - )) + Ok(Box::new(Cursor::new(Vec::new()))) } async fn install_snapshot( &mut self, - _meta: &SnapshotMeta, - _snapshot: Box>>, + meta: &SnapshotMeta, + snapshot: Box>>, ) -> Result<(), StorageError> { - Err(StorageError::from_io_error( - ErrorSubject::None, - ErrorVerb::Write, - io::Error::new(io::ErrorKind::Unsupported, "snapshots not implemented in Stage 1"), - )) + let payload = EngramSnapshot::from_bytes(snapshot.get_ref()) + .map_err(|e| sm_io_err(ErrorVerb::Read, e.to_string()))?; + + // Clone store/graph handles under a short lock so we don't hold it across awaits. + let (short_term, core_memory, knowledge_graph, db) = { + let inner = self.inner.lock().await; + (inner.short_term.clone(), inner.core_memory.clone(), inner.knowledge_graph.clone(), inner.db.clone()) + }; + + let st_sessions = payload.short_term.into_iter().map(|s| (s.session_id, s.messages)).collect(); + short_term.restore_all(st_sessions).await.map_err(|e| sm_io_err(ErrorVerb::Write, e.to_string()))?; + let cm_sessions = payload.core_memory.into_iter().map(|s| (s.session_id, s.facts)).collect(); + core_memory.restore_all(cm_sessions).await.map_err(|e| sm_io_err(ErrorVerb::Write, e.to_string()))?; + *knowledge_graph.write().await = KnowledgeGraph::from_snapshot(payload.knowledge_graph); + + persist_snapshot(&db, meta, snapshot.get_ref())?; + + { + let mut inner = self.inner.lock().await; + inner.last_applied = meta.last_log_id.clone(); + inner.last_membership = meta.last_membership.clone(); + } + Ok(()) } async fn get_current_snapshot( &mut self, ) -> Result>, StorageError> { - Ok(None) - } -} - -/// Stub snapshot builder. Stage 2 will replace this with a real implementation. -pub struct NoOpSnapshotBuilder; - -impl RaftSnapshotBuilder for NoOpSnapshotBuilder { - async fn build_snapshot(&mut self) -> Result, StorageError> { - Err(StorageError::from_io_error( - ErrorSubject::None, - ErrorVerb::Read, - io::Error::new(io::ErrorKind::Unsupported, "snapshot building not implemented in Stage 1"), - )) + let inner = self.inner.lock().await; + load_persisted_snapshot(&inner.db) } } @@ -230,6 +382,8 @@ mod tests { mpsc::Receiver, mpsc::Receiver, Arc>, + Arc, + tempfile::TempDir, ) { let short_term = Arc::new(InMemoryStore::default()); let core_memory = Arc::new(InMemoryCoreMemoryStore::default()); @@ -237,15 +391,18 @@ mod tests { 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 dir = tempfile::tempdir().unwrap(); + let db = Arc::new(redb::Database::create(dir.path().join("sm.redb")).unwrap()); let sm = EngStateMachineStore::new( short_term.clone(), - core_memory, + core_memory.clone(), vector_store as Arc, embed_tx, kg.clone(), know_tx, + db, ); - (sm, short_term, embed_rx, know_rx, kg) + (sm, short_term, embed_rx, know_rx, kg, core_memory, dir) } fn make_entry(index: u64, cmd: MemoryCommand) -> openraft::Entry { @@ -257,7 +414,7 @@ mod tests { #[tokio::test] async fn add_message_writes_to_short_term() { - let (mut sm, short_term, _embed, _know, _kg) = make_sm(); + let (mut sm, short_term, _embed, _know, _kg, _cm, _dir) = make_sm(); sm.apply(vec![make_entry( 0, MemoryCommand::AddMessage { @@ -279,7 +436,7 @@ mod tests { #[tokio::test] async fn add_message_enqueues_embedding_job() { - let (mut sm, _st, mut embed_rx, _know, _kg) = make_sm(); + let (mut sm, _st, mut embed_rx, _know, _kg, _cm, _dir) = make_sm(); sm.apply(vec![make_entry( 0, MemoryCommand::AddMessage { @@ -300,7 +457,7 @@ mod tests { #[tokio::test] async fn delete_session_clears_redis_and_enqueues_lancedb_delete() { - let (mut sm, short_term, mut embed_rx, _know, _kg) = make_sm(); + let (mut sm, short_term, mut embed_rx, _know, _kg, _cm, _dir) = make_sm(); sm.apply(vec![ make_entry( 0, @@ -327,7 +484,7 @@ mod tests { #[tokio::test] async fn noop_command_is_ignored() { - let (mut sm, short_term, _embed, _know, _kg) = make_sm(); + let (mut sm, short_term, _embed, _know, _kg, _cm, _dir) = 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); @@ -335,7 +492,7 @@ mod tests { #[tokio::test] async fn add_message_enqueues_knowledge_job() { - let (mut sm, _st, _embed, mut know_rx, _kg) = make_sm(); + let (mut sm, _st, _embed, mut know_rx, _kg, _cm, _dir) = make_sm(); sm.apply(vec![make_entry(0, MemoryCommand::AddMessage { session_id: "s1".into(), message: MessagePayload { @@ -352,7 +509,7 @@ mod tests { #[tokio::test] async fn add_knowledge_updates_graph() { - let (mut sm, _st, _embed, _know, kg) = make_sm(); + let (mut sm, _st, _embed, _know, kg, _cm, _dir) = make_sm(); sm.apply(vec![make_entry(0, MemoryCommand::AddKnowledge { session_id: "s1".into(), message_id: "m1".into(), @@ -373,7 +530,7 @@ mod tests { #[tokio::test] async fn delete_session_clears_knowledge_graph() { - let (mut sm, _st, _embed, _know, kg) = make_sm(); + let (mut sm, _st, _embed, _know, kg, _cm, _dir) = 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() }], @@ -384,4 +541,98 @@ mod tests { let kg = kg.read().await; assert!(kg.all_entities("s1").is_empty()); } + + #[tokio::test] + async fn install_snapshot_sets_last_applied_to_meta_log_id() { + let (mut src, _st, _e, _k, _kg, _cm, _dir) = make_sm(); + src.apply(vec![make_entry(7, MemoryCommand::AddFact { + session_id: "s1".into(), fact: "f".into(), + })]).await.unwrap(); + let mut builder = src.get_snapshot_builder().await; + let snap = builder.build_snapshot().await.unwrap(); + + let (mut dst, dst_st, _e2, _k2, dst_kg, _cm2, _dir2) = make_sm(); + let mut buf = dst.begin_receiving_snapshot().await.unwrap(); + *buf = std::io::Cursor::new(snap.snapshot.get_ref().clone()); + dst.install_snapshot(&snap.meta, buf).await.unwrap(); + + let (applied, _membership) = dst.applied_state().await.unwrap(); + assert_eq!(applied.unwrap().index, 7); + assert_eq!(dst_st.get_recent("s1", 10).await.unwrap().len(), 0); + let _ = dst_kg.read().await; + } + + #[tokio::test] + async fn apply_build_install_reproduces_state() { + let (mut src, _st, _e, _k, _kg, src_cm, _dir) = make_sm(); + src.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: "Bob".into(), entity_type: "Person".into(), attributes: HashMap::new() }, + ], + relationships: vec![Relationship { from: "Alice".into(), to: "Bob".into(), relationship_type: "knows".into() }], + }), + make_entry(1, MemoryCommand::AddFact { session_id: "s1".into(), fact: "likes tea".into() }), + ]).await.unwrap(); + let src_facts = src_cm.get_facts("s1").await.unwrap(); + + let mut builder = src.get_snapshot_builder().await; + let snap = builder.build_snapshot().await.unwrap(); + + let (mut dst, _st2, _e2, _k2, dst_kg, dst_cm, _dir2) = make_sm(); + let mut buf = dst.begin_receiving_snapshot().await.unwrap(); + *buf = std::io::Cursor::new(snap.snapshot.get_ref().clone()); + dst.install_snapshot(&snap.meta, buf).await.unwrap(); + + assert_eq!(dst_cm.get_facts("s1").await.unwrap(), src_facts); + let path = dst_kg.read().await.find_path("s1", "Alice", "Bob").unwrap(); + assert_eq!(path.len(), 1); + } + + #[tokio::test] + async fn build_snapshot_meta_index_equals_last_applied() { + let (mut sm, _st, _e, _k, _kg, _cm, _dir) = make_sm(); + for i in 0..=4u64 { + sm.apply(vec![make_entry(i, MemoryCommand::AddFact { + session_id: "s1".into(), fact: format!("f{i}"), + })]).await.unwrap(); + } + let mut builder = sm.get_snapshot_builder().await; + let snap = builder.build_snapshot().await.unwrap(); + assert_eq!(snap.meta.last_log_id.unwrap().index, 4); + } + + #[tokio::test] + async fn build_then_get_current_snapshot_returns_same_index() { + let (mut sm, _st, _e, _k, _kg, _cm, _dir) = make_sm(); + sm.apply(vec![make_entry(0, MemoryCommand::AddFact { + session_id: "s1".into(), fact: "f".into(), + })]).await.unwrap(); + let mut builder = sm.get_snapshot_builder().await; + let built = builder.build_snapshot().await.unwrap(); + let current = sm.get_current_snapshot().await.unwrap().expect("snapshot persisted"); + assert_eq!(current.meta.last_log_id, built.meta.last_log_id); + } + + #[tokio::test] + async fn snapshot_payload_contains_applied_state() { + let (mut sm, _st, _e, _k, _kg, _cm, _dir) = 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::AddFact { + session_id: "s1".into(), fact: "likes coffee".into(), + })]).await.unwrap(); + + let mut builder = sm.get_snapshot_builder().await; + let snap = builder.build_snapshot().await.unwrap(); + let payload = crate::raft::snapshot::EngramSnapshot::from_bytes(snap.snapshot.get_ref()).unwrap(); + assert_eq!(payload.version, 1); + assert!(payload.knowledge_graph.sessions.iter().any(|s| s.session_id == "s1")); + assert!(payload.core_memory.iter().any(|s| s.facts.contains(&"likes coffee".to_string()))); + } } diff --git a/src/stores/redis_core_memory.rs b/src/stores/redis_core_memory.rs index afd18ff..e6e0445 100644 --- a/src/stores/redis_core_memory.rs +++ b/src/stores/redis_core_memory.rs @@ -1,6 +1,7 @@ use std::error::Error as StdError; use async_trait::async_trait; +use futures::StreamExt; use redis::{AsyncCommands, Client, aio::MultiplexedConnection}; use crate::core::{CoreMemoryStore, MemoryError}; @@ -57,8 +58,108 @@ impl CoreMemoryStore for RedisCoreMemoryStore { .map_err(memory_error)?; Ok(()) } + + async fn dump_all(&self) -> Result)>, MemoryError> { + let mut connection = self.connection.clone(); + let keys: Vec = { + let mut iter = connection + .scan_match::<_, String>("session:*:core_memory") + .await + .map_err(memory_error)?; + let mut collected = Vec::new(); + while let Some(key) = iter.next().await { + collected.push(key); + } + collected + }; + let mut out = Vec::new(); + for key in keys { + let session_id = core_session_id_from_key(&key); + let facts: Vec = connection.smembers(key.as_str()).await.map_err(memory_error)?; + out.push((session_id, facts)); + } + Ok(out) + } + + async fn restore_all(&self, sessions: Vec<(String, Vec)>) -> Result<(), MemoryError> { + let mut connection = self.connection.clone(); + let existing: Vec = { + let mut iter = connection + .scan_match::<_, String>("session:*:core_memory") + .await + .map_err(memory_error)?; + let mut collected = Vec::new(); + while let Some(key) = iter.next().await { + collected.push(key); + } + collected + }; + for key in existing { + let _: usize = connection.del(key.as_str()).await.map_err(memory_error)?; + } + for (session_id, facts) in sessions { + for fact in facts { + self.add_fact(&session_id, &fact).await?; + } + } + Ok(()) + } +} + +fn core_session_id_from_key(key: &str) -> String { + key.strip_prefix("session:") + .and_then(|s| s.strip_suffix(":core_memory")) + .unwrap_or(key) + .to_string() } fn memory_error(error: impl StdError + Send + Sync + 'static) -> MemoryError { MemoryError::Other(Box::new(error)) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::core::CoreMemoryStore; + use testcontainers::{ + GenericImage, + core::{IntoContainerPort, WaitFor}, + runners::AsyncRunner, + }; + + const REDIS_PORT: u16 = 6379; + + async fn test_store() -> (RedisCoreMemoryStore, testcontainers::ContainerAsync) { + let node = GenericImage::new("redis", "7.2.4") + .with_exposed_port(REDIS_PORT.tcp()) + .with_wait_for(WaitFor::message_on_stdout("Ready to accept connections")) + .start() + .await + .unwrap(); + let host = node.get_host().await.unwrap(); + let port = node.get_host_port_ipv4(REDIS_PORT.tcp()).await.unwrap(); + let url = format!("redis://{host}:{port}/"); + let store = RedisCoreMemoryStore::connect(&url).await.unwrap(); + (store, node) + } + + #[tokio::test] + async fn dump_all_and_restore_all_round_trip() { + let (store, _node) = test_store().await; + store.add_fact("s1", "fact-a").await.unwrap(); + store.add_fact("s1", "fact-b").await.unwrap(); + store.add_fact("s2", "fact-c").await.unwrap(); + + let dump = store.dump_all().await.unwrap(); + assert_eq!(dump.len(), 2); + + store.add_fact("stale", "old").await.unwrap(); + store.restore_all(dump).await.unwrap(); + + assert!(store.get_facts("stale").await.unwrap().is_empty()); + let mut s1_facts = store.get_facts("s1").await.unwrap(); + s1_facts.sort(); + assert_eq!(s1_facts, vec!["fact-a".to_string(), "fact-b".to_string()]); + assert_eq!(store.get_facts("s2").await.unwrap(), vec!["fact-c".to_string()]); + } +} diff --git a/src/stores/redis_shortterm.rs b/src/stores/redis_shortterm.rs index 7317abc..fcbeba3 100644 --- a/src/stores/redis_shortterm.rs +++ b/src/stores/redis_shortterm.rs @@ -1,6 +1,7 @@ use std::error::Error as StdError; use async_trait::async_trait; +use futures::StreamExt; use redis::{AsyncCommands, Client, aio::MultiplexedConnection}; use crate::core::{MemoryError, ShortTermMemory, TokenCounter, trim_messages_to_token_budget}; @@ -197,8 +198,119 @@ impl ShortTermMemory for RedisShortTermMemory { .into_iter() .find(|message| message.id.as_deref() == Some(message_id))) } + + async fn dump_all(&self) -> Result)>, MemoryError> { + let mut connection = self.connection.clone(); + let keys: Vec = { + let mut iter = connection + .scan_match::<_, String>("session:*:messages") + .await + .map_err(memory_error)?; + let mut collected = Vec::new(); + while let Some(key) = iter.next().await { + collected.push(key); + } + collected + }; + let mut out = Vec::new(); + for key in keys { + let session_id = session_id_from_key(&key); + let raw: Vec = connection.lrange(key.as_str(), 0, -1).await.map_err(memory_error)?; + let messages = raw + .into_iter() + .map(|r| serde_json::from_str(&r).map_err(memory_error)) + .collect::, _>>()?; + out.push((session_id, messages)); + } + Ok(out) + } + + async fn restore_all(&self, sessions: Vec<(String, Vec)>) -> Result<(), MemoryError> { + let mut connection = self.connection.clone(); + let existing: Vec = { + let mut iter = connection + .scan_match::<_, String>("session:*:messages") + .await + .map_err(memory_error)?; + let mut collected = Vec::new(); + while let Some(key) = iter.next().await { + collected.push(key); + } + collected + }; + for key in existing { + let _: usize = connection.del(key.as_str()).await.map_err(memory_error)?; + } + for (session_id, messages) in sessions { + self.write_messages(&session_id, &messages).await?; + } + Ok(()) + } +} + +fn session_id_from_key(key: &str) -> String { + key.strip_prefix("session:") + .and_then(|s| s.strip_suffix(":messages")) + .unwrap_or(key) + .to_string() } fn memory_error(error: impl StdError + Send + Sync + 'static) -> MemoryError { MemoryError::Other(Box::new(error)) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::core::ShortTermMemory; + use crate::models::Message; + use testcontainers::{ + GenericImage, + core::{IntoContainerPort, WaitFor}, + runners::AsyncRunner, + }; + + const REDIS_PORT: u16 = 6379; + + async fn test_store() -> (RedisShortTermMemory, testcontainers::ContainerAsync) { + let node = GenericImage::new("redis", "7.2.4") + .with_exposed_port(REDIS_PORT.tcp()) + .with_wait_for(WaitFor::message_on_stdout("Ready to accept connections")) + .start() + .await + .unwrap(); + let host = node.get_host().await.unwrap(); + let port = node.get_host_port_ipv4(REDIS_PORT.tcp()).await.unwrap(); + let url = format!("redis://{host}:{port}/"); + let store = RedisShortTermMemory::connect(&url).await.unwrap(); + (store, node) + } + + fn sample_message(content: &str) -> Message { + Message { + id: None, + role: "user".to_string(), + content: content.to_string(), + timestamp: None, + embedding_status: None, + } + } + + #[tokio::test] + async fn dump_all_and_restore_all_round_trip() { + let (store, _node) = test_store().await; + store.add_message("s1", sample_message("hi")).await.unwrap(); + store.add_message("s2", sample_message("yo")).await.unwrap(); + + let dump = store.dump_all().await.unwrap(); + assert_eq!(dump.len(), 2); + + // Wipe and restore into the same Redis. + store.add_message("stale", sample_message("old")).await.unwrap(); + store.restore_all(dump).await.unwrap(); + + assert!(store.get_recent("stale", 10).await.unwrap().is_empty()); + assert_eq!(store.get_recent("s1", 10).await.unwrap().len(), 1); + assert_eq!(store.get_recent("s2", 10).await.unwrap().len(), 1); + } +} diff --git a/tests/e2e_test.rs b/tests/e2e_test.rs index fdfd11c..4b8fc2f 100644 --- a/tests/e2e_test.rs +++ b/tests/e2e_test.rs @@ -131,6 +131,8 @@ async fn e2e_flow_uses_real_stores_and_background_worker() { knowledge_max_workers: 4, knowledge_channel_size: 500, knowledge_extractor: engram::config::KnowledgeExtractorType::OpenAI, + raft_db_path: std::path::PathBuf::from("./data/raft/engram.redb"), + snapshot_log_threshold: 1000, }; 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 7f014f9..3137a88 100644 --- a/tests/integration_test.rs +++ b/tests/integration_test.rs @@ -94,6 +94,8 @@ async fn setup_test_app() -> TestApp { knowledge_max_workers: 4, knowledge_channel_size: 500, knowledge_extractor: engram::config::KnowledgeExtractorType::OpenAI, + raft_db_path: std::path::PathBuf::from("./data/raft/engram.redb"), + snapshot_log_threshold: 1000, }; let embedding_provider: Arc = Arc::new( diff --git a/tests/raft_write_test.rs b/tests/raft_write_test.rs index 5854473..f53e164 100644 --- a/tests/raft_write_test.rs +++ b/tests/raft_write_test.rs @@ -9,13 +9,16 @@ use engram::app::build_raft_node; use engram::config::Config; use engram::core::{InMemoryCoreMemoryStore, InMemoryStore, InMemoryVectorStore, ShortTermMemory}; use engram::raft::types::{MemoryCommand, MessagePayload}; +use tempfile; #[tokio::test] async fn single_node_raft_write_commits_to_state_machine() { let short_term = Arc::new(InMemoryStore::default()); let (tx, _rx) = mpsc::channel(10); + let raft_dir = tempfile::tempdir().unwrap(); let config = Config { node_id: Some(1), + raft_db_path: raft_dir.path().join("engram.redb"), ..Config::default() }; let knowledge_graph = Arc::new(tokio::sync::RwLock::new(engram::knowledge::graph::KnowledgeGraph::new()));