Optimize handlers to avoid deep clones of KnowledgeGraph, and fix database write concurrency issues
This commit is contained in:
1 parent
7a48fa5d34
commit
9b349e6459
4 files changed
+197
-170
No files matched your search
+150
-146
@@ -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
@@ -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))
|
||||
}
|
||||
@@ -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
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in new issue
Block a user