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> {
|
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>>()
|
||||||
});
|
});
|
||||||
|
|||||||
Reference in new issue
Block a user