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

397 lines
14 KiB
Rust

use crate::models::{Adr, Entity, Snippet, Task};
use std::sync::{Arc, Mutex};
use tantivy::schema::*;
use tantivy::{Index, IndexReader, IndexWriter, ReloadPolicy, doc};
pub type SearchResultTuple = (String, String, String, String, f32);
#[derive(Clone)]
pub struct MemoryIndex {
pub index: Index,
pub reader: IndexReader,
writer: Arc<Mutex<IndexWriter>>,
needs_commit: Arc<std::sync::atomic::AtomicBool>,
// Schema fields
pub id_field: Field,
pub title_field: Field,
pub body_field: Field,
pub type_field: Field,
pub namespace_field: Field,
}
impl MemoryIndex {
pub fn new(store_dir: &std::path::Path) -> tantivy::Result<Self> {
let mut schema_builder = Schema::builder();
let id_field = schema_builder.add_text_field("id", STRING | STORED);
let title_field = schema_builder.add_text_field("title", TEXT | STORED);
let body_field = schema_builder.add_text_field("body", TEXT | STORED);
let type_field = schema_builder.add_text_field("type", STRING | STORED);
let namespace_field = schema_builder.add_text_field("namespace", STRING | STORED);
let schema = schema_builder.build();
let index_dir = store_dir.join("tantivy_index");
std::fs::create_dir_all(&index_dir)
.map_err(|e| tantivy::TantivyError::SystemError(e.to_string()))?;
let index = Index::open_in_dir(&index_dir)
.or_else(|_| Index::create_in_dir(&index_dir, schema.clone()))?;
let mut writer = index.writer(50_000_000)?;
writer.delete_all_documents()?;
writer.commit()?;
let reader = index
.reader_builder()
.reload_policy(ReloadPolicy::OnCommitWithDelay)
.try_into()?;
Ok(Self {
index,
reader,
writer: Arc::new(Mutex::new(writer)),
needs_commit: Arc::new(std::sync::atomic::AtomicBool::new(false)),
id_field,
title_field,
body_field,
type_field,
namespace_field,
})
}
pub fn index_entity(&self, e: &Entity) -> tokio::task::JoinHandle<tantivy::Result<()>> {
let writer = Arc::clone(&self.writer);
let id_field = self.id_field;
let id_val = e.name.clone();
let needs_commit = Arc::clone(&self.needs_commit);
let doc = doc!(
self.id_field => e.name.as_str(),
self.title_field => e.name.as_str(),
self.body_field => e.observations.join(" "),
self.type_field => "entity",
self.namespace_field => e.namespace.as_str()
);
tokio::task::spawn_blocking(move || {
let writer = writer.lock().unwrap_or_else(|e| e.into_inner());
writer.delete_term(tantivy::Term::from_field_text(id_field, &id_val));
writer.add_document(doc)?;
needs_commit.store(true, std::sync::atomic::Ordering::SeqCst);
Ok(())
})
}
pub fn index_task(&self, t: &Task) -> tokio::task::JoinHandle<tantivy::Result<()>> {
let writer = Arc::clone(&self.writer);
let id_field = self.id_field;
let id_val = t.id.clone();
let needs_commit = Arc::clone(&self.needs_commit);
let doc = doc!(
self.id_field => t.id.as_str(),
self.title_field => t.title.as_str(),
self.body_field => format!("{}\n{}", t.description, t.acceptance_criteria.iter().map(|c| c.description.as_str()).collect::<Vec<_>>().join("\n")),
self.type_field => "task",
self.namespace_field => "global"
);
tokio::task::spawn_blocking(move || {
let writer = writer.lock().unwrap_or_else(|e| e.into_inner());
writer.delete_term(tantivy::Term::from_field_text(id_field, &id_val));
writer.add_document(doc)?;
needs_commit.store(true, std::sync::atomic::Ordering::SeqCst);
Ok(())
})
}
pub fn delete_document(&self, id: &str) -> tokio::task::JoinHandle<tantivy::Result<()>> {
let writer = Arc::clone(&self.writer);
let id_field = self.id_field;
let id_val = id.to_string();
let needs_commit = Arc::clone(&self.needs_commit);
tokio::task::spawn_blocking(move || {
let writer = writer.lock().unwrap_or_else(|e| e.into_inner());
writer.delete_term(tantivy::Term::from_field_text(id_field, &id_val));
needs_commit.store(true, std::sync::atomic::Ordering::SeqCst);
Ok(())
})
}
pub async fn commit(&self) -> tantivy::Result<()> {
let writer = Arc::clone(&self.writer);
let needs_commit = Arc::clone(&self.needs_commit);
tokio::task::spawn_blocking(move || {
if needs_commit.swap(false, std::sync::atomic::Ordering::SeqCst) {
let mut writer = writer.lock().unwrap_or_else(|e| e.into_inner());
writer.commit()?;
}
Ok(())
})
.await
.unwrap_or_else(|_| {
Err(tantivy::TantivyError::SystemError(
"Commit task panicked".to_string(),
))
})
}
pub fn search(
&self,
query: &str,
namespace: Option<&str>,
) -> tantivy::Result<Vec<SearchResultTuple>> {
let searcher = self.reader.searcher();
let query_parser = tantivy::query::QueryParser::for_index(
&self.index,
vec![self.title_field, self.body_field],
);
let q = query_parser.parse_query(query)?;
let top_docs = searcher.search(
&q,
&tantivy::collector::TopDocs::with_limit(50).order_by_score(),
)?;
let mut results = Vec::new();
for (score, doc_address) in top_docs {
let retrieved_doc = searcher.doc::<tantivy::TantivyDocument>(doc_address)?;
let id = retrieved_doc
.get_first(self.id_field)
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let doc_type = retrieved_doc
.get_first(self.type_field)
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let title = retrieved_doc
.get_first(self.title_field)
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let body = retrieved_doc
.get_first(self.body_field)
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let doc_ns = retrieved_doc
.get_first(self.namespace_field)
.and_then(|v| v.as_str())
.unwrap_or("");
if let Some(ns) = namespace
&& doc_ns != ns
&& doc_ns != "global"
{
continue;
}
results.push((id, doc_type, title, body, score));
}
Ok(results)
}
pub fn index_snippet(&self, s: &Snippet) -> tokio::task::JoinHandle<tantivy::Result<()>> {
let writer = Arc::clone(&self.writer);
let id_field = self.id_field;
let id_val = s.name.clone();
let needs_commit = Arc::clone(&self.needs_commit);
let doc = doc!(
self.id_field => s.name.as_str(),
self.title_field => s.name.as_str(),
self.body_field => format!("{} {}\n{}", s.language, s.description, s.code),
self.type_field => "snippet",
self.namespace_field => "global"
);
tokio::task::spawn_blocking(move || {
let writer = writer.lock().unwrap_or_else(|e| e.into_inner());
writer.delete_term(tantivy::Term::from_field_text(id_field, &id_val));
writer.add_document(doc)?;
needs_commit.store(true, std::sync::atomic::Ordering::SeqCst);
Ok(())
})
}
pub fn index_adr(&self, a: &Adr) -> tokio::task::JoinHandle<tantivy::Result<()>> {
let writer = Arc::clone(&self.writer);
let id_field = self.id_field;
let id_val = a.id.clone();
let needs_commit = Arc::clone(&self.needs_commit);
let doc = doc!(
self.id_field => a.id.as_str(),
self.title_field => a.title.as_str(),
self.body_field => format!("{} {} {}", a.context, a.decision, a.consequence),
self.type_field => "adr",
self.namespace_field => "global"
);
tokio::task::spawn_blocking(move || {
let writer = writer.lock().unwrap_or_else(|e| e.into_inner());
writer.delete_term(tantivy::Term::from_field_text(id_field, &id_val));
writer.add_document(doc)?;
needs_commit.store(true, std::sync::atomic::Ordering::SeqCst);
Ok(())
})
}
pub fn add_entity_sync(&self, e: &Entity) {
if let Ok(writer) = self.writer.lock() {
let _ = writer.add_document(doc!(
self.id_field => e.name.as_str(),
self.title_field => e.name.as_str(),
self.body_field => e.observations.join(" "),
self.type_field => "entity",
self.namespace_field => e.namespace.as_str()
));
self.needs_commit
.store(true, std::sync::atomic::Ordering::SeqCst);
}
}
pub fn delete_all(&self) {
if let Ok(writer) = self.writer.lock() {
let _ = writer.delete_all_documents();
self.needs_commit
.store(true, std::sync::atomic::Ordering::SeqCst);
}
}
pub fn add_task_sync(&self, t: &Task) {
println!("add_task_sync called for task: {}", t.id);
if let Ok(writer) = self.writer.lock() {
let _res = writer.add_document(doc!(
self.id_field => t.id.as_str(),
self.title_field => t.title.as_str(),
self.body_field => t.description.as_str(),
self.type_field => "task",
self.namespace_field => "global"
));
println!("Writer add_document returned id/result");
self.needs_commit
.store(true, std::sync::atomic::Ordering::SeqCst);
println!("Needs_commit set to true in add_task_sync");
} else {
println!("Failed to acquire writer lock in add_task_sync");
}
}
pub fn add_snippet_sync(&self, s: &Snippet) {
if let Ok(writer) = self.writer.lock() {
let _ = writer.add_document(doc!(
self.id_field => s.name.as_str(),
self.title_field => s.name.as_str(),
self.body_field => format!("{} {}", s.language, s.description),
self.type_field => "snippet",
self.namespace_field => "global"
));
self.needs_commit
.store(true, std::sync::atomic::Ordering::SeqCst);
}
}
pub fn add_adr_sync(&self, a: &Adr) {
if let Ok(writer) = self.writer.lock() {
let _ = writer.add_document(doc!(
self.id_field => a.id.as_str(),
self.title_field => a.title.as_str(),
self.body_field => format!("{} {} {}", a.context, a.decision, a.consequence),
self.type_field => "adr",
self.namespace_field => "global"
));
self.needs_commit
.store(true, std::sync::atomic::Ordering::SeqCst);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::TempDir;
#[tokio::test]
async fn test_search_index_and_retrieve() {
let temp_dir = TempDir::new().unwrap();
let index = MemoryIndex::new(temp_dir.path()).unwrap();
let entity = Entity {
name: "TestEntity".to_string(),
entity_type: "Component".to_string(),
observations: vec!["This is a test observation".to_string()],
namespace: "global".to_string(),
git_branch: None,
};
let _ = index.index_entity(&entity).await.unwrap();
let task = Task {
id: "task-1".to_string(),
title: "Test Task".to_string(),
description: "Test task description".to_string(),
status: "open".to_string(),
created_at: 0,
updated_at: 0,
git_branch: None,
acceptance_criteria: vec![],
dependencies: vec![],
parent_id: None,
expires_at: None,
};
let _ = index.index_task(&task).await.unwrap();
let snippet = Snippet {
name: "test_snippet".to_string(),
code: "fn main() {}".to_string(),
language: "rust".to_string(),
description: "A test snippet".to_string(),
updated_at: 0,
embedding: None,
};
let _ = index.index_snippet(&snippet).await.unwrap();
let adr = Adr {
id: "adr-1".to_string(),
title: "Test ADR".to_string(),
context: "Test context".to_string(),
decision: "Test decision".to_string(),
consequence: "Test consequence".to_string(),
status: "accepted".to_string(),
supersedes: None,
timestamp: 0,
};
let _ = index.index_adr(&adr).await.unwrap();
index.commit().await.unwrap();
index.reader.reload().unwrap();
// Test search
let results = index.search("observation", None).unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].0, "TestEntity");
assert_eq!(results[0].1, "entity");
let results = index.search("task", None).unwrap();
assert!(results.iter().any(|r| r.0 == "task-1"));
let results = index.search("snippet", None).unwrap();
assert!(results.iter().any(|r| r.0 == "test_snippet"));
let results = index.search("decision", None).unwrap();
assert!(results.iter().any(|r| r.0 == "adr-1"));
}
#[test]
fn test_search_malformed_query() {
let temp_dir = TempDir::new().unwrap();
let index = MemoryIndex::new(temp_dir.path()).unwrap();
// Malformed lucene query (unclosed parenthesis)
let result = index.search("title: (unclosed", None);
assert!(result.is_err());
// Another malformed query (unclosed quote)
let result2 = index.search("title: \"unclosed", None);
assert!(result2.is_err());
}
}