diff --git a/linux-nvim/build.rs b/linux-nvim/build.rs index 0e0609c..2a2c799 100644 --- a/linux-nvim/build.rs +++ b/linux-nvim/build.rs @@ -2,19 +2,24 @@ use std::process::Command; fn main() { let git_hash = Command::new("git") - .args(&["rev-parse", "--short", "HEAD"]) - .output() - .ok() - .and_then(|out| String::from_utf8(out.stdout).ok()) - .unwrap_or_else(|| "unknown".to_string()); - - let git_date = Command::new("git") - .args(&["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"]) + .args(["rev-parse", "--short", "HEAD"]) .output() .ok() .and_then(|out| String::from_utf8(out.stdout).ok()) .unwrap_or_else(|| "unknown".to_string()); - let version = format!("{} ({} {})", env!("CARGO_PKG_VERSION"), git_date.trim(), git_hash.trim()); + let git_date = Command::new("git") + .args(["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"]) + .output() + .ok() + .and_then(|out| String::from_utf8(out.stdout).ok()) + .unwrap_or_else(|| "unknown".to_string()); + + let version = format!( + "{} ({} {})", + env!("CARGO_PKG_VERSION"), + 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 6f0a00b..2619403 100644 --- a/linux-nvim/tests/integration_test.rs +++ b/linux-nvim/tests/integration_test.rs @@ -1,6 +1,5 @@ -use serde_json::{json, Value}; +use serde_json::Value; use std::io::{BufRead, BufReader, Read, Write}; -use std::process::{Command, Stdio}; fn send_message(stdin: &mut std::process::ChildStdin, msg: Value) { let s = serde_json::to_string(&msg).unwrap(); @@ -12,7 +11,7 @@ fn send_message(stdin: &mut std::process::ChildStdin, msg: Value) { fn read_message(stdout: &mut std::process::ChildStdout) -> Option { let mut reader = BufReader::new(stdout); let mut length = 0; - + // Read headers loop { let mut line = String::new(); @@ -27,16 +26,16 @@ fn read_message(stdout: &mut std::process::ChildStdout) -> Option { length = len_str.parse().unwrap_or(0); } } - + if length == 0 { return None; } - + // Read body let mut buf = vec![0u8; length]; reader.read_exact(&mut buf).unwrap(); let body_str = String::from_utf8_lossy(&buf); - + Some(serde_json::from_str(&body_str).unwrap()) } @@ -46,7 +45,10 @@ fn test_mcp_initialization_and_tools_list() { let mut nvim_exe = std::env::current_exe().unwrap(); nvim_exe.pop(); nvim_exe.pop(); - nvim_exe.push(format!("mcp-memory-linux-nvim{}", std::env::consts::EXE_SUFFIX)); + nvim_exe.push(format!( + "mcp-memory-linux-nvim{}", + std::env::consts::EXE_SUFFIX + )); let mut child = Command::new(&nvim_exe) .stdin(Stdio::piped()) @@ -88,12 +90,12 @@ fn test_mcp_initialization_and_tools_list() { let s = serde_json::to_string(&init_req).unwrap(); stdin.write_all(format!("{}\n", s).as_bytes()).unwrap(); stdin.flush().unwrap(); - + let init_resp = read_message(&mut stdout).expect("Failed to read initialize response"); - + assert_eq!(init_resp["jsonrpc"], "2.0"); assert_eq!(init_resp["id"], 1); - + // Verify capabilities let capabilities = &init_resp["result"]["capabilities"]; assert_eq!(capabilities["tools"], serde_json::json!({})); @@ -105,17 +107,19 @@ fn test_mcp_initialization_and_tools_list() { "params": {}, "id": 2 }); - + send_message(&mut stdin, tools_req); - + let tools_resp = read_message(&mut stdout).expect("Failed to read tools/list response"); - + assert_eq!(tools_resp["jsonrpc"], "2.0"); assert_eq!(tools_resp["id"], 2); - - let tools = tools_resp["result"]["tools"].as_array().expect("result.tools must be an array"); + + let tools = tools_resp["result"]["tools"] + .as_array() + .expect("result.tools must be an array"); assert!(!tools.is_empty(), "Server must expose at least one tool"); - + let has_get_active_buffer = tools.iter().any(|t| t["name"] == "nvim_get_active_buffer"); assert!(has_get_active_buffer, "Missing nvim_get_active_buffer tool"); diff --git a/nvim-core/src/lib.rs b/nvim-core/src/lib.rs index 3a7bc38..fc1231f 100644 --- a/nvim-core/src/lib.rs +++ b/nvim-core/src/lib.rs @@ -20,19 +20,25 @@ pub struct JsonRpcResponse { pub error: Option, } -pub async fn read_message(stdin: &mut BufReader) -> Option { +pub async fn read_message( + stdin: &mut BufReader, +) -> Option { let mut length = 0; loop { let mut line = String::new(); if stdin.read_line(&mut line).await.unwrap_or(0) == 0 { return None; } - + if line.starts_with('{') { return match serde_json::from_str::(line.trim_end()) { Ok(req) => Some(req), Err(e) => { - tracing::error!("Failed to parse JSON-RPC request from JSONL: {}. Payload: {}", e, line); + tracing::error!( + "Failed to parse JSON-RPC request from JSONL: {}. Payload: {}", + e, + line + ); None } }; @@ -52,13 +58,21 @@ pub async fn read_message(stdin: &mut BufReader } let mut buffer = vec![0; length]; stdin.read_exact(&mut buffer).await.unwrap_or(0); - + serde_json::from_slice(&buffer).ok() } pub async fn send_response(response: JsonRpcResponse) { let msg = serde_json::to_string(&response).unwrap(); - tracing::info!("Sending JSON-RPC response (id: {:?}): {}", response.id, if msg.len() > 500 { format!("{}...", &msg[..500]) } else { msg.clone() }); + tracing::info!( + "Sending JSON-RPC response (id: {:?}): {}", + response.id, + if msg.len() > 500 { + format!("{}...", &msg[..500]) + } else { + msg.clone() + } + ); // CRITICAL ARCHITECTURAL DECISION: // The MCP StdioTransport MUST use Newline-Delimited JSON (NDJSON). // Do NOT use LSP-style HTTP headers (e.g. Content-Length). @@ -75,14 +89,16 @@ pub async fn send_error(id: Value, code: i32, message: &str) { id, result: None, error: Some(serde_json::json!({"code": code, "message": message})), - }).await; + }) + .await; } #[cfg(windows)] async fn get_socket_path() -> Result { - let profile = std::env::var("USERPROFILE").unwrap_or_else(|_| "C:\\Users\\reazul.ashraf".into()); + let profile = + std::env::var("USERPROFILE").unwrap_or_else(|_| "C:\\Users\\reazul.ashraf".into()); let path = format!("{}\\.gemini\\active_nvim.txt", profile); - + if let Ok(content) = std::fs::read_to_string(&path) { let p = content.trim().to_string(); if !p.is_empty() { @@ -120,7 +136,7 @@ async fn get_socket_path() -> Result { } } } - + if let Ok(entries) = std::fs::read_dir("/tmp") { for entry in entries.flatten() { if let Ok(name) = entry.file_name().into_string() { @@ -138,20 +154,28 @@ async fn get_socket_path() -> Result { #[cfg(windows)] async fn call_nvim(req: rmpv::Value) -> Result { use tokio::net::windows::named_pipe::ClientOptions; - + let msgid = if let rmpv::Value::Array(ref arr) = req { - if arr.len() > 1 { arr[1].clone() } else { rmpv::Value::Nil } - } else { rmpv::Value::Nil }; + if arr.len() > 1 { + arr[1].clone() + } else { + rmpv::Value::Nil + } + } else { + rmpv::Value::Nil + }; tracing::info!("Connecting to neovim pipe"); let socket_path = get_socket_path().await?; - let mut client = ClientOptions::new().open(&socket_path).map_err(|e| e.to_string())?; - + let mut client = ClientOptions::new() + .open(&socket_path) + .map_err(|e| e.to_string())?; + let mut buf = Vec::new(); rmpv::encode::write_value(&mut buf, &req).map_err(|e| e.to_string())?; tracing::info!("Sending RPC request to neovim (msgid: {})", msgid); client.write_all(&buf).await.map_err(|e| e.to_string())?; - + let mut resp_buf = Vec::new(); let mut chunk = vec![0u8; 8192]; let mut offset = 0; @@ -161,20 +185,23 @@ async fn call_nvim(req: rmpv::Value) -> Result { match rmpv::decode::read_value(&mut cursor) { Ok(val) => { offset += cursor.position() as usize; - + if let rmpv::Value::Array(ref arr) = val { - if arr.len() >= 4 && arr[0] == rmpv::Value::Integer(1.into()) && arr[1] == msgid { + if arr.len() >= 4 && arr[0] == rmpv::Value::Integer(1.into()) && arr[1] == msgid + { tracing::info!("Received RPC response from neovim (msgid: {})", msgid); return Ok(val); } } continue; - }, + } Err(_) => { let read_future = client.read(&mut chunk); match tokio::time::timeout(tokio::time::Duration::from_secs(5), read_future).await { Ok(Ok(n)) => { - if n == 0 { return Err("Connection closed".into()); } + if n == 0 { + return Err("Connection closed".into()); + } resp_buf.extend_from_slice(&chunk[..n]); } Ok(Err(e)) => return Err(e.to_string()), @@ -193,18 +220,26 @@ async fn call_nvim(req: rmpv::Value) -> Result { use tokio::net::UnixStream; let msgid = if let rmpv::Value::Array(ref arr) = req { - if arr.len() > 1 { arr[1].clone() } else { rmpv::Value::Nil } - } else { rmpv::Value::Nil }; + if arr.len() > 1 { + arr[1].clone() + } else { + rmpv::Value::Nil + } + } else { + rmpv::Value::Nil + }; tracing::info!("Connecting to neovim socket"); let socket_path = get_socket_path().await?; - let mut stream = UnixStream::connect(socket_path).await.map_err(|e| e.to_string())?; - + let mut stream = UnixStream::connect(socket_path) + .await + .map_err(|e| e.to_string())?; + let mut buf = Vec::new(); rmpv::encode::write_value(&mut buf, &req).map_err(|e| e.to_string())?; tracing::info!("Sending RPC request to neovim (msgid: {})", msgid); stream.write_all(&buf).await.map_err(|e| e.to_string())?; - + let mut resp_buf = Vec::new(); let mut chunk = vec![0u8; 8192]; let mut offset = 0; @@ -214,20 +249,23 @@ async fn call_nvim(req: rmpv::Value) -> Result { match rmpv::decode::read_value(&mut cursor) { Ok(val) => { offset += cursor.position() as usize; - + if let rmpv::Value::Array(ref arr) = val { - if arr.len() >= 4 && arr[0] == rmpv::Value::Integer(1.into()) && arr[1] == msgid { + if arr.len() >= 4 && arr[0] == rmpv::Value::Integer(1.into()) && arr[1] == msgid + { tracing::info!("Received RPC response from neovim (msgid: {})", msgid); return Ok(val); } } continue; - }, + } Err(_) => { let read_future = stream.read(&mut chunk); match tokio::time::timeout(tokio::time::Duration::from_secs(5), read_future).await { Ok(Ok(n)) => { - if n == 0 { return Err("Connection closed".into()); } + if n == 0 { + return Err("Connection closed".into()); + } resp_buf.extend_from_slice(&chunk[..n]); } Ok(Err(e)) => return Err(e.to_string()), @@ -248,7 +286,7 @@ async fn send_nvim_command(cmd: &str) -> Result<(), String> { RmpValue::String("nvim_command".into()), RmpValue::Array(vec![RmpValue::String(cmd.into())]), ]); - + let resp = call_nvim(req).await?; if let RmpValue::Array(arr) = resp { if !arr[2].is_nil() { @@ -272,7 +310,7 @@ async fn get_nvim_active_buffer() -> Result { RmpValue::Boolean(true), ]), ]); - + let resp = call_nvim(req).await?; if let RmpValue::Array(arr) = resp { if !arr[2].is_nil() { @@ -300,11 +338,9 @@ async fn get_nvim_cursor() -> Result { RmpValue::Integer(0.into()), RmpValue::Integer(3.into()), // msgid RmpValue::String("nvim_win_get_cursor".into()), - RmpValue::Array(vec![ - RmpValue::Integer(0.into()), - ]), + RmpValue::Array(vec![RmpValue::Integer(0.into())]), ]); - + let resp = call_nvim(req).await?; if let RmpValue::Array(arr) = resp { if !arr[2].is_nil() { @@ -332,7 +368,7 @@ async fn get_nvim_visual_selection() -> Result { return lines end "#; - + use rmpv::Value as RmpValue; let req = RmpValue::Array(vec![ RmpValue::Integer(0.into()), @@ -343,7 +379,7 @@ async fn get_nvim_visual_selection() -> Result { RmpValue::Array(vec![]), ]), ]); - + let resp = call_nvim(req).await?; if let RmpValue::Array(arr) = resp { if !arr[2].is_nil() { @@ -360,7 +396,8 @@ async fn get_nvim_visual_selection() -> Result { async fn set_nvim_diagnostics(line: i64, message: &str) -> Result<(), String> { let escaped_message = message.replace("\\", "\\\\").replace("\"", "\\\""); - let lua_script = format!(r#" + let lua_script = format!( + r#" local ns = vim.api.nvim_create_namespace("gemini_diagnostics") local diagnostics = {{{{ lnum = {} - 1, @@ -369,8 +406,10 @@ async fn set_nvim_diagnostics(line: i64, message: &str) -> Result<(), String> { message = "{}", }}}} vim.diagnostic.set(ns, 0, diagnostics, {{}}) - "#, line, escaped_message); - + "#, + line, escaped_message + ); + use rmpv::Value as RmpValue; let req = RmpValue::Array(vec![ RmpValue::Integer(0.into()), @@ -381,7 +420,7 @@ async fn set_nvim_diagnostics(line: i64, message: &str) -> Result<(), String> { RmpValue::Array(vec![]), ]), ]); - + let resp = call_nvim(req).await?; if let RmpValue::Array(arr) = resp { if !arr[2].is_nil() { @@ -404,7 +443,7 @@ fn rmpv_to_json(val: &rmpv::Value) -> serde_json::Value { } else { serde_json::Value::Null } - }, + } rmpv::Value::F32(f) => serde_json::json!(f), rmpv::Value::F64(f) => serde_json::json!(f), rmpv::Value::String(s) => { @@ -413,11 +452,11 @@ fn rmpv_to_json(val: &rmpv::Value) -> serde_json::Value { } else { serde_json::Value::Null } - }, + } rmpv::Value::Array(arr) => { let vec: Vec = arr.iter().map(rmpv_to_json).collect(); serde_json::Value::Array(vec) - }, + } rmpv::Value::Map(map) => { let mut obj = serde_json::Map::new(); for (k, v) in map { @@ -429,7 +468,7 @@ fn rmpv_to_json(val: &rmpv::Value) -> serde_json::Value { obj.insert(key_str, rmpv_to_json(v)); } serde_json::Value::Object(obj) - }, + } _ => serde_json::json!(format!("{:?}", val)), } } @@ -440,12 +479,9 @@ async fn execute_nvim_lua(code: &str) -> Result { RmpValue::Integer(0.into()), RmpValue::Integer(6.into()), // msgid RmpValue::String("nvim_exec_lua".into()), - RmpValue::Array(vec![ - RmpValue::String(code.into()), - RmpValue::Array(vec![]), - ]), + RmpValue::Array(vec![RmpValue::String(code.into()), RmpValue::Array(vec![])]), ]); - + let resp = call_nvim(req).await?; if let RmpValue::Array(arr) = resp { if !arr[2].is_nil() { @@ -460,7 +496,12 @@ async fn execute_nvim_lua(code: &str) -> Result { } pub async fn run_mcp_loop(app_name: &str, app_version: &str) { if std::env::args().any(|arg| arg == "--version") { - println!("{} {} ({})", app_name, app_version, std::env::var("GIT_HASH").unwrap_or_else(|_| "unknown".to_string())); + println!( + "{} {} ({})", + app_name, + app_version, + std::env::var("GIT_HASH").unwrap_or_else(|_| "unknown".to_string()) + ); return; } let _guard = init_logging(app_name); @@ -471,7 +512,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { Some(m) => { tracing::info!("Received message method: {}", m.method); m - }, + } None => { tracing::info!("Stdin closed, exiting loop"); break; @@ -488,9 +529,14 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { match msg.method.as_str() { "initialize" => { let init = rmcp::model::InitializeResult::new( - rmcp::model::ServerCapabilities::builder().enable_tools().build() + rmcp::model::ServerCapabilities::builder() + .enable_tools() + .build(), ) - .with_server_info(rmcp::model::Implementation::new(app_name.clone(), app_version.clone())) + .with_server_info(rmcp::model::Implementation::new( + app_name.clone(), + app_version.clone(), + )) .with_protocol_version(rmcp::model::ProtocolVersion::V_2024_11_05); send_response(JsonRpcResponse { @@ -498,7 +544,8 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { id, result: Some(serde_json::to_value(init).unwrap()), error: None, - }).await; + }) + .await; } "notifications/initialized" => {} "tools/list" => { @@ -595,7 +642,10 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { match name { "nvim_goto_line" => { - if let (Some(file), Some(line)) = (args.get("file").and_then(|v| v.as_str()), args.get("line").and_then(|v| v.as_i64())) { + if let (Some(file), Some(line)) = ( + args.get("file").and_then(|v| v.as_str()), + args.get("line").and_then(|v| v.as_i64()), + ) { let escaped_file = file.replace("\\", "\\\\"); let cmd = format!("e {} | {} | normal! zz", escaped_file, line); match send_nvim_command(&cmd).await { @@ -615,53 +665,53 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { send_error(id, -32602, "Missing 'file' or 'line'").await; } } - "nvim_get_active_buffer" => { - match get_nvim_active_buffer().await { - Ok(content) => { - send_response(JsonRpcResponse { - jsonrpc: "2.0".to_string(), - id, - result: Some(json!({ - "content": [{"type": "text", "text": content}] - })), - error: None, - }).await; - } - Err(e) => send_error(id, -32603, &e).await, + "nvim_get_active_buffer" => match get_nvim_active_buffer().await { + Ok(content) => { + send_response(JsonRpcResponse { + jsonrpc: "2.0".to_string(), + id, + result: Some(json!({ + "content": [{"type": "text", "text": content}] + })), + error: None, + }) + .await; } - } - "nvim_get_cursor" => { - match get_nvim_cursor().await { - Ok(content) => { - send_response(JsonRpcResponse { - jsonrpc: "2.0".to_string(), - id, - result: Some(json!({ - "content": [{"type": "text", "text": content}] - })), - error: None, - }).await; - } - Err(e) => send_error(id, -32603, &e).await, + Err(e) => send_error(id, -32603, &e).await, + }, + "nvim_get_cursor" => match get_nvim_cursor().await { + Ok(content) => { + send_response(JsonRpcResponse { + jsonrpc: "2.0".to_string(), + id, + result: Some(json!({ + "content": [{"type": "text", "text": content}] + })), + error: None, + }) + .await; } - } - "nvim_get_visual_selection" => { - match get_nvim_visual_selection().await { - Ok(content) => { - send_response(JsonRpcResponse { - jsonrpc: "2.0".to_string(), - id, - result: Some(json!({ - "content": [{"type": "text", "text": content}] - })), - error: None, - }).await; - } - Err(e) => send_error(id, -32603, &e).await, + Err(e) => send_error(id, -32603, &e).await, + }, + "nvim_get_visual_selection" => match get_nvim_visual_selection().await { + Ok(content) => { + send_response(JsonRpcResponse { + jsonrpc: "2.0".to_string(), + id, + result: Some(json!({ + "content": [{"type": "text", "text": content}] + })), + error: None, + }) + .await; } - } + Err(e) => send_error(id, -32603, &e).await, + }, "nvim_set_diagnostics" => { - if let (Some(line), Some(message)) = (args.get("line").and_then(|v| v.as_i64()), args.get("message").and_then(|v| v.as_str())) { + if let (Some(line), Some(message)) = ( + args.get("line").and_then(|v| v.as_i64()), + args.get("message").and_then(|v| v.as_str()), + ) { match set_nvim_diagnostics(line, message).await { Ok(_) => { send_response(JsonRpcResponse { @@ -700,7 +750,8 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { "content": [{"type": "text", "text": result}] })), error: None, - }).await; + }) + .await; } Err(e) => send_error(id, -32603, &e).await, } @@ -730,7 +781,8 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { "content": [{"type": "text", "text": result}] })), error: None, - }).await; + }) + .await; } Err(e) => send_error(id, -32603, &e).await, } @@ -746,7 +798,8 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { "content": [{"type": "text", "text": result}] })), error: None, - }).await; + }) + .await; } Err(e) => send_error(id, -32603, &e).await, } @@ -770,18 +823,20 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { } fn init_logging(app_name: &str) -> tracing_appender::non_blocking::WorkerGuard { - let log_dir = dirs::home_dir().unwrap_or_default().join(".gemini/mcp_memory/logs"); + let log_dir = dirs::home_dir() + .unwrap_or_default() + .join(".gemini/mcp_memory/logs"); std::fs::create_dir_all(&log_dir).unwrap_or_default(); - + let file_appender = tracing_appender::rolling::daily(log_dir, format!("{}.log", app_name)); let (non_blocking, guard) = tracing_appender::non_blocking(file_appender); - + let _ = tracing_subscriber::fmt() .with_writer(non_blocking) .with_ansi(false) .with_max_level(tracing::Level::INFO) .try_init(); - + guard } @@ -794,7 +849,10 @@ mod tests { assert_eq!(rmpv_to_json(&rmpv::Value::Nil), serde_json::Value::Null); assert_eq!(rmpv_to_json(&rmpv::Value::Boolean(true)), json!(true)); assert_eq!(rmpv_to_json(&rmpv::Value::Integer(42.into())), json!(42)); - assert_eq!(rmpv_to_json(&rmpv::Value::String("hello".into())), json!("hello")); + assert_eq!( + rmpv_to_json(&rmpv::Value::String("hello".into())), + json!("hello") + ); } #[test] @@ -809,9 +867,12 @@ mod tests { #[test] fn test_rmpv_to_json_map() { let mut map = vec![]; - map.push((rmpv::Value::String("key1".into()), rmpv::Value::Integer(100.into()))); + map.push(( + 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); assert_eq!(json_map, json!({ "key1": 100 })); } diff --git a/server/build.rs b/server/build.rs index 0e0609c..2a2c799 100644 --- a/server/build.rs +++ b/server/build.rs @@ -2,19 +2,24 @@ use std::process::Command; fn main() { let git_hash = Command::new("git") - .args(&["rev-parse", "--short", "HEAD"]) - .output() - .ok() - .and_then(|out| String::from_utf8(out.stdout).ok()) - .unwrap_or_else(|| "unknown".to_string()); - - let git_date = Command::new("git") - .args(&["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"]) + .args(["rev-parse", "--short", "HEAD"]) .output() .ok() .and_then(|out| String::from_utf8(out.stdout).ok()) .unwrap_or_else(|| "unknown".to_string()); - let version = format!("{} ({} {})", env!("CARGO_PKG_VERSION"), git_date.trim(), git_hash.trim()); + let git_date = Command::new("git") + .args(["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"]) + .output() + .ok() + .and_then(|out| String::from_utf8(out.stdout).ok()) + .unwrap_or_else(|| "unknown".to_string()); + + let version = format!( + "{} ({} {})", + env!("CARGO_PKG_VERSION"), + git_date.trim(), + git_hash.trim() + ); println!("cargo:rustc-env=APP_VERSION={}", version); } diff --git a/server/src/bin_test.rs b/server/src/bin_test.rs index fe21d4e..8538587 100644 --- a/server/src/bin_test.rs +++ b/server/src/bin_test.rs @@ -1,8 +1,10 @@ use rmcp::model::{InitializeResult, ServerCapabilities}; fn main() { - let init = InitializeResult::new( - ServerCapabilities::builder().enable_tools().build() - ).with_server_info(rmcp::model::Implementation::new("gemini-mcp-memory", "3.0.0")); + let init = InitializeResult::new(ServerCapabilities::builder().enable_tools().build()) + .with_server_info(rmcp::model::Implementation::new( + "gemini-mcp-memory", + "3.0.0", + )); println!("{}", serde_json::to_string_pretty(&init).unwrap()); } diff --git a/server/src/handlers.rs b/server/src/handlers.rs index 3852ac4..c150f12 100644 --- a/server/src/handlers.rs +++ b/server/src/handlers.rs @@ -18,8 +18,6 @@ macro_rules! parse_tool { }; } - - use serde::de::DeserializeOwned; use std::collections::HashSet; use std::sync::Arc; @@ -37,7 +35,7 @@ impl MemoryHandler { pub async fn handle_request(&self, req: serde_json::Value) -> Option { let id = req.get("id").cloned().unwrap_or(serde_json::Value::Null); let method = req.get("method").and_then(|m| m.as_str()).unwrap_or(""); - + tracing::debug!(">>> [Server] Handling MCP request method: {}", method); tracing::trace!(">>> [Server] Full MCP Request payload: {}", req.to_string()); let response = match method { @@ -57,92 +55,276 @@ impl MemoryHandler { } } }); - tracing::debug!("<<< [Server] Replying to server/discover with payload: {}", payload.to_string()); + tracing::debug!( + "<<< [Server] Replying to server/discover with payload: {}", + payload.to_string() + ); Some(crate::mcp::success(id, payload)) } "initialize" => { let init = rmcp::model::InitializeResult::new( - rmcp::model::ServerCapabilities::builder().enable_tools().build() - ).with_server_info(rmcp::model::Implementation::new("gemini-mcp-memory", "3.0.0")); + rmcp::model::ServerCapabilities::builder() + .enable_tools() + .build(), + ) + .with_server_info(rmcp::model::Implementation::new( + "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())) + Some(crate::mcp::success( + id, + serde_json::to_value(&init).unwrap(), + )) } - "notifications/initialized" => { - None - } + "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::("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::( + "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::( + "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", "", ), - crate::mcp::tool_def::("register_environment", "Register details about a specific deployment environment."), + crate::mcp::tool_def::( + "register_environment", + "Register details about a specific deployment environment.", + ), crate::mcp::tool_def::( "get_environment_details", "", ), - 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::( + "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", "", ), - 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."), + 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.", + ), ]; Some(crate::mcp::success( id, @@ -157,16 +339,18 @@ crate::mcp::tool_def::("create_entities", "Create new entiti .cloned() .unwrap_or(serde_json::Value::Object(Default::default())); - self.state.broadcast_activity(&format!("Agent executed tool: {}", name)); + self.state + .broadcast_activity(&format!("Agent executed tool: {}", name)); let result: Result = match name { - "query_graph_path" => { + "query_graph_path" => { let req = parse_tool!(args.clone(), id, crate::tools::QueryGraphPathTool); let graph = self.state.get_full_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(); + let mut parents: std::collections::HashMap = + std::collections::HashMap::new(); queue.push_back(req.start_node.clone()); visited.insert(req.start_node.clone()); @@ -186,12 +370,21 @@ crate::mcp::tool_def::("create_entities", "Create new entiti 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())); + 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))); + parents.insert( + rel.from.clone(), + ( + current.clone(), + format!("inverse({})", rel.relation_type), + ), + ); queue.push_back(rel.from.clone()); nodes_at_next_depth += 1; } @@ -215,10 +408,13 @@ crate::mcp::tool_def::("create_entities", "Create new entiti 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)) + Ok(format!( + "No path found between {} and {} within depth {}", + req.start_node, req.end_node, max_depth + )) } } -"create_entities" => { + "create_entities" => { let req = parse_tool!(args.clone(), id, CreateEntitiesTool); self.state.write_to_local_delta(|g| { for entity in req.entities { @@ -248,8 +444,7 @@ crate::mcp::tool_def::("create_entities", "Create new entiti let full = self.state.get_full_graph(); self.state.write_to_local_delta(|g| { for o in req.observations { - if let Some(full_e) = full.entities.get(&o.entity_name) - { + if let Some(full_e) = full.entities.get(&o.entity_name) { let mut e = g.entities.get(&o.entity_name).cloned().unwrap_or_else( || Entity { @@ -284,8 +479,7 @@ crate::mcp::tool_def::("create_entities", "Create new entiti let req = parse_tool!(args.clone(), id, DeleteObservationsTool); self.state.apply_sync_write(|master| { for d in req.deletions { - if let Some(e) = master.entities.get_mut(&d.entity_name) - { + 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)); } @@ -335,9 +529,10 @@ crate::mcp::tool_def::("create_entities", "Create new entiti let full = self.state.get_full_graph(); for (id, doc_type) in matches { if doc_type == "entity" - && let Some(e) = full.entities.get(&id) { - result.entities.insert(id, e.clone()); - } + && let Some(e) = full.entities.get(&id) + { + result.entities.insert(id, e.clone()); + } } let data = serde_json::to_string(&result).unwrap_or_default(); Ok(vec![data.to_string()][0].clone()) @@ -490,10 +685,10 @@ crate::mcp::tool_def::("create_entities", "Create new entiti .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, @@ -502,7 +697,7 @@ crate::mcp::tool_def::("create_entities", "Create new entiti created_at: now, updated_at: now, git_branch: req.git_branch, - parent_id: parent_id, + parent_id, dependencies: deps, acceptance_criteria: vec![], }; @@ -522,26 +717,29 @@ crate::mcp::tool_def::("create_entities", "Create new entiti // 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 { - if to_delete.contains(pid) && !to_delete.contains(&t.id) { + 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()) + Ok(vec![ + format!("Deleted task and its children ({} total).", deleted_count) + .to_string(), + ][0] + .clone()) } else { Ok(vec!["Task not found.".to_string()][0].clone()) } @@ -552,66 +750,84 @@ crate::mcp::tool_def::("create_entities", "Create new entiti 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) { + 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; } + + 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) { - if t.acceptance_criteria.iter().any(|c| !c.is_met) { + 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(); + 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) { - if dep_task.status != "completed" && dep_task.status != "done" { + 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(", ")); + 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)) { + 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(", ")); + 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(); + 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()]; @@ -619,7 +835,9 @@ crate::mcp::tool_def::("create_entities", "Create new entiti 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" { + if t.parent_id.as_ref() == Some(¤t_pid) + && t.status != "completed" + { t.status = target_status.clone(); to_cancel.push(t.id.clone()); } @@ -629,9 +847,15 @@ crate::mcp::tool_def::("create_entities", "Create new entiti } } }); - + if blocked { - Ok(vec![format!("Error: Cannot transition task. {}", blocker_details)].into_iter().next().unwrap()) + Ok(vec![format!( + "Error: Cannot transition task. {}", + blocker_details + )] + .into_iter() + .next() + .unwrap()) } else if found { Ok(vec!["Task status updated.".to_string()][0].clone()) } else { @@ -655,18 +879,30 @@ crate::mcp::tool_def::("create_entities", "Create new entiti let req = parse_tool!(args.clone(), 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(); + 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()) + Ok( + vec!["Acceptance criteria set successfully.".to_string()][0] + .clone(), + ) } else { Ok(vec!["Task not found.".to_string()][0].clone()) } @@ -676,24 +912,37 @@ crate::mcp::tool_def::("create_entities", "Create new entiti 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) { - if let Some(ac) = task.acceptance_criteria.iter_mut().find(|c| c.id == req.criteria || c.description == req.criteria) { + 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(); + 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()) + Ok(vec![format!( + "Acceptance criteria verified with proof: {}", + req.proof + )][0] + .clone()) } else if already_met { Ok(vec!["Acceptance criteria was already met.".to_string()][0].clone()) } else { - Ok(vec!["Acceptance criteria or task not found.".to_string()][0].clone()) + Ok( + vec!["Acceptance criteria or task not found.".to_string()][0] + .clone(), + ) } } "store_snippet" => { @@ -1224,9 +1473,10 @@ crate::mcp::tool_def::("create_entities", "Create new entiti let full = self.state.get_full_graph(); for (id, doc_type) in &matches { if doc_type == "entity" - && let Some(e) = full.entities.get(id) { - kg.entities.insert(id.clone(), e.clone()); - } + && 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") { @@ -1345,23 +1595,19 @@ crate::mcp::tool_def::("create_entities", "Create new entiti }; 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) => 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 }] + }), + )), } } _ => { @@ -1372,7 +1618,7 @@ crate::mcp::tool_def::("create_entities", "Create new entiti } } }; - + tracing::trace!("Returning response from handle_request: {:?}", response); response } @@ -1381,17 +1627,23 @@ crate::mcp::tool_def::("create_entities", "Create new entiti #[cfg(test)] mod tests { use super::*; - use std::sync::Arc; 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())); + 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); @@ -1401,7 +1653,9 @@ mod tests { 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()), + 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()), @@ -1419,7 +1673,8 @@ mod tests { 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, + context_workspaces: crate::store::Store::new("context_workspaces", db.clone()), + activity_tx: tokio::sync::broadcast::channel(100).0, }); let handler = MemoryHandler { state }; @@ -1437,22 +1692,29 @@ mod tests { } }); - let response = handler.handle_request(req).await.expect("Expected a response"); - + 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 + + // 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())); + 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()); @@ -1465,8 +1727,10 @@ mod tests { 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()), + + 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()), @@ -1484,7 +1748,8 @@ mod tests { 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, + context_workspaces: crate::store::Store::new("context_workspaces", db.clone()), + activity_tx: tokio::sync::broadcast::channel(100).0, }); MemoryHandler { state } } @@ -1500,15 +1765,26 @@ mod tests { "params": {} }); - let response = handler.handle_request(req).await.expect("Expected a response"); + 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"); + + 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."); + 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] @@ -1529,12 +1805,20 @@ mod tests { } }); - let response = handler.handle_request(req).await.expect("Expected a response"); + 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: ")); + 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(); @@ -1566,15 +1850,21 @@ mod tests { } }); - let response = handler.handle_request(req).await.expect("Expected a response"); + 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"); + 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()); @@ -1599,7 +1889,10 @@ mod tests { } }); - let response = handler.handle_request(req).await.expect("Expected a response"); + let response = handler + .handle_request(req) + .await + .expect("Expected a response"); assert_eq!(response["id"], 5); let snippets = handler.state.snippets.read(); @@ -1624,7 +1917,10 @@ mod tests { } }); - let response = handler.handle_request(req).await.expect("Expected a response"); + let response = handler + .handle_request(req) + .await + .expect("Expected a response"); assert_eq!(response["id"], 6); let notes = handler.state.sticky.read(); @@ -1666,13 +1962,16 @@ mod tests { 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, - }); + 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", @@ -1700,13 +1999,16 @@ mod tests { 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, - }); + 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.apply_sync_write(|_| {}); @@ -1731,13 +2033,16 @@ mod tests { 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, - }); + 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.apply_sync_write(|_| {}); let req = json!({ @@ -1797,7 +2102,10 @@ mod tests { description: "".to_string(), created_at: 0, updated_at: 0, - git_branch: None, acceptance_criteria: vec![], dependencies: vec![], parent_id: None, + git_branch: None, + acceptance_criteria: vec![], + dependencies: vec![], + parent_id: None, }); tasks.push(crate::models::Task { id: "2".to_string(), @@ -1806,10 +2114,13 @@ mod tests { description: "".to_string(), created_at: 0, updated_at: 0, - git_branch: None, acceptance_criteria: vec![], dependencies: vec![], parent_id: None, + git_branch: None, + acceptance_criteria: vec![], + dependencies: vec![], + parent_id: None, }); }); - + let req = json!({ "jsonrpc": "2.0", "id": 12, @@ -1819,7 +2130,7 @@ mod tests { "arguments": {} } }); - + let response = handler.handle_request(req).await.unwrap(); let content = response["result"]["content"][0]["text"].as_str().unwrap(); assert!(content.contains("Active Task")); @@ -1845,7 +2156,7 @@ mod tests { updated_at: 0, }); }); - + let req = json!({ "jsonrpc": "2.0", "id": 13, @@ -1872,7 +2183,7 @@ mod tests { timestamp: 0, }); }); - + let req = json!({ "jsonrpc": "2.0", "id": 14, @@ -1927,13 +2238,16 @@ mod tests { 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, - }); + 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.apply_sync_write(|_| {}); @@ -1962,7 +2276,9 @@ mod tests { git_branch: None, }; handler.state.graph.modify(|session| { - session.entities.insert("UserRepository".to_string(), entity); + session + .entities + .insert("UserRepository".to_string(), entity); }); handler.state.apply_sync_write(|_| {}); @@ -2322,7 +2638,10 @@ mod tests { description: "".to_string(), created_at: 0, updated_at: 0, - git_branch: None, acceptance_criteria: vec![], dependencies: vec![], parent_id: None, + git_branch: None, + acceptance_criteria: vec![], + dependencies: vec![], + parent_id: None, }); }); @@ -2343,7 +2662,3 @@ mod tests { assert_eq!(tasks[0].status, "in_progress"); } } - - - - diff --git a/server/src/main.rs b/server/src/main.rs index 4c9ec98..18f80f0 100644 --- a/server/src/main.rs +++ b/server/src/main.rs @@ -1,4 +1,7 @@ -#![cfg_attr(not(target_os = "windows"), allow(dead_code, unused_imports, unreachable_code))] +#![cfg_attr( + not(target_os = "windows"), + allow(dead_code, unused_imports, unreachable_code) +)] mod handlers; mod mcp; @@ -13,11 +16,11 @@ use crate::models::*; use crate::state::MemoryState; use crate::store::Store; +use redb::ReadableTable; use std::fs; use std::path::PathBuf; use std::sync::{Arc, RwLock}; use std::time::{Duration, SystemTime, UNIX_EPOCH}; -use redb::ReadableTable; use tokio::time::sleep; use clap::{Parser, Subcommand}; @@ -83,11 +86,10 @@ enum GateCommands { }, } - async fn reconcile_worker(state: Arc) { loop { sleep(Duration::from_secs(5)).await; - + let has_local = { let session = state.graph.read(); !session.entities.is_empty() || !session.relations.is_empty() @@ -106,14 +108,18 @@ async fn reconcile_worker(state: Arc) { let state_clone = state.clone(); let _ = tokio::task::spawn_blocking(move || { state_clone.rebuild_index(); - }).await; + }) + .await; } } } use axum::{ Json, Router, - extract::{Query, State, ws::{WebSocket, Message}}, + extract::{ + Query, State, + ws::{Message, WebSocket}, + }, response::IntoResponse, routing::{get, post}, }; @@ -186,9 +192,11 @@ async fn gate_verify_handler( (axum::http::StatusCode::FORBIDDEN, msg).into_response() } } - None => { - (axum::http::StatusCode::NOT_FOUND, "Action not yet authorized (no gate record found).").into_response() - } + None => ( + axum::http::StatusCode::NOT_FOUND, + "Action not yet authorized (no gate record found).", + ) + .into_response(), } } @@ -229,18 +237,23 @@ fn run_server(state: Arc) -> Result<(), Box> rt.block_on(async { tokio::spawn(reconcile_worker(Arc::clone(&state))); let app_state = Arc::new(AppState { - handler: Arc::new(MemoryHandler { state: Arc::clone(&state) }), + handler: Arc::new(MemoryHandler { + state: Arc::clone(&state), + }), clients: RwLock::new(HashMap::new()), next_id: AtomicUsize::new(1), }); let app = Router::new() - .route("/api/version", get(|| async move { - axum::Json(serde_json::json!({ - "version": env!("APP_VERSION"), - "git_hash": option_env!("GIT_HASH").unwrap_or("unknown") - })) - })) + .route( + "/api/version", + get(|| async move { + axum::Json(serde_json::json!({ + "version": env!("APP_VERSION"), + "git_hash": option_env!("GIT_HASH").unwrap_or("unknown") + })) + }), + ) .route("/ws", get(ws_handler)) .route("/health", get(health_handler)) .route("/nvim/telemetry", post(nvim_telemetry_handler)) @@ -260,65 +273,81 @@ fn run_server(state: Arc) -> Result<(), Box> "/", get(|| async move { axum::response::Html(include_str!("dashboard.html")) }), ) - .route("/api/graph", get({ - let state_clone = app_state.handler.state.clone(); - move || async move { - let graph = state_clone.get_full_graph(); - axum::Json(graph) - } - })) - - .route("/api/tasks/{id}/complete", post({ - let state_clone = app_state.handler.state.clone(); - move |axum::extract::Path(id): axum::extract::Path| async move { - state_clone.tasks.modify(|tasks| { - for t in tasks.iter_mut() { - if t.id == id { - t.status = "completed".to_string(); - break; - } - } - }); - axum::Json(serde_json::json!({"status": "success"})) - } - })) - - .route("/api/tasks", get({ - let state_clone = app_state.handler.state.clone(); - move || async move { - let tasks = state_clone.tasks.read(); - axum::Json(tasks.clone()) - } - })) - .route("/api/sticky", get({ - let state_clone = app_state.handler.state.clone(); - move || async move { - let sticky = state_clone.sticky.read(); - axum::Json(sticky.clone()) - } - })) - .route("/api/search", get({ - let state_clone = app_state.handler.state.clone(); - move |axum::extract::Query(params): axum::extract::Query>| async move { - if let Some(q) = params.get("q") { - if let Ok(idx) = state_clone.search_index.read() { - if let Ok(results) = idx.search(q, None) { - let mut formatted_results = Vec::new(); - for (type_name, content) in results { - formatted_results.push(serde_json::json!({ - "type_name": type_name, - "content": content, - "score": 1.0 - })); - } - return axum::Json(serde_json::json!({ "results": formatted_results })); - } - } + .route( + "/api/graph", + get({ + let state_clone = app_state.handler.state.clone(); + move || async move { + let graph = state_clone.get_full_graph(); + axum::Json(graph) } - axum::Json(serde_json::json!({ "results": [] })) - } - })) - .route("/api/stats", + }), + ) + .route( + "/api/tasks/{id}/complete", + post({ + let state_clone = app_state.handler.state.clone(); + move |axum::extract::Path(id): axum::extract::Path| async move { + state_clone.tasks.modify(|tasks| { + for t in tasks.iter_mut() { + if t.id == id { + t.status = "completed".to_string(); + break; + } + } + }); + axum::Json(serde_json::json!({"status": "success"})) + } + }), + ) + .route( + "/api/tasks", + get({ + let state_clone = app_state.handler.state.clone(); + move || async move { + let tasks = state_clone.tasks.read(); + axum::Json(tasks.clone()) + } + }), + ) + .route( + "/api/sticky", + get({ + let state_clone = app_state.handler.state.clone(); + move || async move { + let sticky = state_clone.sticky.read(); + axum::Json(sticky.clone()) + } + }), + ) + .route( + "/api/search", + get({ + let state_clone = app_state.handler.state.clone(); + move |axum::extract::Query(params): axum::extract::Query< + std::collections::HashMap, + >| async move { + if let Some(q) = params.get("q") + && let Ok(idx) = state_clone.search_index.read() + && let Ok(results) = idx.search(q, None) { + let mut formatted_results = Vec::new(); + for (type_name, content) in results { + formatted_results.push(serde_json::json!({ + "type_name": type_name, + "content": content, + "score": 1.0 + })); + } + return axum::Json( + serde_json::json!({ "results": formatted_results }), + ); + } + axum::Json(serde_json::json!({ "results": [] })) + } + }), + ) + .route( + "/api/stats", get({ let state_clone = app_state.handler.state.clone(); move || async move { @@ -374,7 +403,7 @@ fn run_server(state: Arc) -> Result<(), Box> tracing::info!("MCP Memory Server running on http://127.0.0.1:3000/sse"); let addr = std::env::var("MCP_PORT").unwrap_or_else(|_| "3000".to_string()); let addr: std::net::SocketAddr = format!("127.0.0.1:{}", addr).parse().unwrap(); - + let listener = tokio::net::TcpListener::bind(addr).await.unwrap(); if let Err(e) = axum::serve(listener, app.into_make_service()).await { let log_path = dirs::home_dir() @@ -386,29 +415,39 @@ fn run_server(state: Arc) -> Result<(), Box> }) } - - async fn ws_handler( ws: axum::extract::ws::WebSocketUpgrade, - headers: axum::http::HeaderMap, + _headers: axum::http::HeaderMap, axum::extract::State(state): axum::extract::State>, axum::extract::Query(query): axum::extract::Query>, ) -> axum::response::Response { - let client_type = query.get("client").cloned().unwrap_or_else(|| "unknown".to_string()); - ws.on_upgrade(move |socket| handle_socket(socket, state, client_type)).into_response() + let client_type = query + .get("client") + .cloned() + .unwrap_or_else(|| "unknown".to_string()); + ws.on_upgrade(move |socket| handle_socket(socket, state, client_type)) + .into_response() } async fn handle_socket(socket: WebSocket, state: Arc, client_type: String) { let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst)); let (tx, mut rx) = mpsc::channel::(100); - state.clients.write().unwrap().insert(session_id.clone(), tx.clone()); + state + .clients + .write() + .unwrap() + .insert(session_id.clone(), tx.clone()); let (mut sender, mut receiver) = socket.split(); let mut send_task = tokio::spawn(async move { while let Some(msg) = rx.recv().await { - tracing::trace!("Sending message to websocket (length: {}): {}", msg.len(), msg); + tracing::trace!( + "Sending message to websocket (length: {}): {}", + msg.len(), + msg + ); if sender.send(Message::Text(msg.into())).await.is_err() { tracing::error!("Failed to send message to websocket"); break; @@ -426,75 +465,102 @@ async fn handle_socket(socket: WebSocket, state: Arc, client_type: Str while let Some(msg_result) = receiver.next().await { match msg_result { Ok(Message::Text(text)) => { - tracing::info!("Received text message from websocket (length: {})", text.len()); + tracing::info!( + "Received text message from websocket (length: {})", + text.len() + ); tracing::trace!("Message content: {}", text); if let Ok(payload) = serde_json::from_str::(&text) { if client_type == "proxy" { // Send activity broadcast to UI clients - if let Some(method) = payload.get("method").and_then(|m| m.as_str()) { - if 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 clients_map = state_clone.clients.read().unwrap().clone(); - for (id, client_tx) in clients_map.iter() { - if id != &session_id_clone { - let _ = client_tx.send(event.to_string()).await; - } - } - } - } - } // 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 tx_opt = state_clone.clients.read().unwrap().get(&session_id_clone).cloned(); - if let Some(client_tx) = tx_opt { - if let Err(e) = client_tx.send(res_str).await { - tracing::error!("Failed to send response to client channel for session {}: {}", session_id_clone, e); - } - } else { - tracing::warn!("Could not find client_tx for session_id {} when trying to send response", session_id_clone); - } - } - } // End if let Ok(payload) - else { - tracing::warn!("Failed to parse payload as JSON from websocket message: {}", text); - } - } // End Ok(Message::Text(text)) - Ok(other) => { - tracing::info!("Received non-text message from websocket: {:?}", other); - } - Err(e) => { - tracing::error!("Websocket receive error: {}", e); - break; - } - } - } - tracing::info!("Websocket receiver task ended for session {}", session_id_clone); - }); - - tokio::select! { - _ = (&mut send_task) => { - tracing::info!("Websocket send task finished for session {}", session_id); - recv_task.abort(); - }, - _ = (&mut recv_task) => { - tracing::info!("Websocket recv task finished for session {}", session_id); - send_task.abort(); - }, - }; - - state.clients.write().unwrap().remove(&session_id); - tracing::info!("Websocket session {} closed and removed from state", session_id); - } + 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); + let event = serde_json::json!({ + "type": "activity", + "data": activity_msg + }); + + let clients_map = state_clone.clients.read().unwrap().clone(); + for (id, client_tx) in clients_map.iter() { + if id != &session_id_clone { + let _ = client_tx.send(event.to_string()).await; + } + } + } + } // 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 tx_opt = state_clone + .clients + .read() + .unwrap() + .get(&session_id_clone) + .cloned(); + if let Some(client_tx) = tx_opt { + if let Err(e) = client_tx.send(res_str).await { + tracing::error!( + "Failed to send response to client channel for session {}: {}", + session_id_clone, + e + ); + } + } else { + tracing::warn!( + "Could not find client_tx for session_id {} when trying to send response", + session_id_clone + ); + } + } + } + // End if let Ok(payload) + else { + tracing::warn!( + "Failed to parse payload as JSON from websocket message: {}", + text + ); + } + } // End Ok(Message::Text(text)) + Ok(other) => { + tracing::info!("Received non-text message from websocket: {:?}", other); + } + Err(e) => { + tracing::error!("Websocket receive error: {}", e); + break; + } + } + } + tracing::info!( + "Websocket receiver task ended for session {}", + session_id_clone + ); + }); + + tokio::select! { + _ = (&mut send_task) => { + tracing::info!("Websocket send task finished for session {}", session_id); + recv_task.abort(); + }, + _ = (&mut recv_task) => { + tracing::info!("Websocket recv task finished for session {}", session_id); + send_task.abort(); + }, + }; + + state.clients.write().unwrap().remove(&session_id); + tracing::info!( + "Websocket session {} closed and removed from state", + session_id + ); +} #[derive(serde::Deserialize, serde::Serialize, Debug)] pub struct NvimTelemetry { @@ -510,11 +576,13 @@ async fn nvim_telemetry_handler( axum::Json(payload): axum::Json, ) -> impl axum::response::IntoResponse { // 1. Update active_nvim.txt if FocusGained, BufEnter, or VimEnter - if payload.event == "FocusGained" || payload.event == "BufEnter" || payload.event == "VimEnter" { - let profile = std::env::var("USERPROFILE").unwrap_or_else(|_| "C:\\Users\\reazul.ashraf".into()); + if payload.event == "FocusGained" || payload.event == "BufEnter" || payload.event == "VimEnter" + { + let profile = + std::env::var("USERPROFILE").unwrap_or_else(|_| "C:\\Users\\reazul.ashraf".into()); let win_path = format!("{}\\.gemini\\active_nvim.txt", profile); let _ = std::fs::write(&win_path, &payload.session_id); - + let wsl_path = "\\\\wsl.localhost\\Ubuntu\\home\\riz\\.gemini\\active_nvim.txt"; let _ = std::fs::write(wsl_path, &payload.session_id); } @@ -524,10 +592,10 @@ async fn nvim_telemetry_handler( "type": "nvim_telemetry", "data": payload }); - + let msg_str = ws_msg.to_string(); let clients = state.clients.read().unwrap().clone(); - for (_, tx) in clients.iter() { + for tx in clients.values() { let _ = tx.send(msg_str.clone()).await; } @@ -549,10 +617,10 @@ fn init_logging(app_name: &str) -> Option Option Result<(), Box> { } } + let base_dir = std::env::var("MCP_MEMORY_STORE_DIR").unwrap_or_else(|_| { + dirs::home_dir() + .map(|mut h| { + h.push(".gemini/mcp_memory"); + h.to_string_lossy().into_owned() + }) + .unwrap_or_else(|| ".gemini/mcp_memory".into()) + }); + let base = PathBuf::from(base_dir); + fs::create_dir_all(&base).expect("Failed to create store dir"); - let base_dir = std::env::var("MCP_MEMORY_STORE_DIR").unwrap_or_else(|_| { - dirs::home_dir() - .map(|mut h| { - h.push(".gemini/mcp_memory"); - h.to_string_lossy().into_owned() - }) - .unwrap_or_else(|| ".gemini/mcp_memory".into()) - }); - let base = PathBuf::from(base_dir); - fs::create_dir_all(&base).expect("Failed to create store dir"); + let redb_path = base.join("mcp_store.redb"); + let db = Arc::new(redb::Database::create(&redb_path).unwrap()); - let redb_path = base.join("mcp_store.redb"); - let db = Arc::new(redb::Database::create(&redb_path).unwrap()); - - // Ensure table exists and migrate old JSON files + // Ensure table exists and migrate old JSON files + { + let write_txn = db.begin_write().unwrap(); { - let write_txn = db.begin_write().unwrap(); - { - let mut table = write_txn.open_table(crate::store::STORE_TABLE).unwrap(); - - let stores = vec![ - ("knowledge_graph_master", "knowledge_graph_master.json"), - ("audit_ledger", "audit_ledger.json"), - ("sticky_notes", "sticky_notes.json"), - ("tasks", "tasks.json"), - ("snippets", "snippets.json"), - ("adrs", "adrs.json"), - ("preferences", "preferences.json"), - ("error_fixes", "error_fixes.json"), - ("pinned_files", "pinned_files.json"), - ("session_summaries", "session_summaries.json"), - ("handoff_memos", "handoff_memos.json"), - ("env_fingerprints", "env_fingerprints.json"), - ("env_requirements", "env_requirements.json"), - ("milestones", "milestones.json"), - ("environments", "environments.json"), - ("pr_checklists", "pr_checklists.json"), - ("tech_debts", "tech_debts.json"), - ("gates", "gates.json"), - ("context_workspaces", "context_workspaces.json"), - ]; + let mut table = write_txn.open_table(crate::store::STORE_TABLE).unwrap(); - for (key, file_name) in stores.iter() { - if table.get(*key).unwrap().is_none() { - let json_path = base.join(file_name); - if json_path.exists() { - if let Ok(data) = fs::read(&json_path) { - if 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")); - } + let stores = vec![ + ("knowledge_graph_master", "knowledge_graph_master.json"), + ("audit_ledger", "audit_ledger.json"), + ("sticky_notes", "sticky_notes.json"), + ("tasks", "tasks.json"), + ("snippets", "snippets.json"), + ("adrs", "adrs.json"), + ("preferences", "preferences.json"), + ("error_fixes", "error_fixes.json"), + ("pinned_files", "pinned_files.json"), + ("session_summaries", "session_summaries.json"), + ("handoff_memos", "handoff_memos.json"), + ("env_fingerprints", "env_fingerprints.json"), + ("env_requirements", "env_requirements.json"), + ("milestones", "milestones.json"), + ("environments", "environments.json"), + ("pr_checklists", "pr_checklists.json"), + ("tech_debts", "tech_debts.json"), + ("gates", "gates.json"), + ("context_workspaces", "context_workspaces.json"), + ]; + + for (key, file_name) in stores.iter() { + if table.get(*key).unwrap().is_none() { + 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"), + ); } - } - } } } - write_txn.commit().unwrap(); } + write_txn.commit().unwrap(); + } - let state = Arc::new(MemoryState { - graph: Store::new("knowledge_graph_master", db.clone()), - base_dir: base.clone(), - search_index: RwLock::new(crate::search::MemoryIndex::new(&base).unwrap()), - ledger: Store::new("audit_ledger", db.clone()), - sticky: Store::new("sticky_notes", db.clone()), - tasks: Store::new("tasks", db.clone()), - snippets: Store::new("snippets", db.clone()), - adrs: Store::new("adrs", db.clone()), - prefs: Store::new("preferences", db.clone()), - error_fixes: Store::new("error_fixes", db.clone()), - pinned_files: Store::new("pinned_files", db.clone()), - session_summaries: Store::new("session_summaries", db.clone()), - handoff_memos: Store::new("handoff_memos", db.clone()), - env_fingerprints: Store::new("env_fingerprints", db.clone()), - env_requirements: Store::new("env_requirements", db.clone()), - milestones: Store::new("milestones", db.clone()), - environments: Store::new("environments", db.clone()), - pr_checklists: Store::new("pr_checklists", db.clone()), - tech_debts: Store::new("tech_debts", db.clone()), - gates: Store::new("gates", db.clone()), - context_workspaces: Store::new("context_workspaces", db.clone()), - activity_tx: tokio::sync::broadcast::channel(100).0, - }); + let state = Arc::new(MemoryState { + graph: Store::new("knowledge_graph_master", db.clone()), + base_dir: base.clone(), + search_index: RwLock::new(crate::search::MemoryIndex::new(&base).unwrap()), + ledger: Store::new("audit_ledger", db.clone()), + sticky: Store::new("sticky_notes", db.clone()), + tasks: Store::new("tasks", db.clone()), + snippets: Store::new("snippets", db.clone()), + adrs: Store::new("adrs", db.clone()), + prefs: Store::new("preferences", db.clone()), + error_fixes: Store::new("error_fixes", db.clone()), + pinned_files: Store::new("pinned_files", db.clone()), + session_summaries: Store::new("session_summaries", db.clone()), + handoff_memos: Store::new("handoff_memos", db.clone()), + env_fingerprints: Store::new("env_fingerprints", db.clone()), + env_requirements: Store::new("env_requirements", db.clone()), + milestones: Store::new("milestones", db.clone()), + environments: Store::new("environments", db.clone()), + pr_checklists: Store::new("pr_checklists", db.clone()), + tech_debts: Store::new("tech_debts", db.clone()), + gates: Store::new("gates", db.clone()), + context_workspaces: Store::new("context_workspaces", db.clone()), + activity_tx: tokio::sync::broadcast::channel(100).0, + }); - state.rebuild_index(); + state.rebuild_index(); - run_server(state) + run_server(state) } - - - - diff --git a/server/src/models.rs b/server/src/models.rs index 0dd1b10..12c55f3 100644 --- a/server/src/models.rs +++ b/server/src/models.rs @@ -1,6 +1,6 @@ +use schemars::JsonSchema; use serde::{Deserialize, Serialize}; use std::collections::HashMap; -use schemars::JsonSchema; #[derive(Debug, Clone, Serialize, Deserialize, Default)] pub struct CodeChange { diff --git a/server/src/search.rs b/server/src/search.rs index eb7fdef..3254925 100644 --- a/server/src/search.rs +++ b/server/src/search.rs @@ -28,7 +28,8 @@ impl MemoryIndex { let index_dir = store_dir.join("tantivy_index"); std::fs::create_dir_all(&index_dir).unwrap(); - let index = Index::open_in_dir(&index_dir).unwrap_or_else(|_| Index::create_in_dir(&index_dir, schema.clone()).unwrap()); + let index = Index::open_in_dir(&index_dir) + .unwrap_or_else(|_| Index::create_in_dir(&index_dir, schema.clone()).unwrap()); let writer = index.writer(50_000_000)?; let reader = index @@ -72,7 +73,6 @@ impl MemoryIndex { Ok(()) } - pub fn commit(&self) -> tantivy::Result<()> { let mut writer = self.writer.lock().unwrap(); writer.commit()?; @@ -113,9 +113,11 @@ impl MemoryIndex { .and_then(|v| v.as_str()) .unwrap_or(""); if let Some(ns) = namespace - && doc_ns != ns && doc_ns != "global" { - continue; - } + && doc_ns != ns + && doc_ns != "global" + { + continue; + } results.push((id, doc_type)); } Ok(results) @@ -172,7 +174,10 @@ mod tests { status: "open".to_string(), created_at: 0, updated_at: 0, - git_branch: None, acceptance_criteria: vec![], dependencies: vec![], parent_id: None, + git_branch: None, + acceptance_criteria: vec![], + dependencies: vec![], + parent_id: None, }; index.index_task(&task).unwrap(); @@ -218,11 +223,11 @@ mod tests { fn test_search_malformed_query() { let temp_dir = TempDir::new().unwrap(); let index = MemoryIndex::new(temp_dir.path()).unwrap(); - + // Malformed lucene query (unclosed parenthesis) let result = index.search("title: (unclosed", None); assert!(result.is_err()); - + // Another malformed query (unclosed quote) let result2 = index.search("title: \"unclosed", None); assert!(result2.is_err()); diff --git a/server/src/state.rs b/server/src/state.rs index 80a71a5..7583be3 100644 --- a/server/src/state.rs +++ b/server/src/state.rs @@ -33,19 +33,21 @@ pub struct MemoryState { impl MemoryState { pub fn unique_items(input: Vec) -> Vec { let mut keys = std::collections::HashSet::new(); - input.into_iter().filter(|entry| keys.insert(entry.clone())).collect() + input + .into_iter() + .filter(|entry| keys.insert(entry.clone())) + .collect() } pub fn broadcast_activity(&self, message: &str) { let payload = serde_json::json!({ "type": "activity", "data": message - }).to_string(); + }) + .to_string(); let _ = self.activity_tx.send(payload); } - - pub fn get_full_graph(&self) -> KnowledgeGraph { self.graph.read() } @@ -61,10 +63,10 @@ impl MemoryState { pub fn rebuild_index(&self) { if let Ok(new_idx) = MemoryIndex::new(&self.base_dir) { let graph = self.graph.read(); - for (_, e) in &graph.entities { + for e in graph.entities.values() { let _ = new_idx.index_entity(e); } - + let tasks = self.tasks.read(); for t in tasks { let _ = new_idx.index_task(&t); diff --git a/server/src/store.rs b/server/src/store.rs index b8e5051..a44ff3b 100644 --- a/server/src/store.rs +++ b/server/src/store.rs @@ -22,13 +22,11 @@ impl Store T { let read_txn = db.begin_read().unwrap(); - if let Ok(table) = read_txn.open_table(STORE_TABLE) { - if let Ok(Some(value)) = table.get(key) { - if let Ok(parsed) = serde_json::from_slice::(value.value()) { + 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; } - } - } T::default() } @@ -74,7 +72,7 @@ mod tests { async fn test_store_read_write() { let temp_file = NamedTempFile::new().unwrap(); let db = Database::create(temp_file.path()).unwrap(); - + let write_txn = db.begin_write().unwrap(); { write_txn.open_table(STORE_TABLE).unwrap(); @@ -94,18 +92,30 @@ mod tests { // Need to wait for spawn_blocking to finish tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; - assert_eq!(store.read(), TestData { name: "Hello".to_string(), value: 42 }); + assert_eq!( + store.read(), + TestData { + name: "Hello".to_string(), + value: 42 + } + ); // Load again to verify persistence let store2 = Store::::new("test_key", db.clone()); - assert_eq!(store2.read(), TestData { name: "Hello".to_string(), value: 42 }); + assert_eq!( + store2.read(), + TestData { + name: "Hello".to_string(), + value: 42 + } + ); } #[tokio::test(flavor = "multi_thread", worker_threads = 4)] async fn test_store_concurrency() { let temp_file = NamedTempFile::new().unwrap(); let db = Database::create(temp_file.path()).unwrap(); - + let write_txn = db.begin_write().unwrap(); { write_txn.open_table(STORE_TABLE).unwrap(); @@ -114,7 +124,7 @@ mod tests { let db = Arc::new(db); let store = Arc::new(Store::::new("concurrent_key", db.clone())); - + let mut handles = vec![]; for _ in 0..50 { let s = store.clone(); @@ -124,14 +134,14 @@ mod tests { }); })); } - + for h in handles { h.await.unwrap(); } - + // Wait for all blocking writes to flush tokio::time::sleep(tokio::time::Duration::from_millis(500)).await; - + assert_eq!(store.read().value, 50); } } diff --git a/server/tests/parity_test.rs b/server/tests/parity_test.rs index 34fcafe..0257ed5 100644 --- a/server/tests/parity_test.rs +++ b/server/tests/parity_test.rs @@ -3,53 +3,68 @@ 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"); + let memory_source = + std::fs::read_to_string("src/handlers.rs").expect("Failed to read handlers.rs"); let mut memory_tools = HashSet::new(); for line in memory_source.lines() { - if line.contains("crate::mcp::tool_def") { - if let Some(start) = line.find("(\"") { + if line.contains("crate::mcp::tool_def") + && let Some(start) = line.find("(\"") { let rest = &line[start + 2..]; 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"); + assert!( + !memory_tools.is_empty(), + "Could not find memory tools in handlers.rs" + ); // 2. Read nvim-core/src/lib.rs to get nvim tools - let nvim_source = std::fs::read_to_string("../nvim-core/src/lib.rs").expect("Failed to read nvim lib.rs"); + let nvim_source = + std::fs::read_to_string("../nvim-core/src/lib.rs").expect("Failed to read nvim lib.rs"); let mut nvim_tools = HashSet::new(); for line in nvim_source.lines() { - if line.contains("\"name\": \"nvim_") { - if let Some(start) = line.find("\"name\": \"") { + 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()); } } - } } - assert!(!nvim_tools.is_empty(), "Could not find nvim tools in lib.rs"); + assert!( + !nvim_tools.is_empty(), + "Could not find nvim tools in lib.rs" + ); // 3. Read Windows mcp_config.json - let win_home = std::env::var("USERPROFILE").unwrap_or_else(|_| "C:\\Users\\reazul.ashraf".to_string()); + let win_home = + std::env::var("USERPROFILE").unwrap_or_else(|_| "C:\\Users\\reazul.ashraf".to_string()); let win_config_path = std::path::PathBuf::from(win_home).join(".gemini/config/mcp_config.json"); if win_config_path.exists() { let config_str = std::fs::read_to_string(&win_config_path).unwrap(); let config: serde_json::Value = serde_json::from_str(&config_str).unwrap(); - + if let Some(eager) = config["mcpServers"]["memory"]["eagerTools"].as_array() { for tool in eager { let name = tool.as_str().unwrap(); - assert!(memory_tools.contains(name), "Windows config Memory tool '{}' not implemented in handlers.rs!", name); + assert!( + memory_tools.contains(name), + "Windows config Memory tool '{}' not implemented in handlers.rs!", + name + ); } } if let Some(nvim_eager) = config["mcpServers"]["win-nvim"]["eagerTools"].as_array() { for tool in nvim_eager { let name = tool.as_str().unwrap(); - assert!(nvim_tools.contains(name), "Windows config Nvim tool '{}' not implemented in nvim-core!", name); + assert!( + nvim_tools.contains(name), + "Windows config Nvim tool '{}' not implemented in nvim-core!", + name + ); } } } diff --git a/stub/build.rs b/stub/build.rs index 0e0609c..2a2c799 100644 --- a/stub/build.rs +++ b/stub/build.rs @@ -2,19 +2,24 @@ use std::process::Command; fn main() { let git_hash = Command::new("git") - .args(&["rev-parse", "--short", "HEAD"]) - .output() - .ok() - .and_then(|out| String::from_utf8(out.stdout).ok()) - .unwrap_or_else(|| "unknown".to_string()); - - let git_date = Command::new("git") - .args(&["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"]) + .args(["rev-parse", "--short", "HEAD"]) .output() .ok() .and_then(|out| String::from_utf8(out.stdout).ok()) .unwrap_or_else(|| "unknown".to_string()); - let version = format!("{} ({} {})", env!("CARGO_PKG_VERSION"), git_date.trim(), git_hash.trim()); + let git_date = Command::new("git") + .args(["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"]) + .output() + .ok() + .and_then(|out| String::from_utf8(out.stdout).ok()) + .unwrap_or_else(|| "unknown".to_string()); + + let version = format!( + "{} ({} {})", + env!("CARGO_PKG_VERSION"), + git_date.trim(), + git_hash.trim() + ); println!("cargo:rustc-env=APP_VERSION={}", version); } diff --git a/stub/src/bin/skeletal_client.rs b/stub/src/bin/skeletal_client.rs index 01fac13..ecbaa82 100644 --- a/stub/src/bin/skeletal_client.rs +++ b/stub/src/bin/skeletal_client.rs @@ -1,14 +1,15 @@ +use futures_util::StreamExt; use reqwest::Client; use std::env; use std::time::Duration; -use futures_util::StreamExt; #[tokio::main] async fn main() -> Result<(), Box> { tracing_subscriber::fmt::init(); let target = env::var("MCP_TARGET").unwrap_or_else(|_| "https://127.0.0.1:3000".to_string()); - let token = env::var("MCP_AUTH_TOKEN").unwrap_or_else(|_| "jP76lUJ5DtFRZmcvXH8LKdCTIkp29eAf".to_string()); + let token = env::var("MCP_AUTH_TOKEN") + .unwrap_or_else(|_| "jP76lUJ5DtFRZmcvXH8LKdCTIkp29eAf".to_string()); tracing::info!("Starting skeletal client to {}", target); @@ -17,13 +18,10 @@ async fn main() -> Result<(), Box> { .build()?; let sse_url = format!("{}/sse", target); - + tracing::info!("Connecting to SSE: {}", sse_url); - - let res = client.get(&sse_url) - .bearer_auth(&token) - .send() - .await?; + + let res = client.get(&sse_url).bearer_auth(&token).send().await?; if !res.status().is_success() { tracing::error!("Failed to connect to SSE: {}", res.status()); @@ -40,15 +38,15 @@ async fn main() -> Result<(), Box> { while let Some(chunk) = stream.next().await { let bytes = chunk?; buffer.extend_from_slice(&bytes); - + while let Some(pos) = buffer.windows(2).position(|w| w == b"\n\n" || w == b"\r\n") { let msg_bytes = buffer.drain(..pos).collect::>(); buffer.drain(..2); - + let text = String::from_utf8_lossy(&msg_bytes); let mut is_endpoint = false; let mut data_content = String::new(); - + for line in text.lines() { if line.starts_with("event: endpoint") { is_endpoint = true; @@ -78,7 +76,8 @@ async fn main() -> Result<(), Box> { tracing::info!("Sending test payload to {}", post_url); tracing::info!("Payload: {}", payload); - let post_res = client.post(&post_url) + let post_res = client + .post(&post_url) .bearer_auth(&token) .header("Content-Type", "application/json") .body(payload.to_string()) @@ -91,7 +90,7 @@ async fn main() -> Result<(), Box> { // Wait for the SSE stream to deliver the response tracing::info!("Waiting 2 seconds for SSE response delivery..."); - + let mut timeout = tokio::time::interval(Duration::from_secs(2)); timeout.tick().await; // first tick is immediate diff --git a/stub/src/main.rs b/stub/src/main.rs index 656c66f..3584727 100644 --- a/stub/src/main.rs +++ b/stub/src/main.rs @@ -1,6 +1,6 @@ -use std::sync::Arc; use clap::Parser; use futures_util::{SinkExt, StreamExt}; +use std::sync::Arc; use tokio::io::AsyncBufReadExt; use tokio::sync::mpsc; @@ -23,11 +23,11 @@ async fn read_mcp_message(stdin: &mut tokio::io::BufReader) -> return None; } tracing::info!("Read {} bytes from stdin: {:?}", bytes_read, line); - + if line.starts_with('{') { return Some(line.trim_end().to_string()); } - + let line = line.trim_end(); if line.is_empty() { break; @@ -51,16 +51,16 @@ fn init_logging(app_name: &str) -> Option Result<(), Box> { tracing::info!("Attempting to connect to {}", ws_url); use tokio_tungstenite::tungstenite::client::IntoClientRequest; - let mut request = match ws_url.clone().into_client_request() { + let request = match ws_url.clone().into_client_request() { Ok(req) => req, Err(e) => { tracing::error!("Failed to parse target URL {}: {}", ws_url, e); @@ -165,5 +165,3 @@ fn main() -> Result<(), Box> { Ok(()) }) } - - diff --git a/stub/tests/e2e.rs b/stub/tests/e2e.rs index a20ac19..96bdef4 100644 --- a/stub/tests/e2e.rs +++ b/stub/tests/e2e.rs @@ -1,5 +1,5 @@ -use serde_json::{json, Value}; -use std::io::{BufRead, BufReader, Read, Write}; +use serde_json::{Value, json}; +use std::io::{BufRead, BufReader, Write}; use std::process::{Command, Stdio}; use std::time::Duration; @@ -18,21 +18,26 @@ fn read_message(reader: &mut impl BufRead) -> Option { serde_json::from_str(line.trim()).ok() } - #[tokio::test] async fn test_full_system_e2e_performance() { - let temp_dir = std::env::temp_dir().join(format!("mcp_e2e_{}", std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs())); + let temp_dir = std::env::temp_dir().join(format!( + "mcp_e2e_{}", + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_secs() + )); std::fs::create_dir_all(&temp_dir).unwrap(); let test_port = "3042"; // Use a distinct port let test_auth_token = "test-token-12345"; - - // Since tests run from inside `target/debug/deps`, and `cargo test` does not guarantee + + // Since tests run from inside `target/debug/deps`, and `cargo test` does not guarantee // `env!("CARGO_BIN_EXE_name")` works correctly for binaries compiled in other crates without build dependencies, // we use `CARGO_MANIFEST_DIR` (which points to `stub`) to reliably locate the workspace `target/debug`. let manifest_dir = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR")); let debug_dir = manifest_dir.parent().unwrap().join("target").join("debug"); - + let server_exe = debug_dir.join(format!("mcp-memory-server{}", std::env::consts::EXE_SUFFIX)); let nvim_name = if cfg!(windows) { "mcp-memory-win-nvim" @@ -41,21 +46,24 @@ async fn test_full_system_e2e_performance() { }; let nvim_exe = debug_dir.join(format!("{}{}", nvim_name, std::env::consts::EXE_SUFFIX)); let stub_exe = debug_dir.join(format!("mcp-memory-stub{}", std::env::consts::EXE_SUFFIX)); - + assert!(server_exe.exists(), "Server not found at {:?}", server_exe); assert!(nvim_exe.exists(), "Nvim not found at {:?}", nvim_exe); assert!(stub_exe.exists(), "Stub not found at {:?}", stub_exe); // 1. Start Server - let mut server = Command::new(&server_exe).arg("--daemon") - .env("MCP_PORT", test_port).env("RUST_LOG", "debug") + let mut server = 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") + .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() .danger_accept_invalid_certs(true) @@ -63,28 +71,32 @@ async fn test_full_system_e2e_performance() { .unwrap(); let mut started = false; for _ in 0..30 { - if let Ok(resp) = client.get(format!("http://127.0.0.1:{}/health", test_port)).send().await { - if resp.status().is_success() { - started = true; - break; - } + if let Ok(resp) = client + .get(format!("http://127.0.0.1:{}/health", test_port)) + .send() + .await + && resp.status().is_success() + { + started = true; + break; } tokio::time::sleep(Duration::from_millis(500)).await; } assert!(started, "Server failed to start in time"); - + // 2. Start Stub let mut stub = 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") + .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.stdin.take().unwrap(); let mut stub_stdout = BufReader::new(stub.stdout.take().unwrap()); @@ -111,7 +123,7 @@ async fn test_full_system_e2e_performance() { "params": {}, "id": i }); - + // Alternate between LSP header format and JSONL format if i % 2 == 0 { send_message(&mut stub_stdin, tools_req); @@ -121,7 +133,8 @@ async fn test_full_system_e2e_performance() { stub_stdin.flush().unwrap(); } - let mut resp = read_message(&mut stub_stdout).expect("Failed to read rapid response from stub"); + let mut resp = + read_message(&mut stub_stdout).expect("Failed to read rapid response from stub"); while resp.get("id").is_none() || resp["id"].is_null() { resp = read_message(&mut stub_stdout).expect("Failed to read rapid response from stub"); } @@ -140,7 +153,7 @@ async fn test_full_system_e2e_performance() { "params": {}, "id": i }); - + if i % 2 == 0 { send_message(&mut nvim_stdin, tools_req); } else { @@ -149,9 +162,11 @@ async fn test_full_system_e2e_performance() { nvim_stdin.flush().unwrap(); } - let mut resp = read_message(&mut nvim_stdout).expect("Failed to read rapid response from win-nvim"); + let mut resp = + read_message(&mut nvim_stdout).expect("Failed to read rapid response from win-nvim"); while resp.get("id").is_none() || resp["id"].is_null() { - resp = read_message(&mut nvim_stdout).expect("Failed to read rapid response from win-nvim"); + resp = read_message(&mut nvim_stdout) + .expect("Failed to read rapid response from win-nvim"); } assert_eq!(resp["id"], i); } diff --git a/stub/tests/negative_scenarios.rs b/stub/tests/negative_scenarios.rs index 20e31df..634bc42 100644 --- a/stub/tests/negative_scenarios.rs +++ b/stub/tests/negative_scenarios.rs @@ -10,10 +10,14 @@ fn get_stub_exe() -> std::path::PathBuf { #[tokio::test] async fn test_stub_connection_refused() { - let _ = std::process::Command::new("cargo").arg("build").arg("--bin").arg("mcp-memory-stub").status(); + let _ = std::process::Command::new("cargo") + .arg("build") + .arg("--bin") + .arg("mcp-memory-stub") + .status(); let target = "http://127.0.0.1:49999"; - + let start = Instant::now(); let mut child = Command::new(get_stub_exe()) .arg("--target") @@ -21,19 +25,27 @@ async fn test_stub_connection_refused() { .stdin(Stdio::null()) // close stdin immediately to simulate EOF .spawn() .expect("Failed to execute stub"); - + let res = tokio::time::timeout(Duration::from_secs(5), child.wait()).await; let elapsed = start.elapsed(); - - assert!(res.is_ok(), "Stub hung on connection refused! Took {:?}", elapsed); + + assert!( + res.is_ok(), + "Stub hung on connection refused! Took {:?}", + elapsed + ); } #[tokio::test] async fn test_stub_handles_eof_cleanly() { - let _ = std::process::Command::new("cargo").arg("build").arg("--bin").arg("mcp-memory-stub").status(); + let _ = std::process::Command::new("cargo") + .arg("build") + .arg("--bin") + .arg("mcp-memory-stub") + .status(); let target = "http://127.0.0.1:49998"; - + let mut child = Command::new(get_stub_exe()) .arg("--target") .arg(target) @@ -42,28 +54,32 @@ async fn test_stub_handles_eof_cleanly() { .stderr(Stdio::piped()) .spawn() .expect("Failed to execute stub"); - + if let Some(mut stdin) = child.stdin.take() { use tokio::io::AsyncWriteExt; let msg = "Content-Length: 51\r\n\r\n{\"jsonrpc\":\"2.0\",\"method\":\"tools/list\",\"params\":{},\"id\":1}"; stdin.write_all(msg.as_bytes()).await.unwrap(); } // stdin dropped here - + let start = Instant::now(); let res = tokio::time::timeout(Duration::from_secs(5), child.wait()).await; let elapsed = start.elapsed(); - + assert!(res.is_ok(), "Stub hung after EOF! Took {:?}", elapsed); } #[tokio::test] async fn test_stub_sse_fallback_failure() { - let _ = std::process::Command::new("cargo").arg("build").arg("--bin").arg("mcp-memory-stub").status(); + let _ = std::process::Command::new("cargo") + .arg("build") + .arg("--bin") + .arg("mcp-memory-stub") + .status(); let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let local_addr = listener.local_addr().unwrap(); let target = format!("http://127.0.0.1:{}", local_addr.port()); - + tokio::spawn(async move { while let Ok((mut socket, _)) = listener.accept().await { use tokio::io::AsyncReadExt; @@ -77,12 +93,16 @@ async fn test_stub_sse_fallback_failure() { let mut child = Command::new(get_stub_exe()) .arg("--target") .arg(target) - .stdin(Stdio::null()) + .stdin(Stdio::null()) .spawn() .expect("Failed to execute stub"); - + let res = tokio::time::timeout(Duration::from_secs(5), child.wait()).await; let elapsed = start.elapsed(); - - assert!(res.is_ok(), "Stub hung on fallback failure! Took {:?}", elapsed); + + assert!( + res.is_ok(), + "Stub hung on fallback failure! Took {:?}", + elapsed + ); } diff --git a/win-nvim/build.rs b/win-nvim/build.rs index 0e0609c..2a2c799 100644 --- a/win-nvim/build.rs +++ b/win-nvim/build.rs @@ -2,19 +2,24 @@ use std::process::Command; fn main() { let git_hash = Command::new("git") - .args(&["rev-parse", "--short", "HEAD"]) - .output() - .ok() - .and_then(|out| String::from_utf8(out.stdout).ok()) - .unwrap_or_else(|| "unknown".to_string()); - - let git_date = Command::new("git") - .args(&["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"]) + .args(["rev-parse", "--short", "HEAD"]) .output() .ok() .and_then(|out| String::from_utf8(out.stdout).ok()) .unwrap_or_else(|| "unknown".to_string()); - let version = format!("{} ({} {})", env!("CARGO_PKG_VERSION"), git_date.trim(), git_hash.trim()); + let git_date = Command::new("git") + .args(["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"]) + .output() + .ok() + .and_then(|out| String::from_utf8(out.stdout).ok()) + .unwrap_or_else(|| "unknown".to_string()); + + let version = format!( + "{} ({} {})", + env!("CARGO_PKG_VERSION"), + 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 2c334f7..897c221 100644 --- a/win-nvim/tests/integration_test.rs +++ b/win-nvim/tests/integration_test.rs @@ -12,7 +12,7 @@ fn send_message(stdin: &mut std::process::ChildStdin, msg: Value) { fn read_message(stdout: &mut std::process::ChildStdout) -> Option { let mut reader = BufReader::new(stdout); let mut length = 0; - + // Read headers loop { let mut line = String::new(); @@ -27,16 +27,16 @@ fn read_message(stdout: &mut std::process::ChildStdout) -> Option { length = len_str.parse().unwrap_or(0); } } - + if length == 0 { return None; } - + // Read body let mut buf = vec![0u8; length]; reader.read_exact(&mut buf).unwrap(); let body_str = String::from_utf8_lossy(&buf); - + Some(serde_json::from_str(&body_str).unwrap()) } @@ -87,12 +87,12 @@ fn test_mcp_initialization_and_tools_list() { let s = serde_json::to_string(&init_req).unwrap(); stdin.write_all(format!("{}\n", s).as_bytes()).unwrap(); stdin.flush().unwrap(); - + let init_resp = read_message(&mut stdout).expect("Failed to read initialize response"); - + assert_eq!(init_resp["jsonrpc"], "2.0"); assert_eq!(init_resp["id"], 1); - + // Verify capabilities let capabilities = &init_resp["result"]["capabilities"]; assert_eq!(capabilities["tools"], serde_json::json!({})); @@ -104,17 +104,19 @@ fn test_mcp_initialization_and_tools_list() { "params": {}, "id": 2 }); - + send_message(&mut stdin, tools_req); - + let tools_resp = read_message(&mut stdout).expect("Failed to read tools/list response"); - + assert_eq!(tools_resp["jsonrpc"], "2.0"); assert_eq!(tools_resp["id"], 2); - - let tools = tools_resp["result"]["tools"].as_array().expect("result.tools must be an array"); + + let tools = tools_resp["result"]["tools"] + .as_array() + .expect("result.tools must be an array"); assert!(!tools.is_empty(), "Server must expose at least one tool"); - + let has_get_active_buffer = tools.iter().any(|t| t["name"] == "nvim_get_active_buffer"); assert!(has_get_active_buffer, "Missing nvim_get_active_buffer tool"); @@ -130,14 +132,17 @@ fn test_mcp_initialization_and_tools_list() { }, "id": 3 }); - + send_message(&mut stdin, call_req); - + let call_resp = read_message(&mut stdout).expect("Failed to read tools/call response"); - + assert_eq!(call_resp["jsonrpc"], "2.0"); assert_eq!(call_resp["id"], 3); - assert!(call_resp.get("error").is_some(), "Expected an error response since Neovim shouldn't be running"); + assert!( + call_resp.get("error").is_some(), + "Expected an error response since Neovim shouldn't be running" + ); assert_eq!(call_resp["error"]["code"], -32603); // Internal Error child.kill().expect("Failed to kill child");