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
+197 -170

No files matched your search

+150 -146
View File
@@ -366,74 +366,72 @@ impl MemoryHandler {
let result: Result<String, String> = 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<String, (String, String)> =
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<String, (String, String)> =
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<String> = 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::<Vec<String>>()
});
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);
+8 -8
View File
@@ -214,12 +214,10 @@ async fn gate_set_handler(
(axum::http::StatusCode::OK, "Gate state updated.").into_response()
}
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;
tokio::spawn(index_committer_worker(Arc::clone(&state)));
let app_state = Arc::new(AppState {
async fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>> {
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<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
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<dyn std::error::Error>> {
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<dyn std::error::Error>> {
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()
}
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) {
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 struct Store<T> {
pub key: String,
pub db: Arc<Database>,
pub cache: RwLock<T>,
tx: tokio::sync::mpsc::UnboundedSender<Vec<u8>>,
}
impl<T: DeserializeOwned + Default + Serialize + Clone + Send + 'static> Store<T> {
pub fn new(key: &str, db: Arc<Database>) -> Self {
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 {
key: key.to_string(),
db,
cache: RwLock::new(initial_data),
tx,
}
}
@@ -35,23 +52,22 @@ impl<T: DeserializeOwned + Default + Serialize + Clone + Send + 'static> Store<T
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) {
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);
}
}