Files
mcp-memory/server/src/indexer.rs
T

308 lines
12 KiB
Rust

use crate::embedding::generate_embeddings_async;
use crate::models::Snippet;
use crate::state::MemoryState;
use ignore::WalkBuilder;
use std::sync::Arc;
use tree_sitter::{Node, Parser};
pub async fn start_background_indexer(state: Arc<MemoryState>) {
// Determine workspace root. Since mcp-memory is usually run from its project root,
// current_dir is a good default for the git repository.
let workspace_root = std::env::current_dir().unwrap_or_else(|_| state.base_dir.clone());
tokio::spawn(async move {
tracing::info!("Starting background indexer in {:?}", workspace_root);
let root_clone = workspace_root.clone();
let files_to_process = tokio::task::spawn_blocking(move || {
let walker = WalkBuilder::new(&root_clone)
.hidden(true)
.git_ignore(true)
.build();
let mut files = Vec::new();
for result in walker {
match result {
Ok(entry) => {
if entry.file_type().is_some_and(|ft| ft.is_file()) {
let path = entry.path().to_path_buf();
let ext = path.extension().and_then(|e| e.to_str()).unwrap_or("");
if [
"rs", "ts", "js", "jsx", "tsx", "py", "java", "c", "cpp", "go",
]
.contains(&ext)
{
files.push(path);
}
}
}
Err(e) => {
tracing::warn!("Error walking directory: {}", e);
}
}
}
files
})
.await
.unwrap_or_default();
let idx = state.get_search_index().await;
for file_path in files_to_process {
if let Ok(content) = std::fs::read_to_string(&file_path) {
let ext = file_path.extension().and_then(|e| e.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,
_ => continue,
};
let mut parser = Parser::new();
if parser.set_language(&language.into()).is_err() {
continue;
}
if let Some(tree) = parser.parse(&content, None) {
let mut chunks = Vec::new();
extract_chunks(tree.root_node(), &content, &mut chunks, ext);
// Gold Standard: Batch generate embeddings in chunks of 16 to eliminate sequential HTTP overhead
for chunk_batch in chunks.chunks(16) {
let texts: Vec<String> = chunk_batch
.iter()
.map(|(_, code, _)| code.clone())
.collect();
let embeddings = generate_embeddings_async(texts).await.unwrap_or_default();
let mut new_snippets = Vec::with_capacity(chunk_batch.len());
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
for (i, (name, code, desc)) in chunk_batch.iter().enumerate() {
let embedding = embeddings.get(i).cloned();
let file_name =
file_path.file_name().unwrap_or_default().to_string_lossy();
let snippet_name = format!("{}:{}", file_name, name);
let snippet = Snippet {
name: snippet_name,
language: ext.to_string(),
code: code.clone(),
description: format!("{} in {}", desc, file_path.display()),
updated_at: now,
tags: vec![],
embedding,
..Default::default()
};
new_snippets.push(snippet);
}
// Gold Standard: Modify store ONCE per batch with zero-copy HashSet<&str> lookup
let mut snippets_to_index = Vec::new();
state.code.snippets.modify(|snippets| {
let existing_names: std::collections::HashSet<&str> =
snippets.iter().map(|s| s.name.as_str()).collect();
let mut filtered_new = Vec::with_capacity(new_snippets.len());
let mut seen_in_batch = std::collections::HashSet::new();
for snippet in new_snippets {
if !existing_names.contains(snippet.name.as_str())
&& seen_in_batch.insert(snippet.name.clone())
{
filtered_new.push(snippet);
}
}
for snippet in filtered_new {
snippets.push(snippet.clone());
snippets_to_index.push(snippet);
}
if snippets.len() > 1000 {
let overflow = snippets.len() - 1000;
snippets.drain(0..overflow);
}
});
let threshold: usize = std::env::var("MCP_MEMORY_CONDENSE_THRESHOLD")
.unwrap_or_else(|_| "100".to_string())
.parse()
.unwrap_or(100);
if state.code.snippets.read_with(|s| s.len()) > threshold {
state.condense_notify.notify_one();
}
for snippet in &snippets_to_index {
let _ = idx.index_snippet(snippet).await;
}
}
}
}
}
tracing::info!("Background indexing completed.");
});
}
pub fn extract_chunks(
node: Node,
code: &str,
chunks: &mut Vec<(String, String, String)>,
ext: &str,
) {
extract_chunks_with_parent(node, code, chunks, ext, None, 0);
}
fn extract_chunks_with_parent(
node: Node,
code: &str,
chunks: &mut Vec<(String, String, String)>,
_ext: &str,
parent_scope: Option<&str>,
depth: usize,
) {
// Stack overflow protection: Cap recursion depth at 100
if depth > 100 {
return;
}
let kind = node.kind();
let is_impl_or_class = matches!(kind, "impl_item" | "class_declaration" | "class_definition");
let current_scope: Option<&str> = if is_impl_or_class {
let mut cursor = node.walk();
let mut type_name = None;
for child in node.children(&mut cursor) {
if child.kind() == "type_identifier"
|| child.kind() == "name"
|| child.kind() == "identifier"
{
type_name = child.utf8_text(code.as_bytes()).ok();
break;
}
}
type_name.or(parent_scope)
} else {
parent_scope
};
let is_structural = matches!(
kind,
"function_item"
| "function_declaration"
| "function_definition"
| "method_definition"
| "struct_item"
| "class_declaration"
);
if is_structural {
let mut raw_text = node.utf8_text(code.as_bytes()).unwrap_or("").to_string();
let mut name = "unknown";
let mut cursor = node.walk();
for child in node.children(&mut cursor) {
let child_kind = child.kind();
if child_kind == "identifier" || child_kind == "name" || child_kind == "type_identifier"
{
if let Ok(text) = child.utf8_text(code.as_bytes()) {
name = text;
}
break;
}
}
let mut final_name = name.to_string();
if let Some(scope) = current_scope {
raw_text = format!("// Parent Scope: {}\n{}", scope, raw_text);
final_name = format!("{}::{}", scope, name);
}
let desc = format!("{} AST node", kind);
chunks.push((final_name, raw_text, desc));
} else {
let mut cursor = node.walk();
for child in node.named_children(&mut cursor) {
extract_chunks_with_parent(child, code, chunks, _ext, current_scope, depth + 1);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use tree_sitter::Parser;
#[test]
fn test_extract_chunks_rust_function() {
let code = "fn test_func() { println!(\"hello\"); }";
let mut parser = Parser::new();
parser
.set_language(&tree_sitter_rust::LANGUAGE.into())
.unwrap();
let tree = parser.parse(code, None).unwrap();
let mut chunks = Vec::new();
extract_chunks(tree.root_node(), code, &mut chunks, "rs");
assert_eq!(chunks.len(), 1);
assert_eq!(chunks[0].0, "test_func");
assert!(chunks[0].1.contains("println"));
assert_eq!(chunks[0].2, "function_item AST node");
}
#[test]
fn test_extract_chunks_python_function() {
let code = "def my_python_func():\n pass\n";
let mut parser = Parser::new();
parser
.set_language(&tree_sitter_python::LANGUAGE.into())
.unwrap();
let tree = parser.parse(code, None).unwrap();
let mut chunks = Vec::new();
extract_chunks(tree.root_node(), code, &mut chunks, "py");
assert_eq!(chunks.len(), 1);
assert_eq!(chunks[0].0, "my_python_func");
assert_eq!(chunks[0].2, "function_definition AST node");
}
#[tokio::test]
async fn test_start_background_indexer_lifecycle() {
let temp_dir = tempfile::tempdir().unwrap();
let state = Arc::new(MemoryState::new(temp_dir.path().to_str().unwrap()));
start_background_indexer(state).await;
}
#[test]
fn test_extract_chunks_rust_impl_block() {
let code = "impl MyStruct { fn my_method(&self) {} }";
let mut parser = Parser::new();
parser
.set_language(&tree_sitter_rust::LANGUAGE.into())
.unwrap();
let tree = parser.parse(code, None).unwrap();
let mut chunks = Vec::new();
extract_chunks(tree.root_node(), code, &mut chunks, "rs");
assert_eq!(chunks.len(), 1);
assert_eq!(chunks[0].0, "MyStruct::my_method");
assert!(chunks[0].1.contains("// Parent Scope: MyStruct"));
}
#[tokio::test]
async fn test_start_background_indexer_empty_dir() {
let temp_dir = tempfile::tempdir().unwrap();
let state = Arc::new(MemoryState::new(temp_dir.path().to_str().unwrap()));
start_background_indexer(state).await;
}
}