refactor: bulk AST replacement and fault-tolerant graph handlers

This commit is contained in:
Riz Ashraf committed 2026-10-10 19:05:03 +01:00
1 parent 3ee95f5c39
commit f56750f596
3 files changed
+221 -194

No files matched your search

+154 -145
View File
@@ -204,167 +204,172 @@ impl McpTool for ReplaceAstNodeHandler {
let tool_args: ReplaceAstNodeTool = let tool_args: ReplaceAstNodeTool =
serde_json::from_value(args).map_err(|e| crate::error::AppError::BadRequest(format!("Schema validation failed. Your JSON arguments do not match the expected tool schema: {}", e)))?; serde_json::from_value(args).map_err(|e| crate::error::AppError::BadRequest(format!("Schema validation failed. Your JSON arguments do not match the expected tool schema: {}", e)))?;
let file_path = tool_args.file_path.clone(); let replacements = tool_args.replacements;
validate_safe_path(&file_path)?; if replacements.is_empty() {
return Ok("No replacements provided.".to_string());
}
let result = tokio::task::spawn_blocking(move || -> crate::error::Result<String> { let result = tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
let code = std::fs::read_to_string(&file_path).map_err(|e| { let mut msgs = Vec::new();
crate::error::AppError::Internal(format!("Failed to read file: {}", e)) for rep in replacements {
})?; let file_path = rep.file_path.clone();
validate_safe_path(&file_path)?;
let mut parser = Parser::new(); let code = std::fs::read_to_string(&file_path).map_err(|e| {
crate::error::AppError::Internal(format!("Failed to read file: {}", e))
})?;
let ext = std::path::Path::new(&file_path) let mut parser = Parser::new();
.extension()
.and_then(|s| s.to_str())
.unwrap_or("");
let language = match ext { let ext = std::path::Path::new(&file_path)
"rs" => tree_sitter_rust::LANGUAGE, .extension()
"ts" | "tsx" | "js" | "jsx" => tree_sitter_typescript::LANGUAGE_TYPESCRIPT, .and_then(|s| s.to_str())
"py" => tree_sitter_python::LANGUAGE, .unwrap_or("");
"java" => tree_sitter_java::LANGUAGE,
"c" | "h" => tree_sitter_c::LANGUAGE, let language = match ext {
"cpp" | "cc" | "cxx" | "hpp" | "hxx" => tree_sitter_cpp::LANGUAGE, "rs" => tree_sitter_rust::LANGUAGE,
"go" => tree_sitter_go::LANGUAGE, "ts" | "tsx" | "js" | "jsx" => tree_sitter_typescript::LANGUAGE_TYPESCRIPT,
_ => { "py" => tree_sitter_python::LANGUAGE,
return Err(crate::error::AppError::Internal(format!( "java" => tree_sitter_java::LANGUAGE,
"Unsupported language for AST replacement: {}", "c" | "h" => tree_sitter_c::LANGUAGE,
ext "cpp" | "cc" | "cxx" | "hpp" | "hxx" => tree_sitter_cpp::LANGUAGE,
))); "go" => tree_sitter_go::LANGUAGE,
_ => {
return Err(crate::error::AppError::Internal(format!(
"Unsupported language for AST replacement: {}",
ext
)));
}
};
parser
.set_language(&language.into())
.map_err(|e| crate::error::AppError::Internal(e.to_string()))?;
let tree = parser.parse(&code, None).ok_or_else(|| {
crate::error::AppError::Internal("Failed to parse code".to_string())
})?;
fn matches_node_type(actual_kind: &str, requested_type: &str) -> bool {
if actual_kind == requested_type {
return true;
}
match requested_type.to_lowercase().as_str() {
"function" | "func" | "fn" | "method" | "def" => matches!(
actual_kind,
"function_item"
| "function_declaration"
| "function_definition"
| "method_definition"
| "function"
),
"struct" => matches!(
actual_kind,
"struct_item" | "struct_declaration" | "struct_specifier"
),
"class" => matches!(
actual_kind,
"class_declaration" | "class_definition" | "class_item"
),
"enum" => matches!(actual_kind, "enum_item" | "enum_declaration"),
"trait" | "interface" => matches!(
actual_kind,
"trait_item" | "interface_declaration" | "interface_item"
),
"impl" | "impl_item" => actual_kind == "impl_item",
"type" | "type_alias" => matches!(
actual_kind,
"type_alias_declaration" | "type_item" | "type_definition"
),
_ => false,
}
} }
};
parser fn find_node<'a>(
.set_language(&language.into()) node: Node<'a>,
.map_err(|e| crate::error::AppError::Internal(e.to_string()))?; code: &str,
let tree = parser.parse(&code, None).ok_or_else(|| { target_type: &str,
crate::error::AppError::Internal("Failed to parse code".to_string()) target_name: &str,
})?; ) -> Option<Node<'a>> {
if matches_node_type(node.kind(), target_type) {
fn matches_node_type(actual_kind: &str, requested_type: &str) -> bool { let mut cursor = node.walk();
if actual_kind == requested_type { for child in node.children(&mut cursor) {
return true; let kind = child.kind();
} if kind == "identifier"
match requested_type.to_lowercase().as_str() { || kind == "name"
"function" | "func" | "fn" | "method" | "def" => matches!( || kind == "type_identifier"
actual_kind, || kind == "property_identifier"
"function_item" || kind == "field_identifier"
| "function_declaration" {
| "function_definition" let name = child.utf8_text(code.as_bytes()).unwrap_or("");
| "method_definition" if name == target_name {
| "function" return Some(node);
), }
"struct" => matches!(
actual_kind,
"struct_item" | "struct_declaration" | "struct_specifier"
),
"class" => matches!(
actual_kind,
"class_declaration" | "class_definition" | "class_item"
),
"enum" => matches!(actual_kind, "enum_item" | "enum_declaration"),
"trait" | "interface" => matches!(
actual_kind,
"trait_item" | "interface_declaration" | "interface_item"
),
"impl" | "impl_item" => actual_kind == "impl_item",
"type" | "type_alias" => matches!(
actual_kind,
"type_alias_declaration" | "type_item" | "type_definition"
),
_ => false,
}
}
// Search for the node
fn find_node<'a>(
node: Node<'a>,
code: &str,
target_type: &str,
target_name: &str,
) -> Option<Node<'a>> {
if matches_node_type(node.kind(), target_type) {
// Try to find the name/identifier
let mut cursor = node.walk();
for child in node.children(&mut cursor) {
let kind = child.kind();
if kind == "identifier"
|| kind == "name"
|| kind == "type_identifier"
|| kind == "property_identifier"
|| kind == "field_identifier"
{
let name = child.utf8_text(code.as_bytes()).unwrap_or("");
if name == target_name {
return Some(node);
} }
} }
} }
}
let mut cursor = node.walk(); let mut cursor = node.walk();
for child in node.children(&mut cursor) { for child in node.children(&mut cursor) {
if let Some(found) = find_node(child, code, target_type, target_name) { if let Some(found) = find_node(child, code, target_type, target_name) {
return Some(found); return Some(found);
}
} }
} None
None
}
let target_node = find_node(
tree.root_node(),
&code,
&tool_args.node_type,
&tool_args.node_name,
);
if let Some(node) = target_node {
let start_byte = node.start_byte();
let end_byte = node.end_byte();
if !code.is_char_boundary(start_byte) || !code.is_char_boundary(end_byte) {
return Err(crate::error::AppError::Internal(format!(
"Byte offsets {}..{} do not fall on UTF-8 character boundaries in {}",
start_byte, end_byte, file_path
)));
} }
let mut new_file_content = let target_node = find_node(
String::with_capacity(code.len() + tool_args.new_content.len()); tree.root_node(),
new_file_content.push_str(&code[..start_byte]); &code,
new_file_content.push_str(&tool_args.new_content); &rep.node_type,
new_file_content.push_str(&code[end_byte..]); &rep.node_name,
);
let target_path = std::path::PathBuf::from(&file_path); if let Some(node) = target_node {
let parent_dir = target_path let start_byte = node.start_byte();
.parent() let end_byte = node.end_byte();
.unwrap_or_else(|| std::path::Path::new("."));
let temp_file_path = if !code.is_char_boundary(start_byte) || !code.is_char_boundary(end_byte) {
parent_dir.join(format!(".tmp_ast_{}.tmp", uuid::Uuid::new_v4()));
std::fs::write(&temp_file_path, new_file_content).map_err(|e| crate::error::AppError::Internal(format!("Failed to write temporary file: {}", e)))?;
if let Err(e) = std::fs::rename(&temp_file_path, &target_path) {
// On Windows, std::fs::rename fails if the target file already exists.
// Fall back to copy-and-remove to ensure atomic-like overwrite behavior.
if let Err(copy_err) = std::fs::copy(&temp_file_path, &target_path) {
let _ = std::fs::remove_file(&temp_file_path);
return Err(crate::error::AppError::Internal(format!( return Err(crate::error::AppError::Internal(format!(
"Failed to atomically overwrite {}: rename failed ({}), copy failed ({})", "Byte offsets {}..{} do not fall on UTF-8 character boundaries in {}",
file_path, e, copy_err start_byte, end_byte, file_path
))); )));
} }
let _ = std::fs::remove_file(&temp_file_path);
let mut new_file_content =
String::with_capacity(code.len() + rep.new_content.len());
new_file_content.push_str(&code[..start_byte]);
new_file_content.push_str(&rep.new_content);
new_file_content.push_str(&code[end_byte..]);
let target_path = std::path::PathBuf::from(&file_path);
let parent_dir = target_path
.parent()
.unwrap_or_else(|| std::path::Path::new("."));
let temp_file_path =
parent_dir.join(format!(".tmp_ast_{}.tmp", uuid::Uuid::new_v4()));
std::fs::write(&temp_file_path, new_file_content).map_err(|e| crate::error::AppError::Internal(format!("Failed to write temporary file: {}", e)))?;
if let Err(e) = std::fs::rename(&temp_file_path, &target_path) {
if let Err(copy_err) = std::fs::copy(&temp_file_path, &target_path) {
let _ = std::fs::remove_file(&temp_file_path);
return Err(crate::error::AppError::Internal(format!(
"Failed to atomically overwrite {}: rename failed ({}), copy failed ({})",
file_path, e, copy_err
)));
}
let _ = std::fs::remove_file(&temp_file_path);
}
msgs.push(format!(
"Successfully replaced node {} of type {} in {}",
rep.node_name, rep.node_type, file_path
));
} else {
return Err(crate::error::AppError::Internal(format!(
"Could not find node {} of type {} in {}",
rep.node_name, rep.node_type, file_path
)));
} }
Ok(format!(
"Successfully replaced node {} of type {} in {}",
tool_args.node_name, tool_args.node_type, file_path
))
} else {
Err(crate::error::AppError::Internal(format!(
"Could not find node {} of type {}",
tool_args.node_name, tool_args.node_type
)))
} }
Ok(msgs.join("\n"))
}) })
.await .await
.map_err(|e| crate::error::AppError::Internal(format!("Task panic: {}", e)))??; .map_err(|e| crate::error::AppError::Internal(format!("Task panic: {}", e)))??;
@@ -749,10 +754,14 @@ mod tests {
let handler = ReplaceAstNodeHandler; let handler = ReplaceAstNodeHandler;
let args = json!({ let args = json!({
"file_path": file_path.to_str().unwrap(), "replacements": [
"node_type": "function_item", {
"node_name": "my_func", "file_path": file_path.to_str().unwrap(),
"new_content": "fn my_func() {\n let x = 2;\n}" "node_type": "function_item",
"node_name": "my_func",
"new_content": "fn my_func() {\n let x = 2;\n}"
}
]
}); });
let res = handler.execute(args, state.clone()).await.unwrap(); let res = handler.execute(args, state.clone()).await.unwrap();
+59 -47
View File
@@ -218,19 +218,25 @@ impl McpTool for UpdateEntitiesHandler {
state.modify_graph(|g| { state.modify_graph(|g| {
for update in req.updates { for update in req.updates {
if !g.entities.contains_key(&update.name) { let mut target_name = update.name.clone();
not_found.push(update.name.clone()); if !g.entities.contains_key(&target_name) {
continue; let lower_target = target_name.to_lowercase();
if let Some(matched_key) = g.entities.keys().find(|k| k.to_lowercase() == lower_target).cloned() {
target_name = matched_key;
} else {
not_found.push(update.name.clone());
continue;
}
} }
if let Some(new_name) = &update.new_name { if let Some(new_name) = &update.new_name {
if update.name != *new_name && g.entities.contains_key(new_name) { if target_name != *new_name && g.entities.contains_key(new_name) {
conflict_names.push(new_name.clone()); conflict_names.push(new_name.clone());
continue; continue;
} }
} }
if let Some(mut entity) = g.entities.remove(&update.name) { if let Some(mut entity) = g.entities.remove(&target_name) {
let mut renamed = false; let mut renamed = false;
let old_name = entity.name.clone(); let old_name = entity.name.clone();
@@ -269,18 +275,6 @@ impl McpTool for UpdateEntitiesHandler {
} }
}); });
if !not_found.is_empty() {
return Err(crate::error::AppError::Internal(format!(
"Error: Entities not found: {}", not_found.join(", ")
)));
}
if !conflict_names.is_empty() {
return Err(crate::error::AppError::Internal(format!(
"Error: Cannot rename to existing entity names: {}", conflict_names.join(", ")
)));
}
let idx = state.get_search_index().await; let idx = state.get_search_index().await;
for old_name in deleted_names { for old_name in deleted_names {
drop(idx.delete_document(&old_name)); drop(idx.delete_document(&old_name));
@@ -291,7 +285,16 @@ impl McpTool for UpdateEntitiesHandler {
} }
let names: Vec<String> = updated_entities.iter().map(|e| e.name.clone()).collect(); let names: Vec<String> = updated_entities.iter().map(|e| e.name.clone()).collect();
Ok(format!("Successfully updated {} entity/entities: {}", names.len(), names.join(", "))) let mut msg = format!("Successfully updated {} entity/entities: {}", names.len(), names.join(", "));
if !not_found.is_empty() {
msg.push_str(&format!("\nNote: {} entities were not found and skipped: {}", not_found.len(), not_found.join(", ")));
}
if !conflict_names.is_empty() {
msg.push_str(&format!("\nNote: {} entity renames were skipped due to name conflicts: {}", conflict_names.len(), conflict_names.join(", ")));
}
Ok(msg)
} }
} }
@@ -459,37 +462,43 @@ impl McpTool for DeleteEntitiesHandler {
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> crate::error::Result<String> { async fn execute(&self, args: Value, state: Arc<MemoryState>) -> crate::error::Result<String> {
let req: DeleteEntitiesTool = serde_json::from_value(args).map_err(|e| crate::error::AppError::BadRequest(format!("Schema validation failed. Your JSON arguments do not match the expected tool schema: {}", e)))?; let req: DeleteEntitiesTool = serde_json::from_value(args).map_err(|e| crate::error::AppError::BadRequest(format!("Schema validation failed. Your JSON arguments do not match the expected tool schema: {}", e)))?;
let to_delete: std::collections::HashSet<_> = req.entity_names.into_iter().collect();
let mut missing = Vec::new(); let mut actual_deletes = Vec::new();
state.read_graph(|g| { let mut not_found = Vec::new();
for name in &to_delete {
if !g.entities.contains_key(name) {
missing.push(name.clone());
}
}
});
if !missing.is_empty() {
return Err(crate::error::AppError::Internal(format!(
"Error: Entities not found: {}. Please use the search_nodes or read_graph tools to verify the exact entity names.",
missing.join(", ")
)));
}
state.modify_graph(|master| { state.modify_graph(|master| {
for name in &to_delete { for target in req.entity_names {
if master.entities.contains_key(&target) {
actual_deletes.push(target);
} else {
let lower_target = target.to_lowercase();
if let Some(matched_key) = master.entities.keys().find(|k| k.to_lowercase() == lower_target).cloned() {
actual_deletes.push(matched_key);
} else {
not_found.push(target);
}
}
}
for name in &actual_deletes {
master.entities.remove(name); master.entities.remove(name);
} }
let delete_set: std::collections::HashSet<_> = actual_deletes.iter().cloned().collect();
master master
.relations .relations
.retain(|r| !to_delete.contains(&r.from) && !to_delete.contains(&r.to)); .retain(|r| !delete_set.contains(&r.from) && !delete_set.contains(&r.to));
}); });
let idx = state.get_search_index().await; let idx = state.get_search_index().await;
for name in to_delete { for name in &actual_deletes {
drop(idx.delete_document(&name)); drop(idx.delete_document(name));
} }
Ok("Entities deleted".to_string())
let mut msg = format!("Successfully deleted {} entities.", actual_deletes.len());
if !not_found.is_empty() {
msg.push_str(&format!(" Note: {} entities were not found and skipped: {}", not_found.len(), not_found.join(", ")));
}
Ok(msg)
} }
} }
@@ -581,7 +590,12 @@ impl McpTool for DeleteRelationsHandler {
let initial_len = master.relations.len(); let initial_len = master.relations.len();
master.relations.retain(|r| { master.relations.retain(|r| {
let should_delete = req.relations.iter().any(|target| { let should_delete = req.relations.iter().any(|target| {
target.from == r.from && target.to == r.to && target.relation_type == r.relation_type if target.from == r.from && target.to == r.to && target.relation_type == r.relation_type {
return true;
}
target.from.eq_ignore_ascii_case(&r.from)
&& target.to.eq_ignore_ascii_case(&r.to)
&& target.relation_type.eq_ignore_ascii_case(&r.relation_type)
}); });
!should_delete !should_delete
}); });
@@ -589,13 +603,11 @@ impl McpTool for DeleteRelationsHandler {
}); });
let missing_count = requested_count.saturating_sub(deleted_count); let missing_count = requested_count.saturating_sub(deleted_count);
let mut msg = format!("Successfully deleted {} relations.", deleted_count);
if missing_count > 0 { if missing_count > 0 {
return Err(crate::error::AppError::Internal(format!( msg.push_str(&format!(" Note: {} relations were not found and skipped.", missing_count));
"Error: {} relation(s) not found in graph. Please verify exact relation properties (from, to, relation_type) using read_graph or get_subgraph.",
missing_count
)));
} }
Ok("Relations deleted".to_string()) Ok(msg)
} }
} }
@@ -1601,7 +1613,7 @@ mod tests {
.await .await
.map_err(|e| crate::error::AppError::Internal(e.to_string())) .map_err(|e| crate::error::AppError::Internal(e.to_string()))
.unwrap(); .unwrap();
assert_eq!(res4, "Entities deleted"); assert_eq!(res4, "Successfully deleted 1 entities.");
let res5 = read_graph let res5 = read_graph
.execute(json!({"namespace": "global"}), state.clone()) .execute(json!({"namespace": "global"}), state.clone())
@@ -1774,7 +1786,7 @@ mod tests {
) )
.await .await
.unwrap(); .unwrap();
assert_eq!(del_rel_res, "Relations deleted"); assert_eq!(del_rel_res, "Successfully deleted 1 relations.");
let bcast_handler = AgentSignalsHandler; let bcast_handler = AgentSignalsHandler;
let bcast_res = bcast_handler let bcast_res = bcast_handler
+8 -2
View File
@@ -498,8 +498,8 @@ pub struct ReadFileSkeletonTool {
pub file_path: String, pub file_path: String,
} }
/// Replace a specific AST node in a file (robust structural editing). Use this instead of regex or line-based string replacement to prevent indentation bugs and matching failures. /// Replace a specific AST node in a file (robust structural editing). Use this instead of regex or line-based string replacement to prevent indentation bugs and matching failures.
#[derive(Debug, Deserialize, Serialize, JsonSchema)] #[derive(Debug, Deserialize, Serialize, JsonSchema, Clone)]
pub struct ReplaceAstNodeTool { pub struct AstNodeReplacement {
/// The path of the file to modify. /// The path of the file to modify.
pub file_path: String, pub file_path: String,
/// The AST node type to replace (e.g., 'function_item', 'impl_item'). /// The AST node type to replace (e.g., 'function_item', 'impl_item').
@@ -510,6 +510,12 @@ pub struct ReplaceAstNodeTool {
pub new_content: String, pub new_content: String,
} }
#[derive(Debug, Deserialize, Serialize, JsonSchema)]
pub struct ReplaceAstNodeTool {
/// Array of AST node replacements to execute in bulk.
pub replacements: Vec<AstNodeReplacement>,
}
/// Semantic code search using local vector embeddings. Use this conceptual search instead of raw regex (grep) when trying to locate abstract logic or exploring new patterns. /// Semantic code search using local vector embeddings. Use this conceptual search instead of raw regex (grep) when trying to locate abstract logic or exploring new patterns.
#[derive(Debug, Deserialize, Serialize, JsonSchema)] #[derive(Debug, Deserialize, Serialize, JsonSchema)]
pub struct SemanticCodeSearchTool { pub struct SemanticCodeSearchTool {