perf: fix lifetime issues and optimize remaining string allocations in graph traversal endpoints
This commit is contained in:
1 parent
3c31aeec1f
commit
6b37799cc5
2 files changed
+180
-27
No files matched your search
@@ -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()
|
||||
@@ -337,24 +337,25 @@ impl McpTool for OpenNodesHandler {
|
||||
|
||||
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 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 connected = HashSet::new();
|
||||
state.read_graph(|full| {
|
||||
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<MemoryState>) -> Result<String, 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 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();
|
||||
|
||||
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,34 +402,33 @@ 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 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 out = String::with_capacity(s.len());
|
||||
let mut o = String::with_capacity(s.len());
|
||||
for c in s.chars() {
|
||||
if c != '"' && c != '(' && c != ')' {
|
||||
if id_mode && (c == ' ' || c == '-' || c == '.') {
|
||||
out.push('_');
|
||||
o.push('_');
|
||||
} else {
|
||||
out.push(c);
|
||||
o.push(c);
|
||||
}
|
||||
}
|
||||
}
|
||||
out
|
||||
o
|
||||
};
|
||||
|
||||
for name in &included {
|
||||
let _ = writeln!(
|
||||
output,
|
||||
out,
|
||||
" id_{}[\"{}\"];",
|
||||
sanitize(name, true),
|
||||
sanitize(name, false)
|
||||
@@ -435,13 +436,15 @@ impl McpTool for VisualizeGraphHandler {
|
||||
}
|
||||
for r in to_draw {
|
||||
let _ = writeln!(
|
||||
output,
|
||||
out,
|
||||
" id_{}-->|\"{}\"|id_{};",
|
||||
sanitize(&r.from, true),
|
||||
r.relation_type.replace("\"", ""),
|
||||
sanitize(&r.to, true)
|
||||
);
|
||||
}
|
||||
out
|
||||
});
|
||||
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::<Vec<String>>()
|
||||
});
|
||||
|
||||
Reference in new issue
Block a user