use crate::router::McpTool; use crate::state::MemoryState; use crate::tools::ReadFileSkeletonTool; use async_trait::async_trait; use serde_json::Value; use std::sync::Arc; use tree_sitter::{Node, Parser}; fn validate_safe_path(path_str: &str) -> crate::error::Result<()> { if path_str.trim().is_empty() || path_str.contains('\0') { return Err(crate::error::AppError::Internal( "Invalid file path: path is empty or contains null characters".to_string(), )); } let path = std::path::Path::new(path_str); for component in path.components() { if component == std::path::Component::ParentDir { return Err(crate::error::AppError::Internal(format!( "Path traversal forbidden: '{}' contains parent directory relative components", path_str ))); } } Ok(()) } pub struct ReadFileSkeletonHandler; #[async_trait] impl McpTool for ReadFileSkeletonHandler { fn name(&self) -> &'static str { "read_file_skeleton" } fn schema(&self) -> Value { crate::mcp::tool_def::( "read_file_skeleton", "Read a source file and return only its AST structural skeleton, omitting implementation details to save tokens.", ) } async fn execute(&self, args: Value, _state: Arc) -> crate::error::Result { let tool_args: ReadFileSkeletonTool = 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(); validate_safe_path(&file_path)?; let result = tokio::task::spawn_blocking(move || -> crate::error::Result { 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) .extension() .and_then(|s| s.to_str()) .unwrap_or(""); let language = match ext { "rs" => tree_sitter_rust::LANGUAGE, "ts" | "tsx" | "js" | "jsx" => tree_sitter_typescript::LANGUAGE_TYPESCRIPT, "py" => tree_sitter_python::LANGUAGE, "java" => tree_sitter_java::LANGUAGE, "c" | "h" => tree_sitter_c::LANGUAGE, "cpp" | "cc" | "cxx" | "hpp" | "hxx" => tree_sitter_cpp::LANGUAGE, "go" => tree_sitter_go::LANGUAGE, _ => return Ok(code), }; let mut parser = tree_sitter::Parser::new(); 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()) })?; let mut result_skeleton = String::with_capacity(code.len() / 2); fn extract_skeleton(node: Node, code: &str, out: &mut String, depth: usize) { if depth > 128 { return; } let kind = node.kind(); let is_container = matches!( kind, "impl_item" | "class_declaration" | "class_definition" | "trait_item" | "interface_declaration" ); let is_structural = is_container || matches!( kind, "use_declaration" | "import_statement" | "import_from_statement" | "struct_item" | "enum_item" | "function_item" | "function_declaration" | "function_definition" | "method_definition" | "type_alias_declaration" ); if is_container { let indent = " ".repeat(depth); let node_text = node.utf8_text(code.as_bytes()).unwrap_or(""); let mut header = String::new(); for line in node_text.lines() { let trimmed = line.trim(); if trimmed.ends_with('{') || trimmed.ends_with(':') { header.push_str(line); break; } else { header.push_str(line); header.push('\n'); } } if header.is_empty() && let Some(first_line) = node_text.lines().next() { header = first_line.to_string(); } out.push_str(&indent); out.push_str(header.trim()); out.push('\n'); let mut cursor = node.walk(); for child in node.named_children(&mut cursor) { extract_skeleton(child, code, out, depth + 1); } if header.trim().ends_with('{') { out.push_str(&indent); out.push_str("}\n"); } } else if is_structural { let indent = " ".repeat(depth); let node_text = node.utf8_text(code.as_bytes()).unwrap_or(""); let mut signature = String::new(); for line in node_text.lines() { let trimmed = line.trim(); if trimmed.ends_with('{') || trimmed.ends_with(':') { signature.push_str(line); signature.push_str(" ... }"); break; } else { signature.push_str(line); signature.push('\n'); } } if signature.is_empty() { signature = node_text.to_string(); } out.push_str(&indent); out.push_str(signature.trim()); out.push('\n'); } else if node.is_named() { let mut cursor = node.walk(); for child in node.named_children(&mut cursor) { extract_skeleton(child, code, out, depth); } } } extract_skeleton(tree.root_node(), &code, &mut result_skeleton, 0); if result_skeleton.is_empty() { Ok(code) } else { Ok(result_skeleton) } }) .await .map_err(|e| crate::error::AppError::Internal(format!("Task panic: {}", e)))??; Ok(result) } } use crate::tools::ReplaceAstNodeTool; pub struct ReplaceAstNodeHandler; #[async_trait] impl McpTool for ReplaceAstNodeHandler { fn name(&self) -> &'static str { "replace_ast_node" } fn schema(&self) -> Value { crate::mcp::tool_def::( "replace_ast_node", "Replace a specific AST node (e.g., function, struct, enum, class, trait) entirely using tree-sitter for robust structural editing. Supported node_type values include: 'function_item' (or 'function'/'fn'/'method'), 'struct_item' (or 'struct'), 'class_declaration' (or 'class'), 'enum_item' (or 'enum'), 'trait_item' (or 'trait'/'interface'), 'type_alias_declaration' (or 'type').", ) } async fn execute(&self, args: Value, _state: Arc) -> crate::error::Result { 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)))?; let file_path = tool_args.file_path.clone(); validate_safe_path(&file_path)?; let result = tokio::task::spawn_blocking(move || -> crate::error::Result { let code = std::fs::read_to_string(&file_path).map_err(|e| { crate::error::AppError::Internal(format!("Failed to read file: {}", e)) })?; let mut parser = Parser::new(); let ext = std::path::Path::new(&file_path) .extension() .and_then(|s| s.to_str()) .unwrap_or(""); let language = match ext { "rs" => tree_sitter_rust::LANGUAGE, "ts" | "tsx" | "js" | "jsx" => tree_sitter_typescript::LANGUAGE_TYPESCRIPT, "py" => tree_sitter_python::LANGUAGE, "java" => tree_sitter_java::LANGUAGE, "c" | "h" => tree_sitter_c::LANGUAGE, "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, } } // Search for the node fn find_node<'a>( node: Node<'a>, code: &str, target_type: &str, target_name: &str, ) -> Option> { 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(); for child in node.children(&mut cursor) { if let Some(found) = find_node(child, code, target_type, target_name) { return Some(found); } } 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 = String::with_capacity(code.len() + tool_args.new_content.len()); new_file_content.push_str(&code[..start_byte]); new_file_content.push_str(&tool_args.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) { // 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!( "Failed to atomically overwrite {}: rename failed ({}), copy failed ({})", file_path, e, copy_err ))); } let _ = std::fs::remove_file(&temp_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 ))) } }) .await .map_err(|e| crate::error::AppError::Internal(format!("Task panic: {}", e)))??; Ok(result) } } fn scan_workspace_for_symbol( target_sym: &str, limit: usize, filter_fn_call: bool, workspace_dir: Option, ) -> Vec { let mut results = Vec::new(); let scan_dir = workspace_dir.unwrap_or_else(|| { std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from(".")) }); if !scan_dir.exists() { return results; } let walker = ignore::WalkBuilder::new(&scan_dir) .hidden(true) .git_ignore(true) .build(); let mut scanned_files = 0; let call_pattern = format!("{}(", target_sym); let await_pattern = format!("{}.await", target_sym); for result in walker { let entry = match result { Ok(e) => e, Err(_) => continue, }; if entry.file_type().is_some_and(|ft| ft.is_file()) { let path = entry.path(); let ext = path.extension().and_then(|s| s.to_str()).unwrap_or(""); if matches!( ext, "rs" | "ts" | "tsx" | "js" | "jsx" | "py" | "go" | "java" | "c" | "cpp" | "h" | "hpp" ) { scanned_files += 1; if scanned_files > 500 { break; } if let Ok(meta) = std::fs::metadata(path) && meta.len() > 1024 * 1024 { continue; } if let Ok(content) = std::fs::read_to_string(path) { for (line_num, line) in content.lines().enumerate() { let is_match = if filter_fn_call { line.contains(&call_pattern) || line.contains(&await_pattern) } else { line.contains(target_sym) }; if is_match { results.push(serde_json::json!({ "file_path": crate::handlers::utils::sanitize_path(&path.to_string_lossy()), "line": line_num + 1, "content": line.trim(), })); if results.len() >= limit { return results; } } } } } } } results } pub struct FindSymbolReferencesHandler; #[async_trait] impl McpTool for FindSymbolReferencesHandler { fn name(&self) -> &'static str { "find_symbol_references" } fn schema(&self) -> Value { crate::mcp::tool_def::( "find_symbol_references", "Find all source locations and AST chunks where a specific symbol is referenced or called.", ) } async fn execute(&self, args: Value, state: Arc) -> crate::error::Result { let req: crate::tools::FindSymbolReferencesTool = 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 limit = req.limit.unwrap_or(10); let target_sym = req.symbol.clone(); let mut matches = state.code.snippets.read_with(|snippets| { let mut refs = Vec::new(); for snippet in snippets { if snippet.code.contains(&target_sym) || snippet.name.contains(&target_sym) { refs.push(serde_json::json!({ "source": "snippet", "name": snippet.name, "code": snippet.code, })); if refs.len() >= limit { break; } } } Ok::, crate::error::AppError>(refs) })?; let custom_dir = req.workspace_dir.as_ref().map(std::path::PathBuf::from); if matches.len() < limit { let remaining = limit - matches.len(); let target_sym_clone = target_sym.clone(); let disk_matches = tokio::task::spawn_blocking(move || { scan_workspace_for_symbol(&target_sym_clone, remaining, false, custom_dir) }) .await .unwrap_or_default(); matches.extend(disk_matches); } let mut out = String::new(); out.push_str(&format!("## Symbol References for `{}`\n\n", target_sym)); for match_item in &matches { if let Some(source) = match_item.get("source").and_then(|s| s.as_str()) { if source == "snippet" { let name = match_item.get("name").and_then(|n| n.as_str()).unwrap_or("Unknown"); let code = match_item.get("code").and_then(|c| c.as_str()).unwrap_or(""); out.push_str(&format!("### Snippet: {}\n```rust\n{}\n```\n\n", name, code)); } } else { let file = match_item.get("file_path").and_then(|f| f.as_str()).unwrap_or("Unknown"); let line = match_item.get("line").and_then(|l| l.as_u64()).unwrap_or(0); let content = match_item.get("content").and_then(|c| c.as_str()).unwrap_or(""); out.push_str(&format!("- `{}:{}`: `{}`\n", file, line, content)); } } if matches.is_empty() { out.push_str("No references found.\n"); } Ok(out) } } pub struct GetCallersHandler; #[async_trait] impl McpTool for GetCallersHandler { fn name(&self) -> &'static str { "get_callers" } fn schema(&self) -> Value { crate::mcp::tool_def::( "get_callers", "Find all caller functions or methods that invoke a specified target function name.", ) } async fn execute(&self, args: Value, state: Arc) -> crate::error::Result { let req: crate::tools::GetCallersTool = 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 limit = req.limit.unwrap_or(10); let target_fn = req.function_name.clone(); let call_pattern = format!("{}(", target_fn); let await_pattern = format!("{}.await", target_fn); let mut callers = state.code.snippets.read_with(|snippets| { let mut matching = Vec::new(); for snippet in snippets { if snippet.code.contains(&call_pattern) || snippet.code.contains(&await_pattern) { matching.push(serde_json::json!({ "source": "snippet", "name": snippet.name, "code": snippet.code, })); if matching.len() >= limit { break; } } } Ok::, crate::error::AppError>(matching) })?; let custom_dir = req.workspace_dir.as_ref().map(std::path::PathBuf::from); if callers.len() < limit { let remaining = limit - callers.len(); let target_fn_clone = target_fn.clone(); let disk_callers = tokio::task::spawn_blocking(move || { scan_workspace_for_symbol(&target_fn_clone, remaining, true, custom_dir) }) .await .unwrap_or_default(); callers.extend(disk_callers); } let mut out = String::new(); out.push_str(&format!("## Callers for `{}`\n\n", target_fn)); for caller in &callers { if let Some(source) = caller.get("source").and_then(|s| s.as_str()) { if source == "snippet" { let name = caller.get("name").and_then(|n| n.as_str()).unwrap_or("Unknown"); let code = caller.get("code").and_then(|c| c.as_str()).unwrap_or(""); out.push_str(&format!("### Snippet: {}\n```rust\n{}\n```\n\n", name, code)); } } else { let file = caller.get("file_path").and_then(|f| f.as_str()).unwrap_or("Unknown"); let line = caller.get("line").and_then(|l| l.as_u64()).unwrap_or(0); let content = caller.get("content").and_then(|c| c.as_str()).unwrap_or(""); out.push_str(&format!("- `{}:{}`: `{}`\n", file, line, content)); } } if callers.is_empty() { out.push_str("No callers found.\n"); } Ok(out) } } pub struct AnalyzeImpactHandler; #[async_trait] impl McpTool for AnalyzeImpactHandler { fn name(&self) -> &'static str { "analyze_impact" } fn schema(&self) -> Value { crate::mcp::tool_def::( "analyze_impact", "Analyze the potential downstream breaking impact of modifying a function, struct, or file.", ) } async fn execute(&self, args: Value, state: Arc) -> crate::error::Result { let req: crate::tools::AnalyzeImpactTool = 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 sym = req.target_symbol.clone(); let mut callers = Vec::new(); state.code.snippets.read_with(|snippets| { for snippet in snippets { if snippet.code.contains(&sym) { callers.push(snippet.name.clone()); } } }); let custom_dir = req .file_path .as_ref() .and_then(|p| std::path::Path::new(p).parent().map(|p| p.to_path_buf())); let sym_clone = sym.clone(); let disk_refs = tokio::task::spawn_blocking(move || { scan_workspace_for_symbol(&sym_clone, 20, false, custom_dir) }) .await .unwrap_or_default(); for r in &disk_refs { if let Some(path) = r.get("file_path").and_then(|p| p.as_str()) { let line = r.get("line").and_then(|l| l.as_u64()).unwrap_or(0); callers.push(format!("{}:{}", path, line)); } } let mut kg_connected = Vec::new(); state.read_graph(|g| { for rel in &g.relations { if rel.from == sym { kg_connected.push(format!("Outgoing: {} -> {}", rel.relation_type, rel.to)); } else if rel.to == sym { kg_connected.push(format!("Incoming: {} <- {}", rel.relation_type, rel.from)); } } }); let caller_count = callers.len(); let graph_count = kg_connected.len(); let risk_level = if caller_count > 10 || graph_count > 5 { "CRITICAL" } else if caller_count > 3 || graph_count > 2 { "HIGH" } else if caller_count > 0 || graph_count > 0 { "MEDIUM" } else { "LOW" }; let mut out = String::new(); out.push_str(&format!("## Impact Analysis for `{}`\n\n", sym)); out.push_str(&format!("**Risk Level:** {}\n\n", risk_level)); let rec = match risk_level { "CRITICAL" | "HIGH" => "Requires comprehensive unit test verification and backwards compatibility checks before modifying.", "MEDIUM" => "Verify direct call sites and run affected module tests.", _ => "Safe to modify with standard unit test verification.", }; out.push_str(&format!("**Recommendation:** {}\n\n", rec)); out.push_str(&format!("### AST Callers ({} total, showing up to 10)\n", caller_count)); for c in callers.into_iter().take(10) { out.push_str(&format!("- `{}`\n", c)); } if caller_count == 0 { out.push_str("No callers found.\n"); } out.push_str("\n"); out.push_str(&format!("### Graph Relations ({})\n", graph_count)); for g in kg_connected { out.push_str(&format!("- {}\n", g)); } if graph_count == 0 { out.push_str("No graph relations found.\n"); } Ok(out) } } #[cfg(test)] mod tests { use super::*; use serde_json::json; use tempfile::tempdir; #[tokio::test] async fn test_read_file_skeleton() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); let file_path = dir.path().join("test_skeleton.rs"); let code = "fn my_func() {\n let x = 1;\n}\n\nstruct MyStruct {\n val: i32\n}"; std::fs::write(&file_path, code).unwrap(); let handler = ReadFileSkeletonHandler; let args = json!({ "file_path": file_path.to_str().unwrap() }); let res = handler.execute(args, state.clone()).await.unwrap(); assert!(res.contains("fn my_func()")); assert!(res.contains("struct MyStruct")); } #[tokio::test] async fn test_replace_ast_node() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); let file_path = dir.path().join("test_replace.rs"); let code = "fn my_func() {\n let x = 1;\n}\n\nstruct MyStruct {\n val: i32\n}"; std::fs::write(&file_path, code).unwrap(); let handler = ReplaceAstNodeHandler; let args = json!({ "file_path": file_path.to_str().unwrap(), "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(); assert!(res.contains("Successfully replaced node")); let new_code = std::fs::read_to_string(&file_path).unwrap(); assert!(new_code.contains("let x = 2;")); assert!(!new_code.contains("let x = 1;")); } }