Files
mcp-memory/server/src/handlers/ast.rs
T

576 lines
21 KiB
Rust

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};
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::<ReadFileSkeletonTool>(
"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<MemoryState>) -> crate::error::Result<String> {
let tool_args: ReadFileSkeletonTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
let file_path = tool_args.file_path.clone();
let result = tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
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 Ok(code),
};
parser
.set_language(&language.into())
.map_err(|e| 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_structural = matches!(
kind,
"use_declaration"
| "import_statement"
| "import_from_statement"
| "struct_item"
| "enum_item"
| "trait_item"
| "impl_item"
| "function_item"
| "function_declaration"
| "function_definition"
| "method_definition"
| "interface_declaration"
| "type_alias_declaration"
| "class_declaration"
| "class_definition"
);
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 + 1);
}
}
}
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::<ReplaceAstNodeTool>(
"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<MemoryState>) -> crate::error::Result<String> {
let tool_args: ReplaceAstNodeTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
let file_path = tool_args.file_path.clone();
let result = tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
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| 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"
),
"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 == "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();
let mut new_file_content = String::new();
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..]);
std::fs::write(&file_path, new_file_content).map_err(|e| e.to_string())?;
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) -> Vec<serde_json::Value> {
let mut results = Vec::new();
let cwd = match std::env::current_dir() {
Ok(dir) => dir,
Err(_) => return results,
};
let walker = ignore::WalkBuilder::new(&cwd)
.hidden(true)
.git_ignore(true)
.build();
let mut scanned_files = 0;
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(content) = std::fs::read_to_string(path) {
for (line_num, line) in content.lines().enumerate() {
let is_match = if filter_fn_call {
line.contains(&format!("{}(", target_sym))
|| line.contains(&format!("{}.await", target_sym))
} else {
line.contains(target_sym)
};
if is_match {
results.push(serde_json::json!({
"file_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::<crate::tools::FindSymbolReferencesTool>(
"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<MemoryState>) -> crate::error::Result<String> {
let req: crate::tools::FindSymbolReferencesTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
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::<Vec<serde_json::Value>, crate::error::AppError>(refs)
})?;
if matches.len() < limit {
let remaining = limit - matches.len();
let disk_matches = tokio::task::spawn_blocking(move || {
scan_workspace_for_symbol(&target_sym, remaining, false)
})
.await
.unwrap_or_default();
matches.extend(disk_matches);
}
Ok(serde_json::to_string_pretty(&matches)?)
}
}
pub struct GetCallersHandler;
#[async_trait]
impl McpTool for GetCallersHandler {
fn name(&self) -> &'static str {
"get_callers"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<crate::tools::GetCallersTool>(
"get_callers",
"Find all caller functions or methods that invoke a specified target function name.",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> crate::error::Result<String> {
let req: crate::tools::GetCallersTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
let limit = req.limit.unwrap_or(10);
let target_fn = req.function_name.clone();
let mut callers = state.code.snippets.read_with(|snippets| {
let mut matching = Vec::new();
for snippet in snippets {
if snippet.code.contains(&format!("{}(", target_fn))
|| snippet.code.contains(&format!("{}.await", target_fn))
{
matching.push(serde_json::json!({
"source": "snippet",
"name": snippet.name,
"code": snippet.code,
}));
if matching.len() >= limit {
break;
}
}
}
Ok::<Vec<serde_json::Value>, crate::error::AppError>(matching)
})?;
if callers.len() < limit {
let remaining = limit - callers.len();
let disk_callers = tokio::task::spawn_blocking(move || {
scan_workspace_for_symbol(&target_fn, remaining, true)
})
.await
.unwrap_or_default();
callers.extend(disk_callers);
}
Ok(serde_json::to_string_pretty(&callers)?)
}
}
pub struct AnalyzeImpactHandler;
#[async_trait]
impl McpTool for AnalyzeImpactHandler {
fn name(&self) -> &'static str {
"analyze_impact"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<crate::tools::AnalyzeImpactTool>(
"analyze_impact",
"Analyze the potential downstream breaking impact of modifying a function, struct, or file.",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> crate::error::Result<String> {
let req: crate::tools::AnalyzeImpactTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
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 sym_clone = sym.clone();
let disk_refs = tokio::task::spawn_blocking(move || {
scan_workspace_for_symbol(&sym_clone, 20, false)
})
.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 result = serde_json::json!({
"target_symbol": sym,
"risk_level": risk_level,
"ast_callers_count": caller_count,
"ast_callers_sample": callers.into_iter().take(10).collect::<Vec<_>>(),
"graph_relations_count": graph_count,
"graph_relations": kg_connected,
"recommendation": 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.",
}
});
Ok(serde_json::to_string_pretty(&result)?)
}
}
#[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;"));
}
}