Optimize handlers to avoid deep clones of KnowledgeGraph, and fix database write concurrency issues

This commit is contained in:
Riz Ashraf committed 2026-09-21 06:53:32 +01:00
1 parent 7a48fa5d34
commit 9b349e6459
4 files changed
+87 -60

No files matched your search

+43 -39
View File
@@ -366,7 +366,7 @@ impl MemoryHandler {
let result: Result<String, String> = match name { let result: Result<String, String> = match name {
"query_graph_path" => { "query_graph_path" => {
let req = parse_tool!(args.clone(), id, crate::tools::QueryGraphPathTool); let req = parse_tool!(args.clone(), id, crate::tools::QueryGraphPathTool);
let graph = self.state.get_full_graph(); self.state.read_graph(|graph| {
let max_depth = req.max_depth.unwrap_or(5); let max_depth = req.max_depth.unwrap_or(5);
let mut queue = std::collections::VecDeque::new(); let mut queue = std::collections::VecDeque::new();
let mut visited = std::collections::HashSet::new(); let mut visited = std::collections::HashSet::new();
@@ -422,18 +422,16 @@ impl MemoryHandler {
let mut path = Vec::new(); let mut path = Vec::new();
let mut curr = req.end_node.clone(); let mut curr = req.end_node.clone();
while curr != req.start_node { while curr != req.start_node {
let (parent, rel) = parents.get(&curr).unwrap(); let (parent, rel_type) = parents.get(&curr).unwrap().clone();
path.push(format!("({}) --[{}]--> ({})", parent, rel, curr)); path.push(format!("{} -[{}]-> {}", parent, rel_type, curr));
curr = parent.clone(); curr = parent;
} }
path.reverse(); path.reverse();
Ok(format!("Path found:\n{}", path.join("\n"))) Ok(format!("Path found:\n{}", path.join("\n")))
} else { } else {
Ok(format!( Ok(format!("No path found between {} and {} within depth {}", req.start_node, req.end_node, max_depth))
"No path found between {} and {} within depth {}",
req.start_node, req.end_node, max_depth
))
} }
})
} }
"create_entities" => { "create_entities" => {
let req = parse_tool!(args.clone(), id, CreateEntitiesTool); let req = parse_tool!(args.clone(), id, CreateEntitiesTool);
@@ -462,22 +460,10 @@ impl MemoryHandler {
} }
"add_observations" => { "add_observations" => {
let req = parse_tool!(args.clone(), id, AddObservationsTool); let req = parse_tool!(args.clone(), id, AddObservationsTool);
let full = self.state.get_full_graph();
self.state.modify_graph(|g| { self.state.modify_graph(|g| {
for o in req.observations { for o in req.observations {
if let Some(full_e) = full.entities.get(&o.entity_name) { if let Some(e) = g.entities.get_mut(&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,
},
);
e.observations.extend(o.contents); e.observations.extend(o.contents);
g.entities.insert(o.entity_name, e);
} }
} }
}); });
@@ -518,13 +504,25 @@ impl MemoryHandler {
} }
"read_graph" => { "read_graph" => {
let req = parse_tool!(args.clone(), id, ReadGraphTool); let req = parse_tool!(args.clone(), id, ReadGraphTool);
let mut full = self.state.get_full_graph(); let data = self.state.read_graph(|full| {
if let Some(ns) = req.namespace { if let Some(ns) = req.namespace {
full.entities.retain(|_, e| e.namespace == ns); let mut filtered = KnowledgeGraph::default();
full.relations.retain(|r| r.namespace == ns); for (k, v) in &full.entities {
if v.namespace == ns {
filtered.entities.insert(k.clone(), v.clone());
} }
let data = serde_json::to_string(&full).unwrap_or_default(); }
Ok(data.to_string()) 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" => { "search_nodes" => {
let req = parse_tool!(args.clone(), id, SearchNodesTool); let req = parse_tool!(args.clone(), id, SearchNodesTool);
@@ -536,7 +534,7 @@ impl MemoryHandler {
}; };
let mut result = KnowledgeGraph::default(); let mut result = KnowledgeGraph::default();
let full = self.state.get_full_graph(); self.state.read_graph(|full| {
for (id, doc_type, _, _, _) in matches { for (id, doc_type, _, _, _) in matches {
if doc_type == "entity" if doc_type == "entity"
&& let Some(e) = full.entities.get(&id) && let Some(e) = full.entities.get(&id)
@@ -544,15 +542,16 @@ impl MemoryHandler {
result.entities.insert(id, e.clone()); result.entities.insert(id, e.clone());
} }
} }
});
let data = serde_json::to_string(&result).unwrap_or_default(); let data = serde_json::to_string(&result).unwrap_or_default();
Ok(data.to_string()) Ok(data.to_string())
} }
"open_nodes" => { "open_nodes" => {
let req = parse_tool!(args.clone(), id, OpenNodesTool); let req = parse_tool!(args.clone(), id, OpenNodesTool);
let targets: HashSet<_> = req.names.into_iter().collect(); let targets: HashSet<_> = req.names.into_iter().collect();
let full = self.state.get_full_graph();
let mut result = KnowledgeGraph::default(); let mut result = KnowledgeGraph::default();
let mut connected = HashSet::new(); let mut connected = HashSet::new();
self.state.read_graph(|full| {
for r in &full.relations { for r in &full.relations {
if targets.contains(&r.from) { if targets.contains(&r.from) {
connected.insert(r.to.clone()); connected.insert(r.to.clone());
@@ -562,11 +561,12 @@ impl MemoryHandler {
result.relations.push(r.clone()); result.relations.push(r.clone());
} }
} }
for (name, e) in full.entities { for (name, e) in &full.entities {
if targets.contains(&name) || connected.contains(&name) { if targets.contains(name) || connected.contains(name) {
result.entities.insert(name, e); result.entities.insert(name.clone(), e.clone());
} }
} }
});
let data = serde_json::to_string(&result).unwrap_or_default(); let data = serde_json::to_string(&result).unwrap_or_default();
Ok(data.to_string()) Ok(data.to_string())
} }
@@ -594,8 +594,10 @@ impl MemoryHandler {
"visualize_graph" => { "visualize_graph" => {
let req = parse_tool!(args.clone(), id, VisualizeGraphTool); let req = parse_tool!(args.clone(), id, VisualizeGraphTool);
let query = req.query.unwrap_or_default().to_lowercase(); let query = req.query.unwrap_or_default().to_lowercase();
let full = self.state.get_full_graph();
let mut included = HashSet::new(); let mut included = HashSet::new();
let mut to_draw = Vec::new();
self.state.read_graph(|full| {
for (name, e) in &full.entities { for (name, e) in &full.entities {
if let Some(ns) = &req.namespace if let Some(ns) = &req.namespace
&& e.namespace != *ns && e.namespace != *ns
@@ -609,7 +611,7 @@ impl MemoryHandler {
included.insert(name.clone()); included.insert(name.clone());
} }
} }
let mut to_draw = Vec::new();
for r in &full.relations { for r in &full.relations {
if let Some(ns) = &req.namespace if let Some(ns) = &req.namespace
&& r.namespace != *ns && r.namespace != *ns
@@ -622,9 +624,10 @@ impl MemoryHandler {
{ {
included.insert(r.from.clone()); included.insert(r.from.clone());
included.insert(r.to.clone()); included.insert(r.to.clone());
to_draw.push(r); to_draw.push(r.clone());
} }
} }
});
use std::fmt::Write; use std::fmt::Write;
let mut output = String::with_capacity(included.len() * 40 + to_draw.len() * 60); let mut output = String::with_capacity(included.len() * 40 + to_draw.len() * 60);
output.push_str("graph TD;\n"); output.push_str("graph TD;\n");
@@ -1104,18 +1107,18 @@ impl MemoryHandler {
Ok("Entities merged".to_string()) Ok("Entities merged".to_string())
} }
"find_orphans" => { "find_orphans" => {
let full = self.state.get_full_graph(); let orphans = self.state.read_graph(|full| {
let mut connected = std::collections::HashSet::new(); let mut connected = std::collections::HashSet::new();
for r in &full.relations { for r in &full.relations {
connected.insert(r.from.clone()); connected.insert(r.from.clone());
connected.insert(r.to.clone()); connected.insert(r.to.clone());
} }
let orphans: Vec<String> = full full.entities
.entities
.keys() .keys()
.filter(|k| !connected.contains(*k)) .filter(|k| !connected.contains(*k))
.cloned() .cloned()
.collect(); .collect::<Vec<String>>()
});
let data = serde_json::to_string(&orphans).unwrap_or_default(); let data = serde_json::to_string(&orphans).unwrap_or_default();
Ok(data.to_string()) Ok(data.to_string())
} }
@@ -1522,7 +1525,7 @@ impl MemoryHandler {
let mut snippets = Vec::new(); let mut snippets = Vec::new();
let mut adrs = Vec::new(); let mut adrs = Vec::new();
let full = self.state.get_full_graph(); self.state.read_graph(|full| {
for (id, doc_type, _, _, _) in &matches { for (id, doc_type, _, _, _) in &matches {
if doc_type == "entity" if doc_type == "entity"
&& let Some(e) = full.entities.get(id) && let Some(e) = full.entities.get(id)
@@ -1530,6 +1533,7 @@ impl MemoryHandler {
kg.entities.insert(id.clone(), e.clone()); kg.entities.insert(id.clone(), e.clone());
} }
} }
});
for t in self.state.tasks.read() { for t in self.state.tasks.read() {
if matches.iter().any(|(id, typ, _, _, _)| id == &t.id && typ == "task") { if matches.iter().any(|(id, typ, _, _, _)| id == &t.id && typ == "task") {
tasks.push(t); tasks.push(t);
+5 -5
View File
@@ -214,9 +214,7 @@ async fn gate_set_handler(
(axum::http::StatusCode::OK, "Gate state updated.").into_response() (axum::http::StatusCode::OK, "Gate state updated.").into_response()
} }
fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>> { async fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>> {
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async {
state.rebuild_index().await; state.rebuild_index().await;
tokio::spawn(index_committer_worker(Arc::clone(&state))); tokio::spawn(index_committer_worker(Arc::clone(&state)));
let app_state = Arc::new(AppState { let app_state = Arc::new(AppState {
@@ -427,7 +425,6 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
let _ = tokio::fs::write(&log_path, format!("Server crashed: {}\n", e)).await; let _ = tokio::fs::write(&log_path, format!("Server crashed: {}\n", e)).await;
} }
Ok(()) Ok(())
})
} }
async fn ws_handler( async fn ws_handler(
@@ -796,6 +793,9 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
write_txn.commit().unwrap(); write_txn.commit().unwrap();
} }
let rt = tokio::runtime::Runtime::new().unwrap();
let _guard = rt.enter();
let state = Arc::new(MemoryState { let state = Arc::new(MemoryState {
graph: Store::new("knowledge_graph_master", db.clone()), graph: Store::new("knowledge_graph_master", db.clone()),
base_dir: base.clone(), base_dir: base.clone(),
@@ -830,5 +830,5 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
activity_tx: tokio::sync::broadcast::channel(100).0, activity_tx: tokio::sync::broadcast::channel(100).0,
}); });
run_server(state) rt.block_on(run_server(state))
} }
+7
View File
@@ -49,6 +49,13 @@ impl MemoryState {
self.graph.read() self.graph.read()
} }
pub fn read_graph<F, R>(&self, f: F) -> R
where
F: FnOnce(&KnowledgeGraph) -> R,
{
self.graph.read_with(f)
}
pub fn modify_graph<F: FnOnce(&mut KnowledgeGraph)>(&self, update_fn: F) { pub fn modify_graph<F: FnOnce(&mut KnowledgeGraph)>(&self, update_fn: F) {
self.graph.modify(update_fn); self.graph.modify(update_fn);
} }
+32 -16
View File
@@ -5,18 +5,35 @@ use std::sync::{Arc, RwLock};
pub const STORE_TABLE: TableDefinition<&str, &[u8]> = TableDefinition::new("store"); pub const STORE_TABLE: TableDefinition<&str, &[u8]> = TableDefinition::new("store");
pub struct Store<T> { pub struct Store<T> {
pub key: String,
pub db: Arc<Database>,
pub cache: RwLock<T>, pub cache: RwLock<T>,
tx: tokio::sync::mpsc::UnboundedSender<Vec<u8>>,
} }
impl<T: DeserializeOwned + Default + Serialize + Clone + Send + 'static> Store<T> { impl<T: DeserializeOwned + Default + Serialize + Clone + Send + 'static> Store<T> {
pub fn new(key: &str, db: Arc<Database>) -> Self { pub fn new(key: &str, db: Arc<Database>) -> Self {
let initial_data = Self::load_from_db(key, &db); let initial_data = Self::load_from_db(key, &db);
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::<Vec<u8>>();
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 { Self {
key: key.to_string(),
db,
cache: RwLock::new(initial_data), cache: RwLock::new(initial_data),
tx,
} }
} }
@@ -35,23 +52,22 @@ impl<T: DeserializeOwned + Default + Serialize + Clone + Send + 'static> Store<T
lock.clone() lock.clone()
} }
pub fn read_with<F, R>(&self, f: F) -> R
where
F: FnOnce(&T) -> R,
{
let lock = self.cache.read().unwrap();
f(&lock)
}
pub fn modify<F: FnOnce(&mut T)>(&self, f: F) { pub fn modify<F: FnOnce(&mut T)>(&self, f: F) {
let (key, db, json_data) = { let json_data = {
let mut lock = self.cache.write().unwrap(); let mut lock = self.cache.write().unwrap();
f(&mut lock); f(&mut lock);
// Serialize while holding lock to avoid expensive deep clone of T // Serialize while holding lock to avoid expensive deep clone of T
let json = serde_json::to_vec(&*lock).unwrap(); serde_json::to_vec(&*lock).unwrap()
(self.key.clone(), self.db.clone(), json)
}; };
let _ = self.tx.send(json_data);
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();
});
} }
} }