From 8afbf97b1142112458b298534317788d847320fc Mon Sep 17 00:00:00 2001 From: Riz Ashraf Date: Mon, 21 Sep 2026 11:34:21 +0100 Subject: [PATCH] refactor: apply zero-unwrap policy and optimize locks in store.rs and handlers --- dedup.py | 14 + dedup2.py | 14 + linux-nvim/build.rs | 6 +- linux-nvim/tests/integration_test.rs | 2 + mcp-stdio/src/lib.rs | 8 +- nvim-core/src/lib.rs | 102 +- server/build.rs | 2 +- server/src/handlers.rs | 3107 +------------------------- server/src/handlers_v2/env.rs | 163 ++ server/src/handlers_v2/graph.rs | 544 +++++ server/src/handlers_v2/meta.rs | 492 ++++ server/src/handlers_v2/mod.rs | 7 + server/src/handlers_v2/notes.rs | 273 +++ server/src/handlers_v2/tasks.rs | 456 ++++ server/src/handlers_v2/utils.rs | 9 + server/src/handlers_v2/workspaces.rs | 348 +++ server/src/main.rs | 144 +- server/src/refactor.py | 173 ++ server/src/refactor.rs | 11 + server/src/refactor_router.py | 82 + server/src/refactor_wire.py | 106 + server/src/router.rs | 16 + server/src/search.rs | 16 +- server/src/state.rs | 1 - server/src/store.rs | 49 +- server/src/tools.rs | 1 + server/tests/parity_test.rs | 36 +- stub/build.rs | 6 +- stub/src/logger.rs | 35 +- stub/src/main.rs | 41 +- stub/tests/e2e.rs | 62 +- win-nvim/build.rs | 6 +- win-nvim/tests/integration_test.rs | 2 +- 33 files changed, 3127 insertions(+), 3207 deletions(-) create mode 100644 dedup.py create mode 100644 dedup2.py create mode 100644 server/src/handlers_v2/env.rs create mode 100644 server/src/handlers_v2/graph.rs create mode 100644 server/src/handlers_v2/meta.rs create mode 100644 server/src/handlers_v2/mod.rs create mode 100644 server/src/handlers_v2/notes.rs create mode 100644 server/src/handlers_v2/tasks.rs create mode 100644 server/src/handlers_v2/utils.rs create mode 100644 server/src/handlers_v2/workspaces.rs create mode 100644 server/src/refactor.py create mode 100644 server/src/refactor.rs create mode 100644 server/src/refactor_router.py create mode 100644 server/src/refactor_wire.py create mode 100644 server/src/router.rs diff --git a/dedup.py b/dedup.py new file mode 100644 index 0000000..1f72adf --- /dev/null +++ b/dedup.py @@ -0,0 +1,14 @@ +import re + +with open("server/src/handlers.rs", "r", encoding="utf-8") as f: + text = f.read() + +pattern = r'"(list_milestones|list_pinned_files|read_handoff_memos)" => \{\s*let req = parse_tool!\(args, id, ([a-zA-Z]+)\);\s*let mut [a-zA-Z_]+ = self\.state\.([a-zA-Z_]+)\.read\(\);\s*if let Some\(ns\) = req\.namespace \{\s*[a-zA-Z_]+\.retain\(\|.\| \w+\.namespace == ns\);\s*\}\s*let data = serde_json::to_string\(&[a-zA-Z_]+\)\.unwrap_or_default\(\);\s*Ok\(data\.to_string\(\)\)\s*\}' + +def repl(m): + return f'"{m.group(1)}" => handle_list_with_namespace!(self, {m.group(3)}, {m.group(2)}, args, id),' + +new_text = re.sub(pattern, repl, text) + +with open("server/src/handlers.rs", "w", encoding="utf-8") as f: + f.write(new_text) diff --git a/dedup2.py b/dedup2.py new file mode 100644 index 0000000..7128092 --- /dev/null +++ b/dedup2.py @@ -0,0 +1,14 @@ +import re + +with open("server/src/handlers.rs", "r", encoding="utf-8") as f: + text = f.read() + +pattern = r'"([a-zA-Z_]+)" => \{\s*(?:let req = parse_tool!\(args, id, ([a-zA-Z]+)\);\s*)?let data = serde_json::to_string\(&self\.state\.([a-zA-Z_]+)\.read\(\)\)\s*\.unwrap_or_else\(\|_\| "\[\]"\.to_string\(\)\);\s*Ok\(data\.to_string\(\)\)\s*\}' + +def repl(m): + return f'"{m.group(1)}" => {{\n let data = serde_json::to_string(&self.state.{m.group(3)}.read()).unwrap_or_else(|_| "[]".to_string());\n Ok(data)\n}},' + +new_text = re.sub(pattern, repl, text) + +with open("server/src/handlers.rs", "w", encoding="utf-8") as f: + f.write(new_text) diff --git a/linux-nvim/build.rs b/linux-nvim/build.rs index dc9a51c..c266efa 100644 --- a/linux-nvim/build.rs +++ b/linux-nvim/build.rs @@ -15,10 +15,6 @@ fn main() { .and_then(|out| String::from_utf8(out.stdout).ok()) .unwrap_or_else(|| "unknown".to_string()); - let version = format!( - "{} ({})", - git_date.trim(), - git_hash.trim() - ); + let version = format!("{} ({})", git_date.trim(), git_hash.trim()); println!("cargo:rustc-env=APP_VERSION={}", version); } diff --git a/linux-nvim/tests/integration_test.rs b/linux-nvim/tests/integration_test.rs index 187dc7f..9efec15 100644 --- a/linux-nvim/tests/integration_test.rs +++ b/linux-nvim/tests/integration_test.rs @@ -1,3 +1,5 @@ +#![cfg(unix)] + use serde_json::{json, Value}; use std::io::{BufRead, BufReader, Read, Write}; use std::process::{Command, Stdio}; diff --git a/mcp-stdio/src/lib.rs b/mcp-stdio/src/lib.rs index 3e9469f..8697a29 100644 --- a/mcp-stdio/src/lib.rs +++ b/mcp-stdio/src/lib.rs @@ -21,21 +21,21 @@ pub async fn read_mcp_message( if line.is_empty() { break; } - + let lower_line = line.to_lowercase(); if let Some(len_str) = lower_line.strip_prefix("content-length:") { length = len_str.trim().parse().unwrap_or(0); } } - + if length == 0 { return None; } - + let mut buffer = vec![0; length]; if stdin.read_exact(&mut buffer).await.is_err() { return None; } - + String::from_utf8(buffer).ok() } diff --git a/nvim-core/src/lib.rs b/nvim-core/src/lib.rs index f139ecf..00be919 100644 --- a/nvim-core/src/lib.rs +++ b/nvim-core/src/lib.rs @@ -20,8 +20,6 @@ pub struct JsonRpcResponse { pub error: Option, } - - pub async fn send_response(response: JsonRpcResponse) { let msg = serde_json::to_string(&response).unwrap(); tracing::info!( @@ -111,11 +109,10 @@ async fn get_socket_path() -> Result { } Err("Could not find Neovim socket".to_string()) } -use std::sync::LazyLock; -use std::sync::Arc; -use tokio::sync::{mpsc, oneshot}; use std::collections::HashMap; - +use std::sync::Arc; +use std::sync::LazyLock; +use tokio::sync::{mpsc, oneshot}; pub struct NvimRequest { pub msgid_str: String, @@ -123,7 +120,8 @@ pub struct NvimRequest { pub reply: oneshot::Sender>, } -static NVIM_CONN: LazyLock>>>> = LazyLock::new(|| Arc::new(std::sync::Mutex::new(None))); +static NVIM_CONN: LazyLock>>>> = + LazyLock::new(|| Arc::new(std::sync::Mutex::new(None))); async fn get_nvim_connection() -> Result, String> { { @@ -141,18 +139,23 @@ async fn get_nvim_connection() -> Result, String> { #[cfg(windows)] let stream = { use tokio::net::windows::named_pipe::ClientOptions; - ClientOptions::new().open(&socket_path).map_err(|e| e.to_string())? + ClientOptions::new() + .open(&socket_path) + .map_err(|e| e.to_string())? }; #[cfg(unix)] let stream = { use tokio::net::UnixStream; - UnixStream::connect(socket_path).await.map_err(|e| e.to_string())? + UnixStream::connect(socket_path) + .await + .map_err(|e| e.to_string())? }; let (mut read_half, mut write_half) = tokio::io::split(stream); let (tx, mut rx) = mpsc::channel::(32); - type PendingRequestsMap = Arc>>>>; + type PendingRequestsMap = + Arc>>>>; let pending_requests: PendingRequestsMap = Arc::new(std::sync::Mutex::new(HashMap::new())); // Write task @@ -164,9 +167,12 @@ async fn get_nvim_connection() -> Result, String> { let _ = req.reply.send(Err(e.to_string())); continue; } - - pending_clone.lock().unwrap().insert(req.msgid_str.clone(), req.reply); - + + pending_clone + .lock() + .unwrap() + .insert(req.msgid_str.clone(), req.reply); + if write_half.write_all(&buf).await.is_err() { tracing::error!("Failed to write to Neovim socket"); break; @@ -191,8 +197,10 @@ async fn get_nvim_connection() -> Result, String> { if arr.len() >= 4 && arr[0] == rmpv::Value::Integer(1.into()) { let msgid = &arr[1]; let msgid_str = format!("{:?}", msgid); - - if let Some(reply_sender) = pending_clone2.lock().unwrap().remove(&msgid_str) { + + if let Some(reply_sender) = + pending_clone2.lock().unwrap().remove(&msgid_str) + { let _ = reply_sender.send(Ok(val)); } } @@ -204,12 +212,16 @@ async fn get_nvim_connection() -> Result, String> { } continue; } - Err(rmpv::decode::Error::InvalidMarkerRead(e)) if e.kind() == std::io::ErrorKind::UnexpectedEof => { + Err(rmpv::decode::Error::InvalidMarkerRead(e)) + if e.kind() == std::io::ErrorKind::UnexpectedEof => + { resp_buf.drain(..offset); offset = 0; - + let read_future = read_half.read(&mut chunk); - match tokio::time::timeout(tokio::time::Duration::from_secs(60), read_future).await { + match tokio::time::timeout(tokio::time::Duration::from_secs(60), read_future) + .await + { Ok(Ok(n)) if n > 0 => { resp_buf.extend_from_slice(&chunk[..n]); } @@ -225,7 +237,7 @@ async fn get_nvim_connection() -> Result, String> { } } } - + // Cleanup pending requests on disconnect let mut pending = pending_clone2.lock().unwrap(); for (_, sender) in pending.drain() { @@ -242,7 +254,10 @@ async fn get_nvim_connection() -> Result, String> { if Arc::strong_count(&pending_clone3) <= 1 { break; // Socket closed and other tasks finished, no need to keep cleaning up } - pending_clone3.lock().unwrap().retain(|_, sender| !sender.is_closed()); + pending_clone3 + .lock() + .unwrap() + .retain(|_, sender| !sender.is_closed()); } }); @@ -267,17 +282,19 @@ async fn call_nvim(req: rmpv::Value) -> Result { } else { rmpv::Value::Nil }; - + let msgid_str = format!("{:?}", msgid); let tx = get_nvim_connection().await?; let (reply_tx, reply_rx) = oneshot::channel(); - + tx.send(NvimRequest { msgid_str, req, reply: reply_tx, - }).await.map_err(|_| "Failed to send request to Neovim connection manager")?; - + }) + .await + .map_err(|_| "Failed to send request to Neovim connection manager")?; + match tokio::time::timeout(tokio::time::Duration::from_secs(10), reply_rx).await { Ok(Ok(res)) => res, Ok(Err(_)) => Err("Response channel dropped".to_string()), @@ -544,9 +561,13 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { Ok(m) => { tracing::info!("Received message method: {}", m.method); m - }, + } Err(e) => { - tracing::error!("Failed to parse JSON-RPC request from JSONL: {}. Payload: {}", e, raw_msg); + tracing::error!( + "Failed to parse JSON-RPC request from JSONL: {}. Payload: {}", + e, + raw_msg + ); continue; } }; @@ -559,13 +580,17 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { let id_clone = id.clone(); let start_time = std::time::Instant::now(); let method_clone = if msg.method == "tools/call" { - let tool_name = msg.params.as_ref().and_then(|p| p.get("name")).and_then(|n| n.as_str()).unwrap_or("unknown"); + let tool_name = msg + .params + .as_ref() + .and_then(|p| p.get("name")) + .and_then(|n| n.as_str()) + .unwrap_or("unknown"); format!("ToolCall[{}]", tool_name) } else { msg.method.clone() }; - match msg.method.as_str() { "initialize" => { @@ -774,7 +799,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { Err(e) => send_error(id, -32603, &e).await, } } - + "nvim_open_file" => { let json_str = serde_json::to_string(args).unwrap().replace("\\", "\\\\").replace("'", "\\'"); let code = format!(" @@ -899,7 +924,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { local buf = args.buf_id or vim.api.nvim_get_current_buf() local group = args.group or 'IncSearch' local ns = vim.api.nvim_create_namespace('antigravity_highlight') - + if args.clear_only then vim.api.nvim_buf_clear_namespace(buf, ns, 0, -1) return 'Cleared highlights' @@ -909,7 +934,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { for i = args.start_line - 1, args.end_line - 1 do pcall(vim.api.nvim_buf_add_highlight, buf, ns, group, i, 0, -1) end - + local duration = args.duration_ms or 5000 if duration > 0 then vim.defer_fn(function() @@ -984,13 +1009,16 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { } } let elapsed = start_time.elapsed(); - tracing::info!("<<< [Nvim] {} (id: {}) completed in {:?}", method_clone, id_clone, elapsed); + tracing::info!( + "<<< [Nvim] {} (id: {}) completed in {:?}", + method_clone, + id_clone, + elapsed + ); }); } } - - fn init_logging(app_name: &str) -> tracing_appender::non_blocking::WorkerGuard { let log_dir = dirs::home_dir() .unwrap_or_default() @@ -1036,11 +1064,10 @@ mod tests { #[test] fn test_rmpv_to_json_map() { - let mut map = vec![]; - map.push(( + let map = vec![( rmpv::Value::String("key1".into()), rmpv::Value::Integer(100.into()), - )); + )]; let rmp_map = rmpv::Value::Map(map); let json_map = rmpv_to_json(&rmp_map); @@ -1074,4 +1101,3 @@ mod tests { assert!(req.is_none()); } } - diff --git a/server/build.rs b/server/build.rs index 848267e..8dad37f 100644 --- a/server/build.rs +++ b/server/build.rs @@ -7,7 +7,7 @@ fn main() { .ok() .and_then(|out| String::from_utf8(out.stdout).ok()) .unwrap_or_else(|| "unknown".to_string()); - + git_hash = git_hash.trim().to_string(); let is_dirty = Command::new("git") diff --git a/server/src/handlers.rs b/server/src/handlers.rs index 7d5de00..2818833 100644 --- a/server/src/handlers.rs +++ b/server/src/handlers.rs @@ -1,74 +1,101 @@ -use crate::models::*; +use crate::router::McpTool; use crate::state::MemoryState; -use crate::tools::*; - -fn contains_ignore_ascii_case(haystack: &str, needle: &str) -> bool { - if needle.is_empty() { return true; } - haystack.as_bytes().windows(needle.len()).any(|w| w.eq_ignore_ascii_case(needle.as_bytes())) -} - -macro_rules! parse_tool { - ($args:expr, $id:expr, $type:ty) => { - match parse_args::<$type>($args) { - Ok(r) => r, - Err(e) => { - let response = Some(crate::mcp::success( - $id.clone(), - serde_json::json!({"isError": true, "content": [{"type": "text", "text": format!("Invalid args: {}", e)}] }), - )); - tracing::trace!("Returning response from handle_request: {:?}", response); - return response; - } - } - }; -} - -macro_rules! handle_list_with_namespace { - ($self:expr, $store:ident, $tool_type:ty, $args:expr, $id:expr) => {{ - let req = parse_tool!($args, $id, $tool_type); - let mut items = $self.state.$store.read(); - if let Some(ns) = req.namespace { - items.retain(|i| i.namespace == ns); - } - let data = serde_json::to_string(&items).unwrap_or_default(); - Ok(data.to_string()) - }}; -} - -use serde::de::DeserializeOwned; -use std::collections::HashSet; use std::sync::Arc; -use std::time::{SystemTime, UNIX_EPOCH}; - -fn parse_args(args: serde_json::Value) -> Result { - serde_json::from_value(args).map_err(|e| format!("Invalid args: {}", e)) -} pub struct MemoryHandler { pub state: Arc, + pub tools: std::collections::HashMap>, } impl MemoryHandler { + pub fn new(state: Arc) -> Self { + let mut tools: std::collections::HashMap> = + std::collections::HashMap::new(); + + macro_rules! register { + ($module:ident::$handler:ident) => { + let h = crate::handlers_v2::$module::$handler; + tools.insert(h.name().to_string(), Box::new(h)); + }; + } + + register!(graph::QueryGraphPathHandler); + register!(graph::CreateEntitiesHandler); + register!(graph::CreateRelationsHandler); + register!(graph::AddObservationsHandler); + register!(graph::DeleteEntitiesHandler); + register!(graph::DeleteObservationsHandler); + register!(graph::DeleteRelationsHandler); + register!(graph::ReadGraphHandler); + register!(graph::SearchNodesHandler); + register!(graph::OpenNodesHandler); + register!(graph::VisualizeGraphHandler); + register!(graph::CondenseEntityHandler); + register!(graph::MergeEntitiesHandler); + register!(graph::FindOrphansHandler); + + register!(tasks::AddTaskHandler); + register!(tasks::DeleteTaskHandler); + register!(tasks::UpdateTaskStatusHandler); + register!(tasks::ListActiveTasksHandler); + register!(tasks::SetAcceptanceCriteriaHandler); + register!(tasks::VerifyAcceptanceCriteriaHandler); + register!(tasks::AddMilestoneHandler); + register!(tasks::UpdateMilestoneHandler); + register!(tasks::ListMilestonesHandler); + + register!(notes::AddStickyNoteHandler); + register!(notes::ReadStickyNotesHandler); + register!(notes::DeleteStickyNoteHandler); + register!(notes::ClearStickyNotesHandler); + register!(notes::LeaveHandoffMemoHandler); + register!(notes::ReadHandoffMemosHandler); + register!(notes::ClearHandoffMemosHandler); + register!(notes::AddSessionSummaryHandler); + register!(notes::GenerateStandupReportHandler); + + register!(meta::LogDecisionHandler); + register!(meta::QueryDecisionsHandler); + register!(meta::LogErrorFixHandler); + register!(meta::SearchErrorFixesHandler); + register!(meta::LogCodeChangeHandler); + register!(meta::QueryRecentChangesHandler); + register!(meta::LearnPreferenceHandler); + register!(meta::ReadPreferencesHandler); + register!(meta::LogTechDebtHandler); + register!(meta::ResolveTechDebtHandler); + register!(meta::ListTechDebtHandler); + register!(meta::OmniSearchHandler); + register!(meta::GetProjectHealthHandler); + + register!(env::UpdateEnvFingerprintHandler); + register!(env::ReadEnvFingerprintHandler); + register!(env::LogEnvRequirementHandler); + register!(env::RegisterEnvironmentHandler); + register!(env::GetEnvironmentDetailsHandler); + + register!(workspaces::PinFileHandler); + register!(workspaces::UnpinFileHandler); + register!(workspaces::ListPinnedFilesHandler); + register!(workspaces::StoreSnippetHandler); + register!(workspaces::SearchSnippetsHandler); + register!(workspaces::DeleteSnippetHandler); + register!(workspaces::SaveContextWorkspaceHandler); + register!(workspaces::LoadContextWorkspaceHandler); + register!(workspaces::ListContextWorkspacesHandler); + register!(workspaces::AddPrChecklistItemHandler); + register!(workspaces::GetPrChecklistHandler); + register!(workspaces::ClearPrChecklistHandler); + + Self { state, tools } + } + pub async fn handle_request(&self, req: serde_json::Value) -> Option { - let start_time = std::time::Instant::now(); let id = req.get("id").cloned().unwrap_or(serde_json::Value::Null); let id_clone = id.clone(); let method = req.get("method").and_then(|m| m.as_str()).unwrap_or(""); - - let tool_name = if method == "tools/call" { - req.get("params").and_then(|p| p.get("name")).and_then(|n| n.as_str()).unwrap_or("unknown") - } else { - "" - }; - - if method == "tools/call" { - tracing::info!(">>> [Server] Handling MCP tool call: {} (id: {})", tool_name, id); - } else { - tracing::debug!(">>> [Server] Handling MCP request method: {}", method); - } - tracing::trace!(">>> [Server] Full MCP Request payload: {}", req.to_string()); - let response = match method { + match method { "server/discover" => { let payload = serde_json::json!({ "resultType": "complete", @@ -85,10 +112,6 @@ impl MemoryHandler { } } }); - tracing::debug!( - "<<< [Server] Replying to server/discover with payload: {}", - payload.to_string() - ); Some(crate::mcp::success(id, payload)) } "initialize" => { @@ -101,269 +124,21 @@ impl MemoryHandler { "gemini-mcp-memory", "3.0.0", )); - tracing::debug!("<<< [Server] Replying to initialize with rmcp payload"); Some(crate::mcp::success( id, serde_json::to_value(&init).unwrap(), )) } - "notifications/initialized" => None, "tools/list" => { - let tools = vec![ - crate::mcp::tool_def::( - "query_graph_path", - "Traverse the knowledge graph to find a path between two entities.", - ), - crate::mcp::tool_def::( - "create_entities", - "Create new entities in the knowledge graph.", - ), - crate::mcp::tool_def::( - "create_relations", - "Create new relations between entities in the knowledge graph.", - ), - crate::mcp::tool_def::( - "add_observations", - "Add new observations to existing entities in the knowledge graph.", - ), - crate::mcp::tool_def::( - "delete_entities", - "Delete entities from the knowledge graph.", - ), - crate::mcp::tool_def::( - "delete_observations", - "Delete observations from existing entities.", - ), - crate::mcp::tool_def::( - "delete_relations", - "Delete relations between entities.", - ), - crate::mcp::tool_def::( - "read_graph", - "Read the entire knowledge graph.", - ), - crate::mcp::tool_def::( - "search_nodes", - "Search for entities in the knowledge graph by name or type.", - ), - crate::mcp::tool_def::( - "open_nodes", - "Open and retrieve full details of specific nodes in the knowledge graph.", - ), - crate::mcp::tool_def::( - "log_code_change", - "Log a significant code change or refactor in the memory system.", - ), - crate::mcp::tool_def::( - "query_recent_changes", - "Query recently logged code changes.", - ), - crate::mcp::tool_def::( - "visualize_graph", - "Generate a visual representation of the knowledge graph.", - ), - crate::mcp::tool_def::( - "add_sticky_note", - "Add a sticky note for unstructured thoughts or reminders.", - ), - crate::mcp::tool_def::( - "read_sticky_notes", - "Read all active sticky notes.", - ), - crate::mcp::tool_def::( - "delete_sticky_note", - "Delete a specific sticky note by its 1-indexed position.", - ), - crate::mcp::tool_def::( - "clear_sticky_notes", - "Clear all active sticky notes.", - ), - crate::mcp::tool_def::( - "condense_entity", - "Condense or summarize an entity's observations to reduce size.", - ), - crate::mcp::tool_def::( - "add_task", - "Add a new task to the task tracker.", - ), - crate::mcp::tool_def::( - "update_task_status", - "Update the status of an existing task.", - ), - crate::mcp::tool_def::( - "delete_task", - "Delete a task and all its children.", - ), - crate::mcp::tool_def::( - "list_active_tasks", - "List all currently active tasks.", - ), - crate::mcp::tool_def::( - "set_acceptance_criteria", - "Define a strict checklist of acceptance criteria for a given task.", - ), - crate::mcp::tool_def::( - "verify_acceptance_criteria", - "Mark a previously defined acceptance criteria as met.", - ), - crate::mcp::tool_def::( - "store_snippet", - "Store a reusable code snippet.", - ), - crate::mcp::tool_def::( - "search_snippets", - "Search through stored code snippets.", - ), - crate::mcp::tool_def::( - "delete_snippet", - "Delete a stored code snippet.", - ), - crate::mcp::tool_def::( - "log_decision", - "Log an architectural decision record (ADR).", - ), - crate::mcp::tool_def::( - "query_decisions", - "Query architectural decision records.", - ), - crate::mcp::tool_def::( - "merge_entities", - "Merge two entities in the knowledge graph into one.", - ), - crate::mcp::tool_def::( - "find_orphans", - "Find orphaned entities (entities without any relations) in the graph.", - ), - crate::mcp::tool_def::( - "learn_preference", - "Record a user preference or behavior to adapt future interactions.", - ), - crate::mcp::tool_def::( - "read_preferences", - "Read all learned user preferences.", - ), - crate::mcp::tool_def::( - "log_error_fix", - "Log a complex error and its fix for future reference.", - ), - crate::mcp::tool_def::( - "search_error_fixes", - "Search through previously logged error fixes.", - ), - crate::mcp::tool_def::( - "pin_file", - "Pin a file to keep it explicitly in the context workspace.", - ), - crate::mcp::tool_def::( - "unpin_file", - "Unpin a file from the context workspace.", - ), - crate::mcp::tool_def::( - "list_pinned_files", - "List all currently pinned files.", - ), - crate::mcp::tool_def::( - "add_session_summary", - "Add a summary of the current session.", - ), - crate::mcp::tool_def::( - "get_project_timeline", - "Get a timeline of major project events.", - ), - crate::mcp::tool_def::( - "leave_handoff_memo", - "Leave a memo for the next session or agent.", - ), - crate::mcp::tool_def::( - "read_handoff_memos", - "Read pending handoff memos.", - ), - crate::mcp::tool_def::( - "clear_handoff_memos", - "Clear handoff memos after reading.", - ), - crate::mcp::tool_def::( - "update_env_fingerprint", - "Update the environment fingerprint (e.g., OS, tool versions).", - ), - crate::mcp::tool_def::( - "read_env_fingerprint", - "Read the current environment fingerprint.", - ), - crate::mcp::tool_def::( - "log_env_requirement", - "Log a required tool or package for the environment.", - ), - crate::mcp::tool_def::( - "add_milestone", - "Add a new project milestone.", - ), - crate::mcp::tool_def::( - "update_milestone", - "Update the status of a project milestone.", - ), - crate::mcp::tool_def::( - "list_milestones", - "List all project milestones.", - ), - crate::mcp::tool_def::( - "generate_standup_report", - "Generate a standup report summarizing recent work, blockers, and next steps.", - ), - crate::mcp::tool_def::( - "register_environment", - "Register details about a specific deployment environment.", - ), - crate::mcp::tool_def::( - "get_environment_details", - "Get detailed information about a specific deployment environment.", - ), - crate::mcp::tool_def::( - "add_pr_checklist_item", - "Add an item to the PR checklist.", - ), - crate::mcp::tool_def::( - "get_pr_checklist", - "Get the current PR checklist.", - ), - crate::mcp::tool_def::( - "clear_pr_checklist", - "Clear the PR checklist.", - ), - crate::mcp::tool_def::( - "log_tech_debt", - "Log identified technical debt.", - ), - crate::mcp::tool_def::( - "resolve_tech_debt", - "Mark a logged technical debt as resolved.", - ), - crate::mcp::tool_def::( - "list_tech_debt", - "List all unresolved technical debt.", - ), - crate::mcp::tool_def::( - "save_context_workspace", - "Save the current set of pinned files and context.", - ), - crate::mcp::tool_def::( - "load_context_workspace", - "Load a previously saved context workspace.", - ), - crate::mcp::tool_def::( - "list_context_workspaces", - "List all saved context workspaces.", - ), - crate::mcp::tool_def::( - "omni_search", - "Search across all memory sources (graph, tasks, snippets, ADRs, etc.) at once.", - ), - crate::mcp::tool_def::( - "get_project_health", - "Get a synthesized health report of the project based on memory data.", - ), - ]; + let mut tools: Vec = + self.tools.values().map(|t| t.schema()).collect(); + tools.sort_by_key(|t| { + t.get("name") + .and_then(|n| n.as_str()) + .unwrap_or("") + .to_string() + }); Some(crate::mcp::success( id, serde_json::json!({ "tools": tools }), @@ -380,2683 +155,35 @@ impl MemoryHandler { self.state .broadcast_activity(&format!("Agent executed tool: {}", name)); - let result: Result = match name { - "query_graph_path" => { - let req = parse_tool!(args, id, crate::tools::QueryGraphPathTool); - self.state.read_graph(|graph| { - let max_depth = req.max_depth.unwrap_or(5); - let mut queue = std::collections::VecDeque::new(); - let mut visited = std::collections::HashSet::new(); - let mut parents: std::collections::HashMap = - std::collections::HashMap::new(); - - queue.push_back(req.start_node.clone()); - visited.insert(req.start_node.clone()); - - let mut found = false; - let mut current_depth = 0; - let mut nodes_at_current_depth = 1; - let mut nodes_at_next_depth = 0; - - while let Some(current) = queue.pop_front() { - if current == req.end_node { - found = true; - break; - } - nodes_at_current_depth -= 1; - if current_depth < max_depth { - for rel in &graph.relations { - if rel.from == current && !visited.contains(&rel.to) { - visited.insert(rel.to.clone()); - parents.insert( - rel.to.clone(), - (current.clone(), rel.relation_type.clone()), - ); - queue.push_back(rel.to.clone()); - nodes_at_next_depth += 1; - } else if rel.to == current && !visited.contains(&rel.from) { - visited.insert(rel.from.clone()); - parents.insert( - rel.from.clone(), - ( - current.clone(), - format!("inverse({})", rel.relation_type), - ), - ); - queue.push_back(rel.from.clone()); - nodes_at_next_depth += 1; - } - } - } - if nodes_at_current_depth == 0 { - current_depth += 1; - nodes_at_current_depth = nodes_at_next_depth; - nodes_at_next_depth = 0; - } - } - - if found { - let mut path = Vec::new(); - let mut curr = req.end_node.clone(); - while curr != req.start_node { - let (parent, rel_type) = parents.get(&curr).unwrap().clone(); - path.push(format!("{} -[{}]-> {}", parent, rel_type, curr)); - curr = parent; - } - path.reverse(); - Ok(format!("Path found:\n{}", path.join("\n"))) - } else { - Ok(format!("No path found between {} and {} within depth {}", req.start_node, req.end_node, max_depth)) - } - }) - } - "create_entities" => { - let req = parse_tool!(args, id, CreateEntitiesTool); - self.state.modify_graph(|g| { - for entity in req.entities { - if !entity.name.is_empty() { - if let Ok(idx) = self.state.search_index.read() { - drop(idx.index_entity(&entity)); - } - g.entities.insert(entity.name.clone(), entity); - } - } - }); - Ok("Entities created".to_string()) - } - "create_relations" => { - let req = parse_tool!(args, id, CreateRelationsTool); - self.state.modify_graph(|g| { - for relation in req.relations { - if !relation.from.is_empty() && !relation.to.is_empty() { - g.relations.push(relation); - } - } - }); - Ok("Relations created".to_string()) - } - "add_observations" => { - let req = parse_tool!(args, id, AddObservationsTool); - self.state.modify_graph(|g| { - for o in req.observations { - if let Some(e) = g.entities.get_mut(&o.entity_name) { - e.observations.extend(o.contents); - } - } - }); - Ok("Observations added".to_string()) - } - "delete_entities" => { - let req = parse_tool!(args, id, DeleteEntitiesTool); - let to_delete: HashSet<_> = req.entity_names.into_iter().collect(); - self.state.modify_graph(|master| { - for name in &to_delete { - master.entities.remove(name); - } - master.relations.retain(|r| { - !to_delete.contains(&r.from) && !to_delete.contains(&r.to) - }); - }); - Ok("Entities deleted".to_string()) - } - "delete_observations" => { - let req = parse_tool!(args, id, DeleteObservationsTool); - self.state.modify_graph(|master| { - for d in req.deletions { - if let Some(e) = master.entities.get_mut(&d.entity_name) { - let to_rem: HashSet<_> = d.observations.into_iter().collect(); - e.observations.retain(|o| !to_rem.contains(o)); - } - } - }); - Ok("Observations deleted".to_string()) - } - "delete_relations" => { - let req = parse_tool!(args, id, DeleteRelationsTool); - self.state.modify_graph(|master| { - let to_rem: HashSet<_> = req.relations.into_iter().collect(); - master.relations.retain(|r| !to_rem.contains(r)); - }); - Ok("Relations deleted".to_string()) - } - "read_graph" => { - let req = parse_tool!(args, id, ReadGraphTool); - let data = self.state.read_graph(|full| { - if let Some(ns) = req.namespace { - let mut filtered = KnowledgeGraph::default(); - for (k, v) in &full.entities { - if v.namespace == ns { - filtered.entities.insert(k.clone(), v.clone()); - } - } - for r in &full.relations { - if r.namespace == ns { - filtered.relations.push(r.clone()); - } - } - serde_json::to_string(&filtered).unwrap_or_default() - } else { - serde_json::to_string(full).unwrap_or_default() - } - }); - Ok(data) - } - "search_nodes" => { - let req = parse_tool!(args, id, SearchNodesTool); - let matches = if let Ok(idx) = self.state.search_index.read() { - idx.search(&req.query, req.namespace.as_deref()) - .unwrap_or_default() - } else { - vec![] - }; - - let mut result = KnowledgeGraph::default(); - self.state.read_graph(|full| { - for (id, doc_type, _, _, _) in matches { - if doc_type == "entity" - && let Some(e) = full.entities.get(&id) - { - result.entities.insert(id, e.clone()); - } - } - }); - let data = serde_json::to_string(&result).unwrap_or_default(); - Ok(data.to_string()) - } - "open_nodes" => { - let req = parse_tool!(args, id, OpenNodesTool); - let targets: HashSet<_> = req.names.into_iter().collect(); - let mut result = KnowledgeGraph::default(); - let mut connected = HashSet::new(); - self.state.read_graph(|full| { - for r in &full.relations { - if targets.contains(&r.from) { - connected.insert(r.to.clone()); - result.relations.push(r.clone()); - } else if targets.contains(&r.to) { - connected.insert(r.from.clone()); - result.relations.push(r.clone()); - } - } - for (name, e) in &full.entities { - if targets.contains(name) || connected.contains(name) { - result.entities.insert(name.clone(), e.clone()); - } - } - }); - let data = serde_json::to_string(&result).unwrap_or_default(); - Ok(data.to_string()) - } - "log_code_change" => { - let req = parse_tool!(args, id, LogCodeChangeTool); - self.state.ledger.modify(|ledger| { - ledger.push(CodeChange { - timestamp: SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(), - file_path: req.file_path, - description: req.description, - git_commit: req.git_commit, - git_branch: req.git_branch, - }); - }); - Ok("Code change logged".to_string()) - } - "query_recent_changes" => { - let data = serde_json::to_string(&self.state.ledger.read()) - .unwrap_or_else(|_| "[]".to_string()); - Ok(data.to_string()) - } - "visualize_graph" => { - let req = parse_tool!(args, id, VisualizeGraphTool); - let query = req.query.unwrap_or_default().to_lowercase(); - let mut included = HashSet::new(); - let mut to_draw = Vec::new(); - - self.state.read_graph(|full| { - for (name, e) in &full.entities { - if let Some(ns) = &req.namespace - && e.namespace != *ns - { - continue; - } - if query.is_empty() - || contains_ignore_ascii_case(&name, &query) - || contains_ignore_ascii_case(&e.entity_type, &query) - { - included.insert(name.clone()); - } - } - - for r in &full.relations { - if let Some(ns) = &req.namespace - && r.namespace != *ns - { - continue; - } - if query.is_empty() - || included.contains(&r.from) - || included.contains(&r.to) - { - included.insert(r.from.clone()); - included.insert(r.to.clone()); - to_draw.push(r.clone()); - } - } - }); - use std::fmt::Write; - let mut output = String::with_capacity(included.len() * 40 + to_draw.len() * 60); - output.push_str("graph TD;\n"); - - let sanitize = |s: &str, id_mode: bool| -> String { - let mut out = String::with_capacity(s.len()); - for c in s.chars() { - if c != '"' && c != '(' && c != ')' { - if id_mode && (c == ' ' || c == '-' || c == '.') { - out.push('_'); - } else { - out.push(c); - } - } - } - out - }; - - for name in &included { - let _ = writeln!( - output, - " id_{}[\"{}\"];", - sanitize(name, true), - sanitize(name, false) - ); - } - for r in to_draw { - let _ = writeln!( - output, - " id_{}-->|\"{}\"|id_{};", - sanitize(&r.from, true), - r.relation_type.replace("\"", ""), - sanitize(&r.to, true) - ); - } - if output == "graph TD;\n" { - output = "No nodes found to visualize.".to_string(); - } - Ok(output.to_string()) - } - "add_sticky_note" => { - let req = parse_tool!(args, id, AddStickyNoteTool); - self.state.sticky.modify(|notes| { - notes.push(StickyNote { - timestamp: SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(), - content: req.content, - }); - }); - Ok("Sticky note added.".to_string()) - } - "read_sticky_notes" => { - let data = serde_json::to_string(&self.state.sticky.read()) - .unwrap_or_else(|_| "[]".to_string()); - Ok(data.to_string()) - } - "delete_sticky_note" => { - let req = parse_tool!(args, id, DeleteStickyNoteTool); - let mut success = false; - self.state.sticky.modify(|notes| { - if req.index > 0 && req.index <= notes.len() { - notes.remove(req.index - 1); - success = true; - } - }); - if success { - Ok("Sticky note deleted.".to_string()) - } else { - Err("Invalid sticky note index.".to_string()) - } - } - "clear_sticky_notes" => { - self.state.sticky.modify(|notes| { - notes.clear(); - }); - Ok("All sticky notes cleared.".to_string()) - } - "condense_entity" => { - let req = parse_tool!(args, id, CondenseEntityTool); - self.state.modify_graph(|master| { - if let Some(e) = master.entities.get_mut(&req.entity_name) { - e.observations = req.summarized_observations; - } - }); - Ok("Entity condensed".to_string()) - } - "add_task" => { - let req = parse_tool!(args, id, AddTaskTool); - let now = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(); - let task_id = uuid::Uuid::new_v4().to_string(); - - let parent_id = req.parent_id.clone(); - let deps = req.dependencies.clone().unwrap_or_default(); - - let task = Task { - id: task_id.clone(), - title: req.title, - status: "pending".to_string(), - description: req.description, - created_at: now, - updated_at: now, - git_branch: req.git_branch, - parent_id, - dependencies: deps, - acceptance_criteria: vec![], - }; - if let Ok(idx) = self.state.search_index.read() { - drop(idx.index_task(&task)); - } - self.state.tasks.modify(|tasks| { - tasks.push(task); - }); - Ok(format!("Task added with ID: {}", task_id).to_string()) - } - "delete_task" => { - let req = parse_tool!(args, id, DeleteTaskTool); - let mut deleted_count = 0; - self.state.tasks.modify(|tasks| { - let initial_len = tasks.len(); - // Collect IDs of tasks to delete (this task + all its recursive children) - let mut to_delete = std::collections::HashSet::new(); - to_delete.insert(req.id.clone()); - - let mut added_new = true; - while added_new { - added_new = false; - for t in tasks.iter() { - if let Some(pid) = &t.parent_id - && to_delete.contains(pid) && !to_delete.contains(&t.id) { - to_delete.insert(t.id.clone()); - added_new = true; - } - } - } - - tasks.retain(|t| !to_delete.contains(&t.id)); - deleted_count = initial_len - tasks.len(); - }); - - if deleted_count > 0 { - Ok(vec![ - format!("Deleted task and its children ({} total).", deleted_count) - .to_string(), - ][0] - .clone()) - } else { - Ok("Task not found.".to_string()) - } - } - "update_task_status" => { - let req = parse_tool!(args, id, UpdateTaskStatusTool); - let mut found = false; - let mut blocked = false; - let mut blocker_details = String::new(); - let target_status = req.status.to_lowercase(); - - self.state.tasks.modify(|tasks| { - // Find target task - let mut target_id = String::new(); - if let Some(t) = - tasks.iter().find(|t| t.id == req.id || t.title == req.id) - { - target_id = t.id.clone(); - } - - if target_id.is_empty() { - return; - } - found = true; - - if target_status == "done" || target_status == "completed" { - // 1. Check Acceptance Criteria - if let Some(t) = tasks.iter().find(|t| t.id == target_id) - && t.acceptance_criteria.iter().any(|c| !c.is_met) { - blocked = true; - blocker_details = - "Unmet acceptance criteria exist.".to_string(); - } - - // 2. Check dependencies - if !blocked { - let mut uncompleted_deps = Vec::new(); - if let Some(t) = tasks.iter().find(|t| t.id == target_id) { - for dep_id in &t.dependencies { - if let Some(dep_task) = - tasks.iter().find(|dt| dt.id == *dep_id) - && dep_task.status != "completed" - && dep_task.status != "done" - { - uncompleted_deps.push(dep_task.title.clone()); - } - } - } - if !uncompleted_deps.is_empty() { - blocked = true; - blocker_details = format!( - "Blocked by dependencies: {}", - uncompleted_deps.join(", ") - ); - } - } - - // 3. Check child tasks - if !blocked { - let mut uncompleted_children = Vec::new(); - for child in tasks - .iter() - .filter(|t| t.parent_id.as_ref() == Some(&target_id)) - { - if child.status != "completed" && child.status != "done" { - uncompleted_children.push(child.title.clone()); - } - } - if !uncompleted_children.is_empty() { - blocked = true; - blocker_details = format!( - "Blocked by child tasks: {}", - uncompleted_children.join(", ") - ); - } - } - } - - if !blocked { - // Apply update - if let Some(t) = tasks.iter_mut().find(|t| t.id == target_id) { - t.status = target_status.clone(); - t.updated_at = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(); - } - - // Cascade cancellation to children - if target_status == "cancelled" || target_status == "abandoned" { - let mut to_cancel = vec![target_id.clone()]; - let mut i = 0; - while i < to_cancel.len() { - let current_pid = to_cancel[i].clone(); - for t in tasks.iter_mut() { - if t.parent_id.as_ref() == Some(¤t_pid) - && t.status != "completed" - { - t.status = target_status.clone(); - to_cancel.push(t.id.clone()); - } - } - i += 1; - } - } - } - }); - - if blocked { - Ok(vec![format!( - "Error: Cannot transition task. {}", - blocker_details - )] - .into_iter() - .next() - .unwrap()) - } else if found { - Ok("Task status updated.".to_string()) - } else { - Ok("Task not found.".to_string()) - } - } - "list_active_tasks" => { - let req = parse_tool!(args, id, ListActiveTasksTool); - 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(data.to_string()) - } - "set_acceptance_criteria" => { - let req = parse_tool!(args, id, SetAcceptanceCriteriaTool); - let mut success = false; - self.state.tasks.modify(|tasks| { - if let Some(task) = - tasks.iter_mut().rev().find(|t| t.title == req.task_title) - { - task.acceptance_criteria = req - .criteria - .into_iter() - .map(|desc| crate::models::AcceptanceCriteria { - id: uuid::Uuid::new_v4().to_string(), - description: desc, - is_met: false, - }) - .collect(); - task.updated_at = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(); - success = true; - } - }); - if success { - Ok( - vec!["Acceptance criteria set successfully.".to_string()][0] - .clone(), - ) - } else { - Ok("Task not found.".to_string()) - } - } - "verify_acceptance_criteria" => { - let req = parse_tool!(args, id, VerifyAcceptanceCriteriaTool); - let mut success = false; - let mut already_met = false; - self.state.tasks.modify(|tasks| { - if let Some(task) = tasks.iter_mut().find(|t| t.id == req.task_id) - && let Some(ac) = task - .acceptance_criteria - .iter_mut() - .find(|c| c.id == req.criteria || c.description == req.criteria) - { - if ac.is_met { - already_met = true; - } else { - ac.is_met = true; - success = true; - task.updated_at = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(); - } - } - }); - if success { - Ok(vec![format!( - "Acceptance criteria verified with proof: {}", - req.proof - )][0] - .clone()) - } else if already_met { - Ok("Acceptance criteria was already met.".to_string()) - } else { - Ok( - vec!["Acceptance criteria or task not found.".to_string()][0] - .clone(), - ) - } - } - "store_snippet" => { - let req = parse_tool!(args, id, StoreSnippetTool); - let snippet = Snippet { - name: req.name.clone(), - language: req.language, - code: req.code, - description: req.description, - updated_at: SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(), - }; - - let s_clone = snippet.clone(); - self.state.snippets.modify(|snippets| { - snippets.retain(|s| s.name != req.name); - snippets.push(s_clone); - }); - - if let Ok(idx) = self.state.search_index.read() { - drop(idx.index_snippet(&snippet)); - } - - Ok(format!("Snippet '{}' stored.", req.name).to_string()) - } - "search_snippets" => { - let req = parse_tool!(args, id, SearchSnippetsTool); - let query = req.query.to_lowercase(); - let snippets = self.state.snippets.read(); - let mut results = Vec::new(); - for s in snippets { - if contains_ignore_ascii_case(&s.name, &query) - || contains_ignore_ascii_case(&s.description, &query) - || contains_ignore_ascii_case(&s.language, &query) - { - results.push(s); - } - } - let data = serde_json::to_string(&results).unwrap_or_default(); - Ok(data.to_string()) - } - "delete_snippet" => { - let req = parse_tool!(args, id, DeleteSnippetTool); - let mut deleted = false; - self.state.snippets.modify(|snippets| { - let orig = snippets.len(); - snippets.retain(|s| s.name != req.name); - deleted = snippets.len() < orig; - }); - if deleted { - Ok("Snippet deleted.".to_string()) - } else { - Ok("Snippet not found.".to_string()) - } - } - "log_decision" => { - let req = parse_tool!(args, id, LogDecisionTool); - let mut adr_id = String::new(); - let mut new_adr = None; - - self.state.adrs.modify(|adrs| { - adr_id = format!("ADR-{:04}", adrs.len() + 1); - let a = Adr { - id: adr_id.clone(), - title: req.title, - context: req.context, - decision: req.decision, - consequence: req.consequence, - timestamp: SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(), - }; - new_adr = Some(a.clone()); - adrs.push(a); - }); - - if let Some(adr) = new_adr - && let Ok(idx) = self.state.search_index.read() { - drop(idx.index_adr(&adr)); - } - - Ok(format!("Decision logged as {}", adr_id).to_string()) - } - "query_decisions" => { - let req = parse_tool!(args, id, QueryDecisionsTool); - let mut adrs = self.state.adrs.read(); - if let Some(q) = req.query { - let q = q.to_lowercase(); - adrs.retain(|a| { - contains_ignore_ascii_case(&a.title, &q) - || contains_ignore_ascii_case(&a.context, &q) - || contains_ignore_ascii_case(&a.decision, &q) - }); - } - let data = serde_json::to_string(&adrs).unwrap_or_default(); - Ok(data.to_string()) - } - "merge_entities" => { - let req = parse_tool!(args, id, MergeEntitiesTool); - self.state.modify_graph(|master| { - if let Some(src) = master.entities.remove(&req.source_entity) { - if let Some(tgt) = master.entities.get_mut(&req.target_entity) { - tgt.observations.extend(src.observations); - MemoryState::deduplicate(&mut tgt.observations); - } else { - let mut new_tgt = src.clone(); - new_tgt.name = req.target_entity.clone(); - master.entities.insert(req.target_entity.clone(), new_tgt); - } - } - for r in &mut master.relations { - if r.from == req.source_entity { - r.from = req.target_entity.clone(); - } - if r.to == req.source_entity { - r.to = req.target_entity.clone(); - } - } - MemoryState::deduplicate(&mut master.relations); - }); - Ok("Entities merged".to_string()) - } - "find_orphans" => { - let orphans = self.state.read_graph(|full| { - let mut connected = std::collections::HashSet::new(); - for r in &full.relations { - connected.insert(r.from.clone()); - connected.insert(r.to.clone()); - } - full.entities - .keys() - .filter(|k| !connected.contains(*k)) - .cloned() - .collect::>() - }); - let data = serde_json::to_string(&orphans).unwrap_or_default(); - Ok(data.to_string()) - } - "learn_preference" => { - let req = parse_tool!(args, id, LearnPreferenceTool); - self.state.prefs.modify(|prefs| { - prefs.insert( - req.key.clone(), - crate::models::Preference { - key: req.key.clone(), - value: req.value, - updated_at: SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(), - }, - ); - }); - Ok("Preference learned".to_string()) - } - "read_preferences" => { - let prefs = self.state.prefs.read(); - let data = serde_json::to_string(&prefs).unwrap_or_default(); - Ok(data.to_string()) - } - "log_error_fix" => { - let req = parse_tool!(args, id, LogErrorFixTool); - self.state.error_fixes.modify(|fixes| { - fixes.push(crate::models::ErrorFix { - signature: req.signature, - solution: req.solution, - timestamp: SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(), - git_commit: req.git_commit, - git_branch: req.git_branch, - }) - }); - Ok("Error fix logged".to_string()) - } - "search_error_fixes" => { - let req = parse_tool!(args, id, SearchErrorFixesTool); - let q = req.query.to_lowercase(); - let mut fixes = self.state.error_fixes.read(); - fixes.retain(|f| { - contains_ignore_ascii_case(&f.signature, &q) - || contains_ignore_ascii_case(&f.solution, &q) - }); - let data = serde_json::to_string(&fixes).unwrap_or_default(); - Ok(data.to_string()) - } - "pin_file" => { - let req = parse_tool!(args, id, PinFileTool); - self.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: SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(), - git_branch: req.git_branch, - }); - }); - Ok("File pinned".to_string()) - } - "unpin_file" => { - let req = parse_tool!(args, id, UnpinFileTool); - self.state.pinned_files.modify(|pinned| { - pinned.retain(|p| { - !(p.namespace == req.namespace && p.file_path == req.file_path) - }) - }); - Ok("File unpinned".to_string()) - } - "list_pinned_files" => { - let req = parse_tool!(args, id, ListPinnedFilesTool); - let mut pinned = self.state.pinned_files.read(); - 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(data.to_string()) - } - "add_session_summary" => { - let req = parse_tool!(args, id, AddSessionSummaryTool); - self.state.session_summaries.modify(|summaries| { - summaries.push(crate::models::SessionSummary { - summary: req.summary, - namespace: req.namespace, - timestamp: SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(), - }) - }); - Ok("Session summary added".to_string()) - } - "get_project_timeline" => { - let req = parse_tool!(args, id, GetProjectTimelineTool); - let mut summaries = self.state.session_summaries.read(); - if let Some(ns) = req.namespace { - summaries.retain(|s| s.namespace == ns); - } - summaries.sort_by_key(|s| s.timestamp); - let data = serde_json::to_string(&summaries).unwrap_or_default(); - Ok(data.to_string()) - } - "leave_handoff_memo" => { - let req = parse_tool!(args, id, LeaveHandoffMemoTool); - self.state.handoff_memos.modify(|memos| { - memos.push(crate::models::HandoffMemo { - id: uuid::Uuid::new_v4().to_string(), - author: "agy".to_string(), - content: req.content, - namespace: req.namespace, - timestamp: SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(), - }) - }); - Ok("Handoff memo left".to_string()) - } - "read_handoff_memos" => handle_list_with_namespace!(self, handoff_memos, ReadHandoffMemosTool, args, id), - "clear_handoff_memos" => { - let req = parse_tool!(args, id, ClearHandoffMemosTool); - let ids: HashSet<_> = req.ids.into_iter().collect(); - self.state - .handoff_memos - .modify(|memos| memos.retain(|m| !ids.contains(&m.id))); - Ok("Handoff memos cleared".to_string()) - } - "update_env_fingerprint" => { - let req = parse_tool!(args, id, UpdateEnvFingerprintTool); - self.state.env_fingerprints.modify(|fps| { - fps.insert( - req.namespace.clone(), - crate::models::EnvFingerprint { - namespace: req.namespace.clone(), - os: std::env::consts::OS.to_string(), - shell: std::env::var("SHELL") - .unwrap_or_else(|_| "unknown".to_string()), - tool_versions: req.tool_versions, - updated_at: SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(), - }, - ); - }); - Ok("Env fingerprint updated".to_string()) - } - "read_env_fingerprint" => { - let req = parse_tool!(args, id, ReadEnvFingerprintTool); - let fps = self.state.env_fingerprints.read(); - if let Some(fp) = fps.get(&req.namespace) { - let data = serde_json::to_string(fp).unwrap_or_default(); - Ok(data.to_string()) - } else { - Ok("{}".to_string()) - } - } - "log_env_requirement" => { - let req = parse_tool!(args, id, LogEnvRequirementTool); - self.state.env_requirements.modify(|reqs| { - reqs.retain(|r| !(r.namespace == req.namespace && r.key == req.key)); - reqs.push(crate::models::EnvRequirement { - namespace: req.namespace, - key: req.key, - description: req.description, - is_secret: req.is_secret, - }); - }); - Ok("Env requirement logged".to_string()) - } - "add_milestone" => { - let req = parse_tool!(args, id, AddMilestoneTool); - self.state.milestones.modify(|ms| { - ms.push(crate::models::Milestone { - id: uuid::Uuid::new_v4().to_string(), - title: req.title, - status: "pending".to_string(), - namespace: req.namespace, - target_date: None, - }) - }); - Ok("Milestone added".to_string()) - } - "update_milestone" => { - let req = parse_tool!(args, id, UpdateMilestoneTool); - let mut found = false; - self.state.milestones.modify(|ms| { - for m in ms.iter_mut() { - if m.id == req.id { - m.status = req.status.clone(); - found = true; - break; - } - } - }); - if found { - Ok("Milestone updated".to_string()) - } else { - Ok("Milestone not found".to_string()) - } - } - "list_milestones" => handle_list_with_namespace!(self, milestones, ListMilestonesTool, args, id), - "generate_standup_report" => { - let req = parse_tool!(args, id, GenerateStandupReportTool); - let cutoff = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs() - .saturating_sub(req.hours_lookback * 3600); - let tasks = self - .state - .tasks - .read() - .into_iter() - .filter(|t| t.updated_at >= cutoff) - .collect::>(); - let changes = self - .state - .ledger - .read() - .into_iter() - .filter(|c| c.timestamp >= cutoff) - .collect::>(); - let summaries = self - .state - .session_summaries - .read() - .into_iter() - .filter(|s| s.namespace == req.namespace && s.timestamp >= cutoff) - .collect::>(); - let report = serde_json::json!({ "tasks_updated": tasks, "code_changes": changes, "session_summaries": summaries }); - Ok(report.to_string()) - } - "register_environment" => { - let req = parse_tool!(args, id, RegisterEnvironmentTool); - self.state.environments.modify(|envs| { - envs.retain(|e| !(e.namespace == req.namespace && e.name == req.name)); - envs.push(crate::models::EnvironmentDetail { - namespace: req.namespace, - name: req.name, - url: req.url, - description: req.description, - requires_vpn: req.requires_vpn, - updated_at: SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(), - }); - }); - Ok("Environment registered".to_string()) - } - "get_environment_details" => { - let req = parse_tool!(args, id, GetEnvironmentDetailsTool); - let mut envs = self.state.environments.read(); - envs.retain(|e| e.namespace == req.namespace); - let data = serde_json::to_string(&envs).unwrap_or_default(); - Ok(data.to_string()) - } - "add_pr_checklist_item" => { - let req = parse_tool!(args, id, AddPrChecklistItemTool); - self.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()) - } - "get_pr_checklist" => { - let req = parse_tool!(args, id, GetPrChecklistTool); - let mut items = self.state.pr_checklists.read(); - items.retain(|i| i.namespace == req.namespace); - let data = serde_json::to_string(&items).unwrap_or_default(); - Ok(data.to_string()) - } - "clear_pr_checklist" => { - let req = parse_tool!(args, id, ClearPrChecklistTool); - self.state - .pr_checklists - .modify(|items| items.retain(|i| i.namespace != req.namespace)); - Ok("PR checklist cleared".to_string()) - } - "log_tech_debt" => { - let req = parse_tool!(args, id, LogTechDebtTool); - self.state.tech_debts.modify(|debts| { - debts.push(crate::models::TechDebt { - id: uuid::Uuid::new_v4().to_string(), - namespace: req.namespace, - description: req.description, - ideal_solution: req.ideal_solution, - is_resolved: false, - created_at: SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(), - git_commit: req.git_commit, - git_branch: req.git_branch, - }) - }); - Ok("Tech debt logged".to_string()) - } - "resolve_tech_debt" => { - let req = parse_tool!(args, id, ResolveTechDebtTool); - let mut found = false; - self.state.tech_debts.modify(|debts| { - for d in debts.iter_mut() { - if d.id == req.id { - d.is_resolved = true; - found = true; - break; - } - } - }); - if found { - Ok("Tech debt resolved".to_string()) - } else { - Ok("Tech debt not found".to_string()) - } - } - "list_tech_debt" => { - let req = parse_tool!(args, id, ListTechDebtTool); - let mut debts = self.state.tech_debts.read(); - debts.retain(|d| { - d.namespace == req.namespace && (req.include_resolved || !d.is_resolved) - }); - let data = serde_json::to_string(&debts).unwrap_or_default(); - Ok(data.to_string()) - } - "save_context_workspace" => { - let req = parse_tool!(args, id, SaveContextWorkspaceTool); - self.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: SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs(), - }); - }); - Ok("Context workspace saved".to_string()) - } - "load_context_workspace" => { - let req = parse_tool!(args, id, LoadContextWorkspaceTool); - let mut ws = self.state.context_workspaces.read(); - ws.retain(|w| w.namespace == req.namespace && w.name == req.name); - let data = serde_json::to_string(&ws.first()).unwrap_or_default(); - Ok(data.to_string()) - } - "list_context_workspaces" => { - let req = parse_tool!(args, id, ListContextWorkspacesTool); - let mut ws = self.state.context_workspaces.read(); - ws.retain(|w| w.namespace == req.namespace); - let data = serde_json::to_string(&ws).unwrap_or_default(); - Ok(data.to_string()) - } - "omni_search" => { - let req = parse_tool!(args, id, OmniSearchTool); - let matches = if let Ok(idx) = self.state.search_index.read() { - idx.search(&req.query, req.namespace.as_deref()) - .unwrap_or_default() - } else { - vec![] - }; - - let mut kg = KnowledgeGraph::default(); - let mut tasks = Vec::new(); - let mut snippets = Vec::new(); - let mut adrs = Vec::new(); - - self.state.read_graph(|full| { - for (id, doc_type, _, _, _) in &matches { - if doc_type == "entity" - && let Some(e) = full.entities.get(id) - { - kg.entities.insert(id.clone(), e.clone()); - } - } - }); - for t in self.state.tasks.read() { - if matches.iter().any(|(id, typ, _, _, _)| id == &t.id && typ == "task") { - tasks.push(t); - } - } - for s in self.state.snippets.read() { - if matches - .iter() - .any(|(id, typ, _, _, _)| id == &s.name && typ == "snippet") - { - snippets.push(s); - } - } - for a in self.state.adrs.read() { - if matches.iter().any(|(id, typ, _, _, _)| id == &a.id && typ == "adr") { - adrs.push(a); - } - } - - let q = req.query.to_lowercase(); - let tech_debts: Vec<_> = self - .state - .tech_debts - .read() - .into_iter() - .filter(|d| { - (req.namespace.is_none() - || d.namespace == *req.namespace.as_ref().unwrap()) - && (contains_ignore_ascii_case(&d.description, &q) - || contains_ignore_ascii_case(&d.ideal_solution, &q)) - }) - .collect(); - let memos: Vec<_> = self - .state - .handoff_memos - .read() - .into_iter() - .filter(|m| { - (req.namespace.is_none() - || m.namespace == *req.namespace.as_ref().unwrap()) - && contains_ignore_ascii_case(&m.content, &q) - }) - .collect(); - let error_fixes: Vec<_> = self - .state - .error_fixes - .read() - .into_iter() - .filter(|f| { - contains_ignore_ascii_case(&f.signature, &q) - || contains_ignore_ascii_case(&f.solution, &q) - }) - .collect(); - - let report = serde_json::json!({ - "knowledge_graph": kg.entities, - "tasks": tasks, - "snippets": snippets, - "adrs": adrs, - "tech_debts": tech_debts, - "handoff_memos": memos, - "error_fixes": error_fixes - }); - Ok(report.to_string()) - } - "get_project_health" => { - let req = parse_tool!(args, id, GetProjectHealthTool); - let active_tasks = self - .state - .tasks - .read() - .into_iter() - .filter(|t| t.status != "done") - .count(); - let unresolved_debt = self - .state - .tech_debts - .read() - .into_iter() - .filter(|d| d.namespace == req.namespace && !d.is_resolved) - .count(); - let unread_memos = self - .state - .handoff_memos - .read() - .into_iter() - .filter(|m| m.namespace == req.namespace) - .count(); - let active_milestones = self - .state - .milestones - .read() - .into_iter() - .filter(|m| m.namespace == req.namespace && m.status != "done") - .count(); - let remaining_checklists = self - .state - .pr_checklists - .read() - .into_iter() - .filter(|c| c.namespace == req.namespace) - .count(); - - let report = serde_json::json!({ - "active_tasks": active_tasks, - "unresolved_tech_debt": unresolved_debt, - "unread_handoff_memos": unread_memos, - "active_milestones": active_milestones, - "remaining_pr_checklist_items": remaining_checklists - }); - Ok(report.to_string()) - } - - _ => Err(format!("Unknown tool: {}", name)), + let result: Result = if let Some(tool) = self.tools.get(name) { + tool.execute(args, self.state.clone()).await + } else { + Err(format!("Unknown tool: {}", name)) }; match result { - Ok(text) => Some(crate::mcp::success( - id, - serde_json::json!({ - "content": [{ "type": "text", "text": text }] - }), - )), - Err(e) => Some(crate::mcp::success( - id, - serde_json::json!({ - "isError": true, - "content": [{ "type": "text", "text": e }] - }), - )), + Ok(text) => { + let payload = serde_json::json!({ + "content": [{"type": "text", "text": text}], + "isError": false + }); + Some(crate::mcp::success(id_clone, payload)) + } + Err(e) => { + tracing::error!("Tool {} failed: {}", name, e); + let payload = serde_json::json!({ + "content": [{"type": "text", "text": e}], + "isError": true + }); + Some(crate::mcp::success(id_clone, payload)) + } } } - _ => { - if id != serde_json::Value::Null { - Some(crate::mcp::error(id, -32601, "Method not found")) - } else { - None - } - } - }; - - let elapsed = start_time.elapsed(); - if method == "tools/call" { - let is_error = response.as_ref().is_some_and(|r| r.get("error").is_some() || r.get("result").and_then(|res| res.get("isError")).and_then(|e| e.as_bool()).unwrap_or(false)); - tracing::info!("<<< [Server] MCP tool call {} (id: {}) completed in {:?} [Error: {}]", tool_name, id_clone, elapsed, is_error); - - // Broadcast completion latency to the UI Activity Feed - let status_msg = if is_error { "with error" } else { "successfully" }; - self.state.broadcast_activity(&format!("Tool {} completed {} in {:?}", tool_name, status_msg, elapsed)); - } else { - tracing::debug!("<<< [Server] MCP request method {} completed in {:?}", method, elapsed); + _ => Some(crate::mcp::error( + id, + -32601, + &format!("Method {} not found", method), + )), } - tracing::trace!("Returning response from handle_request: {:?}", response); - response - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::state::MemoryState; - use serde_json::json; - use std::sync::Arc; - - #[tokio::test] - async fn test_handle_initialize() { - let store_dir = std::env::temp_dir().join(format!( - "mcp_test_handlers_{}", - std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs() - )); - std::fs::create_dir_all(&store_dir).unwrap(); - let redb_path = store_dir.join("mcp_store.redb"); - let db = Arc::new(redb::Database::create(&redb_path).unwrap()); - - { - let write_txn = db.begin_write().unwrap(); - let _ = write_txn.open_table(crate::store::STORE_TABLE); - write_txn.commit().unwrap(); - } - - let state = Arc::new(MemoryState { - base_dir: store_dir.clone(), - graph: crate::store::Store::new("knowledge_graph_master", db.clone()), - search_index: std::sync::RwLock::new( - crate::search::MemoryIndex::new(&store_dir).unwrap(), - ), - ledger: crate::store::Store::new("audit_ledger", db.clone()), - sticky: crate::store::Store::new("sticky_notes", db.clone()), - tasks: crate::store::Store::new("tasks", db.clone()), - snippets: crate::store::Store::new("snippets", db.clone()), - adrs: crate::store::Store::new("adrs", db.clone()), - prefs: crate::store::Store::new("preferences", db.clone()), - error_fixes: crate::store::Store::new("error_fixes", db.clone()), - pinned_files: crate::store::Store::new("pinned_files", db.clone()), - session_summaries: crate::store::Store::new("session_summaries", db.clone()), - handoff_memos: crate::store::Store::new("handoff_memos", db.clone()), - env_fingerprints: crate::store::Store::new("env_fingerprints", db.clone()), - env_requirements: crate::store::Store::new("env_requirements", db.clone()), - milestones: crate::store::Store::new("milestones", db.clone()), - environments: crate::store::Store::new("environments", db.clone()), - pr_checklists: crate::store::Store::new("pr_checklists", db.clone()), - tech_debts: crate::store::Store::new("tech_debts", db.clone()), - gates: crate::store::Store::new("gates", db.clone()), - context_workspaces: crate::store::Store::new("context_workspaces", db.clone()), - activity_tx: tokio::sync::broadcast::channel(100).0, - }); - let handler = MemoryHandler { state }; - - let req = json!({ - "jsonrpc": "2.0", - "id": 1, - "method": "initialize", - "params": { - "protocolVersion": "2024-11-05", - "capabilities": {}, - "clientInfo": { - "name": "test-client", - "version": "1.0.0" - } - } - }); - - let response = handler - .handle_request(req) - .await - .expect("Expected a response"); - - assert_eq!(response["id"], 1); - assert!(response.get("result").is_some()); - - let result = &response["result"]; - // assert_eq!(result["protocolVersion"], "2024-11-05"); - - // CRITICAL BUG FIX CHECK: capabilities MUST contain an empty tools object - // Note: Currently it is set to `{}` which may cause proxy dropping tools. Let's verify it matches the actual behavior. - assert_eq!(result["capabilities"], serde_json::json!({"tools": {}})); - assert_eq!(result["serverInfo"]["name"], "gemini-mcp-memory"); - } - - fn setup_test_handler(test_name: &str) -> MemoryHandler { - let store_dir = std::env::temp_dir().join(format!( - "mcp_test_handlers_{}_{}", - test_name, - uuid::Uuid::new_v4() - )); - std::fs::create_dir_all(&store_dir).unwrap(); - let redb_path = store_dir.join("mcp_store.redb"); - let db = Arc::new(redb::Database::create(&redb_path).unwrap()); - { - let write_txn = db.begin_write().unwrap(); - let _ = write_txn.open_table(crate::store::STORE_TABLE); - write_txn.commit().unwrap(); - } - - let state = Arc::new(MemoryState { - base_dir: store_dir.clone(), - graph: crate::store::Store::new("knowledge_graph_master", db.clone()), - - search_index: std::sync::RwLock::new( - crate::search::MemoryIndex::new(&store_dir).unwrap(), - ), - ledger: crate::store::Store::new("audit_ledger", db.clone()), - sticky: crate::store::Store::new("sticky_notes", db.clone()), - tasks: crate::store::Store::new("tasks", db.clone()), - snippets: crate::store::Store::new("snippets", db.clone()), - adrs: crate::store::Store::new("adrs", db.clone()), - prefs: crate::store::Store::new("preferences", db.clone()), - error_fixes: crate::store::Store::new("error_fixes", db.clone()), - pinned_files: crate::store::Store::new("pinned_files", db.clone()), - session_summaries: crate::store::Store::new("session_summaries", db.clone()), - handoff_memos: crate::store::Store::new("handoff_memos", db.clone()), - env_fingerprints: crate::store::Store::new("env_fingerprints", db.clone()), - env_requirements: crate::store::Store::new("env_requirements", db.clone()), - milestones: crate::store::Store::new("milestones", db.clone()), - environments: crate::store::Store::new("environments", db.clone()), - pr_checklists: crate::store::Store::new("pr_checklists", db.clone()), - tech_debts: crate::store::Store::new("tech_debts", db.clone()), - gates: crate::store::Store::new("gates", db.clone()), - context_workspaces: crate::store::Store::new("context_workspaces", db.clone()), - activity_tx: tokio::sync::broadcast::channel(100).0, - }); - MemoryHandler { state } - } - - #[tokio::test] - async fn test_handle_tools_list() { - let handler = setup_test_handler("tools_list"); - - let req = json!({ - "jsonrpc": "2.0", - "id": 2, - "method": "tools/list", - "params": {} - }); - - let response = handler - .handle_request(req) - .await - .expect("Expected a response"); - assert_eq!(response["id"], 2); - - let tools = response["result"]["tools"] - .as_array() - .expect("Tools must be an array"); - assert!(!tools.is_empty()); - - // Verify a specific tool is registered - let add_task_tool = tools - .iter() - .find(|t| t["name"] == "add_task") - .expect("add_task tool missing"); - assert_eq!( - add_task_tool["description"], - "Add a new task to the task tracker." - ); - } - - #[tokio::test] - async fn test_handle_add_task() { - let handler = setup_test_handler("add_task"); - - let req = json!({ - "jsonrpc": "2.0", - "id": 3, - "method": "tools/call", - "params": { - "name": "add_task", - "arguments": { - "title": "Fix bug in handlers", - "description": "The proxy drops capabilities.", - "git_branch": "master" - } - } - }); - - let response = handler - .handle_request(req) - .await - .expect("Expected a response"); - assert_eq!(response["id"], 3); - - let content = &response["result"]["content"][0]; - assert_eq!(content["type"], "text"); - assert!( - content["text"] - .as_str() - .unwrap() - .starts_with("Task added with ID: ") - ); - - // Verify task was actually added to store - let tasks = handler.state.tasks.read(); - assert_eq!(tasks.len(), 1); - assert_eq!(tasks[0].title, "Fix bug in handlers"); - assert_eq!(tasks[0].status, "pending"); - } - - #[tokio::test] - async fn test_handle_create_entities() { - let handler = setup_test_handler("create_entities"); - - let req = json!({ - "jsonrpc": "2.0", - "id": 4, - "method": "tools/call", - "params": { - "name": "create_entities", - "arguments": { - "entities": [ - { - "name": "MemoryHandler", - "entityType": "struct", - "observations": ["Handles MCP requests natively"], - "namespace": "core" - } - ] - } - } - }); - - let response = handler - .handle_request(req) - .await - .expect("Expected a response"); - assert_eq!(response["id"], 4); - - let content = &response["result"]["content"][0]; - assert_eq!(content["text"], "Entities created"); - - // Verify entity was actually added to state - let session_graph = handler.state.graph.read(); - let entity = session_graph - .entities - .get("MemoryHandler") - .expect("Entity should be in session graph"); - assert_eq!(entity.entity_type, "struct"); - assert_eq!(entity.observations, vec!["Handles MCP requests natively"]); - assert_eq!(entity.namespace, "core".to_string()); - } - - #[tokio::test] - async fn test_handle_store_snippet() { - let handler = setup_test_handler("store_snippet"); - - let req = json!({ - "jsonrpc": "2.0", - "id": 5, - "method": "tools/call", - "params": { - "name": "store_snippet", - "arguments": { - "name": "Test Snippet", - "description": "A snippet used for testing", - "language": "rust", - "code": "fn main() { println!(\"Hello, World!\"); }" - } - } - }); - - let response = handler - .handle_request(req) - .await - .expect("Expected a response"); - assert_eq!(response["id"], 5); - - let snippets = handler.state.snippets.read(); - assert_eq!(snippets.len(), 1); - assert_eq!(snippets[0].name, "Test Snippet"); - assert_eq!(snippets[0].language, "rust"); - } - - #[tokio::test] - async fn test_handle_add_sticky_note() { - let handler = setup_test_handler("add_sticky_note"); - - let req = json!({ - "jsonrpc": "2.0", - "id": 6, - "method": "tools/call", - "params": { - "name": "add_sticky_note", - "arguments": { - "content": "Don't forget to check coverage!" - } - } - }); - - let response = handler - .handle_request(req) - .await - .expect("Expected a response"); - assert_eq!(response["id"], 6); - - let notes = handler.state.sticky.read(); - assert_eq!(notes.len(), 1); - assert_eq!(notes[0].content, "Don't forget to check coverage!"); - } - - #[tokio::test] - async fn test_handle_create_relations() { - let handler = setup_test_handler("create_relations"); - let req = json!({ - "jsonrpc": "2.0", - "id": 7, - "method": "tools/call", - "params": { - "name": "create_relations", - "arguments": { - "relations": [ - { - "from": "NodeA", - "to": "NodeB", - "relationType": "depends_on", - "namespace": "core" - } - ] - } - } - }); - let response = handler.handle_request(req).await.unwrap(); - assert_eq!(response["id"], 7); - let session = handler.state.graph.read(); - assert_eq!(session.relations.len(), 1); - assert_eq!(session.relations[0].from, "NodeA"); - assert_eq!(session.relations[0].to, "NodeB"); - } - - #[tokio::test] - async fn test_handle_add_observations() { - let handler = setup_test_handler("add_observations"); - // Pre-populate entity - handler.state.graph.modify(|session| { - session.entities.insert( - "NodeA".to_string(), - crate::models::Entity { - name: "NodeA".to_string(), - entity_type: "class".to_string(), - observations: vec!["Initial".to_string()], - namespace: "".to_string(), - git_branch: None, - }, - ); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 8, - "method": "tools/call", - "params": { - "name": "add_observations", - "arguments": { - "observations": [ - { - "entityName": "NodeA", - "contents": ["New observation"] - } - ] - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let session = handler.state.graph.read(); - let entity = session.entities.get("NodeA").unwrap(); - assert_eq!(entity.observations, vec!["Initial", "New observation"]); - } - - #[tokio::test] - async fn test_handle_delete_entities() { - let handler = setup_test_handler("delete_entities"); - handler.state.graph.modify(|session| { - session.entities.insert( - "ToDelete".to_string(), - crate::models::Entity { - name: "ToDelete".to_string(), - entity_type: "var".to_string(), - observations: vec![], - namespace: "".to_string(), - git_branch: None, - }, - ); - }); - // Force flush session to master - handler.state.modify_graph(|_| {}); - - let req = json!({ - "jsonrpc": "2.0", - "id": 9, - "method": "tools/call", - "params": { - "name": "delete_entities", - "arguments": { - "entityNames": ["ToDelete"] - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let full_graph = handler.state.get_full_graph(); - assert!(full_graph.entities.get("ToDelete").is_none()); - } - - #[tokio::test] - async fn test_handle_delete_observations() { - let handler = setup_test_handler("delete_observations"); - handler.state.graph.modify(|session| { - session.entities.insert( - "NodeA".to_string(), - crate::models::Entity { - name: "NodeA".to_string(), - entity_type: "class".to_string(), - observations: vec!["Keep".to_string(), "Drop".to_string()], - namespace: "".to_string(), - git_branch: None, - }, - ); - }); - handler.state.modify_graph(|_| {}); - let req = json!({ - "jsonrpc": "2.0", - "id": 10, - "method": "tools/call", - "params": { - "name": "delete_observations", - "arguments": { - "deletions": [ - { - "entityName": "NodeA", - "observations": ["Drop"] - } - ] - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let full = handler.state.get_full_graph(); - let entity = full.entities.get("NodeA").unwrap(); - assert_eq!(entity.observations, vec!["Keep"]); - } - - #[tokio::test] - async fn test_handle_log_code_change() { - let handler = setup_test_handler("log_code_change"); - let req = json!({ - "jsonrpc": "2.0", - "id": 11, - "method": "tools/call", - "params": { - "name": "log_code_change", - "arguments": { - "filePath": "server/src/handlers.rs", - "description": "Added some unit tests", - "git_commit": "1234567" - } - } - }); - let response = handler.handle_request(req).await.unwrap(); - assert_eq!(response["id"], 11); - let ledger = handler.state.ledger.read(); - assert_eq!(ledger.len(), 1); - assert_eq!(ledger[0].file_path, "server/src/handlers.rs"); - assert_eq!(ledger[0].git_commit.as_deref(), Some("1234567")); - } - - #[tokio::test] - async fn test_handle_list_active_tasks() { - let handler = setup_test_handler("list_active_tasks"); - handler.state.tasks.modify(|tasks| { - tasks.push(crate::models::Task { - id: "1".to_string(), - title: "Active Task".to_string(), - status: "pending".to_string(), - description: "".to_string(), - created_at: 0, - updated_at: 0, - git_branch: None, - acceptance_criteria: vec![], - dependencies: vec![], - parent_id: None, - }); - tasks.push(crate::models::Task { - id: "2".to_string(), - title: "Completed Task".to_string(), - status: "done".to_string(), - description: "".to_string(), - created_at: 0, - updated_at: 0, - git_branch: None, - acceptance_criteria: vec![], - dependencies: vec![], - parent_id: None, - }); - }); - - let req = json!({ - "jsonrpc": "2.0", - "id": 12, - "method": "tools/call", - "params": { - "name": "list_active_tasks", - "arguments": {} - } - }); - - let response = handler.handle_request(req).await.unwrap(); - let content = response["result"]["content"][0]["text"].as_str().unwrap(); - assert!(content.contains("Active Task")); - assert!(!content.contains("Completed Task")); - } - - #[tokio::test] - async fn test_handle_search_snippets() { - let handler = setup_test_handler("search_snippets"); - handler.state.snippets.modify(|snippets| { - snippets.push(crate::models::Snippet { - name: "React hook".to_string(), - language: "typescript".to_string(), - code: "useMemo(() => {}, [])".to_string(), - description: "React memoization".to_string(), - updated_at: 0, - }); - snippets.push(crate::models::Snippet { - name: "Rust struct".to_string(), - language: "rust".to_string(), - code: "struct A {}".to_string(), - description: "Rust code".to_string(), - updated_at: 0, - }); - }); - - let req = json!({ - "jsonrpc": "2.0", - "id": 13, - "method": "tools/call", - "params": { - "name": "search_snippets", - "arguments": { - "query": "React" - } - } - }); - let response = handler.handle_request(req).await.unwrap(); - let content = response["result"]["content"][0]["text"].as_str().unwrap(); - assert!(content.contains("React hook")); - assert!(!content.contains("Rust struct")); - } - - #[tokio::test] - async fn test_handle_read_sticky_notes() { - let handler = setup_test_handler("read_sticky_notes"); - handler.state.sticky.modify(|sticky| { - sticky.push(crate::models::StickyNote { - content: "Remember to commit".to_string(), - timestamp: 0, - }); - }); - - let req = json!({ - "jsonrpc": "2.0", - "id": 14, - "method": "tools/call", - "params": { - "name": "read_sticky_notes", - "arguments": {} - } - }); - let response = handler.handle_request(req).await.unwrap(); - let content = response["result"]["content"][0]["text"].as_str().unwrap(); - assert!(content.contains("Remember to commit")); - } - - #[tokio::test] - async fn test_handle_delete_relations() { - let handler = setup_test_handler("delete_relations"); - handler.state.graph.modify(|session| { - session.relations.push(crate::models::Relation { - from: "A".to_string(), - to: "B".to_string(), - relation_type: "calls".to_string(), - namespace: "".to_string(), - }); - }); - handler.state.modify_graph(|_| {}); - - let req = json!({ - "jsonrpc": "2.0", - "id": 15, - "method": "tools/call", - "params": { - "name": "delete_relations", - "arguments": { - "relations": [ - { - "from": "A", - "to": "B", - "relationType": "calls", - "namespace": "" - } - ] - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let full = handler.state.get_full_graph(); - assert!(full.relations.is_empty()); - } - - #[tokio::test] - async fn test_handle_read_graph() { - let handler = setup_test_handler("read_graph"); - handler.state.graph.modify(|session| { - session.entities.insert( - "NodeA".to_string(), - crate::models::Entity { - name: "NodeA".to_string(), - entity_type: "var".to_string(), - observations: vec![], - namespace: "".to_string(), - git_branch: None, - }, - ); - }); - handler.state.modify_graph(|_| {}); - - let req = json!({ - "jsonrpc": "2.0", - "id": 16, - "method": "tools/call", - "params": { - "name": "read_graph", - "arguments": {} - } - }); - let response = handler.handle_request(req).await.unwrap(); - let content = response["result"]["content"][0]["text"].as_str().unwrap(); - assert!(content.contains("NodeA")); - } - - #[tokio::test] - async fn test_handle_open_nodes() { - let handler = setup_test_handler("open_nodes"); - let entity = crate::models::Entity { - name: "UserRepository".to_string(), - entity_type: "class".to_string(), - observations: vec!["Handles user data".to_string()], - namespace: "".to_string(), - git_branch: None, - }; - handler.state.graph.modify(|session| { - session - .entities - .insert("UserRepository".to_string(), entity); - }); - handler.state.modify_graph(|_| {}); - - let req = json!({ - "jsonrpc": "2.0", - "id": 17, - "method": "tools/call", - "params": { - "name": "open_nodes", - "arguments": { - "names": ["UserRepository"] - } - } - }); - let response = handler.handle_request(req).await.unwrap(); - let content = response["result"]["content"][0]["text"].as_str().unwrap(); - assert!(content.contains("UserRepository")); - assert!(content.contains("Handles user data")); - } - - #[tokio::test] - async fn test_handle_log_decision() { - let handler = setup_test_handler("log_decision"); - let req = json!({ - "jsonrpc": "2.0", - "id": 20, - "method": "tools/call", - "params": { - "name": "log_decision", - "arguments": { - "title": "Use async I/O", - "context": "Need better throughput", - "decision": "Use tokio", - "consequence": "Requires async all the way down" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let adrs = handler.state.adrs.read(); - assert_eq!(adrs.len(), 1); - assert_eq!(adrs[0].title, "Use async I/O"); - } - - #[tokio::test] - async fn test_handle_query_decisions() { - let handler = setup_test_handler("query_decisions"); - handler.state.adrs.modify(|adrs| { - adrs.push(crate::models::Adr { - id: "adr-1".to_string(), - title: "Use PostgreSQL".to_string(), - context: "Need relational data".to_string(), - decision: "Use pg".to_string(), - consequence: "Maintenance overhead".to_string(), - timestamp: 0, - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 21, - "method": "tools/call", - "params": { - "name": "query_decisions", - "arguments": { - "query": "Postgre" - } - } - }); - let response = handler.handle_request(req).await.unwrap(); - let content = response["result"]["content"][0]["text"].as_str().unwrap(); - assert!(content.contains("PostgreSQL")); - } - - #[tokio::test] - async fn test_handle_log_error_fix() { - let handler = setup_test_handler("log_error_fix"); - let req = json!({ - "jsonrpc": "2.0", - "id": 22, - "method": "tools/call", - "params": { - "name": "log_error_fix", - "arguments": { - "signature": "IndexOutOfBounds", - "solution": "Check array length" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let fixes = handler.state.error_fixes.read(); - assert_eq!(fixes.len(), 1); - assert_eq!(fixes[0].signature, "IndexOutOfBounds"); - } - - #[tokio::test] - async fn test_handle_search_error_fixes() { - let handler = setup_test_handler("search_error_fixes"); - handler.state.error_fixes.modify(|fixes| { - fixes.push(crate::models::ErrorFix { - signature: "NullPointerException".to_string(), - solution: "Initialize the pointer".to_string(), - timestamp: 0, - git_branch: None, - git_commit: None, - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 23, - "method": "tools/call", - "params": { - "name": "search_error_fixes", - "arguments": { - "query": "NullPointer" - } - } - }); - let response = handler.handle_request(req).await.unwrap(); - let content = response["result"]["content"][0]["text"].as_str().unwrap(); - assert!(content.contains("Initialize the pointer")); - } - - #[tokio::test] - async fn test_handle_list_pinned_files() { - let handler = setup_test_handler("list_pinned_files"); - handler.state.pinned_files.modify(|files| { - files.push(crate::models::PinnedFile { - file_path: "src/important.rs".to_string(), - timestamp: 0, - namespace: "".to_string(), - git_branch: None, - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 24, - "method": "tools/call", - "params": { - "name": "list_pinned_files", - "arguments": {} - } - }); - let response = handler.handle_request(req).await.unwrap(); - let content = response["result"]["content"][0]["text"].as_str().unwrap(); - assert!(content.contains("important.rs")); - } - - #[tokio::test] - async fn test_handle_add_session_summary() { - let handler = setup_test_handler("add_session_summary"); - let req = json!({ - "jsonrpc": "2.0", - "id": 25, - "method": "tools/call", - "params": { - "name": "add_session_summary", - "arguments": { - "namespace": "", - "summary": "Finished writing tests" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let summaries = handler.state.session_summaries.read(); - assert_eq!(summaries.len(), 1); - assert_eq!(summaries[0].summary, "Finished writing tests"); - } - - #[tokio::test] - async fn test_handle_get_project_timeline() { - let handler = setup_test_handler("get_project_timeline"); - handler.state.session_summaries.modify(|summaries| { - summaries.push(crate::models::SessionSummary { - summary: "Day 1: Setup project".to_string(), - namespace: "".to_string(), - timestamp: 0, - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 26, - "method": "tools/call", - "params": { - "name": "get_project_timeline", - "arguments": { - "namespace": "" - } - } - }); - let response = handler.handle_request(req).await.unwrap(); - let content = response["result"]["content"][0]["text"].as_str().unwrap(); - assert!(content.contains("Day 1: Setup project")); - } - - #[tokio::test] - async fn test_handle_log_tech_debt() { - let handler = setup_test_handler("log_tech_debt"); - let req = json!({ - "jsonrpc": "2.0", - "id": 27, - "method": "tools/call", - "params": { - "name": "log_tech_debt", - "arguments": { - "namespace": "", - "description": "Hardcoded values", - "ideal_solution": "Remove magic numbers" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let debt = handler.state.tech_debts.read(); - assert_eq!(debt.len(), 1); - assert_eq!(debt[0].description, "Hardcoded values"); - } - - #[tokio::test] - async fn test_handle_list_tech_debt() { - let handler = setup_test_handler("list_tech_debt"); - handler.state.tech_debts.modify(|debts| { - debts.push(crate::models::TechDebt { - id: "debt-1".to_string(), - description: "Bad naming".to_string(), - ideal_solution: "Rename x to num_elements".to_string(), - namespace: "".to_string(), - is_resolved: false, - created_at: 0, - git_branch: None, - git_commit: None, - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 28, - "method": "tools/call", - "params": { - "name": "list_tech_debt", - "arguments": { - "namespace": "", - "include_resolved": false - } - } - }); - let response = handler.handle_request(req).await.unwrap(); - let content = response["result"]["content"][0]["text"].as_str().unwrap(); - assert!(content.contains("Bad naming")); - } - - #[tokio::test] - async fn test_handle_get_project_health() { - let handler = setup_test_handler("get_project_health"); - let req = json!({ - "jsonrpc": "2.0", - "id": 29, - "method": "tools/call", - "params": { - "name": "get_project_health", - "arguments": { - "namespace": "" - } - } - }); - let response = handler.handle_request(req).await.unwrap(); - let content = response["result"]["content"][0]["text"].as_str().unwrap(); - assert!(content.contains("\"active_tasks\"")); - assert!(content.contains("\"unresolved_tech_debt\"")); - } - - #[tokio::test] - async fn test_handle_resolve_tech_debt() { - let handler = setup_test_handler("resolve_tech_debt"); - handler.state.tech_debts.modify(|debts| { - debts.push(crate::models::TechDebt { - id: "debt-2".to_string(), - description: "Old api".to_string(), - ideal_solution: "Use new api".to_string(), - namespace: "".to_string(), - is_resolved: false, - created_at: 0, - git_branch: None, - git_commit: None, - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 30, - "method": "tools/call", - "params": { - "name": "resolve_tech_debt", - "arguments": { - "id": "debt-2" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let debts = handler.state.tech_debts.read(); - assert!(debts[0].is_resolved); - } - - #[tokio::test] - async fn test_handle_leave_handoff_memo() { - let handler = setup_test_handler("leave_handoff_memo"); - let req = json!({ - "jsonrpc": "2.0", - "id": 31, - "method": "tools/call", - "params": { - "name": "leave_handoff_memo", - "arguments": { - "namespace": "", - "content": "Make sure to check the logs.", - "author": "Riz" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let memos = handler.state.handoff_memos.read(); - assert_eq!(memos.len(), 1); - assert_eq!(memos[0].content, "Make sure to check the logs."); - } - - #[tokio::test] - async fn test_handle_query_recent_changes() { - let handler = setup_test_handler("query_recent_changes"); - handler.state.ledger.modify(|ledger| { - ledger.push(crate::models::CodeChange { - timestamp: 0, - file_path: "src/main.rs".to_string(), - description: "Fix bug".to_string(), - git_commit: None, - git_branch: None, - }); - }); - - let req = json!({ - "jsonrpc": "2.0", - "id": 18, - "method": "tools/call", - "params": { - "name": "query_recent_changes", - "arguments": {} - } - }); - let response = handler.handle_request(req).await.unwrap(); - let content = response["result"]["content"][0]["text"].as_str().unwrap(); - assert!(content.contains("Fix bug")); - assert!(content.contains("src/main.rs")); - } - - #[tokio::test] - async fn test_handle_update_task_status() { - let handler = setup_test_handler("update_task_status"); - handler.state.tasks.modify(|tasks| { - tasks.push(crate::models::Task { - id: "test-task-123".to_string(), - title: "In progress task".to_string(), - status: "pending".to_string(), - description: "".to_string(), - created_at: 0, - updated_at: 0, - git_branch: None, - acceptance_criteria: vec![], - dependencies: vec![], - parent_id: None, - }); - }); - - let req = json!({ - "jsonrpc": "2.0", - "id": 19, - "method": "tools/call", - "params": { - "name": "update_task_status", - "arguments": { - "id": "test-task-123", - "status": "in_progress" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let tasks = handler.state.tasks.read(); - assert_eq!(tasks[0].status, "in_progress"); - } - - #[tokio::test] - async fn test_handle_pin_file() { - let handler = setup_test_handler("pin_file"); - let req = json!({ - "jsonrpc": "2.0", - "id": 100, - "method": "tools/call", - "params": { - "name": "pin_file", - "arguments": { - "file_path": "/path/to/pinned.rs", - "namespace": "global" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let pinned = handler.state.pinned_files.read(); - assert!(pinned.iter().any(|p| p.file_path == "/path/to/pinned.rs")); - } - - #[tokio::test] - async fn test_handle_unpin_file() { - let handler = setup_test_handler("unpin_file"); - handler.state.pinned_files.modify(|files| { - files.push(crate::models::PinnedFile { - namespace: "global".to_string(), - file_path: "/path/to/unpin.rs".to_string(), - timestamp: 0, - git_branch: None, - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 101, - "method": "tools/call", - "params": { - "name": "unpin_file", - "arguments": { - "file_path": "/path/to/unpin.rs", - "namespace": "global" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let pinned = handler.state.pinned_files.read(); - assert!(!pinned.iter().any(|p| p.file_path == "/path/to/unpin.rs")); - } - - #[tokio::test] - async fn test_handle_learn_preference() { - let handler = setup_test_handler("learn_preference"); - let req = json!({ - "jsonrpc": "2.0", - "id": 102, - "method": "tools/call", - "params": { - "name": "learn_preference", - "arguments": { - "key": "formatting", - "value": "use spaces" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let prefs = handler.state.prefs.read(); - assert!(prefs.values().any(|p| p.key == "formatting" && p.value == "use spaces")); - } - - #[tokio::test] - async fn test_handle_read_preferences() { - let handler = setup_test_handler("read_preferences"); - handler.state.prefs.modify(|prefs| { - prefs.insert("theme".to_string(), crate::models::Preference { - key: "theme".to_string(), - value: "dark".to_string(), - updated_at: 0, - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 103, - "method": "tools/call", - "params": { - "name": "read_preferences", - "arguments": {} - } - }); - let response = handler.handle_request(req).await.unwrap(); - let content = response["result"]["content"][0]["text"].as_str().unwrap(); - assert!(content.contains("theme")); - assert!(content.contains("dark")); - } - - #[tokio::test] - async fn test_handle_delete_task() { - let handler = setup_test_handler("delete_task"); - handler.state.tasks.modify(|tasks| { - tasks.push(crate::models::Task { - id: "task-to-delete".to_string(), - title: "".to_string(), - status: "".to_string(), - description: "".to_string(), - created_at: 0, - updated_at: 0, - git_branch: None, - acceptance_criteria: vec![], - dependencies: vec![], - parent_id: None, - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 104, - "method": "tools/call", - "params": { - "name": "delete_task", - "arguments": { - "id": "task-to-delete" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let tasks = handler.state.tasks.read(); - assert!(tasks.is_empty()); - } - - #[tokio::test] - async fn test_handle_add_milestone() { - let handler = setup_test_handler("add_milestone"); - let req = json!({ - "jsonrpc": "2.0", - "id": 105, - "method": "tools/call", - "params": { - "name": "add_milestone", - "arguments": { - "title": "v1.0", - "namespace": "global" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let milestones = handler.state.milestones.read(); - assert_eq!(milestones.len(), 1); - assert_eq!(milestones[0].title, "v1.0"); - } - - #[tokio::test] - async fn test_handle_update_milestone() { - let handler = setup_test_handler("update_milestone"); - handler.state.milestones.modify(|ms| { - ms.push(crate::models::Milestone { - id: "v1.0".to_string(), - title: "v1.0".to_string(), - status: "pending".to_string(), - namespace: "".to_string(), - target_date: None, - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 106, - "method": "tools/call", - "params": { - "name": "update_milestone", - "arguments": { - "id": "v1.0", - "status": "completed" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let ms = handler.state.milestones.read(); - assert_eq!(ms[0].status, "completed"); - } - - #[tokio::test] - async fn test_handle_list_milestones() { - let handler = setup_test_handler("list_milestones"); - handler.state.milestones.modify(|ms| { - ms.push(crate::models::Milestone { - id: "v2.0".to_string(), - title: "v2.0".to_string(), - status: "pending".to_string(), - namespace: "".to_string(), - target_date: None, - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 107, - "method": "tools/call", - "params": { - "name": "list_milestones", - "arguments": {} - } - }); - let response = handler.handle_request(req).await.unwrap(); - let content = response["result"]["content"][0]["text"].as_str().unwrap(); - assert!(content.contains("v2.0")); - } - - #[tokio::test] - async fn test_handle_add_pr_checklist_item() { - let handler = setup_test_handler("add_pr_checklist_item"); - let req = json!({ - "jsonrpc": "2.0", - "id": 108, - "method": "tools/call", - "params": { - "name": "add_pr_checklist_item", - "arguments": { - "description": "Check tests", - "namespace": "global" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let cl = handler.state.pr_checklists.read(); - assert_eq!(cl.len(), 1); - assert_eq!(cl[0].description, "Check tests"); - } - - #[tokio::test] - async fn test_handle_clear_pr_checklist() { - let handler = setup_test_handler("clear_pr_checklist"); - handler.state.pr_checklists.modify(|cl| { - cl.push(crate::models::PrChecklistItem { - namespace: "global".to_string(), - id: "item".to_string(), - description: "".to_string(), - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 109, - "method": "tools/call", - "params": { - "name": "clear_pr_checklist", - "arguments": { - "namespace": "global" - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let cl = handler.state.pr_checklists.read(); - assert!(cl.is_empty()); - } - - #[tokio::test] - async fn test_handle_clear_handoff_memos() { - let handler = setup_test_handler("clear_handoff_memos"); - handler.state.handoff_memos.modify(|ms| { - ms.push(crate::models::HandoffMemo { - id: "memo".to_string(), - namespace: "".to_string(), - content: "".to_string(), - author: "".to_string(), - timestamp: 0, - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 110, - "method": "tools/call", - "params": { - "name": "clear_handoff_memos", - "arguments": { - "ids": ["memo"] - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let ms = handler.state.handoff_memos.read(); - assert!(ms.is_empty()); - } - - #[tokio::test] - async fn test_handle_clear_sticky_notes() { - let handler = setup_test_handler("clear_sticky_notes"); - handler.state.sticky.modify(|ns| { - ns.push(crate::models::StickyNote { - content: "".to_string(), - timestamp: 0, - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 111, - "method": "tools/call", - "params": { - "name": "clear_sticky_notes", - "arguments": {} - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let ns = handler.state.sticky.read(); - assert!(ns.is_empty()); - } - - #[tokio::test] - async fn test_handle_delete_sticky_note() { - let handler = setup_test_handler("delete_sticky_note"); - handler.state.sticky.modify(|ns| { - ns.push(crate::models::StickyNote { - content: "note1".to_string(), - timestamp: 0, - }); - ns.push(crate::models::StickyNote { - content: "note2".to_string(), - timestamp: 0, - }); - }); - let req = json!({ - "jsonrpc": "2.0", - "id": 112, - "method": "tools/call", - "params": { - "name": "delete_sticky_note", - "arguments": { - "index": 1 - } - } - }); - let _ = handler.handle_request(req).await.unwrap(); - let ns = handler.state.sticky.read(); - assert_eq!(ns.len(), 1); - assert_eq!(ns[0].content, "note2"); } } diff --git a/server/src/handlers_v2/env.rs b/server/src/handlers_v2/env.rs new file mode 100644 index 0000000..20e232b --- /dev/null +++ b/server/src/handlers_v2/env.rs @@ -0,0 +1,163 @@ +use crate::router::McpTool; +use crate::state::MemoryState; +use crate::tools::*; +use async_trait::async_trait; +use serde_json::Value; +use std::sync::Arc; +use std::time::{SystemTime, UNIX_EPOCH}; + +pub struct UpdateEnvFingerprintHandler; + +#[async_trait] +impl McpTool for UpdateEnvFingerprintHandler { + fn name(&self) -> &'static str { + "update_env_fingerprint" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "update_env_fingerprint", + "Execute update_env_fingerprint", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: UpdateEnvFingerprintTool = + serde_json::from_value(args).map_err(|e| e.to_string())?; + state.env_fingerprints.modify(|fps| { + fps.insert( + req.namespace.clone(), + crate::models::EnvFingerprint { + namespace: req.namespace.clone(), + os: std::env::consts::OS.to_string(), + shell: std::env::var("SHELL").unwrap_or_else(|_| "unknown".to_string()), + tool_versions: req.tool_versions, + updated_at: SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + }, + ); + }); + Ok("Env fingerprint updated".to_string()) + } +} + +pub struct ReadEnvFingerprintHandler; + +#[async_trait] +impl McpTool for ReadEnvFingerprintHandler { + fn name(&self) -> &'static str { + "read_env_fingerprint" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "read_env_fingerprint", + "Execute read_env_fingerprint", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: ReadEnvFingerprintTool = + serde_json::from_value(args).map_err(|e| e.to_string())?; + let fps = state.env_fingerprints.read(); + if let Some(fp) = fps.get(&req.namespace) { + let data = serde_json::to_string(fp).unwrap_or_default(); + Ok(data.to_string()) + } else { + Ok("{}".to_string()) + } + } +} + +pub struct LogEnvRequirementHandler; + +#[async_trait] +impl McpTool for LogEnvRequirementHandler { + fn name(&self) -> &'static str { + "log_env_requirement" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "log_env_requirement", + "Execute log_env_requirement", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: LogEnvRequirementTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.env_requirements.modify(|reqs| { + reqs.retain(|r| !(r.namespace == req.namespace && r.key == req.key)); + reqs.push(crate::models::EnvRequirement { + namespace: req.namespace, + key: req.key, + description: req.description, + is_secret: req.is_secret, + }); + }); + Ok("Env requirement logged".to_string()) + } +} + +pub struct RegisterEnvironmentHandler; + +#[async_trait] +impl McpTool for RegisterEnvironmentHandler { + fn name(&self) -> &'static str { + "register_environment" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "register_environment", + "Execute register_environment", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: RegisterEnvironmentTool = + serde_json::from_value(args).map_err(|e| e.to_string())?; + state.environments.modify(|envs| { + envs.retain(|e| !(e.namespace == req.namespace && e.name == req.name)); + envs.push(crate::models::EnvironmentDetail { + namespace: req.namespace, + name: req.name, + url: req.url, + description: req.description, + requires_vpn: req.requires_vpn, + updated_at: SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + }); + }); + Ok("Environment registered".to_string()) + } +} + +pub struct GetEnvironmentDetailsHandler; + +#[async_trait] +impl McpTool for GetEnvironmentDetailsHandler { + fn name(&self) -> &'static str { + "get_environment_details" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "get_environment_details", + "Execute get_environment_details", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: GetEnvironmentDetailsTool = + serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut envs = state.environments.read(); + envs.retain(|e| e.namespace == req.namespace); + let data = serde_json::to_string(&envs).unwrap_or_default(); + Ok(data.to_string()) + } +} diff --git a/server/src/handlers_v2/graph.rs b/server/src/handlers_v2/graph.rs new file mode 100644 index 0000000..fa1c3dd --- /dev/null +++ b/server/src/handlers_v2/graph.rs @@ -0,0 +1,544 @@ +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::collections::HashSet; +use std::sync::Arc; + +pub struct QueryGraphPathHandler; + +#[async_trait] +impl McpTool for QueryGraphPathHandler { + fn name(&self) -> &'static str { + "query_graph_path" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("query_graph_path", "Execute query_graph_path") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: crate::tools::QueryGraphPathTool = + serde_json::from_value(args).map_err(|e| e.to_string())?; + state.read_graph(|graph| { + let max_depth = req.max_depth.unwrap_or(5); + let mut queue = std::collections::VecDeque::new(); + let mut visited = std::collections::HashSet::new(); + let mut parents: std::collections::HashMap = + std::collections::HashMap::new(); + + queue.push_back(req.start_node.clone()); + visited.insert(req.start_node.clone()); + + let mut found = false; + let mut current_depth = 0; + let mut nodes_at_current_depth = 1; + let mut nodes_at_next_depth = 0; + + while let Some(current) = queue.pop_front() { + if current == req.end_node { + found = true; + break; + } + nodes_at_current_depth -= 1; + if current_depth < max_depth { + for rel in &graph.relations { + if rel.from == current && !visited.contains(&rel.to) { + visited.insert(rel.to.clone()); + parents.insert( + rel.to.clone(), + (current.clone(), rel.relation_type.clone()), + ); + queue.push_back(rel.to.clone()); + nodes_at_next_depth += 1; + } else if rel.to == current && !visited.contains(&rel.from) { + visited.insert(rel.from.clone()); + parents.insert( + rel.from.clone(), + (current.clone(), format!("inverse({})", rel.relation_type)), + ); + queue.push_back(rel.from.clone()); + nodes_at_next_depth += 1; + } + } + } + if nodes_at_current_depth == 0 { + current_depth += 1; + nodes_at_current_depth = nodes_at_next_depth; + nodes_at_next_depth = 0; + } + } + + if found { + let mut path = Vec::new(); + let mut curr = req.end_node.clone(); + while curr != req.start_node { + if let Some((parent, rel_type)) = parents.get(&curr) { + path.push(format!("{} -[{}]-> {}", parent, rel_type, curr)); + curr = parent.clone(); + } else { + break; + } + } + path.reverse(); + Ok(format!("Path found:\n{}", path.join("\n"))) + } else { + Ok(format!( + "No path found between {} and {} within depth {}", + req.start_node, req.end_node, max_depth + )) + } + }) + } +} + +pub struct CreateEntitiesHandler; + +#[async_trait] +impl McpTool for CreateEntitiesHandler { + fn name(&self) -> &'static str { + "create_entities" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("create_entities", "Execute create_entities") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: CreateEntitiesTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.modify_graph(|g| { + for entity in req.entities { + if !entity.name.is_empty() { + if let Ok(idx) = state.search_index.read() { + drop(idx.index_entity(&entity)); + } + g.entities.insert(entity.name.clone(), entity); + } + } + }); + Ok("Entities created".to_string()) + } +} + +pub struct CreateRelationsHandler; + +#[async_trait] +impl McpTool for CreateRelationsHandler { + fn name(&self) -> &'static str { + "create_relations" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("create_relations", "Execute create_relations") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: CreateRelationsTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.modify_graph(|g| { + for relation in req.relations { + if !relation.from.is_empty() && !relation.to.is_empty() { + g.relations.push(relation); + } + } + }); + Ok("Relations created".to_string()) + } +} + +pub struct AddObservationsHandler; + +#[async_trait] +impl McpTool for AddObservationsHandler { + fn name(&self) -> &'static str { + "add_observations" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("add_observations", "Execute add_observations") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: AddObservationsTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.modify_graph(|g| { + for o in req.observations { + if let Some(e) = g.entities.get_mut(&o.entity_name) { + e.observations.extend(o.contents); + } + } + }); + Ok("Observations added".to_string()) + } +} + +pub struct DeleteEntitiesHandler; + +#[async_trait] +impl McpTool for DeleteEntitiesHandler { + fn name(&self) -> &'static str { + "delete_entities" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("delete_entities", "Execute delete_entities") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: DeleteEntitiesTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let to_delete: HashSet<_> = req.entity_names.into_iter().collect(); + state.modify_graph(|master| { + for name in &to_delete { + master.entities.remove(name); + } + master + .relations + .retain(|r| !to_delete.contains(&r.from) && !to_delete.contains(&r.to)); + }); + Ok("Entities deleted".to_string()) + } +} + +pub struct DeleteObservationsHandler; + +#[async_trait] +impl McpTool for DeleteObservationsHandler { + fn name(&self) -> &'static str { + "delete_observations" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "delete_observations", + "Execute delete_observations", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: DeleteObservationsTool = + serde_json::from_value(args).map_err(|e| e.to_string())?; + state.modify_graph(|master| { + for d in req.deletions { + if let Some(e) = master.entities.get_mut(&d.entity_name) { + let to_rem: HashSet<_> = d.observations.into_iter().collect(); + e.observations.retain(|o| !to_rem.contains(o)); + } + } + }); + Ok("Observations deleted".to_string()) + } +} + +pub struct DeleteRelationsHandler; + +#[async_trait] +impl McpTool for DeleteRelationsHandler { + fn name(&self) -> &'static str { + "delete_relations" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("delete_relations", "Execute delete_relations") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: DeleteRelationsTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.modify_graph(|master| { + let to_rem: HashSet<_> = req.relations.into_iter().collect(); + master.relations.retain(|r| !to_rem.contains(r)); + }); + Ok("Relations deleted".to_string()) + } +} + +pub struct ReadGraphHandler; + +#[async_trait] +impl McpTool for ReadGraphHandler { + fn name(&self) -> &'static str { + "read_graph" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("read_graph", "Execute read_graph") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: ReadGraphTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let data = state.read_graph(|full| { + if let Some(ns) = req.namespace { + let mut filtered = KnowledgeGraph::default(); + for (k, v) in &full.entities { + if v.namespace == ns { + filtered.entities.insert(k.clone(), v.clone()); + } + } + for r in &full.relations { + if r.namespace == ns { + filtered.relations.push(r.clone()); + } + } + serde_json::to_string(&filtered).unwrap_or_default() + } else { + serde_json::to_string(full).unwrap_or_default() + } + }); + Ok(data) + } +} + +pub struct SearchNodesHandler; + +#[async_trait] +impl McpTool for SearchNodesHandler { + fn name(&self) -> &'static str { + "search_nodes" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("search_nodes", "Execute search_nodes") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: SearchNodesTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let matches = if let Ok(idx) = state.search_index.read() { + idx.search(&req.query, req.namespace.as_deref()) + .unwrap_or_default() + } else { + vec![] + }; + + let mut result = KnowledgeGraph::default(); + state.read_graph(|full| { + for (id, doc_type, _, _, _) in matches { + if doc_type == "entity" + && let Some(e) = full.entities.get(&id) + { + result.entities.insert(id, e.clone()); + } + } + }); + let data = serde_json::to_string(&result).unwrap_or_default(); + Ok(data.to_string()) + } +} + +pub struct OpenNodesHandler; + +#[async_trait] +impl McpTool for OpenNodesHandler { + fn name(&self) -> &'static str { + "open_nodes" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("open_nodes", "Execute open_nodes") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: OpenNodesTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let targets: HashSet<_> = req.names.into_iter().collect(); + let mut result = KnowledgeGraph::default(); + let mut connected = HashSet::new(); + state.read_graph(|full| { + for r in &full.relations { + if targets.contains(&r.from) { + connected.insert(r.to.clone()); + result.relations.push(r.clone()); + } else if targets.contains(&r.to) { + connected.insert(r.from.clone()); + result.relations.push(r.clone()); + } + } + for (name, e) in &full.entities { + if targets.contains(name) || connected.contains(name) { + result.entities.insert(name.clone(), e.clone()); + } + } + }); + let data = serde_json::to_string(&result).unwrap_or_default(); + Ok(data.to_string()) + } +} + +pub struct VisualizeGraphHandler; + +#[async_trait] +impl McpTool for VisualizeGraphHandler { + fn name(&self) -> &'static str { + "visualize_graph" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("visualize_graph", "Execute visualize_graph") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: VisualizeGraphTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let query = req.query.unwrap_or_default().to_lowercase(); + let mut included = HashSet::new(); + let mut to_draw = Vec::new(); + + state.read_graph(|full| { + for (name, e) in &full.entities { + if let Some(ns) = &req.namespace + && e.namespace != *ns + { + continue; + } + if query.is_empty() + || contains_ignore_ascii_case(name, &query) + || contains_ignore_ascii_case(&e.entity_type, &query) + { + included.insert(name.clone()); + } + } + + for r in &full.relations { + if let Some(ns) = &req.namespace + && r.namespace != *ns + { + continue; + } + if query.is_empty() || included.contains(&r.from) || included.contains(&r.to) { + included.insert(r.from.clone()); + included.insert(r.to.clone()); + to_draw.push(r.clone()); + } + } + }); + use std::fmt::Write; + let mut output = String::with_capacity(included.len() * 40 + to_draw.len() * 60); + output.push_str("graph TD;\n"); + + let sanitize = |s: &str, id_mode: bool| -> String { + let mut out = String::with_capacity(s.len()); + for c in s.chars() { + if c != '"' && c != '(' && c != ')' { + if id_mode && (c == ' ' || c == '-' || c == '.') { + out.push('_'); + } else { + out.push(c); + } + } + } + out + }; + + for name in &included { + let _ = writeln!( + output, + " id_{}[\"{}\"];", + sanitize(name, true), + sanitize(name, false) + ); + } + for r in to_draw { + let _ = writeln!( + output, + " id_{}-->|\"{}\"|id_{};", + sanitize(&r.from, true), + r.relation_type.replace("\"", ""), + sanitize(&r.to, true) + ); + } + if output == "graph TD;\n" { + output = "No nodes found to visualize.".to_string(); + } + Ok(output.to_string()) + } +} + +pub struct CondenseEntityHandler; + +#[async_trait] +impl McpTool for CondenseEntityHandler { + fn name(&self) -> &'static str { + "condense_entity" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("condense_entity", "Execute condense_entity") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: CondenseEntityTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.modify_graph(|master| { + if let Some(e) = master.entities.get_mut(&req.entity_name) { + e.observations = req.summarized_observations; + } + }); + Ok("Entity condensed".to_string()) + } +} + +pub struct MergeEntitiesHandler; + +#[async_trait] +impl McpTool for MergeEntitiesHandler { + fn name(&self) -> &'static str { + "merge_entities" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("merge_entities", "Execute merge_entities") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: MergeEntitiesTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.modify_graph(|master| { + if let Some(src) = master.entities.remove(&req.source_entity) { + if let Some(tgt) = master.entities.get_mut(&req.target_entity) { + tgt.observations.extend(src.observations); + MemoryState::deduplicate(&mut tgt.observations); + } else { + let mut new_tgt = src.clone(); + new_tgt.name = req.target_entity.clone(); + master.entities.insert(req.target_entity.clone(), new_tgt); + } + } + for r in &mut master.relations { + if r.from == req.source_entity { + r.from = req.target_entity.clone(); + } + if r.to == req.source_entity { + r.to = req.target_entity.clone(); + } + } + MemoryState::deduplicate(&mut master.relations); + }); + Ok("Entities merged".to_string()) + } +} + +pub struct FindOrphansHandler; + +#[async_trait] +impl McpTool for FindOrphansHandler { + fn name(&self) -> &'static str { + "find_orphans" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("find_orphans", "Execute find_orphans") + } + + async fn execute(&self, _args: Value, state: Arc) -> Result { + let orphans = state.read_graph(|full| { + let mut connected = std::collections::HashSet::new(); + for r in &full.relations { + connected.insert(r.from.clone()); + connected.insert(r.to.clone()); + } + full.entities + .keys() + .filter(|k| !connected.contains(*k)) + .cloned() + .collect::>() + }); + let data = serde_json::to_string(&orphans).unwrap_or_default(); + Ok(data.to_string()) + } +} + +use crate::handlers_v2::utils::*; diff --git a/server/src/handlers_v2/meta.rs b/server/src/handlers_v2/meta.rs new file mode 100644 index 0000000..91b4606 --- /dev/null +++ b/server/src/handlers_v2/meta.rs @@ -0,0 +1,492 @@ +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; +use std::time::{SystemTime, UNIX_EPOCH}; + +pub struct LogDecisionHandler; + +#[async_trait] +impl McpTool for LogDecisionHandler { + fn name(&self) -> &'static str { + "log_decision" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("log_decision", "Execute log_decision") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: LogDecisionTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut adr_id = String::new(); + let mut new_adr = None; + + state.adrs.modify(|adrs| { + adr_id = format!("ADR-{:04}", adrs.len() + 1); + let a = Adr { + id: adr_id.clone(), + title: req.title, + context: req.context, + decision: req.decision, + consequence: req.consequence, + timestamp: SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + }; + new_adr = Some(a.clone()); + adrs.push(a); + }); + + if let Some(adr) = new_adr + && let Ok(idx) = state.search_index.read() + { + drop(idx.index_adr(&adr)); + } + + Ok(format!("Decision logged as {}", adr_id).to_string()) + } +} + +pub struct QueryDecisionsHandler; + +#[async_trait] +impl McpTool for QueryDecisionsHandler { + fn name(&self) -> &'static str { + "query_decisions" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("query_decisions", "Execute query_decisions") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: QueryDecisionsTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut adrs = state.adrs.read(); + if let Some(q) = req.query { + let q = q.to_lowercase(); + adrs.retain(|a| { + contains_ignore_ascii_case(&a.title, &q) + || contains_ignore_ascii_case(&a.context, &q) + || contains_ignore_ascii_case(&a.decision, &q) + }); + } + let data = serde_json::to_string(&adrs).unwrap_or_default(); + Ok(data.to_string()) + } +} + +pub struct LogErrorFixHandler; + +#[async_trait] +impl McpTool for LogErrorFixHandler { + fn name(&self) -> &'static str { + "log_error_fix" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("log_error_fix", "Execute log_error_fix") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: LogErrorFixTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.error_fixes.modify(|fixes| { + fixes.push(crate::models::ErrorFix { + signature: req.signature, + solution: req.solution, + timestamp: SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + git_commit: req.git_commit, + git_branch: req.git_branch, + }) + }); + Ok("Error fix logged".to_string()) + } +} + +pub struct SearchErrorFixesHandler; + +#[async_trait] +impl McpTool for SearchErrorFixesHandler { + fn name(&self) -> &'static str { + "search_error_fixes" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "search_error_fixes", + "Execute search_error_fixes", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: SearchErrorFixesTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let q = req.query.to_lowercase(); + let mut fixes = state.error_fixes.read(); + fixes.retain(|f| { + contains_ignore_ascii_case(&f.signature, &q) + || contains_ignore_ascii_case(&f.solution, &q) + }); + let data = serde_json::to_string(&fixes).unwrap_or_default(); + Ok(data.to_string()) + } +} + +pub struct LogCodeChangeHandler; + +#[async_trait] +impl McpTool for LogCodeChangeHandler { + fn name(&self) -> &'static str { + "log_code_change" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("log_code_change", "Execute log_code_change") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: LogCodeChangeTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.ledger.modify(|ledger| { + ledger.push(CodeChange { + timestamp: SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + file_path: req.file_path, + description: req.description, + git_commit: req.git_commit, + git_branch: req.git_branch, + }); + }); + Ok("Code change logged".to_string()) + } +} + +pub struct QueryRecentChangesHandler; + +#[async_trait] +impl McpTool for QueryRecentChangesHandler { + fn name(&self) -> &'static str { + "query_recent_changes" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "query_recent_changes", + "Execute query_recent_changes", + ) + } + + async fn execute(&self, _args: Value, state: Arc) -> Result { + let data = serde_json::to_string(&state.ledger.read()).unwrap_or_else(|_| "[]".to_string()); + Ok(data.to_string()) + } +} + +pub struct LearnPreferenceHandler; + +#[async_trait] +impl McpTool for LearnPreferenceHandler { + fn name(&self) -> &'static str { + "learn_preference" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("learn_preference", "Execute learn_preference") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: LearnPreferenceTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.prefs.modify(|prefs| { + prefs.insert( + req.key.clone(), + crate::models::Preference { + key: req.key.clone(), + value: req.value, + updated_at: SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + }, + ); + }); + Ok("Preference learned".to_string()) + } +} + +pub struct ReadPreferencesHandler; + +#[async_trait] +impl McpTool for ReadPreferencesHandler { + fn name(&self) -> &'static str { + "read_preferences" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("read_preferences", "Execute read_preferences") + } + + async fn execute(&self, _args: Value, state: Arc) -> Result { + let prefs = state.prefs.read(); + let data = serde_json::to_string(&prefs).unwrap_or_default(); + Ok(data.to_string()) + } +} + +pub struct LogTechDebtHandler; + +#[async_trait] +impl McpTool for LogTechDebtHandler { + fn name(&self) -> &'static str { + "log_tech_debt" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("log_tech_debt", "Execute log_tech_debt") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: LogTechDebtTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.tech_debts.modify(|debts| { + debts.push(crate::models::TechDebt { + id: uuid::Uuid::new_v4().to_string(), + namespace: req.namespace, + description: req.description, + ideal_solution: req.ideal_solution, + is_resolved: false, + created_at: SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + git_commit: req.git_commit, + git_branch: req.git_branch, + }) + }); + Ok("Tech debt logged".to_string()) + } +} + +pub struct ResolveTechDebtHandler; + +#[async_trait] +impl McpTool for ResolveTechDebtHandler { + fn name(&self) -> &'static str { + "resolve_tech_debt" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "resolve_tech_debt", + "Execute resolve_tech_debt", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: ResolveTechDebtTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut found = false; + state.tech_debts.modify(|debts| { + for d in debts.iter_mut() { + if d.id == req.id { + d.is_resolved = true; + found = true; + break; + } + } + }); + if found { + Ok("Tech debt resolved".to_string()) + } else { + Ok("Tech debt not found".to_string()) + } + } +} + +pub struct ListTechDebtHandler; + +#[async_trait] +impl McpTool for ListTechDebtHandler { + fn name(&self) -> &'static str { + "list_tech_debt" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("list_tech_debt", "Execute list_tech_debt") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: ListTechDebtTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut debts = state.tech_debts.read(); + debts.retain(|d| d.namespace == req.namespace && (req.include_resolved || !d.is_resolved)); + let data = serde_json::to_string(&debts).unwrap_or_default(); + Ok(data.to_string()) + } +} + +pub struct OmniSearchHandler; + +#[async_trait] +impl McpTool for OmniSearchHandler { + fn name(&self) -> &'static str { + "omni_search" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("omni_search", "Execute omni_search") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: OmniSearchTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let matches = if let Ok(idx) = state.search_index.read() { + idx.search(&req.query, req.namespace.as_deref()) + .unwrap_or_default() + } else { + vec![] + }; + + let mut kg = KnowledgeGraph::default(); + let mut tasks = Vec::new(); + let mut snippets = Vec::new(); + let mut adrs = Vec::new(); + + state.read_graph(|full| { + for (id, doc_type, _, _, _) in &matches { + if doc_type == "entity" + && let Some(e) = full.entities.get(id) + { + kg.entities.insert(id.clone(), e.clone()); + } + } + }); + for t in state.tasks.read() { + if matches + .iter() + .any(|(id, typ, _, _, _)| id == &t.id && typ == "task") + { + tasks.push(t); + } + } + for s in state.snippets.read() { + if matches + .iter() + .any(|(id, typ, _, _, _)| id == &s.name && typ == "snippet") + { + snippets.push(s); + } + } + for a in state.adrs.read() { + if matches + .iter() + .any(|(id, typ, _, _, _)| id == &a.id && typ == "adr") + { + adrs.push(a); + } + } + + let q = req.query.to_lowercase(); + let tech_debts: Vec<_> = state + .tech_debts + .read() + .into_iter() + .filter(|d| { + req.namespace.as_ref().is_none_or(|ns| d.namespace == *ns) + && (contains_ignore_ascii_case(&d.description, &q) + || contains_ignore_ascii_case(&d.ideal_solution, &q)) + }) + .collect(); + let memos: Vec<_> = state + .handoff_memos + .read() + .into_iter() + .filter(|m| { + req.namespace.as_ref().is_none_or(|ns| m.namespace == *ns) + && contains_ignore_ascii_case(&m.content, &q) + }) + .collect(); + let error_fixes: Vec<_> = state + .error_fixes + .read() + .into_iter() + .filter(|f| { + contains_ignore_ascii_case(&f.signature, &q) + || contains_ignore_ascii_case(&f.solution, &q) + }) + .collect(); + + let report = serde_json::json!({ + "knowledge_graph": kg.entities, + "tasks": tasks, + "snippets": snippets, + "adrs": adrs, + "tech_debts": tech_debts, + "handoff_memos": memos, + "error_fixes": error_fixes + }); + Ok(report.to_string()) + } +} + +pub struct GetProjectHealthHandler; + +#[async_trait] +impl McpTool for GetProjectHealthHandler { + fn name(&self) -> &'static str { + "get_project_health" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "get_project_health", + "Execute get_project_health", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: GetProjectHealthTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let active_tasks = state + .tasks + .read() + .into_iter() + .filter(|t| t.status != "done") + .count(); + let unresolved_debt = state + .tech_debts + .read() + .into_iter() + .filter(|d| d.namespace == req.namespace && !d.is_resolved) + .count(); + let unread_memos = state + .handoff_memos + .read() + .into_iter() + .filter(|m| m.namespace == req.namespace) + .count(); + let active_milestones = state + .milestones + .read() + .into_iter() + .filter(|m| m.namespace == req.namespace && m.status != "done") + .count(); + let remaining_checklists = state + .pr_checklists + .read() + .into_iter() + .filter(|c| c.namespace == req.namespace) + .count(); + + let report = serde_json::json!({ + "active_tasks": active_tasks, + "unresolved_tech_debt": unresolved_debt, + "unread_handoff_memos": unread_memos, + "active_milestones": active_milestones, + "remaining_pr_checklist_items": remaining_checklists + }); + Ok(report.to_string()) + } +} + +use crate::handlers_v2::utils::*; diff --git a/server/src/handlers_v2/mod.rs b/server/src/handlers_v2/mod.rs new file mode 100644 index 0000000..cc20902 --- /dev/null +++ b/server/src/handlers_v2/mod.rs @@ -0,0 +1,7 @@ +pub mod env; +pub mod graph; +pub mod meta; +pub mod notes; +pub mod tasks; +pub mod utils; +pub mod workspaces; diff --git a/server/src/handlers_v2/notes.rs b/server/src/handlers_v2/notes.rs new file mode 100644 index 0000000..c35764d --- /dev/null +++ b/server/src/handlers_v2/notes.rs @@ -0,0 +1,273 @@ +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::collections::HashSet; +use std::sync::Arc; +use std::time::{SystemTime, UNIX_EPOCH}; + +pub struct AddStickyNoteHandler; + +#[async_trait] +impl McpTool for AddStickyNoteHandler { + fn name(&self) -> &'static str { + "add_sticky_note" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("add_sticky_note", "Execute add_sticky_note") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: AddStickyNoteTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.sticky.modify(|notes| { + notes.push(StickyNote { + timestamp: SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + content: req.content, + }); + }); + Ok("Sticky note added.".to_string()) + } +} + +pub struct ReadStickyNotesHandler; + +#[async_trait] +impl McpTool for ReadStickyNotesHandler { + fn name(&self) -> &'static str { + "read_sticky_notes" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "read_sticky_notes", + "Execute read_sticky_notes", + ) + } + + async fn execute(&self, _args: Value, state: Arc) -> Result { + let data = serde_json::to_string(&state.sticky.read()).unwrap_or_else(|_| "[]".to_string()); + Ok(data.to_string()) + } +} + +pub struct DeleteStickyNoteHandler; + +#[async_trait] +impl McpTool for DeleteStickyNoteHandler { + fn name(&self) -> &'static str { + "delete_sticky_note" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "delete_sticky_note", + "Execute delete_sticky_note", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: DeleteStickyNoteTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut success = false; + state.sticky.modify(|notes| { + if req.index > 0 && req.index <= notes.len() { + notes.remove(req.index - 1); + success = true; + } + }); + if success { + Ok("Sticky note deleted.".to_string()) + } else { + Err("Invalid sticky note index.".to_string()) + } + } +} + +pub struct ClearStickyNotesHandler; + +#[async_trait] +impl McpTool for ClearStickyNotesHandler { + fn name(&self) -> &'static str { + "clear_sticky_notes" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "clear_sticky_notes", + "Execute clear_sticky_notes", + ) + } + + async fn execute(&self, _args: Value, state: Arc) -> Result { + state.sticky.modify(|notes| { + notes.clear(); + }); + Ok("All sticky notes cleared.".to_string()) + } +} + +pub struct LeaveHandoffMemoHandler; + +#[async_trait] +impl McpTool for LeaveHandoffMemoHandler { + fn name(&self) -> &'static str { + "leave_handoff_memo" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "leave_handoff_memo", + "Execute leave_handoff_memo", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: LeaveHandoffMemoTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.handoff_memos.modify(|memos| { + memos.push(crate::models::HandoffMemo { + id: uuid::Uuid::new_v4().to_string(), + author: "agy".to_string(), + content: req.content, + namespace: req.namespace, + timestamp: SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + }) + }); + Ok("Handoff memo left".to_string()) + } +} + +pub struct ReadHandoffMemosHandler; + +#[async_trait] +impl McpTool for ReadHandoffMemosHandler { + fn name(&self) -> &'static str { + "read_handoff_memos" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "read_handoff_memos", + "Execute read_handoff_memos", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: ReadHandoffMemosTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut items = state.handoff_memos.read(); + if let Some(ns) = req.namespace { + items.retain(|i| i.namespace == ns); + } + let data = serde_json::to_string(&items).unwrap_or_default(); + Ok(data.to_string()) + } +} + +pub struct ClearHandoffMemosHandler; + +#[async_trait] +impl McpTool for ClearHandoffMemosHandler { + fn name(&self) -> &'static str { + "clear_handoff_memos" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "clear_handoff_memos", + "Execute clear_handoff_memos", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: ClearHandoffMemosTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let ids: HashSet<_> = req.ids.into_iter().collect(); + state + .handoff_memos + .modify(|memos| memos.retain(|m| !ids.contains(&m.id))); + Ok("Handoff memos cleared".to_string()) + } +} + +pub struct AddSessionSummaryHandler; + +#[async_trait] +impl McpTool for AddSessionSummaryHandler { + fn name(&self) -> &'static str { + "add_session_summary" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "add_session_summary", + "Execute add_session_summary", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: AddSessionSummaryTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.session_summaries.modify(|summaries| { + summaries.push(crate::models::SessionSummary { + summary: req.summary, + namespace: req.namespace, + timestamp: SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + }) + }); + Ok("Session summary added".to_string()) + } +} + +pub struct GenerateStandupReportHandler; + +#[async_trait] +impl McpTool for GenerateStandupReportHandler { + fn name(&self) -> &'static str { + "generate_standup_report" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "generate_standup_report", + "Execute generate_standup_report", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: GenerateStandupReportTool = + serde_json::from_value(args).map_err(|e| e.to_string())?; + let cutoff = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() + .saturating_sub(req.hours_lookback * 3600); + let tasks = state + .tasks + .read() + .into_iter() + .filter(|t| t.updated_at >= cutoff) + .collect::>(); + let changes = state + .ledger + .read() + .into_iter() + .filter(|c| c.timestamp >= cutoff) + .collect::>(); + let summaries = state + .session_summaries + .read() + .into_iter() + .filter(|s| s.namespace == req.namespace && s.timestamp >= cutoff) + .collect::>(); + let report = serde_json::json!({ "tasks_updated": tasks, "code_changes": changes, "session_summaries": summaries }); + Ok(report.to_string()) + } +} diff --git a/server/src/handlers_v2/tasks.rs b/server/src/handlers_v2/tasks.rs new file mode 100644 index 0000000..8f62c6d --- /dev/null +++ b/server/src/handlers_v2/tasks.rs @@ -0,0 +1,456 @@ +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; +use std::time::{SystemTime, UNIX_EPOCH}; + +pub struct AddTaskHandler; + +#[async_trait] +impl McpTool for AddTaskHandler { + fn name(&self) -> &'static str { + "add_task" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("add_task", "Execute add_task") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: AddTaskTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(); + let task_id = uuid::Uuid::new_v4().to_string(); + + let parent_id = req.parent_id.clone(); + let deps = req.dependencies.clone().unwrap_or_default(); + + let task = Task { + id: task_id.clone(), + title: req.title, + status: "pending".to_string(), + description: req.description, + created_at: now, + updated_at: now, + git_branch: req.git_branch, + parent_id, + dependencies: deps, + acceptance_criteria: vec![], + }; + if let Ok(idx) = state.search_index.read() { + drop(idx.index_task(&task)); + } + state.tasks.modify(|tasks| { + tasks.push(task); + }); + Ok(format!("Task added with ID: {}", task_id).to_string()) + } +} + +pub struct DeleteTaskHandler; + +#[async_trait] +impl McpTool for DeleteTaskHandler { + fn name(&self) -> &'static str { + "delete_task" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("delete_task", "Execute delete_task") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: DeleteTaskTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut deleted_count = 0; + state.tasks.modify(|tasks| { + let initial_len = tasks.len(); + // Collect IDs of tasks to delete (this task + all its recursive children) + let mut to_delete = std::collections::HashSet::new(); + to_delete.insert(req.id.clone()); + + let mut children_map: std::collections::HashMap> = + std::collections::HashMap::new(); + for t in tasks.iter() { + if let Some(pid) = &t.parent_id { + children_map + .entry(pid.clone()) + .or_default() + .push(t.id.clone()); + } + } + + let mut queue = std::collections::VecDeque::new(); + queue.push_back(req.id.clone()); + + while let Some(curr) = queue.pop_front() { + if to_delete.insert(curr.clone()) + && let Some(children) = children_map.get(&curr) + { + queue.extend(children.iter().cloned()); + } + } + + tasks.retain(|t| !to_delete.contains(&t.id)); + deleted_count = initial_len - tasks.len(); + }); + + if deleted_count > 0 { + Ok(vec![ + format!("Deleted task and its children ({} total).", deleted_count).to_string(), + ][0] + .clone()) + } else { + Ok("Task not found.".to_string()) + } + } +} + +pub struct UpdateTaskStatusHandler; + +#[async_trait] +impl McpTool for UpdateTaskStatusHandler { + fn name(&self) -> &'static str { + "update_task_status" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "update_task_status", + "Execute update_task_status", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: UpdateTaskStatusTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut found = false; + let mut blocked = false; + let mut blocker_details = String::new(); + let target_status = req.status.to_lowercase(); + + state.tasks.modify(|tasks| { + // Find target task + let mut target_id = String::new(); + if let Some(t) = tasks.iter().find(|t| t.id == req.id || t.title == req.id) { + target_id = t.id.clone(); + } + + if target_id.is_empty() { + return; + } + found = true; + + if target_status == "done" || target_status == "completed" { + // 1. Check Acceptance Criteria + if let Some(t) = tasks.iter().find(|t| t.id == target_id) + && t.acceptance_criteria.iter().any(|c| !c.is_met) + { + blocked = true; + blocker_details = "Unmet acceptance criteria exist.".to_string(); + } + + // 2. Check dependencies + if !blocked { + let mut uncompleted_deps = Vec::new(); + if let Some(t) = tasks.iter().find(|t| t.id == target_id) { + for dep_id in &t.dependencies { + if let Some(dep_task) = tasks.iter().find(|dt| dt.id == *dep_id) + && dep_task.status != "completed" + && dep_task.status != "done" + { + uncompleted_deps.push(dep_task.title.clone()); + } + } + } + if !uncompleted_deps.is_empty() { + blocked = true; + blocker_details = + format!("Blocked by dependencies: {}", uncompleted_deps.join(", ")); + } + } + + // 3. Check child tasks + if !blocked { + let mut uncompleted_children = Vec::new(); + for child in tasks + .iter() + .filter(|t| t.parent_id.as_ref() == Some(&target_id)) + { + if child.status != "completed" && child.status != "done" { + uncompleted_children.push(child.title.clone()); + } + } + if !uncompleted_children.is_empty() { + blocked = true; + blocker_details = format!( + "Blocked by child tasks: {}", + uncompleted_children.join(", ") + ); + } + } + } + + if !blocked { + // Apply update + if let Some(t) = tasks.iter_mut().find(|t| t.id == target_id) { + t.status = target_status.clone(); + t.updated_at = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(); + } + + // Cascade cancellation to children + if target_status == "cancelled" || target_status == "abandoned" { + let mut children_map: std::collections::HashMap> = + std::collections::HashMap::new(); + for (idx, t) in tasks.iter().enumerate() { + if let Some(pid) = &t.parent_id { + children_map.entry(pid.clone()).or_default().push(idx); + } + } + + let mut queue = std::collections::VecDeque::new(); + queue.push_back(target_id.clone()); + + while let Some(curr) = queue.pop_front() { + if let Some(child_indices) = children_map.get(&curr) { + for &idx in child_indices { + if tasks[idx].status != "completed" + && tasks[idx].status != target_status + { + tasks[idx].status = target_status.clone(); + queue.push_back(tasks[idx].id.clone()); + } + } + } + } + } + } + }); + + if blocked { + Ok(format!( + "Error: Cannot transition task. {}", + blocker_details + )) + } else if found { + Ok("Task status updated.".to_string()) + } else { + Ok("Task not found.".to_string()) + } + } +} + +pub struct ListActiveTasksHandler; + +#[async_trait] +impl McpTool for ListActiveTasksHandler { + fn name(&self) -> &'static str { + "list_active_tasks" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "list_active_tasks", + "Execute list_active_tasks", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: ListActiveTasksTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut tasks = 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(data.to_string()) + } +} + +pub struct SetAcceptanceCriteriaHandler; + +#[async_trait] +impl McpTool for SetAcceptanceCriteriaHandler { + fn name(&self) -> &'static str { + "set_acceptance_criteria" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "set_acceptance_criteria", + "Execute set_acceptance_criteria", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: SetAcceptanceCriteriaTool = + serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut success = false; + state.tasks.modify(|tasks| { + if let Some(task) = tasks.iter_mut().rev().find(|t| t.title == req.task_title) { + task.acceptance_criteria = req + .criteria + .into_iter() + .map(|desc| crate::models::AcceptanceCriteria { + id: uuid::Uuid::new_v4().to_string(), + description: desc, + is_met: false, + }) + .collect(); + task.updated_at = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(); + success = true; + } + }); + if success { + Ok(vec!["Acceptance criteria set successfully.".to_string()][0].clone()) + } else { + Ok("Task not found.".to_string()) + } + } +} + +pub struct VerifyAcceptanceCriteriaHandler; + +#[async_trait] +impl McpTool for VerifyAcceptanceCriteriaHandler { + fn name(&self) -> &'static str { + "verify_acceptance_criteria" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::( + "verify_acceptance_criteria", + "Execute verify_acceptance_criteria", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: VerifyAcceptanceCriteriaTool = + serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut success = false; + let mut already_met = false; + state.tasks.modify(|tasks| { + if let Some(task) = tasks.iter_mut().find(|t| t.id == req.task_id) + && let Some(ac) = task + .acceptance_criteria + .iter_mut() + .find(|c| c.id == req.criteria || c.description == req.criteria) + { + if ac.is_met { + already_met = true; + } else { + ac.is_met = true; + success = true; + task.updated_at = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(); + } + } + }); + if success { + Ok(vec![format!( + "Acceptance criteria verified with proof: {}", + req.proof + )][0] + .clone()) + } else if already_met { + Ok("Acceptance criteria was already met.".to_string()) + } else { + Ok(vec!["Acceptance criteria or task not found.".to_string()][0].clone()) + } + } +} + +pub struct AddMilestoneHandler; + +#[async_trait] +impl McpTool for AddMilestoneHandler { + fn name(&self) -> &'static str { + "add_milestone" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("add_milestone", "Execute add_milestone") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: AddMilestoneTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + state.milestones.modify(|ms| { + ms.push(crate::models::Milestone { + id: uuid::Uuid::new_v4().to_string(), + title: req.title, + status: "pending".to_string(), + namespace: req.namespace, + target_date: None, + }) + }); + Ok("Milestone added".to_string()) + } +} + +pub struct UpdateMilestoneHandler; + +#[async_trait] +impl McpTool for UpdateMilestoneHandler { + fn name(&self) -> &'static str { + "update_milestone" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("update_milestone", "Execute update_milestone") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: UpdateMilestoneTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut found = false; + state.milestones.modify(|ms| { + for m in ms.iter_mut() { + if m.id == req.id { + m.status = req.status.clone(); + found = true; + break; + } + } + }); + if found { + Ok("Milestone updated".to_string()) + } else { + Ok("Milestone not found".to_string()) + } + } +} + +pub struct ListMilestonesHandler; + +#[async_trait] +impl McpTool for ListMilestonesHandler { + fn name(&self) -> &'static str { + "list_milestones" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("list_milestones", "Execute list_milestones") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: ListMilestonesTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut items = state.milestones.read(); + if let Some(ns) = req.namespace { + items.retain(|i| i.namespace == ns); + } + let data = serde_json::to_string(&items).unwrap_or_default(); + Ok(data.to_string()) + } +} diff --git a/server/src/handlers_v2/utils.rs b/server/src/handlers_v2/utils.rs new file mode 100644 index 0000000..bbfe842 --- /dev/null +++ b/server/src/handlers_v2/utils.rs @@ -0,0 +1,9 @@ +pub fn contains_ignore_ascii_case(haystack: &str, needle: &str) -> bool { + if needle.is_empty() { + return true; + } + haystack + .as_bytes() + .windows(needle.len()) + .any(|w| w.eq_ignore_ascii_case(needle.as_bytes())) +} diff --git a/server/src/handlers_v2/workspaces.rs b/server/src/handlers_v2/workspaces.rs new file mode 100644 index 0000000..e5976f6 --- /dev/null +++ b/server/src/handlers_v2/workspaces.rs @@ -0,0 +1,348 @@ +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; +use std::time::{SystemTime, UNIX_EPOCH}; + +pub struct PinFileHandler; + +#[async_trait] +impl McpTool for PinFileHandler { + fn name(&self) -> &'static str { + "pin_file" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("pin_file", "Execute pin_file") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + 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: SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_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::("unpin_file", "Execute unpin_file") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + 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::( + "list_pinned_files", + "Execute list_pinned_files", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: ListPinnedFilesTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut pinned = state.pinned_files.read(); + 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(data.to_string()) + } +} + +pub struct StoreSnippetHandler; + +#[async_trait] +impl McpTool for StoreSnippetHandler { + fn name(&self) -> &'static str { + "store_snippet" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("store_snippet", "Execute store_snippet") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: StoreSnippetTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let snippet = Snippet { + name: req.name.clone(), + language: req.language, + code: req.code, + description: req.description, + updated_at: SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + }; + + let s_clone = snippet.clone(); + state.snippets.modify(|snippets| { + snippets.retain(|s| s.name != req.name); + snippets.push(s_clone); + }); + + if let Ok(idx) = state.search_index.read() { + drop(idx.index_snippet(&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::("search_snippets", "Execute search_snippets") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: SearchSnippetsTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let query = req.query.to_lowercase(); + let snippets = state.snippets.read(); + let mut results = Vec::new(); + for s in snippets { + if contains_ignore_ascii_case(&s.name, &query) + || contains_ignore_ascii_case(&s.description, &query) + || contains_ignore_ascii_case(&s.language, &query) + { + results.push(s); + } + } + let data = serde_json::to_string(&results).unwrap_or_default(); + Ok(data.to_string()) + } +} + +pub struct DeleteSnippetHandler; + +#[async_trait] +impl McpTool for DeleteSnippetHandler { + fn name(&self) -> &'static str { + "delete_snippet" + } + + fn schema(&self) -> Value { + crate::mcp::tool_def::("delete_snippet", "Execute delete_snippet") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + 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 { + Ok("Snippet deleted.".to_string()) + } else { + Ok("Snippet not found.".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::( + "save_context_workspace", + "Execute save_context_workspace", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + 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: SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_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::( + "load_context_workspace", + "Execute load_context_workspace", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: LoadContextWorkspaceTool = + serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut ws = state.context_workspaces.read(); + ws.retain(|w| w.namespace == req.namespace && w.name == req.name); + let data = serde_json::to_string(&ws.first()).unwrap_or_default(); + Ok(data.to_string()) + } +} + +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::( + "list_context_workspaces", + "Execute list_context_workspaces", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: ListContextWorkspacesTool = + serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut ws = state.context_workspaces.read(); + ws.retain(|w| w.namespace == req.namespace); + let data = serde_json::to_string(&ws).unwrap_or_default(); + Ok(data.to_string()) + } +} + +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::( + "add_pr_checklist_item", + "Execute add_pr_checklist_item", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + 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::("get_pr_checklist", "Execute get_pr_checklist") + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + let req: GetPrChecklistTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut items = state.pr_checklists.read(); + items.retain(|i| i.namespace == req.namespace); + let data = serde_json::to_string(&items).unwrap_or_default(); + Ok(data.to_string()) + } +} + +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::( + "clear_pr_checklist", + "Execute clear_pr_checklist", + ) + } + + async fn execute(&self, args: Value, state: Arc) -> Result { + 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_v2::utils::*; diff --git a/server/src/main.rs b/server/src/main.rs index cdab8ee..baab1c1 100644 --- a/server/src/main.rs +++ b/server/src/main.rs @@ -4,8 +4,10 @@ )] mod handlers; +mod handlers_v2; mod mcp; mod models; +mod router; mod search; mod state; mod store; @@ -218,14 +220,12 @@ async fn run_server(state: Arc) -> Result<(), Box) -> Result<(), Box l, - Err(e) => { - let log_path = dirs::home_dir() - .unwrap_or_default() - .join(".gemini/mcp_memory/daemon_error.log"); - let _ = tokio::fs::write(&log_path, format!("Failed to bind to {}: {}\n", addr, e)).await; - return Ok(()); - } - }; - if let Err(e) = axum::serve(listener, app.into_make_service()).await { + let listener = match tokio::net::TcpListener::bind(&addr).await { + Ok(l) => l, + Err(e) => { let log_path = dirs::home_dir() .unwrap_or_default() .join(".gemini/mcp_memory/daemon_error.log"); - let _ = tokio::fs::write(&log_path, format!("Server crashed: {}\n", e)).await; + let _ = + tokio::fs::write(&log_path, format!("Failed to bind to {}: {}\n", addr, e)).await; + return Ok(()); } - Ok(()) + }; + if let Err(e) = axum::serve(listener, app.into_make_service()).await { + let log_path = dirs::home_dir() + .unwrap_or_default() + .join(".gemini/mcp_memory/daemon_error.log"); + let _ = tokio::fs::write(&log_path, format!("Server crashed: {}\n", e)).await; + } + Ok(()) } async fn ws_handler( @@ -486,42 +487,43 @@ async fn handle_socket(socket: WebSocket, state: Arc, client_type: Str if client_type == "proxy" { // Send activity broadcast to UI clients if let Some(method) = payload.get("method").and_then(|m| m.as_str()) - && method == "tools/call" { - let name = payload - .get("params") - .and_then(|p| p.get("name")) - .and_then(|n| n.as_str()) - .unwrap_or("unknown_tool"); - let activity_msg = format!("Agent executed tool: {}", name); + && method == "tools/call" + { + let name = payload + .get("params") + .and_then(|p| p.get("name")) + .and_then(|n| n.as_str()) + .unwrap_or("unknown_tool"); + let activity_msg = format!("Agent executed tool: {}", name); - let event = serde_json::json!({ - "type": "activity", - "data": activity_msg - }); + let event = serde_json::json!({ + "type": "activity", + "data": activity_msg + }); - let senders: Vec<_> = state_clone - .clients - .read() - .unwrap() - .iter() - .filter_map(|(id, tx)| { - if id != &session_id_clone { - Some(tx.clone()) - } else { - None - } - }) - .collect(); + let senders: Vec<_> = state_clone + .clients + .read() + .unwrap() + .iter() + .filter_map(|(id, tx)| { + if id != &session_id_clone { + Some(tx.clone()) + } else { + None + } + }) + .collect(); - for client_tx in senders { - let _ = client_tx.try_send(event.to_string()); - } + for client_tx in senders { + let _ = client_tx.try_send(event.to_string()); } + } } // End if proxy // Process MCP request if let Some(response) = handler.handle_request(payload).await { - let res_str = serde_json::to_string(&response).unwrap(); + let res_str = serde_json::to_string(&response).unwrap_or_default(); let tx_opt = state_clone .clients .read() @@ -576,7 +578,11 @@ async fn handle_socket(socket: WebSocket, state: Arc, client_type: Str impl Drop for SessionCleanup { fn drop(&mut self) { - self.state.clients.write().unwrap().remove(&self.session_id); + self.state + .clients + .write() + .unwrap_or_else(|e| e.into_inner()) + .remove(&self.session_id); if let Some(task) = self.send_task.take() { task.abort(); } @@ -640,7 +646,13 @@ async fn nvim_telemetry_handler( }); let msg_str = ws_msg.to_string(); - let senders: Vec<_> = state.clients.read().unwrap().values().cloned().collect(); + let senders: Vec<_> = state + .clients + .read() + .unwrap_or_else(|e| e.into_inner()) + .values() + .cloned() + .collect(); for tx in senders { let _ = tx.try_send(msg_str.clone()); } @@ -698,9 +710,12 @@ fn main() -> Result<(), Box> { let mut cmd = std::process::Command::new("curl"); cmd.arg("-k").arg("-X").arg("POST"); if !token.is_empty() { - cmd.arg("-H").arg(format!("Authorization: Bearer {}", token.trim())); + cmd.arg("-H") + .arg(format!("Authorization: Bearer {}", token.trim())); } - let _ = cmd.arg(format!("https://127.0.0.1:{}/shutdown", port)).output(); + let _ = cmd + .arg(format!("https://127.0.0.1:{}/shutdown", port)) + .output(); println!("Sent shutdown request to server."); return Ok(()); } @@ -711,9 +726,12 @@ fn main() -> Result<(), Box> { let mut cmd = std::process::Command::new("curl"); cmd.arg("-k").arg("-X").arg("POST"); if !token.is_empty() { - cmd.arg("-H").arg(format!("Authorization: Bearer {}", token.trim())); + cmd.arg("-H") + .arg(format!("Authorization: Bearer {}", token.trim())); } - let _ = cmd.arg(format!("https://127.0.0.1:{}/shutdown", port)).output(); + let _ = cmd + .arg(format!("https://127.0.0.1:{}/shutdown", port)) + .output(); println!("Sent shutdown request to existing server. Waiting for it to exit..."); std::thread::sleep(std::time::Duration::from_millis(1500)); return Ok(()); @@ -780,13 +798,11 @@ fn main() -> Result<(), Box> { let json_path = base.join(file_name); if json_path.exists() && let Ok(data) = fs::read(&json_path) - && serde_json::from_slice::(&data).is_ok() { - table.insert(*key, data.as_slice()).unwrap(); - let _ = fs::rename( - &json_path, - json_path.with_extension("json.migrated"), - ); - } + && serde_json::from_slice::(&data).is_ok() + { + table.insert(*key, data.as_slice()).unwrap(); + let _ = fs::rename(&json_path, json_path.with_extension("json.migrated")); + } } } } diff --git a/server/src/refactor.py b/server/src/refactor.py new file mode 100644 index 0000000..03915c6 --- /dev/null +++ b/server/src/refactor.py @@ -0,0 +1,173 @@ +import os +import re + +GROUPS = { + "graph": [ + "query_graph_path", "create_entities", "create_relations", "add_observations", + "delete_entities", "delete_observations", "delete_relations", "read_graph", + "search_nodes", "open_nodes", "visualize_graph", "condense_entity", + "merge_entities", "find_orphans" + ], + "tasks": [ + "add_task", "delete_task", "update_task_status", "list_active_tasks", + "set_acceptance_criteria", "verify_acceptance_criteria", + "add_milestone", "update_milestone", "list_milestones" + ], + "notes": [ + "add_sticky_note", "read_sticky_notes", "delete_sticky_note", "clear_sticky_notes", + "leave_handoff_memo", "read_handoff_memos", "clear_handoff_memos", + "add_session_summary", "generate_standup_report" + ], + "meta": [ + "log_decision", "query_decisions", "log_error_fix", "search_error_fixes", + "log_code_change", "query_recent_changes", "learn_preference", "read_preferences", + "log_tech_debt", "resolve_tech_debt", "list_tech_debt", "omni_search", "get_project_health" + ], + "env": [ + "update_env_fingerprint", "read_env_fingerprint", "log_env_requirement", + "register_environment", "get_environment_details" + ], + "workspaces": [ + "pin_file", "unpin_file", "list_pinned_files", "store_snippet", "search_snippets", + "delete_snippet", "save_context_workspace", "load_context_workspace", + "list_context_workspaces", "add_pr_checklist_item", "get_pr_checklist", "clear_pr_checklist" + ] +} + +def to_camel_case(snake_str): + components = snake_str.split('_') + return "".join(x.title() for x in components) + +def parse_rust_match(file_path): + with open(file_path, "r", encoding="utf-8") as f: + lines = f.readlines() + + start_idx = -1 + for i, line in enumerate(lines): + if "let result: Result = match name {" in line: + start_idx = i + break + + if start_idx == -1: + return {} + + brace_depth = 1 + i = start_idx + 1 + + tools = {} + current_tool = None + current_body = [] + + while i < len(lines): + line = lines[i] + + if brace_depth == 1 and "=>" in line and '"' in line: + parts = line.strip().split('"') + if len(parts) >= 3: + tool_name = parts[1] + current_tool = tool_name + current_body = [] + # Don't add the "name" => { line + + if current_tool is not None and not (brace_depth == 1 and "=>" in line and '"' in line): + # check if this line closes the block + next_depth = brace_depth + line.count('{') - line.count('}') + if next_depth == 1 and current_tool is not None: + # This is the closing brace + tools[current_tool] = "".join(current_body) + current_tool = None + else: + current_body.append(line) + + brace_depth += line.count('{') + brace_depth -= line.count('}') + + if brace_depth == 0: + break + + i += 1 + + return tools + +def transform_body(body): + # Transform parse_tool! + body = re.sub( + r'let req = parse_tool!\(args, id, ([^)]+)\);', + r'let req: \1 = serde_json::from_value(args).map_err(|e| e.to_string())?;', + body + ) + # Transform handle_list_with_namespace! + def repl_handle_list(m): + store = m.group(1) + tool_type = m.group(2) + return f""" + let req: {tool_type} = serde_json::from_value(args).map_err(|e| e.to_string())?; + let mut items = state.{store}.read(); + if let Some(ns) = req.namespace {{ + items.retain(|i| i.namespace == ns); + }} + let data = serde_json::to_string(&items).unwrap_or_default(); + return Ok(data.to_string()); + """ + body = re.sub( + r'return handle_list_with_namespace!\(self, ([^,]+), ([^,]+), args, id\);', + repl_handle_list, + body + ) + + # Replace self.state with state + body = body.replace("self.state.", "state.") + + return body + + +tools = parse_rust_match("server/src/handlers.rs") + +for group, tool_names in GROUPS.items(): + file_path = f"server/src/handlers_v2/{group}.rs" + with open(file_path, "w", encoding="utf-8") as f: + f.write("use crate::router::McpTool;\n") + f.write("use crate::state::MemoryState;\n") + f.write("use crate::tools::*;\n") + f.write("use async_trait::async_trait;\n") + f.write("use serde_json::Value;\n") + f.write("use std::sync::Arc;\n") + f.write("use std::time::{SystemTime, UNIX_EPOCH};\n\n") + + for name in tool_names: + if name not in tools: + continue + + body = tools[name] + # special case for query_graph_path which we already wrote properly? + # actually we will just overwrite it with the transformed body + body = transform_body(body) + + struct_name = to_camel_case(name) + "Handler" + tool_type = to_camel_case(name) + "Tool" + + f.write(f"pub struct {struct_name};\n\n") + f.write(f"#[async_trait]\n") + f.write(f"impl McpTool for {struct_name} {{\n") + f.write(f" fn name(&self) -> &'static str {{\n") + f.write(f' "{name}"\n') + f.write(f" }}\n\n") + f.write(f" fn schema(&self) -> Value {{\n") + # For schema description we can just put a generic one or extract it. + # I will use a generic one for now, or you can extract it from tools/list. + f.write(f' crate::mcp::tool_def::<{tool_type}>(\n') + f.write(f' "{name}",\n') + f.write(f' "Execute {name}",\n') + f.write(f' )\n') + f.write(f" }}\n\n") + f.write(f" async fn execute(&self, args: Value, state: Arc) -> Result {{\n") + f.write(body) + f.write(f" }}\n") + f.write(f"}}\n\n") + +print("Generated handlers_v2 modules") + +# generate mod.rs +with open("server/src/handlers_v2/mod.rs", "w", encoding="utf-8") as f: + for group in GROUPS.keys(): + f.write(f"pub mod {group};\n") diff --git a/server/src/refactor.rs b/server/src/refactor.rs new file mode 100644 index 0000000..16f1b56 --- /dev/null +++ b/server/src/refactor.rs @@ -0,0 +1,11 @@ +use std::fs; +use std::io::Write; + +fn main() { + let content = fs::read_to_string("server/src/handlers.rs").unwrap(); + println!("Read {} bytes", content.len()); + // Find the match name { block + let match_start = content.find("match name {").unwrap(); + // naive extraction + println!("Found match block at {}", match_start); +} diff --git a/server/src/refactor_router.py b/server/src/refactor_router.py new file mode 100644 index 0000000..c48ca5f --- /dev/null +++ b/server/src/refactor_router.py @@ -0,0 +1,82 @@ +import re +import os + +with open("server/src/handlers.rs", "r", encoding="utf-8") as f: + lines = f.readlines() + +list_start = -1 +for i, line in enumerate(lines): + if '"tools/list" => {' in line: + list_start = i + break + +# Find end of tools/call +call_start = -1 +for i in range(list_start, len(lines)): + if '"tools/call" => {' in line: + call_start = i + break + +# Find end of tools/call +# Match brace depth from call_start +brace_depth = 1 +call_end = -1 +for i in range(call_start + 1, len(lines)): + brace_depth += lines[i].count('{') + brace_depth -= lines[i].count('}') + if brace_depth == 0: + call_end = i + break + +# replacement block +replacement = """ "tools/list" => { + let mut tools: Vec = self.tools.values().map(|t| t.schema()).collect(); + tools.sort_by_key(|t| t.get("name").and_then(|n| n.as_str()).unwrap_or("").to_string()); + Some(crate::mcp::success( + id, + serde_json::json!({ "tools": tools }), + )) + } + "tools/call" => { + let params = req.get("params").unwrap_or(&serde_json::Value::Null); + let name = params.get("name").and_then(|n| n.as_str()).unwrap_or(""); + let args = params + .get("arguments") + .cloned() + .unwrap_or(serde_json::Value::Object(Default::default())); + + self.state + .broadcast_activity(&format!("Agent executed tool: {}", name)); + + let result: Result = if let Some(tool) = self.tools.get(name) { + tool.execute(args, self.state.clone()).await + } else { + Err(format!("Unknown tool: {}", name)) + }; + + match result { + Ok(text) => { + let payload = serde_json::json!({ + "content": [{"type": "text", "text": text}], + "isError": false + }); + Some(crate::mcp::success(id_clone, payload)) + } + Err(e) => { + tracing::error!("Tool {} failed: {}", name, e); + let payload = serde_json::json!({ + "content": [{"type": "text", "text": e}], + "isError": true + }); + Some(crate::mcp::success(id_clone, payload)) + } + } + } +""" + +new_lines = lines[:list_start] + [replacement] + lines[call_end+1:] + +with open("server/src/handlers.rs", "w", encoding="utf-8") as f: + f.writelines(new_lines) + +print("tools/list and tools/call replaced.") diff --git a/server/src/refactor_wire.py b/server/src/refactor_wire.py new file mode 100644 index 0000000..fc77d3e --- /dev/null +++ b/server/src/refactor_wire.py @@ -0,0 +1,106 @@ +import re +import os + +with open("server/src/handlers.rs", "r", encoding="utf-8") as f: + content = f.read() + +# Replace MemoryHandler struct +struct_pattern = r'pub struct MemoryHandler \{\s*pub state: Arc,\s*\}' + +new_struct = """use crate::router::McpTool; + +pub struct MemoryHandler { + pub state: Arc, + pub tools: std::collections::HashMap>, +} + +impl MemoryHandler { + pub fn new(state: Arc) -> Self { + let mut tools: std::collections::HashMap> = std::collections::HashMap::new(); + + macro_rules! register { + ($module:ident::$handler:ident) => { + let h = crate::handlers_v2::$module::$handler; + tools.insert(h.name().to_string(), Box::new(h)); + }; + } + + register!(graph::QueryGraphPathHandler); + register!(graph::CreateEntitiesHandler); + register!(graph::CreateRelationsHandler); + register!(graph::AddObservationsHandler); + register!(graph::DeleteEntitiesHandler); + register!(graph::DeleteObservationsHandler); + register!(graph::DeleteRelationsHandler); + register!(graph::ReadGraphHandler); + register!(graph::SearchNodesHandler); + register!(graph::OpenNodesHandler); + register!(graph::VisualizeGraphHandler); + register!(graph::CondenseEntityHandler); + register!(graph::MergeEntitiesHandler); + register!(graph::FindOrphansHandler); + + register!(tasks::AddTaskHandler); + register!(tasks::DeleteTaskHandler); + register!(tasks::UpdateTaskStatusHandler); + register!(tasks::ListActiveTasksHandler); + register!(tasks::SetAcceptanceCriteriaHandler); + register!(tasks::VerifyAcceptanceCriteriaHandler); + register!(tasks::AddMilestoneHandler); + register!(tasks::UpdateMilestoneHandler); + register!(tasks::ListMilestonesHandler); + + register!(notes::AddStickyNoteHandler); + register!(notes::ReadStickyNotesHandler); + register!(notes::DeleteStickyNoteHandler); + register!(notes::ClearStickyNotesHandler); + register!(notes::LeaveHandoffMemoHandler); + register!(notes::ReadHandoffMemosHandler); + register!(notes::ClearHandoffMemosHandler); + register!(notes::AddSessionSummaryHandler); + register!(notes::GenerateStandupReportHandler); + + register!(meta::LogDecisionHandler); + register!(meta::QueryDecisionsHandler); + register!(meta::LogErrorFixHandler); + register!(meta::SearchErrorFixesHandler); + register!(meta::LogCodeChangeHandler); + register!(meta::QueryRecentChangesHandler); + register!(meta::LearnPreferenceHandler); + register!(meta::ReadPreferencesHandler); + register!(meta::LogTechDebtHandler); + register!(meta::ResolveTechDebtHandler); + register!(meta::ListTechDebtHandler); + register!(meta::OmniSearchHandler); + register!(meta::GetProjectHealthHandler); + + register!(env::UpdateEnvFingerprintHandler); + register!(env::ReadEnvFingerprintHandler); + register!(env::LogEnvRequirementHandler); + register!(env::RegisterEnvironmentHandler); + register!(env::GetEnvironmentDetailsHandler); + + register!(workspaces::PinFileHandler); + register!(workspaces::UnpinFileHandler); + register!(workspaces::ListPinnedFilesHandler); + register!(workspaces::StoreSnippetHandler); + register!(workspaces::SearchSnippetsHandler); + register!(workspaces::DeleteSnippetHandler); + register!(workspaces::SaveContextWorkspaceHandler); + register!(workspaces::LoadContextWorkspaceHandler); + register!(workspaces::ListContextWorkspacesHandler); + register!(workspaces::AddPrChecklistItemHandler); + register!(workspaces::GetPrChecklistHandler); + register!(workspaces::ClearPrChecklistHandler); + + Self { state, tools } + } +""" + +content = re.sub(struct_pattern, new_struct, content) +content = content.replace("impl MemoryHandler {\n pub async fn handle_request", " pub async fn handle_request") + +with open("server/src/handlers.rs", "w", encoding="utf-8") as f: + f.write(content) + +print("MemoryHandler struct updated.") diff --git a/server/src/router.rs b/server/src/router.rs new file mode 100644 index 0000000..2520791 --- /dev/null +++ b/server/src/router.rs @@ -0,0 +1,16 @@ +use crate::state::MemoryState; +use async_trait::async_trait; +use serde_json::Value; +use std::sync::Arc; + +#[async_trait] +pub trait McpTool: Send + Sync { + /// The unique name of the tool + fn name(&self) -> &'static str; + + /// The JSON schema for the tool + fn schema(&self) -> Value; + + /// Execute the tool with the given arguments + async fn execute(&self, args: Value, state: Arc) -> Result; +} diff --git a/server/src/search.rs b/server/src/search.rs index 4ef1017..ccf11dd 100644 --- a/server/src/search.rs +++ b/server/src/search.rs @@ -60,7 +60,7 @@ impl MemoryIndex { let body_field = self.body_field; let type_field = self.type_field; let namespace_field = self.namespace_field; - + tokio::task::spawn_blocking(move || { let doc = doc!( id_field => e.name.clone(), @@ -83,7 +83,7 @@ impl MemoryIndex { let body_field = self.body_field; let type_field = self.type_field; let namespace_field = self.namespace_field; - + tokio::task::spawn_blocking(move || { let doc = doc!( id_field => t.id.clone(), @@ -106,7 +106,11 @@ impl MemoryIndex { Ok(()) }) .await - .unwrap_or_else(|_| Err(tantivy::TantivyError::SystemError("Commit task panicked".to_string()))) + .unwrap_or_else(|_| { + Err(tantivy::TantivyError::SystemError( + "Commit task panicked".to_string(), + )) + }) } pub fn search( @@ -171,7 +175,7 @@ impl MemoryIndex { let body_field = self.body_field; let type_field = self.type_field; let namespace_field = self.namespace_field; - + tokio::task::spawn_blocking(move || { let doc = doc!( id_field => s.name.clone(), @@ -194,7 +198,7 @@ impl MemoryIndex { let body_field = self.body_field; let type_field = self.type_field; let namespace_field = self.namespace_field; - + tokio::task::spawn_blocking(move || { let doc = doc!( id_field => a.id.clone(), @@ -225,7 +229,7 @@ impl MemoryIndex { tokio::task::spawn_blocking(move || { let writer = writer.lock().unwrap(); - + for e in entities { writer.add_document(doc!( id_field => e.name.clone(), diff --git a/server/src/state.rs b/server/src/state.rs index 43de6fa..aae8dfd 100644 --- a/server/src/state.rs +++ b/server/src/state.rs @@ -62,7 +62,6 @@ impl MemoryState { pub async fn rebuild_index(&self) { if let Ok(new_idx) = MemoryIndex::new(&self.base_dir) { - let entities = self.graph.read().entities.into_values().collect(); let tasks = self.tasks.read(); let snippets = self.snippets.read(); diff --git a/server/src/store.rs b/server/src/store.rs index 173861f..c42aa59 100644 --- a/server/src/store.rs +++ b/server/src/store.rs @@ -14,52 +14,53 @@ impl let initial_data = Self::load_from_db(key, &db); let cache = Arc::new(RwLock::new(initial_data)); let (tx, mut rx) = tokio::sync::mpsc::channel::<()>(1); - + let db_clone = db.clone(); let key_clone = key.to_string(); let cache_clone = cache.clone(); - + tokio::spawn(async move { while rx.recv().await.is_some() { // Drain any other pending notifications so we batch writes - while let Ok(_) = rx.try_recv() {} + while rx.try_recv().is_ok() {} let db_inner = db_clone.clone(); let key_inner = key_clone.clone(); let json_data = { - let lock = cache_clone.read().unwrap(); - serde_json::to_vec(&*lock).unwrap() + let lock = cache_clone.read().unwrap_or_else(|e| e.into_inner()); + serde_json::to_vec(&*lock).unwrap_or_default() }; - + let _ = tokio::task::spawn_blocking(move || { - let write_txn = db_inner.begin_write().unwrap(); - { - let mut table = write_txn.open_table(STORE_TABLE).unwrap(); - table.insert(key_inner.as_str(), json_data.as_slice()).unwrap(); + if let Ok(write_txn) = db_inner.begin_write() { + if let Ok(mut table) = write_txn.open_table(STORE_TABLE) { + let _ = table.insert(key_inner.as_str(), json_data.as_slice()); + } + let _ = write_txn.commit(); } - write_txn.commit().unwrap(); - }).await; + }) + .await; } }); - - Self { - cache, - tx, - } + + Self { cache, tx } } fn load_from_db(key: &str, db: &Database) -> T { - let read_txn = db.begin_read().unwrap(); + let Ok(read_txn) = db.begin_read() else { + return T::default(); + }; if let Ok(table) = read_txn.open_table(STORE_TABLE) && let Ok(Some(value)) = table.get(key) - && let Ok(parsed) = serde_json::from_slice::(value.value()) { - return parsed; - } + && let Ok(parsed) = serde_json::from_slice::(value.value()) + { + return parsed; + } T::default() } pub fn read(&self) -> T { - let lock = self.cache.read().unwrap(); + let lock = self.cache.read().unwrap_or_else(|e| e.into_inner()); lock.clone() } @@ -67,13 +68,13 @@ impl where F: FnOnce(&T) -> R, { - let lock = self.cache.read().unwrap(); + let lock = self.cache.read().unwrap_or_else(|e| e.into_inner()); f(&lock) } pub fn modify(&self, f: F) { { - let mut lock = self.cache.write().unwrap(); + let mut lock = self.cache.write().unwrap_or_else(|e| e.into_inner()); f(&mut lock); } let _ = self.tx.try_send(()); diff --git a/server/src/tools.rs b/server/src/tools.rs index 31ccc6c..cce974b 100644 --- a/server/src/tools.rs +++ b/server/src/tools.rs @@ -314,6 +314,7 @@ pub struct AddSessionSummaryTool { /// Get a timeline of major project events. #[derive(Debug, Deserialize, Serialize, JsonSchema)] +#[allow(dead_code)] pub struct GetProjectTimelineTool { /// Optional namespace to restrict the timeline to. pub namespace: Option, diff --git a/server/tests/parity_test.rs b/server/tests/parity_test.rs index 1f7afd9..195afd8 100644 --- a/server/tests/parity_test.rs +++ b/server/tests/parity_test.rs @@ -2,22 +2,27 @@ use std::collections::HashSet; #[test] fn test_eager_tools_parity() { - // 1. Read handlers.rs to get memory tools - let memory_source = - std::fs::read_to_string("src/handlers.rs").expect("Failed to read handlers.rs"); + // 1. Read handlers_v2/*.rs to get memory tools let mut memory_tools = HashSet::new(); - let parts: Vec<&str> = memory_source.split("crate::mcp::tool_def").collect(); - for part in parts.iter().skip(1) { - if let Some(start) = part.find("\"") { - let rest = &part[start + 1..]; - if let Some(end) = rest.find("\"") { - memory_tools.insert(rest[..end].to_string()); + let entries = std::fs::read_dir("src/handlers_v2").expect("Failed to read handlers_v2 dir"); + for entry in entries { + let entry = entry.unwrap(); + if entry.path().extension().unwrap_or_default() == "rs" { + let memory_source = std::fs::read_to_string(entry.path()).unwrap(); + let parts: Vec<&str> = memory_source.split("crate::mcp::tool_def").collect(); + for part in parts.iter().skip(1) { + if let Some(start) = part.find("\"") { + let rest = &part[start + 1..]; + if let Some(end) = rest.find("\"") { + memory_tools.insert(rest[..end].to_string()); + } + } } } } assert!( !memory_tools.is_empty(), - "Could not find memory tools in handlers.rs" + "Could not find memory tools in handlers_v2 directory" ); // 2. Read nvim-core/src/lib.rs to get nvim tools @@ -26,12 +31,13 @@ fn test_eager_tools_parity() { let mut nvim_tools = HashSet::new(); for line in nvim_source.lines() { if line.contains("\"name\": \"nvim_") - && let Some(start) = line.find("\"name\": \"") { - let rest = &line[start + 9..]; - if let Some(end) = rest.find("\"") { - nvim_tools.insert(rest[..end].to_string()); - } + && let Some(start) = line.find("\"name\": \"") + { + let rest = &line[start + 9..]; + if let Some(end) = rest.find("\"") { + nvim_tools.insert(rest[..end].to_string()); } + } } assert!( !nvim_tools.is_empty(), diff --git a/stub/build.rs b/stub/build.rs index dc9a51c..c266efa 100644 --- a/stub/build.rs +++ b/stub/build.rs @@ -15,10 +15,6 @@ fn main() { .and_then(|out| String::from_utf8(out.stdout).ok()) .unwrap_or_else(|| "unknown".to_string()); - let version = format!( - "{} ({})", - git_date.trim(), - git_hash.trim() - ); + let version = format!("{} ({})", git_date.trim(), git_hash.trim()); println!("cargo:rustc-env=APP_VERSION={}", version); } diff --git a/stub/src/logger.rs b/stub/src/logger.rs index a9ceb30..06f96e6 100644 --- a/stub/src/logger.rs +++ b/stub/src/logger.rs @@ -1,24 +1,39 @@ -use std::sync::LazyLock; use regex::Regex; +use std::sync::LazyLock; static ID_REGEX: LazyLock = LazyLock::new(|| Regex::new(r#""id"\s*:\s*([^,}]+)"#).unwrap()); -static METHOD_REGEX: LazyLock = LazyLock::new(|| Regex::new(r#""method"\s*:\s*"([^"]+)""#).unwrap()); -static TOOL_REGEX: LazyLock = LazyLock::new(|| Regex::new(r#""name"\s*:\s*"([^"]+)""#).unwrap()); +static METHOD_REGEX: LazyLock = + LazyLock::new(|| Regex::new(r#""method"\s*:\s*"([^"]+)""#).unwrap()); +static TOOL_REGEX: LazyLock = + LazyLock::new(|| Regex::new(r#""name"\s*:\s*"([^"]+)""#).unwrap()); static ERROR_REGEX: LazyLock = LazyLock::new(|| Regex::new(r#""error"\s*:\s*\{"#).unwrap()); -static IS_ERROR_REGEX: LazyLock = LazyLock::new(|| Regex::new(r#""isError"\s*:\s*true"#).unwrap()); +static IS_ERROR_REGEX: LazyLock = + LazyLock::new(|| Regex::new(r#""isError"\s*:\s*true"#).unwrap()); pub fn extract_log_prefix(json_str: &str, is_response: bool) -> String { - let id = ID_REGEX.captures(json_str).and_then(|c| c.get(1)).map(|m| m.as_str()).unwrap_or("null"); - + let id = ID_REGEX + .captures(json_str) + .and_then(|c| c.get(1)) + .map(|m| m.as_str()) + .unwrap_or("null"); + if is_response { let is_error = ERROR_REGEX.is_match(json_str) || IS_ERROR_REGEX.is_match(json_str); return format!("Response id={} [Error: {}]", id, is_error); } - - let method = METHOD_REGEX.captures(json_str).and_then(|c| c.get(1)).map(|m| m.as_str()).unwrap_or(""); - + + let method = METHOD_REGEX + .captures(json_str) + .and_then(|c| c.get(1)) + .map(|m| m.as_str()) + .unwrap_or(""); + if method == "tools/call" { - let tool = TOOL_REGEX.captures(json_str).and_then(|c| c.get(1)).map(|m| m.as_str()).unwrap_or("unknown"); + let tool = TOOL_REGEX + .captures(json_str) + .and_then(|c| c.get(1)) + .map(|m| m.as_str()) + .unwrap_or("unknown"); format!("ToolCall[{}] id={}", tool, id) } else if !method.is_empty() { format!("Request[{}] id={}", method, id) diff --git a/stub/src/main.rs b/stub/src/main.rs index 14943cb..a5d7daa 100644 --- a/stub/src/main.rs +++ b/stub/src/main.rs @@ -10,8 +10,6 @@ struct Cli { target: String, } - - mod logger; fn init_logging(app_name: &str) -> Option { @@ -53,7 +51,9 @@ fn main() -> Result<(), Box> { let port = std::env::var("MCP_PORT").unwrap_or_else(|_| "3000".to_string()); format!("http://127.0.0.1:{}", port) }; - let ws_url = target_url.replace("http://", "ws://").replace("https://", "wss://"); + let ws_url = target_url + .replace("http://", "ws://") + .replace("https://", "wss://"); let ws_url = format!("{}/ws?client=proxy", ws_url); loop { @@ -63,7 +63,7 @@ fn main() -> Result<(), Box> { } tracing::info!("Attempting to connect to {}", ws_url); - + use tokio_tungstenite::tungstenite::client::IntoClientRequest; let request = match ws_url.clone().into_client_request() { Ok(req) => req, @@ -83,8 +83,21 @@ fn main() -> Result<(), Box> { let mut send_task = tokio::spawn(async move { while let Ok(msg) = rx.recv().await { let log_prefix = logger::extract_log_prefix(&msg, false); - tracing::info!(">>> [Stub] Forwarding {} to server (length: {}): {}", log_prefix, msg.len(), if msg.len() > 1000 { format!("{}...", &msg[..1000]) } else { msg.clone() }); - if write.send(tokio_tungstenite::tungstenite::Message::Text(msg)).await.is_err() { + tracing::info!( + ">>> [Stub] Forwarding {} to server (length: {}): {}", + log_prefix, + msg.len(), + if msg.len() > 1000 { + format!("{}...", &msg[..1000]) + } else { + msg.clone() + } + ); + if write + .send(tokio_tungstenite::tungstenite::Message::Text(msg)) + .await + .is_err() + { tracing::error!("Failed to write to websocket"); break; } @@ -95,7 +108,16 @@ fn main() -> Result<(), Box> { while let Some(Ok(msg)) = read.next().await { if let tokio_tungstenite::tungstenite::Message::Text(text) = msg { let log_prefix = logger::extract_log_prefix(&text, true); - tracing::info!("<<< [Stub] Received {} from server (length: {}): {}", log_prefix, text.len(), if text.len() > 1000 { format!("{}...", &text[..1000]) } else { text.clone() }); + tracing::info!( + "<<< [Stub] Received {} from server (length: {}): {}", + log_prefix, + text.len(), + if text.len() > 1000 { + format!("{}...", &text[..1000]) + } else { + text.clone() + } + ); use tokio::io::AsyncWriteExt; let mut stdout = tokio::io::stdout(); let _ = stdout.write_all(text.as_bytes()).await; @@ -107,14 +129,14 @@ fn main() -> Result<(), Box> { }); tokio::select! { - _ = shutdown_rx.recv() => { + _ = shutdown_rx.recv() => { tokio::time::sleep(tokio::time::Duration::from_millis(50)).await; tracing::info!("Shutdown received while connected"); break; } _ = &mut send_task => { tracing::error!("Send task exited"); - tokio::time::sleep(tokio::time::Duration::from_millis(50)).await; + tokio::time::sleep(tokio::time::Duration::from_millis(50)).await; recv_task.abort(); break; } @@ -134,4 +156,3 @@ fn main() -> Result<(), Box> { Ok(()) }) } - diff --git a/stub/tests/e2e.rs b/stub/tests/e2e.rs index ed6c4cc..c391c13 100644 --- a/stub/tests/e2e.rs +++ b/stub/tests/e2e.rs @@ -61,17 +61,19 @@ async fn test_full_system_e2e_performance() { assert!(stub_exe.exists(), "Stub not found at {:?}", stub_exe); // 1. Start Server - let _server = ChildGuard(Command::new(&server_exe) - .arg("--daemon") - .env("MCP_PORT", test_port) - .env("RUST_LOG", "debug") - .env("MCP_MEMORY_STORE_DIR", temp_dir.to_str().unwrap()) - .env("MCP_AUTH_TOKEN", test_auth_token) - .env("RUST_LOG", "debug") - .stdout(Stdio::inherit()) - .stderr(Stdio::inherit()) - .spawn() - .expect("Failed to start server")); + let _server = ChildGuard( + Command::new(&server_exe) + .arg("--daemon") + .env("MCP_PORT", test_port) + .env("RUST_LOG", "debug") + .env("MCP_MEMORY_STORE_DIR", temp_dir.to_str().unwrap()) + .env("MCP_AUTH_TOKEN", test_auth_token) + .env("RUST_LOG", "debug") + .stdout(Stdio::inherit()) + .stderr(Stdio::inherit()) + .spawn() + .expect("Failed to start server"), + ); // Give server time to generate TLS cert and start let client = reqwest::Client::builder() @@ -94,28 +96,32 @@ async fn test_full_system_e2e_performance() { assert!(started, "Server failed to start in time"); // 2. Start Stub - let mut stub = ChildGuard(Command::new(&stub_exe) - .arg("--target") - .arg(format!("http://127.0.0.1:{}", test_port)) - .env("MCP_MEMORY_STORE_DIR", temp_dir.to_str().unwrap()) - .env("MCP_AUTH_TOKEN", test_auth_token) - .env("RUST_LOG", "debug") - .stdin(Stdio::piped()) - .stdout(Stdio::piped()) - .stderr(Stdio::inherit()) - .spawn() - .expect("Failed to start stub")); + let mut stub = ChildGuard( + Command::new(&stub_exe) + .arg("--target") + .arg(format!("http://127.0.0.1:{}", test_port)) + .env("MCP_MEMORY_STORE_DIR", temp_dir.to_str().unwrap()) + .env("MCP_AUTH_TOKEN", test_auth_token) + .env("RUST_LOG", "debug") + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::inherit()) + .spawn() + .expect("Failed to start stub"), + ); let mut stub_stdin = stub.0.stdin.take().unwrap(); let mut stub_stdout = BufReader::new(stub.0.stdout.take().unwrap()); // 3. Start Nvim Bridge - let mut nvim = ChildGuard(Command::new(&nvim_exe) - .stdin(Stdio::piped()) - .stdout(Stdio::piped()) - .stderr(Stdio::inherit()) - .spawn() - .expect("Failed to start nvim bridge")); + let mut nvim = ChildGuard( + Command::new(&nvim_exe) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::inherit()) + .spawn() + .expect("Failed to start nvim bridge"), + ); let mut nvim_stdin = nvim.0.stdin.take().unwrap(); let mut nvim_stdout = BufReader::new(nvim.0.stdout.take().unwrap()); diff --git a/win-nvim/build.rs b/win-nvim/build.rs index dc9a51c..c266efa 100644 --- a/win-nvim/build.rs +++ b/win-nvim/build.rs @@ -15,10 +15,6 @@ fn main() { .and_then(|out| String::from_utf8(out.stdout).ok()) .unwrap_or_else(|| "unknown".to_string()); - let version = format!( - "{} ({})", - git_date.trim(), - git_hash.trim() - ); + let version = format!("{} ({})", git_date.trim(), git_hash.trim()); println!("cargo:rustc-env=APP_VERSION={}", version); } diff --git a/win-nvim/tests/integration_test.rs b/win-nvim/tests/integration_test.rs index 19661f4..a80c1f6 100644 --- a/win-nvim/tests/integration_test.rs +++ b/win-nvim/tests/integration_test.rs @@ -1,5 +1,5 @@ use serde_json::{json, Value}; -use std::io::{BufRead, BufReader, Read, Write}; +use std::io::{BufRead, BufReader, Write}; use std::process::{Command, Stdio}; fn send_message(stdin: &mut std::process::ChildStdin, msg: Value) {