From 6b37799cc5573e342650078f2c4393d02f9bab6f Mon Sep 17 00:00:00 2001 From: Riz Ashraf Date: Mon, 21 Sep 2026 13:47:09 +0100 Subject: [PATCH] perf: fix lifetime issues and optimize remaining string allocations in graph traversal endpoints --- server/src/bin_replace2.py | 150 ++++++++++++++++++++++++++++++++ server/src/handlers_v2/graph.rs | 103 +++++++++++----------- 2 files changed, 203 insertions(+), 50 deletions(-) create mode 100644 server/src/bin_replace2.py diff --git a/server/src/bin_replace2.py b/server/src/bin_replace2.py new file mode 100644 index 0000000..020118a --- /dev/null +++ b/server/src/bin_replace2.py @@ -0,0 +1,150 @@ +import os + +def rewrite_graph(): + with open('server/src/handlers_v2/graph.rs', 'r', encoding='utf-8') as f: + content = f.read() + + old_block = """ let mut included: HashSet<&str> = HashSet::new(); + let mut to_draw = Vec::new(); + + state.read_graph(|full| { + for (name, e) in &full.entities { + if let Some(ns) = &req.namespace + && e.namespace != *ns + { + continue; + } + if query.is_empty() + || contains_ignore_ascii_case(name, &query) + || contains_ignore_ascii_case(&e.entity_type, &query) + { + included.insert(name.as_str()); + } + } + + for r in &full.relations { + if let Some(ns) = &req.namespace + && r.namespace != *ns + { + continue; + } + if query.is_empty() || included.contains(r.from.as_str()) || included.contains(r.to.as_str()) { + included.insert(r.from.as_str()); + included.insert(r.to.as_str()); + 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"); + + let sanitize = |s: &str, id_mode: bool| -> String { + let mut out = String::with_capacity(s.len()); + for c in s.chars() { + if c != '"' && c != '(' && c != ')' { + if id_mode && (c == ' ' || c == '-' || c == '.') { + out.push('_'); + } else { + out.push(c); + } + } + } + out + }; + + for name in &included { + let _ = writeln!( + output, + " id_{}[\\"{}\\"];", + sanitize(name, true), + sanitize(name, false) + ); + } + for r in to_draw { + let _ = writeln!( + output, + " id_{}-->|\\"{}\\"|id_{};", + sanitize(&r.from, true), + r.relation_type.replace("\\"", ""), + sanitize(&r.to, true) + ); + }""" + + new_block = """ use std::fmt::Write; + let mut output = state.read_graph(|full| { + let mut included: HashSet<&str> = HashSet::new(); + let mut to_draw = Vec::new(); + + for (name, e) in &full.entities { + if let Some(ns) = &req.namespace + && e.namespace != *ns + { + continue; + } + if query.is_empty() + || contains_ignore_ascii_case(name, &query) + || contains_ignore_ascii_case(&e.entity_type, &query) + { + included.insert(name.as_str()); + } + } + + for r in &full.relations { + if let Some(ns) = &req.namespace + && r.namespace != *ns + { + continue; + } + if query.is_empty() || included.contains(r.from.as_str()) || included.contains(r.to.as_str()) { + included.insert(r.from.as_str()); + included.insert(r.to.as_str()); + to_draw.push(r.clone()); + } + } + + let mut out = String::with_capacity(included.len() * 40 + to_draw.len() * 60); + out.push_str("graph TD;\\n"); + + let sanitize = |s: &str, id_mode: bool| -> String { + let mut o = String::with_capacity(s.len()); + for c in s.chars() { + if c != '"' && c != '(' && c != ')' { + if id_mode && (c == ' ' || c == '-' || c == '.') { + o.push('_'); + } else { + o.push(c); + } + } + } + o + }; + + for name in &included { + let _ = writeln!( + out, + " id_{}[\\"{}\\"];", + sanitize(name, true), + sanitize(name, false) + ); + } + for r in to_draw { + let _ = writeln!( + out, + " id_{}-->|\\"{}\\"|id_{};", + sanitize(&r.from, true), + r.relation_type.replace("\\"", ""), + sanitize(&r.to, true) + ); + } + out + });""" + + if old_block in content: + with open('server/src/handlers_v2/graph.rs', 'w', encoding='utf-8') as f: + f.write(content.replace(old_block, new_block)) + print("Replaced visualize_graph") + else: + print("Could not find old block") + +rewrite_graph() diff --git a/server/src/handlers_v2/graph.rs b/server/src/handlers_v2/graph.rs index fa1c3dd..4351e37 100644 --- a/server/src/handlers_v2/graph.rs +++ b/server/src/handlers_v2/graph.rs @@ -337,24 +337,25 @@ impl McpTool for OpenNodesHandler { async fn execute(&self, args: Value, state: Arc) -> Result { let req: OpenNodesTool = serde_json::from_value(args).map_err(|e| e.to_string())?; - let targets: HashSet<_> = req.names.into_iter().collect(); - let mut result = KnowledgeGraph::default(); - let mut connected = HashSet::new(); - state.read_graph(|full| { + let result = state.read_graph(|full| { + let targets: HashSet<&str> = req.names.iter().map(|s| s.as_str()).collect(); + let mut result = KnowledgeGraph::default(); + let mut connected: HashSet<&str> = HashSet::new(); for r in &full.relations { - if targets.contains(&r.from) { - connected.insert(r.to.clone()); + if targets.contains(r.from.as_str()) { + connected.insert(r.to.as_str()); result.relations.push(r.clone()); - } else if targets.contains(&r.to) { - connected.insert(r.from.clone()); + } else if targets.contains(r.to.as_str()) { + connected.insert(r.from.as_str()); result.relations.push(r.clone()); } } for (name, e) in &full.entities { - if targets.contains(name) || connected.contains(name) { + if targets.contains(name.as_str()) || connected.contains(name.as_str()) { result.entities.insert(name.clone(), e.clone()); } } + result }); let data = serde_json::to_string(&result).unwrap_or_default(); Ok(data.to_string()) @@ -376,10 +377,11 @@ impl McpTool for VisualizeGraphHandler { async fn execute(&self, args: Value, state: Arc) -> Result { let req: VisualizeGraphTool = serde_json::from_value(args).map_err(|e| e.to_string())?; let query = req.query.unwrap_or_default().to_lowercase(); - let mut included = HashSet::new(); - let mut to_draw = Vec::new(); + use std::fmt::Write; + let mut output = state.read_graph(|full| { + let mut included: HashSet<&str> = HashSet::new(); + let mut to_draw = Vec::new(); - state.read_graph(|full| { for (name, e) in &full.entities { if let Some(ns) = &req.namespace && e.namespace != *ns @@ -390,7 +392,7 @@ impl McpTool for VisualizeGraphHandler { || contains_ignore_ascii_case(name, &query) || contains_ignore_ascii_case(&e.entity_type, &query) { - included.insert(name.clone()); + included.insert(name.as_str()); } } @@ -400,48 +402,49 @@ impl McpTool for VisualizeGraphHandler { { continue; } - if query.is_empty() || included.contains(&r.from) || included.contains(&r.to) { - included.insert(r.from.clone()); - included.insert(r.to.clone()); + if query.is_empty() || included.contains(r.from.as_str()) || included.contains(r.to.as_str()) { + included.insert(r.from.as_str()); + included.insert(r.to.as_str()); 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"); - let sanitize = |s: &str, id_mode: bool| -> String { - let mut out = String::with_capacity(s.len()); - for c in s.chars() { - if c != '"' && c != '(' && c != ')' { - if id_mode && (c == ' ' || c == '-' || c == '.') { - out.push('_'); - } else { - out.push(c); + let mut out = String::with_capacity(included.len() * 40 + to_draw.len() * 60); + out.push_str("graph TD;\n"); + + let sanitize = |s: &str, id_mode: bool| -> String { + let mut o = String::with_capacity(s.len()); + for c in s.chars() { + if c != '"' && c != '(' && c != ')' { + if id_mode && (c == ' ' || c == '-' || c == '.') { + o.push('_'); + } else { + o.push(c); + } } } + o + }; + + for name in &included { + let _ = writeln!( + out, + " id_{}[\"{}\"];", + sanitize(name, true), + sanitize(name, false) + ); + } + for r in to_draw { + let _ = writeln!( + out, + " id_{}-->|\"{}\"|id_{};", + sanitize(&r.from, true), + r.relation_type.replace("\"", ""), + sanitize(&r.to, true) + ); } out - }; - - for name in &included { - let _ = writeln!( - output, - " id_{}[\"{}\"];", - sanitize(name, true), - sanitize(name, false) - ); - } - for r in to_draw { - let _ = writeln!( - output, - " id_{}-->|\"{}\"|id_{};", - sanitize(&r.from, true), - r.relation_type.replace("\"", ""), - sanitize(&r.to, true) - ); - } + }); if output == "graph TD;\n" { output = "No nodes found to visualize.".to_string(); } @@ -527,12 +530,12 @@ impl McpTool for FindOrphansHandler { let orphans = 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()); + connected.insert(r.from.as_str()); + connected.insert(r.to.as_str()); } full.entities .keys() - .filter(|k| !connected.contains(*k)) + .filter(|k| !connected.contains(k.as_str())) .cloned() .collect::>() });