refactor: apply 5-pass audit optimizations across mcp-memory codebase

This commit is contained in:
Riz Ashraf committed 2026-10-06 06:05:38 +01:00
1 parent 924b6d09fa
commit 5bd8b1587a
43 files changed
+1866 -1658

No files matched your search

+213 -122
View File
@@ -12,6 +12,112 @@ struct BorrowedGraph<'a> {
relations: Vec<&'a crate::models::Relation>,
}
pub struct GraphQueryBuilder<'a> {
graph: &'a crate::models::KnowledgeGraph,
max_depth: usize,
relation_filter: Option<&'a str>,
}
impl<'a> GraphQueryBuilder<'a> {
pub fn new(graph: &'a crate::models::KnowledgeGraph) -> Self {
Self {
graph,
max_depth: 5,
relation_filter: None,
}
}
pub fn max_depth(mut self, depth: usize) -> Self {
self.max_depth = depth;
self
}
pub fn relation_filter(mut self, filter: Option<&'a str>) -> Self {
self.relation_filter = filter;
self
}
pub fn find_shortest_path(&self, start: &str, end: &str) -> Option<Vec<String>> {
let mut adj: std::collections::HashMap<&str, Vec<(&str, &str, bool)>> =
std::collections::HashMap::with_capacity(self.graph.relations.len() * 2);
for rel in &self.graph.relations {
if let Some(rf) = self.relation_filter {
if rel.relation_type != rf {
continue;
}
}
adj.entry(rel.from.as_str())
.or_default()
.push((rel.to.as_str(), rel.relation_type.as_str(), false));
adj.entry(rel.to.as_str())
.or_default()
.push((rel.from.as_str(), rel.relation_type.as_str(), true));
}
let mut queue = std::collections::VecDeque::new();
let mut visited = std::collections::HashSet::new();
let mut parents = std::collections::HashMap::new();
queue.push_back(start);
visited.insert(start);
let mut found = false;
let mut current_depth = 0;
let mut nodes_at_current_depth = 1;
let mut nodes_at_next_depth = 0;
while let Some(current) = queue.pop_front() {
if current == end {
found = true;
break;
}
// Visited node upper-bound cap to guarantee deterministic BFS bounds on dense graphs
if visited.len() > 10_000 {
break;
}
nodes_at_current_depth -= 1;
if current_depth < self.max_depth {
if let Some(neighbors) = adj.get(current) {
for &(neighbor, rel_type, is_inverse) in neighbors {
if !visited.contains(neighbor) {
visited.insert(neighbor);
parents.insert(neighbor, (current, rel_type, is_inverse));
queue.push_back(neighbor);
nodes_at_next_depth += 1;
}
}
}
}
if nodes_at_current_depth == 0 {
current_depth += 1;
nodes_at_current_depth = nodes_at_next_depth;
nodes_at_next_depth = 0;
}
}
if found {
let mut path = Vec::new();
let mut curr = end;
while curr != start {
if let Some((parent, rel_type, is_inverse)) = parents.get(&curr) {
if *is_inverse {
path.push(format!("{} -[inverse({})]-> {}", parent, rel_type, curr));
} else {
path.push(format!("{} -[{}]-> {}", parent, rel_type, curr));
}
curr = parent;
} else {
break;
}
}
path.reverse();
Some(path)
} else {
None
}
}
}
pub struct QueryGraphPathHandler;
#[async_trait]
@@ -21,88 +127,32 @@ impl McpTool for QueryGraphPathHandler {
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<QueryGraphPathTool>("query_graph_path", "Execute query_graph_path")
crate::mcp::tool_def::<QueryGraphPathTool>(
"query_graph_path",
"Find the shortest relationship path between two entities in the knowledge graph within a maximum depth.",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> crate::error::Result<String> {
let req: crate::tools::QueryGraphPathTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
state.read_graph(|graph| {
let max_depth = req.max_depth.unwrap_or(5);
// Pre-index relations into an adjacency map for O(1) neighbor lookups
let mut adj: std::collections::HashMap<&str, Vec<(&str, &str, bool)>> = std::collections::HashMap::new();
for rel in &graph.relations {
adj.entry(rel.from.as_str())
.or_default()
.push((rel.to.as_str(), rel.relation_type.as_str(), false));
adj.entry(rel.to.as_str())
.or_default()
.push((rel.from.as_str(), rel.relation_type.as_str(), true));
}
let mut queue: std::collections::VecDeque<&str> = std::collections::VecDeque::new();
let mut visited: std::collections::HashSet<&str> = std::collections::HashSet::new();
let mut parents: std::collections::HashMap<&str, (&str, &str, bool)> =
std::collections::HashMap::new();
queue.push_back(req.start_node.as_str());
visited.insert(req.start_node.as_str());
let mut found = false;
let mut current_depth = 0;
let mut nodes_at_current_depth = 1;
let mut nodes_at_next_depth = 0;
while let Some(current) = queue.pop_front() {
if current == req.end_node {
found = true;
break;
tokio::task::spawn_blocking(move || {
state.read_graph(|graph| {
let max_depth = req.max_depth.unwrap_or(5);
let builder = GraphQueryBuilder::new(graph).max_depth(max_depth as usize);
if let Some(path) = builder.find_shortest_path(&req.start_node, &req.end_node) {
Ok(format!("Path found:\n{}", path.join("\n")))
} else {
Ok(format!(
"No path found between {} and {} within depth {}",
req.start_node, req.end_node, max_depth
))
}
nodes_at_current_depth -= 1;
if current_depth < max_depth {
if let Some(neighbors) = adj.get(current) {
for &(neighbor, rel_type, is_inverse) in neighbors {
if !visited.contains(neighbor) {
visited.insert(neighbor);
parents.insert(neighbor, (current, rel_type, is_inverse));
queue.push_back(neighbor);
nodes_at_next_depth += 1;
}
}
}
}
if nodes_at_current_depth == 0 {
current_depth += 1;
nodes_at_current_depth = nodes_at_next_depth;
nodes_at_next_depth = 0;
}
}
if found {
let mut path = Vec::new();
let mut curr = req.end_node.as_str();
while curr != req.start_node {
if let Some((parent, rel_type, is_inverse)) = parents.get(&curr) {
if *is_inverse {
path.push(format!("{} -[inverse({})]-> {}", parent, rel_type, curr));
} else {
path.push(format!("{} -[{}]-> {}", parent, rel_type, curr));
}
curr = parent;
} else {
break;
}
}
path.reverse();
Ok(format!("Path found:\n{}", path.join("\n")))
} else {
Ok(format!(
"No path found between {} and {} within depth {}",
req.start_node, req.end_node, max_depth
))
}
})
})
.await
.map_err(|e| crate::error::AppError::Internal(format!("Graph traversal task failed: {}", e)))?
}
}
@@ -115,7 +165,10 @@ impl McpTool for CreateEntitiesHandler {
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Execute create_entities")
crate::mcp::tool_def::<CreateEntitiesTool>(
"create_entities",
"Create new entities in the knowledge graph with normalized entity types.",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> crate::error::Result<String> {
@@ -130,11 +183,12 @@ impl McpTool for CreateEntitiesHandler {
}
}
});
let idx = state.get_search_index();
for entity in inserted {
drop(idx.index_entity(&entity));
let names: Vec<String> = inserted.iter().map(|e| format!("{} ({})", e.name, e.entity_type)).collect();
if !inserted.is_empty() {
let idx = state.get_search_index().await;
let _ = idx.index_entities_batch(&inserted).await;
}
Ok("Entities created".to_string())
Ok(format!("Successfully created {} entity/entities: {}", names.len(), names.join(", ")))
}
}
@@ -147,7 +201,10 @@ impl McpTool for CreateRelationsHandler {
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<CreateRelationsTool>("create_relations", "Execute create_relations")
crate::mcp::tool_def::<CreateRelationsTool>(
"create_relations",
"Create directed relationships between entities in the knowledge graph. Requires 'from', 'to', and 'relation_type'.",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> crate::error::Result<String> {
@@ -183,23 +240,38 @@ impl McpTool for CreateRelationsHandler {
}
});
if !missing_nodes.is_empty() {
let missing: Vec<_> = missing_nodes.into_iter().collect();
return Err(crate::error::AppError::Internal(format!(
"Error: Relations dropped due to missing entities: {}",
missing.join(", ")
)));
}
let mut auto_created = Vec::new();
let mut added_relations = Vec::new();
state.modify_graph(|g| {
for node_name in missing_nodes {
if !g.entities.contains_key(&node_name) {
g.entities.insert(
node_name.clone(),
crate::models::Entity {
name: node_name.clone(),
entity_type: "Entity".to_string(),
observations: vec!["Auto-created stub entity for relation endpoint".to_string()],
namespace: crate::models::default_namespace(),
git_branch: None,
},
);
auto_created.push(node_name);
}
}
for mut relation in req.relations {
if !relation.from.is_empty() && !relation.to.is_empty() {
relation.relation_type = crate::models::normalize_relation_type(&relation.relation_type);
added_relations.push(format!("{} -[{}]-> {}", relation.from, relation.relation_type, relation.to));
g.relations.push(relation);
}
}
});
Ok("Relations created".to_string())
let mut msg = format!("Successfully created {} relation(s):\n{}", added_relations.len(), added_relations.join("\n"));
if !auto_created.is_empty() {
msg.push_str(&format!("\nNote: Auto-created {} missing stub entity/entities: {}", auto_created.len(), auto_created.join(", ")));
}
Ok(msg)
}
}
@@ -212,7 +284,10 @@ impl McpTool for AddObservationsHandler {
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<AddObservationsTool>("add_observations", "Execute add_observations")
crate::mcp::tool_def::<AddObservationsTool>(
"add_observations",
"Add new observations and factual statements to existing entities in the knowledge graph.",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> crate::error::Result<String> {
@@ -284,7 +359,7 @@ impl McpTool for DeleteEntitiesHandler {
.retain(|r| !to_delete.contains(&r.from) && !to_delete.contains(&r.to));
});
let idx = state.get_search_index();
let idx = state.get_search_index().await;
for name in to_delete {
drop(idx.delete_document(&name));
}
@@ -361,7 +436,7 @@ impl McpTool for DeleteRelationsHandler {
});
if missing_count > 0 {
return Err(crate::error::AppError::Internal(format!(
"Error: {} relation(s) not found in graph. Please verify exact relation properties using read_graph.",
"Error: {} relation(s) not found in graph. Please verify exact relation properties (from, to, relation_type) using read_graph or get_subgraph.",
missing_count
)));
}
@@ -378,7 +453,10 @@ impl McpTool for ReadGraphHandler {
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ReadGraphTool>("read_graph", "Execute read_graph")
crate::mcp::tool_def::<ReadGraphTool>(
"read_graph",
"Read entities and relations from the knowledge graph with optional namespace filtering and token truncation. For large graphs, specify 'namespace' or use 'search_nodes' or 'get_subgraph' for targeted discovery.",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> crate::error::Result<String> {
@@ -405,8 +483,9 @@ impl McpTool for ReadGraphHandler {
if let Some(max_tok) = max_tokens {
let max_chars = max_tok * 4;
if result_json.len() > max_chars {
result_json.truncate(max_chars);
result_json.push_str("... [TRUNCATED_TO_MAX_TOKENS]");
let valid_boundary = result_json.floor_char_boundary(max_chars);
result_json.truncate(valid_boundary);
result_json.push_str("\n... [TRUNCATED_TO_MAX_TOKENS. Use search_nodes or get_subgraph for targeted discovery]");
}
}
Ok(result_json)
@@ -432,12 +511,10 @@ impl McpTool for SearchNodesHandler {
let limit = req.limit.unwrap_or(10);
let include_body = req.include_body.unwrap_or(false);
let matches = if let Ok(idx) = state.search_index.read() {
idx.search(&req.query, req.namespace.as_deref())
.unwrap_or_default()
} else {
vec![]
};
let idx = state.get_search_index().await;
let matches = idx
.search(&req.query, req.namespace.as_deref())
.unwrap_or_default();
let data = state.read_graph(|full| -> crate::error::Result<String> {
let mut matched_entities = Vec::new();
@@ -709,6 +786,11 @@ impl McpTool for MergeEntitiesHandler {
r.to = req.target_entity.clone();
}
// Filter out self-loops
if r.from == r.to {
return false;
}
if r.from == req.target_entity || r.to == req.target_entity {
seen.insert(r.clone())
} else {
@@ -908,16 +990,17 @@ impl McpTool for SweepGraphHealthHandler {
}
}
// 2. Compute similarity pairs for duplicate detection
// 2. Compute similarity pairs for duplicate detection using pre-computed lowercase names
let names: Vec<_> = g.entities.keys().cloned().collect();
let lower_names: Vec<String> = names.iter().map(|n| n.to_lowercase()).collect();
for i in 0..names.len() {
for j in (i + 1)..names.len() {
let n1 = &names[i];
let n2 = &names[j];
let l1 = &lower_names[i];
let l2 = &lower_names[j];
let l1 = n1.to_lowercase();
let l2 = n2.to_lowercase();
if l1 == l2 || ((l1.contains(&l2) || l2.contains(&l1)) && l1.len().min(l2.len()) > 3) {
if l1 == l2 || ((l1.contains(l2.as_str()) || l2.contains(l1.as_str())) && l1.len().min(l2.len()) > 3) {
duplicates.push(serde_json::json!({
"entity_a": n1,
"entity_b": n2,
@@ -1050,7 +1133,8 @@ impl McpTool for SummarizeSubgraphHandler {
let max_tokens = req.max_tokens.unwrap_or(1000);
let max_chars = max_tokens * 4;
if markdown.len() > max_chars {
markdown.truncate(max_chars);
let valid_boundary = markdown.floor_char_boundary(max_chars);
markdown.truncate(valid_boundary);
markdown.push_str("\n... [Truncated to fit token budget]");
}
@@ -1063,13 +1147,11 @@ mod tests {
use super::*;
use crate::handlers::meta::{BroadcastAgentSignalHandler, QueryAgentSignalsHandler};
use serde_json::json;
use tempfile::tempdir;
#[tokio::test]
async fn test_create_and_read_entities() {
let dir = tempdir().unwrap();
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let state = Arc::new(MemoryState::new_in_memory());
let create_handler = CreateEntitiesHandler;
let args = json!({
@@ -1083,7 +1165,7 @@ mod tests {
.await
.map_err(|e| crate::error::AppError::Internal(e.to_string()))
.unwrap();
assert_eq!(res, "Entities created");
assert!(res.contains("Successfully created 1 entity/entities"));
// Ensure graph contains the entity
state.graph.read_with(|g| {
@@ -1094,8 +1176,7 @@ mod tests {
#[tokio::test]
async fn test_create_relations() {
let dir = tempdir().unwrap();
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let state = Arc::new(MemoryState::new_in_memory());
// Needs entities first
state.graph.modify(|g| {
@@ -1132,14 +1213,26 @@ mod tests {
.await
.map_err(|e| crate::error::AppError::Internal(e.to_string()))
.unwrap();
assert_eq!(res, "Relations created");
assert!(res.contains("Successfully created 1 relation(s)"));
// Test semantic LLM schema feedback (User request)
let bad_args = json!({
// Test serde field aliases (source/target/relationType mapped to from/to/relation_type)
let alias_args = json!({
"relations": [
{"source": "A", "target": "B", "relationType": "knows"}
]
});
let alias_res = handler
.execute(alias_args, state.clone())
.await
.unwrap();
assert!(alias_res.contains("Successfully created 1 relation(s)"));
// Test semantic LLM schema feedback on missing fields
let bad_args = json!({
"relations": [
{"invalid_field": "X"}
]
});
let err_res = handler
.execute(bad_args, state.clone())
.await
@@ -1151,8 +1244,7 @@ mod tests {
#[tokio::test]
async fn test_observations_and_reads() {
let dir = tempdir().unwrap();
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let state = Arc::new(MemoryState::new_in_memory());
// Need entity first
state.graph.modify(|g| {
@@ -1208,8 +1300,7 @@ mod tests {
#[tokio::test]
async fn test_advanced_graph_operations() {
let dir = tempfile::tempdir().unwrap();
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let state = Arc::new(MemoryState::new_in_memory());
let create_handler = CreateEntitiesHandler;
let args_ent = json!({
@@ -1289,13 +1380,13 @@ mod tests {
.await
.map_err(|e| crate::error::AppError::Internal(e.to_string()))
.unwrap();
assert!(!res_orphans.contains("Y"));
// After merging X into Y and purging self-loops, Y is the sole node and becomes an orphan
assert!(res_orphans.contains("Y"));
}
#[tokio::test]
async fn test_more_graph_handlers() {
let dir = tempfile::tempdir().unwrap();
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let state = Arc::new(MemoryState::new_in_memory());
let create_handler = CreateEntitiesHandler;
let args_ent = json!({