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
+87
-60
No files matched your search
+43
-39
@@ -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
@@ -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))
|
||||||
}
|
}
|
||||||
@@ -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
@@ -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();
|
|
||||||
});
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in new issue
Block a user