diff --git a/server/src/handlers.rs b/server/src/handlers.rs index c9f233d..e555c04 100644 --- a/server/src/handlers.rs +++ b/server/src/handlers.rs @@ -366,74 +366,72 @@ impl MemoryHandler { let result: Result = match name { "query_graph_path" => { let req = parse_tool!(args.clone(), id, crate::tools::QueryGraphPathTool); - let graph = self.state.get_full_graph(); - let max_depth = req.max_depth.unwrap_or(5); - let mut queue = std::collections::VecDeque::new(); - let mut visited = std::collections::HashSet::new(); - let mut parents: std::collections::HashMap = - std::collections::HashMap::new(); + self.state.read_graph(|graph| { + let max_depth = req.max_depth.unwrap_or(5); + let mut queue = std::collections::VecDeque::new(); + let mut visited = std::collections::HashSet::new(); + let mut parents: std::collections::HashMap = + std::collections::HashMap::new(); - queue.push_back(req.start_node.clone()); - visited.insert(req.start_node.clone()); + queue.push_back(req.start_node.clone()); + visited.insert(req.start_node.clone()); - let mut found = false; - let mut current_depth = 0; - let mut nodes_at_current_depth = 1; - let mut nodes_at_next_depth = 0; + let mut found = false; + let mut current_depth = 0; + let mut nodes_at_current_depth = 1; + let mut nodes_at_next_depth = 0; - while let Some(current) = queue.pop_front() { - if current == req.end_node { - found = true; - break; - } - nodes_at_current_depth -= 1; - if current_depth < max_depth { - for rel in &graph.relations { - if rel.from == current && !visited.contains(&rel.to) { - visited.insert(rel.to.clone()); - parents.insert( - rel.to.clone(), - (current.clone(), rel.relation_type.clone()), - ); - queue.push_back(rel.to.clone()); - nodes_at_next_depth += 1; - } else if rel.to == current && !visited.contains(&rel.from) { - visited.insert(rel.from.clone()); - parents.insert( - rel.from.clone(), - ( - current.clone(), - format!("inverse({})", rel.relation_type), - ), - ); - queue.push_back(rel.from.clone()); - nodes_at_next_depth += 1; + while let Some(current) = queue.pop_front() { + if current == req.end_node { + found = true; + break; + } + nodes_at_current_depth -= 1; + if current_depth < max_depth { + for rel in &graph.relations { + if rel.from == current && !visited.contains(&rel.to) { + visited.insert(rel.to.clone()); + parents.insert( + rel.to.clone(), + (current.clone(), rel.relation_type.clone()), + ); + queue.push_back(rel.to.clone()); + nodes_at_next_depth += 1; + } else if rel.to == current && !visited.contains(&rel.from) { + visited.insert(rel.from.clone()); + parents.insert( + rel.from.clone(), + ( + current.clone(), + format!("inverse({})", rel.relation_type), + ), + ); + queue.push_back(rel.from.clone()); + nodes_at_next_depth += 1; + } } } + if nodes_at_current_depth == 0 { + current_depth += 1; + nodes_at_current_depth = nodes_at_next_depth; + nodes_at_next_depth = 0; + } } - if nodes_at_current_depth == 0 { - current_depth += 1; - nodes_at_current_depth = nodes_at_next_depth; - nodes_at_next_depth = 0; - } - } - if found { - let mut path = Vec::new(); - let mut curr = req.end_node.clone(); - while curr != req.start_node { - let (parent, rel) = parents.get(&curr).unwrap(); - path.push(format!("({}) --[{}]--> ({})", parent, rel, curr)); - curr = parent.clone(); + if found { + let mut path = Vec::new(); + let mut curr = req.end_node.clone(); + while curr != req.start_node { + let (parent, rel_type) = parents.get(&curr).unwrap().clone(); + path.push(format!("{} -[{}]-> {}", parent, rel_type, curr)); + curr = parent; + } + path.reverse(); + Ok(format!("Path found:\n{}", path.join("\n"))) + } else { + Ok(format!("No path found between {} and {} within depth {}", req.start_node, req.end_node, max_depth)) } - path.reverse(); - Ok(format!("Path found:\n{}", path.join("\n"))) - } else { - Ok(format!( - "No path found between {} and {} within depth {}", - req.start_node, req.end_node, max_depth - )) - } + }) } "create_entities" => { let req = parse_tool!(args.clone(), id, CreateEntitiesTool); @@ -462,22 +460,10 @@ impl MemoryHandler { } "add_observations" => { let req = parse_tool!(args.clone(), id, AddObservationsTool); - let full = self.state.get_full_graph(); self.state.modify_graph(|g| { for o in req.observations { - if let Some(full_e) = full.entities.get(&o.entity_name) { - let mut e = - g.entities.get(&o.entity_name).cloned().unwrap_or_else( - || Entity { - name: o.entity_name.clone(), - entity_type: full_e.entity_type.clone(), - observations: vec![], - namespace: full_e.namespace.clone(), - git_branch: None, - }, - ); + if let Some(e) = g.entities.get_mut(&o.entity_name) { e.observations.extend(o.contents); - g.entities.insert(o.entity_name, e); } } }); @@ -518,13 +504,25 @@ impl MemoryHandler { } "read_graph" => { let req = parse_tool!(args.clone(), id, ReadGraphTool); - let mut full = self.state.get_full_graph(); - if let Some(ns) = req.namespace { - full.entities.retain(|_, e| e.namespace == ns); - full.relations.retain(|r| r.namespace == ns); - } - let data = serde_json::to_string(&full).unwrap_or_default(); - Ok(data.to_string()) + let data = self.state.read_graph(|full| { + if let Some(ns) = req.namespace { + let mut filtered = KnowledgeGraph::default(); + for (k, v) in &full.entities { + if v.namespace == ns { + filtered.entities.insert(k.clone(), v.clone()); + } + } + for r in &full.relations { + if r.namespace == ns { + filtered.relations.push(r.clone()); + } + } + serde_json::to_string(&filtered).unwrap_or_default() + } else { + serde_json::to_string(full).unwrap_or_default() + } + }); + Ok(data) } "search_nodes" => { let req = parse_tool!(args.clone(), id, SearchNodesTool); @@ -536,37 +534,39 @@ impl MemoryHandler { }; let mut result = KnowledgeGraph::default(); - let full = self.state.get_full_graph(); - for (id, doc_type, _, _, _) in matches { - if doc_type == "entity" - && let Some(e) = full.entities.get(&id) - { - result.entities.insert(id, e.clone()); + self.state.read_graph(|full| { + for (id, doc_type, _, _, _) in matches { + if doc_type == "entity" + && let Some(e) = full.entities.get(&id) + { + result.entities.insert(id, e.clone()); + } } - } + }); let data = serde_json::to_string(&result).unwrap_or_default(); Ok(data.to_string()) } "open_nodes" => { let req = parse_tool!(args.clone(), id, OpenNodesTool); let targets: HashSet<_> = req.names.into_iter().collect(); - let full = self.state.get_full_graph(); let mut result = KnowledgeGraph::default(); let mut connected = HashSet::new(); - for r in &full.relations { - if targets.contains(&r.from) { - connected.insert(r.to.clone()); - result.relations.push(r.clone()); - } else if targets.contains(&r.to) { - connected.insert(r.from.clone()); - result.relations.push(r.clone()); + self.state.read_graph(|full| { + for r in &full.relations { + if targets.contains(&r.from) { + connected.insert(r.to.clone()); + result.relations.push(r.clone()); + } else if targets.contains(&r.to) { + connected.insert(r.from.clone()); + result.relations.push(r.clone()); + } } - } - for (name, e) in full.entities { - if targets.contains(&name) || connected.contains(&name) { - result.entities.insert(name, e); + for (name, e) in &full.entities { + if targets.contains(name) || connected.contains(name) { + result.entities.insert(name.clone(), e.clone()); + } } - } + }); let data = serde_json::to_string(&result).unwrap_or_default(); Ok(data.to_string()) } @@ -594,37 +594,40 @@ impl MemoryHandler { "visualize_graph" => { let req = parse_tool!(args.clone(), id, VisualizeGraphTool); let query = req.query.unwrap_or_default().to_lowercase(); - let full = self.state.get_full_graph(); let mut included = HashSet::new(); - for (name, e) in &full.entities { - if let Some(ns) = &req.namespace - && e.namespace != *ns - { - continue; - } - if query.is_empty() - || name.to_lowercase().contains(&query) - || e.entity_type.to_lowercase().contains(&query) - { - included.insert(name.clone()); - } - } let mut to_draw = Vec::new(); - for r in &full.relations { - if let Some(ns) = &req.namespace - && r.namespace != *ns - { - continue; + + self.state.read_graph(|full| { + for (name, e) in &full.entities { + if let Some(ns) = &req.namespace + && e.namespace != *ns + { + continue; + } + if query.is_empty() + || name.to_lowercase().contains(&query) + || e.entity_type.to_lowercase().contains(&query) + { + included.insert(name.clone()); + } } - if query.is_empty() - || included.contains(&r.from) - || included.contains(&r.to) - { - included.insert(r.from.clone()); - included.insert(r.to.clone()); - to_draw.push(r); + + for r in &full.relations { + if let Some(ns) = &req.namespace + && r.namespace != *ns + { + continue; + } + if query.is_empty() + || included.contains(&r.from) + || included.contains(&r.to) + { + included.insert(r.from.clone()); + included.insert(r.to.clone()); + to_draw.push(r.clone()); + } } - } + }); use std::fmt::Write; let mut output = String::with_capacity(included.len() * 40 + to_draw.len() * 60); output.push_str("graph TD;\n"); @@ -1104,18 +1107,18 @@ impl MemoryHandler { Ok("Entities merged".to_string()) } "find_orphans" => { - let full = self.state.get_full_graph(); - let mut connected = std::collections::HashSet::new(); - for r in &full.relations { - connected.insert(r.from.clone()); - connected.insert(r.to.clone()); - } - let orphans: Vec = full - .entities - .keys() - .filter(|k| !connected.contains(*k)) - .cloned() - .collect(); + let orphans = self.state.read_graph(|full| { + let mut connected = std::collections::HashSet::new(); + for r in &full.relations { + connected.insert(r.from.clone()); + connected.insert(r.to.clone()); + } + full.entities + .keys() + .filter(|k| !connected.contains(*k)) + .cloned() + .collect::>() + }); let data = serde_json::to_string(&orphans).unwrap_or_default(); Ok(data.to_string()) } @@ -1522,14 +1525,15 @@ impl MemoryHandler { let mut snippets = Vec::new(); let mut adrs = Vec::new(); - let full = self.state.get_full_graph(); - for (id, doc_type, _, _, _) in &matches { - if doc_type == "entity" - && let Some(e) = full.entities.get(id) - { - kg.entities.insert(id.clone(), e.clone()); + self.state.read_graph(|full| { + for (id, doc_type, _, _, _) in &matches { + if doc_type == "entity" + && let Some(e) = full.entities.get(id) + { + kg.entities.insert(id.clone(), e.clone()); + } } - } + }); for t in self.state.tasks.read() { if matches.iter().any(|(id, typ, _, _, _)| id == &t.id && typ == "task") { tasks.push(t); diff --git a/server/src/main.rs b/server/src/main.rs index c5216d8..36d6d20 100644 --- a/server/src/main.rs +++ b/server/src/main.rs @@ -214,12 +214,10 @@ async fn gate_set_handler( (axum::http::StatusCode::OK, "Gate state updated.").into_response() } -fn run_server(state: Arc) -> Result<(), Box> { - let rt = tokio::runtime::Runtime::new().unwrap(); - rt.block_on(async { - state.rebuild_index().await; - tokio::spawn(index_committer_worker(Arc::clone(&state))); - let app_state = Arc::new(AppState { +async fn run_server(state: Arc) -> Result<(), Box> { + state.rebuild_index().await; + tokio::spawn(index_committer_worker(Arc::clone(&state))); + let app_state = Arc::new(AppState { handler: Arc::new(MemoryHandler { state: Arc::clone(&state), }), @@ -427,7 +425,6 @@ fn run_server(state: Arc) -> Result<(), Box> let _ = tokio::fs::write(&log_path, format!("Server crashed: {}\n", e)).await; } Ok(()) - }) } async fn ws_handler( @@ -796,6 +793,9 @@ fn main() -> Result<(), Box> { write_txn.commit().unwrap(); } + let rt = tokio::runtime::Runtime::new().unwrap(); + let _guard = rt.enter(); + let state = Arc::new(MemoryState { graph: Store::new("knowledge_graph_master", db.clone()), base_dir: base.clone(), @@ -830,5 +830,5 @@ fn main() -> Result<(), Box> { activity_tx: tokio::sync::broadcast::channel(100).0, }); - run_server(state) + rt.block_on(run_server(state)) } diff --git a/server/src/state.rs b/server/src/state.rs index 60e9561..126b0d4 100644 --- a/server/src/state.rs +++ b/server/src/state.rs @@ -49,6 +49,13 @@ impl MemoryState { self.graph.read() } + pub fn read_graph(&self, f: F) -> R + where + F: FnOnce(&KnowledgeGraph) -> R, + { + self.graph.read_with(f) + } + pub fn modify_graph(&self, update_fn: F) { self.graph.modify(update_fn); } diff --git a/server/src/store.rs b/server/src/store.rs index 60d39ee..700fc58 100644 --- a/server/src/store.rs +++ b/server/src/store.rs @@ -5,18 +5,35 @@ use std::sync::{Arc, RwLock}; pub const STORE_TABLE: TableDefinition<&str, &[u8]> = TableDefinition::new("store"); pub struct Store { - pub key: String, - pub db: Arc, pub cache: RwLock, + tx: tokio::sync::mpsc::UnboundedSender>, } impl Store { pub fn new(key: &str, db: Arc) -> Self { let initial_data = Self::load_from_db(key, &db); + let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::>(); + + let db_clone = db.clone(); + let key_clone = key.to_string(); + tokio::spawn(async move { + while let Some(json_data) = rx.recv().await { + let db_inner = db_clone.clone(); + let key_inner = key_clone.clone(); + let _ = tokio::task::spawn_blocking(move || { + let write_txn = db_inner.begin_write().unwrap(); + { + let mut table = write_txn.open_table(STORE_TABLE).unwrap(); + table.insert(key_inner.as_str(), json_data.as_slice()).unwrap(); + } + write_txn.commit().unwrap(); + }).await; + } + }); + Self { - key: key.to_string(), - db, cache: RwLock::new(initial_data), + tx, } } @@ -35,23 +52,22 @@ impl Store(&self, f: F) -> R + where + F: FnOnce(&T) -> R, + { + let lock = self.cache.read().unwrap(); + f(&lock) + } + pub fn modify(&self, f: F) { - let (key, db, json_data) = { + let json_data = { let mut lock = self.cache.write().unwrap(); f(&mut lock); // Serialize while holding lock to avoid expensive deep clone of T - let json = serde_json::to_vec(&*lock).unwrap(); - (self.key.clone(), self.db.clone(), json) + serde_json::to_vec(&*lock).unwrap() }; - - tokio::task::spawn_blocking(move || { - let write_txn = db.begin_write().unwrap(); - { - let mut table = write_txn.open_table(STORE_TABLE).unwrap(); - table.insert(key.as_str(), json_data.as_slice()).unwrap(); - } - write_txn.commit().unwrap(); - }); + let _ = self.tx.send(json_data); } }