perf: fix lifetime issues and optimize remaining string allocations in graph traversal endpoints

This commit is contained in:
Riz Ashraf committed 2026-09-21 13:47:09 +01:00
1 parent 3c31aeec1f
commit 6b37799cc5
2 files changed
+180 -27

No files matched your search

+150
View File
@@ -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()
+30 -27
View File
@@ -337,24 +337,25 @@ impl McpTool for OpenNodesHandler {
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> { async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: OpenNodesTool = serde_json::from_value(args).map_err(|e| e.to_string())?; let req: OpenNodesTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let targets: HashSet<_> = req.names.into_iter().collect(); 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 result = KnowledgeGraph::default();
let mut connected = HashSet::new(); let mut connected: HashSet<&str> = HashSet::new();
state.read_graph(|full| {
for r in &full.relations { for r in &full.relations {
if targets.contains(&r.from) { if targets.contains(r.from.as_str()) {
connected.insert(r.to.clone()); connected.insert(r.to.as_str());
result.relations.push(r.clone()); result.relations.push(r.clone());
} else if targets.contains(&r.to) { } else if targets.contains(r.to.as_str()) {
connected.insert(r.from.clone()); connected.insert(r.from.as_str());
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.as_str()) || connected.contains(name.as_str()) {
result.entities.insert(name.clone(), e.clone()); result.entities.insert(name.clone(), e.clone());
} }
} }
result
}); });
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())
@@ -376,10 +377,11 @@ impl McpTool for VisualizeGraphHandler {
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> { async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: VisualizeGraphTool = serde_json::from_value(args).map_err(|e| e.to_string())?; let req: VisualizeGraphTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let query = req.query.unwrap_or_default().to_lowercase(); let query = req.query.unwrap_or_default().to_lowercase();
let mut included = HashSet::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(); let mut to_draw = Vec::new();
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
@@ -390,7 +392,7 @@ impl McpTool for VisualizeGraphHandler {
|| contains_ignore_ascii_case(name, &query) || contains_ignore_ascii_case(name, &query)
|| contains_ignore_ascii_case(&e.entity_type, &query) || contains_ignore_ascii_case(&e.entity_type, &query)
{ {
included.insert(name.clone()); included.insert(name.as_str());
} }
} }
@@ -400,34 +402,33 @@ impl McpTool for VisualizeGraphHandler {
{ {
continue; continue;
} }
if query.is_empty() || included.contains(&r.from) || included.contains(&r.to) { if query.is_empty() || included.contains(r.from.as_str()) || included.contains(r.to.as_str()) {
included.insert(r.from.clone()); included.insert(r.from.as_str());
included.insert(r.to.clone()); included.insert(r.to.as_str());
to_draw.push(r.clone()); to_draw.push(r.clone());
} }
} }
});
use std::fmt::Write; let mut out = String::with_capacity(included.len() * 40 + to_draw.len() * 60);
let mut output = String::with_capacity(included.len() * 40 + to_draw.len() * 60); out.push_str("graph TD;\n");
output.push_str("graph TD;\n");
let sanitize = |s: &str, id_mode: bool| -> String { let sanitize = |s: &str, id_mode: bool| -> String {
let mut out = String::with_capacity(s.len()); let mut o = String::with_capacity(s.len());
for c in s.chars() { for c in s.chars() {
if c != '"' && c != '(' && c != ')' { if c != '"' && c != '(' && c != ')' {
if id_mode && (c == ' ' || c == '-' || c == '.') { if id_mode && (c == ' ' || c == '-' || c == '.') {
out.push('_'); o.push('_');
} else { } else {
out.push(c); o.push(c);
} }
} }
} }
out o
}; };
for name in &included { for name in &included {
let _ = writeln!( let _ = writeln!(
output, out,
" id_{}[\"{}\"];", " id_{}[\"{}\"];",
sanitize(name, true), sanitize(name, true),
sanitize(name, false) sanitize(name, false)
@@ -435,13 +436,15 @@ impl McpTool for VisualizeGraphHandler {
} }
for r in to_draw { for r in to_draw {
let _ = writeln!( let _ = writeln!(
output, out,
" id_{}-->|\"{}\"|id_{};", " id_{}-->|\"{}\"|id_{};",
sanitize(&r.from, true), sanitize(&r.from, true),
r.relation_type.replace("\"", ""), r.relation_type.replace("\"", ""),
sanitize(&r.to, true) sanitize(&r.to, true)
); );
} }
out
});
if output == "graph TD;\n" { if output == "graph TD;\n" {
output = "No nodes found to visualize.".to_string(); output = "No nodes found to visualize.".to_string();
} }
@@ -527,12 +530,12 @@ impl McpTool for FindOrphansHandler {
let orphans = state.read_graph(|full| { let orphans = 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.as_str());
connected.insert(r.to.clone()); connected.insert(r.to.as_str());
} }
full.entities full.entities
.keys() .keys()
.filter(|k| !connected.contains(*k)) .filter(|k| !connected.contains(k.as_str()))
.cloned() .cloned()
.collect::<Vec<String>>() .collect::<Vec<String>>()
}); });