262 lines
8.8 KiB
Rust
262 lines
8.8 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 struct MemoryIndex {
|
|
index: Index,
|
|
reader: IndexReader,
|
|
writer: Arc<Mutex<IndexWriter>>,
|
|
|
|
// 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).unwrap();
|
|
let index = Index::open_in_dir(&index_dir)
|
|
.unwrap_or_else(|_| Index::create_in_dir(&index_dir, schema.clone()).unwrap());
|
|
|
|
let writer = index.writer(50_000_000)?;
|
|
let reader = index
|
|
.reader_builder()
|
|
.reload_policy(ReloadPolicy::OnCommitWithDelay)
|
|
.try_into()?;
|
|
|
|
Ok(Self {
|
|
index,
|
|
reader,
|
|
writer: Arc::new(Mutex::new(writer)),
|
|
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 doc = doc!(
|
|
self.id_field => e.name.clone(),
|
|
self.title_field => e.name.clone(),
|
|
self.body_field => e.observations.join(" "),
|
|
self.type_field => "entity",
|
|
self.namespace_field => e.namespace.clone()
|
|
);
|
|
tokio::task::spawn_blocking(move || {
|
|
let writer = writer.lock().unwrap();
|
|
writer.add_document(doc)?;
|
|
Ok(())
|
|
})
|
|
}
|
|
|
|
pub fn index_task(&self, t: &Task) -> tokio::task::JoinHandle<tantivy::Result<()>> {
|
|
let writer = Arc::clone(&self.writer);
|
|
let doc = doc!(
|
|
self.id_field => t.id.clone(),
|
|
self.title_field => t.title.clone(),
|
|
self.body_field => t.description.clone(),
|
|
self.type_field => "task",
|
|
self.namespace_field => "global"
|
|
);
|
|
tokio::task::spawn_blocking(move || {
|
|
let writer = writer.lock().unwrap();
|
|
writer.add_document(doc)?;
|
|
Ok(())
|
|
})
|
|
}
|
|
|
|
pub fn commit(&self) -> tantivy::Result<()> {
|
|
let mut writer = self.writer.lock().unwrap();
|
|
writer.commit()?;
|
|
Ok(())
|
|
}
|
|
|
|
pub fn search(
|
|
&self,
|
|
query: &str,
|
|
namespace: Option<&str>,
|
|
) -> tantivy::Result<Vec<(String, String, String, String, f32)>> {
|
|
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 doc = doc!(
|
|
self.id_field => s.name.clone(),
|
|
self.title_field => s.name.clone(),
|
|
self.body_field => format!("{} {}", s.language, s.description),
|
|
self.type_field => "snippet",
|
|
self.namespace_field => "global"
|
|
);
|
|
tokio::task::spawn_blocking(move || {
|
|
let writer = writer.lock().unwrap();
|
|
writer.add_document(doc)?;
|
|
Ok(())
|
|
})
|
|
}
|
|
|
|
pub fn index_adr(&self, a: &Adr) -> tokio::task::JoinHandle<tantivy::Result<()>> {
|
|
let writer = Arc::clone(&self.writer);
|
|
let doc = doc!(
|
|
self.id_field => a.id.clone(),
|
|
self.title_field => a.title.clone(),
|
|
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();
|
|
writer.add_document(doc)?;
|
|
Ok(())
|
|
})
|
|
}
|
|
}
|
|
|
|
#[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,
|
|
};
|
|
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,
|
|
};
|
|
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,
|
|
};
|
|
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(),
|
|
timestamp: 0,
|
|
};
|
|
index.index_adr(&adr).await.unwrap();
|
|
|
|
index.commit().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());
|
|
}
|
|
}
|