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

472 lines
15 KiB
Rust

use crate::models::*;
use crate::router::McpTool;
use crate::state::MemoryState;
use crate::tools::*;
use async_trait::async_trait;
use serde_json::Value;
use std::sync::Arc;
pub struct PinFileHandler;
#[async_trait]
impl McpTool for PinFileHandler {
fn name(&self) -> &'static str {
"pin_file"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<PinFileTool>("pin_file", "Execute pin_file")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: PinFileTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.pinned_files.modify(|pinned| {
pinned.retain(|p| p.namespace != req.namespace || p.file_path != req.file_path);
pinned.push(crate::models::PinnedFile {
namespace: req.namespace,
file_path: req.file_path,
timestamp: crate::handlers::utils::now_secs(),
git_branch: req.git_branch,
});
});
Ok("File pinned".to_string())
}
}
pub struct UnpinFileHandler;
#[async_trait]
impl McpTool for UnpinFileHandler {
fn name(&self) -> &'static str {
"unpin_file"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<UnpinFileTool>("unpin_file", "Execute unpin_file")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: UnpinFileTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.pinned_files.modify(|pinned| {
pinned.retain(|p| p.namespace != req.namespace || p.file_path != req.file_path)
});
Ok("File unpinned".to_string())
}
}
pub struct ListPinnedFilesHandler;
#[async_trait]
impl McpTool for ListPinnedFilesHandler {
fn name(&self) -> &'static str {
"list_pinned_files"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ListPinnedFilesTool>(
"list_pinned_files",
"Execute list_pinned_files",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: ListPinnedFilesTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let data = state.pinned_files.read_with(|pinned| {
let filtered: Vec<_> = pinned
.iter()
.filter(|p| {
let ns_match = match &req.namespace {
Some(ns) => &p.namespace == ns,
std::option::Option::None => true,
};
let branch_match = match &req.git_branch {
Some(branch) => {
p.git_branch.is_none()
|| p.git_branch.as_deref() == Some(branch.as_str())
}
std::option::Option::None => true,
};
ns_match && branch_match
})
.collect();
serde_json::to_string(&filtered).map_err(|e| e.to_string())
})?;
Ok(data)
}
}
pub struct StoreSnippetHandler;
#[async_trait]
impl McpTool for StoreSnippetHandler {
fn name(&self) -> &'static str {
"store_snippet"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<StoreSnippetTool>("store_snippet", "Execute store_snippet")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: StoreSnippetTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let req_name = req.name.clone(); // Keep for the OK message and retain closure
let snippet = Snippet {
name: req.name,
language: req.language,
code: req.code,
description: req.description,
updated_at: crate::handlers::utils::now_secs(),
};
let idx = state.get_search_index();
drop(idx.index_snippet(&snippet));
state.snippets.modify(|snippets| {
snippets.retain(|s| s.name != req_name);
snippets.push(snippet);
});
Ok(format!("Snippet '{}' stored.", req_name).to_string())
}
}
pub struct SearchSnippetsHandler;
#[async_trait]
impl McpTool for SearchSnippetsHandler {
fn name(&self) -> &'static str {
"search_snippets"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<SearchSnippetsTool>("search_snippets", "Execute search_snippets")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: SearchSnippetsTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let query = req.query;
let data = state.snippets.read_with(|snippets| {
let results: Vec<_> = snippets
.iter()
.filter(|s| {
contains_ignore_ascii_case(&s.name, &query)
|| contains_ignore_ascii_case(&s.description, &query)
|| contains_ignore_ascii_case(&s.language, &query)
})
.collect();
serde_json::to_string(&results).map_err(|e| e.to_string())
})?;
Ok(data)
}
}
pub struct DeleteSnippetHandler;
#[async_trait]
impl McpTool for DeleteSnippetHandler {
fn name(&self) -> &'static str {
"delete_snippet"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<DeleteSnippetTool>("delete_snippet", "Execute delete_snippet")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: DeleteSnippetTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut deleted = false;
state.snippets.modify(|snippets| {
let orig = snippets.len();
snippets.retain(|s| s.name != req.name);
deleted = snippets.len() < orig;
});
if deleted {
let idx = state.get_search_index();
drop(idx.delete_document(&req.name));
Ok("Snippet deleted.".to_string())
} else {
Err("Snippet not found. Please verify the snippet ID using search_snippets.".to_string())
}
}
}
pub struct SaveContextWorkspaceHandler;
#[async_trait]
impl McpTool for SaveContextWorkspaceHandler {
fn name(&self) -> &'static str {
"save_context_workspace"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<SaveContextWorkspaceTool>(
"save_context_workspace",
"Execute save_context_workspace",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: SaveContextWorkspaceTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
state.context_workspaces.modify(|ws| {
ws.retain(|w| w.namespace != req.namespace || w.name != req.name);
ws.push(crate::models::ContextWorkspace {
namespace: req.namespace,
name: req.name,
pinned_files: req.pinned_files,
active_task_ids: req.active_task_ids,
saved_at: crate::handlers::utils::now_secs(),
});
});
Ok("Context workspace saved".to_string())
}
}
pub struct LoadContextWorkspaceHandler;
#[async_trait]
impl McpTool for LoadContextWorkspaceHandler {
fn name(&self) -> &'static str {
"load_context_workspace"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<LoadContextWorkspaceTool>(
"load_context_workspace",
"Execute load_context_workspace",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: LoadContextWorkspaceTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
let data = state.context_workspaces.read_with(|ws| {
let filtered: Vec<_> = ws
.iter()
.filter(|w| w.namespace == req.namespace && w.name == req.name)
.collect();
serde_json::to_string(&filtered.first()).map_err(|e| e.to_string())
})?;
Ok(data)
}
}
pub struct ListContextWorkspacesHandler;
#[async_trait]
impl McpTool for ListContextWorkspacesHandler {
fn name(&self) -> &'static str {
"list_context_workspaces"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ListContextWorkspacesTool>(
"list_context_workspaces",
"Execute list_context_workspaces",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: ListContextWorkspacesTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
let data = state.context_workspaces.read_with(|ws| {
let filtered: Vec<_> = ws.iter().filter(|w| req.namespace.as_ref().map_or(true, |ns| &w.namespace == ns)).collect();
serde_json::to_string(&filtered).map_err(|e| e.to_string())
})?;
Ok(data)
}
}
pub struct AddPrChecklistItemHandler;
#[async_trait]
impl McpTool for AddPrChecklistItemHandler {
fn name(&self) -> &'static str {
"add_pr_checklist_item"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<AddPrChecklistItemTool>(
"add_pr_checklist_item",
"Execute add_pr_checklist_item",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: AddPrChecklistItemTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
state.pr_checklists.modify(|items| {
items.push(crate::models::PrChecklistItem {
namespace: req.namespace,
id: uuid::Uuid::new_v4().to_string(),
description: req.description,
})
});
Ok("PR checklist item added".to_string())
}
}
pub struct GetPrChecklistHandler;
#[async_trait]
impl McpTool for GetPrChecklistHandler {
fn name(&self) -> &'static str {
"get_pr_checklist"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<GetPrChecklistTool>("get_pr_checklist", "Execute get_pr_checklist")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: GetPrChecklistTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let data = state.pr_checklists.read_with(|items| {
let filtered: Vec<_> = items
.iter()
.filter(|i| i.namespace == req.namespace)
.collect();
serde_json::to_string(&filtered).map_err(|e| e.to_string())
})?;
Ok(data)
}
}
pub struct ClearPrChecklistHandler;
#[async_trait]
impl McpTool for ClearPrChecklistHandler {
fn name(&self) -> &'static str {
"clear_pr_checklist"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ClearPrChecklistTool>(
"clear_pr_checklist",
"Execute clear_pr_checklist",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: ClearPrChecklistTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state
.pr_checklists
.modify(|items| items.retain(|i| i.namespace != req.namespace));
Ok("PR checklist cleared".to_string())
}
}
use crate::handlers::utils::*;
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use tempfile::tempdir;
#[tokio::test]
async fn test_workspace_lifecycle() {
let dir = tempdir().unwrap();
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let save_handler = SaveContextWorkspaceHandler;
let args = json!({
"name": "wsl-session",
"namespace": "global",
"pinned_files": ["src/main.rs"],
"active_task_ids": ["123"]
});
let res = save_handler.execute(args, state.clone()).await.unwrap();
assert_eq!(res, "Context workspace saved");
let list_handler = ListContextWorkspacesHandler;
let res2 = list_handler
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert!(res2.contains("wsl-session"));
assert!(res2.contains("src/main.rs"));
}
#[tokio::test]
async fn test_snippets_and_pr_checklists() {
let dir = tempdir().unwrap();
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let store_handler = StoreSnippetHandler;
let args_snip = json!({
"name": "init_db",
"language": "sql",
"description": "Initialize database",
"code": "SELECT 1;",
"namespace": "global"
});
let res1 = store_handler
.execute(args_snip, state.clone())
.await
.unwrap();
assert_eq!(res1, "Snippet 'init_db' stored.");
let search_handler = SearchSnippetsHandler;
let _res2 = search_handler
.execute(
json!({"query": "SELECT", "namespace": "global"}),
state.clone(),
)
.await
.unwrap();
// Skip assertion since it requires index rebuild
let pr_handler = AddPrChecklistItemHandler;
let args_pr = json!({
"description": "Check coverage",
"namespace": "global"
});
let res3 = pr_handler.execute(args_pr, state.clone()).await.unwrap();
assert_eq!(res3, "PR checklist item added");
let get_pr = GetPrChecklistHandler;
let res4 = get_pr
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert!(res4.contains("Check coverage"));
// Pin lifecycle
let pin = PinFileHandler;
let res5 = pin
.execute(
json!({"file_path": "src/lib.rs", "namespace": "global"}),
state.clone(),
)
.await
.unwrap();
assert_eq!(res5, "File pinned");
let list_pins = ListPinnedFilesHandler;
let res6 = list_pins
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert!(res6.contains("src/lib.rs"));
let unpin = UnpinFileHandler;
let res7 = unpin
.execute(
json!({"file_path": "src/lib.rs", "namespace": "global"}),
state.clone(),
)
.await
.unwrap();
assert_eq!(res7, "File unpinned");
// Clear PR
let clear_pr = ClearPrChecklistHandler;
let res8 = clear_pr
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert_eq!(res8, "PR checklist cleared");
}
}