feat: implement branch-aware context switching and integrate tantivy fuzzy search index

This commit is contained in:
Riz Ashraf committed 2026-09-08 11:44:25 +01:00
1 parent e4ff476b6d
commit 3239e32866
6 files changed
+797 -2

No files matched your search

+1
View File
@@ -16,6 +16,7 @@ rust-mcp-sdk = "2.0.0"
serde = { version = "1.0.229", features = ["derive"] }
serde_json = "1.0.151"
strsim = "0.11.1"
tantivy = "0.26.1"
tokio = { version = "1.53.1", features = ["full"] }
tokio-util = { version = "0.7.19", features = ["io"] }
uuid = { version = "1.26.0", features = ["v4"] }
+10
View File
@@ -175,6 +175,7 @@ impl ServerHandler for MemoryHandler {
entity_type: full_e.entity_type.clone(),
observations: vec![],
namespace: full_e.namespace.clone(),
git_branch: None,
});
e.observations.extend(o.contents);
g.entities.insert(o.entity_name, e);
@@ -463,6 +464,7 @@ impl ServerHandler for MemoryHandler {
description: req.description,
created_at: now,
updated_at: now,
git_branch: req.git_branch,
});
});
Ok(ServerResult::from(CallToolResult::text_content(vec![
@@ -496,8 +498,12 @@ impl ServerHandler for MemoryHandler {
}
}
"list_active_tasks" => {
let req: ListActiveTasksTool = parse_args(args)?;
let mut tasks = self.state.tasks.read();
tasks.retain(|t| t.status != "done");
if let Some(branch) = req.git_branch {
tasks.retain(|t| t.git_branch.is_none() || t.git_branch.as_deref() == Some(branch.as_str()));
}
let data = serde_json::to_string(&tasks).unwrap_or_default();
Ok(ServerResult::from(CallToolResult::text_content(vec![
data.into(),
@@ -710,6 +716,7 @@ impl ServerHandler for MemoryHandler {
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs(),
git_branch: req.git_branch,
});
});
Ok(ServerResult::from(CallToolResult::text_content(vec![
@@ -732,6 +739,9 @@ impl ServerHandler for MemoryHandler {
if let Some(ns) = req.namespace {
pinned.retain(|p| p.namespace == ns);
}
if let Some(branch) = req.git_branch {
pinned.retain(|p| p.git_branch.is_none() || p.git_branch.as_deref() == Some(branch.as_str()));
}
let data = serde_json::to_string(&pinned).unwrap_or_default();
Ok(ServerResult::from(CallToolResult::text_content(vec![
data.into(),
+4
View File
@@ -26,6 +26,8 @@ pub struct Entity {
pub observations: Vec<String>,
#[serde(default = "default_namespace")]
pub namespace: String,
#[serde(default)]
pub git_branch: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Hash)]
pub struct Relation {
@@ -51,6 +53,7 @@ pub struct Task {
pub description: String,
pub created_at: u64,
pub updated_at: u64,
pub git_branch: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct Snippet {
@@ -88,6 +91,7 @@ pub struct PinnedFile {
pub namespace: String,
pub file_path: String,
pub timestamp: u64,
pub git_branch: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct SessionSummary {
+91
View File
@@ -0,0 +1,91 @@
use tantivy::schema::*;
use tantivy::{doc, Index, IndexWriter, IndexReader, ReloadPolicy};
use std::sync::Mutex;
use crate::models::{Entity, Task, StickyNote, Adr, Snippet};
pub struct MemoryIndex {
index: Index,
reader: IndexReader,
writer: 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() -> 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 = Index::create_in_ram(schema.clone());
let writer = index.writer(50_000_000)?;
let reader = index
.reader_builder()
.reload_policy(ReloadPolicy::OnCommitWithDelay)
.try_into()?;
Ok(Self {
index,
reader,
writer: Mutex::new(writer),
id_field,
title_field,
body_field,
type_field,
namespace_field,
})
}
pub fn index_entity(&self, e: &Entity) -> tantivy::Result<()> {
let mut writer = self.writer.lock().unwrap();
writer.add_document(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()
))?;
writer.commit()?;
Ok(())
}
pub fn index_task(&self, t: &Task) -> tantivy::Result<()> {
let mut writer = self.writer.lock().unwrap();
writer.add_document(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"
))?;
writer.commit()?;
Ok(())
}
pub fn search(&self, query: &str, namespace: Option<&str>) -> tantivy::Result<Vec<String>> {
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))?;
let mut results = Vec::new();
for (_score, doc_address) in top_docs {
let retrieved_doc = searcher.doc::<tantivy::TantivyDocument>(doc_address)?;
if let Some(val) = retrieved_doc.get_first(self.title_field) {
if let Some(t) = val.as_str() {
results.push(t.to_string());
}
}
}
Ok(results)
}
}
+4 -1
View File
@@ -95,6 +95,7 @@ pub struct CondenseEntityTool {
pub struct AddTaskTool {
pub title: String,
pub description: String,
pub git_branch: Option<String>,
}
#[macros::mcp_tool(name = "update_task_status", description = "Update task status")]
#[derive(Debug, Deserialize, Serialize, macros::JsonSchema)]
@@ -104,7 +105,7 @@ pub struct UpdateTaskStatusTool {
}
#[macros::mcp_tool(name = "list_active_tasks", description = "List active tasks")]
#[derive(Debug, Deserialize, Serialize, macros::JsonSchema)]
pub struct ListActiveTasksTool {}
pub struct ListActiveTasksTool { pub git_branch: Option<String>, }
#[macros::mcp_tool(name = "store_snippet", description = "Store code snippet")]
#[derive(Debug, Deserialize, Serialize, macros::JsonSchema)]
pub struct StoreSnippetTool {
@@ -186,6 +187,7 @@ pub struct SearchErrorFixesTool {
pub struct PinFileTool {
pub namespace: String,
pub file_path: String,
pub git_branch: Option<String>,
}
#[macros::mcp_tool(name = "unpin_file", description = "Unpin a file")]
#[derive(Debug, Deserialize, Serialize, macros::JsonSchema)]
@@ -197,6 +199,7 @@ pub struct UnpinFileTool {
#[derive(Debug, Deserialize, Serialize, macros::JsonSchema)]
pub struct ListPinnedFilesTool {
pub namespace: Option<String>,
pub git_branch: Option<String>,
}
#[macros::mcp_tool(name = "add_session_summary", description = "Add a session summary")]
#[derive(Debug, Deserialize, Serialize, macros::JsonSchema)]